unb-server 2.0.3

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

use std::sync::{Arc, OnceLock};

use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
use serde_json::{json, Value};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use unb::{handler, HandlerError, Reply, Request, SendExt};
use unb_server::{
    rustls, Endpoint, EndpointSet, HostConfig, HostError, Node, SelfSignedIdentity, TcpTransport,
    TransportKind, WebTransportConfig,
};

#[derive(Deserialize, Serialize, JsonSchema)]
#[serde(transparent)]
struct EchoPayload(Value);

#[handler]
async fn echo(request: Request<EchoPayload>) -> Result<Reply<Value>, HandlerError> {
    Ok(Reply::new(json!({ "echo": request.into_payload() })))
}

fn echo_node() -> Arc<Node> {
    Node::builder("tls-echo")
        .service(echo)
        .insecure_accept_declared_peer_identities()
        .build()
        .unwrap()
}

fn test_identity() -> (Vec<u8>, Vec<u8>, Vec<u8>) {
    let generated = SelfSignedIdentity::generate(["localhost"]).unwrap();
    let chain_pem = generated
        .identity()
        .certificate_chain()
        .as_slice()
        .iter()
        .map(unb_server::webtransport::wtransport::tls::Certificate::to_pem)
        .collect::<String>();
    let key_pem = generated.identity().private_key().to_secret_pem();
    let cert_der = generated
        .identity()
        .certificate_chain()
        .as_slice()
        .first()
        .unwrap()
        .der()
        .to_vec();
    (chain_pem.into_bytes(), key_pem.into_bytes(), cert_der)
}

fn client_config(cert_der: Vec<u8>) -> Arc<rustls::ClientConfig> {
    let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
    let mut roots = rustls::RootCertStore::empty();
    roots.add(cert_der.into()).unwrap();
    Arc::new(
        rustls::ClientConfig::builder()
            .with_root_certificates(roots)
            .with_no_client_auth(),
    )
}

fn dial_identity() -> &'static (Vec<u8>, Vec<u8>, Vec<u8>) {
    static IDENTITY: OnceLock<(Vec<u8>, Vec<u8>, Vec<u8>)> = OnceLock::new();
    IDENTITY.get_or_init(test_identity)
}

fn install_dial_trust() {
    let (_, _, cert_der) = dial_identity();
    let _ = unb_transport::ws::install_tls_client_config(client_config(cert_der.clone()));
}

async fn rustls_echo_host() -> (Arc<Node>, unb_server::Hosting) {
    install_dial_trust();
    let (chain, key, _) = dial_identity();
    let node = echo_node();
    let hosting = HostConfig::tcp(
        ([127, 0, 0, 1], 0),
        TcpTransport::rustls_pem(chain, key).unwrap(),
    )
    .start(&node)
    .await
    .unwrap();
    (node, hosting)
}

#[tokio::test(flavor = "multi_thread")]
async fn a_node_dials_a_wss_endpoint_through_the_standard_dial_policy() {
    let (node, hosting) = rustls_echo_host().await;
    let port = hosting.websocket_addr().unwrap().port();
    let caller = Node::builder("wss-caller")
        .insecure_accept_declared_peer_identities()
        .build()
        .unwrap();
    caller
        .connect(EndpointSet::from(Endpoint {
            kind: TransportKind::WebSocket,
            address: format!("wss://localhost:{port}"),
            cert_hash: None,
        }))
        .await
        .unwrap();
    let reply: Value = caller
        .request("/tls-echo/echo", json!({ "via": "wss" }))
        .await
        .unwrap();
    assert_eq!(reply["echo"]["via"], "wss");
    caller.shutdown();
    node.shutdown();
    hosting.shutdown().await.unwrap();
}

#[tokio::test(flavor = "multi_thread")]
async fn a_one_shot_https_destination_round_trips() {
    let (node, hosting) = rustls_echo_host().await;
    let port = hosting.websocket_addr().unwrap().port();
    let response = http::Request::builder()
        .uri("/tls-echo/echo")
        .body(json!({ "via": "https" }))
        .send(format!("https://localhost:{port}"))
        .await
        .unwrap();
    assert_eq!(response.status(), 200);
    let body: Value = serde_json::from_slice(response.body()).unwrap();
    assert_eq!(body["echo"]["via"], "https");
    node.shutdown();
    hosting.shutdown().await.unwrap();
}

