#![allow(clippy::expect_used, clippy::panic, clippy::unwrap_used)]
mod support;
use std::collections::HashSet;
use std::time::Duration;
use ant_quic::{EndpointError, Node, NodeError};
use tokio::time::{sleep, timeout};
fn loopback_addr(node: &Node) -> std::net::SocketAddr {
use std::net::{IpAddr, Ipv4Addr};
let addr = node.local_addr().expect("node bound");
if addr.ip().is_unspecified() {
std::net::SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), addr.port())
} else {
addr
}
}
const ECHO_BYTES: usize = 1024 * 1024;
async fn connected_pair() -> (Node, Node, ant_quic::PeerId, ant_quic::PeerId) {
let a = Node::bind("127.0.0.1:0".parse().expect("addr"))
.await
.expect("node a");
let b = Node::bind("127.0.0.1:0".parse().expect("addr"))
.await
.expect("node b");
let b_addr = loopback_addr(&b);
let a_id = a.peer_id();
let b_id = b.peer_id();
let b_for_accept = b.clone();
tokio::spawn(async move { while b_for_accept.accept().await.is_some() {} });
timeout(Duration::from_secs(10), a.connect_addr(b_addr))
.await
.expect("connect timeout")
.expect("connect failed");
(a, b, a_id, b_id)
}
#[tokio::test]
async fn app_bidi_echo_loopback() {
let _g = support::test_guard().await;
let (a, b, a_id, b_id) = connected_pair().await;
let (mut send_a, mut recv_a) = a.open_bi(&b_id).await.expect("a.open_bi");
let (peer, mut send_b, mut recv_b) = b.accept_bi().await.expect("b.accept_bi");
assert_eq!(peer, a_id, "accept_bi must report the opening peer");
let payload = vec![0x5au8; ECHO_BYTES];
let payload_for_write = payload.clone();
let a_write = tokio::spawn(async move {
send_a
.write_all(&payload_for_write)
.await
.expect("a.write_all");
send_a.finish().expect("a.finish");
});
let b_echo = tokio::spawn(async move {
let copied = tokio::io::copy(&mut recv_b, &mut send_b)
.await
.expect("b echo copy");
send_b.finish().expect("b.finish");
copied
});
let echoed = recv_a
.read_to_end(ECHO_BYTES + 64)
.await
.expect("a.read_to_end echo");
a_write.await.expect("a_write join");
let copied = b_echo.await.expect("b_echo join");
assert_eq!(copied, ECHO_BYTES as u64, "b echoed the full payload");
assert_eq!(echoed.len(), ECHO_BYTES, "echo length matches");
assert_eq!(echoed, payload, "echo byte-integrity");
}
#[tokio::test]
async fn app_bidi_concurrent_with_message_traffic() {
let _g = support::test_guard().await;
let (a, b, a_id, b_id) = connected_pair().await;
const MSG_COUNT: u32 = 32;
let a_for_msgs = a.clone();
let msg_sender = tokio::spawn(async move {
for i in 0..MSG_COUNT {
a_for_msgs
.send(&b_id, &i.to_le_bytes())
.await
.expect("message send");
}
});
let b_for_msgs = b.clone();
let msg_receiver = tokio::spawn(async move {
let mut seen: HashSet<u32> = HashSet::new();
while seen.len() < MSG_COUNT as usize {
let (pid, data) = timeout(Duration::from_secs(10), b_for_msgs.recv())
.await
.expect("recv timeout")
.expect("recv");
assert_eq!(pid, a_id, "message from the connected peer");
assert_eq!(data.len(), 4, "message payload is a 4-byte u32");
let i = u32::from_le_bytes(data[..].try_into().expect("u32 bytes"));
assert!(seen.insert(i), "no duplicate message delivery");
}
seen
});
let (mut send_a, mut recv_a) = a.open_bi(&b_id).await.expect("a.open_bi");
let (peer, mut send_b, mut recv_b) = b.accept_bi().await.expect("b.accept_bi");
assert_eq!(peer, a_id);
let app_payload: Vec<u8> = (0..64 * 1024).map(|x| (x & 0xff) as u8).collect();
let app_payload_for_write = app_payload.clone();
let a_app = tokio::spawn(async move {
send_a
.write_all(&app_payload_for_write)
.await
.expect("app write");
send_a.finish().expect("app finish");
});
let b_app = tokio::spawn(async move {
let copied = tokio::io::copy(&mut recv_b, &mut send_b)
.await
.expect("app echo copy");
send_b.finish().expect("app finish");
copied
});
let echoed = recv_a
.read_to_end(64 * 1024 + 64)
.await
.expect("app readback");
a_app.await.expect("a_app join");
let copied = b_app.await.expect("b_app join");
assert_eq!(copied, 64 * 1024u64);
assert_eq!(echoed, app_payload, "app echo integrity");
msg_sender.await.expect("msg_sender join");
let seen = msg_receiver.await.expect("msg_receiver join");
assert_eq!(seen.len(), MSG_COUNT as usize, "all messages delivered");
}
#[tokio::test]
async fn accept_bi_never_yields_internal_stream() {
let _g = support::test_guard().await;
let (a, b, a_id, b_id) = connected_pair().await;
for i in 0u8..5 {
a.send_with_receive_ack(&b_id, &[i; 7], Duration::from_secs(10))
.await
.expect("ACK-v2 send");
}
for i in 0u8..5 {
a.send(&b_id, &[i; 3]).await.expect("plain send");
}
let b_for_drain = b.clone();
let mut delivered = 0usize;
while delivered < 10 {
let (_pid, _data) = b_for_drain.recv().await.expect("drain recv");
delivered += 1;
}
assert_eq!(delivered, 10, "all internal traffic delivered via recv()");
let leaked = timeout(Duration::from_millis(300), b.accept_bi()).await;
assert!(
leaked.is_err(),
"accept_bi must NOT yield an internal stream; got {leaked:?}"
);
let (mut s1, _r1) = a.open_bi(&b_id).await.expect("open_bi 1");
let (p1, mut rs1, mut rr1) = b.accept_bi().await.expect("accept_bi 1");
assert_eq!(p1, a_id);
let (mut s2, _r2) = a.open_bi(&b_id).await.expect("open_bi 2");
let (p2, mut rs2, mut rr2) = b.accept_bi().await.expect("accept_bi 2");
assert_eq!(p2, a_id);
let extra = timeout(Duration::from_millis(300), b.accept_bi()).await;
assert!(
extra.is_err(),
"accept_bi yielded more than the opened app streams; got {extra:?}"
);
s1.write_all(b"stream-1").await.expect("s1 write");
s1.finish().expect("s1 finish");
s2.write_all(b"stream-2").await.expect("s2 write");
s2.finish().expect("s2 finish");
let m1 = rr1.read_to_end(64).await.expect("r1 read");
let m2 = rr2.read_to_end(64).await.expect("r2 read");
assert_eq!(m1, b"stream-1");
assert_eq!(m2, b"stream-2");
rs1.finish().ok();
rs2.finish().ok();
}
#[tokio::test]
async fn node_open_bi_accept_bi_smoke() {
let _g = support::test_guard().await;
let (a, b, a_id, b_id) = connected_pair().await;
let (mut send_a, mut recv_a) = a.open_bi(&b_id).await.expect("Node::open_bi");
let (peer, mut send_b, mut recv_b) = b.accept_bi().await.expect("Node::accept_bi");
assert_eq!(peer, a_id, "Node::accept_bi reports the opening peer");
let payload = b"node-level app byte-stream echo";
send_a.write_all(payload).await.expect("node write");
send_a.finish().expect("node finish");
let b_echo = tokio::spawn(async move {
let copied = tokio::io::copy(&mut recv_b, &mut send_b)
.await
.expect("node echo copy");
send_b.finish().expect("node echo finish");
copied
});
let echoed = recv_a
.read_to_end(payload.len() + 64)
.await
.expect("node readback");
let copied = b_echo.await.expect("node echo join");
assert_eq!(copied, payload.len() as u64);
assert_eq!(&echoed[..], payload, "node echo integrity");
}
#[tokio::test]
async fn accept_bi_drains_queue_before_shutting_down() {
let _g = support::test_guard().await;
let (a, b, _a_id, b_id) = connected_pair().await;
let b_drainer = b.clone();
const DRAIN_STREAMS: usize = 8;
let mut a_streams: Vec<_> = Vec::with_capacity(DRAIN_STREAMS + 1);
a_streams.push(a.open_bi(&b_id).await.expect("primer open_bi"));
timeout(Duration::from_secs(5), b.accept_bi())
.await
.expect("primer timed out — b.accept() may not have run")
.expect("primer accept_bi error");
for _ in 0..DRAIN_STREAMS {
a_streams.push(a.open_bi(&b_id).await.expect("open_bi"));
}
sleep(Duration::from_millis(100)).await;
b.shutdown().await;
for i in 0..DRAIN_STREAMS {
match b_drainer.accept_bi().await {
Ok(_) => {}
Err(e) => panic!(
"accept_bi[{i}] returned {e:?} before the queue was drained; \
regression: `biased;` may be missing from the internal select"
),
}
}
match b_drainer.accept_bi().await {
Err(NodeError::Endpoint(EndpointError::ShuttingDown)) => {}
other => panic!("expected ShuttingDown after full drain, got {other:?}"),
}
drop(a_streams);
}