diff options
Diffstat (limited to 'src/transport')
| -rw-r--r-- | src/transport/connection.rs | 36 | ||||
| -rw-r--r-- | src/transport/error.rs | 19 | ||||
| -rw-r--r-- | src/transport/listener.rs | 21 | ||||
| -rw-r--r-- | src/transport/receiver.rs | 5 | ||||
| -rw-r--r-- | src/transport/sender.rs | 5 | ||||
| -rw-r--r-- | src/transport/tcp.rs | 143 | ||||
| -rw-r--r-- | src/transport/udp.rs | 203 |
7 files changed, 432 insertions, 0 deletions
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<Uid, Data, S, R> +where + Uid: PartialEq, + S: PacketSender<Data> + Clone, + R: PacketReceiver<Data, Uid>, +{ + #[get = "pub"] + sender: S, + #[get = "pub"] + receiver: R, + _uid: PhantomData<Uid>, + _data: PhantomData<Data>, +} + +impl<Uid, Data, S, R> Connection<Uid, Data, S, R> +where + Uid: PartialEq, + S: PacketSender<Data> + Clone, + R: PacketReceiver<Data, Uid>, +{ + 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<Uid, L> = Connection< + Uid, + <L as Listener<Uid>>::Data, + <L as Listener<Uid>>::Sender, + <L as Listener<Uid>>::Receiver, +>; + +pub trait Listener<Uid> +where + Uid: PartialEq, +{ + type Data; + type Sender: PacketSender<Self::Data> + Clone; + type Receiver: PacketReceiver<Self::Data, Uid>; + fn accept(&mut self) -> Result<ListenerConnection<Uid, Self>>; +} 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<Data, Uid: PartialEq>: Send + Sync { + fn recv<P: Codec<Data>>(&self, ctx: &Context<Data>) -> 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<Data>: Send + Sync { + fn send<P: Codec<Data>>(&self, packet: P, ctx: &Context<Data>) -> 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<Mutex<TcpStream>>, +} + +impl<Data> PacketSender<Data> for TcpSender { + fn send<P: Codec<Data>>(&self, packet: P, ctx: &Context<Data>) -> 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<Uid> { + stream: Arc<Mutex<TcpStream>>, + uid: Uid, +} + +impl<Data, Uid> PacketReceiver<Data, Uid> for TcpReceiver<Uid> +where + Uid: PartialEq + Clone + Send + Sync, +{ + fn recv<P: Codec<Data>>(&self, ctx: &Context<Data>) -> 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<Uid, Data> { + listener: Arc<StdTcpListener>, + factory: Arc<dyn Fn() -> Uid + Send + Sync>, + _p: PhantomData<Data>, +} + +impl<Uid, Data> Listener<Uid> for TcpListener<Uid, Data> +where + Uid: PartialEq + Eq + Hash + Clone + Send + Sync, + Data: Send + Sync, +{ + type Data = Data; + type Sender = TcpSender; + type Receiver = TcpReceiver<Uid>; + + fn accept(&mut self) -> Result<Connection<Uid, Data, TcpSender, TcpReceiver<Uid>>> { + 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<Uid, Data> { + factory: Arc<dyn Fn() -> Uid + Send + Sync>, + _p: PhantomData<(Uid, Data)>, +} + +impl<Uid, Data> TcpTransport<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 TcpTransport<u64, ()> { + fn default() -> Self { + let counter = Arc::new(AtomicU64::new(0)); + TcpTransport::new(move || counter.fetch_add(1, Ordering::Relaxed)) + } +} + +impl<Uid, Data> Transport<Uid> for TcpTransport<Uid, Data> +where + Uid: PartialEq + Eq + Hash + Clone + Send + Sync, + Data: Send + Sync, +{ + type Data = Data; + type Addr = SocketAddr; + type Sender = TcpSender; + type Receiver = TcpReceiver<Uid>; + type Listener = TcpListener<Uid, Data>; + + fn connect( + &self, + addr: SocketAddr, + ) -> Result<Connection<Uid, Data, TcpSender, TcpReceiver<Uid>>> { + 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<TcpListener<Uid, Data>> { + 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<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, + }) + } +} |
