From bf9b81e669f4ce98c8d58df5d776ae90726d7ec2 Mon Sep 17 00:00:00 2001 From: zirkonya Date: Sat, 29 Aug 2026 22:11:23 +0200 Subject: Add transport layer --- src/transport/connection.rs | 36 ++++++++ src/transport/error.rs | 19 +++++ src/transport/listener.rs | 21 +++++ src/transport/receiver.rs | 5 ++ src/transport/sender.rs | 5 ++ src/transport/tcp.rs | 143 +++++++++++++++++++++++++++++++ src/transport/udp.rs | 203 ++++++++++++++++++++++++++++++++++++++++++++ 7 files changed, 432 insertions(+) create mode 100644 src/transport/connection.rs create mode 100644 src/transport/error.rs create mode 100644 src/transport/listener.rs create mode 100644 src/transport/receiver.rs create mode 100644 src/transport/sender.rs create mode 100644 src/transport/tcp.rs create mode 100644 src/transport/udp.rs (limited to 'src/transport') diff --git a/src/transport/connection.rs b/src/transport/connection.rs new file mode 100644 index 0000000..34366a2 --- /dev/null +++ b/src/transport/connection.rs @@ -0,0 +1,36 @@ +use std::marker::PhantomData; + +use getset::Getters; + +use crate::transport::{receiver::PacketReceiver, sender::PacketSender}; + +#[derive(Getters)] +pub struct Connection +where + Uid: PartialEq, + S: PacketSender + Clone, + R: PacketReceiver, +{ + #[get = "pub"] + sender: S, + #[get = "pub"] + receiver: R, + _uid: PhantomData, + _data: PhantomData, +} + +impl Connection +where + Uid: PartialEq, + S: PacketSender + Clone, + R: PacketReceiver, +{ + pub fn new(sender: S, receiver: R) -> Self { + Self { + sender, + receiver, + _uid: PhantomData, + _data: PhantomData, + } + } +} diff --git a/src/transport/error.rs b/src/transport/error.rs new file mode 100644 index 0000000..41f9cf2 --- /dev/null +++ b/src/transport/error.rs @@ -0,0 +1,19 @@ +use thiserror::Error; + +#[derive(Debug, Error)] +pub enum TransportError { + #[error("io error: {0}")] + Io(#[from] std::io::Error), + + #[error("codec error: {0}")] + Codec(#[from] crate::codec::error::CodecError), + + #[error("receiver channel closed")] + ChannelClosed, + + #[error("receiver lock poisoned")] + LockPoisoned, + + #[error("{0}")] + Other(String), +} diff --git a/src/transport/listener.rs b/src/transport/listener.rs new file mode 100644 index 0000000..e9cb198 --- /dev/null +++ b/src/transport/listener.rs @@ -0,0 +1,21 @@ +use crate::transport::{ + Result, connection::Connection, receiver::PacketReceiver, sender::PacketSender, +}; + +/// Convenience alias for the [`Connection`] returned by a [`Listener`]. +pub type ListenerConnection = Connection< + Uid, + >::Data, + >::Sender, + >::Receiver, +>; + +pub trait Listener +where + Uid: PartialEq, +{ + type Data; + type Sender: PacketSender + Clone; + type Receiver: PacketReceiver; + fn accept(&mut self) -> Result>; +} diff --git a/src/transport/receiver.rs b/src/transport/receiver.rs new file mode 100644 index 0000000..3a7c948 --- /dev/null +++ b/src/transport/receiver.rs @@ -0,0 +1,5 @@ +use crate::{codec::Codec, context::Context, transport::Result}; + +pub trait PacketReceiver: Send + Sync { + fn recv>(&self, ctx: &Context) -> Result<(Uid, P)>; +} diff --git a/src/transport/sender.rs b/src/transport/sender.rs new file mode 100644 index 0000000..804914f --- /dev/null +++ b/src/transport/sender.rs @@ -0,0 +1,5 @@ +use crate::{codec::Codec, context::Context, transport::Result}; + +pub trait PacketSender: Send + Sync { + fn send>(&self, packet: P, ctx: &Context) -> Result<()>; +} diff --git a/src/transport/tcp.rs b/src/transport/tcp.rs new file mode 100644 index 0000000..19d02d6 --- /dev/null +++ b/src/transport/tcp.rs @@ -0,0 +1,143 @@ +use std::hash::Hash; +use std::io::Write; +use std::marker::PhantomData; +use std::net::{SocketAddr, TcpListener as StdTcpListener, TcpStream}; +use std::sync::atomic::{AtomicU64, Ordering}; +use std::sync::{Arc, Mutex}; + +use crate::codec::Codec; +use crate::context::Context; +use crate::transport::connection::Connection; +use crate::transport::error::TransportError; +use crate::transport::listener::Listener; +use crate::transport::receiver::PacketReceiver; +use crate::transport::sender::PacketSender; +use crate::transport::{Result, Transport}; + +// TODO : check Tcp Transport + +#[derive(Clone)] +pub struct TcpSender { + stream: Arc>, +} + +impl PacketSender for TcpSender { + fn send>(&self, packet: P, ctx: &Context) -> Result<()> { + let mut buf = Vec::new(); + packet.encode(&mut buf, ctx)?; + let mut guard = self + .stream + .lock() + .map_err(|_| TransportError::LockPoisoned)?; + guard.write_all(&buf)?; + guard.flush()?; + Ok(()) + } +} + +pub struct TcpReceiver { + stream: Arc>, + uid: Uid, +} + +impl PacketReceiver for TcpReceiver +where + Uid: PartialEq + Clone + Send + Sync, +{ + fn recv>(&self, ctx: &Context) -> Result<(Uid, P)> { + let mut guard = self + .stream + .lock() + .map_err(|_| TransportError::LockPoisoned)?; + let packet = P::decode(&mut *guard, ctx)?; + Ok((self.uid.clone(), packet)) + } +} + +pub struct TcpListener { + listener: Arc, + factory: Arc Uid + Send + Sync>, + _p: PhantomData, +} + +impl Listener for TcpListener +where + Uid: PartialEq + Eq + Hash + Clone + Send + Sync, + Data: Send + Sync, +{ + type Data = Data; + type Sender = TcpSender; + type Receiver = TcpReceiver; + + fn accept(&mut self) -> Result>> { + let (stream, _peer) = self.listener.accept()?; + let stream = Arc::new(Mutex::new(stream)); + let uid = (self.factory)(); + + let sender = TcpSender { + stream: stream.clone(), + }; + let receiver = TcpReceiver { stream, uid }; + Ok(Connection::new(sender, receiver)) + } +} + +pub struct TcpTransport { + factory: Arc Uid + Send + Sync>, + _p: PhantomData<(Uid, Data)>, +} + +impl TcpTransport +where + Uid: PartialEq + Eq + Hash + Clone + Send + Sync, +{ + pub fn new Uid + Send + Sync + 'static>(factory: F) -> Self { + Self { + factory: Arc::new(factory), + _p: PhantomData, + } + } +} + +impl Default for TcpTransport { + fn default() -> Self { + let counter = Arc::new(AtomicU64::new(0)); + TcpTransport::new(move || counter.fetch_add(1, Ordering::Relaxed)) + } +} + +impl Transport for TcpTransport +where + Uid: PartialEq + Eq + Hash + Clone + Send + Sync, + Data: Send + Sync, +{ + type Data = Data; + type Addr = SocketAddr; + type Sender = TcpSender; + type Receiver = TcpReceiver; + type Listener = TcpListener; + + fn connect( + &self, + addr: SocketAddr, + ) -> Result>> { + let stream = TcpStream::connect(addr)?; + let stream = Arc::new(Mutex::new(stream)); + let uid = (self.factory)(); + + let sender = TcpSender { + stream: stream.clone(), + }; + let receiver = TcpReceiver { stream, uid }; + Ok(Connection::new(sender, receiver)) + } + + fn listen(&self, addr: SocketAddr) -> Result> { + let listener = StdTcpListener::bind(addr)?; + Ok(TcpListener { + listener: Arc::new(listener), + factory: self.factory.clone(), + _p: PhantomData, + }) + } +} diff --git a/src/transport/udp.rs b/src/transport/udp.rs new file mode 100644 index 0000000..acb36f2 --- /dev/null +++ b/src/transport/udp.rs @@ -0,0 +1,203 @@ +use std::collections::HashMap; +use std::hash::Hash; +use std::marker::PhantomData; +use std::net::{SocketAddr, UdpSocket}; +use std::sync::atomic::{AtomicU64, Ordering}; +use std::sync::mpsc; +use std::sync::{Arc, Mutex}; + +use crate::context::Context; +use crate::transport::Result; +use crate::transport::connection::Connection; +use crate::transport::listener::Listener; +use crate::transport::receiver::PacketReceiver; +use crate::transport::sender::PacketSender; +use crate::transport::{Transport, TransportError}; + +// TODO : check Udp Transport + +#[derive(Clone)] +pub struct UdpPeerSender { + sock: Arc, + peer: SocketAddr, +} + +impl PacketSender for UdpPeerSender { + fn send>(&self, packet: P, ctx: &Context) -> Result<()> { + let mut buf = Vec::new(); + packet.encode(&mut buf, ctx)?; + self.sock.send_to(&buf, self.peer)?; + Ok(()) + } +} + +enum UdpReceiverInner { + Channel(Mutex>>), + Socket(Arc), +} + +pub struct UdpReceiver { + uid: Uid, + inner: UdpReceiverInner, +} + +impl UdpReceiver { + fn channel(rx: Mutex>>, uid: Uid) -> Self { + Self { + uid, + inner: UdpReceiverInner::Channel(rx), + } + } + + fn socket(sock: Arc, uid: Uid) -> Self { + Self { + uid, + inner: UdpReceiverInner::Socket(sock), + } + } +} + +impl PacketReceiver for UdpReceiver +where + Uid: PartialEq + Clone + Send + Sync, +{ + fn recv>(&self, ctx: &Context) -> Result<(Uid, P)> { + let bytes = match &self.inner { + UdpReceiverInner::Channel(rx) => rx + .lock() + .map_err(|_| TransportError::LockPoisoned)? + .recv() + .map_err(|_| TransportError::ChannelClosed)?, + UdpReceiverInner::Socket(sock) => { + let mut buf = vec![0u8; 65535]; + let (n, _) = sock.recv_from(&mut buf)?; + buf.truncate(n); + buf + } + }; + let mut reader = bytes.as_slice(); + let packet = P::decode(&mut reader, ctx)?; + Ok((self.uid.clone(), packet)) + } +} + +pub struct UdpListener { + sock: Arc, + factory: Arc Uid + Send + Sync>, + addr_to_uid: HashMap, + uid_to_tx: HashMap>>, + _p: PhantomData, +} + +impl Listener for UdpListener +where + Uid: PartialEq + Eq + Hash + Clone + Send + Sync, + Data: Send + Sync, +{ + type Data = Data; + type Sender = UdpPeerSender; + type Receiver = UdpReceiver; + + fn accept(&mut self) -> Result>> { + loop { + let mut buf = vec![0u8; 65535]; + let (n, src) = self.sock.recv_from(&mut buf)?; + buf.truncate(n); + + if let Some(uid) = self.addr_to_uid.get(&src).cloned() { + match self.uid_to_tx.get(&uid) { + Some(tx) => { + if tx.send(buf).is_err() { + self.uid_to_tx.remove(&uid); + self.addr_to_uid.remove(&src); + } + } + None => { + self.addr_to_uid.remove(&src); + } + } + continue; + } + + let uid = (self.factory)(); + let (tx, rx) = mpsc::channel(); + self.addr_to_uid.insert(src, uid.clone()); + self.uid_to_tx.insert(uid.clone(), tx); + + if self.uid_to_tx.get(&uid).unwrap().send(buf).is_err() { + self.addr_to_uid.remove(&src); + self.uid_to_tx.remove(&uid); + continue; + } + + let sender = UdpPeerSender { + sock: self.sock.clone(), + peer: src, + }; + let receiver = UdpReceiver::channel(Mutex::new(rx), uid.clone()); + return Ok(Connection::new(sender, receiver)); + } + } +} + +pub struct UdpTransport { + factory: Arc Uid + Send + Sync>, + _p: PhantomData<(Uid, Data)>, +} + +impl UdpTransport +where + Uid: PartialEq + Eq + Hash + Clone + Send + Sync, +{ + pub fn new Uid + Send + Sync + 'static>(factory: F) -> Self { + Self { + factory: Arc::new(factory), + _p: PhantomData, + } + } +} + +impl Default for UdpTransport { + fn default() -> Self { + let counter = Arc::new(AtomicU64::new(0)); + UdpTransport::new(move || counter.fetch_add(1, Ordering::Relaxed)) + } +} + +impl Transport for UdpTransport +where + Uid: PartialEq + Eq + Hash + Clone + Send + Sync, + Data: Send + Sync, +{ + type Data = Data; + type Addr = SocketAddr; + type Sender = UdpPeerSender; + type Receiver = UdpReceiver; + type Listener = UdpListener; + + fn connect( + &self, + addr: SocketAddr, + ) -> Result>> { + let sock = UdpSocket::bind("0.0.0.0:0")?; + let sock = Arc::new(sock); + let uid = (self.factory)(); + let sender = UdpPeerSender { + sock: sock.clone(), + peer: addr, + }; + let receiver = UdpReceiver::socket(sock, uid); + Ok(Connection::new(sender, receiver)) + } + + fn listen(&self, addr: SocketAddr) -> Result> { + let sock = UdpSocket::bind(addr)?; + Ok(UdpListener { + sock: Arc::new(sock), + factory: self.factory.clone(), + addr_to_uid: HashMap::new(), + uid_to_tx: HashMap::new(), + _p: PhantomData, + }) + } +} -- cgit v1.2.3