use std::sync::{
Arc as StdArc,
atomic::{AtomicUsize, Ordering},
};
use hyperion_framework::containerisation::client_broker::ClientBroker;
use hyperion_framework::containerisation::container_state::ContainerState;
use hyperion_framework::messages::client_broker_message::ClientBrokerMessage;
use hyperion_framework::network::network_topology::{
ClientConnections, Connection, NetworkTopology,
};
use tokio::io::AsyncReadExt;
use tokio::net::TcpListener;
use tokio::sync::Notify;
use tokio::time::{Duration, sleep, timeout};
async fn start_test_server() -> (TcpListener, String) {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("Bind test server");
let addr = listener.local_addr().unwrap();
(listener, addr.to_string())
}
async fn read_one_framed_string(listener: TcpListener) -> String {
let (mut socket, _) = listener.accept().await.expect("Accept connection");
let mut len_buf = [0u8; 4];
socket.read_exact(&mut len_buf).await.expect("Read len");
let len = u32::from_be_bytes(len_buf) as usize;
let mut payload = vec![0u8; len];
socket.read_exact(&mut payload).await.expect("Read payload");
serde_json::from_slice::<String>(&payload).expect("Deserialize string")
}
async fn read_two_framed_strings(listener: TcpListener) -> (String, Option<String>) {
let (mut socket, _) = listener.accept().await.expect("Accept connection");
async fn read_msg(socket: &mut tokio::net::TcpStream) -> Option<String> {
let mut len_buf = [0u8; 4];
if socket.read_exact(&mut len_buf).await.is_err() {
return None;
}
let len = u32::from_be_bytes(len_buf) as usize;
let mut payload = vec![0u8; len];
if socket.read_exact(&mut payload).await.is_err() {
return None;
}
Some(serde_json::from_slice::<String>(&payload).expect("Deserialize string"))
}
let first = read_msg(&mut socket).await.expect("First message");
let second = timeout(Duration::from_millis(500), read_msg(&mut socket))
.await
.unwrap_or(None);
(first, second)
}
fn build_topology_with_client(name: &str, address: String) -> StdArc<NetworkTopology> {
StdArc::new(NetworkTopology {
container_name: "test-container".to_string(),
server_address: "127.0.0.1:0".to_string(),
client_connections: ClientConnections {
client_connection_vec: vec![Connection {
name: name.to_string(),
address,
}],
},
})
}
fn new_state_and_notify() -> (StdArc<AtomicUsize>, StdArc<Notify>) {
(
StdArc::new(AtomicUsize::new(ContainerState::Running as usize)),
StdArc::new(Notify::new()),
)
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn test_handle_message_forwards_to_correct_client() {
let (listener, addr) = start_test_server().await;
let topology = build_topology_with_client("Alpha", addr.clone());
let (container_state, container_notify) = new_state_and_notify();
let broker =
ClientBroker::<String>::init(topology, container_state.clone(), container_notify.clone());
sleep(Duration::from_millis(100)).await;
let msg = ClientBrokerMessage::new(vec!["Alpha"], "Beta".to_string());
let read_task = tokio::spawn(read_one_framed_string(listener));
broker.handle_message(msg).await;
let received = timeout(Duration::from_secs(2), read_task)
.await
.expect("No timeout")
.expect("Task ok");
assert_eq!(received, "Beta");
container_state.store(ContainerState::ShuttingDown as usize, Ordering::SeqCst);
container_notify.notify_waiters();
let mut broker = broker;
broker.shutdown().await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn test_forward_shutdown_blocks_future_messages() {
let (listener, addr) = start_test_server().await;
let topology = build_topology_with_client("Alpha", addr.clone());
let (container_state, container_notify) = new_state_and_notify();
let mut broker =
ClientBroker::<String>::init(topology, container_state.clone(), container_notify.clone());
sleep(Duration::from_millis(100)).await;
let read_task = tokio::spawn(read_two_framed_strings(listener));
broker.forward_shutdown("Beta".to_string()).await;
let later = ClientBrokerMessage::new(vec!["Alpha"], "Charlie".to_string());
broker.handle_message(later).await;
let (first, second) = timeout(Duration::from_secs(3), read_task)
.await
.expect("No timeout")
.expect("Task ok");
assert_eq!(first, "Beta");
assert!(
second.is_none(),
"Expected no second message after forward_shutdown, got: {:?}",
second
);
container_state.store(ContainerState::ShuttingDown as usize, Ordering::SeqCst);
container_notify.notify_waiters();
broker.shutdown().await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn test_unknown_target_does_not_panic() {
let topology = StdArc::new(NetworkTopology {
container_name: "test".to_string(),
server_address: "127.0.0.1:0".to_string(),
client_connections: ClientConnections {
client_connection_vec: vec![],
},
});
let (container_state, container_notify) = new_state_and_notify();
let broker = ClientBroker::<String>::init(topology, container_state, container_notify);
broker
.handle_message(ClientBrokerMessage::new(vec!["Ghost"], "Msg".to_string()))
.await;
let mut broker = broker;
broker.shutdown().await;
}