support connection management

This commit is contained in:
sijie.sun
2025-01-24 17:28:11 +08:00
parent 0057f666b7
commit a5f4b4e3a0
10 changed files with 1074 additions and 124 deletions
+9 -2
View File
@@ -12,11 +12,18 @@ keywords = ["kcp", "bindings"]
[dependencies]
bytes = "1"
tokio = { version = "1", features = ["full"] }
tokio-util = { version = "0.7.13" }
dashmap = "6.1.0"
zerocopy = { version = "0.7", features = ["derive", "simd"] }
bitflags = "2.5"
bitflags = "2.8.0"
parking_lot = "0.12.3"
auto_impl = "1.2.1"
thiserror = "2.0.11"
anyhow = "1.0.95"
tracing = "0.1.41"
tracing-subscriber = "0.3.19"
rand = "0.8.5"
[build-dependencies]
bindgen = "0.71.1"
cc = "1.2.7"
cc = "1.2.10"
-1
View File
@@ -22,7 +22,6 @@ fn main() {
.header("wrapper.h")
.parse_callbacks(Box::new(bindgen::CargoCallbacks::new()))
.allowlist_function("ikcp_.*")
.opaque_type("ikcpcb")
.generate()
.expect("Unable to generate bindings");
+628 -96
View File
@@ -3,15 +3,18 @@ use std::sync::{
Arc,
};
use bytes::BytesMut;
use anyhow::Context;
use bytes::{Bytes, BytesMut};
use dashmap::DashMap;
use parking_lot::Mutex;
use tokio::{select, sync::Notify, task::JoinSet};
use tokio::{select, sync::Notify, task::JoinSet, time::timeout};
use tracing::Instrument;
use crate::{
error::Error,
ffi_safe::{Kcp, KcpConfig},
packet_def::KcpPacket,
state::{KcpConnectionFSM, PacketHeaderFlagManipulator},
};
pub type Sender<T> = tokio::sync::mpsc::Sender<T>;
@@ -23,26 +26,34 @@ pub type KcpPacketReceiver = Receiver<KcpPacket>;
pub type KcpStreamSender = Sender<BytesMut>;
pub type KcpStreamReceiver = Receiver<BytesMut>;
enum KcpClientState {
SynSent,
Established,
Fin,
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct ConnId {
conv: u32,
src_session_id: u32,
dst_session_id: u32,
}
enum KcpServerState {
SynReceived,
Established,
Fin,
impl From<&KcpPacket> for ConnId {
fn from(packet: &KcpPacket) -> Self {
Self {
conv: packet.header().conv(),
src_session_id: packet.header().src_session_id(),
dst_session_id: packet.header().dst_session_id(),
}
}
}
enum KcpConnectionState {
Client(KcpClientState),
Server(KcpServerState),
impl ConnId {
fn fill_packet_header(&self, packet: &mut KcpPacket) {
packet
.mut_header()
.set_conv(self.conv)
.set_src_session_id(self.src_session_id)
.set_dst_session_id(self.dst_session_id);
}
}
struct KcpConnectionInner {
conv: u32,
update_notifier: Notify,
recv_notifier: Notify,
send_notifier: Notify,
@@ -52,34 +63,35 @@ struct KcpConnectionInner {
}
struct KcpConnection {
conv: u32,
conn_id: ConnId,
kcp: Arc<Mutex<Box<Kcp>>>,
inner: Arc<KcpConnectionInner>,
send_sender: Sender<BytesMut>,
send_sender: Option<Sender<BytesMut>>,
send_receiver: Option<Receiver<BytesMut>>,
recv_sender: Sender<BytesMut>,
recv_sender: Option<Sender<BytesMut>>,
recv_receiver: Option<Receiver<BytesMut>>,
send_close_notifier: Arc<Notify>,
recv_closed: Arc<AtomicBool>,
tasks: JoinSet<()>,
}
impl KcpConnection {
fn new(conv: u32) -> Result<Self, Error> {
let kcp = Kcp::new(KcpConfig::new(conv))?;
pub fn new(conn_id: ConnId) -> Result<Self, Error> {
let kcp = Kcp::new(KcpConfig::new_turbo(conn_id.conv))?;
let (send_sender, send_receiver) = tokio::sync::mpsc::channel(16);
let (recv_sender, recv_receiver) = tokio::sync::mpsc::channel(16);
let (send_sender, send_receiver) = tokio::sync::mpsc::channel(128);
let (recv_sender, recv_receiver) = tokio::sync::mpsc::channel(128);
Ok(Self {
conv,
conn_id,
kcp: Arc::new(Mutex::new(kcp)),
inner: Arc::new(KcpConnectionInner {
conv,
update_notifier: Notify::new(),
recv_notifier: Notify::new(),
send_notifier: Notify::new(),
@@ -88,34 +100,38 @@ impl KcpConnection {
waiting_new_send_window: AtomicBool::new(false),
}),
send_sender,
send_sender: Some(send_sender),
send_receiver: Some(send_receiver),
recv_sender,
recv_sender: Some(recv_sender),
recv_receiver: Some(recv_receiver),
send_close_notifier: Arc::new(Notify::new()),
recv_closed: Arc::new(AtomicBool::new(false)),
tasks: JoinSet::new(),
})
}
fn run(&mut self, output_sender: KcpPakcetSender) {
pub fn run(&mut self, output_sender: KcpPakcetSender) {
let conn_id = self.conn_id;
self.kcp
.lock()
.set_output_cb(Box::new(move |conv, data: BytesMut| {
let mut kcp_packet = KcpPacket::new_with_payload(&data);
kcp_packet
.mut_header()
.set_conv(conv)
.set_len(data.len() as u16)
.set_data(true);
println!("conv {} send output data: {:?}", conv, kcp_packet);
let _ = output_sender.try_send(kcp_packet);
conn_id.fill_packet_header(&mut kcp_packet);
kcp_packet.mut_header().set_data(true).set_ack(true);
tracing::trace!(?conv, "sending output data: {:?}", kcp_packet);
if let Err(e) = output_sender.try_send(kcp_packet) {
tracing::debug!(?e, ?conn_id, "send output data failed");
}
Ok(())
}));
// kcp updater
let inner = self.inner.clone();
let kcp = self.kcp.clone();
let recv_closed = self.recv_closed.clone();
self.tasks.spawn(async move {
loop {
let next_update_ms = kcp.lock().next_update_delay_ms();
@@ -133,6 +149,10 @@ impl KcpConnection {
if inner.waiting_new_send_window.swap(false, std::sync::atomic::Ordering::SeqCst) {
inner.send_notifier.notify_one();
}
if recv_closed.load(std::sync::atomic::Ordering::Relaxed) {
inner.recv_notifier.notify_one();
}
}
});
@@ -140,72 +160,247 @@ impl KcpConnection {
let kcp = self.kcp.clone();
let inner = self.inner.clone();
let mut send_receiver = self.send_receiver.take().unwrap();
self.tasks.spawn(async move {
while let Some(data) = send_receiver.recv().await {
loop {
let (waitsnd, sndwnd) = {
let kcp = kcp.lock();
(kcp.waitsnd(), kcp.sendwnd())
};
if waitsnd > 2 * sndwnd {
inner
.waiting_new_send_window
.store(true, std::sync::atomic::Ordering::SeqCst);
inner.send_notifier.notified().await;
} else {
break;
let send_close_notifier = self.send_close_notifier.clone();
self.tasks.spawn(
async move {
while let Some(data) = send_receiver.recv().await {
loop {
let (waitsnd, sndwnd) = {
let kcp = kcp.lock();
(kcp.waitsnd(), kcp.sendwnd())
};
if waitsnd > 2 * sndwnd {
inner
.waiting_new_send_window
.store(true, std::sync::atomic::Ordering::SeqCst);
inner.send_notifier.notified().await;
} else {
break;
}
}
kcp.lock().send(data.freeze()).unwrap();
kcp.lock().flush();
inner.update_notifier.notify_one();
}
kcp.lock().send(data.freeze()).unwrap();
tracing::debug!(
?conn_id,
"connection packet sender close, waiting for waitsnd to be 0"
);
// waiting for waitsnd to be 0
while kcp.lock().waitsnd() > 0 {
inner
.waiting_new_send_window
.store(true, std::sync::atomic::Ordering::SeqCst);
inner.send_notifier.notified().await;
}
send_close_notifier.notify_one();
tracing::debug!(?conn_id, "connection packet send task done");
}
});
.instrument(tracing::trace_span!("send_task", conn = ?conn_id)),
);
// handle packet recv
let kcp = self.kcp.clone();
let inner = self.inner.clone();
let recv_sender = self.recv_sender.clone();
self.tasks.spawn(async move {
let mut buf = BytesMut::new();
loop {
if buf.capacity() < 1024 {
buf.reserve(4096);
}
let ret = kcp.lock().recv(&mut buf);
if let Err(_) = ret {
println!("recv error, conv: {}", inner.conv);
inner.recv_notifier.notified().await;
} else {
println!("recv data: {:?}", buf);
let conn_id = self.conn_id;
let recv_sender = self.recv_sender.take().unwrap();
let recv_closed = self.recv_closed.clone();
self.tasks.spawn(
async move {
let mut buf = BytesMut::new();
while !recv_closed.load(std::sync::atomic::Ordering::Relaxed) {
let peeksize = kcp.lock().peeksize();
if peeksize <= 0 {
tracing::trace!("recv nothing, wait for next update");
inner.recv_notifier.notified().await;
continue;
};
if buf.capacity() < peeksize as usize {
buf.reserve(std::cmp::max(peeksize as usize, 4096));
}
kcp.lock().recv(&mut buf).unwrap();
tracing::trace!("recv data ({}): {:?}", buf.len(), buf);
assert_ne!(0, buf.len());
let send_ret = recv_sender.send(buf.split()).await;
if let Err(_) = send_ret {
break;
}
}
tracing::debug!(?conn_id, "connection packet recv task done");
}
});
.instrument(tracing::trace_span!("recv_task", conn = ?conn_id)),
);
}
fn handle_input(&mut self, packet: KcpPacket) {
let _ = self.kcp.lock().handle_input(packet.payload());
fn handle_input(&mut self, packet: &KcpPacket) -> Result<(), Error> {
self.kcp.lock().handle_input(packet.payload())?;
self.inner
.has_new_input
.store(true, std::sync::atomic::Ordering::SeqCst);
self.inner.update_notifier.notify_one();
Ok(())
}
fn send_sender(&self) -> KcpStreamSender {
self.send_sender.clone()
fn send_sender(&mut self) -> KcpStreamSender {
self.send_sender.take().unwrap()
}
fn recv_receiver(&mut self) -> KcpStreamReceiver {
self.recv_receiver.take().unwrap()
}
fn send_close_notifier(&self) -> Arc<Notify> {
self.send_close_notifier.clone()
}
fn close_recv(&self) {
self.recv_closed
.store(true, std::sync::atomic::Ordering::SeqCst);
self.inner.recv_notifier.notify_one();
}
}
impl Drop for KcpConnection {
fn drop(&mut self) {
self.send_close_notifier.notify_one();
}
}
impl PacketHeaderFlagManipulator for KcpPacket {
fn has_syn(&self) -> bool {
self.header().is_syn()
}
fn has_ack(&self) -> bool {
self.header().is_ack()
}
fn has_fin(&self) -> bool {
self.header().is_fin()
}
fn has_rst(&self) -> bool {
self.header().is_rst()
}
fn has_data(&self) -> bool {
self.header().is_data()
}
fn set_syn(&mut self, value: bool) {
self.mut_header().set_syn(value);
}
fn set_ack(&mut self, value: bool) {
self.mut_header().set_ack(value);
}
fn set_fin(&mut self, value: bool) {
self.mut_header().set_fin(value);
}
fn set_rst(&mut self, value: bool) {
self.mut_header().set_rst(value);
}
fn set_data(&mut self, value: bool) {
self.mut_header().set_data(value);
}
}
struct KcpConnectionState {
fsm: KcpConnectionFSM,
notify: Arc<Notify>,
conn_data: Bytes,
last_pong: std::time::Instant,
}
impl std::fmt::Debug for KcpConnectionState {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("KcpConnectionState")
.field("fsm", &self.fsm)
.finish()
}
}
impl KcpConnectionState {
fn new(fsm: KcpConnectionFSM) -> Self {
Self {
fsm,
notify: Arc::new(Notify::new()),
conn_data: Bytes::new(),
last_pong: std::time::Instant::now(),
}
}
fn handle_packet(&mut self, packet: &KcpPacket) -> Result<Option<KcpPacket>, Error> {
self.notify_pong();
let mut out_packet = None;
let old_state = self.fsm.clone();
let _ = self.fsm.handle_packet(packet, &mut out_packet);
if old_state != self.fsm {
self.notify.notify_one();
return Ok(out_packet);
}
Ok(None)
}
fn notify(&self) -> Arc<Notify> {
self.notify.clone()
}
fn is_established(&self) -> bool {
matches!(self.fsm, KcpConnectionFSM::Established)
}
fn is_peer_closed(&self) -> bool {
matches!(
self.fsm,
KcpConnectionFSM::PeerClosed | KcpConnectionFSM::Closed
)
}
fn is_closed(&self) -> bool {
matches!(self.fsm, KcpConnectionFSM::Closed)
}
fn set_data(&mut self, data: Bytes) {
self.conn_data = data;
}
fn notify_pong(&mut self) {
self.last_pong = std::time::Instant::now();
}
fn is_pong_timeout(&self) -> bool {
self.last_pong.elapsed() > std::time::Duration::from_secs(60)
}
}
struct KcpEndpointData {
cur_conv: AtomicU32,
conn_map: DashMap<ConnId, KcpConnection>,
state_map: DashMap<ConnId, KcpConnectionState>,
}
impl KcpEndpointData {
fn new() -> Self {
Self {
cur_conv: AtomicU32::new(rand::random()),
conn_map: DashMap::new(),
state_map: DashMap::new(),
}
}
}
pub struct KcpEndpoint {
cur_conv: AtomicU32,
established_map: Arc<DashMap<u32, KcpConnection>>,
id: u64,
data: Arc<KcpEndpointData>,
input_sender: KcpPakcetSender,
input_receiver: Option<KcpPacketReceiver>,
@@ -213,17 +408,27 @@ pub struct KcpEndpoint {
output_sender: KcpPakcetSender,
output_receiver: Option<KcpPacketReceiver>,
new_conn_sender: tokio::sync::mpsc::Sender<ConnId>,
new_conn_receiver: Arc<tokio::sync::Mutex<tokio::sync::mpsc::Receiver<ConnId>>>,
tasks: JoinSet<()>,
}
impl std::fmt::Debug for KcpEndpoint {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("KcpEndpoint").field("id", &self.id).finish()
}
}
impl KcpEndpoint {
pub fn new() -> Self {
let (input_sender, input_receiver) = tokio::sync::mpsc::channel(1024);
let (output_sender, output_receiver) = tokio::sync::mpsc::channel(1024);
let (new_conn_sender, new_conn_receiver) = tokio::sync::mpsc::channel(4);
Self {
cur_conv: AtomicU32::new(0),
established_map: Arc::new(DashMap::new()),
id: rand::random(),
data: Arc::new(KcpEndpointData::new()),
input_sender,
input_receiver: Some(input_receiver),
@@ -231,32 +436,219 @@ impl KcpEndpoint {
output_sender,
output_receiver: Some(output_receiver),
new_conn_sender,
new_conn_receiver: Arc::new(tokio::sync::Mutex::new(new_conn_receiver)),
tasks: JoinSet::new(),
}
}
async fn try_handle_pingpong(
data: &KcpEndpointData,
packet: &KcpPacket,
output_sender: &KcpPakcetSender,
) -> bool {
if !packet.header().is_ping() {
return false;
}
if !packet.header().is_pong() {
let mut out_packet = packet.clone();
out_packet.mut_header().set_pong(true);
tracing::trace!("sending pong packet: {:?}", out_packet);
let ret = output_sender.send(out_packet).await;
if let Err(e) = ret {
tracing::error!(?e, "send pong packet failed");
}
} else {
let conv = ConnId::from(packet);
if let Some(mut state) = data.state_map.get_mut(&conv) {
state.notify_pong();
}
}
true
}
pub async fn run(&mut self) {
let mut input_receiver = self.input_receiver.take().unwrap();
let established_map = self.established_map.clone();
let data = self.data.clone();
let output_sender = self.output_sender.clone();
let new_conn_sender = self.new_conn_sender.clone();
self.tasks.spawn(async move {
while let Some(packet) = input_receiver.recv().await {
let conv = packet.header().conv();
if packet.header().is_data() {
let Some(mut conn) = established_map.get_mut(&conv) else {
self.tasks.spawn(
async move {
while let Some(packet) = input_receiver.recv().await {
tracing::trace!("recv packet: {:?}", packet);
if Self::try_handle_pingpong(&data, &packet, &output_sender).await {
continue;
};
let _ = conn.handle_input(packet);
}
let conv = ConnId::from(&packet);
if packet.header().is_data() && packet.payload().len() > 0 {
if let Some(mut conn) = data.conn_map.get_mut(&conv) {
if let Err(e) = conn.handle_input(&packet) {
tracing::error!(?e, ?conv, "handle input on connection failed");
} else {
tracing::trace!(?conv, "handle input on connection done");
}
} else {
tracing::debug!(
?conv,
?packet,
"no conn for conv when handling data packet"
);
}
}
let mut state_ref = data.state_map.get_mut(&conv);
let state = state_ref.as_deref_mut();
let mut out_packet: Option<KcpPacket> = None;
if state.is_none() {
if packet.header().is_rst() {
tracing::debug!(?conv, "reset packet for conn, but no state");
continue;
}
let mut tmp_fsm = KcpConnectionFSM::listen();
let res = tmp_fsm.handle_packet(&packet, &mut out_packet);
tracing::trace!(
?conv,
?state,
?out_packet,
"handle first packet for conn, ret: {:?}",
res
);
if res.is_ok() {
let mut conn_state = KcpConnectionState::new(tmp_fsm);
conn_state.set_data(packet.payload().to_vec().into());
data.state_map.insert(conv, conn_state);
}
} else {
let state = state.unwrap();
let prev_established = state.is_established();
let ret = state.handle_packet(&packet);
tracing::trace!(?conv, ?state, "handle packet for conn, ret: {:?}", ret);
if ret.is_ok() {
out_packet = ret.unwrap();
}
if !prev_established && state.is_established() {
let _ = new_conn_sender.try_send(conv);
}
if state.is_peer_closed() {
tracing::debug!(?conv, "peer half closed, close recv");
data.conn_map.get_mut(&conv).map(|conn| conn.close_recv());
}
if state.is_closed() {
// state map will be cleaned by periodic task
tracing::debug!(?conv, "connection closed, remove state");
data.conn_map.remove(&conv);
}
}
drop(state_ref);
if let Some(mut out_packet) = out_packet {
conv.fill_packet_header(&mut out_packet);
tracing::trace!(?conv, ?out_packet, "sending output packet");
let ret = output_sender.send(out_packet).await;
if let Err(e) = ret {
tracing::error!(?e, "send output packet failed");
}
}
}
}
.instrument(tracing::trace_span!("recv_task", id = self.id)),
);
// conn clean task
let data = self.data.clone();
self.tasks.spawn(async move {
loop {
data.state_map.retain(|_, state| {
!matches!(state.fsm, KcpConnectionFSM::Closed) && !state.is_pong_timeout()
});
data.conn_map
.retain(|conn_id, _| data.state_map.contains_key(conn_id));
tokio::time::sleep(std::time::Duration::from_secs(10)).await;
}
});
// conn ping task
let data = self.data.clone();
let output_sender = self.output_sender.clone();
self.tasks.spawn(async move {
loop {
let packets = data
.state_map
.iter()
.filter_map(|item| {
let (conn_id, state) = item.pair();
if state.is_closed() {
return None;
}
let mut out_packet = KcpPacket::new(0);
conn_id.fill_packet_header(&mut out_packet);
out_packet.mut_header().set_ping(true);
Some(out_packet)
})
.collect::<Vec<_>>();
for packet in packets {
let ret = output_sender.send(packet).await;
if let Err(e) = ret {
tracing::error!(?e, "send ping packet failed");
}
tokio::time::sleep(std::time::Duration::from_millis(5)).await;
}
tokio::time::sleep(std::time::Duration::from_secs(10)).await;
}
});
}
pub fn add_established(&mut self, conv: u32) -> Result<(), Error> {
let mut conn = KcpConnection::new(conv)?;
fn add_conn(&self, conn_id: ConnId) -> Result<(), Error> {
let mut conn = KcpConnection::new(conn_id)?;
conn.run(self.output_sender.clone());
self.established_map.insert(conv, conn);
let data = self.data.clone();
let close_notifier = conn.send_close_notifier();
data.conn_map.insert(conn_id, conn);
let output_sender = self.output_sender.clone();
let data = Arc::downgrade(&data);
tokio::spawn(async move {
close_notifier.notified().await;
let Some(data) = data.upgrade() else {
return;
};
let mut out_packet = KcpPacket::new(0);
let Some(mut state) = data.state_map.get_mut(&conn_id) else {
return;
};
let close_ret = state.fsm.close(&mut out_packet);
let cur_state = state.fsm.clone();
let is_closed = state.is_closed();
drop(state);
match close_ret {
Ok(_) => {
conn_id.fill_packet_header(&mut out_packet);
output_sender.send(out_packet).await.unwrap();
}
Err(e) => {
tracing::error!(?e, ?conn_id, "close connection failed");
}
}
if is_closed {
data.conn_map.remove(&conn_id);
}
tracing::debug!(?conn_id, ?cur_state, "connection close watcher done");
});
Ok(())
}
@@ -269,30 +661,131 @@ impl KcpEndpoint {
self.input_sender.clone()
}
pub fn conn_sender_receiver(&self, conv: u32) -> Option<(KcpStreamSender, KcpStreamReceiver)> {
let mut conn = self.established_map.get_mut(&conv)?;
pub fn input_sender_ref(&self) -> &KcpPakcetSender {
&self.input_sender
}
pub fn conn_sender_receiver(
&self,
conn_id: ConnId,
) -> Option<(KcpStreamSender, KcpStreamReceiver)> {
let mut conn = self.data.conn_map.get_mut(&conn_id)?;
Some((conn.send_sender(), conn.recv_receiver()))
}
pub fn conn_data(&self, conn_id: &ConnId) -> Option<Bytes> {
let state = self.data.state_map.get(conn_id)?;
Some(state.conn_data.clone())
}
#[tracing::instrument(ret)]
pub async fn connect(
&self,
timeout_dur: std::time::Duration,
src_session_id: u32,
dst_session_id: u32,
conn_data: Bytes,
) -> Result<ConnId, Error> {
let mut out_packet = KcpPacket::new_with_payload(&conn_data);
let conn_id = loop {
let conv_cand = self
.data
.cur_conv
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
let conn_id = ConnId {
conv: conv_cand,
src_session_id,
dst_session_id,
};
if !self.data.state_map.contains_key(&conn_id) {
break conn_id;
}
};
let fsm = KcpConnectionFSM::connect(&mut out_packet);
let mut state = KcpConnectionState::new(fsm);
state.set_data(conn_data);
let notify = state.notify();
self.data.state_map.insert(conn_id, state);
conn_id.fill_packet_header(&mut out_packet);
tracing::trace!(?conn_id, "connect packet: {:?}", out_packet);
self.output_sender
.send(out_packet)
.await
.with_context(|| "send connect packet failed")?;
if timeout(timeout_dur, notify.notified()).await.is_err() {
self.data.state_map.remove(&conn_id);
return Err(Error::ConnectTimeout);
}
if let Some(state) = self.data.state_map.get(&conn_id) {
tracing::debug!(?conn_id, ?state, "connect done, checkin state");
if matches!(state.fsm, KcpConnectionFSM::Established) {
self.add_conn(conn_id)?;
return Ok(conn_id);
} else {
drop(state);
self.data.state_map.remove(&conn_id);
}
// if task aborted, the state map will be cleaned by periodic task
}
return Err(anyhow::anyhow!("connect failed").into());
}
pub async fn accept(&self) -> Result<ConnId, Error> {
let conn_receiver = self.new_conn_receiver.clone();
loop {
let Some(conn_id) = conn_receiver.lock().await.recv().await else {
return Err(Error::Shutdown);
};
let Some(state) = self.data.state_map.get(&conn_id) else {
tracing::debug!(?conn_id, "no state for conn, ignore");
continue;
};
if matches!(state.fsm, KcpConnectionFSM::Established) {
self.add_conn(conn_id)?;
return Ok(conn_id);
}
}
}
}
#[cfg(test)]
mod tests {
use tracing::level_filters::LevelFilter;
use tracing_subscriber::{layer::SubscriberExt, util::SubscriberInitExt, Layer as _};
use super::*;
#[tokio::test]
async fn test_kcp_endpoint() {
fn _enable_log() {
let console_layer = tracing_subscriber::fmt::layer()
.pretty()
.with_writer(std::io::stderr)
.with_filter(LevelFilter::TRACE);
tracing_subscriber::Registry::default()
.with(console_layer)
.init();
}
async fn prepare_test() -> (KcpEndpoint, KcpEndpoint, JoinSet<()>) {
let mut client_endpoint = KcpEndpoint::new();
let mut server_endpoint = KcpEndpoint::new();
let mut t = JoinSet::new();
client_endpoint.run().await;
server_endpoint.run().await;
let _ = client_endpoint.add_established(1).unwrap();
let _ = server_endpoint.add_established(1).unwrap();
let client_input_sender = client_endpoint.input_sender();
let mut server_output_receiver = server_endpoint.output_receiver().unwrap();
let t1 = tokio::spawn(async move {
t.spawn(async move {
while let Some(packet) = server_output_receiver.recv().await {
let _ = client_input_sender.send(packet).await;
}
@@ -300,14 +793,41 @@ mod tests {
let server_input_sender = server_endpoint.input_sender();
let mut client_output_receiver = client_endpoint.output_receiver().unwrap();
let t2 = tokio::spawn(async move {
t.spawn(async move {
while let Some(packet) = client_output_receiver.recv().await {
let _ = server_input_sender.send(packet).await;
}
});
let (client_sender, mut client_receiver) = client_endpoint.conn_sender_receiver(1).unwrap();
let (server_sender, mut server_receiver) = server_endpoint.conn_sender_receiver(1).unwrap();
(client_endpoint, server_endpoint, t)
}
#[tokio::test]
async fn test_kcp_connect_and_close() {
let mut p = KcpPacket::new(0);
let _ = p.mut_header().conv();
let (client_endpoint, server_endpoint, t) = prepare_test().await;
let (connect_ret, accept_ret) = tokio::join!(
client_endpoint.connect(std::time::Duration::from_secs(1), 1, 3, Bytes::from("conn")),
server_endpoint.accept()
);
assert_eq!(*connect_ret.as_ref().unwrap(), accept_ret.unwrap());
let conv = connect_ret.unwrap();
let client_conn_data = client_endpoint.conn_data(&conv).unwrap();
assert_eq!("conn", String::from_utf8_lossy(&client_conn_data));
let server_conn_data = server_endpoint.conn_data(&conv).unwrap();
assert_eq!("conn", String::from_utf8_lossy(&server_conn_data));
let (client_sender, mut client_receiver) =
client_endpoint.conn_sender_receiver(conv).unwrap();
let (server_sender, mut server_receiver) =
server_endpoint.conn_sender_receiver(conv).unwrap();
client_sender.send(BytesMut::from("hello")).await.unwrap();
let data = server_receiver.recv().await.unwrap();
@@ -317,9 +837,21 @@ mod tests {
let data = client_receiver.recv().await.unwrap();
assert_eq!("world", String::from_utf8_lossy(&data));
// test half close
drop(client_sender);
assert!(server_receiver.recv().await.is_none());
// server can still send data
server_sender.send(BytesMut::from("world")).await.unwrap();
let data = client_receiver.recv().await.unwrap();
assert_eq!("world", String::from_utf8_lossy(&data));
// full close
drop(server_sender);
assert!(client_receiver.recv().await.is_none());
drop(client_endpoint);
drop(server_endpoint);
let _ = tokio::join!(t1, t2);
t.join_all().await;
}
}
+16 -2
View File
@@ -1,5 +1,19 @@
#[derive(Debug)]
#[derive(Debug, thiserror::Error)]
pub enum Error {
#[error("Invalid state")]
InvalidState,
#[error("Invalid state need reset")]
InvalidStateNeedRst,
#[error("Connection reset")]
ConnectioinReset,
#[error("Create connection failed")]
CreateConnectionFailed,
Unknown,
#[error("Anyhow error")]
AnyhowError(#[from] anyhow::Error),
#[error("Connect timeout")]
ConnectTimeout,
#[error("Shutdown")]
Shutdown,
}
+1
View File
@@ -1,3 +1,4 @@
#![allow(unused)]
#![allow(non_upper_case_globals)]
#![allow(non_camel_case_types)]
#![allow(non_snake_case)]
+31 -6
View File
@@ -30,6 +30,19 @@ impl KcpConfig {
nc: None,
}
}
pub fn new_turbo(conv: IUINT32) -> Self {
Self {
conv,
mtu: Some(1200),
sndwnd: Some(1024),
rcvwnd: Some(1024),
nodelay: Some(1),
interval: Some(10),
resend: Some(2),
nc: Some(1),
}
}
}
pub type OutputCb = Box<dyn Fn(u32, BytesMut) -> Result<(), Error>>;
@@ -81,6 +94,8 @@ impl Kcp {
return Err(Error::CreateConnectionFailed);
}
(*kcp).stream = 1;
ret.kcp = kcp;
ikcp_setoutput(kcp, Some(ikcp_output));
@@ -98,7 +113,7 @@ impl Kcp {
pub fn handle_input(&mut self, data: &[u8]) -> Result<(), Error> {
let ret = unsafe { ikcp_input(self.kcp, data.as_ptr() as *const i8, data.len() as i64) };
if ret < 0 {
return Err(Error::Unknown);
return Err(anyhow::anyhow!("input failed, return: {}", ret).into());
} else {
return Ok(());
}
@@ -119,17 +134,27 @@ impl Kcp {
pub fn send(&mut self, data: Bytes) -> Result<usize, Error> {
let ret = unsafe { ikcp_send(self.kcp, data.as_ptr() as *const i8, data.len() as i32) };
if ret < 0 {
return Err(Error::Unknown);
return Err(anyhow::anyhow!("send failed, return: {}", ret).into());
} else {
return Ok(ret as usize);
}
}
pub fn flush(&mut self) {
unsafe {
ikcp_flush(self.kcp);
}
}
pub fn peeksize(&self) -> i32 {
unsafe { ikcp_peeksize(self.kcp) }
}
pub fn recv(&mut self, buf: &mut BytesMut) -> Result<(), Error> {
let ret =
unsafe { ikcp_recv(self.kcp, buf.as_mut_ptr() as *mut i8, buf.capacity() as i32) };
if ret < 0 {
return Err(Error::Unknown);
return Err(anyhow::anyhow!("recv failed, return: {}", ret).into());
} else {
unsafe {
buf.set_len(ret as usize);
@@ -155,7 +180,7 @@ impl Kcp {
unsafe {
let ret = ikcp_setmtu(self.kcp, self.config.mtu.unwrap_or(1200));
if ret < 0 {
return Err(Error::Unknown);
return Err(anyhow::anyhow!("setmtu failed, return: {}", ret).into());
}
let ret = ikcp_wndsize(
@@ -164,7 +189,7 @@ impl Kcp {
self.config.rcvwnd.unwrap_or(-1),
);
if ret < 0 {
return Err(Error::Unknown);
return Err(anyhow::anyhow!("wndsize failed, return: {}", ret).into());
}
let ret = ikcp_nodelay(
@@ -175,7 +200,7 @@ impl Kcp {
self.config.nc.unwrap_or(-1),
);
if ret < 0 {
return Err(Error::Unknown);
return Err(anyhow::anyhow!("nodelay failed, return: {}", ret).into());
}
}
+4 -2
View File
@@ -1,6 +1,8 @@
mod ffi;
mod packet_def;
mod state;
pub mod error;
pub mod endpoint;
pub mod error;
pub mod ffi_safe;
pub mod packet_def;
pub mod stream;
+113 -15
View File
@@ -1,33 +1,45 @@
use std::fmt::Formatter;
use bytes::{Bytes, BytesMut};
use zerocopy::{AsBytes, FromBytes, FromZeroes};
use zerocopy::{AsBytes, FromBytes, FromZeroes, LittleEndian, U32};
bitflags::bitflags! {
#[derive(Debug)]
struct KcpPacketHeaderFlags: u8 {
const SYN = 0b0000_0001;
const ACK = 0b0000_0010;
const FIN = 0b0000_0100;
const DATA = 0b0000_1000;
const RST = 0b0001_0000;
const PING = 0b0010_0000;
const PONG = 0b0100_0000;
const _ = !0;
}
}
#[repr(C, packed)]
#[derive(AsBytes, FromBytes, FromZeroes, Clone, Debug, Default)]
#[derive(AsBytes, FromBytes, FromZeroes, Clone, Default)]
pub struct KcpPacketHeader {
conv: u32,
len: u16,
conv: U32<LittleEndian>,
src_session_id: U32<LittleEndian>,
dst_session_id: U32<LittleEndian>,
flag: u8,
rsv: u8,
}
impl KcpPacketHeader {
pub fn conv(&self) -> u32 {
self.conv
self.conv.into()
}
pub fn len(&self) -> u16 {
self.len
pub fn src_session_id(&self) -> u32 {
self.src_session_id.into()
}
pub fn dst_session_id(&self) -> u32 {
self.dst_session_id.into()
}
pub fn is_syn(&self) -> bool {
@@ -54,13 +66,36 @@ impl KcpPacketHeader {
.contains(KcpPacketHeaderFlags::DATA)
}
pub fn is_rst(&self) -> bool {
KcpPacketHeaderFlags::from_bits(self.flag)
.unwrap()
.contains(KcpPacketHeaderFlags::RST)
}
pub fn is_ping(&self) -> bool {
KcpPacketHeaderFlags::from_bits(self.flag)
.unwrap()
.contains(KcpPacketHeaderFlags::PING)
}
pub fn is_pong(&self) -> bool {
KcpPacketHeaderFlags::from_bits(self.flag)
.unwrap()
.contains(KcpPacketHeaderFlags::PONG)
}
pub fn set_conv(&mut self, conv: u32) -> &mut Self {
self.conv = conv;
self.conv = conv.into();
self
}
pub fn set_len(&mut self, len: u16) -> &mut Self {
self.len = len;
pub fn set_src_session_id(&mut self, session_id: u32) -> &mut Self {
self.src_session_id = session_id.into();
self
}
pub fn set_dst_session_id(&mut self, session_id: u32) -> &mut Self {
self.dst_session_id = session_id.into();
self
}
@@ -107,19 +142,78 @@ impl KcpPacketHeader {
self.flag = flags.bits();
self
}
pub fn set_rst(&mut self, rst: bool) -> &mut Self {
let mut flags = KcpPacketHeaderFlags::from_bits(self.flag).unwrap();
if rst {
flags.insert(KcpPacketHeaderFlags::RST);
} else {
flags.remove(KcpPacketHeaderFlags::RST);
}
self.flag = flags.bits();
self
}
pub fn set_ping(&mut self, ping: bool) -> &mut Self {
let mut flags = KcpPacketHeaderFlags::from_bits(self.flag).unwrap();
if ping {
flags.insert(KcpPacketHeaderFlags::PING);
} else {
flags.remove(KcpPacketHeaderFlags::PING);
}
self.flag = flags.bits();
self
}
pub fn set_pong(&mut self, pong: bool) -> &mut Self {
let mut flags = KcpPacketHeaderFlags::from_bits(self.flag).unwrap();
if pong {
flags.insert(KcpPacketHeaderFlags::PONG);
} else {
flags.remove(KcpPacketHeaderFlags::PONG);
}
self.flag = flags.bits();
self
}
}
#[derive(Debug)]
impl std::fmt::Debug for KcpPacketHeader {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
f.debug_struct("KcpPacketHeader")
.field("conv", &self.conv())
.field("src_session_id", &self.src_session_id())
.field("dst_session_id", &self.dst_session_id())
.field("flag", &KcpPacketHeaderFlags::from_bits(self.flag).unwrap())
.finish()
}
}
#[derive(Clone)]
pub struct KcpPacket {
inner: BytesMut,
}
impl Default for KcpPacket {
fn default() -> Self {
Self::new(0)
}
}
impl From<BytesMut> for KcpPacket {
fn from(inner: BytesMut) -> Self {
Self { inner }
}
}
impl std::fmt::Debug for KcpPacket {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
f.debug_struct("KcpPacket")
.field("header", &self.header())
.field("payload", &self.payload())
.finish()
}
}
impl Into<BytesMut> for KcpPacket {
fn into(self) -> BytesMut {
self.inner
@@ -133,10 +227,10 @@ impl Into<Bytes> for KcpPacket {
}
impl KcpPacket {
pub fn new(cap: Option<usize>) -> Self {
Self {
inner: BytesMut::with_capacity(cap.unwrap_or(std::mem::size_of::<KcpPacketHeader>())),
}
pub fn new(body_size: usize) -> Self {
let mut inner = BytesMut::with_capacity(std::mem::size_of::<KcpPacketHeader>() + body_size);
inner.resize(inner.capacity(), 0);
Self { inner }
}
pub fn new_with_payload(payload: &[u8]) -> Self {
@@ -158,4 +252,8 @@ impl KcpPacket {
pub fn payload(&self) -> &[u8] {
&self.inner[std::mem::size_of::<KcpPacketHeader>()..]
}
pub fn inner(self) -> BytesMut {
self.inner
}
}
+179
View File
@@ -0,0 +1,179 @@
use crate::error::Error;
#[auto_impl::auto_impl(&mut)]
pub trait PacketHeaderFlagManipulator: Default {
fn has_syn(&self) -> bool;
fn has_ack(&self) -> bool;
fn has_fin(&self) -> bool;
fn has_rst(&self) -> bool;
fn has_data(&self) -> bool;
fn set_syn(&mut self, value: bool);
fn set_ack(&mut self, value: bool);
fn set_fin(&mut self, value: bool);
fn set_rst(&mut self, value: bool);
fn set_data(&mut self, value: bool);
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum KcpConnectionFSM {
Closed,
// server start state
Listen,
SynReceived,
// client start state
SynSent,
// common states
Established,
// common active closing states
LocalClosed,
// common passive closing states
PeerClosed,
}
impl KcpConnectionFSM {
pub fn listen() -> Self {
KcpConnectionFSM::Listen
}
pub fn connect<P: PacketHeaderFlagManipulator>(out_packet: &mut P) -> Self {
out_packet.set_syn(true);
KcpConnectionFSM::SynSent
}
pub fn close<P: PacketHeaderFlagManipulator>(
&mut self,
out_packet: &mut P,
) -> Result<(), Error> {
out_packet.set_fin(true);
if matches!(self, KcpConnectionFSM::Established) {
*self = KcpConnectionFSM::LocalClosed;
} else {
*self = KcpConnectionFSM::Closed;
}
Ok(())
}
fn check_packet_flag<P: PacketHeaderFlagManipulator>(
packet: &P,
has_syn: bool,
has_ack: bool,
has_fin: bool,
has_rst: bool,
has_data: bool,
) -> bool {
packet.has_syn() == has_syn
&& packet.has_ack() == has_ack
&& packet.has_fin() == has_fin
&& packet.has_rst() == has_rst
&& packet.has_data() == has_data
}
fn do_handle_packet<P: PacketHeaderFlagManipulator>(
&mut self,
packet: &P,
out_packet: &mut Option<P>,
) -> Result<(), Error> {
let mut p = P::default();
match self {
KcpConnectionFSM::Closed => Err(Error::InvalidStateNeedRst),
KcpConnectionFSM::Listen => {
if Self::check_packet_flag(packet, true, false, false, false, false) {
p.set_syn(true);
p.set_ack(true);
out_packet.replace(p);
*self = KcpConnectionFSM::SynReceived;
Ok(())
} else {
Err(Error::InvalidStateNeedRst)
}
}
KcpConnectionFSM::SynReceived => {
// when client receives the syn-ack packet, all following packets should have ack+data flag
if Self::check_packet_flag(packet, false, true, false, false, true) {
*self = KcpConnectionFSM::Established;
Ok(())
} else if packet.has_rst() {
*self = KcpConnectionFSM::Closed;
Err(Error::InvalidState)
} else if packet.has_fin() {
*self = KcpConnectionFSM::Closed;
Err(Error::InvalidStateNeedRst)
} else {
Err(Error::InvalidStateNeedRst)
}
}
KcpConnectionFSM::SynSent => {
if Self::check_packet_flag(packet, true, true, false, false, false) {
p.set_ack(true);
p.set_data(true);
out_packet.replace(p);
*self = KcpConnectionFSM::Established;
Ok(())
} else if packet.has_rst() {
*self = KcpConnectionFSM::Closed;
Err(Error::InvalidState)
} else if packet.has_fin() {
*self = KcpConnectionFSM::Closed;
Err(Error::InvalidStateNeedRst)
} else {
Err(Error::InvalidStateNeedRst)
}
}
KcpConnectionFSM::Established => {
if Self::check_packet_flag(packet, false, true, false, false, true) {
Ok(())
} else if packet.has_rst() {
*self = KcpConnectionFSM::Closed;
Err(Error::InvalidState)
} else if packet.has_fin() {
*self = KcpConnectionFSM::PeerClosed;
Ok(())
} else {
Err(Error::InvalidStateNeedRst)
}
}
KcpConnectionFSM::LocalClosed => {
if packet.has_fin() {
*self = KcpConnectionFSM::Closed;
Ok(())
} else if packet.has_rst() {
*self = KcpConnectionFSM::Closed;
Err(Error::InvalidState)
} else if packet.has_data() {
Ok(())
} else {
Err(Error::InvalidState)
}
}
KcpConnectionFSM::PeerClosed => {
if packet.has_rst() {
*self = KcpConnectionFSM::Closed;
Err(Error::InvalidState)
} else {
Err(Error::InvalidState)
}
}
}
}
pub fn handle_packet<P: PacketHeaderFlagManipulator>(
&mut self,
packet: &P,
out_packet: &mut Option<P>,
) -> Result<(), Error> {
let ret = self.do_handle_packet(packet, out_packet);
if matches!(ret, Err(Error::InvalidStateNeedRst)) {
let mut p = P::default();
p.set_rst(true);
out_packet.replace(p);
}
ret
}
}
+93
View File
@@ -0,0 +1,93 @@
use std::{
pin::Pin,
task::ready,
task::{Context, Poll},
};
use bytes::{Bytes, BytesMut};
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
use crate::endpoint::{ConnId, KcpEndpoint, KcpStreamReceiver};
pub struct KcpStream {
sender: tokio_util::sync::PollSender<BytesMut>,
receiver: KcpStreamReceiver,
conn_id: ConnId,
conn_data: Bytes,
}
impl std::fmt::Debug for KcpStream {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("KcpStream")
.field("conn_id", &self.conn_id)
.finish()
}
}
impl KcpStream {
pub fn new(endpoint: &KcpEndpoint, conn_id: ConnId) -> Option<Self> {
let (sender, receiver) = endpoint.conn_sender_receiver(conn_id)?;
let conn_data = endpoint.conn_data(&conn_id)?;
Some(Self {
sender: tokio_util::sync::PollSender::new(sender),
receiver,
conn_id,
conn_data,
})
}
pub fn conn_data(&self) -> &Bytes {
&self.conn_data
}
pub fn conn_id(&self) -> ConnId {
self.conn_id
}
}
impl AsyncRead for KcpStream {
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut Context,
buf: &mut ReadBuf,
) -> Poll<std::io::Result<()>> {
let Some(read_buf) = ready!(self.receiver.poll_recv(cx)) else {
return Poll::Ready(Err(std::io::Error::new(
std::io::ErrorKind::BrokenPipe,
"stream closed",
)));
};
buf.put_slice(&read_buf);
Poll::Ready(Ok(()))
}
}
impl AsyncWrite for KcpStream {
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context,
buf: &[u8],
) -> Poll<std::io::Result<usize>> {
let mut ret = ready!(self.sender.poll_reserve(cx));
if ret.is_ok() {
ret = self.sender.send_item(BytesMut::from(buf));
}
match ret {
Ok(_) => Poll::Ready(Ok(buf.len())),
Err(_) => Poll::Ready(Err(std::io::Error::new(
std::io::ErrorKind::BrokenPipe,
"stream closed",
))),
}
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context) -> Poll<std::io::Result<()>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(mut self: Pin<&mut Self>, _cx: &mut Context) -> Poll<std::io::Result<()>> {
self.sender.close();
Poll::Ready(Ok(()))
}
}