#[tokio::test(flavor = "multi_thread")]
async fn a_rustls_host_terminates_https_ingress_and_wss_upgrades() {
    let (chain, key, cert_der) = test_identity();
    let node = echo_node();
    let hosting = HostConfig::tcp(
        ([127, 0, 0, 1], 0),
        TcpTransport::rustls_pem(&chain, &key).unwrap(),
    )
    .start(&node)
    .await
    .unwrap();
    let addr = hosting.websocket_addr().unwrap();
    let config = client_config(cert_der);

    let connector = tokio_rustls::TlsConnector::from(config.clone());
    let tcp = tokio::net::TcpStream::connect(addr).await.unwrap();
    let mut tls = connector
        .connect("localhost".try_into().unwrap(), tcp)
        .await
        .unwrap();
    let body = r#"{"hello":"tls"}"#;
    let request = format!(
        "POST /tls-echo/echo HTTP/1.1\r\nhost: localhost\r\nconnection: close\r\ncontent-length: {}\r\n\r\n{body}",
        body.len()
    );
    tls.write_all(request.as_bytes()).await.unwrap();
    let mut response = Vec::new();
    let _ = tls.read_to_end(&mut response).await;
    let response = String::from_utf8_lossy(&response);
    assert!(response.starts_with("HTTP/1.1 200"), "{response}");
    assert!(response.contains(r#""hello":"tls""#), "{response}");

    let ws_config = tokio_tungstenite::tungstenite::protocol::WebSocketConfig::default();
    let (_stream, upgrade) = tokio_tungstenite::connect_async_tls_with_config(
        format!("wss://localhost:{}", addr.port()),
        Some(ws_config),
        false,
        Some(tokio_tungstenite::Connector::Rustls(config)),
    )
    .await
    .unwrap();
    assert_eq!(
        upgrade.status(),
        tokio_tungstenite::tungstenite::http::StatusCode::SWITCHING_PROTOCOLS
    );
    hosting.shutdown().await.unwrap();
}

#[tokio::test(flavor = "multi_thread")]
async fn an_alpn_config_without_http1_is_rejected_before_binding() {
    let (chain, key, _) = test_identity();
    let identity = unb_server::ServerIdentity::from_pem(&chain, &key).unwrap();
    let mut config = (*identity.tcp_rustls().unwrap()).clone();
    config.alpn_protocols = vec![b"h2".to_vec()];
    let error = HostConfig::tcp(([127, 0, 0, 1], 0), TcpTransport::rustls(Arc::new(config)))
        .start(&echo_node())
        .await
        .err()
        .unwrap();
    assert!(matches!(error, HostError::AlpnMissingHttp1));
}

#[tokio::test(flavor = "multi_thread")]
async fn an_empty_alpn_config_is_normalized_to_http1() {
    let (chain, key, cert_der) = test_identity();
    let identity = unb_server::ServerIdentity::from_pem(&chain, &key).unwrap();
    let mut config = (*identity.tcp_rustls().unwrap()).clone();
    config.alpn_protocols = Vec::new();
    let node = echo_node();
    let hosting = HostConfig::tcp(([127, 0, 0, 1], 0), TcpTransport::rustls(Arc::new(config)))
        .start(&node)
        .await
        .unwrap();
    let addr = hosting.websocket_addr().unwrap();
    let mut dialer = (*client_config(cert_der)).clone();
    dialer.alpn_protocols = vec![b"http/1.1".to_vec()];
    let connector = tokio_rustls::TlsConnector::from(Arc::new(dialer));
    let tcp = tokio::net::TcpStream::connect(addr).await.unwrap();
    let tls = connector
        .connect("localhost".try_into().unwrap(), tcp)
        .await
        .unwrap();
    let (_, session) = tls.get_ref();
    assert_eq!(session.alpn_protocol(), Some(b"http/1.1".as_ref()));
    hosting.shutdown().await.unwrap();
}

#[tokio::test(flavor = "multi_thread")]
async fn an_alpn_config_advertising_h2_is_normalized_to_http1_only() {
    let (chain, key, cert_der) = test_identity();
    let identity = unb_server::ServerIdentity::from_pem(&chain, &key).unwrap();
    let mut config = (*identity.tcp_rustls().unwrap()).clone();
    config.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec()];
    let node = echo_node();
    let hosting = HostConfig::tcp(([127, 0, 0, 1], 0), TcpTransport::rustls(Arc::new(config)))
        .start(&node)
        .await
        .unwrap();
    let addr = hosting.websocket_addr().unwrap();
    let mut dialer = (*client_config(cert_der)).clone();
    dialer.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec()];
    let connector = tokio_rustls::TlsConnector::from(Arc::new(dialer));
    let tcp = tokio::net::TcpStream::connect(addr).await.unwrap();
    let tls = connector
        .connect("localhost".try_into().unwrap(), tcp)
        .await
        .unwrap();
    let (_, session) = tls.get_ref();
    assert_eq!(session.alpn_protocol(), Some(b"http/1.1".as_ref()));
    hosting.shutdown().await.unwrap();
}

