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/udp.rs | 203 +++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 203 insertions(+) create mode 100644 src/transport/udp.rs (limited to 'src/transport/udp.rs') 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