use std::sync::Arc as StdArc;
use std::sync::atomic::{AtomicUsize, Ordering};
use serde::{Serialize, de::DeserializeOwned};
use tokio::io::AsyncReadExt;
use tokio::net::{TcpListener, TcpStream};
use tokio::sync::{Notify, mpsc};
use tokio::task::JoinSet;
use crate::containerisation::container_state::ContainerState;
use crate::network::serialiser;
use crate::utilities::tx_sender::add_to_tx_with_retry;
#[derive(Debug, Clone)]
pub struct Server<T> {
address: String,
server_tx: mpsc::Sender<T>,
container_state: StdArc<AtomicUsize>,
container_state_notify: StdArc<Notify>,
}
impl<T> Server<T>
where
T: Clone + Send + Serialize + DeserializeOwned + Sync + 'static,
{
pub fn new(
address: String,
server_tx: mpsc::Sender<T>,
container_state: StdArc<AtomicUsize>,
container_state_notify: StdArc<Notify>,
) -> StdArc<Self> {
StdArc::new(Self {
address,
server_tx,
container_state,
container_state_notify,
})
}
pub async fn run(arc_server: StdArc<Self>) -> Result<(), Box<dyn std::error::Error>> {
let listener = TcpListener::bind(arc_server.address.clone()).await?; log::trace!("Server listening on {}", arc_server.address);
let mut join_set = JoinSet::new();
loop {
tokio::select! { result = listener.accept() => {
match result {
Ok((stream, addr)) => {
log::info!("Accepted connection from {addr}");
let handler = StdArc::clone(&arc_server);
join_set.spawn(async move {
Server::handle_stream_notification(&handler, stream).await;
});
},
Err(e) => {
log::error!("Failed to accept connection: {e:?}");
continue; }
}
},
_ = arc_server.container_state_notify.notified() => {
if ContainerState::from(arc_server.container_state.load(Ordering::SeqCst)) == ContainerState::ShuttingDown {
log::info!("Server {} graceful shutdown initiated", arc_server.address);
drop(listener); break; }
}
}
}
log::info!("Waiting for ongoing tasks to complete...");
while join_set.join_next().await.is_some() {} log::info!("Server {} shut down gracefully.", arc_server.address);
arc_server
.container_state
.store(ContainerState::ShuttingDown as usize, Ordering::SeqCst);
arc_server.container_state_notify.notify_waiters();
Ok(())
}
async fn handle_stream_notification(arc_server: &StdArc<Self>, mut stream: TcpStream) {
let mut buf = vec![0u8; 65_536]; let mut message_buf = Vec::new();
loop {
tokio::select! {
read_result = stream.read(&mut buf) => {
match read_result {
Ok(0) => {
log::info!("Client disconnected gracefully.");
break;
}
Ok(n) => {
message_buf.extend_from_slice(&buf[..n]);
while message_buf.len() >= 4 {
let len_bytes = &message_buf[..4];
let msg_len = u32::from_be_bytes(len_bytes.try_into().unwrap()) as usize;
if message_buf.len() < 4 + msg_len {
break; }
let msg_bytes = message_buf[4..4 + msg_len].to_vec();
message_buf.drain(..4 + msg_len);
match serialiser::deserialise_message(&msg_bytes) {
Ok(msg) => {
add_to_tx_with_retry(&arc_server.server_tx, &msg, "Server", "Main").await;
}
Err(e) => {
log::warn!("Message (raw): {}", String::from_utf8_lossy(&msg_bytes));
log::error!("Failed to deserialise message: {e:?}");
continue;
}
}
}
}
Err(e) => {
log::error!("Failed to read from socket; err = {e:?}");
break;
}
}
}
}
}
}
}