#[tokio::test(flavor = "multi_thread")]
async fn a_user_identity_exposes_no_development_cert_hash() {
    let (chain, key, _) = test_identity();
    let node = echo_node();
    let hosting = HostConfig::new(([127, 0, 0, 1], 0))
        .with_webtransport(WebTransportConfig::pem(&chain, &key).unwrap())
        .start(&node)
        .await
        .unwrap();
    assert!(hosting.development_cert_hash().is_none());
    hosting.shutdown().await.unwrap();
}

#[tokio::test(flavor = "multi_thread")]
async fn a_doubly_configured_webtransport_is_rejected() {
    let (chain, key, _) = test_identity();
    let identity = unb_server::ServerIdentity::from_pem(&chain, &key).unwrap();
    let endpoint_config = unb_server::webtransport::wtransport::ServerConfig::builder()
        .with_bind_address("127.0.0.1:0".parse().unwrap())
        .with_identity(identity.webtransport().unwrap())
        .build();
    let endpoint = unb_server::webtransport::wtransport::Endpoint::server(endpoint_config).unwrap();
    let error = HostConfig::new(([127, 0, 0, 1], 0))
        .with_webtransport(WebTransportConfig::self_signed_for_development(["localhost"]).unwrap())
        .webtransport_endpoint(endpoint)
        .start(&echo_node())
        .await
        .err()
        .unwrap();
    assert!(matches!(error, HostError::WebTransportConflict));
}

#[tokio::test(flavor = "multi_thread")]
async fn mismatched_listener_and_endpoint_ports_are_rejected_atomically() {
    let (chain, key, _) = test_identity();
    let identity = unb_server::ServerIdentity::from_pem(&chain, &key).unwrap();
    let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
    let endpoint_config = unb_server::webtransport::wtransport::ServerConfig::builder()
        .with_bind_address("127.0.0.1:0".parse().unwrap())
        .with_identity(identity.webtransport().unwrap())
        .build();
    let endpoint = unb_server::webtransport::wtransport::Endpoint::server(endpoint_config).unwrap();
    let error = HostConfig::new(([127, 0, 0, 1], 0))
        .tcp_listener(listener, TcpTransport::plain())
        .webtransport_endpoint(endpoint)
        .start(&echo_node())
        .await
        .err()
        .unwrap();
    assert!(matches!(error, HostError::ListenerAddressMismatch { .. }));
}

#[tokio::test(flavor = "multi_thread")]
async fn an_external_endpoint_serves_webtransport_with_user_owned_quic() {
    let generated = SelfSignedIdentity::generate(["localhost"]).unwrap();
    let hash = generated.cert_hash();
    let endpoint_config = unb_server::webtransport::wtransport::ServerConfig::builder()
        .with_bind_address("127.0.0.1:0".parse().unwrap())
        .with_identity(generated.identity().clone_identity())
        .build();
    let endpoint = unb_server::webtransport::wtransport::Endpoint::server(endpoint_config).unwrap();
    let node = echo_node();
    let hosting = HostConfig::new(([127, 0, 0, 1], 0))
        .webtransport_endpoint(endpoint)
        .start(&node)
        .await
        .unwrap();
    assert!(hosting.development_cert_hash().is_none());
    let url = format!("https://{}", hosting.webtransport_addr().unwrap());
    let (pipe, initiator, _bodies) = unb_transport::webtransport::dial(&url, Some(hash))
        .await
        .unwrap();
    assert!(initiator);
    let wire = unb_runtime::Wire::open(unb_runtime::Pipe::Piped { pipe, initiator });
    common::ready_client(&wire).await;
    wire.shutdown();
    hosting.shutdown().await.unwrap();
}