mod common;
use common::TestCall;
use std::sync::Arc;
use std::time::Duration;
use common::connect_nodes;
use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
use serde_json::{json, Value};
use unb::{handler, HandlerError, Node, Origin, Reply, Request, Streaming};
use unb_client::pair;
use unb_core::{ErrorCode, Kind};
use unb_runtime::{ClientError, Wire};
#[derive(Deserialize, JsonSchema)]
struct Probe {
tag: String,
}
#[derive(Serialize, JsonSchema)]
struct Seen {
tag: String,
stamp: Option<u64>,
origin: String,
}
#[derive(Clone, Copy)]
struct Stamp(u64);
fn origin_name(origin: &Origin) -> String {
match origin {
Origin::Client { .. } => "client".into(),
Origin::Peer { .. } => "peer".into(),
Origin::Local => "local".into(),
Origin::Nested => "nested".into(),
}
}
#[handler]
async fn observe(request: Request<Probe>) -> Result<Reply<Seen>, HandlerError> {
Ok(Reply::new(Seen {
stamp: request.extensions().get::<Stamp>().map(|stamp| stamp.0),
origin: origin_name(request.origin()),
tag: request.into_payload().tag,
}))
}
#[derive(Deserialize, JsonSchema)]
struct Outer {
tag: String,
}
#[handler]
async fn relay_inward(request: Request<Outer>) -> Result<Reply<Value>, HandlerError> {
let tag = request.payload().tag.clone();
let value = request
.call("/ordered-1/api.observe", json!({ "tag": tag }))
.await?;
Ok(Reply::new(value))
}
#[handler]
async fn watch_guarded(
request: Request<Probe>,
) -> Result<Streaming<Seen, HandlerError>, HandlerError> {
let stamp = request.extensions().get::<Stamp>().map(|stamp| stamp.0);
let origin = origin_name(request.origin());
Ok(Streaming::new(futures_util::stream::iter(vec![Ok(Seen {
tag: request.into_payload().tag,
stamp,
origin,
})])))
}
#[tokio::test(flavor = "multi_thread")]
async fn layers_run_in_order_around_local_nested_and_client_dispatch() {
let trace = Arc::new(parking_lot::Mutex::new(Vec::<String>::new()));
let global_trace = trace.clone();
let scope_trace = trace.clone();
let node = Node::builder("ordered-1")
.layer_fn(move |request, next| {
let trace = global_trace.clone();
async move {
trace.lock().push(format!(
"global:{}",
unb_core::Envelope::subject_of(request.uri())
));
let response = next.run(request).await;
trace.lock().push("global:after".into());
response
}
})
.scope("api", |scope| {
scope
.layer_fn(move |request, next| {
let trace = scope_trace.clone();
async move {
trace.lock().push(format!(
"scope:{}",
unb_core::Envelope::subject_of(request.uri())
));
next.run(request).await
}
})
.service(observe)
.service(relay_inward)
})
.insecure_accept_declared_peer_identities()
.build()
.unwrap();
let value = node
.request("/ordered-1/api.observe", json!({ "tag": "direct" }))
.await
.unwrap();
assert_eq!(value["origin"], "local");
assert_eq!(
trace.lock().drain(..).collect::<Vec<_>>(),
vec!["global:api.observe", "scope:api.observe", "global:after"]
);
let nested = node
.request("/ordered-1/api.relay_inward", json!({ "tag": "deep" }))
.await
.unwrap();
assert_eq!(nested["origin"], "nested", "{nested}");
let order = trace.lock().drain(..).collect::<Vec<_>>();
assert_eq!(
order,
vec![
"global:api.relay_inward",
"scope:api.relay_inward",
"global:api.observe",
"scope:api.observe",
"global:after",
"global:after"
],
"the nested target runs its own full chain inside the outer chain"
);
}
#[tokio::test(flavor = "multi_thread")]
async fn a_rejecting_layer_surfaces_the_error_frame_to_clients() {
let node = Node::builder("guarded-1")
.layer_fn(|request, next| async move {
if request.headers().get("authorization").is_none() {
return Err(HandlerError::new(
ErrorCode::Unauthorized,
"an authorization header is required",
));
}
next.run(request).await
})
.service(observe)
.service(watch_guarded)
.insecure_accept_declared_peer_identities()
.build()
.unwrap();
let (client_side, server) = pair();
node.serve_transport(server).await;
let client = Wire::open(client_side);
for kind in [Kind::Request, Kind::Subscribe] {
let subject = if kind == Kind::Request {
"/guarded-1/observe"
} else {
"/guarded-1/watch_guarded"
};
let mut call = common::stream(&client, subject, kind, json!({ "tag": "x" })).await;
let error = tokio::time::timeout(Duration::from_secs(5), call.next())
.await
.unwrap()
.unwrap_err();
assert!(matches!(
error,
ClientError::Protocol {
code: ErrorCode::Unauthorized,
..
}
));
}
}
#[tokio::test(flavor = "multi_thread")]
async fn a_layer_error_status_response_reaches_the_client_as_an_error_frame() {
let node = Node::builder("status-guarded-1")
.layer_fn(
|request: unb::http::Request<bytes::Bytes>, _next| async move {
let _ = request;
Ok(unb::http::Response::builder()
.status(401)
.body(unb::ServiceBody::Unary(unb_core::Envelope::encode_payload(
&json!({ "message": "no entry" }),
)))
.expect("test response is well formed"))
},
)
.service(observe)
.insecure_accept_declared_peer_identities()
.build()
.unwrap();
let (client_side, server) = pair();
node.serve_transport(server).await;
let client = Wire::open(client_side);
let mut call = common::stream(
&client,
"/status-guarded-1/observe",
Kind::Request,
json!({ "tag": "x" }),
)
.await;
let error = tokio::time::timeout(Duration::from_secs(5), call.next())
.await
.unwrap()
.unwrap_err();
assert!(
matches!(error, ClientError::Protocol { code: ErrorCode::Unauthorized, ref message, .. } if message == "no entry")
);
}
#[tokio::test(flavor = "multi_thread")]
async fn typed_extensions_stay_isolated_between_requests() {
let counter = Arc::new(std::sync::atomic::AtomicU64::new(0));
let stamping = counter.clone();
let node = Node::builder("isolated-1")
.layer_fn(move |mut request, next| {
let counter = stamping.clone();
async move {
let stamp = counter.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
request.extensions_mut().insert(Stamp(stamp));
next.run(request).await
}
})
.service(observe)
.insecure_accept_declared_peer_identities()
.build()
.unwrap();
let first = node
.request("/isolated-1/observe", json!({ "tag": "a" }))
.await
.unwrap();
let second = node
.request("/isolated-1/observe", json!({ "tag": "b" }))
.await
.unwrap();
assert_eq!(first["stamp"], 0);
assert_eq!(
second["stamp"], 1,
"each request carries only its own typed facts"
);
}
#[tokio::test(flavor = "multi_thread")]
async fn a_relay_forwards_without_running_its_own_application_layers() {
let owner = Node::builder("owner-1")
.service(observe)
.insecure_accept_declared_peer_identities()
.build()
.unwrap();
let relay = Node::builder("relay-1")
.layer_fn(|_request, _next| async move {
Err(HandlerError::new(
ErrorCode::Unauthorized,
"the relay must never gate forwarded frames",
))
})
.insecure_accept_declared_peer_identities()
.build()
.unwrap();
connect_nodes(&relay, "owner-1", &owner).await;
let (client_side, server) = pair();
relay.serve_transport(server).await;
let client = Wire::open(client_side);
let mut call = common::stream(
&client,
"/owner-1/observe",
Kind::Request,
json!({ "tag": "through" }),
)
.await;
let response = tokio::time::timeout(Duration::from_secs(5), call.next())
.await
.unwrap()
.unwrap()
.unwrap();
assert_eq!(
response.kind,
Kind::Response,
"{:?}",
response.payload_json()
);
assert_eq!(response.payload_json()["tag"], "through");
assert_eq!(
response.payload_json()["origin"],
"peer",
"the owner sees its verified immediate peer, the relay, as the frame's origin"
);
}
#[tokio::test(flavor = "multi_thread")]
async fn streaming_setup_runs_the_chain_and_carries_typed_facts() {
let node = Node::builder("stream-guard-1")
.layer_fn(|mut request, next| async move {
request.extensions_mut().insert(Stamp(7));
next.run(request).await
})
.service(watch_guarded)
.insecure_accept_declared_peer_identities()
.build()
.unwrap();
let (client_side, server) = pair();
node.serve_transport(server).await;
let client = Wire::open(client_side);
let mut call = common::stream(
&client,
"/stream-guard-1/watch_guarded",
Kind::Subscribe,
json!({ "tag": "s" }),
)
.await;
let event = tokio::time::timeout(Duration::from_secs(5), call.next())
.await
.unwrap()
.unwrap()
.unwrap();
assert_eq!(event.payload_json()["stamp"], 7);
assert_eq!(event.payload_json()["origin"], "client");
}