diff --git a/Cargo.toml b/Cargo.toml index abb5e83..ceec139 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -10,6 +10,12 @@ license = "MIT" keywords = ["kcp", "bindings"] [dependencies] +bytes = "1" +tokio = { version = "1", features = ["full"] } +dashmap = "6.1.0" +zerocopy = { version = "0.7", features = ["derive", "simd"] } +bitflags = "2.5" +parking_lot = "0.12.3" [build-dependencies] bindgen = "0.71.1" diff --git a/src/endpoint.rs b/src/endpoint.rs new file mode 100644 index 0000000..7a50fdb --- /dev/null +++ b/src/endpoint.rs @@ -0,0 +1,325 @@ +use std::sync::{ + atomic::{AtomicBool, AtomicU32}, + Arc, +}; + +use bytes::BytesMut; +use dashmap::DashMap; +use parking_lot::Mutex; +use tokio::{select, sync::Notify, task::JoinSet}; + +use crate::{ + error::Error, + ffi_safe::{Kcp, KcpConfig}, + packet_def::KcpPacket, +}; + +pub type Sender = tokio::sync::mpsc::Sender; +pub type Receiver = tokio::sync::mpsc::Receiver; + +pub type KcpPakcetSender = Sender; +pub type KcpPacketReceiver = Receiver; + +pub type KcpStreamSender = Sender; +pub type KcpStreamReceiver = Receiver; + +enum KcpClientState { + SynSent, + Established, + Fin, +} + +enum KcpServerState { + SynReceived, + Established, + Fin, +} + +enum KcpConnectionState { + Client(KcpClientState), + Server(KcpServerState), +} + +struct KcpConnectionInner { + conv: u32, + + update_notifier: Notify, + recv_notifier: Notify, + send_notifier: Notify, + + has_new_input: AtomicBool, + waiting_new_send_window: AtomicBool, +} + +struct KcpConnection { + conv: u32, + kcp: Arc>>, + + inner: Arc, + + send_sender: Sender, + send_receiver: Option>, + + recv_sender: Sender, + recv_receiver: Option>, + + tasks: JoinSet<()>, +} + +impl KcpConnection { + fn new(conv: u32) -> Result { + let kcp = Kcp::new(KcpConfig::new(conv))?; + + let (send_sender, send_receiver) = tokio::sync::mpsc::channel(16); + let (recv_sender, recv_receiver) = tokio::sync::mpsc::channel(16); + + Ok(Self { + conv, + kcp: Arc::new(Mutex::new(kcp)), + + inner: Arc::new(KcpConnectionInner { + conv, + + update_notifier: Notify::new(), + recv_notifier: Notify::new(), + send_notifier: Notify::new(), + + has_new_input: AtomicBool::new(false), + waiting_new_send_window: AtomicBool::new(false), + }), + + send_sender, + send_receiver: Some(send_receiver), + + recv_sender, + recv_receiver: Some(recv_receiver), + + tasks: JoinSet::new(), + }) + } + + fn run(&mut self, output_sender: KcpPakcetSender) { + 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); + Ok(()) + })); + + // kcp updater + let inner = self.inner.clone(); + let kcp = self.kcp.clone(); + self.tasks.spawn(async move { + loop { + let next_update_ms = kcp.lock().next_update_delay_ms(); + select! { + _ = tokio::time::sleep(tokio::time::Duration::from_millis(next_update_ms as u64)) => {} + _ = inner.update_notifier.notified() => {} + } + + kcp.lock().update(); + + if inner.has_new_input.swap(false, std::sync::atomic::Ordering::SeqCst) { + inner.recv_notifier.notify_one(); + } + + if inner.waiting_new_send_window.swap(false, std::sync::atomic::Ordering::SeqCst) { + inner.send_notifier.notify_one(); + } + } + }); + + // handle packet send + 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; + } + } + kcp.lock().send(data.freeze()).unwrap(); + } + }); + + // 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); + assert_ne!(0, buf.len()); + let send_ret = recv_sender.send(buf.split()).await; + if let Err(_) = send_ret { + break; + } + } + } + }); + } + + fn handle_input(&mut self, packet: KcpPacket) { + let _ = 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(); + } + + fn send_sender(&self) -> KcpStreamSender { + self.send_sender.clone() + } + + fn recv_receiver(&mut self) -> KcpStreamReceiver { + self.recv_receiver.take().unwrap() + } +} + +pub struct KcpEndpoint { + cur_conv: AtomicU32, + established_map: Arc>, + + input_sender: KcpPakcetSender, + input_receiver: Option, + + output_sender: KcpPakcetSender, + output_receiver: Option, + + tasks: JoinSet<()>, +} + +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); + + Self { + cur_conv: AtomicU32::new(0), + established_map: Arc::new(DashMap::new()), + + input_sender, + input_receiver: Some(input_receiver), + + output_sender, + output_receiver: Some(output_receiver), + + tasks: JoinSet::new(), + } + } + + pub async fn run(&mut self) { + let mut input_receiver = self.input_receiver.take().unwrap(); + let established_map = self.established_map.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 { + continue; + }; + let _ = conn.handle_input(packet); + } + } + }); + } + + pub fn add_established(&mut self, conv: u32) -> Result<(), Error> { + let mut conn = KcpConnection::new(conv)?; + conn.run(self.output_sender.clone()); + + self.established_map.insert(conv, conn); + + Ok(()) + } + + pub fn output_receiver(&mut self) -> Option { + self.output_receiver.take() + } + + pub fn input_sender(&self) -> KcpPakcetSender { + self.input_sender.clone() + } + + pub fn conn_sender_receiver(&self, conv: u32) -> Option<(KcpStreamSender, KcpStreamReceiver)> { + let mut conn = self.established_map.get_mut(&conv)?; + Some((conn.send_sender(), conn.recv_receiver())) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn test_kcp_endpoint() { + let mut client_endpoint = KcpEndpoint::new(); + let mut server_endpoint = KcpEndpoint::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 { + while let Some(packet) = server_output_receiver.recv().await { + let _ = client_input_sender.send(packet).await; + } + }); + + let server_input_sender = server_endpoint.input_sender(); + let mut client_output_receiver = client_endpoint.output_receiver().unwrap(); + let t2 = tokio::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_sender.send(BytesMut::from("hello")).await.unwrap(); + let data = server_receiver.recv().await.unwrap(); + assert_eq!("hello", String::from_utf8_lossy(&data)); + + server_sender.send(BytesMut::from("world")).await.unwrap(); + let data = client_receiver.recv().await.unwrap(); + assert_eq!("world", String::from_utf8_lossy(&data)); + + drop(client_endpoint); + drop(server_endpoint); + + let _ = tokio::join!(t1, t2); + } +} diff --git a/src/error.rs b/src/error.rs new file mode 100644 index 0000000..bb8a9a1 --- /dev/null +++ b/src/error.rs @@ -0,0 +1,5 @@ +#[derive(Debug)] +pub enum Error { + CreateConnectionFailed, + Unknown, +} diff --git a/src/ffi.rs b/src/ffi.rs new file mode 100644 index 0000000..a38a13a --- /dev/null +++ b/src/ffi.rs @@ -0,0 +1,5 @@ +#![allow(non_upper_case_globals)] +#![allow(non_camel_case_types)] +#![allow(non_snake_case)] + +include!(concat!(env!("OUT_DIR"), "/bindings.rs")); diff --git a/src/ffi_safe.rs b/src/ffi_safe.rs new file mode 100644 index 0000000..14c5308 --- /dev/null +++ b/src/ffi_safe.rs @@ -0,0 +1,192 @@ +use crate::{error::Error, ffi::*}; +use std::time::Instant; + +use bytes::{Bytes, BytesMut}; + +#[derive(Debug, Clone, Copy)] +pub struct KcpConfig { + pub conv: IUINT32, + pub mtu: Option, + + pub sndwnd: Option, + pub rcvwnd: Option, + + pub nodelay: Option, + pub interval: Option, + pub resend: Option, + pub nc: Option, +} + +impl KcpConfig { + pub fn new(conv: IUINT32) -> Self { + Self { + conv, + mtu: None, + sndwnd: None, + rcvwnd: None, + nodelay: None, + interval: None, + resend: None, + nc: None, + } + } +} + +pub type OutputCb = Box Result<(), Error>>; + +pub struct Kcp { + kcp: *mut ikcpcb, + config: KcpConfig, + now: Instant, + output_cb: Option Result<(), Error>>>, + + _marker: core::marker::PhantomData<(*mut u8, core::marker::PhantomPinned)>, +} + +unsafe impl Send for Kcp {} + +unsafe extern "C" fn ikcp_output( + buf: *const ::std::os::raw::c_char, + len: ::std::os::raw::c_int, + kcp: *mut ikcpcb, + this: *mut ::std::os::raw::c_void, +) -> i32 { + // convert this to KcpConnection + let kcp_connection = &mut *(this as *mut Kcp); + assert_eq!(kcp_connection.kcp, kcp); + + let buf = BytesMut::from(std::slice::from_raw_parts(buf as *const u8, len as usize)); + + // TODO: handle output error + let _ = kcp_connection.handle_output_callback(buf); + + // kcp doesn't care about the return value + 0 +} + +impl Kcp { + pub fn new(config: KcpConfig) -> Result, Error> { + unsafe { + let conv = config.conv; + let mut ret = Box::new(Self { + kcp: std::ptr::null_mut(), + config, + now: Instant::now(), + output_cb: None, + _marker: core::marker::PhantomData, + }); + + let kcp = ikcp_create(conv, &mut *ret as *mut Kcp as *mut ::std::os::raw::c_void); + if kcp.is_null() { + return Err(Error::CreateConnectionFailed); + } + + ret.kcp = kcp; + + ikcp_setoutput(kcp, Some(ikcp_output)); + + ret.apply_config()?; + + return Ok(ret); + } + } + + pub fn set_output_cb(&mut self, output_cb: OutputCb) { + self.output_cb = Some(output_cb); + } + + 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); + } else { + return Ok(()); + } + } + + pub fn update(&mut self) { + unsafe { + ikcp_update(self.kcp, self.now.elapsed().as_millis() as IUINT32); + } + } + + pub fn next_update_delay_ms(&mut self) -> IUINT32 { + let current = self.now.elapsed().as_millis() as IUINT32; + let next = unsafe { ikcp_check(self.kcp, current) }; + next - current + } + + 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); + } else { + return Ok(ret as usize); + } + } + + 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); + } else { + unsafe { + buf.set_len(ret as usize); + } + return Ok(()); + } + } + + pub fn waitsnd(&self) -> i32 { + unsafe { ikcp_waitsnd(self.kcp) } + } + + pub fn sendwnd(&self) -> i32 { + // see IKCP_WND_SND + self.config.sndwnd.unwrap_or(32) + } + + fn handle_output_callback(&self, buf: BytesMut) -> Result<(), Error> { + (self.output_cb.as_ref().unwrap())(self.config.conv, buf) + } + + fn apply_config(&mut self) -> Result<(), Error> { + unsafe { + let ret = ikcp_setmtu(self.kcp, self.config.mtu.unwrap_or(1200)); + if ret < 0 { + return Err(Error::Unknown); + } + + let ret = ikcp_wndsize( + self.kcp, + self.config.sndwnd.unwrap_or(-1), + self.config.rcvwnd.unwrap_or(-1), + ); + if ret < 0 { + return Err(Error::Unknown); + } + + let ret = ikcp_nodelay( + self.kcp, + self.config.nodelay.unwrap_or(-1), + self.config.interval.unwrap_or(-1), + self.config.resend.unwrap_or(-1), + self.config.nc.unwrap_or(-1), + ); + if ret < 0 { + return Err(Error::Unknown); + } + } + + return Ok(()); + } +} + +impl Drop for Kcp { + fn drop(&mut self) { + unsafe { + ikcp_release(self.kcp); + } + } +} diff --git a/src/lib.rs b/src/lib.rs index a38a13a..c6d3f71 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,5 +1,6 @@ -#![allow(non_upper_case_globals)] -#![allow(non_camel_case_types)] -#![allow(non_snake_case)] +mod ffi; +mod packet_def; -include!(concat!(env!("OUT_DIR"), "/bindings.rs")); +pub mod error; +pub mod endpoint; +pub mod ffi_safe; diff --git a/src/packet_def.rs b/src/packet_def.rs new file mode 100644 index 0000000..e92ef9b --- /dev/null +++ b/src/packet_def.rs @@ -0,0 +1,161 @@ +use bytes::{Bytes, BytesMut}; +use zerocopy::{AsBytes, FromBytes, FromZeroes}; + +bitflags::bitflags! { + struct KcpPacketHeaderFlags: u8 { + const SYN = 0b0000_0001; + const ACK = 0b0000_0010; + const FIN = 0b0000_0100; + const DATA = 0b0000_1000; + + const _ = !0; + } +} + +#[repr(C, packed)] +#[derive(AsBytes, FromBytes, FromZeroes, Clone, Debug, Default)] +pub struct KcpPacketHeader { + conv: u32, + len: u16, + flag: u8, + rsv: u8, +} + +impl KcpPacketHeader { + pub fn conv(&self) -> u32 { + self.conv + } + + pub fn len(&self) -> u16 { + self.len + } + + pub fn is_syn(&self) -> bool { + KcpPacketHeaderFlags::from_bits(self.flag) + .unwrap() + .contains(KcpPacketHeaderFlags::SYN) + } + + pub fn is_ack(&self) -> bool { + KcpPacketHeaderFlags::from_bits(self.flag) + .unwrap() + .contains(KcpPacketHeaderFlags::ACK) + } + + pub fn is_fin(&self) -> bool { + KcpPacketHeaderFlags::from_bits(self.flag) + .unwrap() + .contains(KcpPacketHeaderFlags::FIN) + } + + pub fn is_data(&self) -> bool { + KcpPacketHeaderFlags::from_bits(self.flag) + .unwrap() + .contains(KcpPacketHeaderFlags::DATA) + } + + pub fn set_conv(&mut self, conv: u32) -> &mut Self { + self.conv = conv; + self + } + + pub fn set_len(&mut self, len: u16) -> &mut Self { + self.len = len; + self + } + + pub fn set_syn(&mut self, syn: bool) -> &mut Self { + let mut flags = KcpPacketHeaderFlags::from_bits(self.flag).unwrap(); + if syn { + flags.insert(KcpPacketHeaderFlags::SYN); + } else { + flags.remove(KcpPacketHeaderFlags::SYN); + } + self.flag = flags.bits(); + self + } + + pub fn set_ack(&mut self, ack: bool) -> &mut Self { + let mut flags = KcpPacketHeaderFlags::from_bits(self.flag).unwrap(); + if ack { + flags.insert(KcpPacketHeaderFlags::ACK); + } else { + flags.remove(KcpPacketHeaderFlags::ACK); + } + self.flag = flags.bits(); + self + } + + pub fn set_fin(&mut self, fin: bool) -> &mut Self { + let mut flags = KcpPacketHeaderFlags::from_bits(self.flag).unwrap(); + if fin { + flags.insert(KcpPacketHeaderFlags::FIN); + } else { + flags.remove(KcpPacketHeaderFlags::FIN); + } + self.flag = flags.bits(); + self + } + + pub fn set_data(&mut self, data: bool) -> &mut Self { + let mut flags = KcpPacketHeaderFlags::from_bits(self.flag).unwrap(); + if data { + flags.insert(KcpPacketHeaderFlags::DATA); + } else { + flags.remove(KcpPacketHeaderFlags::DATA); + } + self.flag = flags.bits(); + self + } +} + +#[derive(Debug)] +pub struct KcpPacket { + inner: BytesMut, +} + +impl From for KcpPacket { + fn from(inner: BytesMut) -> Self { + Self { inner } + } +} + +impl Into for KcpPacket { + fn into(self) -> BytesMut { + self.inner + } +} + +impl Into for KcpPacket { + fn into(self) -> Bytes { + self.inner.freeze() + } +} + +impl KcpPacket { + pub fn new(cap: Option) -> Self { + Self { + inner: BytesMut::with_capacity(cap.unwrap_or(std::mem::size_of::())), + } + } + + pub fn new_with_payload(payload: &[u8]) -> Self { + let mut inner = + BytesMut::with_capacity(std::mem::size_of::() + payload.len()); + inner.resize(std::mem::size_of::(), 0); + inner.extend_from_slice(payload); + Self { inner } + } + + pub fn mut_header(&mut self) -> &mut KcpPacketHeader { + KcpPacketHeader::mut_from_prefix(&mut self.inner).unwrap() + } + + pub fn header(&self) -> &KcpPacketHeader { + KcpPacketHeader::ref_from_prefix(&self.inner).unwrap() + } + + pub fn payload(&self) -> &[u8] { + &self.inner[std::mem::size_of::()..] + } +}