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();
}