kcp pingpong works

This commit is contained in:
sijie.sun
2025-01-12 01:23:30 +08:00
parent 519386124c
commit b30e472539
7 changed files with 699 additions and 4 deletions
+6
View File
@@ -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"
+325
View File
@@ -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<T> = tokio::sync::mpsc::Sender<T>;
pub type Receiver<T> = tokio::sync::mpsc::Receiver<T>;
pub type KcpPakcetSender = Sender<KcpPacket>;
pub type KcpPacketReceiver = Receiver<KcpPacket>;
pub type KcpStreamSender = Sender<BytesMut>;
pub type KcpStreamReceiver = Receiver<BytesMut>;
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<Mutex<Box<Kcp>>>,
inner: Arc<KcpConnectionInner>,
send_sender: Sender<BytesMut>,
send_receiver: Option<Receiver<BytesMut>>,
recv_sender: Sender<BytesMut>,
recv_receiver: Option<Receiver<BytesMut>>,
tasks: JoinSet<()>,
}
impl KcpConnection {
fn new(conv: u32) -> Result<Self, Error> {
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<DashMap<u32, KcpConnection>>,
input_sender: KcpPakcetSender,
input_receiver: Option<KcpPacketReceiver>,
output_sender: KcpPakcetSender,
output_receiver: Option<KcpPacketReceiver>,
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<KcpPacketReceiver> {
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);
}
}
+5
View File
@@ -0,0 +1,5 @@
#[derive(Debug)]
pub enum Error {
CreateConnectionFailed,
Unknown,
}
+5
View File
@@ -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"));
+192
View File
@@ -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<i32>,
pub sndwnd: Option<i32>,
pub rcvwnd: Option<i32>,
pub nodelay: Option<i32>,
pub interval: Option<i32>,
pub resend: Option<i32>,
pub nc: Option<i32>,
}
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<dyn Fn(u32, BytesMut) -> Result<(), Error>>;
pub struct Kcp {
kcp: *mut ikcpcb,
config: KcpConfig,
now: Instant,
output_cb: Option<Box<dyn Fn(u32, BytesMut) -> 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<Box<Self>, 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<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);
} 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);
}
}
}
+5 -4
View File
@@ -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;
+161
View File
@@ -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<BytesMut> for KcpPacket {
fn from(inner: BytesMut) -> Self {
Self { inner }
}
}
impl Into<BytesMut> for KcpPacket {
fn into(self) -> BytesMut {
self.inner
}
}
impl Into<Bytes> for KcpPacket {
fn into(self) -> Bytes {
self.inner.freeze()
}
}
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_with_payload(payload: &[u8]) -> Self {
let mut inner =
BytesMut::with_capacity(std::mem::size_of::<KcpPacketHeader>() + payload.len());
inner.resize(std::mem::size_of::<KcpPacketHeader>(), 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::<KcpPacketHeader>()..]
}
}