diff --git a/Cargo.lock b/Cargo.lock index e8fd941..b608bf9 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -820,6 +820,7 @@ dependencies = [ "p256", "rand 0.9.2", "reqwest", + "rustls", "serde", "serde_json", "serial_test", @@ -833,6 +834,7 @@ dependencies = [ "tracing-subscriber", "urlencoding", "uuid", + "webpki-roots 0.26.11", "wiremock", ] diff --git a/Cargo.toml b/Cargo.toml index 3788d15..067c9cf 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -18,6 +18,7 @@ p256 = { version = "0.13", features = ["pkcs8"] } uuid = { version = "1", features = ["v4"] } rand = "0.9" reqwest = { version = "0.12", features = ["json"] } +rustls = { version = "0.23", default-features = false } serde = { version = "1", features = ["derive"] } serde_json = "1" sha2 = "0.10" @@ -33,6 +34,7 @@ mlua = { version = "0.11", features = ["lua54", "async", "serialize", "vendored" tracing = "0.1" tracing-subscriber = { version = "0.3", features = ["env-filter", "json"] } urlencoding = "2.1.3" +webpki-roots = "0.26" [dev-dependencies] wiremock = "0.6" diff --git a/src/labeler.rs b/src/labeler.rs index 1451db0..0902f59 100644 --- a/src/labeler.rs +++ b/src/labeler.rs @@ -5,6 +5,7 @@ use futures_util::StreamExt; use serde::Deserialize; use tokio::sync::watch; use tokio::task::JoinHandle; +use tokio_tungstenite::Connector; use tokio_tungstenite::tungstenite::Message; use tokio_tungstenite::tungstenite::client::IntoClientRequest; @@ -101,17 +102,22 @@ pub fn spawn(state: AppState, mut subscriptions_rx: watch::Receiver<()>) { // --------------------------------------------------------------------------- async fn run_subscription(state: AppState, did: String) { + let mut backoff_secs: u64 = 2; + const MAX_BACKOFF_SECS: u64 = 300; // 5 minutes + loop { match run_subscription_once(&state, &did).await { Ok(()) => { tracing::info!(did = %did, "labeler subscription ended cleanly, reconnecting"); + backoff_secs = 2; // reset on clean disconnect } Err(e) => { - tracing::warn!(did = %did, "labeler subscription error: {e}"); + tracing::warn!(did = %did, backoff = backoff_secs, "labeler subscription error: {e}"); } } - tokio::time::sleep(std::time::Duration::from_secs(2)).await; + tokio::time::sleep(std::time::Duration::from_secs(backoff_secs)).await; tracing::info!(did = %did, "reconnecting to labeler"); + backoff_secs = (backoff_secs * 2).min(MAX_BACKOFF_SECS); } } @@ -150,12 +156,21 @@ async fn run_subscription_once( let request = url.into_client_request()?; - let (ws, _): ( - tokio_tungstenite::WebSocketStream< - tokio_tungstenite::MaybeTlsStream, - >, - _, - ) = tokio_tungstenite::connect_async(request).await?; + // Build a rustls config that only advertises HTTP/1.1 in ALPN. + // Without this, rustls negotiates h2 and the WebSocket upgrade fails + // ("HTTP version must be 1.1 or higher"). + let root_store = + rustls::RootCertStore::from_iter(webpki_roots::TLS_SERVER_ROOTS.iter().cloned()); + let mut tls_config = rustls::ClientConfig::builder() + .with_root_certificates(root_store) + .with_no_client_auth(); + tls_config.alpn_protocols = vec![b"http/1.1".to_vec()]; + + let connector = Connector::Rustls(Arc::new(tls_config)); + + let (ws, _) = + tokio_tungstenite::connect_async_tls_with_config(request, None, false, Some(connector)) + .await?; tracing::info!(did = %did, "connected to labeler");