use super::command_channel::*;
use super::timeout::*;
use crate::crypto::*;
use crate::error::{Error, Result};
use crate::util::*;
use serde::de::DeserializeOwned;
use serde::ser::Serialize;
use std::collections::HashMap;
use std::future::Future;
use std::marker::PhantomData;
use std::net::SocketAddr;
use std::pin::Pin;
use std::sync::Arc;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream, ToSocketAddrs};
use tokio::sync::mpsc::{channel, Receiver, Sender};
use tokio::task::JoinHandle;
#[allow(clippy::type_complexity)]
#[must_use = "event callbacks do nothing unless you configure them for a server"]
pub struct ServerEventCallbacks<R>
where
R: DeserializeOwned + 'static,
{
connect: Option<Arc<dyn Fn(usize) -> Pin<Box<dyn Future<Output = ()> + Send>> + Send + Sync>>,
disconnect:
Option<Arc<dyn Fn(usize) -> Pin<Box<dyn Future<Output = ()> + Send>> + Send + Sync>>,
receive:
Option<Arc<dyn Fn(usize, R) -> Pin<Box<dyn Future<Output = ()> + Send>> + Send + Sync>>,
stop: Option<Arc<dyn Fn() -> Pin<Box<dyn Future<Output = ()> + Send>> + Send + Sync>>,
}
impl<R> ServerEventCallbacks<R>
where
R: DeserializeOwned + 'static,
{
pub const fn new() -> Self {
Self {
connect: None,
disconnect: None,
receive: None,
stop: None,
}
}
pub fn on_connect<C, F>(mut self, callback: C) -> Self
where
C: Fn(usize) -> F + Send + Sync + 'static,
F: Future<Output = ()> + Send + 'static,
{
self.connect = Some(Arc::new(move |client_id| Box::pin((callback)(client_id))));
self
}
pub fn on_disconnect<C, F>(mut self, callback: C) -> Self
where
C: Fn(usize) -> F + Send + Sync + 'static,
F: Future<Output = ()> + Send + 'static,
{
self.disconnect = Some(Arc::new(move |client_id| Box::pin((callback)(client_id))));
self
}
pub fn on_receive<C, F>(mut self, callback: C) -> Self
where
C: Fn(usize, R) -> F + Send + Sync + 'static,
F: Future<Output = ()> + Send + 'static,
{
self.receive = Some(Arc::new(move |client_id, data| {
Box::pin((callback)(client_id, data))
}));
self
}
pub fn on_stop<C, F>(mut self, callback: C) -> Self
where
C: Fn() -> F + Send + Sync + 'static,
F: Future<Output = ()> + Send + 'static,
{
self.stop = Some(Arc::new(move || Box::pin((callback)())));
self
}
}
impl<R> Default for ServerEventCallbacks<R>
where
R: DeserializeOwned + 'static,
{
fn default() -> Self {
Self::new()
}
}
pub trait ServerEventHandler<R>
where
Self: Send + Sync,
R: DeserializeOwned + 'static,
{
#[allow(unused_variables)]
fn on_connect(&self, client_id: usize) -> impl Future<Output = ()> + Send {
async {}
}
#[allow(unused_variables)]
fn on_disconnect(&self, client_id: usize) -> impl Future<Output = ()> + Send {
async {}
}
#[allow(unused_variables)]
fn on_receive(&self, client_id: usize, data: R) -> impl Future<Output = ()> + Send {
async {}
}
fn on_stop(&self) -> impl Future<Output = ()> + Send {
async {}
}
}
pub struct ServerSendingUnknown;
pub struct ServerSending<S>(PhantomData<fn() -> S>)
where
S: Serialize + 'static;
trait ServerSendingConfig {}
impl ServerSendingConfig for ServerSendingUnknown {}
impl<S> ServerSendingConfig for ServerSending<S> where S: Serialize + 'static {}
pub struct ServerReceivingUnknown;
pub struct ServerReceiving<R>(PhantomData<fn() -> R>)
where
R: DeserializeOwned + 'static;
trait ServerReceivingConfig {}
impl ServerReceivingConfig for ServerReceivingUnknown {}
impl<R> ServerReceivingConfig for ServerReceiving<R> where R: DeserializeOwned + 'static {}
pub struct ServerEventReportingUnknown;
pub struct ServerEventReporting<E>(E);
pub struct ServerEventReportingCallbacks<R>(ServerEventCallbacks<R>)
where
R: DeserializeOwned + 'static;
pub struct ServerEventReportingHandler<R, H>
where
R: DeserializeOwned + 'static,
H: ServerEventHandler<R>,
{
handler: H,
phantom_receive: PhantomData<fn() -> R>,
}
pub struct ServerEventReportingChannel;
trait ServerEventReportingConfig {}
impl ServerEventReportingConfig for ServerEventReportingUnknown {}
impl<R> ServerEventReportingConfig for ServerEventReporting<ServerEventReportingCallbacks<R>> where
R: DeserializeOwned + 'static
{
}
impl<R, H> ServerEventReportingConfig for ServerEventReporting<ServerEventReportingHandler<R, H>>
where
R: DeserializeOwned + 'static,
H: ServerEventHandler<R>,
{
}
impl ServerEventReportingConfig for ServerEventReporting<ServerEventReportingChannel> {}
#[allow(private_bounds)]
#[must_use = "server builders do nothing unless `start` is called"]
pub struct ServerBuilder<SC, RC, EC>
where
SC: ServerSendingConfig,
RC: ServerReceivingConfig,
EC: ServerEventReportingConfig,
{
marker: PhantomData<fn() -> (SC, RC)>,
event_reporting: EC,
}
impl ServerBuilder<ServerSendingUnknown, ServerReceivingUnknown, ServerEventReportingUnknown> {
pub const fn new() -> Self {
Self {
marker: PhantomData,
event_reporting: ServerEventReportingUnknown,
}
}
}
impl Default
for ServerBuilder<ServerSendingUnknown, ServerReceivingUnknown, ServerEventReportingUnknown>
{
fn default() -> Self {
Self::new()
}
}
#[allow(private_bounds)]
impl<RC, EC> ServerBuilder<ServerSendingUnknown, RC, EC>
where
RC: ServerReceivingConfig,
EC: ServerEventReportingConfig,
{
pub fn sending<S>(self) -> ServerBuilder<ServerSending<S>, RC, EC>
where
S: Serialize + 'static,
{
ServerBuilder {
marker: PhantomData,
event_reporting: self.event_reporting,
}
}
}
#[allow(private_bounds)]
impl<SC, EC> ServerBuilder<SC, ServerReceivingUnknown, EC>
where
SC: ServerSendingConfig,
EC: ServerEventReportingConfig,
{
pub fn receiving<R>(self) -> ServerBuilder<SC, ServerReceiving<R>, EC>
where
R: DeserializeOwned + 'static,
{
ServerBuilder {
marker: PhantomData,
event_reporting: self.event_reporting,
}
}
}
impl<S, R> ServerBuilder<ServerSending<S>, ServerReceiving<R>, ServerEventReportingUnknown>
where
S: Serialize + 'static,
R: DeserializeOwned + 'static,
{
pub fn with_event_callbacks(
self,
callbacks: ServerEventCallbacks<R>,
) -> ServerBuilder<
ServerSending<S>,
ServerReceiving<R>,
ServerEventReporting<ServerEventReportingCallbacks<R>>,
>
where
R: DeserializeOwned + 'static,
{
ServerBuilder {
marker: PhantomData,
event_reporting: ServerEventReporting(ServerEventReportingCallbacks(callbacks)),
}
}
pub fn with_event_handler<H>(
self,
handler: H,
) -> ServerBuilder<
ServerSending<S>,
ServerReceiving<R>,
ServerEventReporting<ServerEventReportingHandler<R, H>>,
>
where
H: ServerEventHandler<R>,
{
ServerBuilder {
marker: PhantomData,
event_reporting: ServerEventReporting(ServerEventReportingHandler {
handler,
phantom_receive: PhantomData,
}),
}
}
pub fn with_event_channel(
self,
) -> ServerBuilder<
ServerSending<S>,
ServerReceiving<R>,
ServerEventReporting<ServerEventReportingChannel>,
> {
ServerBuilder {
marker: PhantomData,
event_reporting: ServerEventReporting(ServerEventReportingChannel),
}
}
}
impl<S, R>
ServerBuilder<
ServerSending<S>,
ServerReceiving<R>,
ServerEventReporting<ServerEventReportingCallbacks<R>>,
>
where
S: Serialize + 'static,
R: DeserializeOwned + 'static,
{
#[allow(clippy::future_not_send)]
pub async fn start<A>(self, addr: A) -> Result<ServerHandle<S>>
where
A: ToSocketAddrs,
{
let (server, mut server_events) = Server::<S, R>::start(addr).await?;
let callbacks = self.event_reporting.0 .0;
tokio::spawn(async move {
while let Ok(event) = server_events.next_raw().await {
match event {
ServerEventRawSafe::Connect { client_id } => {
if let Some(ref connect) = callbacks.connect {
let connect = Arc::clone(connect);
tokio::spawn(async move {
(*connect)(client_id).await;
});
}
}
ServerEventRawSafe::Disconnect { client_id } => {
if let Some(ref disconnect) = callbacks.disconnect {
let disconnect = Arc::clone(disconnect);
tokio::spawn(async move {
(*disconnect)(client_id).await;
});
}
}
ServerEventRawSafe::Receive { client_id, data } => {
if let Some(ref receive) = callbacks.receive {
let receive = Arc::clone(receive);
tokio::spawn(async move {
let data = data.deserialize();
(*receive)(client_id, data).await;
});
}
}
ServerEventRawSafe::Stop => {
if let Some(ref stop) = callbacks.stop {
let stop = Arc::clone(stop);
tokio::spawn(async move {
(*stop)().await;
});
}
}
}
}
});
Ok(server)
}
}
impl<S, R, H>
ServerBuilder<
ServerSending<S>,
ServerReceiving<R>,
ServerEventReporting<ServerEventReportingHandler<R, H>>,
>
where
S: Serialize + 'static,
R: DeserializeOwned + 'static,
H: ServerEventHandler<R> + 'static,
{
#[allow(clippy::future_not_send)]
pub async fn start<A>(self, addr: A) -> Result<ServerHandle<S>>
where
A: ToSocketAddrs,
{
let (server, mut server_events) = Server::<S, R>::start(addr).await?;
let handler = Arc::new(self.event_reporting.0.handler);
tokio::spawn(async move {
while let Ok(event) = server_events.next_raw().await {
match event {
ServerEventRawSafe::Connect { client_id } => {
let handler = Arc::clone(&handler);
tokio::spawn(async move {
handler.on_connect(client_id).await;
});
}
ServerEventRawSafe::Disconnect { client_id } => {
let handler = Arc::clone(&handler);
tokio::spawn(async move {
handler.on_disconnect(client_id).await;
});
}
ServerEventRawSafe::Receive { client_id, data } => {
let handler = Arc::clone(&handler);
tokio::spawn(async move {
let data = data.deserialize();
handler.on_receive(client_id, data).await;
});
}
ServerEventRawSafe::Stop => {
let handler = Arc::clone(&handler);
tokio::spawn(async move {
handler.on_stop().await;
});
}
}
}
});
Ok(server)
}
}
impl<S, R>
ServerBuilder<
ServerSending<S>,
ServerReceiving<R>,
ServerEventReporting<ServerEventReportingChannel>,
>
where
S: Serialize + 'static,
R: DeserializeOwned + 'static,
{
#[allow(clippy::future_not_send)]
pub async fn start<A>(self, addr: A) -> Result<(ServerHandle<S>, ServerEventStream<R>)>
where
A: ToSocketAddrs,
{
Server::<S, R>::start(addr).await
}
}
pub enum ServerCommand {
Stop,
Send {
client_id: usize,
data: Vec<u8>,
},
SendAll {
data: Vec<u8>,
},
GetAddr,
GetClientAddr {
client_id: usize,
},
RemoveClient {
client_id: usize,
},
}
pub enum ServerCommandReturn {
Stop(Result<()>),
Send(Result<()>),
SendAll(Result<()>),
GetAddr(Result<SocketAddr>),
GetClientAddr(Result<SocketAddr>),
RemoveClient(Result<()>),
}
pub enum ServerClientCommand {
Send {
data: Arc<[u8]>,
},
GetAddr,
Remove,
}
pub enum ServerClientCommandReturn {
Send(Result<()>),
GetAddr(Result<SocketAddr>),
Remove(Result<()>),
}
#[derive(Debug, Clone)]
pub enum ServerEvent<R>
where
R: DeserializeOwned + 'static,
{
Connect {
client_id: usize,
},
Disconnect {
client_id: usize,
},
Receive {
client_id: usize,
data: R,
},
Stop,
}
#[derive(Debug, Clone)]
enum ServerEventRaw {
Connect {
client_id: usize,
},
Disconnect {
client_id: usize,
},
Receive {
client_id: usize,
data: Vec<u8>,
},
Stop,
}
impl ServerEventRaw {
fn deserialize<R>(&self) -> Result<ServerEvent<R>>
where
R: DeserializeOwned + 'static,
{
match self {
Self::Connect { client_id } => Ok(ServerEvent::Connect {
client_id: *client_id,
}),
Self::Disconnect { client_id } => Ok(ServerEvent::Disconnect {
client_id: *client_id,
}),
Self::Receive { client_id, data } => {
Ok(
serde_json::from_slice(data).map(|data| ServerEvent::Receive {
client_id: *client_id,
data,
})?,
)
}
Self::Stop => Ok(ServerEvent::Stop),
}
}
}
#[derive(Debug, Clone)]
struct ServerEventRawSafeData<R>
where
R: DeserializeOwned + 'static,
{
data: Vec<u8>,
marker: PhantomData<fn() -> R>,
}
#[derive(Debug, Clone)]
enum ServerEventRawSafe<R>
where
R: DeserializeOwned + 'static,
{
Connect {
client_id: usize,
},
Disconnect {
client_id: usize,
},
Receive {
client_id: usize,
data: ServerEventRawSafeData<R>,
},
Stop,
}
impl<R> TryFrom<ServerEventRaw> for ServerEventRawSafe<R>
where
R: DeserializeOwned + 'static,
{
type Error = Error;
fn try_from(value: ServerEventRaw) -> std::result::Result<Self, Self::Error> {
value.deserialize::<R>()?;
Ok(match value {
ServerEventRaw::Connect { client_id } => Self::Connect { client_id },
ServerEventRaw::Disconnect { client_id } => Self::Disconnect { client_id },
ServerEventRaw::Receive { client_id, data } => Self::Receive {
client_id,
data: ServerEventRawSafeData {
data,
marker: PhantomData,
},
},
ServerEventRaw::Stop => Self::Stop,
})
}
}
impl<R> ServerEventRawSafeData<R>
where
R: DeserializeOwned + 'static,
{
fn deserialize(&self) -> R {
serde_json::from_slice(&self.data).unwrap()
}
}
impl<R> ServerEventRawSafe<R>
where
R: DeserializeOwned + 'static,
{
#[allow(dead_code)]
fn deserialize(&self) -> ServerEvent<R> {
match self {
Self::Connect { client_id } => ServerEvent::Connect {
client_id: *client_id,
},
Self::Disconnect { client_id } => ServerEvent::Disconnect {
client_id: *client_id,
},
Self::Receive { client_id, data } => ServerEvent::Receive {
client_id: *client_id,
data: data.deserialize(),
},
Self::Stop => ServerEvent::Stop,
}
}
}
pub struct ServerEventStream<R>
where
R: DeserializeOwned + 'static,
{
event_receiver: Receiver<ServerEventRaw>,
marker: PhantomData<fn() -> R>,
}
impl<R> ServerEventStream<R>
where
R: DeserializeOwned + 'static,
{
pub async fn next(&mut self) -> Result<ServerEvent<R>> {
match self.event_receiver.recv().await {
Some(serialized_event) => serialized_event.deserialize(),
None => Err(Error::ConnectionClosed),
}
}
async fn next_raw(&mut self) -> Result<ServerEventRawSafe<R>> {
match self.event_receiver.recv().await {
Some(serialized_event) => serialized_event.try_into(),
None => Err(Error::ConnectionClosed),
}
}
}
pub struct ServerHandle<S>
where
S: Serialize + 'static,
{
server_command_sender: CommandChannelSender<ServerCommand, ServerCommandReturn>,
server_task_handle: JoinHandle<Result<()>>,
marker: PhantomData<fn() -> S>,
}
impl<S> ServerHandle<S>
where
S: Serialize + 'static,
{
#[allow(clippy::missing_panics_doc)]
pub async fn stop(mut self) -> Result<()> {
let value = self
.server_command_sender
.send_command(ServerCommand::Stop)
.await?;
self.server_task_handle.await.unwrap()?;
unwrap_enum!(value, ServerCommandReturn::Stop)
}
#[allow(clippy::future_not_send)]
pub async fn send(&mut self, client_id: usize, data: S) -> Result<()> {
let data_serialized = serde_json::to_vec(&data)?;
let value = self
.server_command_sender
.send_command(ServerCommand::Send {
client_id,
data: data_serialized,
})
.await?;
unwrap_enum!(value, ServerCommandReturn::Send)
}
#[allow(clippy::future_not_send)]
pub async fn send_all(&mut self, data: S) -> Result<()> {
let data_serialized = serde_json::to_vec(&data)?;
let value = self
.server_command_sender
.send_command(ServerCommand::SendAll {
data: data_serialized,
})
.await?;
unwrap_enum!(value, ServerCommandReturn::SendAll)
}
pub async fn get_addr(&mut self) -> Result<SocketAddr> {
let value = self
.server_command_sender
.send_command(ServerCommand::GetAddr)
.await?;
unwrap_enum!(value, ServerCommandReturn::GetAddr)
}
pub async fn get_client_addr(&mut self, client_id: usize) -> Result<SocketAddr> {
let value = self
.server_command_sender
.send_command(ServerCommand::GetClientAddr { client_id })
.await?;
unwrap_enum!(value, ServerCommandReturn::GetClientAddr)
}
pub async fn remove_client(&mut self, client_id: usize) -> Result<()> {
let value = self
.server_command_sender
.send_command(ServerCommand::RemoveClient { client_id })
.await?;
unwrap_enum!(value, ServerCommandReturn::RemoveClient)
}
}
pub struct Server<S, R>
where
S: Serialize + 'static,
R: DeserializeOwned + 'static,
{
marker: PhantomData<fn() -> (S, R)>,
}
impl Server<(), ()> {
pub const fn builder(
) -> ServerBuilder<ServerSendingUnknown, ServerReceivingUnknown, ServerEventReportingUnknown>
{
ServerBuilder::new()
}
}
impl<S, R> Server<S, R>
where
S: Serialize + 'static,
R: DeserializeOwned + 'static,
{
#[allow(clippy::future_not_send)]
pub async fn start<A>(addr: A) -> Result<(ServerHandle<S>, ServerEventStream<R>)>
where
A: ToSocketAddrs,
{
let listener = TcpListener::bind(addr).await?;
let (server_command_sender, server_command_receiver) = command_channel();
let (server_event_sender, server_event_receiver) = channel(CHANNEL_BUFFER_SIZE);
let server_task_handle = tokio::spawn(server_handler(
listener,
server_event_sender,
server_command_receiver,
));
let server_handle = ServerHandle {
server_command_sender,
server_task_handle,
marker: PhantomData,
};
let server_event_stream = ServerEventStream {
event_receiver: server_event_receiver,
marker: PhantomData,
};
Ok((server_handle, server_event_stream))
}
}
#[allow(clippy::too_many_lines)]
async fn server_client_loop(
client_id: usize,
mut socket: TcpStream,
server_client_event_sender: Sender<ServerEventRaw>,
mut client_command_receiver: CommandChannelReceiver<
ServerClientCommand,
ServerClientCommandReturn,
>,
) -> Result<()> {
let (public_key, secret_key) = dh_key_pair().await;
socket.write_all(public_key.as_bytes()).await?;
socket.flush().await?;
let mut other_public_key = [0; PUBLIC_KEY_SIZE];
handshake_timeout! {
socket.read_exact(&mut other_public_key)
}??;
let aes_key = dh_shared_key(secret_key, other_public_key).await;
let mut size_buffer = [0; LEN_SIZE];
loop {
tokio::select! {
read_value = socket.read(&mut size_buffer[..]) => {
let n_size = read_value?;
if n_size != LEN_SIZE {
socket.shutdown().await?;
break;
}
let encrypted_data_size = decode_message_size(&size_buffer);
let mut encrypted_data_buffer = vec![0; encrypted_data_size];
let n_data = data_read_timeout! {
socket.read_exact(&mut encrypted_data_buffer[..])
}??;
if n_data != encrypted_data_size {
socket.shutdown().await?;
break;
}
let data_serialized = aes_decrypt(aes_key, encrypted_data_buffer.into()).await?;
if let Err(_e) = server_client_event_sender.send(ServerEventRaw::Receive { client_id, data: data_serialized }).await {
socket.shutdown().await?;
break;
}
}
client_command_value = client_command_receiver.recv_command() => {
match client_command_value {
Ok(client_command) => {
match client_command {
ServerClientCommand::Send { data } => {
let value = 'val: {
let encrypted_data_buffer = break_on_err!(aes_encrypt(aes_key, data).await, 'val);
let size_buffer = encode_message_size(encrypted_data_buffer.len());
let mut buffer = vec![];
buffer.extend_from_slice(&size_buffer);
buffer.extend(&encrypted_data_buffer);
break_on_err!(socket.write_all(&buffer).await, 'val);
break_on_err!(socket.flush().await, 'val);
Ok(())
};
let error_occurred = value.is_err();
if let Err(_e) = client_command_receiver.command_return(ServerClientCommandReturn::Send(value)).await {
socket.shutdown().await?;
break;
}
if error_occurred {
socket.shutdown().await?;
break;
}
},
ServerClientCommand::GetAddr => {
let addr = socket.peer_addr();
if let Err(_e) = client_command_receiver.command_return(ServerClientCommandReturn::GetAddr(addr.map_err(Into::into))).await {
socket.shutdown().await?;
break;
}
},
ServerClientCommand::Remove => {
let value = socket.shutdown().await;
_ = client_command_receiver.command_return(ServerClientCommandReturn::Remove(value.map_err(Into::into))).await;
break;
},
}
},
Err(_e) => {
socket.shutdown().await?;
break;
},
}
}
}
}
Ok(())
}
fn server_client_handler(
client_id: usize,
socket: TcpStream,
server_client_event_sender: Sender<ServerEventRaw>,
client_cleanup_sender: Sender<usize>,
) -> (
CommandChannelSender<ServerClientCommand, ServerClientCommandReturn>,
JoinHandle<Result<()>>,
) {
let (client_command_sender, client_command_receiver) = command_channel();
let client_task_handle = tokio::spawn(async move {
let res = server_client_loop(
client_id,
socket,
server_client_event_sender,
client_command_receiver,
)
.await;
_ = client_cleanup_sender.send(client_id).await;
res
});
(client_command_sender, client_task_handle)
}
#[allow(clippy::too_many_lines)]
async fn server_loop(
listener: TcpListener,
server_event_sender: Sender<ServerEventRaw>,
mut server_command_receiver: CommandChannelReceiver<ServerCommand, ServerCommandReturn>,
client_command_senders: &mut HashMap<
usize,
CommandChannelSender<ServerClientCommand, ServerClientCommandReturn>,
>,
client_join_handles: &mut HashMap<usize, JoinHandle<Result<()>>>,
) -> Result<()> {
let mut next_client_id = 0usize;
let (server_client_cleanup_sender, mut server_client_cleanup_receiver) =
channel::<usize>(CHANNEL_BUFFER_SIZE);
loop {
tokio::select! {
accept_value = listener.accept() => {
let (socket, _) = accept_value?;
let client_id = next_client_id;
next_client_id += 1;
let server_client_event_sender = server_event_sender.clone();
let client_cleanup_sender = server_client_cleanup_sender.clone();
let (client_command_sender, client_task_handle) = server_client_handler(client_id, socket, server_client_event_sender, client_cleanup_sender);
client_command_senders.insert(client_id, client_command_sender);
client_join_handles.insert(client_id, client_task_handle);
if let Err(_e) = server_event_sender
.send(ServerEventRaw::Connect { client_id })
.await
{
break;
}
},
command_value = server_command_receiver.recv_command() => {
match command_value {
Ok(command) => {
match command {
ServerCommand::Stop => {
_ = server_command_receiver.command_return(ServerCommandReturn::Stop(Ok(()))).await;
break;
},
ServerCommand::Send { client_id, data } => {
let value = match client_command_senders.get_mut(&client_id) {
Some(client_command_sender) => {
let shareable_data = Arc::<[u8]>::from(data);
match client_command_sender.send_command(ServerClientCommand::Send { data: shareable_data }).await {
Ok(return_value) => unwrap_enum!(return_value, ServerClientCommandReturn::Send),
Err(_e) => {
Ok(())
},
}
},
None => Err(Error::InvalidClientId(client_id)),
};
_ = server_command_receiver.command_return(ServerCommandReturn::Send(value)).await;
},
ServerCommand::SendAll { data } => {
let value = {
let shareable_data = Arc::<[u8]>::from(data);
let send_futures = client_command_senders.iter_mut().map(|(_client_id, client_command_sender)| async {
match client_command_sender.send_command(ServerClientCommand::Send { data: Arc::clone(&shareable_data) }).await {
Ok(return_value) => unwrap_enum!(return_value, ServerClientCommandReturn::Send),
Err(_e) => {
Ok(())
}
}
});
let resolved = futures::future::join_all(send_futures).await;
resolved.into_iter().collect::<Result<Vec<_>>>().map(|_| ())
};
_ = server_command_receiver.command_return(ServerCommandReturn::SendAll(value)).await;
},
ServerCommand::GetAddr => {
let addr = listener.local_addr();
_ = server_command_receiver.command_return(ServerCommandReturn::GetAddr(addr.map_err(Into::into))).await;
},
ServerCommand::GetClientAddr { client_id } => {
let value = match client_command_senders.get_mut(&client_id) {
Some(client_command_sender) => match client_command_sender.send_command(ServerClientCommand::GetAddr).await {
Ok(return_value) => unwrap_enum!(return_value, ServerClientCommandReturn::GetAddr),
Err(_e) => {
Err(Error::InvalidClientId(client_id))
},
},
None => Err(Error::InvalidClientId(client_id)),
};
_ = server_command_receiver.command_return(ServerCommandReturn::GetClientAddr(value)).await;
},
ServerCommand::RemoveClient { client_id } => {
let value = match client_command_senders.get_mut(&client_id) {
Some(client_command_sender) => match client_command_sender.send_command(ServerClientCommand::Remove).await {
Ok(return_value) => unwrap_enum!(return_value, ServerClientCommandReturn::Remove),
Err(_e) => {
Ok(())
},
},
None => Err(Error::InvalidClientId(client_id)),
};
_ = server_command_receiver.command_return(ServerCommandReturn::RemoveClient(value)).await;
},
}
},
Err(_e) => {
break;
},
}
}
disconnecting_client_id = server_client_cleanup_receiver.recv() => {
match disconnecting_client_id {
Some(client_id) => {
client_command_senders.remove(&client_id);
if let Some(handle) = client_join_handles.remove(&client_id) {
if let Err(e) = handle.await.unwrap() {
if cfg!(test) {
Err(e)?;
} else {
}
}
}
if let Err(_e) = server_event_sender.send(ServerEventRaw::Disconnect { client_id }).await {
break;
}
},
None => {
break;
},
}
}
}
}
Ok(())
}
async fn server_handler(
listener: TcpListener,
server_event_sender: Sender<ServerEventRaw>,
server_command_receiver: CommandChannelReceiver<ServerCommand, ServerCommandReturn>,
) -> Result<()> {
let mut client_command_senders: HashMap<
usize,
CommandChannelSender<ServerClientCommand, ServerClientCommandReturn>,
> = HashMap::new();
let mut client_join_handles: HashMap<usize, JoinHandle<Result<()>>> = HashMap::new();
let server_exit = server_loop(
listener,
server_event_sender.clone(),
server_command_receiver,
&mut client_command_senders,
&mut client_join_handles,
)
.await;
futures::future::join_all(client_command_senders.into_values().map(
|mut client_command_sender| async move {
_ = client_command_sender
.send_command(ServerClientCommand::Remove)
.await;
},
))
.await;
futures::future::join_all(client_join_handles.into_values().map(|handle| async move {
if let Err(e) = handle.await.unwrap() {
if cfg!(test) {
Err(e)?;
} else {
}
}
Ok(())
}))
.await
.into_iter()
.collect::<Result<Vec<_>>>()?;
_ = server_event_sender.send(ServerEventRaw::Stop).await;
server_exit
}