mod common;
use std::sync::Arc;
use std::time::Duration;
use futures_util::stream;
use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
use serde_json::{json, Value};
use unb::{handler, Handler, HandlerError, Reply, Request, Streaming};
use unb_client::pair;
use unb_core::{Envelope, Kind};
use unb_runtime::{ClientSession, Wire};
use unb_server::Node;
const CALLS: usize = 100;
const EVENTS: u64 = 5000;
#[derive(Deserialize, JsonSchema)]
struct FirehoseSpec {
count: u64,
}
#[derive(Deserialize, Serialize, JsonSchema)]
#[serde(transparent)]
struct EchoPayload(Value);
#[handler]
async fn firehose(
request: Request<FirehoseSpec>,
) -> Result<Streaming<Value, HandlerError>, HandlerError> {
let count = request.payload().count;
let value = json!({ "data": "x".repeat(1024) });
let payload = Envelope::encode_payload(&value);
let items = (0..count).map(move |_| Ok(payload.clone()));
Ok(Streaming::raw(stream::iter(items)))
}
#[handler]
async fn echo(request: Request<EchoPayload>) -> Result<Reply<EchoPayload>, HandlerError> {
Ok(Reply::new(request.into_payload()))
}
fn firehose_node(name: &str) -> Arc<Node> {
Node::builder(name)
.service(firehose.at_subject(format!("{name}-firehose")))
.insecure_accept_declared_peer_identities()
.build()
.unwrap()
}
async fn consume_firehose(client: ClientSession, subject: &str, count: u64) -> Result<(), String> {
let mut stream = client
.start(
subject,
Kind::Subscribe,
Envelope::encode_payload(&json!({ "count": count })),
None,
Default::default(),
)
.await
.map_err(|error| format!("{subject}: open failed: {error:?}"))?;
let mut events = 0u64;
loop {
match stream.next().await {
Ok(Some(envelope)) => match envelope.kind {
Kind::Event => events += 1,
Kind::Response => {
return (events == count)
.then_some(())
.ok_or(format!("{subject}: terminal after {events}/{count} events"));
}
Kind::Error => {
return Err(format!(
"{subject}: error after {events}/{count} events: {}",
envelope.payload_json()
));
}
_ => {}
},
Ok(None) => return Err(format!("{subject}: wire closed after {events} events")),
Err(error) => return Err(format!("{subject}: {error}")),
}
}
}
#[tokio::test(flavor = "multi_thread")]
async fn relayed_firehoses_in_both_directions_complete() {
let a = firehose_node("a");
let b = firehose_node("b");
let (ta, tb) = pair();
let (connected, _) = tokio::join!(a.connect_transport("b", ta), b.serve_transport(tb));
connected.unwrap();
let (ca_side, a_side) = pair();
a.serve_transport(a_side).await;
let client_a = Wire::open(ca_side);
common::ready_client(&client_a).await;
let (cb_side, b_side) = pair();
b.serve_transport(b_side).await;
let client_b = Wire::open(cb_side);
common::ready_client(&client_b).await;
let downstream = tokio::spawn(async move {
let result = consume_firehose(client_a.client_session(), "/b/b-firehose", EVENTS).await;
client_a.shutdown();
result
});
let upstream = tokio::spawn(async move {
let result = consume_firehose(client_b.client_session(), "/a/a-firehose", EVENTS).await;
client_b.shutdown();
result
});
let completed = tokio::time::timeout(Duration::from_secs(30), async {
let (down, up) = tokio::join!(downstream, upstream);
down.expect("downstream consumer panicked")
.expect("downstream firehose failed");
up.expect("upstream consumer panicked")
.expect("upstream firehose failed");
})
.await;
assert!(
completed.is_ok(),
"duplex firehoses wedged: opposing streams over one link deadlocked within 30s"
);
}
#[tokio::test(flavor = "multi_thread")]
async fn bidirectional_load_completes_without_wedging() {
let node = Node::builder("echo-node")
.service(echo)
.insecure_accept_declared_peer_identities()
.build()
.unwrap();
let (client_side, node_side) = pair();
node.serve_transport(node_side).await;
let wire = Arc::new(Wire::open(client_side));
common::ready_client(&wire).await;
let client = Arc::new(wire.client_session());
let payload = json!({ "data": "x".repeat(8 * 1024) });
let mut tasks = Vec::new();
for _ in 0..CALLS {
let client = client.clone();
let payload = payload.clone();
tasks.push(tokio::spawn(async move {
let mut call = client
.start(
"/echo-node/echo",
Kind::Request,
Envelope::encode_payload(&payload),
None,
Default::default(),
)
.await
.unwrap();
call.next().await
}));
}
let completed = tokio::time::timeout(Duration::from_secs(30), async {
for task in tasks {
assert_eq!(task.await.unwrap().unwrap().unwrap().kind, Kind::Response);
}
})
.await;
assert!(
completed.is_ok(),
"duplex load wedged: request/response flow deadlocked within 30s"
);
wire.shutdown();
}