From a5f4b4e3a0891b36ba4df77e412cc6c1bec0d9d8 Mon Sep 17 00:00:00 2001 From: "sijie.sun" Date: Tue, 14 Jan 2025 14:10:18 +0800 Subject: [PATCH] support connection management --- Cargo.toml | 11 +- build.rs | 1 - src/endpoint.rs | 724 ++++++++++++++++++++++++++++++++++++++++------ src/error.rs | 18 +- src/ffi.rs | 1 + src/ffi_safe.rs | 37 ++- src/lib.rs | 6 +- src/packet_def.rs | 128 +++++++- src/state.rs | 179 ++++++++++++ src/stream.rs | 93 ++++++ 10 files changed, 1074 insertions(+), 124 deletions(-) create mode 100644 src/state.rs create mode 100644 src/stream.rs diff --git a/Cargo.toml b/Cargo.toml index ceec139..84f22d8 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -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" diff --git a/build.rs b/build.rs index adf26a8..a94ae8a 100644 --- a/build.rs +++ b/build.rs @@ -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"); diff --git a/src/endpoint.rs b/src/endpoint.rs index 7a50fdb..0126378 100644 --- a/src/endpoint.rs +++ b/src/endpoint.rs @@ -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 = tokio::sync::mpsc::Sender; @@ -23,26 +26,34 @@ pub type KcpPacketReceiver = Receiver; pub type KcpStreamSender = Sender; pub type KcpStreamReceiver = Receiver; -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>>, inner: Arc, - send_sender: Sender, + send_sender: Option>, send_receiver: Option>, - recv_sender: Sender, + recv_sender: Option>, recv_receiver: Option>, + send_close_notifier: Arc, + recv_closed: Arc, + tasks: JoinSet<()>, } impl KcpConnection { - fn new(conv: u32) -> Result { - let kcp = Kcp::new(KcpConfig::new(conv))?; + pub fn new(conn_id: ConnId) -> Result { + 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 { + 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, + 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, 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 { + 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, + state_map: DashMap, +} + +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>, + id: u64, + data: Arc, input_sender: KcpPakcetSender, input_receiver: Option, @@ -213,17 +408,27 @@ pub struct KcpEndpoint { output_sender: KcpPakcetSender, output_receiver: Option, + new_conn_sender: tokio::sync::mpsc::Sender, + new_conn_receiver: Arc>>, + 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 = 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::>(); + + 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 { + 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 { + 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 { + 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; } } diff --git a/src/error.rs b/src/error.rs index bb8a9a1..fc49f72 100644 --- a/src/error.rs +++ b/src/error.rs @@ -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, } diff --git a/src/ffi.rs b/src/ffi.rs index a38a13a..81258b6 100644 --- a/src/ffi.rs +++ b/src/ffi.rs @@ -1,3 +1,4 @@ +#![allow(unused)] #![allow(non_upper_case_globals)] #![allow(non_camel_case_types)] #![allow(non_snake_case)] diff --git a/src/ffi_safe.rs b/src/ffi_safe.rs index 14c5308..d4b7ab7 100644 --- a/src/ffi_safe.rs +++ b/src/ffi_safe.rs @@ -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 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 { 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()); } } diff --git a/src/lib.rs b/src/lib.rs index c6d3f71..6cb89fb 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -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; diff --git a/src/packet_def.rs b/src/packet_def.rs index e92ef9b..731bb0e 100644 --- a/src/packet_def.rs +++ b/src/packet_def.rs @@ -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, + src_session_id: U32, + dst_session_id: U32, 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 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 for KcpPacket { fn into(self) -> BytesMut { self.inner @@ -133,10 +227,10 @@ impl Into for KcpPacket { } impl KcpPacket { - pub fn new(cap: Option) -> Self { - Self { - inner: BytesMut::with_capacity(cap.unwrap_or(std::mem::size_of::())), - } + pub fn new(body_size: usize) -> Self { + let mut inner = BytesMut::with_capacity(std::mem::size_of::() + 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::()..] } + + pub fn inner(self) -> BytesMut { + self.inner + } } diff --git a/src/state.rs b/src/state.rs new file mode 100644 index 0000000..09c20d4 --- /dev/null +++ b/src/state.rs @@ -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(out_packet: &mut P) -> Self { + out_packet.set_syn(true); + KcpConnectionFSM::SynSent + } + + pub fn close( + &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( + 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( + &mut self, + packet: &P, + out_packet: &mut Option

, + ) -> 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( + &mut self, + packet: &P, + out_packet: &mut Option

, + ) -> 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 + } +} diff --git a/src/stream.rs b/src/stream.rs new file mode 100644 index 0000000..ba6c397 --- /dev/null +++ b/src/stream.rs @@ -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, + 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 { + 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> { + 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> { + 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> { + Poll::Ready(Ok(())) + } + + fn poll_shutdown(mut self: Pin<&mut Self>, _cx: &mut Context) -> Poll> { + self.sender.close(); + Poll::Ready(Ok(())) + } +}