use std::hash::Hash; use std::io::{Read, 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::context::Context; use crate::transport::connection::Connection; use crate::transport::error::TransportError; use crate::transport::listener::Listener; use crate::transport::receiver::PacketReceiverBuf; use crate::transport::sender::PacketSenderBuf; use crate::transport::{Result, Transport}; #[derive(Clone)] pub struct TcpSender { stream: Arc>, } impl PacketSenderBuf for TcpSender { fn send_buf(&self, buf: &[u8]) -> Result<()> { let mut guard = self .stream .lock() .map_err(|_| TransportError::LockPoisoned)?; guard.write_all(buf)?; guard.flush()?; Ok(()) } } pub struct TcpReceiver { stream: Arc>, uid: Uid, } impl PacketReceiverBuf for TcpReceiver where Uid: PartialEq + Clone + Send + Sync, { fn recv_buf(&self, _ctx: &Context) -> Result<(Uid, Vec)> { let mut guard = self .stream .lock() .map_err(|_| TransportError::LockPoisoned)?; // Read 4-byte big-endian length prefix let mut len_buf = [0u8; 4]; guard.read_exact(&mut len_buf)?; let len = u32::from_be_bytes(len_buf) as usize; let mut buf = vec![0u8; len]; guard.read_exact(&mut buf)?; Ok((self.uid.clone(), buf)) } } pub struct TcpListener { listener: Arc, factory: Arc Uid + Send + Sync>, _p: PhantomData, } impl Listener for TcpListener where Uid: PartialEq + Eq + Hash + Clone + Send + Sync, Data: Send + Sync, { type Data = Data; type Sender = TcpSender; type Receiver = TcpReceiver; fn accept(&mut self) -> Result>> { 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 { factory: Arc Uid + Send + Sync>, _p: PhantomData<(Uid, Data)>, } impl TcpTransport 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 TcpTransport { fn default() -> Self { let counter = Arc::new(AtomicU64::new(0)); TcpTransport::new(move || counter.fetch_add(1, Ordering::Relaxed)) } } impl Transport for TcpTransport where Uid: PartialEq + Eq + Hash + Clone + Send + Sync, Data: Send + Sync, { type Data = Data; type Addr = SocketAddr; type Sender = TcpSender; type Receiver = TcpReceiver; type Listener = TcpListener; fn connect( &self, addr: SocketAddr, ) -> Result>> { 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> { let listener = StdTcpListener::bind(addr)?; Ok(TcpListener { listener: Arc::new(listener), factory: self.factory.clone(), _p: PhantomData, }) } }