unb-server 2.0.3

unb inbound server: Node, request/subscribe handlers, catalog, relay orchestration, accept
Documentation
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();
}