mirror of
https://github.com/EasyTier/kcp-sys.git
synced 2025-05-19 10:31:09 +00:00
support connection management
This commit is contained in:
+9
-2
@@ -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"
|
||||
|
||||
@@ -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
@@ -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
@@ -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,3 +1,4 @@
|
||||
#![allow(unused)]
|
||||
#![allow(non_upper_case_globals)]
|
||||
#![allow(non_camel_case_types)]
|
||||
#![allow(non_snake_case)]
|
||||
|
||||
+31
-6
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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(()))
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user