mod common;
use std::time::Duration;
use bytes::Bytes;
use common::{connect_nodes, start, stream, wait_until};
use schemars::JsonSchema;
use serde::Deserialize;
use serde_json::{json, Value};
use unb::{handler, Handler, HandlerError, Reply, Request};
use unb_client::{dial_transport, pair};
use unb_core::{ErrorCode, Kind};
use unb_runtime::{ClientError, Wire};
use unb_server::{Node, SendExt, ServiceBody};
#[derive(Deserialize, JsonSchema)]
struct CityQuery {
city: Option<String>,
}
#[derive(Deserialize, JsonSchema)]
struct Doubling {
n: i64,
}
#[derive(Deserialize, JsonSchema)]
struct Probe {}
#[handler]
async fn forecast(request: Request<CityQuery>) -> Result<Reply<Value>, HandlerError> {
let city = request.into_payload().city;
Ok(Reply::new(json!({ "temp_c": 21, "city": city })))
}
#[handler]
async fn plan(request: Request<CityQuery>) -> Result<Reply<Value>, HandlerError> {
let city = request.payload().city.clone();
let outlook = request
.call("/tools-1/forecast", json!({ "city": city }))
.await?;
Ok(Reply::new(
json!({ "reply": format!("{}C", outlook["temp_c"]), "via": outlook["city"] }),
))
}
#[handler]
async fn double(request: Request<Doubling>) -> Result<Reply<Value>, HandlerError> {
let n = request.payload().n;
Ok(Reply::new(json!({ "n": n * 2 })))
}
#[handler]
async fn quad(request: Request<Doubling>) -> Result<Reply<Value>, HandlerError> {
let n = request.payload().n;
let once = request.call("/mono-1/double", json!({ "n": n })).await?;
let twice = request
.call("/mono-1/double", json!({ "n": once["n"] }))
.await?;
Ok(Reply::new(json!({ "n": twice["n"] })))
}
#[handler]
async fn failing_forecast(_request: Request<Probe>) -> Result<Reply<Value>, HandlerError> {
Err(HandlerError::new(ErrorCode::InvalidInput, "no such city"))
}
#[handler]
async fn echo_plan(request: Request<Probe>) -> Result<Reply<Value>, HandlerError> {
let outlook = request.call("/tools-1/forecast", json!({})).await?;
Ok(Reply::new(json!({ "reply": outlook })))
}
#[tokio::test(flavor = "multi_thread")]
async fn a_handler_calls_a_downstream_subject_and_returns_its_result_upstream() {
let tools = Node::builder("tools-1")
.service(forecast)
.insecure_accept_declared_peer_identities()
.build()
.unwrap();
let ai = Node::builder("ai-1")
.service(plan.at_subject("assistant.plan"))
.insecure_accept_declared_peer_identities()
.build()
.unwrap();
connect_nodes(&ai, "tools-1", &tools).await;
let (client_side, ai_server) = pair();
ai.serve_transport(ai_server).await;
let client = Wire::open(client_side);
let mut call = stream(
&client,
"/ai-1/assistant/plan",
Kind::Request,
json!({ "city": "hue" }),
)
.await;
let corr = call.operation().as_str().to_owned();
let response = call.next().await.unwrap().unwrap();
assert_eq!(response.kind, Kind::Response);
assert_eq!(response.corr.as_deref(), Some(corr.as_str()));
assert_eq!(response.payload_json()["reply"], "21C");
assert_eq!(
response.payload_json()["via"],
"hue",
"the downstream result flowed back into the AI handler"
);
}
#[tokio::test(flavor = "multi_thread")]
async fn a_handler_can_call_a_sibling_local_subject_in_process() {
let node = Node::builder("mono-1")
.service(double)
.service(quad)
.insecure_accept_declared_peer_identities()
.build()
.unwrap();
let url = start(node).await;
let wire = Wire::open(dial_transport(&url).await.unwrap());
let mut call = stream(&wire, "/mono-1/quad", Kind::Request, json!({ "n": 3 })).await;
let response = call.next().await.unwrap().unwrap();
assert_eq!(
response.payload_json()["n"],
12,
"3 -> double -> double = 12, all in-process"
);
}
#[tokio::test(flavor = "multi_thread")]
async fn a_downstream_error_propagates_back_through_the_calling_handler() {
let tools = Node::builder("tools-1")
.service(failing_forecast.at_subject("forecast"))
.insecure_accept_declared_peer_identities()
.build()
.unwrap();
let ai = Node::builder("ai-1")
.service(echo_plan.at_subject("assistant.plan"))
.insecure_accept_declared_peer_identities()
.build()
.unwrap();
connect_nodes(&ai, "tools-1", &tools).await;
let (client_side, ai_server) = pair();
ai.serve_transport(ai_server).await;
let client = Wire::open(client_side);
let mut call = stream(&client, "/ai-1/assistant/plan", Kind::Request, json!({})).await;
let error = call.next().await.unwrap_err();
assert!(
matches!(
error,
ClientError::Protocol {
code: ErrorCode::InvalidInput,
..
}
),
"the downstream code survives the hop back: {error}"
);
}
#[tokio::test(flavor = "multi_thread")]
async fn routed_non_json_response_bytes_pass_through_fetch_verbatim() {
let caller = Node::builder("caller")
.insecure_accept_declared_peer_identities()
.build()
.unwrap();
let (to_owner, at_owner) = pair();
let owner = std::sync::Arc::new(Wire::open(at_owner));
let (connected, ()) = tokio::join!(
caller.connect_transport("owner", to_owner),
common::ghost_establish_with(&owner, "owner", &["remote"])
);
connected.unwrap();
wait_until("caller sees the ghost owner's route", || {
caller.reachable_names().contains(&"owner".to_string())
})
.await;
let response_owner = owner.clone();
let responder = tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(10)).await;
response_owner
.respond_with("s1", Bytes::from_static(b"{not-json"), Default::default())
.await
.unwrap();
});
let fetch = tokio::time::timeout(
Duration::from_secs(5),
caller.fetch(
http::Request::builder()
.method(http::Method::POST)
.uri("/owner/remote")
.body(Bytes::new())
.unwrap(),
),
)
.await
.expect("non-json routed response must complete");
responder.abort();
let response = fetch.unwrap();
let ServiceBody::Unary(payload) = response.into_body() else {
panic!("expected a unary response");
};
assert_eq!(
payload,
Bytes::from_static(b"{not-json"),
"peer output bytes reach the fetch caller verbatim"
);
let response_owner = owner.clone();
let responder = tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(10)).await;
response_owner
.respond_with("s3", Bytes::new(), Default::default())
.await
.unwrap();
});
let empty_fetch = tokio::time::timeout(
Duration::from_secs(5),
caller.fetch(
http::Request::builder()
.method(http::Method::POST)
.uri("/owner/remote")
.body(Bytes::new())
.unwrap(),
),
)
.await
.expect("empty routed response must complete");
responder.abort();
let response = empty_fetch.unwrap();
assert!(matches!(
response.into_body(),
ServiceBody::Unary(payload) if payload.is_empty()
));
}
#[tokio::test(flavor = "multi_thread")]
async fn an_outbound_response_larger_than_the_frame_limit_is_refused() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let mut scratch = [0u8; 2048];
let _ = stream.read(&mut scratch).await;
let oversized = 16 * 1024 * 1024 + 1;
let head = format!("HTTP/1.1 200 OK\r\ncontent-length: {oversized}\r\n\r\n");
if stream.write_all(head.as_bytes()).await.is_err() {
return;
}
let chunk = vec![0u8; 64 * 1024];
let mut written = 0;
while written < oversized {
let take = (oversized - written).min(chunk.len());
if stream.write_all(&chunk[..take]).await.is_err() {
break;
}
written += take;
}
});
let outcome = http::Request::builder()
.uri("/remote/oversized")
.body(json!({}))
.send(format!("http://{address}"))
.await;
let error = outcome.expect_err("an oversized outbound response must be refused, not buffered");
assert_eq!(error.code, ErrorCode::Protocol);
server.abort();
}