diff options
| author | zirkonya <zirkonya@iridium.lan> | 2026-08-29 22:11:23 +0200 |
|---|---|---|
| committer | zirkonya <zirkonya@iridium.lan> | 2026-08-29 22:11:23 +0200 |
| commit | bf9b81e669f4ce98c8d58df5d776ae90726d7ec2 (patch) | |
| tree | e0d5bdbd8a3423f1303577ddebbfa34696d336bb | |
| parent | bb41e396aca01fee9f81785681fac0a79d0fd9d8 (diff) | |
Add transport layer
| -rw-r--r-- | src/backend.rs | 47 | ||||
| -rw-r--r-- | src/codec/encode.rs | 304 | ||||
| -rw-r--r-- | src/codec/error.rs | 35 | ||||
| -rw-r--r-- | src/event.rs | 8 | ||||
| -rw-r--r-- | src/event/connection.rs | 24 | ||||
| -rw-r--r-- | src/event/disconnect.rs | 4 | ||||
| -rw-r--r-- | src/event/handler.rs | 6 | ||||
| -rw-r--r-- | src/event/listener.rs | 5 | ||||
| -rw-r--r-- | src/event/received.rs | 2 | ||||
| -rw-r--r-- | src/event/sent.rs | 2 | ||||
| -rw-r--r-- | src/lib.rs | 8 | ||||
| -rw-r--r-- | src/transport.rs | 39 | ||||
| -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 | ||||
| -rw-r--r-- | src/types/prefix/count.rs | 8 | ||||
| -rw-r--r-- | src/types/prefix/length.rs | 10 |
21 files changed, 840 insertions, 94 deletions
diff --git a/src/backend.rs b/src/backend.rs index 8b13789..fce8a26 100644 --- a/src/backend.rs +++ b/src/backend.rs @@ -1 +1,48 @@ +use crate::{ + codec::Codec, + context::Context, + event::handler::EventHandler, + transport::{ConnectionOf, Result, Transport, receiver::PacketReceiver, sender::PacketSender}, +}; +// TODO : check backend ! maybe better implementations +// TODO : maybe add default functions implementation +// TODO : not sure about send/recv functions + +pub trait Backend<Packet, Data = (), Uid = u64> +where + Uid: PartialEq, + Data: Send + Sync + 'static, + Packet: Codec<Data> + Send + Sync + 'static, + Self::Transport: Transport<Uid, Data = Data>, + <Self::Transport as Transport<Uid>>::Sender: PacketSender<Data> + Clone, + <Self::Transport as Transport<Uid>>::Receiver: PacketReceiver<Data, Uid>, +{ + type Transport: Transport<Uid, Data = Data>; + + fn transport(&self) -> &Self::Transport; + fn event_handler(&self) -> &EventHandler<Packet, Data, Uid>; + fn make_context(&self) -> Context<Data>; + + fn recv( + &self, + conn: &ConnectionOf<Uid, Self::Transport>, + ctx: &Context<Data>, + ) -> Result<(Uid, Packet)>; + + fn send( + &self, + conn: &ConnectionOf<Uid, Self::Transport>, + packet: &Packet, + ctx: &Context<Data>, + ) -> Result<()>; + + fn connect( + &self, + addr: <Self::Transport as Transport<Uid>>::Addr, + ) -> Result<ConnectionOf<Uid, Self::Transport>>; + + fn serve<A>(&self, addr: <Self::Transport as Transport<Uid>>::Addr, on_accept: A) -> Result<()> + where + A: Fn(&Self, &ConnectionOf<Uid, Self::Transport>, &Context<Data>) + Send + Sync + 'static; +} diff --git a/src/codec/encode.rs b/src/codec/encode.rs index b9ac961..ae5eb4e 100644 --- a/src/codec/encode.rs +++ b/src/codec/encode.rs @@ -1,83 +1,270 @@ //! Encoding: writing protocol values into a byte stream. -use crate::codec::error::Error; +use crate::codec::error::Result; use crate::context::Context; -use std::{io::Write, sync::Arc}; +use std::{ + collections::{BTreeMap, BTreeSet, BinaryHeap, HashMap, HashSet, LinkedList, VecDeque}, + io::Write, +}; macro_rules! impl_encode { + (u8) => { + impl<Data> Encode<Data> for u8 { + /// write number using big endian + fn encode(&self, buffer: &mut dyn Write, _: &Context<Data>) -> Result<usize> { + let size = buffer.write(&self.to_be_bytes())?; + Ok(size) + } + + /// write raw bytes + fn encode_slice( + slice: &[Self], + buffer: &mut dyn Write, + _: &Context<Data>, + ) -> Result<usize> { + buffer.write(slice).map_err(Into::into) + } + } + }; ($t: ty) => { impl<Data> Encode<Data> for $t { /// write number using big endian - fn encode(&self, buffer: &mut dyn Write, _: &Context<Data>) -> Result<usize, Error> { + fn encode(&self, buffer: &mut dyn Write, _: &Context<Data>) -> Result<usize> { let size = buffer.write(&self.to_be_bytes())?; Ok(size) } + + fn encode_slice( + slice: &[Self], + buffer: &mut dyn Write, + _: &Context<Data>, + ) -> Result<usize> { + let len = buffer.write( + &slice + .iter() + .flat_map(|n| n.to_be_bytes()) + .collect::<Box<[u8]>>(), + )?; + Ok(len) + } + } + }; +} + +macro_rules! impl_encode_tuples { + ($($generic: ident),+) => { + impl<Data, $($generic),+> Encode<Data> for ($($generic,)+) + where + $($generic: Encode<Data>),+ + { + #[allow(non_snake_case)] + fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { + let ($($generic,)+): &($($generic,)+) = self; + let mut len = 0; + $(len += $generic.encode(buffer, ctx)?;)+ + Ok(len) + } } }; } /// Write a value into a byte stream pub trait Encode<Data> { - fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize, Error>; + fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize>; + fn encode_slice(slice: &[Self], buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> + where + Self: Sized, + { + let mut len = 0; + for item in slice { + len += item.encode(buffer, ctx)?; + } + Ok(len) + } } -impl<Data> Encode<Data> for bool { - fn encode(&self, buffer: &mut dyn Write, _: &Context<Data>) -> Result<usize, Error> { - let size = buffer.write(&[*self as u8])?; - Ok(size) +// ~ Encode arrays + +impl<Data, T: Encode<Data>> Encode<Data> for Vec<T> { + /// Encode each element of Vec + /// use `CountPrefix` to prefix the vector with number of elements + /// use `LenPrefix` to prefix the vector with encoded byte size + fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { + T::encode_slice(self, buffer, ctx) } } -impl<Data, T: Encode<Data>> Encode<Data> for Option<T> { - /// Encode `T` if Some(T) or do nothing if None - fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize, Error> { - match self { - Some(val) => { - let size = val.encode(buffer, ctx)?; - Ok(size) - } - None => Ok(0), - } +impl<Data, T: Encode<Data>> Encode<Data> for VecDeque<T> { + /// Encode each element of VecDeque + /// use `CountPrefix` to prefix the VecDeque with number of elements + /// use `LenPrefix` to prefix the VecDeque with encoded byte size + fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { + let (front, _) = self.as_slices(); + T::encode_slice(front, buffer, ctx) } } -impl<Data, T: Encode<Data>> Encode<Data> for Vec<T> { - /// Encode each element of vector - /// use `CountPrefix` to prefix the vector with number of elements - /// use `LenPrefix` to prefix the vector with encoded byte size - fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize, Error> { - let mut size = 0; - for item in self { - size += item.encode(buffer, ctx)?; - } - Ok(size) +impl<Data, T: Encode<Data>> Encode<Data> for LinkedList<T> +where + T: Clone, +{ + /// Encode each element of LinkedList (use clone..) + /// use `CountPrefix` to prefix the LinkedList with number of elements + /// use `LenPrefix` to prefix the LinkedList with encoded byte size + fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { + let mut view = Vec::with_capacity(self.len()); + view.extend(self.iter().cloned()); + T::encode_slice(&view, buffer, ctx) + } +} + +// ~ Encode slices + +impl<Data, T: Encode<Data>> Encode<Data> for &[T] { + /// Encode each element of slice (use clone..) + /// use `CountPrefix` to prefix the slice with number of elements + /// use `LenPrefix` to prefix the slice with encoded byte size + fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { + T::encode_slice(self, buffer, ctx) + } +} + +impl<Data, T: Encode<Data>, const S: usize> Encode<Data> for [T; S] { + /// Encode each element of slice (use clone..) + /// use `CountPrefix` to prefix the slice with number of elements + /// use `LenPrefix` to prefix the slice with encoded byte size + fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { + T::encode_slice(self, buffer, ctx) + } +} + +impl<Data, T: Encode<Data>> Encode<Data> for [T] { + /// Encode each element of slice (use clone..) + /// use `CountPrefix` to prefix the slice with number of elements + /// use `LenPrefix` to prefix the slice with encoded byte size + fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { + T::encode_slice(self, buffer, ctx) } } +// ~ Encode set + +impl<Data, T: Encode<Data>> Encode<Data> for HashSet<T> +where + T: Clone, +{ + /// Encode each element of HashSet (use clone..) + /// use `CountPrefix` to prefix the HashSet with number of elements + /// use `LenPrefix` to prefix the HashSet with encoded byte size + fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { + let mut view = Vec::with_capacity(self.len()); + view.extend(self.iter().cloned()); + T::encode_slice(&view, buffer, ctx) + } +} + +impl<Data, T: Encode<Data>> Encode<Data> for BTreeSet<T> +where + T: Clone, +{ + /// Encode each element of BTreeSet (use clone..) + /// use `CountPrefix` to prefix the BTreeSet with number of elements + /// use `LenPrefix` to prefix the BTreeSet with encoded byte size + fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { + let mut view = Vec::with_capacity(self.len()); + view.extend(self.iter().cloned()); + T::encode_slice(&view, buffer, ctx) + } +} + +// ~ Encode misc + +impl<Data, T> Encode<Data> for BinaryHeap<T> +where + T: Encode<Data>, +{ + fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { + T::encode_slice(self.as_slice(), buffer, ctx) + } +} + +// ~ Encode maps + +impl<Data, K, V> Encode<Data> for HashMap<K, V> +where + K: Encode<Data> + Clone, + V: Encode<Data> + Clone, +{ + /// Encode each element of HashMap (use clone..) + /// use `CountPrefix` to prefix the HashMap with number of pairs + /// use `LenPrefix` to prefix the HashMap with encoded byte size + fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { + let slice: Vec<(K, V)> = self.clone().into_iter().collect(); + <(K, V) as Encode<Data>>::encode_slice(&slice, buffer, ctx) + } +} + +impl<Data, K, V> Encode<Data> for BTreeMap<K, V> +where + K: Encode<Data> + Clone, + V: Encode<Data> + Clone, +{ + /// Encode each element of BTreeMap (use clone..) + /// use `CountPrefix` to prefix the BTreeMap with number of pairs + /// use `LenPrefix` to prefix the BTreeMap with encoded byte size + fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { + let slice: Vec<(K, V)> = self.clone().into_iter().collect(); + <(K, V) as Encode<Data>>::encode_slice(&slice, buffer, ctx) + } +} + +// ~ Encode string impl<Data> Encode<Data> for String { /// Encode the string using utf8 /// use `CountPrefix` to prefix the string with char length /// use `LenPrefix` to prefix the string with utf8 bytes length - fn encode(&self, buffer: &mut dyn Write, _: &Context<Data>) -> Result<usize, Error> { + fn encode(&self, buffer: &mut dyn Write, _: &Context<Data>) -> Result<usize> { let utf8 = self.as_bytes(); buffer.write_all(utf8)?; Ok(utf8.len()) } } -impl<Data> Encode<Data> for Arc<[u8]> { - /// write raw slice into the buffer - /// use `LenPrefix` to prefix with byte length - fn encode(&self, buffer: &mut dyn Write, _: &Context<Data>) -> Result<usize, Error> { - let len = buffer.write(self)?; - Ok(len) +impl<Data> Encode<Data> for &str { + /// Encode the string using utf8 + /// use `CountPrefix` to prefix the string with char length + /// use `LenPrefix` to prefix the string with utf8 bytes length + fn encode(&self, buffer: &mut dyn Write, _: &Context<Data>) -> Result<usize> { + let utf8 = self.as_bytes(); + buffer.write_all(utf8)?; + Ok(utf8.len()) } } -impl<Data, const S: usize> Encode<Data> for [u8; S] { - /// write raw slice into the buffer - /// use `LenPrefix` to prefix with byte length - fn encode(&self, buffer: &mut dyn Write, _: &Context<Data>) -> Result<usize, Error> { - let len = buffer.write(self)?; +impl<Data, T: Encode<Data>> Encode<Data> for Option<T> { + /// Encode `T` if Some(T) or do nothing if None + fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { + match self { + Some(val) => { + let size = val.encode(buffer, ctx)?; + Ok(size) + } + None => Ok(0), + } + } +} + +// ~ Encode primitive +impl<Data> Encode<Data> for bool { + fn encode(&self, buffer: &mut dyn Write, _: &Context<Data>) -> Result<usize> { + let size = buffer.write(&[*self as u8])?; + Ok(size) + } + + fn encode_slice(slice: &[Self], buffer: &mut dyn Write, _: &Context<Data>) -> Result<usize> + where + Self: Sized, + { + let len = buffer.write(&slice.iter().map(|n| *n as u8).collect::<Box<[u8]>>())?; Ok(len) } } @@ -96,5 +283,42 @@ impl_encode!(i64); impl_encode!(i128); impl_encode!(isize); +#[cfg(feature = "f16")] +impl_encode!(f16); impl_encode!(f32); impl_encode!(f64); +#[cfg(feature = "f128")] +impl_encode!(f128); + +// ~ Encode tuple + +impl<Data> Encode<Data> for () { + fn encode(&self, _: &mut dyn Write, _: &Context<Data>) -> Result<usize> { + Ok(0) + } + + fn encode_slice(_: &[Self], _: &mut dyn Write, _: &Context<Data>) -> Result<usize> + where + Self: Sized, + { + Ok(0) + } +} + +impl_encode_tuples!(A); +impl_encode_tuples!(A, B); +impl_encode_tuples!(A, B, C); +impl_encode_tuples!(A, B, C, D); +impl_encode_tuples!(A, B, C, D, E); +impl_encode_tuples!(A, B, C, D, E, F); +impl_encode_tuples!(A, B, C, D, E, F, G); +impl_encode_tuples!(A, B, C, D, E, F, G, H); +impl_encode_tuples!(A, B, C, D, E, F, G, H, I); +impl_encode_tuples!(A, B, C, D, E, F, G, H, I, J); +impl_encode_tuples!(A, B, C, D, E, F, G, H, I, J, K); +impl_encode_tuples!(A, B, C, D, E, F, G, H, I, J, K, L); +impl_encode_tuples!(A, B, C, D, E, F, G, H, I, J, K, L, M); +impl_encode_tuples!(A, B, C, D, E, F, G, H, I, J, K, L, M, N); +impl_encode_tuples!(A, B, C, D, E, F, G, H, I, J, K, L, M, N, O); +impl_encode_tuples!(A, B, C, D, E, F, G, H, I, J, K, L, M, N, O, P); +impl_encode_tuples!(A, B, C, D, E, F, G, H, I, J, K, L, M, N, O, P, Q); diff --git a/src/codec/error.rs b/src/codec/error.rs index c1bd4f7..7d88d55 100644 --- a/src/codec/error.rs +++ b/src/codec/error.rs @@ -1,37 +1,24 @@ -use std::fmt::Display; +use thiserror::Error; -// TODO : better error (using thiserror) +#[derive(Debug, Error)] +pub enum CodecError { + #[error("io error: {0}")] + IoError(#[from] std::io::Error), -#[derive(Debug)] -pub enum Error { - IoError(std::io::Error), + #[error("{0}")] Custom(String), } -impl std::error::Error for Error {} +pub type Result<T> = core::result::Result<T, CodecError>; -impl Display for Error { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - write!(f, "{self:?}") - } -} - -pub type Result<T> = core::result::Result<T, Error>; - -impl From<std::io::Error> for Error { - fn from(value: std::io::Error) -> Self { - Self::IoError(value) +impl From<&str> for CodecError { + fn from(value: &str) -> Self { + Self::Custom(value.to_string()) } } -impl From<String> for Error { +impl From<String> for CodecError { fn from(value: String) -> Self { Self::Custom(value) } } - -impl From<&str> for Error { - fn from(value: &str) -> Self { - Self::Custom(value.to_string()) - } -} diff --git a/src/event.rs b/src/event.rs index 95d5a12..6339afe 100644 --- a/src/event.rs +++ b/src/event.rs @@ -19,7 +19,7 @@ pub enum EventKind<Packet, Uid> where Uid: PartialEq, { - Connection(ConnectionEvent), + Connection(ConnectionEvent<Uid>), PacketReceived(PacketReceivedEvent<Packet>), PacketSent(PacketSentEvent<Packet>), Disconnect(DisconnectEvent<Uid>), @@ -30,13 +30,13 @@ pub struct Event<Packet, Uid> where Uid: PartialEq, { - #[getset(get = "pub")] + #[get = "pub"] instant: Instant, - #[getset(get = "pub")] + #[get = "pub"] connection_id: Uid, canceled: Cell<bool>, - #[getset(get = "pub")] + #[get = "pub"] kind: EventKind<Packet, Uid>, } diff --git a/src/event/connection.rs b/src/event/connection.rs index 99813fd..4af2164 100644 --- a/src/event/connection.rs +++ b/src/event/connection.rs @@ -1,13 +1,19 @@ -use std::net::SocketAddr; - use getset::Getters; -// TODO : identification -// TODO : reason type #[derive(Getters)] -pub struct ConnectionEvent { - #[getset(get = "pub")] - peer_addr: SocketAddr, - #[getset(get = "pub")] - local_addr: SocketAddr, +pub struct ConnectionEvent<Uid> +where + Uid: PartialEq, +{ + #[get = "pub"] + client_uid: Uid, +} + +impl<Uid> ConnectionEvent<Uid> +where + Uid: PartialEq, +{ + pub fn new(uid: Uid) -> Self { + Self { client_uid: uid } + } } diff --git a/src/event/disconnect.rs b/src/event/disconnect.rs index 6511027..2183180 100644 --- a/src/event/disconnect.rs +++ b/src/event/disconnect.rs @@ -18,8 +18,8 @@ pub struct DisconnectEvent<Uid> where Uid: PartialEq, { - #[getset(get = "pub")] + #[get = "pub"] connection_id: Uid, - #[getset(get = "pub")] + #[get = "pub"] reason: DisconnectReason, } diff --git a/src/event/handler.rs b/src/event/handler.rs index 3f9827e..f801d12 100644 --- a/src/event/handler.rs +++ b/src/event/handler.rs @@ -1,7 +1,5 @@ -// TODO : dispatch event through all listener -// TODO : maybe compile listener into one ? - use crate::{ + codec::Codec, context::Context, event::{ Event, @@ -12,6 +10,7 @@ use crate::{ pub struct EventHandler<Packet, Data = (), Uid = u64> where Uid: PartialEq, + Packet: Codec<Data>, { listeners: Vec<Box<dyn Listener<Packet, Data, Uid>>>, } @@ -19,6 +18,7 @@ where impl<Packet, Data, Uid> EventHandler<Packet, Data, Uid> where Uid: PartialEq + Clone, + Packet: Codec<Data>, { pub fn register<L: Listener<Packet, Data, Uid> + 'static>(&mut self, listener: L) { self.listeners.push(Box::new(listener)); diff --git a/src/event/listener.rs b/src/event/listener.rs index 81f4b57..8ff6719 100644 --- a/src/event/listener.rs +++ b/src/event/listener.rs @@ -1,9 +1,9 @@ use crate::{ + codec::Codec, context::Context, event::{ConnectionEvent, DisconnectEvent, PacketReceivedEvent, PacketSentEvent}, }; -// TODO : Listener ; interface to perform action when event occured // TODO async trait for async runtime pub enum ListenerResult { @@ -21,8 +21,9 @@ pub enum ListenerResult { pub trait Listener<Packet, Data = (), Uid = u64> where Uid: PartialEq, + Packet: Codec<Data>, { - fn on_connection(&self, event: &ConnectionEvent, ctx: &Context<Data>) -> ListenerResult { + fn on_connection(&self, event: &ConnectionEvent<Uid>, ctx: &Context<Data>) -> ListenerResult { let _ = (event, ctx); ListenerResult::Continue } diff --git a/src/event/received.rs b/src/event/received.rs index 0ef6d0e..46e1f63 100644 --- a/src/event/received.rs +++ b/src/event/received.rs @@ -2,6 +2,6 @@ use getset::Getters; #[derive(Getters)] pub struct PacketReceivedEvent<Packet> { - #[getset(get = "pub")] + #[get = "pub"] packet: Packet, } diff --git a/src/event/sent.rs b/src/event/sent.rs index 5278842..7e88a7f 100644 --- a/src/event/sent.rs +++ b/src/event/sent.rs @@ -2,6 +2,6 @@ use getset::Getters; #[derive(Getters)] pub struct PacketSentEvent<Packet> { - #[getset(get = "pub")] + #[get = "pub"] packet: Packet, } @@ -1,8 +1,11 @@ -#[doc = include_str!("../README.md")] +#![doc = include_str!("../README.md")] +#![cfg_attr(feature = "f16", feature(f16))] +#![cfg_attr(feature = "f128", feature(f128))] pub mod backend; pub mod codec; pub mod context; pub mod event; +pub mod transport; pub mod types; #[cfg(feature = "macros")] @@ -10,3 +13,6 @@ pub use zr_protocol_macros as macros; /// The default buffer length used to allocate chunks pub(crate) const DEFAULT_BUFFER_LEN: usize = 2048; + +// TODO : rewrite tests +// TODO : benchmark with some client / server example diff --git a/src/transport.rs b/src/transport.rs new file mode 100644 index 0000000..a5cf55d --- /dev/null +++ b/src/transport.rs @@ -0,0 +1,39 @@ +use crate::transport::{ + connection::Connection, listener::Listener, receiver::PacketReceiver, sender::PacketSender, +}; + +#[cfg(feature = "tcp")] +pub mod tcp; +#[cfg(feature = "udp")] +pub mod udp; + +pub mod connection; +pub mod error; +pub mod listener; +pub mod receiver; +pub mod sender; + +pub use crate::transport::error::TransportError; + +pub type Result<T> = std::result::Result<T, TransportError>; + +pub trait Transport<Uid = u64>: Send + Sync +where + Uid: PartialEq, +{ + type Data; + type Addr; + type Sender: PacketSender<Self::Data> + Clone; + type Receiver: PacketReceiver<Self::Data, Uid>; + type Listener: Listener<Uid, Data = Self::Data, Sender = Self::Sender, Receiver = Self::Receiver>; + fn connect(&self, addr: Self::Addr) -> Result<ConnectionOf<Uid, Self>>; + fn listen(&self, addr: Self::Addr) -> Result<Self::Listener>; +} + +/// Convenience alias for the [`Connection`] produced by a [`Transport`]. +pub type ConnectionOf<Uid, T> = Connection< + Uid, + <T as Transport<Uid>>::Data, + <T as Transport<Uid>>::Sender, + <T as Transport<Uid>>::Receiver, +>; 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, + }) + } +} diff --git a/src/types/prefix/count.rs b/src/types/prefix/count.rs index e00c18c..7099cec 100644 --- a/src/types/prefix/count.rs +++ b/src/types/prefix/count.rs @@ -13,7 +13,7 @@ pub struct CountPrefix<I, L, D> where D: IntoIterator<Item = I>, { - #[getset(get = "pub")] + #[get = "pub"] data: D, _len: PhantomData<L>, } @@ -70,13 +70,13 @@ where &self, buffer: &mut dyn std::io::prelude::Write, ctx: &Context<Data>, - ) -> Result<usize, crate::codec::error::Error> + ) -> Result<usize, crate::codec::error::CodecError> where Self: Sized, { let vec = self.data.clone().into_iter().collect::<Vec<I>>(); let len: L = vec.len().try_into().map_err(|_| { - crate::codec::error::Error::Custom("count exceeds prefix capacity".into()) + crate::codec::error::CodecError::Custom("count exceeds prefix capacity".into()) })?; let mut l = len.encode(buffer, ctx)?; for item in vec { @@ -107,4 +107,4 @@ where _len: PhantomData, }) } -}
\ No newline at end of file +} diff --git a/src/types/prefix/length.rs b/src/types/prefix/length.rs index 750aa23..8704d01 100644 --- a/src/types/prefix/length.rs +++ b/src/types/prefix/length.rs @@ -11,7 +11,7 @@ use crate::{ #[derive(Getters)] pub struct LenPrefixed<L, D> { - #[getset(get = "pub")] + #[get = "pub"] data: D, _len: PhantomData<L>, } @@ -54,7 +54,11 @@ where L: Codec<Data> + Size, D: Codec<Data>, { - fn encode(&self, writer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize, codec::error::Error> { + fn encode( + &self, + writer: &mut dyn Write, + ctx: &Context<Data>, + ) -> Result<usize, codec::error::CodecError> { let mut buf = Vec::with_capacity(DEFAULT_BUFFER_LEN); self.data.encode(&mut buf, ctx)?; let len = L::from_size(buf.len()); @@ -85,4 +89,4 @@ where _len: PhantomData, }) } -}
\ No newline at end of file +} |
