diff --git a/Cargo.lock b/Cargo.lock index ec912aba6e..b185771c43 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4040,6 +4040,7 @@ dependencies = [ "web-transport-proto", "web-transport-quiche", "web-transport-quinn", + "web-transport-trait", ] [[package]] diff --git a/rs/libmoq/src/error.rs b/rs/libmoq/src/error.rs index 2e7be9b054..d43a146c65 100644 --- a/rs/libmoq/src/error.rs +++ b/rs/libmoq/src/error.rs @@ -115,6 +115,14 @@ pub enum Error { #[error("offline")] Offline, + /// Connection was rejected as unauthorized by the server. + #[error("unauthorized")] + Unauthorized, + + /// Connection was forbidden by the server. + #[error("forbidden")] + Forbidden, + /// Error from the hang media layer. #[error("hang error: {0}")] Hang(#[from] hang::Error), @@ -195,6 +203,8 @@ impl ffi::ReturnCode for Error { Error::Audio(_) => -30, Error::GroupNotFound => -31, Error::Native(_) => -32, + Error::Unauthorized => -33, + Error::Forbidden => -34, } } } diff --git a/rs/libmoq/src/session.rs b/rs/libmoq/src/session.rs index 56ed2bfe3d..1ecad59604 100644 --- a/rs/libmoq/src/session.rs +++ b/rs/libmoq/src/session.rs @@ -74,25 +74,24 @@ impl Session { .with_consume(consume) .reconnect(url); - // report() runs until the reconnect loop gives up; map its terminal error to Connect. - Self::report(callback, reconnect) - .await - .map_err(|err| Error::Connect(Arc::new(err))) + Self::report(callback, reconnect).await } /// Forward connection epochs to the status callback until the reconnect loop stops. /// /// Returns the terminal error via `?`. Disconnects aren't reported: status 0 is reserved for a /// clean close (delivered as the terminal callback once the task ends). - async fn report(callback: ffi::OnStatus, mut reconnect: moq_native::Reconnect) -> anyhow::Result<()> { + async fn report(callback: ffi::OnStatus, mut reconnect: moq_native::Reconnect) -> Result<(), Error> { let mut connects: u64 = 0; loop { - if let moq_native::Status::Connected = reconnect.status().await? { + if let moq_native::Status::Connected = reconnect.status().await.map_err(map_connect_error)? { connects += 1; // Positive status carries the connection epoch, so callers can tell a // reconnect (>1) from the first connect (1). No lock is held, so the C // callback is free to re-enter libmoq. - let code = i32::try_from(connects).context("connection epoch exceeded i32::MAX")?; + let code = i32::try_from(connects) + .context("connection epoch exceeded i32::MAX") + .map_err(|err| Error::Connect(Arc::new(err)))?; callback.call(code); } } @@ -110,3 +109,40 @@ impl Session { Ok(()) } } + +fn map_connect_error(err: moq_native::Error) -> Error { + match err.connect_error() { + Some(moq_native::ConnectError::Unauthorized) => Error::Unauthorized, + Some(moq_native::ConnectError::Forbidden) => Error::Forbidden, + _ => Error::Connect(Arc::new(err.into())), + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::ffi::ReturnCode; + + #[test] + fn maps_native_auth_connect_errors() { + assert!(matches!( + map_connect_error(moq_native::ConnectError::Unauthorized.into()), + Error::Unauthorized + )); + assert!(matches!( + map_connect_error(moq_native::ConnectError::Forbidden.into()), + Error::Forbidden + )); + assert!(matches!( + map_connect_error(moq_net::Error::Unauthorized.into()), + Error::Unauthorized + )); + assert!(matches!( + map_connect_error(moq_native::Error::ConnectFailed), + Error::Connect(_) + )); + assert_eq!(Error::Unauthorized.code(), -33); + assert_eq!(Error::Forbidden.code(), -34); + assert_eq!(map_connect_error(moq_native::Error::ConnectFailed).code(), -5); + } +} diff --git a/rs/moq-ffi/src/error.rs b/rs/moq-ffi/src/error.rs index d0f2b784c3..be497c4b8c 100644 --- a/rs/moq-ffi/src/error.rs +++ b/rs/moq-ffi/src/error.rs @@ -54,6 +54,9 @@ pub enum MoqError { #[error("unauthorized")] Unauthorized, + #[error("forbidden")] + Forbidden, + #[error("log: {0}")] Log(String), } diff --git a/rs/moq-ffi/src/session.rs b/rs/moq-ffi/src/session.rs index e832b087c0..de083dc2bc 100644 --- a/rs/moq-ffi/src/session.rs +++ b/rs/moq-ffi/src/session.rs @@ -15,11 +15,7 @@ struct Client { impl Client { async fn connect(&self, url: Url) -> Result, MoqError> { - let client = self - .config - .clone() - .init() - .map_err(|err| MoqError::Connect(format!("{err}")))?; + let client = self.config.clone().init().map_err(map_connect_error)?; let publish = self.publish.as_ref().map(|o| o.inner().consume()); let consume = self.consume.as_ref().map(|o| o.inner().clone()); @@ -29,12 +25,37 @@ impl Client { .with_consume(consume) .connect(url) .await - .map_err(|err| MoqError::Connect(format!("{err}")))?; + .map_err(map_connect_error)?; Ok(Arc::new(MoqSession::new(session))) } } +fn map_connect_error(err: moq_native::Error) -> MoqError { + match err.connect_error() { + Some(moq_native::ConnectError::Unauthorized) => MoqError::Unauthorized, + Some(moq_native::ConnectError::Forbidden) => MoqError::Forbidden, + _ => MoqError::Connect(format!("{err}")), + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn maps_native_auth_connect_errors() { + assert!(matches!( + map_connect_error(moq_native::ConnectError::Unauthorized.into()), + MoqError::Unauthorized + )); + assert!(matches!( + map_connect_error(moq_native::ConnectError::Forbidden.into()), + MoqError::Forbidden + )); + } +} + #[derive(uniffi::Object)] pub struct MoqClient { task: Task, diff --git a/rs/moq-native/Cargo.toml b/rs/moq-native/Cargo.toml index a04267a03e..d5bca2a288 100644 --- a/rs/moq-native/Cargo.toml +++ b/rs/moq-native/Cargo.toml @@ -62,6 +62,7 @@ web-transport-noq = { workspace = true, optional = true } web-transport-proto = { workspace = true, optional = true } web-transport-quiche = { workspace = true, optional = true } web-transport-quinn = { workspace = true, optional = true } +web-transport-trait = { workspace = true } [target.'cfg(target_os = "android")'.dependencies] tracing-android = { version = "0.2", optional = true } diff --git a/rs/moq-native/src/client.rs b/rs/moq-native/src/client.rs index 1fe2249990..10779b2028 100644 --- a/rs/moq-native/src/client.rs +++ b/rs/moq-native/src/client.rs @@ -1,4 +1,6 @@ use crate::{Backoff, Error, QuicBackend, Reconnect}; +#[cfg(feature = "websocket")] +use std::future::Future; use std::net; use url::Url; @@ -266,24 +268,11 @@ impl Client { if let Some(noq) = self.noq.as_ref() { let tls = self.tls.clone(); let quic_url = url.clone(); - let quic_handle = async { - let res = noq.connect(&tls, quic_url).await; - if let Err(err) = &res { - tracing::warn!(%err, "QUIC connection failed"); - } - res - }; + let quic_handle = async { noq.connect(&tls, quic_url).await.map_err(Error::from) }; #[cfg(feature = "websocket")] { - let alpns = self.versions.alpns(); - let ws_handle = crate::websocket::race_handle(&self.websocket, &self.tls, url, &alpns); - - return Ok(tokio::select! { - Ok(quic) = quic_handle => self.moq.connect(quic).await?, - Some(Ok(ws)) = ws_handle => self.moq.connect(ws).await?, - else => return Err(Error::ConnectFailed), - }); + return self.race_moq_connect(url, quic_handle).await; } #[cfg(not(feature = "websocket"))] @@ -297,24 +286,11 @@ impl Client { if let Some(quinn) = self.quinn.as_ref() { let tls = self.tls.clone(); let quic_url = url.clone(); - let quic_handle = async { - let res = quinn.connect(&tls, quic_url).await; - if let Err(err) = &res { - tracing::warn!(%err, "QUIC connection failed"); - } - res - }; + let quic_handle = async { quinn.connect(&tls, quic_url).await.map_err(Error::from) }; #[cfg(feature = "websocket")] { - let alpns = self.versions.alpns(); - let ws_handle = crate::websocket::race_handle(&self.websocket, &self.tls, url, &alpns); - - return Ok(tokio::select! { - Ok(quic) = quic_handle => self.moq.connect(quic).await?, - Some(Ok(ws)) = ws_handle => self.moq.connect(ws).await?, - else => return Err(Error::ConnectFailed), - }); + return self.race_moq_connect(url, quic_handle).await; } #[cfg(not(feature = "websocket"))] @@ -327,24 +303,11 @@ impl Client { #[cfg(feature = "quiche")] if let Some(quiche) = self.quiche.as_ref() { let quic_url = url.clone(); - let quic_handle = async { - let res = quiche.connect(quic_url).await; - if let Err(err) = &res { - tracing::warn!(%err, "QUIC connection failed"); - } - res - }; + let quic_handle = async { quiche.connect(quic_url).await.map_err(Error::from) }; #[cfg(feature = "websocket")] { - let alpns = self.versions.alpns(); - let ws_handle = crate::websocket::race_handle(&self.websocket, &self.tls, url, &alpns); - - return Ok(tokio::select! { - Ok(quic) = quic_handle => self.moq.connect(quic).await?, - Some(Ok(ws)) = ws_handle => self.moq.connect(ws).await?, - else => return Err(Error::ConnectFailed), - }); + return self.race_moq_connect(url, quic_handle).await; } #[cfg(not(feature = "websocket"))] @@ -364,6 +327,93 @@ impl Client { #[cfg(not(feature = "websocket"))] return Err(Error::NoBackend("no QUIC backend matched; this should not happen")); } + + #[cfg(feature = "websocket")] + async fn race_moq_connect(&self, url: Url, quic: Q) -> crate::Result + where + Q: Future>, + S: web_transport_trait::Session, + { + let alpns = self.versions.alpns(); + let ws_config = self.websocket.clone(); + let ws_tls = self.tls.clone(); + let websocket = async move { + crate::websocket::race_handle(&ws_config, &ws_tls, url, &alpns) + .await + .map(|res| res.map_err(Error::from)) + }; + + match race_transport_connect(quic, websocket).await? { + TransportRace::Quic(quic) => Ok(self.moq.connect(quic).await?), + TransportRace::WebSocket(websocket) => Ok(self.moq.connect(websocket).await?), + } + } +} + +#[cfg(feature = "websocket")] +#[derive(Debug, PartialEq, Eq)] +enum TransportRace { + Quic(Q), + WebSocket(W), +} + +#[cfg(feature = "websocket")] +async fn race_transport_connect(quic: Q, websocket: W) -> crate::Result> +where + Q: Future>, + W: Future>>, +{ + tokio::pin!(quic); + tokio::pin!(websocket); + + let mut quic_err = None; + let mut websocket_err = None; + let mut quic_done = false; + let mut websocket_done = false; + + loop { + tokio::select! { + res = &mut quic, if !quic_done => { + match res { + Ok(session) => return Ok(TransportRace::Quic(session)), + Err(err) if err.is_auth() => return Err(err), + Err(err) => { + tracing::warn!(%err, "QUIC connection failed"); + quic_err = Some(err); + quic_done = true; + } + } + } + res = &mut websocket, if !websocket_done => { + match res { + Some(Ok(session)) => return Ok(TransportRace::WebSocket(session)), + Some(Err(err)) if err.is_auth() => return Err(err), + Some(Err(err)) => { + tracing::warn!(%err, "WebSocket connection failed"); + websocket_err = Some(err); + websocket_done = true; + } + None => { + websocket_done = true; + } + } + } + else => break, + } + + if quic_done && websocket_done { + break; + } + } + + match (quic_err, websocket_err) { + (Some(quic), Some(websocket)) => Err(Error::TransportRace { + quic: std::sync::Arc::new(quic), + websocket: std::sync::Arc::new(websocket), + }), + (Some(err), None) | (None, Some(err)) => Err(err), + (None, None) => Err(Error::ConnectFailed), + } } #[cfg(test)] @@ -436,4 +486,44 @@ mod tests { // versions() helper returns all when none specified assert_eq!(config.versions().alpns().len(), moq_net::ALPNS.len()); } + + #[cfg(feature = "websocket")] + #[tokio::test] + async fn race_transport_connect_stops_on_quic_auth_error() { + let quic = async { Err::(crate::ConnectError::Unauthorized.into()) }; + let websocket = async { + // This only needs to complete later than the immediately ready QUIC auth error. + tokio::task::yield_now().await; + Some(Ok(1usize)) + }; + + let err = super::race_transport_connect(quic, websocket).await.unwrap_err(); + assert_eq!(err.connect_error(), Some(crate::ConnectError::Unauthorized)); + } + + #[cfg(feature = "websocket")] + #[tokio::test] + async fn race_transport_connect_keeps_websocket_after_quic_non_auth_error() { + let quic = async { Err::(Error::ConnectFailed) }; + let websocket = async { Some(Ok(7usize)) }; + + let value = super::race_transport_connect(quic, websocket).await.unwrap(); + assert_eq!(value, super::TransportRace::WebSocket(7)); + } + + #[cfg(feature = "websocket")] + #[tokio::test] + async fn race_transport_connect_returns_when_quic_transport_connects() { + let quic = async { Ok("quic") }; + let websocket = std::future::pending::>>(); + + let value = tokio::time::timeout( + std::time::Duration::from_secs(1), + super::race_transport_connect(quic, websocket), + ) + .await + .expect("race waited for WebSocket after QUIC transport connected") + .unwrap(); + assert_eq!(value, super::TransportRace::Quic("quic")); + } } diff --git a/rs/moq-native/src/connect.rs b/rs/moq-native/src/connect.rs new file mode 100644 index 0000000000..e5cbd651bd --- /dev/null +++ b/rs/moq-native/src/connect.rs @@ -0,0 +1,42 @@ +/// Error returned when connection setup fails for a terminal auth reason. +#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)] +#[non_exhaustive] +pub enum ConnectError { + #[error("unauthorized")] + Unauthorized, + + #[error("forbidden")] + Forbidden, +} + +impl ConnectError { + pub(crate) fn from_status_u16(status: u16) -> Option { + match status { + 401 => Some(Self::Unauthorized), + 403 => Some(Self::Forbidden), + _ => None, + } + } + + pub fn is_auth(&self) -> bool { + matches!(self, Self::Unauthorized | Self::Forbidden) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn auth_statuses_are_terminal() { + assert_eq!(ConnectError::from_status_u16(401), Some(ConnectError::Unauthorized)); + assert_eq!(ConnectError::from_status_u16(403), Some(ConnectError::Forbidden)); + } + + #[test] + fn non_auth_statuses_are_not_terminal() { + for status in [400, 404, 500] { + assert_eq!(ConnectError::from_status_u16(status), None); + } + } +} diff --git a/rs/moq-native/src/error.rs b/rs/moq-native/src/error.rs index 18981e9cd2..887e6a13f1 100644 --- a/rs/moq-native/src/error.rs +++ b/rs/moq-native/src/error.rs @@ -29,6 +29,13 @@ pub enum Error { #[error("failed to connect to server")] ConnectFailed, + #[error(transparent)] + Connect(#[from] crate::ConnectError), + + #[cfg(feature = "websocket")] + #[error("failed to connect to server: QUIC failed: {quic}; WebSocket failed: {websocket}")] + TransportRace { quic: Arc, websocket: Arc }, + #[cfg(feature = "iroh")] #[error("Iroh support is not enabled")] IrohDisabled, @@ -66,6 +73,30 @@ pub enum Error { WebSocket(Arc), } +impl Error { + pub fn connect_error(&self) -> Option { + match self { + Self::Connect(err) => Some(*err), + Self::MoqNet(moq_net::Error::Unauthorized) => Some(crate::ConnectError::Unauthorized), + #[cfg(feature = "quinn")] + Self::Quinn(err) => err.connect_error(), + #[cfg(feature = "noq")] + Self::Noq(err) => err.connect_error(), + #[cfg(feature = "quiche")] + Self::Quiche(err) => err.connect_error(), + #[cfg(feature = "websocket")] + Self::TransportRace { quic, websocket } => quic.connect_error().or_else(|| websocket.connect_error()), + #[cfg(feature = "websocket")] + Self::WebSocket(err) => err.connect_error(), + _ => None, + } + } + + pub fn is_auth(&self) -> bool { + self.connect_error().is_some_and(|err| err.is_auth()) + } +} + // The wrapped sources aren't `Clone`, so `#[from]` can't store them behind `Arc` // directly. These hand-written conversions keep `?` ergonomic at the call sites. impl From for Error { @@ -89,6 +120,10 @@ impl From for Error { #[cfg(feature = "quinn")] impl From for Error { fn from(err: crate::quinn::Error) -> Self { + if let Some(err) = err.connect_error() { + return Self::Connect(err); + } + Self::Quinn(Arc::new(err)) } } @@ -96,6 +131,10 @@ impl From for Error { #[cfg(feature = "noq")] impl From for Error { fn from(err: crate::noq::Error) -> Self { + if let Some(err) = err.connect_error() { + return Self::Connect(err); + } + Self::Noq(Arc::new(err)) } } @@ -103,6 +142,10 @@ impl From for Error { #[cfg(feature = "quiche")] impl From for Error { fn from(err: crate::quiche::Error) -> Self { + if let Some(err) = err.connect_error() { + return Self::Connect(err); + } + Self::Quiche(Arc::new(err)) } } @@ -117,9 +160,33 @@ impl From for Error { #[cfg(feature = "websocket")] impl From for Error { fn from(err: crate::websocket::Error) -> Self { + if let Some(err) = err.connect_error() { + return Self::Connect(err); + } + Self::WebSocket(Arc::new(err)) } } /// Convenience alias for results produced by this crate. pub type Result = std::result::Result; + +#[cfg(all(test, feature = "websocket"))] +mod tests { + use super::*; + + #[test] + fn transport_race_propagates_nested_connect_errors() { + let quic = Error::TransportRace { + quic: Arc::new(crate::ConnectError::Unauthorized.into()), + websocket: Arc::new(crate::ConnectError::Forbidden.into()), + }; + assert_eq!(quic.connect_error(), Some(crate::ConnectError::Unauthorized)); + + let websocket = Error::TransportRace { + quic: Arc::new(Error::ConnectFailed), + websocket: Arc::new(crate::ConnectError::Forbidden.into()), + }; + assert_eq!(websocket.connect_error(), Some(crate::ConnectError::Forbidden)); + } +} diff --git a/rs/moq-native/src/lib.rs b/rs/moq-native/src/lib.rs index cb9cdaa99b..89934269d7 100644 --- a/rs/moq-native/src/lib.rs +++ b/rs/moq-native/src/lib.rs @@ -12,6 +12,7 @@ pub(crate) const DEFAULT_MAX_STREAMS: u64 = 1024; mod client; +mod connect; mod crypto; mod error; #[cfg(feature = "jemalloc")] @@ -32,6 +33,7 @@ mod watch; pub mod websocket; pub use client::*; +pub use connect::ConnectError; pub use error::{Error, Result}; pub use log::*; pub use reconnect::*; diff --git a/rs/moq-native/src/noq.rs b/rs/moq-native/src/noq.rs index 55fb5999fd..61b2890eb1 100644 --- a/rs/moq-native/src/noq.rs +++ b/rs/moq-native/src/noq.rs @@ -92,6 +92,9 @@ pub enum Error { #[error(transparent)] Client(#[from] web_transport_noq::ClientError), + #[error(transparent)] + ConnectRejected(#[from] crate::ConnectError), + #[error(transparent)] Server(#[from] web_transport_noq::ServerError), @@ -217,7 +220,9 @@ impl NoqClient { } let session = match url.scheme() { - "https" => web_transport_noq::Session::connect(connection, request).await?, + "https" => web_transport_noq::Session::connect(connection, request) + .await + .map_err(map_client_error)?, "moqt" | "moql" => { let handshake = connection .handshake_data() @@ -238,6 +243,49 @@ impl NoqClient { } } +impl Error { + pub(crate) fn connect_error(&self) -> Option { + match self { + Self::ConnectRejected(err) => Some(*err), + Self::Client(err) => classify_client_error(err), + _ => None, + } + } +} + +fn map_client_error(err: web_transport_noq::ClientError) -> Error { + if let Some(err) = classify_client_error(&err) { + return err.into(); + } + + err.into() +} + +fn classify_client_error(err: &web_transport_noq::ClientError) -> Option { + match err { + web_transport_noq::ClientError::HttpError(err) => classify_connect_error(err), + _ => None, + } +} + +fn classify_connect_error(err: &web_transport_noq::ConnectError) -> Option { + match err { + web_transport_noq::ConnectError::ErrorStatus(status) => crate::ConnectError::from_status_u16(status.as_u16()), + web_transport_noq::ConnectError::ProtoError(err) => classify_proto_error(err), + _ => None, + } +} + +fn classify_proto_error(err: &web_transport_noq::proto::ConnectError) -> Option { + match err { + web_transport_noq::proto::ConnectError::ErrorStatus(status) + | web_transport_noq::proto::ConnectError::WrongStatus(Some(status)) => { + crate::ConnectError::from_status_u16(status.as_u16()) + } + _ => None, + } +} + // ── Server ────────────────────────────────────────────────────────── pub(crate) struct NoqServer { diff --git a/rs/moq-native/src/quiche.rs b/rs/moq-native/src/quiche.rs index 20fb4277e8..df4e2474d6 100644 --- a/rs/moq-native/src/quiche.rs +++ b/rs/moq-native/src/quiche.rs @@ -58,7 +58,10 @@ pub enum Error { Establish(#[source] web_transport_quiche::ez::ConnectionError), #[error("failed to connect to quiche server")] - ClientConnect(#[source] web_transport_quiche::ClientError), + ClientConnect(#[from] web_transport_quiche::ClientError), + + #[error(transparent)] + ConnectRejected(#[from] crate::ConnectError), #[error("failed to create quiche server")] ServerBuild(#[source] std::io::Error), @@ -150,7 +153,7 @@ impl QuicheClient { .map_err(Error::Establish)?; let session = web_transport_quiche::Connection::connect(conn, request) .await - .map_err(Error::ClientConnect)?; + .map_err(map_client_error)?; Ok(session) } "moqt" | "moql" => { @@ -174,6 +177,49 @@ impl QuicheClient { } } +impl Error { + pub(crate) fn connect_error(&self) -> Option { + match self { + Self::ConnectRejected(err) => Some(*err), + Self::ClientConnect(err) => classify_client_error(err), + _ => None, + } + } +} + +fn map_client_error(err: web_transport_quiche::ClientError) -> Error { + if let Some(err) = classify_client_error(&err) { + return err.into(); + } + + err.into() +} + +fn classify_client_error(err: &web_transport_quiche::ClientError) -> Option { + match err { + web_transport_quiche::ClientError::Connect(err) => classify_connect_error(err), + _ => None, + } +} + +fn classify_connect_error(err: &web_transport_quiche::h3::ConnectError) -> Option { + match err { + web_transport_quiche::h3::ConnectError::Status(status) => crate::ConnectError::from_status_u16(status.as_u16()), + web_transport_quiche::h3::ConnectError::Proto(err) => classify_proto_error(err), + _ => None, + } +} + +fn classify_proto_error(err: &web_transport_quiche::proto::ConnectError) -> Option { + match err { + web_transport_quiche::proto::ConnectError::ErrorStatus(status) + | web_transport_quiche::proto::ConnectError::WrongStatus(Some(status)) => { + crate::ConnectError::from_status_u16(status.as_u16()) + } + _ => None, + } +} + // ── Server ────────────────────────────────────────────────────────── pub(crate) struct QuicheServer { diff --git a/rs/moq-native/src/quinn.rs b/rs/moq-native/src/quinn.rs index 6a8360a862..de01612cac 100644 --- a/rs/moq-native/src/quinn.rs +++ b/rs/moq-native/src/quinn.rs @@ -93,6 +93,9 @@ pub enum Error { #[error(transparent)] Client(#[from] web_transport_quinn::ClientError), + #[error(transparent)] + ConnectRejected(#[from] crate::ConnectError), + #[error(transparent)] Server(#[from] web_transport_quinn::ServerError), @@ -219,7 +222,9 @@ impl QuinnClient { } let session = match url.scheme() { - "https" => web_transport_quinn::Session::connect(connection, request).await?, + "https" => web_transport_quinn::Session::connect(connection, request) + .await + .map_err(map_client_error)?, "moqt" | "moql" => { let handshake = connection .handshake_data() @@ -240,6 +245,49 @@ impl QuinnClient { } } +impl Error { + pub(crate) fn connect_error(&self) -> Option { + match self { + Self::ConnectRejected(err) => Some(*err), + Self::Client(err) => classify_client_error(err), + _ => None, + } + } +} + +fn map_client_error(err: web_transport_quinn::ClientError) -> Error { + if let Some(err) = classify_client_error(&err) { + return err.into(); + } + + err.into() +} + +fn classify_client_error(err: &web_transport_quinn::ClientError) -> Option { + match err { + web_transport_quinn::ClientError::HttpError(err) => classify_connect_error(err), + _ => None, + } +} + +fn classify_connect_error(err: &web_transport_quinn::ConnectError) -> Option { + match err { + web_transport_quinn::ConnectError::ErrorStatus(status) => crate::ConnectError::from_status_u16(status.as_u16()), + web_transport_quinn::ConnectError::ProtoError(err) => classify_proto_error(err), + _ => None, + } +} + +fn classify_proto_error(err: &web_transport_quinn::proto::ConnectError) -> Option { + match err { + web_transport_quinn::proto::ConnectError::ErrorStatus(status) + | web_transport_quinn::proto::ConnectError::WrongStatus(Some(status)) => { + crate::ConnectError::from_status_u16(status.as_u16()) + } + _ => None, + } +} + // ── Server ────────────────────────────────────────────────────────── pub(crate) struct QuinnServer { diff --git a/rs/moq-native/src/reconnect.rs b/rs/moq-native/src/reconnect.rs index ff084e0b57..5e52791a99 100644 --- a/rs/moq-native/src/reconnect.rs +++ b/rs/moq-native/src/reconnect.rs @@ -146,6 +146,9 @@ impl Reconnect { retry_start = tokio::time::Instant::now(); } Err(err) => { + if err.is_auth() { + return Err(err); + } tracing::warn!(%url, %err, ?delay, "connection failed, retrying"); last_error = Some(err); tokio::time::sleep(delay).await; diff --git a/rs/moq-native/src/websocket.rs b/rs/moq-native/src/websocket.rs index f62774faf4..0793401825 100644 --- a/rs/moq-native/src/websocket.rs +++ b/rs/moq-native/src/websocket.rs @@ -1,4 +1,5 @@ use qmux::tokio_tungstenite; +use qmux::tokio_tungstenite::tungstenite::{self, client::IntoClientRequest, http}; use std::collections::HashSet; use std::sync::{Arc, LazyLock, Mutex}; use std::{net, time}; @@ -23,6 +24,18 @@ pub enum Error { #[error("failed to connect WebSocket")] Connect(#[source] qmux::Error), + #[error("failed to build WebSocket request")] + BuildRequest(#[source] tungstenite::Error), + + #[error("failed to build WebSocket protocols header")] + ProtocolHeader(#[source] http::header::InvalidHeaderValue), + + #[error("failed to connect WebSocket")] + WebSocketConnect(#[source] tungstenite::Error), + + #[error(transparent)] + ConnectRejected(#[from] crate::ConnectError), + #[error("WebSocket accept failed")] Accept(#[source] qmux::Error), } @@ -141,20 +154,41 @@ pub(crate) async fn connect( tracing::debug!(%url, "connecting via WebSocket"); - // Use the existing TLS config (which respects tls-disable-verify) for secure connections + // Use the existing TLS config (which respects tls-disable-verify) for secure connections. let connector = if needs_tls { tokio_tungstenite::Connector::Rustls(Arc::new(tls.clone())) } else { tokio_tungstenite::Connector::Plain }; - let session = qmux::Client::new() - .with_protocols(alpns) - .with_connector(connector) - .with_keep_alive(qmux::KeepAlive::default()) // 5s ping / 30s deadline — parity with QUIC - .connect(url.as_str()) - .await - .map_err(Error::Connect)?; + let mut request = url.as_str().into_client_request().map_err(Error::BuildRequest)?; + let protocols = websocket_subprotocols(alpns).join(", "); + request.headers_mut().insert( + http::header::SEC_WEBSOCKET_PROTOCOL, + http::HeaderValue::from_str(&protocols).map_err(Error::ProtocolHeader)?, + ); + + let (socket, response) = if needs_tls { + tokio_tungstenite::connect_async_tls_with_config(request, None, false, Some(connector)) + .await + .map_err(map_websocket_error)? + } else { + tokio_tungstenite::connect_async_with_config(request, None, false) + .await + .map_err(map_websocket_error)? + }; + + let alpn = response + .headers() + .get(http::header::SEC_WEBSOCKET_PROTOCOL) + .and_then(|header| header.to_str().ok()) + .map(str::to_owned); + let bare = qmux::ws::Bare::new(socket).with_keep_alive(qmux::KeepAlive::default()); + let bare = match alpn.as_deref() { + Some(alpn) => bare.with_alpn(alpn), + None => bare, + }; + let session = bare.connect(); tracing::warn!(%url, "using WebSocket fallback"); WEBSOCKET_WON.lock().unwrap().insert(key); @@ -162,6 +196,34 @@ pub(crate) async fn connect( Ok(session) } +fn websocket_subprotocols(alpns: &[&str]) -> Vec { + let mut protocols = Vec::with_capacity(qmux::ALPNS.len() + qmux::PREFIXES.len() * alpns.len()); + for (&bare, &prefix) in qmux::ALPNS.iter().zip(qmux::PREFIXES) { + protocols.push(bare.to_string()); + protocols.extend(alpns.iter().map(|alpn| format!("{prefix}{alpn}"))); + } + protocols +} + +impl Error { + pub(crate) fn connect_error(&self) -> Option { + match self { + Self::ConnectRejected(err) => Some(*err), + _ => None, + } + } +} + +fn map_websocket_error(err: tungstenite::Error) -> Error { + if let tungstenite::Error::Http(response) = &err + && let Some(err) = crate::ConnectError::from_status_u16(response.status().as_u16()) + { + return err.into(); + } + + Error::WebSocketConnect(err) +} + /// Listens for incoming WebSocket connections on a TCP port. /// /// Use with [`crate::Server::with_websocket`] to accept WebSocket connections diff --git a/rs/moq-native/tests/broadcast.rs b/rs/moq-native/tests/broadcast.rs index e89340d9c1..98763f7a1d 100644 --- a/rs/moq-native/tests/broadcast.rs +++ b/rs/moq-native/tests/broadcast.rs @@ -568,7 +568,7 @@ async fn broadcast_websocket_fallback() { let mut client_config = moq_native::ClientConfig::default(); client_config.tls.disable_verify = Some(true); - // No delay — race QUIC and WebSocket simultaneously. + // No delay. Race QUIC and WebSocket simultaneously. client_config.websocket.delay = None; let client = client_config.init().expect("failed to init client"); @@ -780,3 +780,89 @@ async fn broadcast_race_quic_wins() { .expect("server task panicked") .expect("server task failed"); } + +#[tracing_test::traced_test] +#[tokio::test] +async fn websocket_unauthorized_handshake_is_explicit() { + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("failed to bind TCP listener"); + let addr = listener.local_addr().expect("failed to get local addr"); + + let server_handle = tokio::spawn(async move { + let (mut stream, _) = listener.accept().await?; + let mut buf = [0; 1024]; + let _ = stream.read(&mut buf).await?; + stream + .write_all(b"HTTP/1.1 401 Unauthorized\r\nContent-Length: 0\r\nConnection: close\r\n\r\n") + .await?; + Ok::<_, anyhow::Error>(()) + }); + + let mut client_config = moq_native::ClientConfig::default(); + client_config.websocket.delay = None; + let client = client_config.init().expect("failed to init client"); + let url: url::Url = format!("ws://{addr}").parse().unwrap(); + + let err = tokio::time::timeout(TIMEOUT, client.connect(url)) + .await + .expect("client connect timed out"); + let err = expect_connect_err(err); + assert_connect_error(&err, moq_native::ConnectError::Unauthorized); + + server_handle + .await + .expect("server task panicked") + .expect("server task failed"); +} + +#[tracing_test::traced_test] +#[tokio::test] +async fn reconnect_stops_on_websocket_unauthorized() { + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("failed to bind TCP listener"); + let addr = listener.local_addr().expect("failed to get local addr"); + + let server_handle = tokio::spawn(async move { + let (mut stream, _) = listener.accept().await?; + let mut buf = [0; 1024]; + let _ = stream.read(&mut buf).await?; + stream + .write_all(b"HTTP/1.1 401 Unauthorized\r\nContent-Length: 0\r\nConnection: close\r\n\r\n") + .await?; + Ok::<_, anyhow::Error>(()) + }); + + let mut client_config = moq_native::ClientConfig::default(); + client_config.websocket.delay = None; + let client = client_config.init().expect("failed to init client"); + let url: url::Url = format!("ws://{addr}").parse().unwrap(); + + let reconnect = client.reconnect(url); + let err = tokio::time::timeout(TIMEOUT, reconnect.closed()) + .await + .expect("reconnect close timed out") + .expect_err("reconnect unexpectedly succeeded"); + assert_connect_error(&err, moq_native::ConnectError::Unauthorized); + + server_handle + .await + .expect("server task panicked") + .expect("server task failed"); +} + +fn assert_connect_error(err: &moq_native::Error, expected: moq_native::ConnectError) { + assert_eq!(err.connect_error(), Some(expected), "unexpected error: {err}",); +} + +fn expect_connect_err(result: moq_native::Result) -> moq_native::Error { + match result { + Ok(_) => panic!("client connect unexpectedly succeeded"), + Err(err) => err, + } +}