diff --git a/Cargo.lock b/Cargo.lock index a65f33ae..597a01ce 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -8334,6 +8334,7 @@ dependencies = [ "serde_json", "tempfile", "tokio", + "tokio-rustls", "tokio-tungstenite 0.30.0", "url", "web-transport-quinn", diff --git a/crates/cli/src/cli.rs b/crates/cli/src/cli.rs index 26548582..5b70cd9a 100644 --- a/crates/cli/src/cli.rs +++ b/crates/cli/src/cli.rs @@ -539,6 +539,15 @@ pub enum Command { required = true )] allow_client: Vec, + + /// What carries the relay session: auto (WebTransport, else WebSocket + /// where UDP is blocked), webtransport or websocket + #[arg( + long, + env = "YAS_UPLINK_TRANSPORT", + value_parser = ["auto", "webtransport", "websocket"] + )] + transport: Option, }, /// Print an X25519 key pair as JSON (private_key and public_key, base64url) diff --git a/crates/cli/src/main.rs b/crates/cli/src/main.rs index 80547857..492994f6 100644 --- a/crates/cli/src/main.rs +++ b/crates/cli/src/main.rs @@ -947,8 +947,9 @@ async fn async_main() { url, identity, allow_client, + transport, } => { - if let Err(e) = uplink::cmd_uplink(url, identity, allow_client).await { + if let Err(e) = uplink::cmd_uplink(url, identity, allow_client, transport).await { eprintln!("yas: {e}"); std::process::exit(1); } diff --git a/crates/cli/src/uplink.rs b/crates/cli/src/uplink.rs index c1d403c7..e4d8f517 100644 --- a/crates/cli/src/uplink.rs +++ b/crates/cli/src/uplink.rs @@ -1,25 +1,12 @@ //! `yas uplink` — expose the local yas server through a relay. //! //! Authenticates to an HTTPS control endpoint with `YAS_UPLINK_TOKEN`, -//! receives a pool of WebTransport relays, -//! establishes a session with one, authenticates each consumer using pinned -//! X25519 keys over Noise IK, then bridges decrypted streams to local YAS. +//! receives a pool of relays (WebTransport, WebSocket), holds a session with +//! one, authenticates each consumer using pinned X25519 keys over Noise IK, +//! then bridges decrypted streams to local YAS. The producer itself is +//! [`yas_proxy::uplink_producer`]; this is its command line. -use std::sync::{Arc, Mutex}; -use std::time::Duration; - -use tokio::io::AsyncWriteExt; -use web_transport_quinn as wt; - -/// Stream application error codes. -const CODE_SHUTDOWN: u32 = 2; - -const INITIAL_BACKOFF: Duration = Duration::from_secs(1); -const MAX_BACKOFF: Duration = Duration::from_secs(60); -const MAX_DATAGRAM_ROUTES: usize = 4_096; -const DATAGRAM_QUEUE: usize = 64; - -type DatagramRoutes = yas_composite_transport::RoutedDatagramRoutes; +use yas_proxy::uplink_producer::{Local, Producer, Transport}; /// Derive only the public identity; private input never appears in diagnostics. pub fn public_key_from_env() -> Result { @@ -62,21 +49,13 @@ pub fn connection_url(control: &str, client_token: &str) -> Result>, -} - +/// Publish the local server until a fatal error (the control endpoint +/// refusing the token) or Ctrl-C, which closes the relay session. pub async fn cmd_uplink( url: String, identity: String, allow_client: Vec, + transport: Option, ) -> Result<(), String> { let identity = yas_uplink::Identity::from_base64(&identity)?; let allowed = allow_client @@ -84,871 +63,27 @@ pub async fn cmd_uplink( .map(|key| key.parse()) .collect::, _>>()?; let crypto = identity.server_config(allowed)?; - let control = url::Url::parse(&url).map_err(|_| "invalid uplink control URL")?; - if control.scheme() != "https" - || !control.username().is_empty() - || control.password().is_some() - || control.fragment().is_some() - { - return Err("uplink control URL must be HTTPS without userinfo or a fragment".into()); - } let token = std::env::var("YAS_UPLINK_TOKEN").unwrap_or_default(); - if token.is_empty() { - return Err("YAS_UPLINK_TOKEN is not set".into()); - } - - // The active session, shared with the ctrl-c arm so shutdown can close - // it with CODE_SHUTDOWN instead of letting it idle out on the relay. - let current: Arc>> = Arc::new(Mutex::new(None)); - let current2 = current.clone(); - - tokio::select! { - _ = tokio::signal::ctrl_c() => { - if let Some(session) = current2.lock().unwrap().take() { - session.close(CODE_SHUTDOWN, b"uplink shutting down"); - } - Ok(()) - } - result = run_loop(&url, &token, current, crypto) => result, - } -} - -async fn run_loop( - url: &str, - token: &str, - current: Arc>>, - crypto: Arc, -) -> Result<(), String> { - let http = yas_proxy::uplink_http_client()?; - let mut backoff = INITIAL_BACKOFF; - - loop { - let mut pool = match fetch_pool(&http, url, token).await? { - FetchOutcome::Pool(pool) => pool, - FetchOutcome::Retry { after, reason } => { - let delay = after.unwrap_or(backoff); - eprintln!("[uplink] {reason}; retrying in {}s", delay.as_secs()); - tokio::time::sleep(jittered(delay)).await; - backoff = (backoff * 2).min(MAX_BACKOFF); - continue; - } - }; - - { - use rand::seq::SliceRandom; - pool.shuffle(&mut rand::rng()); - } - - // Try relays in (shuffled) order; a session that was actually - // established ends with a fresh control-plane query, per - // docs/uplink.md. - let mut established = false; - for relay in &pool { - match run_session(relay, ¤t, crypto.clone()).await { - SessionEnd::Ended(reason) => { - eprintln!("[uplink] session ended: {reason}"); - established = true; - break; - } - SessionEnd::NeverConnected(e) => { - eprintln!("[uplink] relay {}: {e}", relay.label); - } - } - } - - if established { - backoff = INITIAL_BACKOFF; - } else { - eprintln!( - "[uplink] relay pool exhausted; re-querying in {}s", - backoff.as_secs() - ); - tokio::time::sleep(jittered(backoff)).await; - backoff = (backoff * 2).min(MAX_BACKOFF); - } - } -} - -// --------------------------------------------------------------------------- -// Control plane -// --------------------------------------------------------------------------- - -enum FetchOutcome { - Pool(Vec), - Retry { - after: Option, - reason: String, - }, -} - -/// Query the control endpoint. `Err` is fatal (bad token); every other -/// failure is a retryable `FetchOutcome::Retry`. -async fn fetch_pool( - http: &reqwest::Client, - url: &str, - token: &str, -) -> Result { - let retry = |after, reason| Ok(FetchOutcome::Retry { after, reason }); - - let resp = match http - .get(url) - .header("authorization", format!("Bearer {token}")) - .header("accept", "application/json") - .send() - .await - { - Ok(resp) => resp, - Err(e) => return retry(None, format!("control endpoint unreachable: {e}")), - }; - - let status = resp.status(); - if status.as_u16() == 401 || status.as_u16() == 403 { - return Err(format!( - "control endpoint rejected YAS_UPLINK_TOKEN ({status})" - )); - } - if !status.is_success() { - let after = resp - .headers() - .get("retry-after") - .and_then(|v| v.to_str().ok()) - .and_then(|s| s.trim().parse::().ok()) - .map(Duration::from_secs); - return retry(after, format!("control endpoint returned {status}")); - } - - let body = match resp.text().await { - Ok(body) => body, - Err(e) => return retry(None, format!("error reading relay pool: {e}")), - }; - match parse_pool(&body) { - Ok(pool) => Ok(FetchOutcome::Pool(pool)), - Err(e) => retry(None, format!("bad relay pool: {e}")), - } -} - -fn parse_pool(body: &str) -> Result, String> { - let v: serde_json::Value = - serde_json::from_str(body).map_err(|e| format!("invalid JSON: {e}"))?; - let relays = v - .get("relays") - .and_then(|r| r.as_array()) - .ok_or("missing \"relays\" array")?; - - let mut pool = Vec::new(); - for relay in relays { - let s = relay.as_str().ok_or("relay entries must be URL strings")?; - pool.push(parse_relay(s)?); - } - if pool.is_empty() { - return Err("empty \"relays\" array".into()); - } - Ok(pool) -} - -fn parse_relay(s: &str) -> Result { - let mut url = url::Url::parse(s).map_err(|e| format!("bad relay URL: {e}"))?; - if url.scheme() != "https" { - return Err(format!( - "relay URL scheme must be https, got {}", - url.scheme() - )); - } - let host = url.host_str().ok_or("relay URL has no host")?.to_string(); - let label = format!("{}:{}", host, url.port().unwrap_or(443)); - - // A malformed pin must be an error, never a silent fall-back to system - // roots — that would defeat the pinning. - let cert_hash = match url.fragment() { - None | Some("") => None, - Some(frag) => { - let hash = frag - .strip_prefix("sha256=") - .ok_or("relay URL fragment must be sha256=")?; - let bytes = - base64url_decode(hash).ok_or("relay certificate pin is not valid base64url")?; - if bytes.len() != 32 { - return Err("relay certificate pin must be a SHA-256 (32 bytes)".into()); - } - Some(bytes) - } - }; - // Fragments are client-side only; strip before connecting. - url.set_fragment(None); - - Ok(Relay { - url, - label, - cert_hash, - }) -} - -fn base64url_decode(s: &str) -> Option> { - let s = s.trim_end_matches('='); - let mut out = Vec::with_capacity(s.len() * 3 / 4); - let mut acc: u32 = 0; - let mut bits = 0u32; - for b in s.bytes() { - let v = match b { - b'A'..=b'Z' => b - b'A', - b'a'..=b'z' => b - b'a' + 26, - b'0'..=b'9' => b - b'0' + 52, - b'-' => 62, - b'_' => 63, - _ => return None, - }; - acc = (acc << 6) | u32::from(v); - bits += 6; - if bits >= 8 { - bits -= 8; - out.push((acc >> bits) as u8); - } - } - // Leftover bits must be padding zeros of a valid encoding. - if bits > 0 && (acc & ((1 << bits) - 1)) != 0 { - return None; - } - Some(out) -} - -// --------------------------------------------------------------------------- -// Relay session -// --------------------------------------------------------------------------- - -enum SessionEnd { - /// Handshake or CONNECT failed — try the next relay in the pool. - NeverConnected(String), - /// The session was established and later died — re-query the control - /// plane for a fresh pool. - Ended(String), -} - -async fn run_session( - relay: &Relay, - current: &Arc>>, - crypto: Arc, -) -> SessionEnd { - let client = match build_client(relay.cert_hash.as_deref()) { - Ok(client) => client, - Err(e) => return SessionEnd::NeverConnected(e), - }; - // Careful: the URL is the credential — log `relay.label` only. - let session = match client.connect(relay.url.clone()).await { - Ok(session) => session, - Err(e) => return SessionEnd::NeverConnected(format!("connect failed: {e}")), - }; - eprintln!("[uplink] connected to relay {}", relay.label); - *current.lock().unwrap() = Some(session.clone()); - - let routes = DatagramRoutes::new(MAX_DATAGRAM_ROUTES); - tokio::spawn(distribute_datagrams(session.clone(), routes.clone())); - let pending = Arc::new(tokio::sync::Semaphore::new(64)); - let local_socket = crate::transport::default_local_socket(); - - let reason = loop { - match session.accept_bi().await { - Ok((send, recv)) => { - let Ok(permit) = pending.clone().try_acquire_owned() else { - continue; - }; - tokio::spawn(bridge( - session.clone(), - routes.clone(), - send, - recv, - crypto.clone(), - permit, - local_socket.clone(), - )); - } - Err(e) => break format!("{e}"), - } - }; - current.lock().unwrap().take(); - SessionEnd::Ended(reason) -} - -/// Bridge one relay-initiated stream to one local yas server connection. -/// Authentication precedes both ingress classification and local IPC. The -/// relay sees only Noise records; plaintext selectors cannot bypass admission. -async fn bridge( - session: wt::Session, - routes: DatagramRoutes, - send: wt::SendStream, - recv: wt::RecvStream, - crypto: Arc, - permit: tokio::sync::OwnedSemaphorePermit, - local_socket: String, -) { - let relay = tokio::io::join(recv, send); - let relay = match yas_uplink::accept(relay, crypto).await { - Ok(relay) => relay, - Err(_) => return, - }; - // Bind each sideband's keys to its authenticated main stream. The random - // route token remains authenticated as AEAD AAD on every datagram. - let material = relay.datagram_key_material(); - let ingress = tokio::time::timeout( - Duration::from_secs(5), - yas_composite_transport::classify(relay), - ) - .await; - drop(permit); - match ingress { - Ok(Ok(yas_composite_transport::Ingress::Direct(relay))) => { - bridge_direct(relay, &local_socket).await - } - Ok(Ok(yas_composite_transport::Ingress::Composite { offer, stream })) - if offer.role == yas_composite_transport::Role::Main => - { - let (sender, receiver) = yas_uplink::datagram_pair(material, offer.token, false); - bridge_composite( - session, - routes, - offer, - stream, - sender, - receiver, - &local_socket, - ) - .await; - } - _ => {} - } -} - -async fn bridge_direct(relay: S, path: &str) -where - S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin, -{ - let transport = match crate::transport::connect_ipc(path).await { - Ok(transport) => transport, - Err(e) => { - eprintln!("[uplink] local yas server unavailable at {path}: {e}"); - return; - } - }; - let (mut sock_read, mut sock_write) = transport.split(); - let (mut recv, mut send) = tokio::io::split(relay); - - let down = async move { - tokio::io::copy(&mut recv, &mut sock_write).await?; - sock_write.shutdown().await - }; - let up = async move { - tokio::io::copy(&mut sock_read, &mut send).await?; - send.shutdown().await - }; - let _ = tokio::try_join!(down, up); -} - -async fn distribute_datagrams(session: wt::Session, routes: DatagramRoutes) { - while let Ok(bytes) = session.read_datagram().await { - let _ = routes.route(&bytes); - } - routes.clear(); -} - -async fn bridge_composite( - session: wt::Session, - routes: DatagramRoutes, - offer: yas_composite_transport::Offer, - relay: S, - mut sender: yas_uplink::DatagramSender, - mut receiver: yas_uplink::DatagramReceiver, - path: &str, -) where - S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin, -{ - let physical_maximum = session - .max_datagram_size() - .saturating_sub(yas_composite_transport::ROUTED_DATAGRAM_HEADER) - .saturating_sub(yas_uplink::DATAGRAM_OVERHEAD) - .min(yas_composite_transport::HARD_MAX_DATAGRAM as usize - yas_uplink::DATAGRAM_OVERHEAD); - if offer.max_datagram as usize > physical_maximum { - eprintln!( - "[uplink] rejected composite maximum {} above path maximum {physical_maximum}", - offer.max_datagram - ); - return; - } - - let main = match crate::transport::connect_ipc(path).await { - Ok(transport) => transport, - Err(error) => { - eprintln!("[uplink] local yas server unavailable at {path}: {error}"); - return; - } - }; - let side = match crate::transport::connect_ipc(path).await { - Ok(transport) => transport, - Err(error) => { - eprintln!("[uplink] local YAS datagram sideband unavailable at {path}: {error}"); - return; - } - }; - let (mut main_read, mut main_write) = main.split(); - let (mut side_read, mut side_write) = side.split(); - let main_offer = yas_composite_transport::Offer::new( - yas_composite_transport::Role::Main, - offer.token, - offer.max_datagram, - ) - .expect("the classified offer is valid"); - let side_offer = yas_composite_transport::Offer::new( - yas_composite_transport::Role::Datagram, - offer.token, - offer.max_datagram, - ) - .expect("the classified offer is valid"); - if yas_composite_transport::write_offer(&mut main_write, main_offer) - .await - .is_err() - || yas_composite_transport::write_offer(&mut side_write, side_offer) - .await - .is_err() - { - return; - } - - let encrypted_maximum = offer.max_datagram + yas_uplink::DATAGRAM_OVERHEAD as u32; - let Ok(mut route_rx) = routes.register(offer.token, encrypted_maximum, DATAGRAM_QUEUE) else { - return; - }; - - let (mut relay_read, mut relay_write) = tokio::io::split(relay); - let down = async { - tokio::io::copy(&mut relay_read, &mut main_write).await?; - main_write.shutdown().await - }; - let up = async { - tokio::io::copy(&mut main_read, &mut relay_write).await?; - relay_write.shutdown().await - }; - let side_routes = routes.clone(); - let side_token = offer.token; - let side_task = tokio::spawn(async move { - let side_out = async { - loop { - let Ok(frame) = - yas_composite_transport::read_datagram(&mut side_read, offer.max_datagram) - .await - else { - break; - }; - let Some(encrypted) = sender.seal(&frame) else { - break; - }; - let Ok(routed) = yas_composite_transport::encode_routed_datagram( - offer.token, - &encrypted, - encrypted_maximum, - ) else { - break; - }; - // Congestion is ordinary datagram loss; never await reliable - // transport capacity or fall back to the main stream here. - let _ = session.send_datagram(routed.into()); - } - }; - let side_in = async { - while let Some(frame) = route_rx.recv().await { - let Some(frame) = receiver.open(&frame) else { - continue; - }; - if yas_composite_transport::write_datagram( - &mut side_write, - &frame, - offer.max_datagram, - ) + let transport: Transport = match transport { + Some(transport) => transport.parse()?, + None => Transport::default(), + }; + let socket = crate::transport::default_local_socket(); + let local = Local::custom(socket.clone(), move || { + let socket = socket.clone(); + async move { + let transport = crate::transport::connect_ipc(&socket) .await - .is_err() - { - break; - } - } - }; - tokio::select! { - _ = side_out => {} - _ = side_in => {} + .map_err(std::io::Error::other)?; + let (reader, writer) = transport.split(); + Ok(tokio::io::join(reader, writer)) } - side_routes.remove(side_token); }); - - // The authoritative reliable splice owns the connection lifetime. An - // optional sideband ending only removes its route and cannot drop `main`. - let _ = tokio::try_join!(down, up); - routes.remove(offer.token); - side_task.abort(); -} - -/// Build a WebTransport client with the liveness settings from -/// docs/uplink.md (10s keepalive, 30s idle timeout) and either -/// system-root or pinned TLS verification. `wt::ClientBuilder` doesn't expose the quinn transport -/// config, so this mirrors its setup by hand. -fn build_client(cert_hash: Option<&[u8]>) -> Result { - let provider = Arc::new(rustls::crypto::ring::default_provider()); - let builder = rustls::ClientConfig::builder_with_provider(provider.clone()) - .with_protocol_versions(&[&rustls::version::TLS13]) - .map_err(|e| format!("TLS config: {e}"))?; - - let mut crypto = match cert_hash { - Some(hash) => builder - .dangerous() - .with_custom_certificate_verifier(Arc::new(PinnedCert { - hash: hash.to_vec(), - provider, - })) - .with_no_client_auth(), - None => { - let mut roots = rustls::RootCertStore::empty(); - for cert in rustls_native_certs::load_native_certs().certs { - let _ = roots.add(cert); - } - builder.with_root_certificates(roots).with_no_client_auth() - } - }; - crypto.alpn_protocols = vec![wt::ALPN.as_bytes().to_vec()]; - - let quic_crypto = wt::quinn::crypto::rustls::QuicClientConfig::try_from(crypto) - .map_err(|e| format!("QUIC TLS config: {e}"))?; - let mut config = wt::quinn::ClientConfig::new(Arc::new(quic_crypto)); - - let mut transport = wt::quinn::TransportConfig::default(); - transport.keep_alive_interval(Some(Duration::from_secs(10))); - transport.max_idle_timeout(Some( - wt::quinn::IdleTimeout::try_from(Duration::from_secs(30)) - .expect("30s fits in an idle timeout"), - )); - config.transport_config(Arc::new(transport)); - - let endpoint = wt::quinn::Endpoint::client("[::]:0".parse().unwrap()) - .or_else(|_| wt::quinn::Endpoint::client("0.0.0.0:0".parse().unwrap())) - .map_err(|e| format!("UDP socket: {e}"))?; - Ok(wt::Client::new(endpoint, config)) -} - -/// Pins the relay's end-entity certificate to a SHA-256 hash from the pool -/// (the edge's `serverCertificateHashes` flow, client side). Chain and -/// expiry are deliberately not checked — the hash is the trust anchor. -#[derive(Debug)] -struct PinnedCert { - hash: Vec, - provider: Arc, -} - -impl rustls::client::danger::ServerCertVerifier for PinnedCert { - fn verify_server_cert( - &self, - end_entity: &rustls::pki_types::CertificateDer<'_>, - _intermediates: &[rustls::pki_types::CertificateDer<'_>], - _server_name: &rustls::pki_types::ServerName<'_>, - _ocsp_response: &[u8], - _now: rustls::pki_types::UnixTime, - ) -> Result { - let digest = ring::digest::digest(&ring::digest::SHA256, end_entity.as_ref()); - if digest.as_ref() == self.hash.as_slice() { - Ok(rustls::client::danger::ServerCertVerified::assertion()) - } else { - Err(rustls::Error::InvalidCertificate( - rustls::CertificateError::ApplicationVerificationFailure, - )) - } - } - - fn verify_tls12_signature( - &self, - _message: &[u8], - _cert: &rustls::pki_types::CertificateDer<'_>, - _dss: &rustls::DigitallySignedStruct, - ) -> Result { - Err(rustls::Error::PeerIncompatible( - rustls::PeerIncompatible::Tls12NotOffered, - )) - } - - fn verify_tls13_signature( - &self, - message: &[u8], - cert: &rustls::pki_types::CertificateDer<'_>, - dss: &rustls::DigitallySignedStruct, - ) -> Result { - rustls::crypto::verify_tls13_signature( - message, - cert, - dss, - &self.provider.signature_verification_algorithms, - ) - } - - fn supported_verify_schemes(&self) -> Vec { - self.provider - .signature_verification_algorithms - .supported_schemes() - } -} - -fn jittered(base: Duration) -> Duration { - use rand::RngExt as _; - let ms = base.as_millis().max(1) as u64; - // 0.75x–1.25x - Duration::from_millis(rand::rng().random_range(ms * 3 / 4..=ms * 5 / 4)) -} - -#[cfg(test)] -mod tests { - use super::*; - - #[cfg(unix)] - #[tokio::test] - async fn producer_authenticates_before_ipc_and_encrypts_datagrams() { - use tokio::io::AsyncReadExt; - use yas_composite_transport::{Offer, Role}; - let _ = rustls::crypto::ring::default_provider().install_default(); - tokio::time::timeout(Duration::from_secs(15), async { - let dir = tempfile::tempdir().unwrap(); - let (server_key, server_public) = yas_uplink::Identity::generate().unwrap(); - let (client_key, client_public) = yas_uplink::Identity::generate().unwrap(); - let (attacker_key, _) = yas_uplink::Identity::generate().unwrap(); - let crypto = yas_uplink::Identity::from_base64(&server_key) - .unwrap() - .server_config(vec![client_public]) - .unwrap(); - let client_crypto = yas_uplink::Identity::from_base64(&client_key) - .unwrap() - .client_config(server_public) - .unwrap(); - let attacker_crypto = yas_uplink::Identity::from_base64(&attacker_key) - .unwrap() - .client_config(server_public) - .unwrap(); - let socket = dir.path().join("local.sock"); - let listener = tokio::net::UnixListener::bind(&socket).unwrap(); - - let cert = rcgen::generate_simple_self_signed(vec!["localhost".into()]).unwrap(); - let hash = ring::digest::digest(&ring::digest::SHA256, cert.cert.der()); - let mut worker = wt::ServerBuilder::new() - .with_addr("127.0.0.1:0".parse().unwrap()) - .with_certificate( - vec![cert.cert.der().clone()], - rustls::pki_types::PrivatePkcs8KeyDer::from(cert.signing_key.serialize_der()) - .into(), - ) - .unwrap(); - let outer = build_client(Some(hash.as_ref())).unwrap(); - let url: url::Url = - format!("https://127.0.0.1:{}/", worker.local_addr().unwrap().port()) - .parse() - .unwrap(); - let (local, remote) = tokio::join!(outer.connect(url), async { - worker.accept().await.unwrap().ok().await.unwrap() - }); - let local = local.unwrap(); - let routes = DatagramRoutes::new(16); - let datagrams = tokio::spawn(distribute_datagrams(local.clone(), routes.clone())); - - for case in 0..4 { - let local = local.clone(); - let routes = routes.clone(); - let crypto = crypto.clone(); - let socket = socket.to_str().unwrap().to_string(); - let producer = tokio::spawn(async move { - let (send, recv) = local.accept_bi().await.unwrap(); - let permit = Arc::new(tokio::sync::Semaphore::new(1)) - .acquire_owned() - .await - .unwrap(); - bridge(local, routes, send, recv, crypto, permit, socket).await; - }); - let (send, recv) = remote.open_bi().await.unwrap(); - let mut stream = tokio::io::join(recv, send); - if case == 0 { - // A relay can synthesize protocol bytes, but those bytes - // must never open a local socket before Noise authentication. - stream - .write_all(yas_wire::PREFACE.as_slice()) - .await - .unwrap(); - stream.shutdown().await.unwrap(); - drop(stream); - producer.await.unwrap(); - assert!( - tokio::time::timeout(Duration::from_millis(25), listener.accept()) - .await - .is_err() - ); - continue; - } - if case == 1 { - assert!( - yas_uplink::connect(stream, attacker_crypto.clone()) - .await - .is_err() - ); - producer.await.unwrap(); - assert!( - tokio::time::timeout(Duration::from_millis(25), listener.accept()) - .await - .is_err() - ); - continue; - } - let mut stream = yas_uplink::connect(stream, client_crypto.clone()) - .await - .unwrap(); - // Authentication alone also does not open IPC: the encrypted - // YAS preface or composite selector must follow first. - assert!( - tokio::time::timeout(Duration::from_millis(25), listener.accept()) - .await - .is_err() - ); - if case == 2 { - stream - .write_all(yas_wire::PREFACE.as_slice()) - .await - .unwrap(); - stream.write_all(b"command").await.unwrap(); - stream.flush().await.unwrap(); - let (mut ipc, _) = listener.accept().await.unwrap(); - let mut received = vec![0; yas_wire::PREFACE.len() + 7]; - ipc.read_exact(&mut received).await.unwrap(); - assert_eq!( - received, - [yas_wire::PREFACE.as_slice(), b"command"].concat() - ); - ipc.write_all(b"reply").await.unwrap(); - ipc.shutdown().await.unwrap(); - let mut reply = Vec::new(); - stream.read_to_end(&mut reply).await.unwrap(); - assert_eq!(reply, b"reply"); - stream.shutdown().await.unwrap(); - producer.await.unwrap(); - continue; - } - - let token = [0x35; 16]; - let maximum = 512; - let material = stream.datagram_key_material(); - let (mut sender, mut receiver) = yas_uplink::datagram_pair(material, token, true); - yas_composite_transport::write_offer( - &mut stream, - Offer::new(Role::Main, token, maximum).unwrap(), - ) - .await - .unwrap(); - stream.flush().await.unwrap(); - let (mut main, _) = listener.accept().await.unwrap(); - let (mut side, _) = listener.accept().await.unwrap(); - for (socket, expected) in [(&mut main, Role::Main), (&mut side, Role::Datagram)] { - match yas_composite_transport::classify(socket).await.unwrap() { - yas_composite_transport::Ingress::Composite { offer, .. } => { - assert_eq!(offer.role, expected); - assert_eq!(offer.token, token); - assert_eq!(offer.max_datagram, maximum); - } - _ => panic!("missing local composite selector"), - } - } - // A reliable reply also synchronizes route registration. - main.write_all(b"ready").await.unwrap(); - let mut ready = [0; 5]; - stream.read_exact(&mut ready).await.unwrap(); - assert_eq!(&ready, b"ready"); - - let encrypted = sender.seal(b"private datagram").unwrap(); - let wire_maximum = maximum + yas_uplink::DATAGRAM_OVERHEAD as u32; - let mut forged = encrypted.clone(); - *forged.last_mut().unwrap() ^= 1; - for bytes in [&forged, &encrypted, &encrypted] { - let routed = - yas_composite_transport::encode_routed_datagram(token, bytes, wire_maximum) - .unwrap(); - remote.send_datagram(routed.into()).unwrap(); - } - assert_eq!( - yas_composite_transport::read_datagram(&mut side, maximum) - .await - .unwrap(), - b"private datagram" - ); - assert!( - tokio::time::timeout( - Duration::from_millis(50), - yas_composite_transport::read_datagram(&mut side, maximum) - ) - .await - .is_err() - ); - - yas_composite_transport::write_datagram(&mut side, b"private reply", maximum) - .await - .unwrap(); - let packet = remote.read_datagram().await.unwrap(); - let (received_token, ciphertext) = - yas_composite_transport::split_routed_datagram(&packet).unwrap(); - assert_eq!(received_token, token); - assert!(!ciphertext.windows(13).any(|w| w == b"private reply")); - assert_eq!(receiver.open(ciphertext).unwrap(), b"private reply"); - stream.shutdown().await.unwrap(); - main.shutdown().await.unwrap(); - producer.await.unwrap(); - } - datagrams.abort(); - local.close(0, b"test complete"); + Producer::new(&url, &token, crypto, local)? + .transport(transport) + .on_event(|event| eprintln!("[uplink] {event}")) + .run_until(async { + let _ = tokio::signal::ctrl_c().await; }) .await - .expect("producer uplink test stalled"); - } - - #[test] - fn parse_pool_accepts_plain_and_pinned_relay_urls() { - let pool = parse_pool( - r#"{"relays":[ - "https://relay-1.indent.com:4443/t/kfV3aB", - "https://[2001:db8::7]/session?key=x#sha256=AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA" - ],"ttl":60}"#, - ) - .unwrap(); - assert_eq!(pool.len(), 2); - assert_eq!( - pool[0].url.as_str(), - "https://relay-1.indent.com:4443/t/kfV3aB" - ); - assert_eq!(pool[0].label, "relay-1.indent.com:4443"); - assert!(pool[0].cert_hash.is_none()); - assert_eq!(pool[1].label, "[2001:db8::7]:443"); - assert_eq!(pool[1].cert_hash.as_ref().unwrap().len(), 32); - // The pin must be stripped before the URL is used to connect. - assert_eq!(pool[1].url.fragment(), None); - assert_eq!(pool[1].url.query(), Some("key=x")); - } - - #[test] - fn parse_pool_rejects_bad_input() { - assert!(parse_pool("not json").is_err()); - assert!(parse_pool(r#"{"relays":[]}"#).is_err()); - assert!(parse_pool(r#"{"relays":[{"host":"h"}]}"#).is_err()); - assert!( - parse_relay("http://relay.example/t/x").is_err(), - "non-https scheme must be rejected" - ); - assert!( - parse_relay("https://relay.example/t/x#sha256=AAAA").is_err(), - "a pin of the wrong length must be rejected, not ignored" - ); - assert!( - parse_relay("https://relay.example/t/x#pin=abc").is_err(), - "an unrecognized fragment must be rejected, not ignored" - ); - } - - #[test] - fn base64url_decodes() { - assert_eq!(base64url_decode("aGVsbG8").unwrap(), b"hello"); - assert_eq!(base64url_decode("aGVsbG8=").unwrap(), b"hello"); - assert_eq!(base64url_decode("_-8").unwrap(), vec![0xff, 0xef]); - assert!(base64url_decode("a+b").is_none()); - assert_eq!(base64url_decode("").unwrap(), Vec::::new()); - } } diff --git a/crates/cli/tests/client_host.rs b/crates/cli/tests/client_host.rs index f5befd7e..eca8e90c 100644 --- a/crates/cli/tests/client_host.rs +++ b/crates/cli/tests/client_host.rs @@ -1465,12 +1465,17 @@ exec sleep 600"#; std::fs::canonicalize(String::from_utf8(cwd).unwrap()).unwrap(), std::fs::canonicalize(directory.path()).unwrap() ); - // Nothing else starts: waiting for the next command finds none. + // Nothing else starts: waiting for the next command times out (TIMEOUT + // since #54; NOT_FOUND is for an exited terminal or an evicted index). let waited = client .wait_terminal_command(id, None, Duration::from_millis(200)) .await .unwrap_err(); - assert!(waited.is_not_found(), "{waited:?}"); + assert_eq!( + waited.status(), + Some(yas_client::wire::core::Status::Timeout), + "{waited:?}" + ); client.close_terminal(id).await.unwrap(); } @@ -1485,6 +1490,11 @@ async fn surfaces_are_none_without_the_compositor() { .await .unwrap_err(); assert!(missing.is_not_found(), "{missing:?}"); + let missing = client + .capture_surface_at(1, 1, CaptureFormat::Png) + .await + .unwrap_err(); + assert!(missing.is_not_found(), "{missing:?}"); } /// Against a real window: build the compositor's probe client @@ -1569,7 +1579,46 @@ async fn surfaces_capture_take_input_and_close_with_the_paste_probe() { .await .unwrap(); client.type_surface_text(surface.id, "é").await.unwrap(); + // At the revision listed, the capture goes at once. + let listed = client.surface(surface.id).await.unwrap(); + let png = client + .capture_surface_at(surface.id, listed.revision, CaptureFormat::Png) + .await + .unwrap(); + assert!(png.starts_with(b"\x89PNG\r\n\x1a\n"), "{} bytes", png.len()); client.resize_surface(surface.id, 320, 240).await.unwrap(); + // Once the window changed, that revision is stale: the server says so, and + // capture_surface_at looks the window up again rather than failing. + let deadline = Instant::now() + TIMEOUT; + while client.surface(surface.id).await.unwrap().revision == listed.revision { + assert!( + Instant::now() < deadline, + "the resize never changed the window" + ); + tokio::time::sleep(Duration::from_millis(50)).await; + } + let stale = client + .request_raw( + yas_client::wire::family::SURFACE, + yas_client::wire::surface::request_kind::CAPTURE, + yas_client::wire::Encode::encode(&yas_client::wire::surface::Capture { + surface_handle: surface.id, + revision: listed.revision, + initial_receive_credit: 0, + formats: vec![yas_client::wire::schema::surface::CAPTURE_PNG as u8], + extensions: Default::default(), + }) + .unwrap(), + None, + ) + .await + .unwrap(); + assert_eq!(stale.status, yas_client::wire::core::Status::Stale); + let png = client + .capture_surface_at(surface.id, listed.revision, CaptureFormat::Png) + .await + .unwrap(); + assert!(png.starts_with(b"\x89PNG\r\n\x1a\n"), "{} bytes", png.len()); // A read-only session sees the window but cannot send it input. let viewer = server diff --git a/crates/cli/tests/support/uplink.rs b/crates/cli/tests/support/uplink.rs index 0116544f..92194b18 100644 --- a/crates/cli/tests/support/uplink.rs +++ b/crates/cli/tests/support/uplink.rs @@ -1,5 +1,9 @@ +// Shared by the uplink E2E test and the browser fixture, which use different parts. +#![allow(dead_code)] + use futures_util::{SinkExt, StreamExt}; use std::{ + collections::HashMap, path::Path, process::Stdio, sync::{Arc, Mutex}, @@ -9,11 +13,18 @@ use tokio::{ io::{AsyncReadExt, AsyncWriteExt}, net::TcpListener, process::Command, + sync::{mpsc, oneshot}, }; use tokio_rustls::TlsAcceptor; -use tokio_tungstenite::tungstenite::Message; +use tokio_tungstenite::tungstenite::{ + Message, + handshake::server::{Request, Response}, +}; use web_transport_quinn as wt; +/// The WebSocket relay protocol's subprotocol (docs/uplink.md). +const WEBSOCKET_SUBPROTOCOL: &str = "yas-uplink.v1"; + pub fn cli(binary: &Path, directory: &Path, ca: &Path) -> Command { let mut command = Command::new(binary); command @@ -42,10 +53,41 @@ pub fn cli(binary: &Path, directory: &Path, ca: &Path) -> Command { .env("YAS_CHANNEL", "0") .env("YAS_REMOTES", directory.join("remotes")) .env_remove("YAS_TARGET") + .env_remove("YAS_UPLINK_TRANSPORT") .stdin(Stdio::null()); command } +/// What carries the producer's relay session in a fixture. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum Carrier { + /// A pool of WebTransport relays only, as older control endpoints give. + WebTransport, + /// A pool of both kinds, the producer forced onto WebSockets. + WebSocket, + /// A pool of both kinds whose WebTransport relay's UDP goes nowhere: the + /// producer (on its default transport) must fall back to WebSockets. + Fallback, +} + +/// The producer's relay session, as the relay holds it. +#[derive(Clone)] +enum ProducerLink { + WebTransport(Box), + /// Stream requests to send on the session WebSocket. + WebSocket(mpsc::UnboundedSender), +} + +type StreamSocket = + tokio_tungstenite::WebSocketStream>; + +#[derive(Default)] +struct RelayState { + producer: Option, + next_stream: u64, + streams: HashMap>, +} + pub struct Fixture { pub directory: tempfile::TempDir, pub ca: std::path::PathBuf, @@ -55,12 +97,19 @@ pub struct Fixture { pub capture: Arc>>, pub requests: Arc>>, pub producer: tokio::process::Child, + /// What the producer printed on stderr (it goes to the test's stderr too). + pub producer_log: Arc>>, _server: tokio::process::Child, + _black_hole: Option, tasks: Vec>, } impl Fixture { pub async fn start(binary: &Path) -> Self { + Self::start_with(binary, Carrier::WebTransport).await + } + + pub async fn start_with(binary: &Path, carrier: Carrier) -> Self { let _ = rustls::crypto::ring::default_provider().install_default(); let directory = tempfile::tempdir().unwrap(); let root = directory.path(); @@ -82,14 +131,35 @@ impl Fixture { .with_addr("127.0.0.1:0".parse().unwrap()) .with_certificate(vec![cert.cert.der().clone()], key()) .unwrap(); - let relay_url = format!( - "https://127.0.0.1:{}/producer", - worker.local_addr().unwrap().port() - ); + // A UDP socket that never answers stands for a network dropping UDP. + let black_hole = (carrier == Carrier::Fallback) + .then(|| std::net::UdpSocket::bind("127.0.0.1:0").unwrap()) + .map(|socket| { + socket.set_nonblocking(true).unwrap(); + tokio::net::UdpSocket::from_std(socket).unwrap() + }); + let relay_url = match &black_hole { + Some(hole) => format!( + "https://127.0.0.1:{}/producer", + hole.local_addr().unwrap().port() + ), + None => format!( + "https://127.0.0.1:{}/producer", + worker.local_addr().unwrap().port() + ), + }; let control = TcpListener::bind("127.0.0.1:0").await.unwrap(); let control_url = format!("https://{}", control.local_addr().unwrap()); let websocket = TcpListener::bind("127.0.0.1:0").await.unwrap(); - let ws_url = format!("wss://{}/consumer", websocket.local_addr().unwrap()); + let websocket_origin = format!("wss://{}", websocket.local_addr().unwrap()); + let ws_url = format!("{websocket_origin}/consumer"); + let pool = match carrier { + Carrier::WebTransport => serde_json::json!({ "relays": [relay_url] }), + Carrier::WebSocket | Carrier::Fallback => serde_json::json!({ + "relays": [relay_url], + "websockets": [format!("{websocket_origin}/producer")], + }), + }; let requests = Arc::new(Mutex::new(Vec::new())); let seen = requests.clone(); let control_tls = tls.clone(); @@ -97,7 +167,7 @@ impl Fixture { loop { let (stream, _) = control.accept().await.unwrap(); let tls = control_tls.clone(); - let relay_url = relay_url.clone(); + let pool = pool.to_string(); let ws_url = ws_url.clone(); let seen = seen.clone(); tokio::spawn(async move { @@ -114,10 +184,7 @@ impl Fixture { && header.contains("authorization: bearer producer-token\r\n") { seen.lock().unwrap().push("allocate"); - ( - "200 OK", - serde_json::json!({"relays": [relay_url]}).to_string(), - ) + ("200 OK", pool) } else if header.starts_with("get /attach ") && header.contains("authorization: bearer consumer-token\r\n") { @@ -158,7 +225,21 @@ impl Fixture { ); tokio::time::sleep(Duration::from_millis(20)).await; } - let mut producer = cli(binary, root, &ca) + + let state = Arc::new(Mutex::new(RelayState::default())); + let connected = Arc::new(tokio::sync::Notify::new()); + let capture = Arc::new(Mutex::new(Vec::new())); + let relay_task = tokio::spawn(serve_websockets( + websocket, + tls, + websocket_origin, + state.clone(), + connected.clone(), + capture.clone(), + )); + + let mut command = cli(binary, root, &ca); + command .env("YAS_UPLINK_IDENTITY", &producer_key) .env("YAS_UPLINK_TOKEN", "producer-token") .args([ @@ -168,61 +249,47 @@ impl Fixture { &consumer_public.to_string(), ]) .stdout(Stdio::null()) - .stderr(Stdio::inherit()) - .spawn() - .unwrap(); - let session = tokio::select! { - session = async { worker.accept().await.unwrap().ok().await.unwrap() } => session, - status = producer.wait() => panic!("uplink producer exited before connecting: {status:?}"), - }; - let capture = Arc::new(Mutex::new(Vec::new())); - let captured = capture.clone(); - let relay_task = tokio::spawn(async move { - loop { - let (stream, _) = websocket.accept().await.unwrap(); - let tls = tls.clone(); - let session = session.clone(); - let captured = captured.clone(); - tokio::spawn(async move { - let stream = tls.accept(stream).await.unwrap(); - let mut ws = tokio_tungstenite::accept_async(stream).await.unwrap(); - assert_eq!( - ws.next().await.unwrap().unwrap(), - Message::Text("consumer-token".into()) - ); - ws.send(Message::Text("ok".into())).await.unwrap(); - let (mut send, mut recv) = session.open_bi().await.unwrap(); - let (mut ws_send, mut ws_recv) = ws.split(); - let upstream = async { - while let Some(Ok(Message::Binary(bytes))) = ws_recv.next().await { - captured.lock().unwrap().extend_from_slice(&bytes); - if send.write_all(&bytes).await.is_err() { - break; - } - } - let _ = send.shutdown().await; - }; - let downstream = async { - let mut buffer = [0; 16384]; - while let Ok(Some(count)) = recv.read(&mut buffer).await { - captured.lock().unwrap().extend_from_slice(&buffer[..count]); - if count == 0 { - break; - } - if ws_send - .send(Message::Binary(buffer[..count].to_vec().into())) - .await - .is_err() - { - break; - } - } - let _ = ws_send.close().await; - }; - tokio::join!(upstream, downstream); - }); + .stderr(Stdio::piped()); + if carrier == Carrier::WebSocket { + command.env("YAS_UPLINK_TRANSPORT", "websocket"); + } + let mut producer = command.spawn().unwrap(); + let producer_log = Arc::new(Mutex::new(Vec::new())); + let log = producer_log.clone(); + let stderr = producer.stderr.take().unwrap(); + let log_task = tokio::spawn(async move { + use tokio::io::AsyncBufReadExt; + let mut lines = tokio::io::BufReader::new(stderr).lines(); + while let Ok(Some(line)) = lines.next_line().await { + eprintln!("{line}"); + log.lock().unwrap().push(line); } }); + match carrier { + Carrier::WebTransport => { + let session = tokio::select! { + session = async { worker.accept().await.unwrap().ok().await.unwrap() } => session, + status = producer.wait() => panic!("uplink producer exited before connecting: {status:?}"), + }; + state.lock().unwrap().producer = + Some(ProducerLink::WebTransport(Box::new(session))); + } + Carrier::WebSocket | Carrier::Fallback => { + let waiting = async { + loop { + let notified = connected.notified(); + if state.lock().unwrap().producer.is_some() { + break; + } + notified.await; + } + }; + tokio::select! { + () = waiting => {} + status = producer.wait() => panic!("uplink producer exited before connecting: {status:?}"), + } + } + } // Invoke the documented URL generator; the consumer resolves /attach // and uses the normal WSS connector, with no test-only transport hooks. let generated = cli(binary, root, &ca) @@ -255,12 +322,203 @@ impl Fixture { capture, requests, producer, + producer_log, _server: server, - tasks: vec![control_task, relay_task], + _black_hole: black_hole, + tasks: vec![control_task, relay_task, log_task], } } } +/// The relay's WebSocket side: consumers (`/consumer`), the producer's session +/// (`/producer`) and the stream WebSockets it opens (`/stream/`). +async fn serve_websockets( + listener: TcpListener, + tls: TlsAcceptor, + origin: String, + state: Arc>, + connected: Arc, + capture: Arc>>, +) { + loop { + let (stream, _) = listener.accept().await.unwrap(); + let tls = tls.clone(); + let origin = origin.clone(); + let state = state.clone(); + let connected = connected.clone(); + let capture = capture.clone(); + tokio::spawn(async move { + let Ok(stream) = tls.accept(stream).await else { + return; + }; + let path = Arc::new(Mutex::new(String::new())); + let seen = path.clone(); + // tungstenite's callback signature: its error is an HTTP response. + #[allow(clippy::result_large_err)] + let callback = move |request: &Request, mut response: Response| { + *seen.lock().unwrap() = request.uri().path().to_owned(); + let offered = request + .headers() + .get("sec-websocket-protocol") + .and_then(|value| value.to_str().ok()) + .unwrap_or_default(); + if offered + .split(',') + .any(|each| each.trim() == WEBSOCKET_SUBPROTOCOL) + { + response.headers_mut().insert( + "sec-websocket-protocol", + WEBSOCKET_SUBPROTOCOL.parse().unwrap(), + ); + } + Ok(response) + }; + let Ok(ws) = tokio_tungstenite::accept_hdr_async(stream, callback).await else { + return; + }; + let path = path.lock().unwrap().clone(); + if path == "/producer" { + let (mut sink, mut source) = ws.split(); + let (requests, mut pending) = mpsc::unbounded_channel(); + state.lock().unwrap().producer = Some(ProducerLink::WebSocket(requests)); + connected.notify_waiters(); + loop { + tokio::select! { + request = pending.recv() => match request { + Some(message) => if sink.send(message).await.is_err() { break }, + None => break, + }, + message = source.next() => match message { + Some(Ok(_)) => {} + _ => break, + }, + } + } + } else if let Some(id) = path.strip_prefix("/stream/") { + let waiting = state.lock().unwrap().streams.remove(id); + if let Some(waiting) = waiting { + let _ = waiting.send(ws); + } + } else if path == "/consumer" { + serve_consumer(ws, &origin, &state, &capture).await; + } + }); + } +} + +async fn serve_consumer( + mut ws: StreamSocket, + origin: &str, + state: &Arc>, + captured: &Arc>>, +) { + assert_eq!( + ws.next().await.unwrap().unwrap(), + Message::Text("consumer-token".into()) + ); + ws.send(Message::Text("ok".into())).await.unwrap(); + let link = state.lock().unwrap().producer.clone().expect("a producer"); + let (mut recv, mut send): ( + Box, + Box, + ) = match link { + ProducerLink::WebTransport(session) => { + let (send, recv) = session.open_bi().await.unwrap(); + (Box::new(recv), Box::new(send)) + } + ProducerLink::WebSocket(requests) => { + let (tx, rx) = oneshot::channel(); + let id = { + let mut state = state.lock().unwrap(); + state.next_stream += 1; + let id = state.next_stream.to_string(); + state.streams.insert(id.clone(), tx); + id + }; + let url = format!("{origin}/stream/{id}"); + requests + .send(Message::Text( + serde_json::json!({ "stream": url }).to_string().into(), + )) + .unwrap(); + let stream = tokio::time::timeout(Duration::from_secs(10), rx) + .await + .expect("the producer opened no stream WebSocket") + .unwrap(); + let (ours, theirs) = tokio::io::duplex(128 << 10); + tokio::spawn(pump(stream, ours)); + let (read, write) = tokio::io::split(theirs); + (Box::new(read), Box::new(write)) + } + }; + let (mut ws_send, mut ws_recv) = ws.split(); + let upstream = async { + while let Some(Ok(Message::Binary(bytes))) = ws_recv.next().await { + captured.lock().unwrap().extend_from_slice(&bytes); + if send.write_all(&bytes).await.is_err() { + break; + } + } + let _ = send.shutdown().await; + }; + let downstream = async { + let mut buffer = [0; 16384]; + loop { + match recv.read(&mut buffer).await { + Ok(0) | Err(_) => break, + Ok(count) => { + captured.lock().unwrap().extend_from_slice(&buffer[..count]); + if ws_send + .send(Message::Binary(buffer[..count].to_vec().into())) + .await + .is_err() + { + break; + } + } + } + } + let _ = ws_send.close().await; + }; + tokio::join!(upstream, downstream); +} + +/// Bytes between a stream WebSocket and a duplex pipe, as a relay forwards them. +async fn pump(socket: StreamSocket, pipe: tokio::io::DuplexStream) { + let (mut sink, mut source) = socket.split(); + let (mut read, mut write) = tokio::io::split(pipe); + let up = async { + while let Some(Ok(message)) = source.next().await { + match message { + Message::Binary(bytes) => { + if write.write_all(&bytes).await.is_err() { + break; + } + } + Message::Close(_) => break, + _ => {} + } + } + let _ = write.shutdown().await; + }; + let down = async { + let mut buffer = vec![0; 16 << 10]; + loop { + match read.read(&mut buffer).await { + Ok(0) | Err(_) => break, + Ok(count) => { + let chunk = buffer[..count].to_vec(); + if sink.send(Message::Binary(chunk.into())).await.is_err() { + break; + } + } + } + } + let _ = sink.close().await; + }; + tokio::join!(up, down); +} + impl Drop for Fixture { fn drop(&mut self) { let _ = self.producer.start_kill(); diff --git a/crates/cli/tests/uplink_e2e.rs b/crates/cli/tests/uplink_e2e.rs index 1c08c541..61fe12d8 100644 --- a/crates/cli/tests/uplink_e2e.rs +++ b/crates/cli/tests/uplink_e2e.rs @@ -4,18 +4,64 @@ #[path = "support/uplink.rs"] mod uplink; use std::{path::Path, time::Duration}; -use uplink::{Fixture, cli}; +use uplink::{Carrier, Fixture, cli}; #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn uplink_cli_remote_execution_and_authentication() { - tokio::time::timeout(Duration::from_secs(45), exercise_uplink()) + tokio::time::timeout( + Duration::from_secs(45), + exercise_uplink(Carrier::WebTransport), + ) + .await + .expect("uplink E2E stalled"); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn uplink_over_websockets_when_forced() { + tokio::time::timeout(Duration::from_secs(45), exercise_uplink(Carrier::WebSocket)) + .await + .expect("uplink E2E over WebSockets stalled"); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn uplink_falls_back_to_websockets_where_udp_goes_nowhere() { + tokio::time::timeout(Duration::from_secs(60), exercise_uplink(Carrier::Fallback)) .await - .expect("uplink E2E stalled"); + .expect("uplink E2E falling back to WebSockets stalled"); } -async fn exercise_uplink() { +async fn exercise_uplink(carrier: Carrier) { let binary = Path::new(env!("CARGO_BIN_EXE_yas")); - let mut fixture = Fixture::start(binary).await; + let mut fixture = Fixture::start_with(binary, carrier).await; + { + // The relay sees the session before the producer's line about it is read. + let deadline = std::time::Instant::now() + Duration::from_secs(10); + let log = loop { + let log = fixture.producer_log.lock().unwrap().clone(); + let connected = log + .iter() + .any(|line| line.starts_with("[uplink] connected to relay")); + if connected || std::time::Instant::now() > deadline { + assert!(connected, "producer log: {log:?}"); + break log; + } + tokio::time::sleep(Duration::from_millis(50)).await; + }; + let over_websocket = log.iter().any(|line| { + line.starts_with("[uplink] connected to relay") && line.ends_with(" over WebSocket") + }); + assert_eq!( + over_websocket, + carrier != Carrier::WebTransport, + "producer log: {log:?}" + ); + let gave_up_on_udp = log.iter().any(|line| line.contains("(is UDP blocked?)")); + assert_eq!( + gave_up_on_udp, + carrier == Carrier::Fallback, + "producer log: {log:?}" + ); + } let root = fixture.directory.path(); let ca = &fixture.ca; let uri = &fixture.uri; diff --git a/crates/client/src/lib.rs b/crates/client/src/lib.rs index 2ecfeed3..67ae4488 100644 --- a/crates/client/src/lib.rs +++ b/crates/client/src/lib.rs @@ -75,6 +75,9 @@ pub use client::{Client, DEFAULT_REQUEST_TIMEOUT}; pub use error::{Error, Result, format_result_detail}; pub use options::{ConnectOptions, HelloOptions, read_only_extension}; +/// The producer side of YAS uplinks (`yas uplink`): publish a YAS server +/// through a relay, over WebTransport or WebSockets. +pub use yas_proxy::uplink_producer; /// The wire codecs this crate speaks (`yas-wire`), for the escape hatches /// ([`Client::request`], [`Client::request_raw`]). pub use yas_wire as wire; diff --git a/crates/client/src/surface.rs b/crates/client/src/surface.rs index 3049b7e8..e6ab0ea4 100644 --- a/crates/client/src/surface.rs +++ b/crates/client/src/surface.rs @@ -4,7 +4,8 @@ //! `YAS_SKIP_COMPOSITOR=1`; [`crate::host::HostOptions::compositor`] for a //! hosted one), the programs it starts get a `WAYLAND_DISPLAY`, and each //! window they map is a surface in the catalogue ([`Client::surfaces`]). A -//! client can capture one as an image ([`Client::capture_surface`]), click, +//! client can capture one as an image ([`Client::capture_surface`], or +//! [`Client::capture_surface_at`] with the revision it listed), click, //! scroll and type into it, resize, focus and close it: what `yas surface` //! does, for programs that drive GUIs (a screenshot, then a click). //! @@ -116,9 +117,34 @@ impl Client { Ok(SurfaceInfo::from_record(&self.surface_record(id).await?)) } - /// What the window shows now, as an image. + /// What the window shows now, as an image: its revision looked up, then `CAPTURE`, two + /// round trips. [`Client::capture_surface_at`] takes one with a revision already listed. pub async fn capture_surface(&self, id: u64, format: CaptureFormat) -> Result> { let record = self.surface_record(id).await?; + self.capture(id, record.revision, format).await + } + + /// What the window shows now, as an image, given the revision the caller last saw it at + /// ([`SurfaceInfo::revision`], from [`Client::surfaces`] or [`Client::surface`]): one + /// round trip. When the window changed since (the server answers `STALE`), it is looked up + /// again and captured at its current revision, as [`Client::capture_surface`] does. + pub async fn capture_surface_at( + &self, + id: u64, + revision: u64, + format: CaptureFormat, + ) -> Result> { + match self.capture(id, revision, format).await { + Err(error) if error.status() == Some(yas_wire::core::Status::Stale) => { + self.capture_surface(id, format).await + } + captured => captured, + } + } + + /// `CAPTURE` of the window at `revision`, which the server answers `STALE` unless it + /// is the window's current one. + async fn capture(&self, id: u64, revision: u64, format: CaptureFormat) -> Result> { let hook: Hook = Box::new(|prefix: &ResultPrefix| { InlineOrTransfer::decode(&prefix.body) .map(|result| delivery_routes(&result)) @@ -130,7 +156,7 @@ impl Client { request_kind::CAPTURE, wire::Capture { surface_handle: id, - revision: record.revision, + revision, initial_receive_credit: CAPTURE_CREDIT, formats: vec![match format { CaptureFormat::Png => schema::CAPTURE_PNG as u8, diff --git a/crates/client/src/terminal.rs b/crates/client/src/terminal.rs index 72bca5d1..d0317995 100644 --- a/crates/client/src/terminal.rs +++ b/crates/client/src/terminal.rs @@ -381,7 +381,8 @@ impl Client { /// one with this journal `index`, or the one running now (else the next /// one to start). Answers its journal record (exit code, times, command /// line). When it still runs at the timeout: a `Timeout` error; when - /// none started by then (or it left the backlog): a `NotFound` status. + /// none started by then: a `Timeout` status; when it left the backlog, + /// or the terminal exited: a `NotFound` status. /// Needs shell integration. pub async fn wait_terminal_command( &self, diff --git a/crates/compositor/src/imp.rs b/crates/compositor/src/imp.rs index fdcf2dc4..31759d6d 100644 --- a/crates/compositor/src/imp.rs +++ b/crates/compositor/src/imp.rs @@ -12228,6 +12228,13 @@ impl CompositorCommandSender { ) -> Result<(), mpsc::SendError> { send_command_with_wake(&self.command_tx, command, || self.loop_signal.wakeup()) } + + /// Wake the compositor loop, after a command admitted on the raw channel + /// (`try_send`): an idle loop otherwise sees it only at its next dispatch + /// timeout, up to a second later. + pub fn wake(&self) { + self.loop_signal.wakeup(); + } } fn send_command_with_wake( diff --git a/crates/compositor/src/lib.rs b/crates/compositor/src/lib.rs index 50b37ad0..1854954f 100644 --- a/crates/compositor/src/lib.rs +++ b/crates/compositor/src/lib.rs @@ -709,6 +709,9 @@ mod stub { ) -> Result<(), mpsc::SendError> { self.command_tx.send(command) } + + /// Wake the compositor event loop immediately. + pub fn wake(&self) {} } impl CompositorHandle { diff --git a/crates/proxy/Cargo.toml b/crates/proxy/Cargo.toml index 0ad40c22..6eb4f038 100644 --- a/crates/proxy/Cargo.toml +++ b/crates/proxy/Cargo.toml @@ -54,3 +54,4 @@ libc = "0.2" [dev-dependencies] tempfile = "3" rcgen = "0.14" +tokio-rustls = { version = "0.26", default-features = false, features = ["tls12"] } diff --git a/crates/proxy/src/lib.rs b/crates/proxy/src/lib.rs index 42cb43eb..fe4e860a 100644 --- a/crates/proxy/src/lib.rs +++ b/crates/proxy/src/lib.rs @@ -782,6 +782,8 @@ fn parse_uplink_uri(rest: &str) -> Result { }) } +pub mod uplink_producer; + /// HTTPS control client with the same explicit CA override semantics as the /// WSS and WebTransport legs. Reqwest's platform verifier otherwise ignores /// SSL_CERT_FILE/SSL_CERT_DIR on macOS. @@ -1055,10 +1057,12 @@ async fn connect_ws_mode( .max_message_size(Some(64 * 1024)) .max_frame_size(Some(64 * 1024)) }); + // Nagle's algorithm off (`disable_nagle`), as for tcp: upstreams: a request's small frames + // must not wait for the ACK of the one before, which the relay delays (40 ms on Linux). let (mut ws, response) = tokio_tungstenite::connect_async_tls_with_config( request, config, - false, + true, Some(yas_webrtc_forwarder::tls::websocket_connector()), ) .await diff --git a/crates/proxy/src/uplink_producer.rs b/crates/proxy/src/uplink_producer.rs new file mode 100644 index 00000000..773db6a4 --- /dev/null +++ b/crates/proxy/src/uplink_producer.rs @@ -0,0 +1,1361 @@ +//! The producer side of a YAS uplink (`yas uplink`, docs/uplink.md), as a +//! library: publish a local YAS server through an untrusted relay. +//! +//! A [`Producer`] authenticates to an HTTPS control endpoint, receives a pool +//! of relays, holds a session with one of them, and bridges each consumer the +//! relay sends it to a fresh connection to the local server, once the consumer +//! has passed end-to-end Noise IK authentication against pinned keys. The +//! session is carried either by WebTransport (HTTP/3 over UDP) or by +//! WebSockets (over TCP, where networks block UDP); [`Transport`] chooses, and +//! [`Transport::Auto`] falls back from one to the other by itself. +//! +//! ```no_run +//! # async fn run() -> Result<(), String> { +//! use yas_proxy::uplink_producer::{Local, Producer, Transport}; +//! +//! let identity = yas_uplink::Identity::from_base64(&std::env::var("YAS_UPLINK_IDENTITY").unwrap())?; +//! let allowed = vec!["CLIENT_PUBLIC_KEY_43_CHARACTERS_OF_BASE64URL".parse()?]; +//! Producer::new( +//! "https://relay.example/uplink/control", +//! &std::env::var("YAS_UPLINK_TOKEN").unwrap(), +//! identity.server_config(allowed)?, +//! Local::socket("/run/user/1000/yas/yas.sock"), +//! )? +//! .transport(Transport::from_env()?) +//! .on_event(|event| eprintln!("[uplink] {event}")) +//! .run_until(async { tokio::signal::ctrl_c().await.ok(); }) +//! .await +//! # } +//! ``` + +use std::fmt; +use std::future::Future; +use std::pin::Pin; +use std::sync::{Arc, Mutex}; +use std::time::{Duration, Instant}; + +use futures_util::{SinkExt, StreamExt}; +use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt}; +use tokio_tungstenite::tungstenite::Message; +use tokio_tungstenite::tungstenite::protocol::frame::coding::CloseCode; +use web_transport_quinn as wt; + +/// The keys a producer is configured with. +pub use yas_uplink::{Identity, PublicKey, ServerConfig}; + +/// The WebSocket subprotocol of the uplink's WebSocket relay sessions and +/// streams (docs/uplink.md, "WebSocket relay session"), with its version. +pub const WEBSOCKET_SUBPROTOCOL: &str = "yas-uplink.v1"; + +/// The environment variable [`Transport::from_env`] reads. +pub const TRANSPORT_ENV: &str = "YAS_UPLINK_TRANSPORT"; + +/// WebTransport application code closing a session on shutdown. +const CODE_SHUTDOWN: u32 = 2; +const INITIAL_BACKOFF: Duration = Duration::from_secs(1); +const MAX_BACKOFF: Duration = Duration::from_secs(60); +const MAX_DATAGRAM_ROUTES: usize = 4_096; +const DATAGRAM_QUEUE: usize = 64; +/// Consumer streams authenticating at once on one session. +const MAX_PENDING: usize = 64; +/// Liveness of a session: a keepalive this often, dead after this long silent. +const KEEPALIVE: Duration = Duration::from_secs(10); +const IDLE_TIMEOUT: Duration = Duration::from_secs(30); +/// How long a WebSocket (session or stream) may take to open. +const WEBSOCKET_CONNECT: Duration = Duration::from_secs(10); +/// How long [`Transport::Auto`] waits for a WebTransport session before it +/// tries the pool's WebSocket relays. Where UDP is dropped rather than +/// refused, QUIC would otherwise wait out its whole idle timeout. +const WEBTRANSPORT_CONNECT_AUTO: Duration = Duration::from_secs(5); +/// A WebTransport session shorter than this counts as WebTransport failing. +const SHORT_SESSION: Duration = Duration::from_secs(60); +/// How long [`Transport::Auto`] tries WebSockets first after WebTransport failed. +const PREFER_WEBSOCKET_FOR: Duration = Duration::from_secs(600); +/// The largest WebSocket message either way. +const MAX_WEBSOCKET_MESSAGE: usize = 64 << 10; + +type DatagramRoutes = yas_composite_transport::RoutedDatagramRoutes; +type BoxRead = Box; +type BoxWrite = Box; +type Socket = + tokio_tungstenite::WebSocketStream>; + +// --------------------------------------------------------------------------- +// Configuration +// --------------------------------------------------------------------------- + +/// What carries the relay session. +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +pub enum Transport { + /// WebTransport when it works, else WebSockets (when the control endpoint + /// offers WebSocket relays): the default. + #[default] + Auto, + /// WebTransport (HTTP/3 over UDP) only. + WebTransport, + /// WebSockets (over TCP) only. + WebSocket, +} + +impl Transport { + /// `YAS_UPLINK_TRANSPORT` (`auto`, `webtransport` or `websocket`), else + /// [`Transport::Auto`]. + pub fn from_env() -> Result { + match std::env::var(TRANSPORT_ENV) { + Ok(value) if !value.trim().is_empty() => value.parse(), + Ok(_) | Err(std::env::VarError::NotPresent) => Ok(Self::Auto), + Err(std::env::VarError::NotUnicode(_)) => Err(format!( + "{TRANSPORT_ENV} must be auto, webtransport or websocket" + )), + } + } +} + +impl std::str::FromStr for Transport { + type Err = String; + + fn from_str(value: &str) -> Result { + match value.trim().to_ascii_lowercase().as_str() { + "auto" => Ok(Self::Auto), + "webtransport" => Ok(Self::WebTransport), + "websocket" => Ok(Self::WebSocket), + other => Err(format!( + "unknown uplink transport {other:?}: auto, webtransport or websocket" + )), + } + } +} + +impl fmt::Display for Transport { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(match self { + Self::Auto => "auto", + Self::WebTransport => "webtransport", + Self::WebSocket => "websocket", + }) + } +} + +/// What carries one relay session. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum Carrier { + WebTransport, + WebSocket, +} + +impl fmt::Display for Carrier { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(match self { + Self::WebTransport => "WebTransport", + Self::WebSocket => "WebSocket", + }) + } +} + +type Connect = dyn Fn() -> Pin> + Send>> + + Send + + Sync; + +/// How the producer reaches the local YAS server, once per consumer (twice for +/// a consumer with a datagram lane). +#[derive(Clone)] +pub struct Local { + label: String, + connect: Arc, +} + +impl Local { + /// The YAS server listening at `path`: a Unix socket, or a named pipe on + /// Windows. + pub fn socket(path: impl Into) -> Self { + let path = path.into(); + let label = path.clone(); + Self { + label, + connect: Arc::new(move || { + let path = path.clone(); + Box::pin(async move { connect_socket_split(&path).await }) + }), + } + } + + /// Any other way to reach the server (an in-process one, say): `connect` + /// opens a fresh byte stream to it each time; `label` names it in events. + pub fn custom(label: impl Into, connect: F) -> Self + where + F: Fn() -> Fut + Send + Sync + 'static, + Fut: Future> + Send + 'static, + S: AsyncRead + AsyncWrite + Unpin + Send + 'static, + { + Self { + label: label.into(), + connect: Arc::new(move || { + let connecting = connect(); + Box::pin(async move { + let stream = connecting.await.map_err(|error| error.to_string())?; + let (reader, writer) = tokio::io::split(stream); + Ok((Box::new(reader) as BoxRead, Box::new(writer) as BoxWrite)) + }) + }), + } + } + + async fn connect(&self) -> Result<(BoxRead, BoxWrite), String> { + (self.connect)().await + } +} + +impl fmt::Debug for Local { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("Local").field("label", &self.label).finish() + } +} + +async fn connect_socket_split(path: &str) -> Result<(BoxRead, BoxWrite), String> { + #[cfg(unix)] + { + let stream = tokio::net::UnixStream::connect(path) + .await + .map_err(|error| error.to_string())?; + let (reader, writer) = tokio::io::split(stream); + Ok((Box::new(reader), Box::new(writer))) + } + #[cfg(windows)] + { + // Every instance may be taken for a while (the server makes its next + // one only once it accepts the last): wait out ERROR_PIPE_BUSY. + use tokio::net::windows::named_pipe::ClientOptions; + const ERROR_PIPE_BUSY: i32 = 231; + let deadline = Instant::now() + Duration::from_secs(10); + let pipe = loop { + match ClientOptions::new().open(path) { + Err(error) + if error.raw_os_error() == Some(ERROR_PIPE_BUSY) + && Instant::now() < deadline => + { + tokio::time::sleep(Duration::from_millis(20)).await; + } + opened => break opened.map_err(|error| error.to_string())?, + } + }; + let (reader, writer) = tokio::io::split(pipe); + Ok((Box::new(reader), Box::new(writer))) + } +} + +/// What happens to a producer, for logs and status. Displayed, each is the +/// text `yas uplink` prints after `[uplink] `. +#[derive(Clone, Debug)] +#[non_exhaustive] +pub enum Event { + /// A relay session is up: consumers can reach the server. `relay` is the + /// relay's `host:port` (its URL is a credential and never shown). + Connected { relay: String, carrier: Carrier }, + /// The session that was up ended; the producer asks the control endpoint + /// for a new pool at once. + SessionEnded { carrier: Carrier, reason: String }, + /// A relay of the pool didn't take a session; the next one is tried. + RelayFailed { + relay: String, + carrier: Carrier, + error: String, + }, + /// No relay of the pool took a session. + PoolExhausted { retry_in: Duration }, + /// The control endpoint couldn't give a pool (unreachable, an error, a + /// pool naming no relay the transport may use). + ControlRetry { reason: String, retry_in: Duration }, + /// An authenticated consumer couldn't reach the local server. + LocalUnavailable { local: String, error: String }, + /// A consumer stream couldn't be served (a WebSocket stream the relay + /// asked for that didn't open, a datagram lane this session can't carry). + StreamFailed { error: String }, +} + +impl fmt::Display for Event { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + let over = |carrier: &Carrier| match carrier { + Carrier::WebTransport => "", + Carrier::WebSocket => " over WebSocket", + }; + match self { + Self::Connected { relay, carrier } => { + write!(f, "connected to relay {relay}{}", over(carrier)) + } + Self::SessionEnded { reason, .. } => write!(f, "session ended: {reason}"), + Self::RelayFailed { + relay, + carrier, + error, + } => write!(f, "relay {relay}{}: {error}", over(carrier)), + Self::PoolExhausted { retry_in } => write!( + f, + "relay pool exhausted; re-querying in {}s", + retry_in.as_secs() + ), + Self::ControlRetry { reason, retry_in } => { + write!(f, "{reason}; retrying in {}s", retry_in.as_secs()) + } + Self::LocalUnavailable { local, error } => { + write!(f, "local yas server unavailable at {local}: {error}") + } + Self::StreamFailed { error } => write!(f, "consumer stream failed: {error}"), + } + } +} + +type Events = Arc; + +/// A YAS server published through a relay. See the module documentation. +#[derive(Clone)] +pub struct Producer { + control: url::Url, + shared: Arc, + local: Local, + transport: Transport, + events: Events, +} + +/// What a running producer takes from its [`Handle`]. +struct Shared { + token: Mutex, + crypto: Mutex>, +} + +impl Shared { + fn token(&self) -> String { + self.token.lock().unwrap_or_else(|p| p.into_inner()).clone() + } + + fn crypto(&self) -> Arc { + self.crypto + .lock() + .unwrap_or_else(|p| p.into_inner()) + .clone() + } +} + +/// Changes a running [`Producer`] takes without restarting its session. +#[derive(Clone)] +pub struct Handle(Arc); + +impl Handle { + /// The token the next control request presents (a renewed one, say). + pub fn set_token(&self, token: impl Into) { + *self.0.token.lock().unwrap_or_else(|p| p.into_inner()) = token.into(); + } + + /// Who may connect from now on: consumers that start their handshake + /// afterwards are checked against `crypto`'s allowlist (and answered with + /// its identity). Consumers already connected keep their authority until + /// they disconnect, as they do across a restart. + pub fn set_server_config(&self, crypto: Arc) { + *self.0.crypto.lock().unwrap_or_else(|p| p.into_inner()) = crypto; + } +} + +impl fmt::Debug for Handle { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("Handle").finish_non_exhaustive() + } +} + +impl fmt::Debug for Producer { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + // The token and keys are credentials: never shown. + f.debug_struct("Producer") + .field("control", &label(&self.control)) + .field("local", &self.local) + .field("transport", &self.transport) + .finish_non_exhaustive() + } +} + +impl Producer { + /// A producer authenticating to `control` (HTTPS, without userinfo or a + /// fragment) with `token`, admitting the consumers `crypto` allows + /// ([`yas_uplink::Identity::server_config`]), bridged to `local`. + pub fn new( + control: &str, + token: &str, + crypto: Arc, + local: Local, + ) -> Result { + let control = url::Url::parse(control).map_err(|_| "invalid uplink control URL")?; + if control.scheme() != "https" + || control.host_str().is_none() + || !control.username().is_empty() + || control.password().is_some() + || control.fragment().is_some() + { + return Err("uplink control URL must be HTTPS without userinfo or a fragment".into()); + } + if token.is_empty() { + return Err("YAS_UPLINK_TOKEN is not set".into()); + } + Ok(Self { + control, + shared: Arc::new(Shared { + token: Mutex::new(token.to_owned()), + crypto: Mutex::new(crypto), + }), + local, + transport: Transport::Auto, + events: Arc::new(|_| {}), + }) + } + + /// Where to change the token and the allowlist while it runs. + pub fn handle(&self) -> Handle { + Handle(self.shared.clone()) + } + + /// What may carry the relay session ([`Transport::Auto`] by default). + pub fn transport(mut self, transport: Transport) -> Self { + self.transport = transport; + self + } + + /// Hear what happens (nothing is printed otherwise). + pub fn on_event(mut self, events: impl Fn(&Event) + Send + Sync + 'static) -> Self { + self.events = Arc::new(events); + self + } + + /// Publish the server until a fatal error: the control endpoint refusing + /// the token. Every other failure is retried. + pub async fn run(self) -> Result<(), String> { + self.run_until(std::future::pending()).await + } + + /// [`Producer::run`] until `shutdown` completes, then close the relay + /// session rather than let it idle out on the relay. + pub async fn run_until(self, shutdown: impl Future) -> Result<(), String> { + let active = Active::default(); + tokio::select! { + () = shutdown => { + active.close().await; + Ok(()) + } + result = self.run_loop(&active) => result, + } + } + + fn emit(&self, event: Event) { + (self.events)(&event); + } + + async fn run_loop(&self, active: &Active) -> Result<(), String> { + let http = crate::uplink_http_client()?; + let mut backoff = INITIAL_BACKOFF; + let mut prefer_websocket_until: Option = None; + loop { + let token = self.shared.token(); + let pool = match fetch_pool(&http, &self.control, &token).await? { + FetchOutcome::Pool(pool) => pool, + FetchOutcome::Retry { after, reason } => { + let delay = after.unwrap_or(backoff); + self.emit(Event::ControlRetry { + reason, + retry_in: delay, + }); + tokio::time::sleep(jittered(delay)).await; + backoff = (backoff * 2).min(MAX_BACKOFF); + continue; + } + }; + let websocket_first = + prefer_websocket_until.is_some_and(|until| Instant::now() < until); + let relays = order(pool, self.transport, websocket_first); + if relays.is_empty() { + let reason = format!( + "the relay pool has no relay for the {} transport", + self.transport + ); + self.emit(Event::ControlRetry { + reason, + retry_in: backoff, + }); + tokio::time::sleep(jittered(backoff)).await; + backoff = (backoff * 2).min(MAX_BACKOFF); + continue; + } + let has_websocket = relays + .iter() + .any(|relay| relay.carrier == Carrier::WebSocket); + // A session that was actually established ends with a fresh + // control-plane query, per docs/uplink.md. + let mut established = false; + for relay in &relays { + let end = match relay.carrier { + Carrier::WebTransport => { + let limit = (self.transport == Transport::Auto && has_websocket) + .then_some(WEBTRANSPORT_CONNECT_AUTO); + self.webtransport_session(relay, limit, active).await + } + Carrier::WebSocket => self.websocket_session(relay, active).await, + }; + match end { + SessionEnd::Ended { reason, lasted } => { + self.emit(Event::SessionEnded { + carrier: relay.carrier, + reason, + }); + if relay.carrier == Carrier::WebTransport && lasted < SHORT_SESSION { + prefer_websocket_until = Some(Instant::now() + PREFER_WEBSOCKET_FOR); + } + established = true; + break; + } + SessionEnd::NeverConnected(error) => { + if relay.carrier == Carrier::WebTransport { + prefer_websocket_until = Some(Instant::now() + PREFER_WEBSOCKET_FOR); + } + self.emit(Event::RelayFailed { + relay: relay.label.clone(), + carrier: relay.carrier, + error, + }); + } + } + } + if established { + backoff = INITIAL_BACKOFF; + } else { + self.emit(Event::PoolExhausted { retry_in: backoff }); + tokio::time::sleep(jittered(backoff)).await; + backoff = (backoff * 2).min(MAX_BACKOFF); + } + } + } + + async fn webtransport_session( + &self, + relay: &Relay, + limit: Option, + active: &Active, + ) -> SessionEnd { + let client = match webtransport_client(relay.cert_hash.as_deref()) { + Ok(client) => client, + Err(error) => return SessionEnd::NeverConnected(error), + }; + // Careful: the URL is the credential — show `relay.label` only. + let connecting = client.connect(relay.url.clone()); + let connected = match limit { + Some(limit) => match tokio::time::timeout(limit, connecting).await { + Ok(connected) => connected, + Err(_) => { + return SessionEnd::NeverConnected(format!( + "no WebTransport session within {}s (is UDP blocked?)", + limit.as_secs() + )); + } + }, + None => connecting.await, + }; + let session = match connected { + Ok(session) => session, + Err(error) => return SessionEnd::NeverConnected(format!("connect failed: {error}")), + }; + let started = Instant::now(); + self.emit(Event::Connected { + relay: relay.label.clone(), + carrier: Carrier::WebTransport, + }); + active.set(Session::WebTransport(Box::new(session.clone()))); + + let routes = DatagramRoutes::new(MAX_DATAGRAM_ROUTES); + let distributing = tokio::spawn(distribute_datagrams(session.clone(), routes.clone())); + let pending = Arc::new(tokio::sync::Semaphore::new(MAX_PENDING)); + let reason = loop { + match session.accept_bi().await { + Ok((send, recv)) => { + let Ok(permit) = pending.clone().try_acquire_owned() else { + continue; + }; + let lane = Lane::WebTransport { + session: Box::new(session.clone()), + routes: routes.clone(), + }; + tokio::spawn(bridge( + tokio::io::join(recv, send), + lane, + self.shared.clone(), + permit, + self.local.clone(), + self.events.clone(), + )); + } + Err(error) => break error.to_string(), + } + }; + distributing.abort(); + active.clear(); + SessionEnd::Ended { + reason, + lasted: started.elapsed(), + } + } + + async fn websocket_session(&self, relay: &Relay, active: &Active) -> SessionEnd { + let socket = match connect_websocket(&relay.url, relay.cert_hash.as_deref()).await { + Ok(socket) => socket, + Err(error) => return SessionEnd::NeverConnected(error), + }; + let started = Instant::now(); + self.emit(Event::Connected { + relay: relay.label.clone(), + carrier: Carrier::WebSocket, + }); + let (sink, mut source) = socket.split(); + let sink = Arc::new(tokio::sync::Mutex::new(sink)); + active.set(Session::WebSocket(sink.clone())); + + let pending = Arc::new(tokio::sync::Semaphore::new(MAX_PENDING)); + // Stream WebSockets are connections of their own: they end with the + // session (dropped, the set aborts them), as WebTransport streams end + // with their QUIC connection. + let mut streams = tokio::task::JoinSet::new(); + let mut keepalive = tokio::time::interval(KEEPALIVE); + keepalive.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay); + keepalive.tick().await; + let mut heard = Instant::now(); + let reason = loop { + tokio::select! { + message = source.next() => match message { + Some(Ok(Message::Text(text))) => { + heard = Instant::now(); + let Some(stream) = stream_request(&text, &relay.url) else { + continue; + }; + let Ok(permit) = pending.clone().try_acquire_owned() else { + continue; + }; + streams.spawn(websocket_stream( + stream, + relay.cert_hash.clone(), + self.shared.clone(), + permit, + self.local.clone(), + self.events.clone(), + )); + } + Some(Ok(Message::Close(frame))) => { + break match frame { + Some(frame) if !frame.reason.is_empty() => { + format!("closed by the relay: {}", frame.reason) + } + _ => "closed by the relay".to_owned(), + }; + } + Some(Ok(_)) => heard = Instant::now(), + Some(Err(error)) => break error.to_string(), + None => break "the relay closed the connection".to_owned(), + }, + Some(_) = streams.join_next(), if !streams.is_empty() => {} + _ = keepalive.tick() => { + if heard.elapsed() >= IDLE_TIMEOUT { + break format!("nothing from the relay for {}s", IDLE_TIMEOUT.as_secs()); + } + let ping = Message::Ping(Default::default()); + if sink.lock().await.send(ping).await.is_err() { + break "the connection to the relay failed".to_owned(); + } + } + } + }; + active.clear(); + SessionEnd::Ended { + reason, + lasted: started.elapsed(), + } + } +} + +/// The relays to try, in order: WebTransport's then WebSocket's for +/// [`Transport::Auto`] (the other way round while WebTransport is failing), +/// each carrier's shuffled. +fn order(pool: Pool, transport: Transport, websocket_first: bool) -> Vec { + use rand::seq::SliceRandom; + let Pool { + mut webtransport, + mut websocket, + } = pool; + webtransport.shuffle(&mut rand::rng()); + websocket.shuffle(&mut rand::rng()); + match transport { + Transport::WebTransport => webtransport, + Transport::WebSocket => websocket, + Transport::Auto if websocket_first => websocket.into_iter().chain(webtransport).collect(), + Transport::Auto => webtransport.into_iter().chain(websocket).collect(), + } +} + +// --------------------------------------------------------------------------- +// The active session, for a graceful shutdown +// --------------------------------------------------------------------------- + +type WsSink = futures_util::stream::SplitSink; + +enum Session { + WebTransport(Box), + WebSocket(Arc>), +} + +#[derive(Clone, Default)] +struct Active(Arc>>); + +impl Active { + fn set(&self, session: Session) { + *self.0.lock().unwrap_or_else(|p| p.into_inner()) = Some(session); + } + + fn clear(&self) { + self.0.lock().unwrap_or_else(|p| p.into_inner()).take(); + } + + async fn close(&self) { + let session = self.0.lock().unwrap_or_else(|p| p.into_inner()).take(); + match session { + Some(Session::WebTransport(session)) => { + session.close(CODE_SHUTDOWN, b"uplink shutting down"); + } + Some(Session::WebSocket(sink)) => { + let frame = tokio_tungstenite::tungstenite::protocol::CloseFrame { + code: CloseCode::Away, + reason: "uplink shutting down".into(), + }; + let _ = tokio::time::timeout(Duration::from_secs(1), async { + sink.lock().await.send(Message::Close(Some(frame))).await + }) + .await; + } + None => {} + } + } +} + +// --------------------------------------------------------------------------- +// Control plane +// --------------------------------------------------------------------------- + +struct Relay { + /// Connection URL as minted by the control plane (fragment stripped). + /// The URL is the session credential — never show it; use `label`. + url: url::Url, + /// `host:port`, for events. + label: String, + /// SHA-256 of the relay's certificate (DER), from the URL's `#sha256=` + /// fragment; pins TLS verification instead of using system trust roots. + cert_hash: Option>, + carrier: Carrier, +} + +#[derive(Default)] +struct Pool { + webtransport: Vec, + websocket: Vec, +} + +enum FetchOutcome { + Pool(Pool), + Retry { + after: Option, + reason: String, + }, +} + +/// Query the control endpoint. `Err` is fatal (bad token); every other +/// failure is a retryable `FetchOutcome::Retry`. +async fn fetch_pool( + http: &reqwest::Client, + url: &url::Url, + token: &str, +) -> Result { + let retry = |after, reason| Ok(FetchOutcome::Retry { after, reason }); + let response = match http + .get(url.clone()) + .header("authorization", format!("Bearer {token}")) + .header("accept", "application/json") + .send() + .await + { + Ok(response) => response, + Err(error) => return retry(None, format!("control endpoint unreachable: {error}")), + }; + let status = response.status(); + if status.as_u16() == 401 || status.as_u16() == 403 { + return Err(format!( + "control endpoint rejected YAS_UPLINK_TOKEN ({status})" + )); + } + if !status.is_success() { + let after = response + .headers() + .get("retry-after") + .and_then(|value| value.to_str().ok()) + .and_then(|value| value.trim().parse::().ok()) + .map(Duration::from_secs); + return retry(after, format!("control endpoint returned {status}")); + } + let body = match response.text().await { + Ok(body) => body, + Err(error) => return retry(None, format!("error reading relay pool: {error}")), + }; + match parse_pool(&body) { + Ok(pool) => Ok(FetchOutcome::Pool(pool)), + Err(error) => retry(None, format!("bad relay pool: {error}")), + } +} + +/// `relays` (WebTransport, `https`) and `websockets` (`wss`): either may be +/// absent or empty, not both. +fn parse_pool(body: &str) -> Result { + let value: serde_json::Value = + serde_json::from_str(body).map_err(|error| format!("invalid JSON: {error}"))?; + let list = |name: &str| -> Result, String> { + match value.get(name) { + None | Some(serde_json::Value::Null) => Ok(Vec::new()), + Some(serde_json::Value::Array(items)) => items + .iter() + .map(|item| { + item.as_str() + .map(str::to_owned) + .ok_or_else(|| format!("\"{name}\" entries must be URL strings")) + }) + .collect(), + Some(_) => Err(format!("\"{name}\" must be an array")), + } + }; + let pool = Pool { + webtransport: list("relays")? + .iter() + .map(|url| parse_relay(url, Carrier::WebTransport)) + .collect::>()?, + websocket: list("websockets")? + .iter() + .map(|url| parse_relay(url, Carrier::WebSocket)) + .collect::>()?, + }; + if pool.webtransport.is_empty() && pool.websocket.is_empty() { + return Err("the pool names no relay (\"relays\" and \"websockets\" are empty)".into()); + } + Ok(pool) +} + +fn parse_relay(value: &str, carrier: Carrier) -> Result { + let mut url = url::Url::parse(value).map_err(|error| format!("bad relay URL: {error}"))?; + let scheme = match carrier { + Carrier::WebTransport => "https", + Carrier::WebSocket => "wss", + }; + if url.scheme() != scheme { + return Err(format!( + "{carrier} relay URL scheme must be {scheme}, got {}", + url.scheme() + )); + } + if url.host_str().is_none() { + return Err("relay URL has no host".into()); + } + if !url.username().is_empty() || url.password().is_some() { + return Err("relay URL must not have userinfo".into()); + } + let label = label(&url); + // A malformed pin must be an error, never a silent fall-back to system + // roots — that would defeat the pinning. + let cert_hash = match url.fragment() { + None | Some("") => None, + Some(fragment) => { + let hash = fragment + .strip_prefix("sha256=") + .ok_or("relay URL fragment must be sha256=")?; + let bytes = + base64url_decode(hash).ok_or("relay certificate pin is not valid base64url")?; + if bytes.len() != 32 { + return Err("relay certificate pin must be a SHA-256 (32 bytes)".into()); + } + Some(bytes) + } + }; + // Fragments are client-side only; strip before connecting. + url.set_fragment(None); + Ok(Relay { + url, + label, + cert_hash, + carrier, + }) +} + +/// `host:port`: what may be shown of a URL that is a credential. +fn label(url: &url::Url) -> String { + let host = match url.host() { + Some(url::Host::Ipv6(address)) => format!("[{address}]"), + Some(host) => host.to_string(), + None => String::new(), + }; + format!("{host}:{}", url.port_or_known_default().unwrap_or(443)) +} + +fn base64url_decode(value: &str) -> Option> { + let value = value.trim_end_matches('='); + let mut out = Vec::with_capacity(value.len() * 3 / 4); + let mut acc: u32 = 0; + let mut bits = 0u32; + for byte in value.bytes() { + let digit = match byte { + b'A'..=b'Z' => byte - b'A', + b'a'..=b'z' => byte - b'a' + 26, + b'0'..=b'9' => byte - b'0' + 52, + b'-' => 62, + b'_' => 63, + _ => return None, + }; + acc = (acc << 6) | u32::from(digit); + bits += 6; + if bits >= 8 { + bits -= 8; + out.push((acc >> bits) as u8); + } + } + // Leftover bits must be padding zeros of a valid encoding. + if bits > 0 && (acc & ((1 << bits) - 1)) != 0 { + return None; + } + Some(out) +} + +fn jittered(base: Duration) -> Duration { + use rand::RngExt as _; + let ms = base.as_millis().max(1) as u64; + // 0.75x–1.25x + Duration::from_millis(rand::rng().random_range(ms * 3 / 4..=ms * 5 / 4)) +} + +// --------------------------------------------------------------------------- +// Relay sessions +// --------------------------------------------------------------------------- + +enum SessionEnd { + /// Handshake or CONNECT failed — try the next relay in the pool. + NeverConnected(String), + /// The session was established and later died — re-query the control + /// plane for a fresh pool. + Ended { reason: String, lasted: Duration }, +} + +/// A WebTransport client with the liveness settings of docs/uplink.md (10 s +/// keepalive, 30 s idle timeout), on YAS's rustls provider, verifying the +/// relay with the platform's roots or a certificate pin. +fn webtransport_client(cert_hash: Option<&[u8]>) -> Result { + let provider = yas_webrtc_forwarder::tls::provider(); + let builder = rustls::ClientConfig::builder_with_provider(provider.clone()) + .with_protocol_versions(&[&rustls::version::TLS13]) + .map_err(|error| format!("TLS config: {error}"))?; + let mut crypto = match cert_hash { + Some(hash) => builder + .dangerous() + .with_custom_certificate_verifier(Arc::new(crate::CertificateHash { + provider, + hash: hash.to_vec(), + })) + .with_no_client_auth(), + None => builder + .with_root_certificates(yas_webrtc_forwarder::tls::native_roots()) + .with_no_client_auth(), + }; + crypto.alpn_protocols = vec![wt::ALPN.as_bytes().to_vec()]; + let crypto = wt::quinn::crypto::rustls::QuicClientConfig::try_from(crypto) + .map_err(|error| format!("QUIC TLS config: {error}"))?; + let mut config = wt::quinn::ClientConfig::new(Arc::new(crypto)); + let mut transport = wt::quinn::TransportConfig::default(); + transport.keep_alive_interval(Some(KEEPALIVE)); + transport.max_idle_timeout(Some( + wt::quinn::IdleTimeout::try_from(IDLE_TIMEOUT).expect("30s fits in an idle timeout"), + )); + config.transport_config(Arc::new(transport)); + let endpoint = wt::quinn::Endpoint::client((std::net::Ipv6Addr::UNSPECIFIED, 0).into()) + .or_else(|_| wt::quinn::Endpoint::client((std::net::Ipv4Addr::UNSPECIFIED, 0).into())) + .map_err(|error| format!("UDP socket: {error}"))?; + Ok(wt::Client::new(endpoint, config)) +} + +/// A WebSocket to `url` (`wss`, a session's or a stream's) speaking +/// [`WEBSOCKET_SUBPROTOCOL`], TLS verified by the platform's roots or a pin. +async fn connect_websocket(url: &url::Url, cert_hash: Option<&[u8]>) -> Result { + use tokio_tungstenite::tungstenite::client::IntoClientRequest; + let mut request = url + .as_str() + .into_client_request() + .map_err(|_| "bad WebSocket URL".to_owned())?; + request.headers_mut().insert( + "sec-websocket-protocol", + tokio_tungstenite::tungstenite::http::HeaderValue::from_static(WEBSOCKET_SUBPROTOCOL), + ); + let connector = match cert_hash { + Some(hash) => { + let provider = yas_webrtc_forwarder::tls::provider(); + let config = rustls::ClientConfig::builder_with_provider(provider.clone()) + .with_safe_default_protocol_versions() + .map_err(|error| format!("TLS config: {error}"))? + .dangerous() + .with_custom_certificate_verifier(Arc::new(crate::CertificateHash { + provider, + hash: hash.to_vec(), + })) + .with_no_client_auth(); + tokio_tungstenite::Connector::Rustls(Arc::new(config)) + } + None => yas_webrtc_forwarder::tls::websocket_connector(), + }; + let config = tokio_tungstenite::tungstenite::protocol::WebSocketConfig::default() + .max_message_size(Some(MAX_WEBSOCKET_MESSAGE)) + .max_frame_size(Some(MAX_WEBSOCKET_MESSAGE)); + let connecting = tokio_tungstenite::connect_async_tls_with_config( + request, + Some(config), + true, + Some(connector), + ); + // Errors may quote the URL, a credential: say what failed, not where. + let (socket, response) = match tokio::time::timeout(WEBSOCKET_CONNECT, connecting).await { + Err(_) => { + return Err(format!( + "no WebSocket within {}s", + WEBSOCKET_CONNECT.as_secs() + )); + } + Ok(Err(tokio_tungstenite::tungstenite::Error::Http(response))) => { + return Err(format!("the relay answered HTTP {}", response.status())); + } + // tungstenite checks the answer's subprotocol against the one offered. + Ok(Err(tokio_tungstenite::tungstenite::Error::Protocol( + tokio_tungstenite::tungstenite::error::ProtocolError::SecWebSocketSubProtocolError(_), + ))) => { + return Err(format!("the relay doesn't speak {WEBSOCKET_SUBPROTOCOL}")); + } + Ok(Err(error)) => return Err(websocket_error(&error)), + Ok(Ok(connected)) => connected, + }; + let selected = response + .headers() + .get("sec-websocket-protocol") + .and_then(|value| value.to_str().ok()); + if selected != Some(WEBSOCKET_SUBPROTOCOL) { + return Err(format!("the relay doesn't speak {WEBSOCKET_SUBPROTOCOL}")); + } + Ok(socket) +} + +fn websocket_error(error: &tokio_tungstenite::tungstenite::Error) -> String { + use tokio_tungstenite::tungstenite::Error; + match error { + Error::Io(error) => format!("connect failed: {error}"), + Error::Tls(error) => format!("TLS: {error}"), + Error::Url(_) => "bad WebSocket URL".to_owned(), + Error::Protocol(error) => format!("WebSocket handshake: {error}"), + other => format!("WebSocket: {other}"), + } +} + +/// The stream URL of a relay's request (`{"stream": "wss://…"}`), when it is +/// one: a `wss` URL on the session's own origin (so the same TLS trust +/// applies), without userinfo or a fragment. Anything else is ignored, for +/// relays newer than this producer. +fn stream_request(text: &str, session: &url::Url) -> Option { + let value: serde_json::Value = serde_json::from_str(text).ok()?; + let stream = url::Url::parse(value.get("stream")?.as_str()?).ok()?; + (stream.scheme() == "wss" + && stream.host() == session.host() + && stream.port_or_known_default() == session.port_or_known_default() + && stream.username().is_empty() + && stream.password().is_none() + && stream.fragment().is_none()) + .then_some(stream) +} + +/// One consumer on a WebSocket session: open the stream the relay asked for, +/// then serve it as a WebTransport stream is served, without datagrams. +async fn websocket_stream( + url: url::Url, + cert_hash: Option>, + shared: Arc, + permit: tokio::sync::OwnedSemaphorePermit, + local: Local, + events: Events, +) { + let socket = match connect_websocket(&url, cert_hash.as_deref()).await { + Ok(socket) => socket, + Err(error) => { + events(&Event::StreamFailed { error }); + return; + } + }; + let (sink, source) = socket.split(); + // Opaque chunks of Noise records both ways. A stream's write side + // finishing doesn't close the WebSocket (Noise carries the half-close). + let relay = tokio::io::join( + crate::WsFrameReader { + inner: source, + buf: bytes::Bytes::new(), + frame_lengths: false, + }, + crate::WsByteWriter { inner: sink }, + ); + bridge(relay, Lane::None, shared, permit, local, events).await; +} + +async fn distribute_datagrams(session: wt::Session, routes: DatagramRoutes) { + while let Ok(bytes) = session.read_datagram().await { + let _ = routes.route(&bytes); + } + routes.clear(); +} + +// --------------------------------------------------------------------------- +// Bridging a consumer to the local server +// --------------------------------------------------------------------------- + +/// Where a consumer's datagrams go: a WebTransport session's datagrams, or +/// nowhere (WebSocket sessions carry none). +enum Lane { + WebTransport { + session: Box, + routes: DatagramRoutes, + }, + None, +} + +/// Bridge one consumer stream to one local YAS server connection. +/// Authentication precedes both ingress classification and local IPC. The +/// relay sees only Noise records; plaintext selectors cannot bypass admission. +async fn bridge( + relay: S, + lane: Lane, + shared: Arc, + permit: tokio::sync::OwnedSemaphorePermit, + local: Local, + events: Events, +) where + S: AsyncRead + AsyncWrite + Unpin, +{ + let relay = match yas_uplink::accept(relay, shared.crypto()).await { + Ok(relay) => relay, + Err(_) => return, + }; + // Bind each sideband's keys to its authenticated main stream. The random + // route token remains authenticated as AEAD AAD on every datagram. + let material = relay.datagram_key_material(); + let ingress = tokio::time::timeout( + Duration::from_secs(5), + yas_composite_transport::classify(relay), + ) + .await; + drop(permit); + match ingress { + Ok(Ok(yas_composite_transport::Ingress::Direct(relay))) => { + bridge_direct(relay, &local, &events).await + } + Ok(Ok(yas_composite_transport::Ingress::Composite { offer, stream })) + if offer.role == yas_composite_transport::Role::Main => + { + let Lane::WebTransport { session, routes } = lane else { + events(&Event::StreamFailed { + error: "rejected a datagram lane: a WebSocket session carries no datagrams" + .into(), + }); + return; + }; + let (sender, receiver) = yas_uplink::datagram_pair(material, offer.token, false); + let lane = Sideband { + session: *session, + routes, + sender, + receiver, + }; + bridge_composite(lane, offer, stream, &local, &events).await; + } + _ => {} + } +} + +async fn bridge_direct(relay: S, local: &Local, events: &Events) +where + S: AsyncRead + AsyncWrite + Unpin, +{ + let (mut sock_read, mut sock_write) = match local.connect().await { + Ok(halves) => halves, + Err(error) => { + events(&Event::LocalUnavailable { + local: local.label.clone(), + error, + }); + return; + } + }; + let (mut recv, mut send) = tokio::io::split(relay); + let down = async move { + tokio::io::copy(&mut recv, &mut sock_write).await?; + sock_write.shutdown().await + }; + let up = async move { + tokio::io::copy(&mut sock_read, &mut send).await?; + send.shutdown().await + }; + let _ = tokio::try_join!(down, up); +} + +/// A consumer's datagram lane on a WebTransport session. +struct Sideband { + session: wt::Session, + routes: DatagramRoutes, + sender: yas_uplink::DatagramSender, + receiver: yas_uplink::DatagramReceiver, +} + +async fn bridge_composite( + lane: Sideband, + offer: yas_composite_transport::Offer, + relay: S, + local: &Local, + events: &Events, +) where + S: AsyncRead + AsyncWrite + Unpin, +{ + let Sideband { + session, + routes, + mut sender, + mut receiver, + } = lane; + let physical_maximum = session + .max_datagram_size() + .saturating_sub(yas_composite_transport::ROUTED_DATAGRAM_HEADER) + .saturating_sub(yas_uplink::DATAGRAM_OVERHEAD) + .min(yas_composite_transport::HARD_MAX_DATAGRAM as usize - yas_uplink::DATAGRAM_OVERHEAD); + if offer.max_datagram as usize > physical_maximum { + events(&Event::StreamFailed { + error: format!( + "rejected composite maximum {} above path maximum {physical_maximum}", + offer.max_datagram + ), + }); + return; + } + let (mut main_read, mut main_write) = match local.connect().await { + Ok(halves) => halves, + Err(error) => { + events(&Event::LocalUnavailable { + local: local.label.clone(), + error, + }); + return; + } + }; + let (mut side_read, mut side_write) = match local.connect().await { + Ok(halves) => halves, + Err(error) => { + events(&Event::LocalUnavailable { + local: format!("{} (datagram sideband)", local.label), + error, + }); + return; + } + }; + let main_offer = yas_composite_transport::Offer::new( + yas_composite_transport::Role::Main, + offer.token, + offer.max_datagram, + ) + .expect("the classified offer is valid"); + let side_offer = yas_composite_transport::Offer::new( + yas_composite_transport::Role::Datagram, + offer.token, + offer.max_datagram, + ) + .expect("the classified offer is valid"); + if yas_composite_transport::write_offer(&mut main_write, main_offer) + .await + .is_err() + || yas_composite_transport::write_offer(&mut side_write, side_offer) + .await + .is_err() + { + return; + } + + let encrypted_maximum = offer.max_datagram + yas_uplink::DATAGRAM_OVERHEAD as u32; + let Ok(mut route_rx) = routes.register(offer.token, encrypted_maximum, DATAGRAM_QUEUE) else { + return; + }; + + let (mut relay_read, mut relay_write) = tokio::io::split(relay); + let down = async { + tokio::io::copy(&mut relay_read, &mut main_write).await?; + main_write.shutdown().await + }; + let up = async { + tokio::io::copy(&mut main_read, &mut relay_write).await?; + relay_write.shutdown().await + }; + let side_routes = routes.clone(); + let side_token = offer.token; + let side_task = tokio::spawn(async move { + let side_out = async { + loop { + let Ok(frame) = + yas_composite_transport::read_datagram(&mut side_read, offer.max_datagram) + .await + else { + break; + }; + let Some(encrypted) = sender.seal(&frame) else { + break; + }; + let Ok(routed) = yas_composite_transport::encode_routed_datagram( + offer.token, + &encrypted, + encrypted_maximum, + ) else { + break; + }; + // Congestion is ordinary datagram loss; never await reliable + // transport capacity or fall back to the main stream here. + let _ = session.send_datagram(routed.into()); + } + }; + let side_in = async { + while let Some(frame) = route_rx.recv().await { + let Some(frame) = receiver.open(&frame) else { + continue; + }; + if yas_composite_transport::write_datagram( + &mut side_write, + &frame, + offer.max_datagram, + ) + .await + .is_err() + { + break; + } + } + }; + tokio::select! { + _ = side_out => {} + _ = side_in => {} + } + side_routes.remove(side_token); + }); + + // The authoritative reliable splice owns the connection lifetime. An + // optional sideband ending only removes its route and cannot drop `main`. + let _ = tokio::try_join!(down, up); + routes.remove(offer.token); + side_task.abort(); +} + +#[cfg(test)] +#[path = "uplink_producer_tests.rs"] +mod tests; diff --git a/crates/proxy/src/uplink_producer_tests.rs b/crates/proxy/src/uplink_producer_tests.rs new file mode 100644 index 00000000..35550e69 --- /dev/null +++ b/crates/proxy/src/uplink_producer_tests.rs @@ -0,0 +1,915 @@ +// The WebSocket and local-socket tests are Unix-only: on Windows the helpers +// only they use would fail CI's `-D warnings`. +#![cfg_attr(not(unix), allow(dead_code, unused_imports))] + +use super::*; +use tokio::io::AsyncReadExt; + +fn events() -> (Events, Arc>>) { + let seen = Arc::new(Mutex::new(Vec::new())); + let log = seen.clone(); + let events: Events = Arc::new(move |event: &Event| { + log.lock().unwrap().push(event.to_string()); + }); + (events, seen) +} + +fn shared(crypto: Arc) -> Arc { + Arc::new(Shared { + token: Mutex::new("token".into()), + crypto: Mutex::new(crypto), + }) +} + +struct Keys { + server: Arc, + server_public: yas_uplink::PublicKey, + server_identity: yas_uplink::Identity, + client: Arc, + client_public: yas_uplink::PublicKey, + other: Arc, + other_public: yas_uplink::PublicKey, +} + +fn keys() -> Keys { + let (server_key, server_public) = yas_uplink::Identity::generate().unwrap(); + let (client_key, client_public) = yas_uplink::Identity::generate().unwrap(); + let (other_key, other_public) = yas_uplink::Identity::generate().unwrap(); + let server_identity = yas_uplink::Identity::from_base64(&server_key).unwrap(); + Keys { + server: server_identity.server_config(vec![client_public]).unwrap(), + server_public, + server_identity, + client: yas_uplink::Identity::from_base64(&client_key) + .unwrap() + .client_config(server_public) + .unwrap(), + client_public, + other: yas_uplink::Identity::from_base64(&other_key) + .unwrap() + .client_config(server_public) + .unwrap(), + other_public, + } +} + +#[cfg(unix)] +fn local_socket(dir: &std::path::Path) -> (Local, tokio::net::UnixListener) { + let socket = dir.join("local.sock"); + let listener = tokio::net::UnixListener::bind(&socket).unwrap(); + (Local::socket(socket.to_str().unwrap()), listener) +} + +#[cfg(unix)] +async fn nothing_accepted(listener: &tokio::net::UnixListener) -> bool { + tokio::time::timeout(Duration::from_millis(25), listener.accept()) + .await + .is_err() +} + +#[cfg(unix)] +#[tokio::test] +async fn producer_authenticates_before_ipc_and_encrypts_datagrams() { + use yas_composite_transport::{Offer, Role}; + yas_webrtc_forwarder::tls::install_default_provider(); + tokio::time::timeout(Duration::from_secs(15), async { + let dir = tempfile::tempdir().unwrap(); + let keys = keys(); + let (local, listener) = local_socket(dir.path()); + let (events, _) = events(); + + let cert = rcgen::generate_simple_self_signed(vec!["localhost".into()]).unwrap(); + let hash = wt::crypto::sha256(&yas_webrtc_forwarder::tls::provider(), cert.cert.der()); + let mut worker = wt::ServerBuilder::new() + .with_addr("127.0.0.1:0".parse().unwrap()) + .with_certificate( + vec![cert.cert.der().clone()], + rustls::pki_types::PrivatePkcs8KeyDer::from(cert.signing_key.serialize_der()) + .into(), + ) + .unwrap(); + let outer = webtransport_client(Some(hash.as_ref())).unwrap(); + let url: url::Url = format!("https://127.0.0.1:{}/", worker.local_addr().unwrap().port()) + .parse() + .unwrap(); + let (connected, remote) = tokio::join!(outer.connect(url), async { + worker.accept().await.unwrap().ok().await.unwrap() + }); + let session = connected.unwrap(); + let routes = DatagramRoutes::new(16); + let datagrams = tokio::spawn(distribute_datagrams(session.clone(), routes.clone())); + + for case in 0..4 { + let lane = Lane::WebTransport { + session: Box::new(session.clone()), + routes: routes.clone(), + }; + let shared = shared(keys.server.clone()); + let (local, events) = (local.clone(), events.clone()); + let accepting = session.clone(); + let producer = tokio::spawn(async move { + let (send, recv) = accepting.accept_bi().await.unwrap(); + let permit = Arc::new(tokio::sync::Semaphore::new(1)) + .acquire_owned() + .await + .unwrap(); + bridge( + tokio::io::join(recv, send), + lane, + shared, + permit, + local, + events, + ) + .await; + }); + let (send, recv) = remote.open_bi().await.unwrap(); + let mut stream = tokio::io::join(recv, send); + if case == 0 { + // A relay can synthesize protocol bytes, but those bytes must + // never open a local socket before Noise authentication. + stream + .write_all(yas_wire::PREFACE.as_slice()) + .await + .unwrap(); + stream.shutdown().await.unwrap(); + drop(stream); + producer.await.unwrap(); + assert!(nothing_accepted(&listener).await); + continue; + } + if case == 1 { + assert!( + yas_uplink::connect(stream, keys.other.clone()) + .await + .is_err() + ); + producer.await.unwrap(); + assert!(nothing_accepted(&listener).await); + continue; + } + let mut stream = yas_uplink::connect(stream, keys.client.clone()) + .await + .unwrap(); + // Authentication alone also does not open IPC: the encrypted YAS + // preface or composite selector must follow first. + assert!(nothing_accepted(&listener).await); + if case == 2 { + stream + .write_all(yas_wire::PREFACE.as_slice()) + .await + .unwrap(); + stream.write_all(b"command").await.unwrap(); + stream.flush().await.unwrap(); + let (mut ipc, _) = listener.accept().await.unwrap(); + let mut received = vec![0; yas_wire::PREFACE.len() + 7]; + ipc.read_exact(&mut received).await.unwrap(); + assert_eq!( + received, + [yas_wire::PREFACE.as_slice(), b"command"].concat() + ); + ipc.write_all(b"reply").await.unwrap(); + ipc.shutdown().await.unwrap(); + let mut reply = Vec::new(); + stream.read_to_end(&mut reply).await.unwrap(); + assert_eq!(reply, b"reply"); + stream.shutdown().await.unwrap(); + producer.await.unwrap(); + continue; + } + + let token = [0x35; 16]; + let maximum = 512; + let material = stream.datagram_key_material(); + let (mut sender, mut receiver) = yas_uplink::datagram_pair(material, token, true); + yas_composite_transport::write_offer( + &mut stream, + Offer::new(Role::Main, token, maximum).unwrap(), + ) + .await + .unwrap(); + stream.flush().await.unwrap(); + let (mut main, _) = listener.accept().await.unwrap(); + let (mut side, _) = listener.accept().await.unwrap(); + for (socket, expected) in [(&mut main, Role::Main), (&mut side, Role::Datagram)] { + match yas_composite_transport::classify(socket).await.unwrap() { + yas_composite_transport::Ingress::Composite { offer, .. } => { + assert_eq!(offer.role, expected); + assert_eq!(offer.token, token); + assert_eq!(offer.max_datagram, maximum); + } + _ => panic!("missing local composite selector"), + } + } + // A reliable reply also synchronizes route registration. + main.write_all(b"ready").await.unwrap(); + let mut ready = [0; 5]; + stream.read_exact(&mut ready).await.unwrap(); + assert_eq!(&ready, b"ready"); + + let encrypted = sender.seal(b"private datagram").unwrap(); + let wire_maximum = maximum + yas_uplink::DATAGRAM_OVERHEAD as u32; + let mut forged = encrypted.clone(); + *forged.last_mut().unwrap() ^= 1; + for bytes in [&forged, &encrypted, &encrypted] { + let routed = + yas_composite_transport::encode_routed_datagram(token, bytes, wire_maximum) + .unwrap(); + remote.send_datagram(routed.into()).unwrap(); + } + assert_eq!( + yas_composite_transport::read_datagram(&mut side, maximum) + .await + .unwrap(), + b"private datagram" + ); + assert!( + tokio::time::timeout( + Duration::from_millis(50), + yas_composite_transport::read_datagram(&mut side, maximum) + ) + .await + .is_err() + ); + + yas_composite_transport::write_datagram(&mut side, b"private reply", maximum) + .await + .unwrap(); + let packet = remote.read_datagram().await.unwrap(); + let (received_token, ciphertext) = + yas_composite_transport::split_routed_datagram(&packet).unwrap(); + assert_eq!(received_token, token); + assert!(!ciphertext.windows(13).any(|w| w == b"private reply")); + assert_eq!(receiver.open(ciphertext).unwrap(), b"private reply"); + stream.shutdown().await.unwrap(); + main.shutdown().await.unwrap(); + producer.await.unwrap(); + } + datagrams.abort(); + session.close(0, b"test complete"); + }) + .await + .expect("producer uplink test stalled"); +} + +/// A relay for WebSocket sessions, pinned by its certificate's hash: it +/// takes one session, then answers `/stream/…` WebSockets, handing each to the +/// test as a byte stream the consumer's side of Noise runs on. +#[cfg(unix)] +struct WebSocketRelay { + url: url::Url, + hash: Vec, + /// Stream requests sent on the session: what the test writes there. + requests: tokio::sync::mpsc::UnboundedSender, + /// Stream WebSockets the producer opened, as duplex byte streams. + streams: tokio::sync::mpsc::UnboundedReceiver<(String, tokio::io::DuplexStream)>, + /// What the relay offered as the session's subprotocol, and saw. + session_protocols: Arc>>, + _task: tokio::task::JoinHandle<()>, +} + +#[cfg(unix)] +impl WebSocketRelay { + async fn start() -> Self { + use tokio_tungstenite::tungstenite::handshake::server::{Request, Response}; + yas_webrtc_forwarder::tls::install_default_provider(); + let cert = rcgen::generate_simple_self_signed(vec!["127.0.0.1".into()]).unwrap(); + let hash = wt::crypto::sha256(&yas_webrtc_forwarder::tls::provider(), cert.cert.der()) + .as_ref() + .to_vec(); + let tls = tokio_rustls::TlsAcceptor::from(Arc::new( + rustls::ServerConfig::builder() + .with_no_client_auth() + .with_single_cert( + vec![cert.cert.der().clone()], + rustls::pki_types::PrivatePkcs8KeyDer::from(cert.signing_key.serialize_der()) + .into(), + ) + .unwrap(), + )); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let url: url::Url = format!( + "wss://127.0.0.1:{}/session", + listener.local_addr().unwrap().port() + ) + .parse() + .unwrap(); + let (requests, mut pending) = tokio::sync::mpsc::unbounded_channel::(); + let (streams_tx, streams) = tokio::sync::mpsc::unbounded_channel(); + let session_protocols = Arc::new(Mutex::new(Vec::new())); + let protocols = session_protocols.clone(); + let task = tokio::spawn(async move { + let mut session_taken = false; + loop { + let (tcp, _) = listener.accept().await.unwrap(); + let Ok(tls) = tls.accept(tcp).await else { + continue; + }; + let path = Arc::new(Mutex::new(String::new())); + let seen = path.clone(); + let offered = protocols.clone(); + // tungstenite's callback signature: its error is an HTTP response. + #[allow(clippy::result_large_err)] + let callback = move |request: &Request, mut response: Response| { + *seen.lock().unwrap() = request.uri().path().to_owned(); + let protocol = request + .headers() + .get("sec-websocket-protocol") + .and_then(|value| value.to_str().ok()) + .unwrap_or_default() + .to_owned(); + offered.lock().unwrap().push(protocol.clone()); + if protocol + .split(',') + .any(|each| each.trim() == WEBSOCKET_SUBPROTOCOL) + { + response.headers_mut().insert( + "sec-websocket-protocol", + WEBSOCKET_SUBPROTOCOL.parse().unwrap(), + ); + } + Ok(response) + }; + let Ok(socket) = tokio_tungstenite::accept_hdr_async(tls, callback).await else { + continue; + }; + let path = path.lock().unwrap().clone(); + if path == "/session" && !session_taken { + session_taken = true; + let (mut sink, mut source) = socket.split(); + let mut requests = + std::mem::replace(&mut pending, tokio::sync::mpsc::unbounded_channel().1); + tokio::spawn(async move { + loop { + tokio::select! { + request = requests.recv() => match request { + Some(message) => { + if sink.send(message).await.is_err() { + break; + } + } + None => break, + }, + message = source.next() => match message { + Some(Ok(_)) => {} + _ => break, + }, + } + } + }); + } else if let Some(stream) = path.strip_prefix("/stream/") { + let (ours, theirs) = tokio::io::duplex(128 << 10); + tokio::spawn(pump(socket, ours)); + let _ = streams_tx.send((stream.to_owned(), theirs)); + } + } + }); + Self { + url, + hash, + requests, + streams, + session_protocols, + _task: task, + } + } + + fn relay(&self) -> Relay { + Relay { + url: self.url.clone(), + label: label(&self.url), + cert_hash: Some(self.hash.clone()), + carrier: Carrier::WebSocket, + } + } + + fn ask(&self, stream: &str) { + let url = format!( + "wss://127.0.0.1:{}/stream/{stream}", + self.url.port().unwrap() + ); + self.requests + .send(Message::text( + serde_json::json!({ "stream": url }).to_string(), + )) + .unwrap(); + } + + async fn stream(&mut self) -> (String, tokio::io::DuplexStream) { + tokio::time::timeout(Duration::from_secs(5), self.streams.recv()) + .await + .expect("the producer opened no stream WebSocket") + .unwrap() + } +} + +/// Bytes between a stream WebSocket and a duplex pipe, as a relay forwards them. +#[cfg(unix)] +async fn pump(socket: tokio_tungstenite::WebSocketStream, pipe: tokio::io::DuplexStream) +where + S: AsyncRead + AsyncWrite + Unpin, +{ + let (mut sink, mut source) = socket.split(); + let (mut read, mut write) = tokio::io::split(pipe); + let up = async { + while let Some(Ok(message)) = source.next().await { + match message { + Message::Binary(bytes) => { + if write.write_all(&bytes).await.is_err() { + break; + } + } + Message::Close(_) => break, + _ => {} + } + } + let _ = write.shutdown().await; + }; + let down = async { + let mut buffer = vec![0; 16 << 10]; + loop { + match read.read(&mut buffer).await { + Ok(0) | Err(_) => break, + Ok(count) => { + let chunk = bytes::Bytes::copy_from_slice(&buffer[..count]); + if sink.send(Message::Binary(chunk)).await.is_err() { + break; + } + } + } + } + let _ = sink.close().await; + }; + tokio::join!(up, down); +} + +#[cfg(unix)] +#[tokio::test] +async fn websocket_session_serves_consumers_and_takes_allowlist_changes() { + use yas_composite_transport::{Offer, Role}; + tokio::time::timeout(Duration::from_secs(20), async { + let dir = tempfile::tempdir().unwrap(); + let keys = keys(); + let (local, listener) = local_socket(dir.path()); + let mut relay = WebSocketRelay::start().await; + let producer = Producer::new( + "https://127.0.0.1:1/control", + "token", + keys.server.clone(), + local, + ) + .unwrap(); + let (events, seen) = events(); + let producer = Producer { events, ..producer }; + let handle = producer.handle(); + let active = Active::default(); + let session = { + let (producer, active, relay) = (producer.clone(), active.clone(), relay.relay()); + tokio::spawn(async move { producer.websocket_session(&relay, &active).await }) + }; + // The session is up once it reports it. + for _ in 0..200 { + if seen + .lock() + .unwrap() + .iter() + .any(|line| line.starts_with("connected")) + { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + assert_eq!( + seen.lock().unwrap()[0], + format!( + "connected to relay 127.0.0.1:{} over WebSocket", + relay.url.port().unwrap() + ) + ); + assert_eq!( + relay.session_protocols.lock().unwrap()[0], + WEBSOCKET_SUBPROTOCOL + ); + + // A stream elsewhere than the session's origin is ignored. + relay + .requests + .send(Message::text( + r#"{"stream":"wss://127.0.0.2:1/stream/elsewhere"}"#.to_owned(), + )) + .unwrap(); + relay + .requests + .send(Message::text("{\"hello\":1}".to_owned())) + .unwrap(); + + // A consumer: authenticated, then its bytes reach the local server. + relay.ask("first"); + let (name, stream) = relay.stream().await; + assert_eq!(name, "first"); + assert_eq!( + relay.session_protocols.lock().unwrap()[1], + WEBSOCKET_SUBPROTOCOL + ); + assert!(nothing_accepted(&listener).await); + let mut first = yas_uplink::connect(stream, keys.client.clone()) + .await + .unwrap(); + first.write_all(yas_wire::PREFACE.as_slice()).await.unwrap(); + first.write_all(b"command").await.unwrap(); + first.flush().await.unwrap(); + let (mut ipc, _) = listener.accept().await.unwrap(); + let mut received = vec![0; yas_wire::PREFACE.len() + 7]; + ipc.read_exact(&mut received).await.unwrap(); + assert_eq!( + received, + [yas_wire::PREFACE.as_slice(), b"command"].concat() + ); + + // A key it doesn't allow is refused before the local server hears of it. + relay.ask("stranger"); + let (_, stream) = relay.stream().await; + assert!( + yas_uplink::connect(stream, keys.other.clone()) + .await + .is_err() + ); + assert!(nothing_accepted(&listener).await); + + // The allowlist changes while it runs: the other key only, now. + handle.set_server_config( + keys.server_identity + .server_config(vec![keys.other_public]) + .unwrap(), + ); + relay.ask("revoked"); + let (_, stream) = relay.stream().await; + assert!( + yas_uplink::connect(stream, keys.client.clone()) + .await + .is_err() + ); + relay.ask("added"); + let (_, stream) = relay.stream().await; + let mut added = yas_uplink::connect(stream, keys.other.clone()) + .await + .unwrap(); + added.write_all(yas_wire::PREFACE.as_slice()).await.unwrap(); + added.flush().await.unwrap(); + let (mut second, _) = listener.accept().await.unwrap(); + let mut preface = vec![0; yas_wire::PREFACE.len()]; + second.read_exact(&mut preface).await.unwrap(); + assert_eq!(preface, yas_wire::PREFACE.as_slice()); + drop(second); + drop(added); + let _ = keys.client_public; + + // The consumer connected before keeps its session, both ways, + // including the local server finishing first (Noise's half-close). + ipc.write_all(b"reply").await.unwrap(); + ipc.shutdown().await.unwrap(); + let mut reply = [0; 5]; + first.read_exact(&mut reply).await.unwrap(); + assert_eq!(&reply, b"reply"); + first.write_all(b" and more").await.unwrap(); + first.flush().await.unwrap(); + let mut more = [0; 9]; + ipc.read_exact(&mut more).await.unwrap(); + assert_eq!(&more, b" and more"); + first.shutdown().await.unwrap(); + drop(ipc); + + // A datagram lane can't be carried: refused, nothing opened locally. + relay.ask("composite"); + let (_, stream) = relay.stream().await; + let mut composite = yas_uplink::connect(stream, keys.other.clone()) + .await + .unwrap(); + yas_composite_transport::write_offer( + &mut composite, + Offer::new(Role::Main, [7; 16], 512).unwrap(), + ) + .await + .unwrap(); + composite.flush().await.unwrap(); + for _ in 0..200 { + if seen + .lock() + .unwrap() + .iter() + .any(|line| line.contains("datagram lane")) + { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + assert!(nothing_accepted(&listener).await); + assert!( + seen.lock() + .unwrap() + .iter() + .any(|line| line.contains("a WebSocket session carries no datagrams")) + ); + + // A consumer still connected when the session ends goes with it, as + // a WebTransport stream goes with its connection. + relay.ask("lingering"); + let (_, stream) = relay.stream().await; + let mut lingering = yas_uplink::connect(stream, keys.other.clone()) + .await + .unwrap(); + lingering + .write_all(yas_wire::PREFACE.as_slice()) + .await + .unwrap(); + lingering.flush().await.unwrap(); + let (mut held, _) = listener.accept().await.unwrap(); + let mut preface = vec![0; yas_wire::PREFACE.len()]; + held.read_exact(&mut preface).await.unwrap(); + + // Shutting down closes the session with a close frame. + active.close().await; + match tokio::time::timeout(Duration::from_secs(5), session).await { + Ok(Ok(SessionEnd::Ended { reason, .. })) => { + assert!(!reason.is_empty()); + } + _ => panic!("the session didn't end"), + } + let mut rest = Vec::new(); + assert!( + tokio::time::timeout(Duration::from_secs(5), held.read_to_end(&mut rest)) + .await + .is_ok(), + "the lingering consumer's local connection closes with the session" + ); + drop(lingering); + let _ = keys.server_public; + }) + .await + .expect("WebSocket session test stalled"); +} + +#[tokio::test] +async fn websocket_relay_must_speak_the_subprotocol() { + // A WebSocket server that selects no subprotocol is not an uplink relay. + yas_webrtc_forwarder::tls::install_default_provider(); + let cert = rcgen::generate_simple_self_signed(vec!["127.0.0.1".into()]).unwrap(); + let hash = wt::crypto::sha256(&yas_webrtc_forwarder::tls::provider(), cert.cert.der()) + .as_ref() + .to_vec(); + let tls = tokio_rustls::TlsAcceptor::from(Arc::new( + rustls::ServerConfig::builder() + .with_no_client_auth() + .with_single_cert( + vec![cert.cert.der().clone()], + rustls::pki_types::PrivatePkcs8KeyDer::from(cert.signing_key.serialize_der()) + .into(), + ) + .unwrap(), + )); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let port = listener.local_addr().unwrap().port(); + tokio::spawn(async move { + let (tcp, _) = listener.accept().await.unwrap(); + let tls = tls.accept(tcp).await.unwrap(); + let _ = tokio_tungstenite::accept_async(tls).await; + }); + let url: url::Url = format!("wss://127.0.0.1:{port}/session").parse().unwrap(); + let error = connect_websocket(&url, Some(&hash)).await.err().unwrap(); + assert!(error.contains("doesn't speak yas-uplink.v1"), "{error}"); + // A wrong pin fails TLS, and the error never quotes the URL. + let error = connect_websocket(&url, Some(&[0; 32])).await.err().unwrap(); + assert!(!error.contains("/session"), "{error}"); +} + +#[tokio::test] +async fn webtransport_gives_up_quickly_where_udp_goes_nowhere() { + yas_webrtc_forwarder::tls::install_default_provider(); + // A UDP socket that never answers: what a network dropping UDP looks like. + let hole = tokio::net::UdpSocket::bind("127.0.0.1:0").await.unwrap(); + let url: url::Url = format!( + "https://127.0.0.1:{}/t/x", + hole.local_addr().unwrap().port() + ) + .parse() + .unwrap(); + let relay = Relay { + label: label(&url), + url, + cert_hash: Some(vec![0; 32]), + carrier: Carrier::WebTransport, + }; + let keys = keys(); + let dir = tempfile::tempdir().unwrap(); + let producer = Producer::new( + "https://127.0.0.1:1/control", + "token", + keys.server, + Local::socket(dir.path().join("none").to_str().unwrap()), + ) + .unwrap(); + let started = Instant::now(); + let end = producer + .webtransport_session(&relay, Some(Duration::from_millis(300)), &Active::default()) + .await; + match end { + SessionEnd::NeverConnected(error) => assert!(error.contains("is UDP blocked?"), "{error}"), + SessionEnd::Ended { .. } => panic!("a session through a black hole"), + } + assert!(started.elapsed() < Duration::from_secs(5)); +} + +#[test] +fn parse_pool_accepts_plain_and_pinned_relay_urls() { + let pool = parse_pool( + r#"{"relays":[ + "https://relay-1.indent.com:4443/t/kfV3aB", + "https://[2001:db8::7]/session?key=x#sha256=AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA" + ],"ttl":60}"#, + ) + .unwrap(); + assert_eq!(pool.webtransport.len(), 2); + assert!(pool.websocket.is_empty()); + let relays = &pool.webtransport; + assert_eq!( + relays[0].url.as_str(), + "https://relay-1.indent.com:4443/t/kfV3aB" + ); + assert_eq!(relays[0].label, "relay-1.indent.com:4443"); + assert!(relays[0].cert_hash.is_none()); + assert_eq!(relays[1].label, "[2001:db8::7]:443"); + assert_eq!(relays[1].cert_hash.as_ref().unwrap().len(), 32); + // The pin must be stripped before the URL is used to connect. + assert_eq!(relays[1].url.fragment(), None); + assert_eq!(relays[1].url.query(), Some("key=x")); +} + +#[test] +fn parse_pool_takes_websocket_relays() { + let pool = parse_pool( + r#"{"relays":["https://relay.example:4433/t/a"], + "websockets":["wss://relay.example/uplink/producer/b", + "wss://relay.example:8443/p#sha256=AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"]}"#, + ) + .unwrap(); + assert_eq!(pool.webtransport.len(), 1); + assert_eq!(pool.websocket.len(), 2); + assert_eq!(pool.websocket[0].label, "relay.example:443"); + assert_eq!(pool.websocket[0].carrier, Carrier::WebSocket); + assert_eq!(pool.websocket[1].cert_hash.as_ref().unwrap().len(), 32); + assert_eq!(pool.websocket[1].url.fragment(), None); + // WebSocket relays alone make a pool; so does an empty WebTransport list. + let pool = parse_pool(r#"{"relays":[],"websockets":["wss://relay.example/p"]}"#).unwrap(); + assert!(pool.webtransport.is_empty()); + let pool = parse_pool(r#"{"websockets":["wss://relay.example/p"]}"#).unwrap(); + assert_eq!(pool.websocket.len(), 1); +} + +#[test] +fn parse_pool_rejects_bad_input() { + assert!(parse_pool("not json").is_err()); + assert!(parse_pool(r#"{"relays":[]}"#).is_err()); + assert!(parse_pool(r#"{"relays":[],"websockets":[]}"#).is_err()); + assert!(parse_pool(r#"{"relays":[{"host":"h"}]}"#).is_err()); + assert!(parse_pool(r#"{"relays":"https://relay.example/t"}"#).is_err()); + assert!( + parse_relay("http://relay.example/t/x", Carrier::WebTransport).is_err(), + "non-https scheme must be rejected" + ); + assert!( + parse_relay( + "https://relay.example/t/x#sha256=AAAA", + Carrier::WebTransport + ) + .is_err(), + "a pin of the wrong length must be rejected, not ignored" + ); + assert!( + parse_relay("https://relay.example/t/x#pin=abc", Carrier::WebTransport).is_err(), + "an unrecognized fragment must be rejected, not ignored" + ); + assert!( + parse_relay("ws://relay.example/p", Carrier::WebSocket).is_err(), + "a plaintext WebSocket relay must be rejected: its URL is a credential" + ); + assert!(parse_relay("https://relay.example/p", Carrier::WebSocket).is_err()); + assert!(parse_relay("wss://user:pass@relay.example/p", Carrier::WebSocket).is_err()); + // One bad WebSocket relay spoils the pool, as a bad WebTransport one does. + assert!( + parse_pool( + r#"{"relays":["https://relay.example/t"],"websockets":["ws://relay.example/p"]}"# + ) + .is_err() + ); +} + +fn pool(webtransport: usize, websocket: usize) -> Pool { + let relay = |carrier, n| { + let url = match carrier { + Carrier::WebTransport => format!("https://wt{n}.example/t"), + Carrier::WebSocket => format!("wss://ws{n}.example/p"), + }; + parse_relay(&url, carrier).unwrap() + }; + Pool { + webtransport: (0..webtransport) + .map(|n| relay(Carrier::WebTransport, n)) + .collect(), + websocket: (0..websocket) + .map(|n| relay(Carrier::WebSocket, n)) + .collect(), + } +} + +fn carriers(relays: &[Relay]) -> Vec { + relays.iter().map(|relay| relay.carrier).collect() +} + +#[test] +fn order_puts_webtransport_first_unless_it_is_failing() { + use Carrier::{WebSocket as S, WebTransport as T}; + assert_eq!( + carriers(&order(pool(2, 1), Transport::Auto, false)), + [T, T, S] + ); + assert_eq!( + carriers(&order(pool(2, 1), Transport::Auto, true)), + [S, T, T] + ); + assert_eq!( + carriers(&order(pool(2, 1), Transport::WebTransport, true)), + [T, T] + ); + assert_eq!( + carriers(&order(pool(2, 1), Transport::WebSocket, false)), + [S] + ); + assert!(order(pool(0, 1), Transport::WebTransport, false).is_empty()); + assert!(order(pool(1, 0), Transport::WebSocket, false).is_empty()); +} + +#[test] +fn transport_parses() { + assert_eq!("auto".parse::().unwrap(), Transport::Auto); + assert_eq!( + " WebSocket ".parse::().unwrap(), + Transport::WebSocket + ); + assert_eq!( + "webtransport".parse::().unwrap(), + Transport::WebTransport + ); + assert!("udp".parse::().is_err()); + assert_eq!(Transport::WebSocket.to_string(), "websocket"); +} + +#[test] +fn stream_requests_stay_on_the_session_origin() { + let session: url::Url = "wss://relay.example/uplink/producer/abc".parse().unwrap(); + let ask = + |url: &str| stream_request(&serde_json::json!({ "stream": url }).to_string(), &session); + assert_eq!( + ask("wss://relay.example/uplink/stream/1").unwrap().as_str(), + "wss://relay.example/uplink/stream/1" + ); + assert!(ask("wss://relay.example:443/uplink/stream/1").is_some()); + assert!(ask("wss://relay.example:8443/uplink/stream/1").is_none()); + assert!(ask("wss://elsewhere.example/uplink/stream/1").is_none()); + assert!(ask("ws://relay.example/uplink/stream/1").is_none()); + assert!(ask("wss://user@relay.example/uplink/stream/1").is_none()); + assert!(ask("wss://relay.example/uplink/stream/1#x").is_none()); + assert!(stream_request("not json", &session).is_none()); + assert!(stream_request(r#"{"other":"wss://relay.example/x"}"#, &session).is_none()); +} + +#[test] +fn events_read_as_yas_uplink_prints_them() { + let connected = |carrier| Event::Connected { + relay: "relay.example:443".into(), + carrier, + }; + assert_eq!( + connected(Carrier::WebTransport).to_string(), + "connected to relay relay.example:443" + ); + assert_eq!( + connected(Carrier::WebSocket).to_string(), + "connected to relay relay.example:443 over WebSocket" + ); + assert_eq!( + Event::PoolExhausted { + retry_in: Duration::from_secs(4) + } + .to_string(), + "relay pool exhausted; re-querying in 4s" + ); +} + +#[test] +fn base64url_decodes() { + assert_eq!(base64url_decode("aGVsbG8").unwrap(), b"hello"); + assert_eq!(base64url_decode("aGVsbG8=").unwrap(), b"hello"); + assert_eq!(base64url_decode("_-8").unwrap(), vec![0xff, 0xef]); + assert!(base64url_decode("a+b").is_none()); + assert_eq!(base64url_decode("").unwrap(), Vec::::new()); +} diff --git a/crates/server/src/lib.rs b/crates/server/src/lib.rs index 2f99755a..c8d9d2c9 100644 --- a/crates/server/src/lib.rs +++ b/crates/server/src/lib.rs @@ -2419,8 +2419,12 @@ fn downscale_target_color_mode( } } +/// Ask the compositor for a surface's pixels: the command, then `wake` for its +/// loop (an idle one would see the command only at its next dispatch timeout, a +/// second later), then the reply, or None after `timeout`. async fn request_surface_capture_with_timeout( command_tx: std::sync::mpsc::SyncSender, + wake: impl FnOnce(), surface_id: u16, scale_120: u16, timeout: Duration, @@ -2433,6 +2437,7 @@ async fn request_surface_capture_with_timeout( reply: tx, }) .ok()?; + wake(); // The compositor replies through a blocking std::sync::mpsc channel. // Wait for it off the async runtime so this request never stalls the @@ -15718,12 +15723,22 @@ mod tests { }) .unwrap(); + let woken = Arc::new(std::sync::atomic::AtomicBool::new(false)); + let wake = { + let woken = woken.clone(); + move || woken.store(true, std::sync::atomic::Ordering::SeqCst) + }; let result = - request_surface_capture_with_timeout(command_tx, 7, 0, Duration::from_millis(50)).await; + request_surface_capture_with_timeout(command_tx, wake, 7, 0, Duration::from_millis(50)) + .await; let (w, h, pixels) = result.unwrap(); assert_eq!((w, h), (2, 3)); assert_eq!(pixels.to_rgba(w, h), vec![1, 2, 3, 4]); + assert!( + woken.load(std::sync::atomic::Ordering::SeqCst), + "the compositor loop is woken for the command" + ); } #[tokio::test] @@ -15736,8 +15751,14 @@ mod tests { }) .unwrap(); - let result = - request_surface_capture_with_timeout(command_tx, 7, 0, Duration::from_millis(50)).await; + let result = request_surface_capture_with_timeout( + command_tx, + || {}, + 7, + 0, + Duration::from_millis(50), + ) + .await; assert!(result.is_none()); } diff --git a/crates/server/src/yas.rs b/crates/server/src/yas.rs index 09d20bc0..2798aa6e 100644 --- a/crates/server/src/yas.rs +++ b/crates/server/src/yas.rs @@ -9753,16 +9753,19 @@ impl Session { .max_by_key(|(_, pixels)| u64::from(pixels.width) * u64::from(pixels.height)) .map(|(_, pixels)| (pixels.width, pixels.height, pixels.pixels.clone())) }); - let command_tx = shared - .compositor - .as_ref() - .map(|compositor| compositor.handle.command_tx.clone()); + let command_tx = shared.compositor.as_ref().map(|compositor| { + ( + compositor.handle.command_tx.clone(), + compositor.handle.command_sender(), + ) + }); (snapshot, command_tx) }; let mut captured = match command_tx { - Some(command_tx) => { + Some((command_tx, loop_waker)) => { super::request_surface_capture_with_timeout( command_tx, + move || loop_waker.wake(), surface_id, 0, Duration::from_secs(5), diff --git a/crates/ssh/src/lib.rs b/crates/ssh/src/lib.rs index 64a4352d..744ee4a9 100644 --- a/crates/ssh/src/lib.rs +++ b/crates/ssh/src/lib.rs @@ -1053,6 +1053,10 @@ async fn establish_connection( // every 15 s and give up after 3 consecutive misses (~45 s). keepalive_interval: Some(std::time::Duration::from_secs(15)), keepalive_max: 3, + // YAS requests are small frames, often several in a row: with Nagle's algorithm, one + // waits for the ACK of the one before, which the server delays (40 ms on Linux) since it + // has nothing to answer until the whole request is in. + nodelay: true, ..Default::default() }; diff --git a/docs/uplink.md b/docs/uplink.md index 49207b1a..fa1003f5 100644 --- a/docs/uplink.md +++ b/docs/uplink.md @@ -2,13 +2,15 @@ `yas uplink --allow-client PUBLIC_KEY`, with the private key in `YAS_UPLINK_IDENTITY`, makes -the local YAS server reachable outside NAT. It holds an outbound WebTransport -session to a relay. Each relay-initiated stream must complete end-to-end +the local YAS server reachable outside NAT. It holds an outbound session to a +relay: WebTransport (HTTP/3 over UDP), or WebSockets (over TCP) where a network +blocks UDP. Each relay-initiated stream must complete end-to-end mutual authentication before it can reach the local server socket. This document specifies the protocol between the uplink and its control endpoint and relay. It leaves the relay side abstract: a WebTransport server that opens -one bidirectional stream per consumer and -forwards opaque bytes can act as a relay. The inner stream carries Noise records. +one bidirectional stream per consumer, or a WebSocket server that asks for +one WebSocket per consumer, and forwards opaque bytes can act as a relay. The +inner stream carries Noise records. ## Roles @@ -16,7 +18,7 @@ forwards opaque bytes can act as a relay. The inner stream carries Noise records | ---------------- | -------------------------------------------------------------------- | | uplink | `yas uplink` — connects out, bridges streams to the local YAS server | | control endpoint | HTTPS URL that authenticates the uplink and allocates it a relay | -| relay | WebTransport server the uplink stays connected to | +| relay | WebTransport or WebSocket server the uplink stays connected to | | consumer | A YAS client reaching the server through the relay | ## End-to-end trust and setup @@ -150,11 +152,20 @@ Accept: application/json A success response is the **relay pool**: ```json -{ "relays": ["https://relay-1.example.com:4443/t/kfV3aB#sha256="] } +{ + "relays": ["https://relay-1.example.com:4443/t/kfV3aB#sha256="], + "websockets": ["wss://relay-1.example.com/uplink/producer/Qm9vdA"] +} ``` -- `relays` is a non-empty array of `https` URLs. Any other scheme is an - error. +- `relays` is an array of `https` URLs: WebTransport relays. Any other + scheme is an error. +- `websockets` is an array of `wss` URLs: WebSocket relays (see + [WebSocket relay session](#websocket-relay-session-version-1)). Any other + scheme, plaintext `ws` included, is an error. Uplinks older than this + field ignore it, so a control endpoint can offer both kinds to every uplink. +- Either array may be empty or absent, not both. An uplink only uses the + relays its transport allows (see [Choosing a transport](#choosing-a-transport)). - **A relay URL is a credential.** Whatever authenticates the uplink to the relay (a token in the path, a capability URL) is embedded in it. Implementations MUST NOT log relay URLs; log `host:port` instead. @@ -229,10 +240,33 @@ Both producers and consumers must upgrade and configure keys. Old `uplink:https://relay.example#CLIENT_TOKEN` URIs and unencrypted consumer streams are intentionally rejected. Relays must forward Noise bytes unchanged. +## Choosing a transport + +`yas uplink --transport auto|webtransport|websocket` (or +`YAS_UPLINK_TRANSPORT`) chooses what carries the relay session: + +- `auto` (the default) tries the pool's WebTransport relays first, then its + WebSocket relays. While the pool also names WebSocket relays, a WebTransport + session that isn't up within **5 seconds** counts as failed: a network that + drops UDP gives no error, and QUIC would otherwise wait out its idle + timeout. After WebTransport fails to connect, or a WebTransport session + ends within 60 seconds of starting, the uplink tries WebSocket relays first + for the next 10 minutes, then WebTransport first again. +- `webtransport` uses only the pool's `relays`; `websocket` only its + `websockets`. A pool with none of them is retried like a failed control + request. + +Both carry the same Noise streams to the same local server; only WebTransport +carries datagrams. + ## Relay session -The uplink shuffles the pool and tries each relay in order: a -WebTransport (HTTP/3 CONNECT) session to the relay URL. Liveness settings +The uplink shuffles the pool's relays of each transport and tries them in the +order above. + +### WebTransport relay session + +A WebTransport (HTTP/3 CONNECT) session to the relay URL. Liveness settings are a **10s keepalive** and a **30s idle timeout**, so a dead relay is noticed within 30 seconds without any application-level pings. @@ -241,6 +275,52 @@ per consumer**. After Noise authentication, the uplink bridges decrypted bytes to a fresh local YAS socket. Direct streams carry the normal YAS preface and length-prefixed frames, unparsed and unreframed inside Noise. +### WebSocket relay session (version 1) + +A relay can't open streams over one WebSocket, so it asks the uplink to open +one WebSocket per consumer instead. Both kinds of WebSocket use the +subprotocol `yas-uplink.v1`: the uplink offers it in +`Sec-WebSocket-Protocol` and fails the relay unless the relay selects it. A +later version of this carrier gets a new subprotocol, which the uplink offers +alongside. + +- **TLS.** `wss` only. A `#sha256=` fragment pins the relay's certificate + exactly as for WebTransport relays (and is stripped before connecting); + otherwise the system roots (or `SSL_CERT_FILE`/`SSL_CERT_DIR`) verify it. + The relay URL is a credential: implementations MUST NOT log it. +- **Session.** The uplink holds one WebSocket to the relay URL. The relay + sends a **text message per consumer**, a JSON object naming a stream URL: + + ```json + { "stream": "wss://relay-1.example.com/uplink/stream/c2Vzc2lvbg" } + ``` + + The stream URL is a single-use credential, and must be `wss` on the + session's own origin (scheme, host and port), without userinfo or a + fragment, so the session's TLS trust applies to it. The uplink ignores + messages that aren't such an object (unknown fields included, for later + relays) and never sends data messages on the session. + +- **Streams.** For each request the uplink opens a WebSocket to the stream + URL (same subprotocol, same pin) and treats it as a WebTransport stream: + Noise authentication, then the local socket. Binary messages carry opaque + chunks of at most 64 KiB whose boundaries mean nothing; text messages are + ignored. Noise's authenticated FIN carries the half-close, so neither end + closes a stream WebSocket before both directions are done; closing it ends + both. At most 64 streams may be authenticating per session; the uplink + ignores requests beyond that, and the relay gives up on a stream the uplink + doesn't open within 10 seconds. +- **Liveness.** The uplink pings every **10 seconds** and gives the session up + after **30 seconds** without hearing anything from the relay. A relay should + likewise drop a session it hears nothing on for 30 seconds. +- **Datagrams.** None: an encrypted composite selector asking for a datagram + lane is refused on a WebSocket session (the stream closes), and consumers + on reliable carriers never send one. + +A relay pairs a consumer's attachment (`/attach` and its WSS worker) with a +stream WebSocket just as it would with a WebTransport stream, forwarding the +bytes unchanged either way. + ### Encrypted native datagrams A consumer with an unreliable carrier sends a composite-main selector **inside @@ -293,8 +373,9 @@ move audio onto unreliable datagrams. ### Closure Failed authentication or unavailable local IPC closes the consumer stream. -On SIGINT the outer session closes with application code 2. The worker should -close its consumer connection when the corresponding stream closes. +On SIGINT the outer session closes: a WebTransport session with application +code 2, a WebSocket session with close code 1001 (going away). The worker +should close its consumer connection when the corresponding stream closes. ## Reconnection @@ -305,13 +386,47 @@ close its consumer connection when the corresponding stream closes. re-balanced on every reconnect — and reset the backoff. - Pool exhausted with no session established: back off (same schedule as the control endpoint) and re-query. -- On SIGINT the uplink closes the active session with code 2 instead of - letting it idle out on the relay. +- On SIGINT the uplink closes the active session (code 2, or WebSocket close + code 1001) instead of letting it idle out on the relay. Relay URLs stay valid for as long as their embedded credential does; the uplink treats each pool response as single-use and re-queries rather than caching it. +## Embedding the uplink + +The uplink is a library as well as a command: +`yas_proxy::uplink_producer` (re-exported as `yas_client::uplink_producer`). +It installs no process-wide TLS provider, changes no environment and never +exits the process. + +```rust +use yas_client::uplink_producer::{Event, Identity, Local, Producer, Transport}; + +let identity = Identity::from_base64(&private_key)?; +let producer = Producer::new( + "https://relay.example/uplink/control", + &token, + identity.server_config(allowed_client_keys)?, + Local::socket(socket_path), // a Unix socket, or a named pipe on Windows +)? +.transport(Transport::Auto) +.on_event(|event: &Event| eprintln!("[uplink] {event}")); +let handle = producer.handle(); +tokio::spawn(producer.run_until(shutdown)); +// Later, without restarting the session: +handle.set_server_config(identity.server_config(new_allowlist)?); +handle.set_token(renewed_token); +``` + +`Local::custom` reaches a server some other way. `Event` values (a relay +session up, with its carrier; a session ended; a relay failing…) display as +the lines `yas uplink` prints. A new allowlist applies to consumers that start +their handshake afterwards; consumers already connected keep their authority +until they disconnect, as across a restart. A new token applies from the next +control request. `run_until` returns an error only when the control endpoint +refuses the token. + ## Browser embedding Serve the browser client from an origin trusted independently of the relay. @@ -372,7 +487,12 @@ in process arguments. Both endpoints must upgrade together. `direnv exec . cargo test -p yas-cli --test uplink_e2e` exercises the complete CLI path through a local HTTPS control endpoint and WSS/WebTransport relay to -an isolated YAS server. CI includes this test in the Rust workspace test suite. +an isolated YAS server: over WebTransport, over WebSockets when forced, and +falling back to WebSockets when the WebTransport relay's UDP goes nowhere. CI +includes this test in the Rust workspace test suite. +`cargo test -p yas-proxy uplink_producer` covers the library: pools, the +transport order, WebSocket sessions and streams, allowlist changes while +running, and refusals before local IPC. For private relay infrastructure, `SSL_CERT_FILE` and `SSL_CERT_DIR` select outer TLS trust roots consistently for HTTPS, WSS, and WebTransport (unless a