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, }) } }