diff options
Diffstat (limited to 'src/transport/udp.rs')
| -rw-r--r-- | src/transport/udp.rs | 203 |
1 files changed, 203 insertions, 0 deletions
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<UdpSocket>, + peer: SocketAddr, +} + +impl<Data> PacketSender<Data> for UdpPeerSender { + fn send<P: crate::codec::Codec<Data>>(&self, packet: P, ctx: &Context<Data>) -> Result<()> { + let mut buf = Vec::new(); + packet.encode(&mut buf, ctx)?; + self.sock.send_to(&buf, self.peer)?; + Ok(()) + } +} + +enum UdpReceiverInner { + Channel(Mutex<mpsc::Receiver<Vec<u8>>>), + Socket(Arc<UdpSocket>), +} + +pub struct UdpReceiver<Uid> { + uid: Uid, + inner: UdpReceiverInner, +} + +impl<Uid> UdpReceiver<Uid> { + fn channel(rx: Mutex<mpsc::Receiver<Vec<u8>>>, uid: Uid) -> Self { + Self { + uid, + inner: UdpReceiverInner::Channel(rx), + } + } + + fn socket(sock: Arc<UdpSocket>, uid: Uid) -> Self { + Self { + uid, + inner: UdpReceiverInner::Socket(sock), + } + } +} + +impl<Data, Uid> PacketReceiver<Data, Uid> for UdpReceiver<Uid> +where + Uid: PartialEq + Clone + Send + Sync, +{ + fn recv<P: crate::codec::Codec<Data>>(&self, ctx: &Context<Data>) -> 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<Uid, Data> { + sock: Arc<UdpSocket>, + factory: Arc<dyn Fn() -> Uid + Send + Sync>, + addr_to_uid: HashMap<SocketAddr, Uid>, + uid_to_tx: HashMap<Uid, mpsc::Sender<Vec<u8>>>, + _p: PhantomData<Data>, +} + +impl<Uid, Data> Listener<Uid> for UdpListener<Uid, Data> +where + Uid: PartialEq + Eq + Hash + Clone + Send + Sync, + Data: Send + Sync, +{ + type Data = Data; + type Sender = UdpPeerSender; + type Receiver = UdpReceiver<Uid>; + + fn accept(&mut self) -> Result<Connection<Uid, Data, UdpPeerSender, UdpReceiver<Uid>>> { + 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<Uid, Data> { + factory: Arc<dyn Fn() -> Uid + Send + Sync>, + _p: PhantomData<(Uid, Data)>, +} + +impl<Uid, Data> UdpTransport<Uid, Data> +where + Uid: PartialEq + Eq + Hash + Clone + Send + Sync, +{ + pub fn new<F: Fn() -> Uid + Send + Sync + 'static>(factory: F) -> Self { + Self { + factory: Arc::new(factory), + _p: PhantomData, + } + } +} + +impl Default for UdpTransport<u64, ()> { + fn default() -> Self { + let counter = Arc::new(AtomicU64::new(0)); + UdpTransport::new(move || counter.fetch_add(1, Ordering::Relaxed)) + } +} + +impl<Uid, Data> Transport<Uid> for UdpTransport<Uid, Data> +where + Uid: PartialEq + Eq + Hash + Clone + Send + Sync, + Data: Send + Sync, +{ + type Data = Data; + type Addr = SocketAddr; + type Sender = UdpPeerSender; + type Receiver = UdpReceiver<Uid>; + type Listener = UdpListener<Uid, Data>; + + fn connect( + &self, + addr: SocketAddr, + ) -> Result<Connection<Uid, Data, UdpPeerSender, UdpReceiver<Uid>>> { + 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<UdpListener<Uid, Data>> { + 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, + }) + } +} |
