mod common;
use std::time::Duration;
use common::{Certs, Server, raw};
use weida::{Error, Runtime, RuntimeConfig, TransferMeta, codes};
use weida_protocol::header::{GuaranteeSet, OrderingMode};
use weida_protocol::{
CreditHeader, DataHeader, ErrorHeader, FrameKind, Hello, MAGIC, SubscriptionHeader,
encode_frame, encode_preamble,
};
const DEADLINE: Duration = Duration::from_secs(15);
async fn within<F: Future>(f: F) -> F::Output {
tokio::time::timeout(DEADLINE, f)
.await
.expect("operation timed out")
}
#[tokio::test]
async fn garbage_first_bytes_close_the_connection() {
let server = Server::start().await;
let _replier = server.listener.replier("/t").expect("replier");
let endpoint = raw::client_endpoint(&server.certs);
let conn = within(endpoint.connect(server.addr, "127.0.0.1").expect("connect"))
.await
.expect("handshake");
raw::send_raw(&conn, &[0xff, 0xff, 0xff, 0xff]).await;
let code = within(raw::closed_code(&conn)).await;
assert_eq!(
code,
codes::PROTOCOL_VIOLATION,
"expected a violation close"
);
let replier = server.listener.replier("/ok").expect("replier");
let handler = tokio::spawn(async move {
let mut request = replier.accept().await.expect("accept");
let body = request.body().read_capped(64).await.expect("body");
assert_eq!(body, b"ping");
let reply = request.reply(TransferMeta::default()).await.expect("reply");
reply.finish().expect("finish");
});
let client = server.client_runtime();
let requester = client.requester(server.trust());
within(requester.connect(&server.url("/ok")))
.await
.expect("connect");
let reply = within(requester.request(b"ping")).await.expect("request");
assert!(within(reply.collect(64)).await.expect("collect").is_empty());
handler.await.expect("handler");
client.shutdown().await;
}
#[tokio::test]
async fn an_unknown_frame_kind_closes_the_connection() {
let server = Server::start().await;
let endpoint = raw::client_endpoint(&server.certs);
let conn = within(endpoint.connect(server.addr, "127.0.0.1").expect("connect"))
.await
.expect("handshake");
raw::send_raw(&conn, &[MAGIC, 9, 0x00]).await;
assert_eq!(
within(raw::closed_code(&conn)).await,
codes::PROTOCOL_VIOLATION
);
}
#[tokio::test]
async fn a_credit_frame_for_a_path_with_no_queue_is_ignored() {
let server = Server::start().await;
let _replier = server.listener.replier("/rpc").expect("replier");
let endpoint = raw::client_endpoint(&server.certs);
let conn = within(endpoint.connect(server.addr, "127.0.0.1").expect("connect"))
.await
.expect("handshake");
raw::send_hello(&conn).await;
let credit = CreditHeader::new("/rpc", "", 7);
raw::send_raw(&conn, &encode_frame(FrameKind::Credit, &credit.encode())).await;
let (mut send, mut recv) = raw::open_exchange(&conn, &DataHeader::addressed("/rpc")).await;
send.write_all(b"ping").await.expect("write request");
send.finish().expect("finish request");
let served = tokio::spawn(async move {
let request = _replier.accept().await.expect("accept");
let mut reply = request
.reply(weida::TransferMeta::default())
.await
.expect("reply");
reply.write_all(b"pong").await.expect("write reply");
reply.finish().expect("finish reply");
});
let mut answer = Vec::new();
within(async {
let mut scratch = [0u8; 256];
while let Ok(Some(n)) = recv.read(&mut scratch).await {
answer.extend_from_slice(&scratch[..n]);
}
})
.await;
within(served).await.expect("served");
assert!(
answer.ends_with(b"pong"),
"the connection did not survive the credit frame: {answer:?}"
);
}
#[tokio::test]
async fn a_malformed_credit_header_closes_the_connection() {
let server = Server::start().await;
let endpoint = raw::client_endpoint(&server.certs);
let conn = within(endpoint.connect(server.addr, "127.0.0.1").expect("connect"))
.await
.expect("handshake");
raw::send_hello(&conn).await;
raw::send_raw(
&conn,
&encode_frame(FrameKind::Credit, &[0xA1, 0x00, 0x62, 0x2F, 0x71]),
)
.await;
assert_eq!(
within(raw::closed_code(&conn)).await,
codes::PROTOCOL_VIOLATION
);
}
#[tokio::test]
async fn hello_with_an_unsupported_version_fails_negotiation() {
let server = Server::start().await;
let endpoint = raw::client_endpoint(&server.certs);
let conn = within(endpoint.connect(server.addr, "127.0.0.1").expect("connect"))
.await
.expect("handshake");
let hello = Hello {
versions: vec![99],
..Hello::v0(16 * 1024, 8)
};
raw::send_frame(&conn, FrameKind::Hello, &hello.encode()).await;
assert_eq!(
within(raw::closed_code(&conn)).await,
codes::NEGOTIATION_FAILED
);
}
#[tokio::test]
async fn a_required_capability_we_lack_fails_negotiation() {
let server = Server::start().await;
let endpoint = raw::client_endpoint(&server.certs);
let conn = within(endpoint.connect(server.addr, "127.0.0.1").expect("connect"))
.await
.expect("handshake");
let hello = Hello {
capabilities: vec![7],
required_capabilities: vec![7],
..Hello::v0(16 * 1024, 8)
};
raw::send_frame(&conn, FrameKind::Hello, &hello.encode()).await;
assert_eq!(
within(raw::closed_code(&conn)).await,
codes::NEGOTIATION_FAILED
);
}
#[tokio::test]
async fn a_required_guarantee_level_we_do_not_offer_fails_negotiation() {
let server = Server::start().await;
let endpoint = raw::client_endpoint(&server.certs);
let conn = within(endpoint.connect(server.addr, "127.0.0.1").expect("connect"))
.await
.expect("handshake");
let wants_ordering = GuaranteeSet {
ordering: OrderingMode::PerProducerDetect,
..GuaranteeSet::CORE
};
let hello = Hello {
guarantees_offered: Some(wants_ordering),
guarantees_required: Some(wants_ordering),
..Hello::v0(16 * 1024, 8)
};
raw::send_frame(&conn, FrameKind::Hello, &hello.encode()).await;
assert_eq!(
within(raw::closed_code(&conn)).await,
codes::NEGOTIATION_FAILED
);
}
#[tokio::test]
async fn a_malformed_guarantee_set_is_a_framing_violation() {
for header in [
raw::hello_with_guarantee_map(&[(4, 99)]),
raw::hello_with_guarantee_map(&[(2, 1)]),
raw::hello_with_guarantee_map(&[(3, 1)]),
raw::hello_with_guarantee_map(&[(5, 1)]),
raw::hello_with_guarantee_map(&[(6, 1_000)]),
raw::hello_with_guarantee_map(&[(9, 2)]),
] {
let server = Server::start().await;
let endpoint = raw::client_endpoint(&server.certs);
let conn = within(endpoint.connect(server.addr, "127.0.0.1").expect("connect"))
.await
.expect("handshake");
raw::send_frame(&conn, FrameKind::Hello, &header).await;
assert_eq!(
within(raw::closed_code(&conn)).await,
codes::PROTOCOL_VIOLATION
);
}
}
#[tokio::test]
async fn an_oversized_header_length_closes_the_connection_before_allocating() {
let server = Server::start().await;
let endpoint = raw::client_endpoint(&server.certs);
let conn = within(endpoint.connect(server.addr, "127.0.0.1").expect("connect"))
.await
.expect("handshake");
let mut bytes = Vec::new();
encode_preamble(FrameKind::Data, 1024 * 1024, &mut bytes);
raw::send_raw(&conn, &bytes).await;
assert_eq!(
within(raw::closed_code(&conn)).await,
codes::PROTOCOL_VIOLATION
);
}
#[tokio::test]
async fn a_malformed_cbor_header_closes_the_connection() {
let server = Server::start().await;
let endpoint = raw::client_endpoint(&server.certs);
let conn = within(endpoint.connect(server.addr, "127.0.0.1").expect("connect"))
.await
.expect("handshake");
raw::send_hello(&conn).await;
let mut bytes = Vec::new();
encode_preamble(FrameKind::Data, 3, &mut bytes);
bytes.extend_from_slice(&[0xff, 0xff, 0xff]);
raw::send_raw(&conn, &bytes).await;
assert_eq!(
within(raw::closed_code(&conn)).await,
codes::PROTOCOL_VIOLATION
);
}
#[tokio::test]
async fn uni_data_without_an_endpoint_is_a_violation() {
let server = Server::start().await;
let _puller = server.listener.puller("/jobs").expect("puller");
let endpoint = raw::client_endpoint(&server.certs);
let conn = within(endpoint.connect(server.addr, "127.0.0.1").expect("connect"))
.await
.expect("handshake");
raw::send_hello(&conn).await;
let header = DataHeader::reply().encode();
let mut bytes = Vec::new();
encode_preamble(FrameKind::Data, header.len() as u64, &mut bytes);
bytes.extend_from_slice(&header);
raw::send_raw(&conn, &bytes).await;
assert_eq!(
within(raw::closed_code(&conn)).await,
codes::PROTOCOL_VIOLATION
);
}
#[tokio::test]
async fn bidi_data_without_an_endpoint_is_a_violation() {
let server = Server::start().await;
let _replier = server.listener.replier("/t").expect("replier");
let endpoint = raw::client_endpoint(&server.certs);
let conn = within(endpoint.connect(server.addr, "127.0.0.1").expect("connect"))
.await
.expect("handshake");
raw::send_hello(&conn).await;
let (mut send, _recv) = conn.open_bi().await.expect("open bi");
let header = DataHeader::reply().encode();
let mut bytes = Vec::new();
encode_preamble(FrameKind::Data, header.len() as u64, &mut bytes);
bytes.extend_from_slice(&header);
send.write_all(&bytes).await.expect("write header");
assert_eq!(
within(raw::closed_code(&conn)).await,
codes::PROTOCOL_VIOLATION
);
}
#[tokio::test]
async fn error_frame_on_a_uni_stream_is_a_violation() {
let server = Server::start().await;
let _replier = server.listener.replier("/t").expect("replier");
let endpoint = raw::client_endpoint(&server.certs);
let conn = within(endpoint.connect(server.addr, "127.0.0.1").expect("connect"))
.await
.expect("handshake");
raw::send_hello(&conn).await;
let header = ErrorHeader::new(weida::ErrorCode::Internal).encode();
raw::send_frame(&conn, FrameKind::Error, &header).await;
assert_eq!(
within(raw::closed_code(&conn)).await,
codes::PROTOCOL_VIOLATION
);
}
#[tokio::test]
async fn a_non_data_frame_may_not_open_a_bidirectional_stream() {
let server = Server::start().await;
let _replier = server.listener.replier("/t").expect("replier");
let endpoint = raw::client_endpoint(&server.certs);
let conn = within(endpoint.connect(server.addr, "127.0.0.1").expect("connect"))
.await
.expect("handshake");
raw::send_hello(&conn).await;
let (mut send, _recv) = conn.open_bi().await.expect("open bi");
send.write_all(&weida_protocol::encode_frame(
FrameKind::Subscribe,
&SubscriptionHeader::new("/t", "").encode(),
))
.await
.expect("write frame");
assert_eq!(
within(raw::closed_code(&conn)).await,
codes::PROTOCOL_VIOLATION
);
}
#[tokio::test]
async fn bidi_request_to_a_pull_path_is_refused_with_unsupported() {
let server = Server::start().await;
let _puller = server.listener.puller("/jobs").expect("puller");
let endpoint = raw::client_endpoint(&server.certs);
let conn = within(endpoint.connect(server.addr, "127.0.0.1").expect("connect"))
.await
.expect("handshake");
raw::send_hello(&conn).await;
let (mut send, mut recv) = conn.open_bi().await.expect("open bi");
let header = DataHeader::addressed("/jobs").encode();
let mut bytes = Vec::new();
encode_preamble(FrameKind::Data, header.len() as u64, &mut bytes);
bytes.extend_from_slice(&header);
send.write_all(&bytes).await.expect("write header");
let (preamble, error_header) = within(raw::read_frame(&mut recv)).await;
assert_eq!(preamble.kind, FrameKind::Error);
let error = ErrorHeader::decode(&error_header).expect("decode ERROR");
assert_eq!(error.code, weida::ErrorCode::Unsupported.to_wire());
let stopped = within(async {
let payload = vec![0u8; 1024 * 1024];
loop {
if let Err(e) = send.write_all(&payload).await {
return e;
}
}
})
.await;
match stopped {
quinn::WriteError::Stopped(code) => assert_eq!(code.into_inner(), codes::UNSUPPORTED),
other => panic!("expected STOP_SENDING(UNSUPPORTED), got {other}"),
}
assert!(conn.close_reason().is_none());
}
#[tokio::test]
async fn bidi_request_to_an_unknown_path_is_refused_with_unknown_endpoint() {
let server = Server::start().await;
let _replier = server.listener.replier("/known").expect("replier");
let endpoint = raw::client_endpoint(&server.certs);
let conn = within(endpoint.connect(server.addr, "127.0.0.1").expect("connect"))
.await
.expect("handshake");
raw::send_hello(&conn).await;
let (mut send, mut recv) = conn.open_bi().await.expect("open bi");
let header = DataHeader::addressed("/nowhere").encode();
let mut bytes = Vec::new();
encode_preamble(FrameKind::Data, header.len() as u64, &mut bytes);
bytes.extend_from_slice(&header);
send.write_all(&bytes).await.expect("write header");
let (preamble, error_header) = within(raw::read_frame(&mut recv)).await;
assert_eq!(preamble.kind, FrameKind::Error);
let error = ErrorHeader::decode(&error_header).expect("decode ERROR");
assert_eq!(error.code, weida::ErrorCode::UnknownEndpoint.to_wire());
assert!(conn.close_reason().is_none());
}
fn raw_server<F, Fut>(certs: &Certs, behaviour: F) -> std::net::SocketAddr
where
F: FnOnce(quinn::Connection) -> Fut + Send + 'static,
Fut: Future<Output = ()> + Send,
{
let (endpoint, addr) = raw::server_endpoint(certs);
tokio::spawn(async move {
let incoming = endpoint.accept().await.expect("a connection");
let conn = incoming.await.expect("handshake");
behaviour(conn).await;
drop(endpoint);
});
addr
}
fn client_for() -> Runtime {
Runtime::new(RuntimeConfig::default()).expect("client runtime")
}
#[tokio::test]
async fn a_server_that_never_answers_yields_indeterminate() {
let certs = Certs::generate();
let addr = raw_server(&certs, |conn| async move {
raw::send_hello(&conn).await;
let (_send, mut recv, header) = raw::accept_exchange(&conn).await;
assert_eq!(header.endpoint.as_deref(), Some("/t"));
let mut sink = vec![0u8; 64 * 1024];
while let Ok(Some(n)) = recv.read(&mut sink).await {
if n == 0 {
break;
}
}
conn.close(quinn::VarInt::from_u32(0), b"bye");
});
let client = client_for();
let requester = client.requester(certs.client_tls());
within(requester.connect(&format!("weida://127.0.0.1:{}/t", addr.port())))
.await
.expect("connect");
let (mut transfer, reply) = within(requester.open(TransferMeta::default()))
.await
.expect("open");
within(transfer.write_all(b"payload")).await.expect("write");
let delivery = transfer.finish().expect("finish");
match within(delivery.delivered()).await {
Ok(()) | Err(Error::Indeterminate) => {}
Err(e) => panic!("unexpected receipt outcome: {e:?}"),
}
let err = within(reply.recv()).await.expect_err("must not succeed");
assert!(
matches!(err, Error::Indeterminate),
"expected Indeterminate, got {err:?}"
);
}
#[tokio::test]
async fn a_server_that_disappears_mid_stream_yields_connection_lost() {
let certs = Certs::generate();
let addr = raw_server(&certs, |conn| async move {
raw::send_hello(&conn).await;
let (_send, mut recv, _header) = raw::accept_exchange(&conn).await;
let mut sink = vec![0u8; 4096];
let _ = recv.read(&mut sink).await;
conn.close(quinn::VarInt::from_u32(0), b"gone");
});
let client = client_for();
let requester = client.requester(certs.client_tls());
within(requester.connect(&format!("weida://127.0.0.1:{}/t", addr.port())))
.await
.expect("connect");
let (mut transfer, _reply) = within(requester.open(TransferMeta::default()))
.await
.expect("open");
let payload = vec![0x5au8; 4 * 1024 * 1024];
let err = within(async {
loop {
if let Err(e) = transfer.write_all(&payload).await {
return e;
}
}
})
.await;
assert!(
matches!(err, Error::ConnectionLost(_)),
"expected ConnectionLost, got {err:?}"
);
assert!(err.is_definite_failure());
}
#[tokio::test]
async fn a_server_that_never_sends_hello_is_dropped_after_the_timeout() {
let certs = Certs::generate();
let addr = raw_server(&certs, |conn| async move {
conn.closed().await;
});
let limits = weida::Limits {
hello_timeout_ms: 500,
..weida::Limits::default()
};
let client = Runtime::new(RuntimeConfig {
limits,
..RuntimeConfig::default()
})
.expect("client runtime");
let requester = client.requester(certs.client_tls());
let err = within(requester.connect(&format!("weida://127.0.0.1:{}/t", addr.port())))
.await
.expect_err("connect must not succeed without negotiation");
assert!(
matches!(err, Error::Negotiation(_) | Error::ConnectionLost(_)),
"expected a negotiation failure, got {err:?}"
);
}
#[tokio::test]
async fn subscribe_with_an_oversized_filter_closes_the_connection() {
let server = Server::start().await;
let _publisher = server.listener.publisher("/md").expect("publisher");
let endpoint = raw::client_endpoint(&server.certs);
let conn = within(endpoint.connect(server.addr, "127.0.0.1").expect("connect"))
.await
.expect("handshake");
raw::send_hello(&conn).await;
let filter = "a".repeat(weida_protocol::header_limits::MAX_FILTER_BYTES + 1);
let header = SubscriptionHeader::new("/md", filter).encode();
raw::send_frame(&conn, FrameKind::Subscribe, &header).await;
assert_eq!(
within(raw::closed_code(&conn)).await,
codes::PROTOCOL_VIOLATION
);
}
#[tokio::test]
async fn a_subscribe_flood_closes_the_connection_with_limit_exceeded() {
let limits = weida::Limits {
max_subscriptions: 8,
..weida::Limits::default()
};
let server = Server::start_with(limits).await;
let _publisher = server.listener.publisher("/md").expect("publisher");
let endpoint = raw::client_endpoint(&server.certs);
let conn = within(endpoint.connect(server.addr, "127.0.0.1").expect("connect"))
.await
.expect("handshake");
raw::send_hello(&conn).await;
for i in 0..=limits.max_subscriptions {
let header = SubscriptionHeader::new("/md", format!("f{i}.")).encode();
let Ok(mut stream) = conn.open_uni().await else {
break;
};
if stream
.write_all(&weida_protocol::encode_frame(FrameKind::Subscribe, &header))
.await
.is_err()
{
break;
}
let _ = stream.finish();
}
assert_eq!(within(raw::closed_code(&conn)).await, codes::LIMIT_EXCEEDED);
}
#[tokio::test]
async fn unsubscribing_an_unknown_filter_is_ignored() {
let server = Server::start().await;
let publisher = server.listener.publisher("/md").expect("publisher");
let endpoint = raw::client_endpoint(&server.certs);
let conn = within(endpoint.connect(server.addr, "127.0.0.1").expect("connect"))
.await
.expect("handshake");
raw::send_hello(&conn).await;
let header = SubscriptionHeader::new("/md", "never.").encode();
raw::send_frame(&conn, FrameKind::Unsubscribe, &header).await;
let header = SubscriptionHeader::new("/nope", "x").encode();
raw::send_frame(&conn, FrameKind::Unsubscribe, &header).await;
let header = SubscriptionHeader::new("/md", "px.").encode();
raw::send_frame(&conn, FrameKind::Subscribe, &header).await;
within(async {
while publisher.filter_count() != 1 {
tokio::time::sleep(Duration::from_millis(1)).await;
}
})
.await;
assert!(conn.close_reason().is_none(), "the connection must survive");
}
#[tokio::test]
async fn data_before_hello_is_parked_not_rejected() {
let server = Server::start().await;
let puller = server.listener.puller("/jobs").expect("puller");
let endpoint = raw::client_endpoint(&server.certs);
let conn = within(endpoint.connect(server.addr, "127.0.0.1").expect("connect"))
.await
.expect("handshake");
let header = DataHeader::addressed("/jobs").encode();
let mut stream = conn.open_uni().await.expect("open uni");
let mut bytes = Vec::new();
encode_preamble(FrameKind::Data, header.len() as u64, &mut bytes);
bytes.extend_from_slice(&header);
stream.write_all(&bytes).await.expect("write header");
tokio::time::sleep(Duration::from_millis(50)).await;
raw::send_hello(&conn).await;
stream
.write_all(b"parked payload")
.await
.expect("write body");
stream.finish().expect("finish");
let transfer = within(puller.recv()).await.expect("recv");
assert_eq!(transfer.meta().endpoint.as_deref(), Some("/jobs"));
let body = within(transfer.collect(1024)).await.expect("collect");
assert_eq!(body, b"parked payload");
assert!(conn.close_reason().is_none(), "the connection must survive");
}
#[tokio::test]
async fn an_exchange_before_hello_is_parked_not_rejected() {
let server = Server::start().await;
let replier = server.listener.replier("/t").expect("replier");
let handler = tokio::spawn(async move {
let mut request = replier.accept().await.expect("accept");
let body = request.body().read_capped(1024).await.expect("body");
assert_eq!(body, b"parked");
let mut reply = request.reply(TransferMeta::default()).await.expect("reply");
reply.write_all(b"ok").await.expect("write reply");
reply.finish().expect("finish");
});
let endpoint = raw::client_endpoint(&server.certs);
let conn = within(endpoint.connect(server.addr, "127.0.0.1").expect("connect"))
.await
.expect("handshake");
let (mut send, mut recv) = conn.open_bi().await.expect("open bi");
let header = DataHeader::addressed("/t").encode();
let mut bytes = Vec::new();
encode_preamble(FrameKind::Data, header.len() as u64, &mut bytes);
bytes.extend_from_slice(&header);
send.write_all(&bytes).await.expect("write header");
tokio::time::sleep(Duration::from_millis(50)).await;
raw::send_hello(&conn).await;
send.write_all(b"parked").await.expect("write body");
send.finish().expect("finish");
let (preamble, _header) = within(raw::read_frame(&mut recv)).await;
assert_eq!(preamble.kind, FrameKind::Data);
handler.await.expect("handler");
assert!(conn.close_reason().is_none(), "the connection must survive");
}