summaryrefslogtreecommitdiff
path: root/src/transport/udp.rs
diff options
context:
space:
mode:
Diffstat (limited to 'src/transport/udp.rs')
-rw-r--r--src/transport/udp.rs203
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,
+ })
+ }
+}