mod common;
use std::sync::Arc;
use std::time::Duration;
use common::{Certs, Server, Xorshift};
use tokio::io::AsyncReadExt;
use weida::{Error, TraceContext, TransferMeta};
async fn spawn_uppercase_handler(replier: weida::Replier) -> tokio::task::JoinHandle<()> {
tokio::spawn(async move {
let mut request = replier.accept().await.expect("accept");
let mut body = request.take_body();
let mut out = request
.reply(TransferMeta::default())
.await
.expect("open reply");
let mut chunk = vec![0u8; 64 * 1024];
loop {
let n = body.read(&mut chunk).await.expect("read body");
if n == 0 {
break;
}
chunk[..n].make_ascii_uppercase();
out.write_all(&chunk[..n]).await.expect("write reply");
}
out.finish().expect("finish reply");
})
}
#[tokio::test]
async fn echo_roundtrip() {
let server = Server::start().await;
let replier = server.listener.replier("/transform").expect("replier");
let handler = spawn_uppercase_handler(replier).await;
let client = server.client_runtime();
let requester = client.requester(server.trust());
requester
.connect(&server.url("/transform"))
.await
.expect("connect");
let body = "hello weida ".repeat(64);
assert_eq!(body.len(), 768);
let (mut transfer, reply) = requester.open(TransferMeta::default()).await.expect("open");
transfer.write_all(body.as_bytes()).await.expect("write");
let delivery = transfer.finish().expect("finish");
delivery.delivered().await.expect("delivered");
let received = reply
.recv()
.await
.expect("recv reply")
.collect(1 << 20)
.await
.expect("collect");
assert_eq!(received, body.to_uppercase().into_bytes());
handler.await.expect("handler");
client.shutdown().await;
}
#[tokio::test]
async fn request_convenience_returns_the_reply_stream() {
let server = Server::start().await;
let replier = server.listener.replier("/transform").expect("replier");
let handler = spawn_uppercase_handler(replier).await;
let client = server.client_runtime();
let requester = client.requester(server.trust());
requester
.connect(&server.url("/transform"))
.await
.expect("connect");
let reply = requester.request(b"abc").await.expect("request");
assert_eq!(reply.collect(16).await.expect("collect"), b"ABC");
handler.await.expect("handler");
client.shutdown().await;
}
#[tokio::test]
async fn streaming_overlap() {
const HALF: usize = 4 * 1024 * 1024;
let server = Server::start().await;
let replier = server.listener.replier("/transform").expect("replier");
let handler = spawn_uppercase_handler(replier).await;
let client = server.client_runtime();
let requester = client.requester(server.trust());
requester
.connect(&server.url("/transform"))
.await
.expect("connect");
let mut rng = Xorshift::new(0x51ee_d0d0_5eed_0001);
let mut payload = vec![0u8; HALF * 2];
rng.fill(&mut payload);
for b in payload.iter_mut() {
*b = b'a' + (*b % 26);
}
let (mut transfer, reply) = requester
.open(TransferMeta::default().with_content_len((HALF * 2) as u64))
.await
.expect("open");
let (first_byte_tx, first_byte_rx) = tokio::sync::oneshot::channel();
let reader = tokio::spawn(async move {
let mut body = reply.recv().await.expect("early reply header");
let mut received = Vec::with_capacity(HALF * 2);
let mut chunk = vec![0u8; 64 * 1024];
let mut announce = Some(first_byte_tx);
loop {
let n = body.read(&mut chunk).await.expect("read reply");
if n == 0 {
break;
}
received.extend_from_slice(&chunk[..n]);
if let Some(tx) = announce.take() {
let _ = tx.send(received[0]);
}
}
received
});
transfer.write_all(&payload[..HALF]).await.expect("write 1");
let first = tokio::time::timeout(Duration::from_secs(30), first_byte_rx)
.await
.expect("a reply byte must arrive before the request is finished")
.expect("reader task alive");
assert_eq!(first, payload[0].to_ascii_uppercase());
transfer.write_all(&payload[HALF..]).await.expect("write 2");
transfer.finish().expect("finish");
let received = reader.await.expect("reader task");
assert_eq!(received.len(), payload.len());
let expected: Vec<u8> = payload.iter().map(|b| b.to_ascii_uppercase()).collect();
assert_eq!(received, expected);
handler.await.expect("handler");
client.shutdown().await;
}
#[tokio::test]
async fn cancel_mid_transfer() {
let server = Server::start().await;
let replier = server.listener.replier("/sink").expect("replier");
let observed = Arc::new(tokio::sync::Mutex::new(None::<Error>));
let server_side = Arc::clone(&observed);
let handler = tokio::spawn(async move {
let mut request = replier.accept().await.expect("accept");
let mut chunk = vec![0u8; 64 * 1024];
loop {
match request.body().read(&mut chunk).await {
Ok(0) => break,
Ok(_) => continue,
Err(e) => {
*server_side.lock().await = Some(Error::Io(e));
break;
}
}
}
drop(request);
});
let client = server.client_runtime();
let requester = client.requester(server.trust());
requester
.connect(&server.url("/sink"))
.await
.expect("connect");
let (mut transfer, reply) = requester.open(TransferMeta::default()).await.expect("open");
let chunk = vec![0x61u8; 256 * 1024];
for _ in 0..4 {
transfer.write_all(&chunk).await.expect("write");
}
transfer.cancel();
drop(reply);
handler.await.expect("handler");
assert!(
observed.lock().await.is_some(),
"the server must observe the reset instead of a clean EOF"
);
let replier2 = server.listener.replier("/again").expect("replier");
let handler2 = spawn_uppercase_handler(replier2).await;
let requester2 = client.requester(server.trust());
requester2
.connect(&server.url("/again"))
.await
.expect("reconnect");
let reply = requester2.request(b"ok").await.expect("second request");
assert_eq!(reply.collect(16).await.expect("collect"), b"OK");
handler2.await.expect("handler2");
client.shutdown().await;
}
#[tokio::test]
async fn reply_abort_when_reply_stream_dropped() {
let server = Server::start().await;
let replier = server.listener.replier("/firehose").expect("replier");
let saw_cancel = Arc::new(tokio::sync::Mutex::new(false));
let server_side = Arc::clone(&saw_cancel);
let handler = tokio::spawn(async move {
let mut request = replier.accept().await.expect("accept");
let _ = request.body().read_capped(1024).await;
let canceled = request.canceled();
tokio::pin!(canceled);
let mut reply = match request.reply(TransferMeta::default()).await {
Ok(reply) => reply,
Err(e) => {
assert!(matches!(e, Error::Canceled), "{e:?}");
*server_side.lock().await = true;
return;
}
};
let chunk = vec![0x7au8; 64 * 1024];
loop {
tokio::select! {
_ = &mut canceled => {
*server_side.lock().await = true;
break;
}
written = reply.write_all(&chunk) => {
if let Err(e) = written {
assert!(matches!(e, Error::Canceled), "{e:?}");
*server_side.lock().await = true;
break;
}
}
}
}
drop(reply);
});
let client = server.client_runtime();
let requester = client.requester(server.trust());
requester
.connect(&server.url("/firehose"))
.await
.expect("connect");
let (mut transfer, reply) = requester.open(TransferMeta::default()).await.expect("open");
transfer.write_all(b"start").await.expect("write");
transfer.finish().expect("finish");
drop(reply);
tokio::time::timeout(Duration::from_secs(10), handler)
.await
.expect("handler must notice the cancellation")
.expect("handler");
assert!(*saw_cancel.lock().await);
let replier2 = server.listener.replier("/after").expect("replier");
let handler2 = spawn_uppercase_handler(replier2).await;
let requester2 = client.requester(server.trust());
requester2
.connect(&server.url("/after"))
.await
.expect("reconnect");
let reply = requester2.request(b"live").await.expect("request");
assert_eq!(reply.collect(16).await.expect("collect"), b"LIVE");
handler2.await.expect("handler2");
client.shutdown().await;
}
#[tokio::test]
async fn canceled_resolves_when_the_requester_walks_away() {
let server = Server::start().await;
let replier = server.listener.replier("/watch").expect("replier");
let handler = tokio::spawn(async move {
let mut request = replier.accept().await.expect("accept");
let _ = request.body().read_capped(1024).await;
let canceled = request.canceled();
let mut reply = request
.reply(TransferMeta::default())
.await
.expect("open reply");
reply.write_all(b"live").await.expect("write reply");
canceled.await;
drop(reply);
});
let client = server.client_runtime();
let requester = client.requester(server.trust());
requester
.connect(&server.url("/watch"))
.await
.expect("connect");
let (mut transfer, reply) = requester.open(TransferMeta::default()).await.expect("open");
transfer.write_all(b"watch me").await.expect("write");
transfer.finish().expect("finish");
let mut body = reply.recv().await.expect("recv reply");
let mut marker = [0u8; 4];
body.read_exact(&mut marker)
.await
.expect("read the liveness marker");
assert_eq!(&marker, b"live");
drop(body);
tokio::time::timeout(Duration::from_secs(10), handler)
.await
.expect("canceled() must resolve when the requester walks away")
.expect("handler");
client.shutdown().await;
}
#[tokio::test]
async fn a_transfer_carries_no_trace_context_unless_the_sender_supplied_one() {
let server = Server::start().await;
let replier = server.listener.replier("/untraced").expect("replier");
let seen = Arc::new(tokio::sync::Mutex::new(Some(())));
let server_side = Arc::clone(&seen);
let handler = tokio::spawn(async move {
let mut request = replier.accept().await.expect("accept");
*server_side.lock().await = request.meta().trace.map(|_| ());
let _ = request.body().read_capped(64).await.expect("read");
let mut reply = request.reply(TransferMeta::default()).await.expect("reply");
reply.write_all(b"y").await.expect("write");
reply.finish().expect("finish");
});
let client = server.client_runtime();
let requester = client.requester(server.trust());
requester
.connect(&server.url("/untraced"))
.await
.expect("connect");
let (mut transfer, reply) = requester.open(TransferMeta::default()).await.expect("open");
assert_eq!(
transfer.trace(),
None,
"nothing mints a context on the caller's behalf"
);
transfer.write_all(b"x").await.expect("write");
transfer.finish().expect("finish");
let body = reply.recv().await.expect("recv reply");
assert!(
body.meta().trace.is_none(),
"and the reply half carries none either"
);
handler.await.expect("handler");
assert!(
seen.lock().await.is_none(),
"the responder saw no trace context on a request that started no trace"
);
client.shutdown().await;
}
#[tokio::test]
async fn trace_propagation() {
let server = Server::start().await;
let replier = server.listener.replier("/traced").expect("replier");
let seen = Arc::new(tokio::sync::Mutex::new(None::<TraceContext>));
let server_side = Arc::clone(&seen);
let handler = tokio::spawn(async move {
let mut request = replier.accept().await.expect("accept");
*server_side.lock().await = request.meta().trace;
let _ = request.body().read_capped(1024).await;
let reply = request
.reply(TransferMeta::default())
.await
.expect("open reply");
reply.finish().expect("finish reply");
});
let client = server.client_runtime();
let requester = client.requester(server.trust());
requester
.connect(&server.url("/traced"))
.await
.expect("connect");
let trace =
TraceContext::parse_traceparent("00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01")
.expect("valid traceparent");
let (mut transfer, reply) = requester
.open(TransferMeta::default().with_trace(trace))
.await
.expect("open");
assert_eq!(transfer.trace(), Some(trace));
transfer.write_all(b"x").await.expect("write");
transfer.finish().expect("finish");
let body = reply.recv().await.expect("recv reply");
let reply_trace = body.meta().trace.expect("reply carries a trace context");
assert_eq!(reply_trace.trace_id, trace.trace_id);
assert_eq!(reply_trace.span_id, trace.span_id);
handler.await.expect("handler");
assert_eq!(seen.lock().await.expect("server saw a trace"), trace);
client.shutdown().await;
}
#[tokio::test]
async fn unknown_endpoint_is_reported() {
let server = Server::start().await;
let _replier = server.listener.replier("/known").expect("replier");
let client = server.client_runtime();
let requester = client.requester(server.trust());
requester
.connect(&server.url("/nowhere"))
.await
.expect("connect");
let (mut transfer, reply) = requester.open(TransferMeta::default()).await.expect("open");
let err = match transfer.write_all(&vec![0u8; 4096]).await {
Err(e) => e,
Ok(()) => {
let _ = transfer.finish();
reply.recv().await.expect_err("reply must fail")
}
};
assert!(
matches!(err, Error::UnknownEndpoint),
"expected UnknownEndpoint, got {err:?}"
);
client.shutdown().await;
}
#[tokio::test]
async fn dropping_a_request_without_replying_reports_no_reply() {
let server = Server::start().await;
let replier = server.listener.replier("/silent").expect("replier");
let handler = tokio::spawn(async move {
let mut request = replier.accept().await.expect("accept");
let _ = request.body().read_capped(1024).await;
drop(request);
});
let client = server.client_runtime();
let requester = client.requester(server.trust());
requester
.connect(&server.url("/silent"))
.await
.expect("connect");
let (mut transfer, reply) = requester.open(TransferMeta::default()).await.expect("open");
transfer.write_all(b"anyone there?").await.expect("write");
transfer.finish().expect("finish");
let err = reply.recv().await.expect_err("reply must fail");
assert!(
matches!(err, Error::NoReply),
"expected NoReply, got {err:?}"
);
handler.await.expect("handler");
client.shutdown().await;
}
#[tokio::test]
async fn concurrent_exchanges_do_not_interfere() {
const COUNT: usize = 16;
let server = Server::start().await;
let replier = server.listener.replier("/transform").expect("replier");
tokio::spawn(async move {
while let Ok(mut request) = replier.accept().await {
tokio::spawn(async move {
let mut body = request.take_body();
let payload = body.read_capped(1024).await.expect("read body");
let mut out = request
.reply(TransferMeta::default())
.await
.expect("open reply");
out.write_all(&payload.to_ascii_uppercase())
.await
.expect("write reply");
out.finish().expect("finish reply");
});
}
});
let client = server.client_runtime();
let requester = Arc::new(client.requester(server.trust()));
requester
.connect(&server.url("/transform"))
.await
.expect("connect");
let mut tasks = Vec::with_capacity(COUNT);
for i in 0..COUNT {
let requester = Arc::clone(&requester);
tasks.push(tokio::spawn(async move {
let body = format!("exchange-{i}");
let reply = requester
.request(body.as_bytes())
.await
.expect("request")
.collect(1024)
.await
.expect("collect");
assert_eq!(reply, body.to_uppercase().into_bytes());
}));
}
for task in tasks {
task.await.expect("exchange");
}
client.shutdown().await;
}
#[tokio::test]
async fn multi_peer_requests_round_robin() {
let a = Server::start().await;
let b = Server::start().await;
let handler_a = spawn_uppercase_handler(a.listener.replier("/rr").expect("replier")).await;
let handler_b = spawn_uppercase_handler(b.listener.replier("/rr").expect("replier")).await;
let client = weida::Runtime::new(weida::RuntimeConfig::default()).expect("client runtime");
let requester = client
.requester(weida::Trust::anchor_file(&a.certs.cert_pem).and_anchor_file(&b.certs.cert_pem));
requester.connect(&a.url("/rr")).await.expect("connect a");
requester.connect(&b.url("/rr")).await.expect("connect b");
assert_eq!(requester.peer_count(), 2);
for _ in 0..2 {
let reply = requester.request(b"peer").await.expect("request");
assert_eq!(reply.collect(16).await.expect("collect"), b"PEER");
}
handler_a.await.expect("handler a");
handler_b.await.expect("handler b");
client.shutdown().await;
}
#[tokio::test]
async fn endpoints_with_different_trust_do_not_share_a_connection() {
let server = Server::start().await;
let handler = spawn_uppercase_handler(server.listener.replier("/t").expect("replier")).await;
let client = server.client_runtime();
let trusting = client.requester(server.trust());
trusting.connect(&server.url("/t")).await.expect("connect");
let reply = trusting.request(b"hi").await.expect("request");
assert_eq!(reply.collect(64).await.expect("collect"), b"HI");
let stranger = Certs::generate();
let distrusting = client.requester(stranger.client_tls());
let err = distrusting
.connect(&server.url("/t"))
.await
.expect_err("a peer outside our trust must not be accepted");
assert_eq!(
match err {
Error::Untrusted(presented) => presented,
other => panic!("expected Untrusted, got {other:?}"),
},
server.certs.fingerprint(),
"the refusal names the identity that answered"
);
handler.await.expect("handler");
client.shutdown().await;
}
#[tokio::test]
async fn duplicate_endpoint_registration_is_rejected() {
let server = Server::start().await;
let _first = server.listener.replier("/dup").expect("first");
assert!(matches!(
server.listener.replier("/dup"),
Err(Error::AlreadyRegistered)
));
assert!(matches!(
server.listener.replier("no-slash"),
Err(Error::InvalidEndpointPath)
));
}