unb-server 1.0.0

unb inbound server: Node, request/subscribe handlers, catalog, relay orchestration, accept
Documentation
mod common;
use common::TestCall;

use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::Duration;

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};
use schemars::JsonSchema;
use serde::Deserialize;
use serde_json::{json, Value};

#[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("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("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(
            "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(),
        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, "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.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, "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, "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.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(
                &request.subject,
                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(
            "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");
}