diff options
Diffstat (limited to 'src/transport/tcp.rs')
| -rw-r--r-- | src/transport/tcp.rs | 143 |
1 files changed, 143 insertions, 0 deletions
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, + }) + } +} |
