mod common;
use common::TestCall;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::Duration;
use schemars::JsonSchema;
use serde::Deserialize;
use serde_json::{json, Value};
use unb::{handler, Handler, HandlerError, Node, Reply, Request, State, Streaming};
use unb_client::pair;
use unb_core::{Envelope, ErrorCode, Kind, PROTOCOL_VERSION};
use unb_runtime::{ClientError, Pipe, Wire};
#[derive(Deserialize, JsonSchema)]
struct Probe {}
#[derive(Clone)]
struct Count(u64);
#[derive(Clone)]
struct Order(Vec<String>);
#[derive(Clone)]
struct Actor(String);
struct HandlerRan(Arc<AtomicBool>);
#[handler]
async fn inner(request: Request<Probe>) -> Result<Reply<Value>, HandlerError> {
Ok(Reply::new(json!({
"count": request.extensions().get::<Count>().map(|count| count.0),
"actor": request.headers().get("actor").and_then(|value| value.to_str().ok()),
})))
}
#[handler]
async fn outer(request: Request<Probe>) -> Result<Reply<Value>, HandlerError> {
let value = request.call("/local-middleware/inner", json!({})).await?;
Ok(Reply::new(value))
}
#[handler]
async fn audit(request: Request<Probe>) -> Result<Reply<Value>, HandlerError> {
let order = request
.extensions()
.get::<Order>()
.map(|order| order.0.clone())
.unwrap();
Ok(Reply::new(json!({ "order": order })))
}
#[handler]
async fn whoami(request: Request<Probe>) -> Result<Reply<Value>, HandlerError> {
Ok(Reply::new(json!({
"actor": request.extensions().get::<Actor>().map(|actor| actor.0.clone()),
})))
}
#[handler]
async fn secure(
ran: State<HandlerRan>,
_request: Request<Probe>,
) -> Result<Reply<Value>, HandlerError> {
ran.0.store(true, Ordering::SeqCst);
Ok(Reply::new(json!({})))
}
#[handler]
async fn feed(_request: Request<Probe>) -> Result<Streaming<Value, HandlerError>, HandlerError> {
panic!("a rejected subscribe must never reach its handler")
}
#[handler]
async fn whoami_authorization(request: Request<Probe>) -> Result<Reply<Value>, HandlerError> {
Ok(Reply::new(
json!({ "actor": request.headers().get("authorization").and_then(|value| value.to_str().ok()) }),
))
}
#[tokio::test(flavor = "multi_thread")]
async fn local_and_nested_calls_run_middleware_with_fresh_extensions_and_propagated_headers() {
let node = Node::builder("local-middleware")
.layer_fn(|mut request, next| async move {
let count = request
.extensions()
.get::<Count>()
.map(|count| count.0)
.unwrap_or(0);
request.extensions_mut().insert(Count(count + 1));
next.run(request).await
})
.service(inner)
.service(outer)
.insecure_accept_declared_peer_identities()
.build()
.unwrap();
let (client_side, node_side) = pair();
node.serve_transport(node_side).await;
let client = Wire::open(client_side);
common::ready_client(&client).await;
let direct = node
.request("/local-middleware/inner", json!({}))
.await
.unwrap();
assert_eq!(direct["count"], 1);
let mut headers = serde_json::Map::new();
headers.insert("actor".into(), json!("nested"));
let mut call = client
.client_session()
.start(
"/local-middleware/outer",
Kind::Request,
Envelope::encode_payload(&json!({})),
None,
headers,
)
.await
.unwrap();
let nested = call.next().await.unwrap().unwrap();
assert_eq!(nested.kind, Kind::Response);
assert_eq!(
nested.payload_json(),
json!({"count": 1, "actor": "nested"})
);
}
fn envelope(kind: Kind, subject: &str, corr: Option<&str>, payload: serde_json::Value) -> Envelope {
Envelope {
v: PROTOCOL_VERSION,
id: "f1".into(),
target: String::new(),
subject: subject.into(),
kind,
corr: corr.map(str::to_owned),
seq: None,
hops: None,
body_token: None,
payload: Envelope::encode_payload(&payload),
path: Vec::new(),
headers: Default::default(),
}
}
#[tokio::test(flavor = "multi_thread")]
async fn middleware_chain_runs_in_order_before_the_handler() {
let node = Node::builder("gate-1")
.layer_fn(|mut request, next| async move {
request.extensions_mut().insert(Order(vec!["first".into()]));
next.run(request).await
})
.layer_fn(|mut request, next| async move {
let mut order = request.extensions().get::<Order>().cloned().unwrap();
order.0.push("second".into());
request.extensions_mut().insert(order);
next.run(request).await
})
.service(audit)
.insecure_accept_declared_peer_identities()
.build()
.unwrap();
let (client_side, node_side) = pair();
node.serve_transport(node_side).await;
let client = Wire::open(client_side);
common::ready_client(&client).await;
let mut call = common::stream(&client, "/gate-1/audit", Kind::Request, json!({})).await;
let response = call.next().await.unwrap().unwrap();
assert_eq!(response.kind, Kind::Response);
assert_eq!(response.payload_json()["order"], json!(["first", "second"]));
}
#[tokio::test(flavor = "multi_thread")]
async fn middleware_enrichment_from_a_header_is_visible_to_the_handler() {
let node = Node::builder("gate-2")
.layer_fn(|mut request, next| async move {
let actor = request
.headers()
.get("actor")
.and_then(|value| value.to_str().ok())
.map(str::to_owned);
let Some(actor) = actor else {
return Err(HandlerError::new(
ErrorCode::Unauthorized,
"an actor header is required",
));
};
let vetted = format!("verified:{actor}");
request.extensions_mut().insert(Actor(vetted));
next.run(request).await
})
.service(whoami)
.insecure_accept_declared_peer_identities()
.build()
.unwrap();
let (manual_side, node_side) = pair();
node.serve_transport(node_side).await;
let Pipe::Local {
rx: mut manual_rx,
tx: manual_tx,
..
} = manual_side
else {
panic!("pair() returns local transports");
};
manual_tx
.send(envelope(Kind::Hello, "", None, json!({ "versions": [1] })))
.await
.unwrap();
assert_eq!(manual_rx.recv().await.unwrap().kind, Kind::Welcome);
let mut request = envelope(Kind::Request, "whoami", Some("s1"), json!({}));
request.target = "gate-2".into();
request.headers.insert("actor".into(), json!("jwt-abc"));
manual_tx.send(request).await.unwrap();
let response = manual_rx.recv().await.unwrap();
assert_eq!(response.kind, Kind::Response);
assert_eq!(response.payload_json()["actor"], "verified:jwt-abc");
}
#[tokio::test(flavor = "multi_thread")]
async fn a_middleware_reject_fails_the_stream_and_the_handler_never_runs() {
let handler_ran = Arc::new(AtomicBool::new(false));
let node = Node::builder("gate-3")
.layer_fn(|_request, _next| async move {
Err(HandlerError::new(ErrorCode::Unauthorized, "no credential"))
})
.state(HandlerRan(handler_ran.clone()))
.service(secure)
.insecure_accept_declared_peer_identities()
.build()
.unwrap();
let (client_side, node_side) = pair();
node.serve_transport(node_side).await;
let client = Wire::open(client_side);
common::ready_client(&client).await;
let mut call = common::stream(&client, "/gate-3/secure", Kind::Request, json!({})).await;
assert!(matches!(
call.next().await,
Err(ClientError::Protocol {
code: ErrorCode::Unauthorized,
..
})
));
assert!(
!handler_ran.load(Ordering::SeqCst),
"a rejected request must never reach its handler"
);
}
#[tokio::test(flavor = "multi_thread")]
async fn a_rejected_subscribe_opens_no_event_stream() {
let node = Node::builder("gate-4")
.layer_fn(|_request, _next| async move {
Err(HandlerError::new(ErrorCode::Unauthorized, "no credential"))
})
.service(feed)
.insecure_accept_declared_peer_identities()
.build()
.unwrap();
let (client_side, node_side) = pair();
node.serve_transport(node_side).await;
let client = Wire::open(client_side);
common::ready_client(&client).await;
let mut call = common::stream(&client, "/gate-4/feed", Kind::Subscribe, json!({})).await;
assert!(matches!(
call.next().await,
Err(ClientError::Protocol {
code: ErrorCode::Unauthorized,
..
})
));
}
#[tokio::test(flavor = "multi_thread")]
async fn relayed_frames_bypass_middleware_and_keep_custom_headers() {
let stage = Arc::new(std::sync::atomic::AtomicU8::new(0));
let observed_stage = stage.clone();
let result = tokio::time::timeout(Duration::from_secs(3), async move {
let middleware_ran = Arc::new(AtomicBool::new(false));
let observed = middleware_ran.clone();
let relay = Node::builder("relay-1")
.layer_fn(move |request, next| {
let middleware_ran = observed.clone();
async move {
middleware_ran.store(true, Ordering::SeqCst);
next.run(request).await
}
})
.insecure_accept_declared_peer_identities()
.build()
.unwrap();
let (to_owner, at_owner) = pair();
let owner = Arc::new(Wire::open(at_owner));
let (connected, ()) = tokio::join!(
relay.connect_transport("owner", to_owner),
common::ghost_establish_with(&owner, "owner", &["remote.echo"])
);
connected.unwrap();
observed_stage.store(1, Ordering::SeqCst);
let (client_side, relay_side) = pair();
relay.serve_transport(relay_side).await;
let client = Wire::open(client_side);
common::ready_client(&client).await;
observed_stage.store(2, Ordering::SeqCst);
let mut request = envelope(Kind::Request, "remote.echo", Some("s1"), json!({ "n": 1 }));
request.target = "owner".into();
request.headers.insert("actor".into(), json!("jwt-abc"));
request.headers.insert("trace".into(), json!("span-7"));
let source_payload = request.payload.clone();
let mut observer = owner.observe();
let mut call = client
.client_session()
.start(
"/owner/remote/echo",
request.kind,
request.payload,
request.hops,
request.headers,
)
.await
.unwrap_or_else(|error| panic!("relay request opening failed: {error}"));
let client_corr = call.operation().as_str().to_owned();
observed_stage.store(3, Ordering::SeqCst);
let forwarded = common::deliver_observed(&mut observer).await;
observed_stage.store(4, Ordering::SeqCst);
assert_eq!(forwarded.subject, "remote.echo");
assert_eq!(
forwarded.headers["actor"], "jwt-abc",
"a relay forwards custom headers unchanged"
);
assert_eq!(forwarded.headers["trace"], "span-7");
assert_eq!(
forwarded.payload.as_ptr(),
source_payload.as_ptr(),
"a relay forwards the body without decoding or copying it"
);
assert!(
!middleware_ran.load(Ordering::SeqCst),
"middleware must not intercept relayed frames"
);
owner
.respond(forwarded.corr.as_deref().unwrap(), json!({ "ok": true }))
.await
.unwrap();
observed_stage.store(5, Ordering::SeqCst);
let response = call.next().await.unwrap().unwrap();
assert_eq!(response.kind, Kind::Response);
assert_eq!(response.payload_json(), json!({ "ok": true }));
assert_eq!(response.corr.as_deref(), Some(client_corr.as_str()));
observed_stage.store(6, Ordering::SeqCst);
})
.await;
if result.is_err() {
let stalled = match stage.load(Ordering::SeqCst) {
0 => "owner establishment and route propagation",
1 => "client establishment",
2 => "relay request opening and splice registration",
3 => "request forwarding",
4 => "owner terminal response submission",
5 => "terminal response forwarding and cleanup",
_ => "completed relay lifecycle",
};
panic!("relayed request lifecycle stalled during {stalled}");
}
}
#[tokio::test(flavor = "multi_thread")]
async fn open_stream_with_carries_headers_a_handler_reads() {
let node = Node::builder("gate-5")
.service(whoami_authorization.at_subject("whoami"))
.insecure_accept_declared_peer_identities()
.build()
.unwrap();
let (client_side, node_side) = pair();
node.serve_transport(node_side).await;
let client = Wire::open(client_side);
common::ready_client(&client).await;
let mut headers = serde_json::Map::new();
headers.insert("authorization".into(), json!("jwt-abc"));
let mut call = client
.client_session()
.start(
"/gate-5/whoami",
Kind::Request,
Envelope::encode_payload(&json!({})),
None,
headers,
)
.await
.unwrap();
let response = call.next().await.unwrap().unwrap();
assert_eq!(response.kind, Kind::Response);
assert_eq!(response.payload_json()["actor"], "jwt-abc");
}