use async_trait::async_trait;
use flagset::FlagSet;
use std::net::SocketAddr;
use std::sync::Arc;
use crate::net::message::MessageFlags;
use crate::net::{Connection, NetworkingApi, NetworkingError};
#[async_trait]
pub trait NetDriver: Send + Sync + 'static {
async fn close(&self);
}
#[async_trait]
pub trait NetDriverConnect: NetDriver {
async fn connect(
&self,
connection: Connection,
socket_addr: SocketAddr,
) -> Result<(), NetworkingError>;
}
#[async_trait]
pub trait NetDriverCloseConnection: NetDriver {
async fn close_connection(
&self,
connection: Connection,
);
}
#[async_trait]
pub trait NetDriverListen: NetDriver {
async fn listen(
&self,
socket_addr: SocketAddr,
) -> Result<(), NetworkingError>;
}
pub trait NetDriverSend: NetDriver {
fn send(
&self,
msg: Vec<u8>,
flags: FlagSet<MessageFlags>,
) -> Result<(), NetworkingError>;
}
pub trait NetDriverSendTo: NetDriver {
fn send_to(
&self,
connection: Connection,
msg: Vec<u8>,
flags: FlagSet<MessageFlags>,
) -> Result<(), NetworkingError>;
}
pub trait NetDriverBroadcast: NetDriver {
fn broadcast(
&self,
msg: Vec<u8>,
flags: FlagSet<MessageFlags>,
) -> Result<(), NetworkingError>;
}
pub trait NetDriverFactory<D: NetDriver> {
fn create(&self, api: Arc<NetworkingApi>) -> Arc<D>;
}
impl<D, F> NetDriverFactory<D> for F
where
D: NetDriver,
F: Fn(Arc<NetworkingApi>) -> Arc<D>,
{
fn create(&self, api: Arc<NetworkingApi>) -> Arc<D> {
self(api)
}
}
pub struct NoNetDriver;
impl NoNetDriver {
#[inline]
pub fn new() -> Arc<Self> {
Arc::new(Self {})
}
}
#[async_trait]
impl NetDriver for NoNetDriver {
async fn close(&self) {
}
}
impl NetDriverFactory<NoNetDriver> for Arc<NoNetDriver> {
#[inline]
fn create(&self, _api: Arc<NetworkingApi>) -> Arc<NoNetDriver> {
self.clone()
}
}
#[cfg(test)]
pub mod tests {
use ahash::HashMap;
use bitcode::{Decode, Encode};
use educe::Educe;
use parking_lot::{Mutex, RwLock};
use smol::future;
use smol::future::FutureExt;
use std::convert::Infallible;
use std::fmt;
use std::fmt::{Debug, Display, Formatter};
use std::future::Future;
use std::net::Ipv4Addr;
use std::pin::Pin;
use std::sync::atomic::{AtomicU64, Ordering};
use std::task::{Context, Poll, Waker};
use std::time::{Duration, Instant};
use tracing::debug;
use crate::event::{Event, EventHandler};
use crate::net::message::{Message, MessageBody};
use crate::net::{Networking, NetworkingConnectEvent, NetworkingDisconnectEvent};
use super::*;
pub use gtether_derive::{
test_net_driver_client_server_broadcast,
test_net_driver_client_server_core,
test_net_driver_client_server_send,
test_net_driver_client_server_send_to,
};
pub mod prelude {
pub use super::{
ClientServerStackFactoryExt,
ConnectionPairSliceExt,
NetTestBroadcastSliceExt,
NetTestConnectSliceExt,
NetTestSliceExt,
TestConnectionSendSliceExt,
TestConnectionSendToSliceExt,
TestConnectionSliceExt,
};
}
struct Timeout {
time: Duration,
deadline: Instant,
}
impl Timeout {
fn new(time: Duration) -> Self {
let deadline = Instant::now() + time;
Self {
time,
deadline,
}
}
fn run<T>(
&self,
fut: impl Future<Output = T>,
context: impl AsRef<str>,
extra: Option<fmt::Arguments>,
) -> T {
let timeout_fn = async move {
smol::Timer::at(self.deadline).await;
match extra {
Some(extra) =>
panic!("Timeout ({:?}) reached while: {}; {}", self.time, context.as_ref(), extra),
None =>
panic!("Timeout ({:?}) reached while: {}", self.time, context.as_ref()),
}
};
future::block_on(fut.or(timeout_fn))
}
}
#[derive(Encode, Decode, MessageBody, Debug)]
#[message_flag(Reliable)]
pub struct TestMessage {
pub value: String,
}
impl TestMessage {
#[inline]
pub fn new(value: impl Into<String>) -> Self {
Self {
value: value.into(),
}
}
}
#[derive(Encode, Decode, MessageBody, Debug)]
#[message_flag(Reliable)]
#[message_reply(TestMessage)]
pub struct TestMessageRepliable {
pub value: String
}
impl TestMessageRepliable {
#[inline]
pub fn new(value: impl Into<String>) -> Self {
Self {
value: value.into(),
}
}
}
async fn execute_futures(futures: impl Iterator<Item=impl Future>) {
let ex = smol::LocalExecutor::new();
let tasks = futures
.map(|fut| ex.spawn(fut))
.collect::<Vec<_>>();
ex.run(async {
for task in tasks {
task.await;
}
}).await;
}
#[inline]
pub fn assert_networking_error_eq(actual: NetworkingError, expected: &NetworkingError) {
match expected {
NetworkingError::InvalidAddress(expected_socket_addr) =>
assert_matches!(actual, NetworkingError::InvalidAddress(socket_addr) if socket_addr == *expected_socket_addr),
NetworkingError::InvalidConnection(expected_connection) =>
assert_matches!(actual, NetworkingError::InvalidConnection(connection) if connection == *expected_connection),
NetworkingError::Closed =>
assert_matches!(actual, NetworkingError::Closed),
NetworkingError::Cancelled =>
assert_matches!(actual, NetworkingError::Cancelled),
other => panic!("Unsupported expected error type: {other:?}"),
}
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Hash)]
pub struct NetTestId(u64);
static NEXT_NET_TEST_ID: AtomicU64 = AtomicU64::new(1);
impl NetTestId {
fn new() -> Self {
Self(NEXT_NET_TEST_ID.fetch_add(1, Ordering::SeqCst))
}
}
impl Display for NetTestId {
#[inline]
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
write!(f, "net-test-{}", self.0)
}
}
fn fmt_record_data(
data: &RwLock<HashMap<Connection, HashMap<String, usize>>>,
f: &mut Formatter<'_>,
) -> fmt::Result {
let data = data.read();
f.debug_map()
.entries(data.iter().map(|(connection, stats)| {
(connection.to_string(), stats)
}))
.finish()
}
#[derive(Educe)]
#[educe(Debug)]
struct Records {
#[educe(Debug(method(std::fmt::Display::fmt)))]
id: NetTestId,
#[educe(Debug(method(fmt_record_data)))]
inner: RwLock<HashMap<Connection, HashMap<String, usize>>>,
#[educe(Debug(ignore))]
wakers: Mutex<Vec<Waker>>,
}
impl Records {
fn new(id: NetTestId) -> Self {
Self {
id,
inner: RwLock::default(),
wakers: Mutex::default(),
}
}
fn record_message(&self, connection: Connection, msg: String) {
debug!(id = %self.id, from = %connection, ?msg, "Recording message");
let mut wakers = self.wakers.lock();
let mut records = self.inner.write();
let values: &mut HashMap<String, usize> = records.entry(connection).or_default();
let value_count = values.entry(msg).or_default();
*value_count += 1;
let waker_count = wakers.len();
for waker in wakers.drain(0..waker_count) {
waker.wake();
}
}
fn message_count_for_connection(
&self,
connection: Connection,
value: impl AsRef<str>,
) -> usize {
let records = self.inner.read();
records.get(&connection)
.map(|values| {
values.get(value.as_ref()).map(|v| *v).unwrap_or(0)
})
.unwrap_or(0)
}
fn wait_for_message_count_for_connection(
self: &Arc<Self>,
connection: Connection,
value: impl AsRef<str>,
count: usize,
) -> RecordCountFuture {
let msg = value.as_ref().to_string();
RecordCountFuture {
count_type: RecordCountFutureType::ForConnection { connection, msg },
threshold: count,
records: self.clone(),
}
}
fn message_count(&self, value: impl AsRef<str>) -> usize {
let records = self.inner.read();
records.values().map(|values| {
values.get(value.as_ref()).map(|v| *v).unwrap_or(0)
}).sum()
}
fn wait_for_message_count(
self: &Arc<Self>,
value: impl AsRef<str>,
count: usize,
) -> RecordCountFuture {
let msg = value.as_ref().to_string();
RecordCountFuture {
count_type: RecordCountFutureType::ForMessage { msg },
threshold: count,
records: self.clone(),
}
}
fn total_message_count(&self) -> usize {
let records = self.inner.read();
records.values().map(|values| {
values.values().sum::<usize>()
}).sum()
}
fn wait_for_total_message_count(
self: &Arc<Self>,
count: usize,
) -> RecordCountFuture {
RecordCountFuture {
count_type: RecordCountFutureType::Total,
threshold: count,
records: self.clone(),
}
}
fn clear(&self) {
self.inner.write().clear();
}
}
enum RecordCountFutureType {
ForConnection {
connection: Connection,
msg: String,
},
ForMessage {
msg: String,
},
Total,
}
struct RecordCountFuture {
count_type: RecordCountFutureType,
threshold: usize,
records: Arc<Records>,
}
impl Future for RecordCountFuture {
type Output = usize;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let mut wakers = self.records.wakers.lock();
let count = match &self.count_type {
RecordCountFutureType::ForConnection { connection, msg } =>
self.records.message_count_for_connection(*connection, msg),
RecordCountFutureType::ForMessage { msg } =>
self.records.message_count(msg),
RecordCountFutureType::Total =>
self.records.total_message_count(),
};
if count >= self.threshold {
Poll::Ready(count)
} else {
wakers.push(cx.waker().clone());
Poll::Pending
}
}
}
struct NetworkingClosedFuture<D: NetDriver> {
net: Arc<Networking<D>>,
connection: Connection,
#[allow(unused)]
event_handler: Arc<dyn EventHandler<NetworkingDisconnectEvent>>,
waker: Arc<Mutex<Waker>>,
}
impl<D: NetDriver> Future for NetworkingClosedFuture<D> {
type Output = ();
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let mut waker = self.waker.lock();
if self.net.connection_info(&self.connection).is_none() {
Poll::Ready(())
} else {
*waker = cx.waker().clone();
Poll::Pending
}
}
}
#[derive(Educe)]
#[educe(Debug)]
pub struct NetTest<D: NetDriver> {
#[educe(Debug(method(std::fmt::Display::fmt)))]
id: NetTestId,
#[educe(Debug(ignore))]
net: Arc<Networking<D>>,
records: Arc<Records>,
}
impl<D: NetDriver> NetTest<D> {
pub fn new(driver_factory: impl NetDriverFactory<D>) -> Self {
let id = NetTestId::new();
let net = Arc::new(Networking::new(driver_factory));
let records = Arc::new(Records::new(id));
let handler_records = records.clone();
net.insert_msg_handler(move |connection, msg: Message<TestMessage>| {
let value = msg.body().value.clone();
handler_records.record_message(connection, value);
Ok::<_, Infallible>(())
});
Self {
id,
net,
records,
}
}
#[inline]
pub fn net(&self) -> &Arc<Networking<D>> {
&self.net
}
pub async fn wait_for_connection_closed(&self, connection: Connection) {
let waker = Arc::new(Mutex::new(Waker::noop().clone()));
let event_waker = waker.clone();
let event_handler = Arc::new(move |_: &mut Event<NetworkingDisconnectEvent>| {
event_waker.lock().wake_by_ref();
});
self.net().event_bus().register(Arc::downgrade(&event_handler))
.expect("NetworkingDisconnectEvent type should be valid");
NetworkingClosedFuture {
net: self.net.clone(),
connection,
event_handler: event_handler as _,
waker,
}.await;
}
pub fn clear_records(&self) {
self.records.clear();
}
pub fn assert_total_message_count(&self, count: usize) {
assert_eq!(self.records.total_message_count(), count);
}
pub async fn wait_for_total_message_count(&self, count: usize) {
self.records.wait_for_total_message_count(count).await;
}
pub fn assert_message_count(&self, msg: impl AsRef<str>, count: usize) {
assert_eq!(self.records.message_count(msg), count);
}
pub async fn wait_for_message_count(
&self,
msg: impl AsRef<str>,
count: usize,
) {
self.records.wait_for_message_count(msg, count).await;
}
pub fn assert_connection_message_count(
&self,
connection: Connection,
msg: impl AsRef<str>,
count: usize,
) {
assert_eq!(self.records.message_count_for_connection(connection, msg), count);
}
pub async fn wait_for_connection_message_count(
&self,
connection: Connection,
msg: impl AsRef<str>,
count: usize,
) {
self.records.wait_for_message_count_for_connection(connection, msg, count).await;
}
}
#[async_trait(?Send)]
pub trait NetTestSliceExt {
fn clear_records(&self);
fn assert_total_message_count_for_each(&self, count: usize);
async fn wait_for_total_message_count_for_each(&self, count: usize);
fn assert_message_count_for_each(&self, msg: impl AsRef<str>, count: usize);
async fn wait_for_message_count_for_each(&self, msg: impl AsRef<str>, count: usize);
fn assert_connection_message_count_for_each(
&self,
connection: Connection,
msg: impl AsRef<str>,
count: usize,
);
async fn wait_for_connection_message_count_for_each(
&self,
connection: Connection,
msg: impl AsRef<str>,
count: usize,
);
}
#[async_trait(?Send)]
impl<D: NetDriver> NetTestSliceExt for [NetTest<D>] {
fn clear_records(&self) {
for net in self {
net.clear_records();
}
}
fn assert_total_message_count_for_each(&self, count: usize) {
for net in self {
net.assert_total_message_count(count);
}
}
async fn wait_for_total_message_count_for_each(&self, count: usize) {
execute_futures(self.iter().map(|net| {
net.wait_for_total_message_count(count)
})).await;
}
fn assert_message_count_for_each(&self, msg: impl AsRef<str>, count: usize) {
for net in self {
net.assert_message_count(msg.as_ref(), count);
}
}
async fn wait_for_message_count_for_each(&self, msg: impl AsRef<str>, count: usize) {
execute_futures(self.iter().map(|net| {
net.wait_for_message_count(msg.as_ref(), count)
})).await;
}
fn assert_connection_message_count_for_each(
&self,
connection: Connection,
msg: impl AsRef<str>,
count: usize,
) {
for net in self {
net.assert_connection_message_count(connection, msg.as_ref(), count);
}
}
async fn wait_for_connection_message_count_for_each(
&self,
connection: Connection,
msg: impl AsRef<str>,
count: usize,
) {
execute_futures(self.iter().map(|net| {
net.wait_for_connection_message_count(connection, msg.as_ref(), count)
})).await;
}
}
pub fn testing_ports_from_env() -> impl Iterator<Item=u16> {
let var_str = std::env::var("GTETHER_TEST_NET_DRIVER_PORTS")
.unwrap_or("59000:60000".to_owned());
let (start, end) = var_str.split_once(':')
.unwrap_or(("59000", "60000"));
let start = start.parse::<u16>().unwrap_or(59000);
let end = end.parse::<u16>().unwrap_or(60000);
start..end
}
impl<D: NetDriverListen> NetTest<D> {
pub async fn assert_listen(&self) -> SocketAddr {
for port in testing_ports_from_env() {
let socket_addr = SocketAddr::new(Ipv4Addr::LOCALHOST.into(), port);
match self.net.listen(socket_addr).await {
Ok(_) => return socket_addr,
Err(NetworkingError::InvalidAddress(_)) => continue,
result => result.expect("Listen should succeed"),
}
}
panic!("No port available to bind/listen on")
}
pub async fn assert_listen_error(
&self,
socket_addr: SocketAddr,
expected_error: NetworkingError,
) {
let error = self.net.listen(socket_addr).await
.expect_err("Listen should fail");
assert_networking_error_eq(error, &expected_error);
}
}
impl<S: NetDriverConnect> NetTest<S> {
pub async fn assert_connect<'a, D: NetDriver>(
&'a self,
dst: &'a NetTest<D>,
socket_addr: SocketAddr,
) -> ConnectionPair<'a, S, D> {
let server_connection = Arc::new(Mutex::new(None));
let event_server_connection = server_connection.clone();
let fut = dst.net().event_bus().register_once(move |event: &mut Event<NetworkingConnectEvent>| {
*event_server_connection.lock() = Some(event.connection());
}).expect("NetworkingConnectEvent type should be valid");
let src_to_dst = self.net.connect(socket_addr).await
.expect("src should connect")
.into_connection();
fut.await.expect("NetworkingConnectEvent should fire successfully");
let dst_to_src = server_connection.lock().take()
.expect("server_connection should be set");
ConnectionPair {
src: TestConnection {
net: self,
connection: src_to_dst,
},
dst: TestConnection {
net: dst,
connection: dst_to_src,
},
}
}
}
#[async_trait(?Send)]
pub trait NetTestConnectSliceExt<S: NetDriverConnect> {
async fn assert_connect<'a, D: NetDriver>(
&'a self,
dst: &'a NetTest<D>,
socket_addr: SocketAddr,
) -> Vec<ConnectionPair<'a, S, D>>;
async fn assert_connect_error(
&self,
socket_addr: SocketAddr,
expected_error: NetworkingError,
);
}
#[async_trait(?Send)]
impl<S: NetDriverConnect> NetTestConnectSliceExt<S> for [NetTest<S>] {
async fn assert_connect<'a, D: NetDriver>(
&'a self,
dst: &'a NetTest<D>,
socket_addr: SocketAddr,
) -> Vec<ConnectionPair<'a, S, D>> {
let mut srcs_out = Vec::with_capacity(self.len());
for src in self {
let pair = src.assert_connect(dst, socket_addr).await;
srcs_out.push(pair);
}
srcs_out
}
async fn assert_connect_error(
&self,
socket_addr: SocketAddr,
expected_error: NetworkingError,
) {
for src in self {
let error = src.net().connect(socket_addr).await
.expect_err("Connect should fail");
assert_networking_error_eq(error, &expected_error);
}
}
}
impl<D: NetDriverSend> NetTest<D> {
pub fn init_test_message_repliable_send_handler(
&self,
reply_fn: impl (Fn(String) -> String) + Send + Sync + 'static,
) {
let handler_net = self.net.clone();
let handler_records = self.records.clone();
self.net.insert_msg_handler(move |connection, msg: Message<TestMessageRepliable>| {
let value = msg.body().value.clone();
handler_records.record_message(connection, value.clone());
let reply = msg.reply(TestMessage {
value: reply_fn(value),
});
handler_net.send(reply)
});
}
}
impl<D: NetDriverSendTo> NetTest<D> {
pub fn init_test_message_repliable_send_to_handler(
&self,
reply_fn: impl (Fn(String) -> String) + Send + Sync + 'static,
) {
let handler_net = self.net.clone();
let handler_records = self.records.clone();
self.net.insert_msg_handler(move |connection, msg: Message<TestMessageRepliable>| {
let value = msg.body().value.clone();
handler_records.record_message(connection, value.clone());
let reply = msg.reply(TestMessage {
value: reply_fn(value),
});
handler_net.send_to(connection, reply)
});
}
}
impl<D: NetDriverBroadcast> NetTest<D> {
pub fn assert_broadcast(&self, msg: impl Into<String>) {
let msg = TestMessage::new(msg.into());
self.net.broadcast(msg)
.expect("Message should broadcast");
}
pub fn assert_broadcast_error(&self, msg: impl Into<String>, expected_error: &NetworkingError) {
let msg = TestMessage::new(msg.into());
let error = self.net.broadcast(msg)
.expect_err("Message should NOT broadcast");
assert_networking_error_eq(error, expected_error);
}
}
pub trait NetTestBroadcastSliceExt {
fn assert_broadcast_all(&self, msg: impl Into<String>);
fn assert_broadcast_error_all(&self, msg: impl Into<String>, expected_error: &NetworkingError);
}
impl<D: NetDriverBroadcast> NetTestBroadcastSliceExt for [NetTest<D>] {
fn assert_broadcast_all(&self, msg: impl Into<String>) {
let msg = msg.into();
for net in self {
net.assert_broadcast(&msg);
}
}
fn assert_broadcast_error_all(&self, msg: impl Into<String>, expected_error: &NetworkingError) {
let msg = msg.into();
for net in self {
net.assert_broadcast_error(&msg, expected_error);
}
}
}
pub struct TestConnection<'a, D: NetDriver> {
pub net: &'a NetTest<D>,
pub connection: Connection,
}
impl<'a, D: NetDriver> Debug for TestConnection<'a, D> {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
f.debug_struct("TestConnection")
.field("net", &self.net)
.field("connection", &self.connection)
.field("closed", &self.net.net.connection_info(&self.connection).is_none())
.finish()
}
}
impl<'a, D: NetDriver> Clone for TestConnection<'a, D> {
fn clone(&self) -> Self {
Self {
net: self.net,
connection: self.connection,
}
}
}
impl<'a, D: NetDriver> TestConnection<'a, D> {
#[inline]
pub fn assert_closed(&self) {
assert!(self.net.net.connection_info(&self.connection).is_none());
}
#[inline]
pub async fn wait_for_closed(&self) {
self.net.wait_for_connection_closed(self.connection).await;
}
}
#[async_trait(?Send)]
pub trait TestConnectionSliceExt {
fn assert_all_closed(&self);
async fn wait_for_all_closed(&self);
}
#[async_trait(?Send)]
impl<'a, D: NetDriver> TestConnectionSliceExt for [TestConnection<'a, D>] {
#[inline]
fn assert_all_closed(&self) {
for connection in self {
connection.assert_closed();
}
}
#[inline]
async fn wait_for_all_closed(&self) {
execute_futures(self.iter().map(TestConnection::wait_for_closed)).await;
}
}
impl<'a, D: NetDriverSend> TestConnection<'a, D> {
pub fn assert_send(&self, msg: impl Into<String>) {
let msg = TestMessage::new(msg.into());
self.net.net.send(msg)
.expect("Message should send");
}
pub fn assert_send_error(&self, msg: impl Into<String>, expected_error: &NetworkingError) {
let msg = TestMessage::new(msg.into());
let error = self.net.net.send(msg)
.expect_err("Message should NOT send");
assert_networking_error_eq(error, expected_error);
}
pub async fn assert_send_recv(
&self,
msg: impl Into<String>,
expected_reply: impl Into<String>,
) {
let msg = TestMessageRepliable::new(msg.into());
let reply = self.net.net.send_recv(msg)
.expect("Message should send")
.await;
let msg_reply = reply.into_body().value;
let expected_reply = expected_reply.into();
assert_eq!(msg_reply, expected_reply);
self.net.records.record_message(self.connection, msg_reply);
}
}
#[async_trait(?Send)]
pub trait TestConnectionSendSliceExt {
fn assert_send_all(&self, msg: impl Into<String>);
fn assert_send_error_all(&self, msg: impl Into<String>, expected_error: &NetworkingError);
async fn assert_send_recv_all(
&self,
msg: impl Into<String>,
expected_reply: impl Into<String>,
);
}
#[async_trait(?Send)]
impl <'a, D: NetDriverSend> TestConnectionSendSliceExt for [TestConnection<'a, D>] {
fn assert_send_all(&self, msg: impl Into<String>) {
let msg = msg.into();
for connection in self {
connection.assert_send(&msg);
}
}
fn assert_send_error_all(&self, msg: impl Into<String>, expected_error: &NetworkingError) {
let msg = msg.into();
for connection in self {
connection.assert_send_error(&msg, expected_error);
}
}
async fn assert_send_recv_all(
&self,
msg: impl Into<String>,
expected_reply: impl Into<String>,
) {
let msg = msg.into();
let expected_reply = expected_reply.into();
execute_futures(self.iter().map(|cn| {
cn.assert_send_recv(&msg, &expected_reply)
})).await;
}
}
impl<'a, D: NetDriverSendTo> TestConnection<'a, D> {
pub fn assert_send_to(&self, msg: impl Into<String>) {
let msg = TestMessage::new(msg.into());
self.net.net.send_to(self.connection, msg)
.expect("Message should send");
}
pub fn assert_send_to_error(&self, msg: impl Into<String>, expected_error: &NetworkingError) {
let msg = TestMessage::new(msg.into());
let error = self.net.net.send_to(self.connection, msg)
.expect_err("Message should NOT send");
assert_networking_error_eq(error, expected_error);
}
pub async fn assert_send_recv_to(
&self,
msg: impl Into<String>,
expected_reply: impl Into<String>,
) {
let msg = TestMessageRepliable::new(msg.into());
let reply = self.net.net.send_recv_to(self.connection, msg)
.expect("Message should send")
.await;
let msg_reply = reply.into_body().value;
let expected_reply = expected_reply.into();
assert_eq!(msg_reply, expected_reply);
self.net.records.record_message(self.connection, msg_reply);
}
}
#[async_trait(?Send)]
pub trait TestConnectionSendToSliceExt {
fn assert_send_to_all(&self, msg: impl Into<String>);
fn assert_send_to_error_all(&self, msg: impl Into<String>, expected_error: &NetworkingError);
async fn assert_send_recv_to_all(
&self,
msg: impl Into<String>,
expected_reply: impl Into<String>,
);
}
#[async_trait(?Send)]
impl <'a, D: NetDriverSendTo> TestConnectionSendToSliceExt for [TestConnection<'a, D>] {
fn assert_send_to_all(&self, msg: impl Into<String>) {
let msg = msg.into();
for connection in self {
connection.assert_send_to(&msg);
}
}
fn assert_send_to_error_all(&self, msg: impl Into<String>, expected_error: &NetworkingError) {
let msg = msg.into();
for connection in self {
connection.assert_send_to_error(&msg, expected_error);
}
}
async fn assert_send_recv_to_all(
&self,
msg: impl Into<String>,
expected_reply: impl Into<String>,
) {
let msg = msg.into();
let expected_reply = expected_reply.into();
execute_futures(self.iter().map(|cn| {
cn.assert_send_recv_to(&msg, &expected_reply)
})).await;
}
}
#[derive(Clone, Educe)]
#[educe(Debug)]
pub struct ConnectionPair<'a, S: NetDriver, D: NetDriver> {
pub src: TestConnection<'a, S>,
pub dst: TestConnection<'a, D>,
}
impl<'a, S: NetDriver, D: NetDriver> ConnectionPair<'a, S, D> {
pub async fn close_src_and_wait(&self) {
self.src.net.net.close().await;
self.src.assert_closed();
self.dst.wait_for_closed().await;
}
pub async fn close_dst_and_wait(&self) {
self.dst.net.net.close().await;
self.dst.assert_closed();
self.src.wait_for_closed().await;
}
pub fn flip(&self) -> ConnectionPair<'a, D, S> {
ConnectionPair {
src: self.dst.clone(),
dst: self.src.clone(),
}
}
}
#[async_trait(?Send)]
pub trait ConnectionPairSliceExt<'a, S: NetDriver, D: NetDriver> {
fn srcs(&self) -> Vec<TestConnection<'a, S>>;
async fn close_all_srcs_and_wait(&self);
fn dsts(&self) -> Vec<TestConnection<'a, D>>;
async fn close_all_dsts_and_wait(&self);
fn flip(&self) -> Vec<ConnectionPair<'a, D, S>>;
}
#[async_trait(?Send)]
impl<'a, S: NetDriver, D: NetDriver> ConnectionPairSliceExt<'a, S, D> for [ConnectionPair<'a, S, D>] {
#[inline]
fn srcs(&self) -> Vec<TestConnection<'a, S>> {
self.iter()
.map(|entry| entry.src.clone())
.collect()
}
#[inline]
async fn close_all_srcs_and_wait(&self) {
execute_futures(self.iter().map(ConnectionPair::close_src_and_wait)).await;
}
#[inline]
fn dsts(&self) -> Vec<TestConnection<'a, D>> {
self.iter()
.map(|entry| entry.dst.clone())
.collect()
}
#[inline]
async fn close_all_dsts_and_wait(&self) {
execute_futures(self.iter().map(ConnectionPair::close_dst_and_wait)).await;
}
#[inline]
fn flip(&self) -> Vec<ConnectionPair<'a, D, S>> {
self.iter()
.map(ConnectionPair::flip)
.collect()
}
}
pub trait ClientServerStackFactory: Default
{
type ClientDriver: NetDriverConnect;
type ClientDriverFactory: NetDriverFactory<Self::ClientDriver>;
type ServerDriver: NetDriverListen;
type ServerDriverFactory: NetDriverFactory<Self::ServerDriver>;
fn client_factory(&self) -> Self::ClientDriverFactory;
fn server_factory(&self) -> Self::ServerDriverFactory;
}
pub trait ClientServerStackFactoryExt: ClientServerStackFactory {
#[inline]
fn create_client(&self) -> NetTest<Self::ClientDriver> {
NetTest::new(self.client_factory())
}
#[inline]
fn create_server(&self) -> NetTest<Self::ServerDriver> {
NetTest::new(self.server_factory())
}
fn create_stacks<const CLIENT_COUNT: usize>(
&self,
) -> ([NetTest<Self::ClientDriver>; CLIENT_COUNT], NetTest<Self::ServerDriver>) {
let clients: [NetTest<Self::ClientDriver>; CLIENT_COUNT] = core::array::from_fn(|_| {
NetTest::new(self.client_factory())
});
let server = NetTest::new(self.server_factory());
(clients, server)
}
}
impl<F: ClientServerStackFactory> ClientServerStackFactoryExt for F {}
pub mod suites {
use super::*;
#[derive(Default)]
pub struct ClientServer<F>
where
F: ClientServerStackFactory,
{
stack_factory: F,
}
impl<F> ClientServer<F>
where
F: ClientServerStackFactory,
{
pub fn test_connect(&self, timeout: Duration) {
let timeout = Timeout::new(timeout);
let (clients, server) = self.stack_factory.create_stacks::<1>();
let socket_addr = timeout.run(
server.assert_listen(),
"Server listening",
Some(format_args!("{server:#?}")),
);
timeout.run(
clients.assert_connect(&server, socket_addr),
"Clients connecting",
Some(format_args!("{clients:#?}")),
);
}
pub fn test_listen_already_bound(&self, timeout: Duration) {
let timeout = Timeout::new(timeout);
let server = self.stack_factory.create_server();
let server2 = self.stack_factory.create_server();
let socket_addr = timeout.run(
server.assert_listen(),
"Server 1 listening",
Some(format_args!("{server:#?}")),
);
timeout.run(
server2.assert_listen_error(
socket_addr,
NetworkingError::InvalidAddress(socket_addr),
),
"Server 2 listening",
Some(format_args!("{server2:#?}")),
);
}
pub fn test_connect_cancelled(&self, timeout: Duration) {
let timeout = Timeout::new(timeout);
let (clients, server) = self.stack_factory.create_stacks::<1>();
let socket_addr = timeout.run(
server.assert_listen(),
"Server listening",
Some(format_args!("{server:#?}")),
);
for client in &clients {
client.net().event_bus().register(move |event: &mut Event<NetworkingConnectEvent>| {
event.cancel();
}).expect("NetworkingConnectEvent type should be valid");
}
timeout.run(
clients.assert_connect_error(
socket_addr,
NetworkingError::Cancelled,
),
"Clients connecting",
Some(format_args!("{clients:#?}")),
);
}
pub fn test_connect_after_close(&self, timeout: Duration) {
let timeout = Timeout::new(timeout);
let (clients, server) = self.stack_factory.create_stacks::<1>();
let socket_addr = timeout.run(
server.assert_listen(),
"Server listening",
Some(format_args!("{server:#?}")),
);
let connections = timeout.run(
clients.assert_connect(&server, socket_addr),
"Clients connecting",
Some(format_args!("{clients:#?}")),
);
timeout.run(
connections.close_all_srcs_and_wait(),
"Closing connections from clients",
Some(format_args!("{connections:#?}")),
);
}
}
impl<F> ClientServer<F>
where
F: ClientServerStackFactory,
<F as ClientServerStackFactory>::ClientDriver: NetDriverSend,
{
pub fn test_send(&self, timeout: Duration) {
let timeout = Timeout::new(timeout);
let (clients, server) = self.stack_factory.create_stacks::<1>();
let socket_addr = timeout.run(
server.assert_listen(),
"Server listening",
None,
);
let connections = timeout.run(
clients.assert_connect(&server, socket_addr),
"Clients connecting",
None,
);
connections.srcs().assert_send_all("client->server");
timeout.run(
server.wait_for_total_message_count(clients.len()),
"Waiting for messages",
Some(format_args!("{connections:#?}")),
);
server.assert_message_count("client->server", clients.len());
clients.assert_total_message_count_for_each(0);
}
pub fn test_send_many(&self, timeout: Duration) {
let timeout = Timeout::new(timeout);
let (clients, server) = self.stack_factory.create_stacks::<3>();
let socket_addr = timeout.run(
server.assert_listen(),
"Server listening",
None,
);
let connections = timeout.run(
clients.assert_connect(&server, socket_addr),
"Clients connecting",
None,
);
connections.srcs().assert_send_all("client->server");
timeout.run(
server.wait_for_total_message_count(clients.len()),
"Waiting for messages",
Some(format_args!("{connections:#?}")),
);
server.assert_message_count("client->server", clients.len());
clients.assert_total_message_count_for_each(0);
}
pub fn test_send_closed_src(&self, timeout: Duration) {
let timeout = Timeout::new(timeout);
let (clients, server) = self.stack_factory.create_stacks::<1>();
let socket_addr = timeout.run(
server.assert_listen(),
"Server listening",
None,
);
let connections = timeout.run(
clients.assert_connect(&server, socket_addr),
"Clients connecting",
None,
);
timeout.run(
connections.close_all_srcs_and_wait(),
"Closing clients",
Some(format_args!("{clients:#?}")),
);
connections.srcs().assert_send_error_all("client->server", &NetworkingError::Closed);
server.assert_total_message_count(0);
clients.assert_total_message_count_for_each(0);
}
pub fn test_send_closed_dst(&self, timeout: Duration) {
let timeout = Timeout::new(timeout);
let (clients, server) = self.stack_factory.create_stacks::<1>();
let socket_addr = timeout.run(
server.assert_listen(),
"Server listening",
None,
);
let connections = timeout.run(
clients.assert_connect(&server, socket_addr),
"Clients connecting",
None,
);
timeout.run(
connections.close_all_dsts_and_wait(),
"Closing server",
Some(format_args!("{server:#?}")),
);
connections.srcs().assert_send_error_all("client->server", &NetworkingError::Closed);
server.assert_total_message_count(0);
clients.assert_total_message_count_for_each(0);
}
pub fn test_send_closed_partial(&self, timeout: Duration) {
let timeout = Timeout::new(timeout);
let (clients, server) = self.stack_factory.create_stacks::<3>();
let socket_addr = timeout.run(
server.assert_listen(),
"Server listening",
None,
);
let connections = timeout.run(
clients.assert_connect(&server, socket_addr),
"Clients connecting",
None,
);
timeout.run(
connections[0].close_src_and_wait(),
"Closing client 0",
Some(format_args!("{clients:#?}")),
);
connections[1..].srcs().assert_send_all("client->server");
timeout.run(
server.wait_for_total_message_count(clients.len() - 1),
"Waiting for messages",
Some(format_args!("{connections:#?}")),
);
server.assert_message_count("client->server", clients.len() - 1);
clients.assert_total_message_count_for_each(0);
}
}
impl<F> ClientServer<F>
where
F: ClientServerStackFactory,
<F as ClientServerStackFactory>::ServerDriver: NetDriverSendTo,
{
pub fn test_send_to(&self, timeout: Duration) {
let timeout = Timeout::new(timeout);
let (clients, server) = self.stack_factory.create_stacks::<1>();
let socket_addr = timeout.run(
server.assert_listen(),
"Server listening",
None,
);
let connections = timeout.run(
clients.assert_connect(&server, socket_addr),
"Clients connecting",
None,
).flip();
connections.srcs().assert_send_to_all("server->client");
timeout.run(
clients.wait_for_total_message_count_for_each(1),
"Waiting for messages",
Some(format_args!("{connections:#?}")),
);
clients.assert_message_count_for_each("server->client", 1);
server.assert_total_message_count(0);
}
pub fn test_send_to_many(&self, timeout: Duration) {
let timeout = Timeout::new(timeout);
let (clients, server) = self.stack_factory.create_stacks::<3>();
let socket_addr = timeout.run(
server.assert_listen(),
"Server listening",
None,
);
let connections = timeout.run(
clients.assert_connect(&server, socket_addr),
"Clients connecting",
None,
).flip();
connections.srcs().assert_send_to_all("server->client");
timeout.run(
clients.wait_for_total_message_count_for_each(1),
"Waiting for messages",
Some(format_args!("{connections:#?}")),
);
clients.assert_message_count_for_each("server->client", 1);
server.assert_total_message_count(0);
}
pub fn test_send_to_closed_src(&self, timeout: Duration) {
let timeout = Timeout::new(timeout);
let (clients, server) = self.stack_factory.create_stacks::<1>();
let socket_addr = timeout.run(
server.assert_listen(),
"Server listening",
None,
);
let connections = timeout.run(
clients.assert_connect(&server, socket_addr),
"Clients connecting",
None,
).flip();
timeout.run(
connections.close_all_srcs_and_wait(),
"Closing server",
Some(format_args!("{clients:#?}")),
);
connections.srcs().assert_send_to_error_all("server->client", &NetworkingError::Closed);
clients.assert_total_message_count_for_each(0);
server.assert_total_message_count(0);
}
pub fn test_send_to_closed_dst(&self, timeout: Duration) {
let timeout = Timeout::new(timeout);
let (clients, server) = self.stack_factory.create_stacks::<1>();
let socket_addr = timeout.run(
server.assert_listen(),
"Server listening",
None,
);
let connections = timeout.run(
clients.assert_connect(&server, socket_addr),
"Clients connecting",
None,
).flip();
timeout.run(
connections.close_all_dsts_and_wait(),
"Closing clients",
Some(format_args!("{clients:#?}")),
);
for src in connections.srcs() {
let connection = src.connection;
src.assert_send_to_error("server->client", &NetworkingError::InvalidConnection(connection));
}
clients.assert_total_message_count_for_each(0);
server.assert_total_message_count(0);
}
pub fn test_send_to_closed_partial(&self, timeout: Duration) {
let timeout = Timeout::new(timeout);
let (clients, server) = self.stack_factory.create_stacks::<1>();
let socket_addr = timeout.run(
server.assert_listen(),
"Server listening",
None,
);
let connections = timeout.run(
clients.assert_connect(&server, socket_addr),
"Clients connecting",
None,
).flip();
timeout.run(
connections[0].close_dst_and_wait(),
"Closing client 0",
Some(format_args!("{clients:#?}")),
);
connections[1..].srcs().assert_send_to_all("server->client");
timeout.run(
clients[1..].wait_for_total_message_count_for_each(1),
"Waiting for messages",
Some(format_args!("{connections:#?}")),
);
clients[0].assert_total_message_count(0);
clients[1..].assert_message_count_for_each("server->client", 1);
server.assert_total_message_count(0);
}
pub fn test_send_to_invalid_connection(&self, timeout: Duration) {
let timeout = Timeout::new(timeout);
let (clients, server) = self.stack_factory.create_stacks::<1>();
let socket_addr = timeout.run(
server.assert_listen(),
"Server listening",
None,
);
let mut connections = timeout.run(
clients.assert_connect(&server, socket_addr),
"Clients connecting",
None,
).flip();
for connection in &mut connections {
connection.src.connection = Connection::INVALID;
}
connections.srcs().assert_send_to_error_all(
"server->client",
&NetworkingError::InvalidConnection(Connection::INVALID),
);
server.assert_total_message_count(0);
clients.assert_total_message_count_for_each(0);
}
}
impl<F> ClientServer<F>
where
F: ClientServerStackFactory,
<F as ClientServerStackFactory>::ClientDriver: NetDriverSend,
<F as ClientServerStackFactory>::ServerDriver: NetDriverSendTo,
{
pub fn test_send_recv(&self, timeout: Duration) {
let timeout = Timeout::new(timeout);
let (clients, server) = self.stack_factory.create_stacks::<1>();
let socket_addr = timeout.run(
server.assert_listen(),
"Server listening",
None,
);
let connections = timeout.run(
clients.assert_connect(&server, socket_addr),
"Clients connecting",
None,
);
server.init_test_message_repliable_send_to_handler(|value| value + "-REPLY");
timeout.run(
connections.srcs().assert_send_recv_all(
"client->server",
"client->server-REPLY",
),
"Waiting for message replies",
Some(format_args!("{connections:#?}")),
);
timeout.run(
server.wait_for_total_message_count(clients.len()),
"Waiting for server-side messages",
Some(format_args!("{connections:#?}")),
);
server.assert_message_count("client->server", clients.len());
timeout.run(
clients.wait_for_total_message_count_for_each(1),
"Waiting for client-side messages",
Some(format_args!("{connections:#?}")),
);
clients.assert_message_count_for_each("client->server-REPLY", 1);
}
pub fn test_send_recv_many(&self, timeout: Duration) {
let timeout = Timeout::new(timeout);
let (clients, server) = self.stack_factory.create_stacks::<3>();
let socket_addr = timeout.run(
server.assert_listen(),
"Server listening",
None,
);
let connections = timeout.run(
clients.assert_connect(&server, socket_addr),
"Clients connecting",
None,
);
server.init_test_message_repliable_send_to_handler(|value| value + "-REPLY");
timeout.run(
connections.srcs().assert_send_recv_all(
"client->server",
"client->server-REPLY",
),
"Waiting for message replies",
Some(format_args!("{connections:#?}")),
);
timeout.run(
server.wait_for_total_message_count(clients.len()),
"Waiting for server-side messages",
Some(format_args!("{connections:#?}")),
);
server.assert_message_count("client->server", clients.len());
timeout.run(
clients.wait_for_total_message_count_for_each(1),
"Waiting for client-side messages",
Some(format_args!("{connections:#?}")),
);
clients.assert_message_count_for_each("client->server-REPLY", 1);
}
pub fn test_send_recv_to(&self, timeout: Duration) {
let timeout = Timeout::new(timeout);
let (clients, server) = self.stack_factory.create_stacks::<1>();
let socket_addr = timeout.run(
server.assert_listen(),
"Server listening",
None,
);
let connections = timeout.run(
clients.assert_connect(&server, socket_addr),
"Clients connecting",
None,
).flip();
for client in &clients {
client.init_test_message_repliable_send_handler(|value| value + "-REPLY");
}
timeout.run(
connections.srcs().assert_send_recv_to_all(
"server->client",
"server->client-REPLY",
),
"Waiting for message replies",
Some(format_args!("{connections:#?}")),
);
timeout.run(
clients.wait_for_total_message_count_for_each(1),
"Waiting for client-side messages",
Some(format_args!("{connections:#?}")),
);
clients.assert_message_count_for_each("server->client", 1);
timeout.run(
server.wait_for_total_message_count(clients.len()),
"Waiting for server-side messages",
Some(format_args!("{connections:#?}")),
);
server.assert_message_count("server->client-REPLY", clients.len());
}
pub fn test_send_recv_to_many(&self, timeout: Duration) {
let timeout = Timeout::new(timeout);
let (clients, server) = self.stack_factory.create_stacks::<3>();
let socket_addr = timeout.run(
server.assert_listen(),
"Server listening",
None,
);
let connections = timeout.run(
clients.assert_connect(&server, socket_addr),
"Clients connecting",
None,
).flip();
for client in &clients {
client.init_test_message_repliable_send_handler(|value| value + "-REPLY");
}
timeout.run(
connections.srcs().assert_send_recv_to_all(
"server->client",
"server->client-REPLY",
),
"Waiting for message replies",
Some(format_args!("{connections:#?}")),
);
timeout.run(
clients.wait_for_total_message_count_for_each(1),
"Waiting for client-side messages",
Some(format_args!("{connections:#?}")),
);
clients.assert_message_count_for_each("server->client", 1);
timeout.run(
server.wait_for_total_message_count(clients.len()),
"Waiting for server-side messages",
Some(format_args!("{connections:#?}")),
);
server.assert_message_count("server->client-REPLY", clients.len());
}
}
impl<F> ClientServer<F>
where
F: ClientServerStackFactory,
<F as ClientServerStackFactory>::ServerDriver: NetDriverBroadcast,
{
pub fn test_broadcast(&self, timeout: Duration) {
let timeout = Timeout::new(timeout);
let (clients, server) = self.stack_factory.create_stacks::<3>();
let socket_addr = timeout.run(
server.assert_listen(),
"Server listening",
None,
);
let _connections = timeout.run(
clients.assert_connect(&server, socket_addr),
"Clients connecting",
None,
).flip();
server.assert_broadcast("server->BROADCAST");
timeout.run(
clients.wait_for_total_message_count_for_each(1),
"Waiting for messages",
Some(format_args!("{clients:#?}")),
);
clients.assert_message_count_for_each("server->BROADCAST", 1);
server.assert_total_message_count(0);
}
pub fn test_broadcast_closed_src(&self, timeout: Duration) {
let timeout = Timeout::new(timeout);
let (clients, server) = self.stack_factory.create_stacks::<3>();
let socket_addr = timeout.run(
server.assert_listen(),
"Server listening",
None,
);
let connections = timeout.run(
clients.assert_connect(&server, socket_addr),
"Clients connecting",
None,
).flip();
timeout.run(
connections.close_all_srcs_and_wait(),
"Closing connections from server",
Some(format_args!("{connections:#?}")),
);
server.assert_broadcast_error("server->BROADCAST", &NetworkingError::Closed);
clients.assert_total_message_count_for_each(0);
server.assert_total_message_count(0);
}
pub fn test_broadcast_closed_dst(&self, timeout: Duration) {
let timeout = Timeout::new(timeout);
let (clients, server) = self.stack_factory.create_stacks::<3>();
let socket_addr = timeout.run(
server.assert_listen(),
"Server listening",
None,
);
let connections = timeout.run(
clients.assert_connect(&server, socket_addr),
"Clients connecting",
None,
).flip();
timeout.run(
connections.close_all_dsts_and_wait(),
"Closing connections from clients",
Some(format_args!("{connections:#?}")),
);
server.assert_broadcast("server->BROADCAST");
clients.assert_total_message_count_for_each(0);
server.assert_total_message_count(0);
}
pub fn test_broadcast_closed_partial(&self, timeout: Duration) {
let timeout = Timeout::new(timeout);
let (clients, server) = self.stack_factory.create_stacks::<3>();
let socket_addr = timeout.run(
server.assert_listen(),
"Server listening",
None,
);
let connections = timeout.run(
clients.assert_connect(&server, socket_addr),
"Clients connecting",
None,
).flip();
timeout.run(
connections[0].close_dst_and_wait(),
"Closing connection 0",
Some(format_args!("{connections:#?}")),
);
server.assert_broadcast("server->BROADCAST");
timeout.run(
clients[1..].wait_for_total_message_count_for_each(1),
"Waiting for messages",
Some(format_args!("{connections:#?}")),
);
clients[0].assert_total_message_count(0);
clients[1..].assert_message_count_for_each("server->BROADCAST", 1);
server.assert_total_message_count(0);
}
}
}
}