use std::sync::Arc as StdArc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Instant;
use hyperion_framework::containerisation::container_state::ContainerState;
use hyperion_framework::messages::container_directive::ContainerDirective;
use hyperion_framework::network::serialiser::serialise_message;
use hyperion_framework::network::server::Server;
use serde::{Deserialize, Serialize};
use tokio::io::AsyncWriteExt;
use tokio::net::{TcpListener, TcpStream};
use tokio::sync::{Notify, mpsc};
use tokio::time::{Duration, sleep};
#[derive(Serialize, Deserialize, Debug, Clone)]
pub enum ContainerMessage {
ContainerDirectiveMsg(ContainerDirective),
}
async fn ephemeral_addr() -> String {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
listener.local_addr().unwrap().to_string()
}
async fn _client_task(id: usize, address: String) {
let message = ContainerMessage::ContainerDirectiveMsg(ContainerDirective::RetryAllConnections);
let payload = serialise_message(&message).expect("Message serialisation failed");
let len_prefix = (payload.len() as u32).to_be_bytes();
let mut framed = Vec::with_capacity(4 + payload.len());
framed.extend_from_slice(&len_prefix);
framed.extend_from_slice(&payload);
match TcpStream::connect(&address).await {
Ok(mut stream) => match stream.write_all(&framed).await {
Ok(_) => log::debug!("Message sent to server! ID: {id}"),
Err(e) => log::error!("Failed to send message: {e:?} ID: {id}"),
},
Err(e) => log::error!("Couldn't connect to {address}: {e:?} ID: {id}"),
}
}
fn new_state_and_notify() -> (StdArc<AtomicUsize>, StdArc<Notify>) {
(
StdArc::new(AtomicUsize::new(ContainerState::Running as usize)),
StdArc::new(Notify::new()),
)
}
#[tokio::test]
async fn test_server_high_loading() {
let addr = ephemeral_addr().await;
let (server_tx, rx) = mpsc::channel::<ContainerMessage>(120);
let (container_state, container_state_notify) = new_state_and_notify();
let arc_server = Server::new(
addr.clone(),
server_tx,
container_state.clone(),
container_state_notify.clone(),
);
tokio::spawn(async move {
if let Err(e) = Server::run(arc_server).await {
log::error!("Server error: {e:?}");
}
});
sleep(Duration::from_millis(50)).await;
let mut handles = Vec::new();
let start_time = Instant::now();
let client_count = 100;
for i in 0..client_count {
let addr_clone = addr.clone();
handles.push(tokio::spawn(
async move { _client_task(i, addr_clone).await },
));
}
for handle in handles {
handle.await.expect("client task");
}
sleep(Duration::from_millis(500)).await;
let duration = start_time.elapsed();
log::debug!("Server strain test completed in {duration:?}");
assert_eq!(
rx.len(),
client_count,
"Some messages were not received by the server"
);
container_state.store(ContainerState::ShuttingDown as usize, Ordering::SeqCst);
container_state_notify.notify_waiters();
}
#[tokio::test]
async fn test_server_shutdown_command() {
let addr = ephemeral_addr().await;
let (server_tx, _rx) = mpsc::channel::<ContainerMessage>(10);
let (container_state, container_state_notify) = new_state_and_notify();
let arc_server = Server::new(
addr,
server_tx,
container_state.clone(),
container_state_notify.clone(),
);
let server_handle = tokio::spawn(async move {
if let Err(e) = Server::run(arc_server).await {
panic!("Server error: {e:?}");
}
});
sleep(Duration::from_millis(50)).await;
container_state.store(ContainerState::ShuttingDown as usize, Ordering::SeqCst);
container_state_notify.notify_waiters();
assert!(
server_handle.await.is_ok(),
"Server task did not terminate cleanly"
);
}
#[tokio::test]
async fn test_server_does_not_blow_up_on_invalid_message_deserialisation() {
let addr = ephemeral_addr().await;
let (server_tx, _rx) = mpsc::channel::<ContainerMessage>(10);
let (container_state, container_state_notify) = new_state_and_notify();
let arc_server = Server::new(
addr.clone(),
server_tx,
container_state.clone(),
container_state_notify.clone(),
);
tokio::spawn(async move {
if let Err(e) = Server::run(arc_server).await {
log::error!("Server error: {e:?}");
}
});
sleep(Duration::from_millis(50)).await;
if let Ok(mut stream) = TcpStream::connect(&addr).await {
let _ = stream.write_all(&[0xFF, 0xFF, 0xFF, 0xFF]).await;
}
sleep(Duration::from_millis(100)).await;
assert_eq!(
container_state.load(Ordering::SeqCst),
ContainerState::Running as usize,
"Server crashed after receiving an invalid message"
);
container_state.store(ContainerState::ShuttingDown as usize, Ordering::SeqCst);
container_state_notify.notify_waiters();
}