From 59a37b660ab358c05dc7c1beaaa689c9bf17781f Mon Sep 17 00:00:00 2001 From: DanConwayDev Date: Sat, 12 Sep 2026 14:48:26 +0000 Subject: [PATCH] test: cancel proxy connections when fixtures stop Stopping accept loops left detached relay/proxy sessions alive. Track HTTP, WebSocket upgrade and forwarding tasks under their owning fixture, cancel them on shutdown, and drain cancellation before explicit stop returns. Preserve censoring, rate limits, authentication and simulated disconnect behavior. Add regressions that observe a live protocol exchange before asserting the connection closes on stop, all under bounded deadlines. Validation: fixture lifecycle checks passed for censoring, REQ/NEG limiting, flapping and setup-drop relays. Auth-gating and upload proxy shutdown regressions passed through relay_identity's common helper tests. Assisted-by: Codex (GPT-6) --- tests/common/auth_gating_relay.rs | 69 ++++++++++++++++------ tests/common/censoring_proxy.rs | 10 +++- tests/common/flapping_relay.rs | 10 +++- tests/common/neg_limiting_proxy.rs | 10 +++- tests/common/req_limiting_proxy.rs | 10 +++- tests/common/setup_drop_relay.rs | 10 +++- tests/common/upload_pack_counting_proxy.rs | 46 ++++++++++++++- tests/fixture_lifecycle.rs | 58 ++++++++++++++++++ 8 files changed, 198 insertions(+), 25 deletions(-) create mode 100644 tests/fixture_lifecycle.rs diff --git a/tests/common/auth_gating_relay.rs b/tests/common/auth_gating_relay.rs index 647ba23..93cd681 100644 --- a/tests/common/auth_gating_relay.rs +++ b/tests/common/auth_gating_relay.rs @@ -35,6 +35,7 @@ use hyper_util::rt::TokioIo; use nostr_sdk::prelude::{Event, Keys, Kind, PublicKey}; use tokio::net::TcpListener; use tokio::sync::oneshot; +use tokio::task::JoinSet; use tokio_tungstenite::tungstenite::protocol::Role; use tokio_tungstenite::tungstenite::Message; use tokio_tungstenite::WebSocketStream; @@ -82,6 +83,8 @@ impl AuthGatingRelay { let backend_url = backend_url.to_string(); let accept_state = state.clone(); let handle = tokio::spawn(async move { + let mut connections = JoinSet::new(); + let upgrades = Arc::new(Mutex::new(JoinSet::new())); loop { tokio::select! { accepted = listener.accept() => { @@ -89,11 +92,13 @@ impl AuthGatingRelay { let backend_url = backend_url.clone(); let state = accept_state.clone(); let io = TokioIo::new(stream); - tokio::spawn(async move { + let upgrades = upgrades.clone(); + connections.spawn(async move { let service = service_fn(move |req| { + let upgrades = upgrades.clone(); let backend_url = backend_url.clone(); let state = state.clone(); - async move { handle_request(req, backend_url, mode, state).await } + async move { handle_request(req, backend_url, mode, state, upgrades).await } }); // Errors are expected when clients disconnect. let _ = http1::Builder::new() @@ -103,8 +108,12 @@ impl AuthGatingRelay { }); } _ = &mut shutdown_rx => break, + _ = connections.join_next(), if !connections.is_empty() => {} } } + connections.shutdown().await; + let mut upgrades = std::mem::take(&mut *upgrades.lock().unwrap()); + upgrades.shutdown().await; }); Self { @@ -140,7 +149,10 @@ impl AuthGatingRelay { let _ = tx.send(()); } if let Some(handle) = self.handle.take() { - let _ = handle.await; + tokio::time::timeout(std::time::Duration::from_secs(5), handle) + .await + .expect("auth gating relay tasks should stop") + .expect("auth gating relay server task should not panic"); } } } @@ -158,6 +170,7 @@ async fn handle_request( backend_url: String, mode: GateMode, state: GateState, + upgrades: Arc>>, ) -> Result>, hyper::Error> { let is_websocket = req .headers() @@ -173,7 +186,9 @@ async fn handle_request( .map(str::to_string) { let accept_key = derive_accept_key(key.as_bytes()); - tokio::spawn(async move { + let mut upgrades = upgrades.lock().unwrap(); + while upgrades.try_join_next().is_some() {} + upgrades.spawn(async move { match hyper::upgrade::on(req).await { Ok(upgraded) => { let ws = WebSocketStream::from_raw_socket( @@ -251,7 +266,7 @@ async fn run_session( let mut authenticated = false; let mut backend: Option> = None; - let mut backend_task: Option> = None; + let mut backend_tasks = JoinSet::new(); while let Some(message) = client_rx.next().await { let message = message.map_err(|e| format!("client read: {e}"))?; @@ -273,7 +288,7 @@ async fn run_session( .map_err(|e| format!("backend connect failed: {e}"))?; let (backend_tx, backend_rx) = backend_ws.split(); backend = Some(backend_tx); - backend_task = Some(spawn_backend_forwarder(backend_rx, client_tx.clone())); + backend_tasks.spawn(forward_backend(backend_rx, client_tx.clone())); } continue; } @@ -342,9 +357,7 @@ async fn run_session( } } - if let Some(task) = backend_task { - task.abort(); - } + backend_tasks.shutdown().await; Ok(()) } @@ -395,20 +408,18 @@ async fn handle_auth( .await } -fn spawn_backend_forwarder( +async fn forward_backend( mut backend_rx: SplitStream< WebSocketStream>, >, client_tx: SharedClientSink, -) -> tokio::task::JoinHandle<()> { - tokio::spawn(async move { - while let Some(message) = backend_rx.next().await { - let Ok(message) = message else { break }; - if client_tx.lock().await.send(message).await.is_err() { - break; - } +) { + while let Some(message) = backend_rx.next().await { + let Ok(message) = message else { break }; + if client_tx.lock().await.send(message).await.is_err() { + break; } - }) + } } /// Derive the Sec-WebSocket-Accept key from the request key. @@ -427,3 +438,25 @@ fn derive_accept_key(request_key: &[u8]) -> String { hash.as_byte_array(), ) } + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn stop_closes_active_websocket_sessions() { + tokio::time::timeout(std::time::Duration::from_secs(5), async { + let gate = AuthGatingRelay::start("ws://127.0.0.1:1", GateMode::Restricted).await; + let (mut client, _) = tokio_tungstenite::connect_async(gate.url()).await.unwrap(); + let challenge = client.next().await.unwrap().unwrap(); + assert!(challenge.to_text().unwrap().contains("AUTH")); + gate.stop().await; + assert!(matches!( + client.next().await, + None | Some(Err(_)) | Some(Ok(Message::Close(_))) + )); + }) + .await + .expect("gate stop must close active sessions"); + } +} diff --git a/tests/common/censoring_proxy.rs b/tests/common/censoring_proxy.rs index a986e79..179f2bf 100644 --- a/tests/common/censoring_proxy.rs +++ b/tests/common/censoring_proxy.rs @@ -61,6 +61,7 @@ impl CensoringProxy { let accept_dropped = dropped.clone(); let handle = tokio::spawn(async move { + let mut connections = tokio::task::JoinSet::new(); loop { tokio::select! { accepted = listener.accept() => { @@ -68,7 +69,7 @@ impl CensoringProxy { let backend_url = backend_url.clone(); let withheld = accept_withheld.clone(); let dropped = accept_dropped.clone(); - tokio::spawn(async move { + connections.spawn(async move { if let Err(error) = proxy_connection(stream, &backend_url, withheld, dropped).await { @@ -77,9 +78,13 @@ impl CensoringProxy { } }); } + result = connections.join_next(), if !connections.is_empty() => { + result.expect("connection task").expect("fixture connection panicked"); + } _ = &mut shutdown_rx => break, } } + connections.shutdown().await; }); Self { @@ -131,6 +136,9 @@ impl CensoringProxy { impl Drop for CensoringProxy { fn drop(&mut self) { + if let Some(handle) = self.handle.take() { + handle.abort(); + } if let Some(tx) = self.shutdown_tx.take() { let _ = tx.send(()); } diff --git a/tests/common/flapping_relay.rs b/tests/common/flapping_relay.rs index cd08627..897a5fb 100644 --- a/tests/common/flapping_relay.rs +++ b/tests/common/flapping_relay.rs @@ -31,13 +31,14 @@ impl FlappingRelay { let notify = connection_observed.clone(); let handle = tokio::spawn(async move { + let mut connections = tokio::task::JoinSet::new(); loop { tokio::select! { accepted = listener.accept() => { let Ok((stream, _)) = accepted else { break }; let observed = observed.clone(); let notify = notify.clone(); - tokio::spawn(async move { + connections.spawn(async move { let Ok(mut websocket) = tokio_tungstenite::accept_async(stream).await else { // NIP-11 HTTP probes are expected to fail this @@ -61,9 +62,13 @@ impl FlappingRelay { } }); } + result = connections.join_next(), if !connections.is_empty() => { + result.expect("connection task").expect("fixture connection panicked"); + } _ = &mut shutdown_rx => break, } } + connections.shutdown().await; }); Self { @@ -112,6 +117,9 @@ impl FlappingRelay { impl Drop for FlappingRelay { fn drop(&mut self) { + if let Some(handle) = self.handle.take() { + handle.abort(); + } if let Some(tx) = self.shutdown_tx.take() { let _ = tx.send(()); } diff --git a/tests/common/neg_limiting_proxy.rs b/tests/common/neg_limiting_proxy.rs index aa4e8b6..f94a7a0 100644 --- a/tests/common/neg_limiting_proxy.rs +++ b/tests/common/neg_limiting_proxy.rs @@ -74,6 +74,7 @@ impl NegLimitingProxy { let accept_opened = opened.clone(); let handle = tokio::spawn(async move { + let mut connections = tokio::task::JoinSet::new(); loop { tokio::select! { accepted = listener.accept() => { @@ -82,7 +83,7 @@ impl NegLimitingProxy { let peak = accept_peak.clone(); let rejected = accept_rejected.clone(); let opened = accept_opened.clone(); - tokio::spawn(async move { + connections.spawn(async move { if let Err(error) = proxy_connection( stream, &backend_url, @@ -98,9 +99,13 @@ impl NegLimitingProxy { } }); } + result = connections.join_next(), if !connections.is_empty() => { + result.expect("connection task").expect("fixture connection panicked"); + } _ = &mut shutdown_rx => break, } } + connections.shutdown().await; }); Self { @@ -146,6 +151,9 @@ impl NegLimitingProxy { impl Drop for NegLimitingProxy { fn drop(&mut self) { + if let Some(handle) = self.handle.take() { + handle.abort(); + } if let Some(tx) = self.shutdown_tx.take() { let _ = tx.send(()); } diff --git a/tests/common/req_limiting_proxy.rs b/tests/common/req_limiting_proxy.rs index 22eaa35..1d40dd3 100644 --- a/tests/common/req_limiting_proxy.rs +++ b/tests/common/req_limiting_proxy.rs @@ -83,6 +83,7 @@ impl ReqLimitingProxy { let accept_opened = opened.clone(); let handle = tokio::spawn(async move { + let mut connections = tokio::task::JoinSet::new(); loop { tokio::select! { accepted = listener.accept() => { @@ -91,7 +92,7 @@ impl ReqLimitingProxy { let peak = accept_peak.clone(); let rejected = accept_rejected.clone(); let opened = accept_opened.clone(); - tokio::spawn(async move { + connections.spawn(async move { if let Err(error) = proxy_connection( stream, &backend_url, @@ -107,9 +108,13 @@ impl ReqLimitingProxy { } }); } + result = connections.join_next(), if !connections.is_empty() => { + result.expect("connection task").expect("fixture connection panicked"); + } _ = &mut shutdown_rx => break, } } + connections.shutdown().await; }); Self { @@ -156,6 +161,9 @@ impl ReqLimitingProxy { impl Drop for ReqLimitingProxy { fn drop(&mut self) { + if let Some(handle) = self.handle.take() { + handle.abort(); + } if let Some(tx) = self.shutdown_tx.take() { let _ = tx.send(()); } diff --git a/tests/common/setup_drop_relay.rs b/tests/common/setup_drop_relay.rs index ab27b5e..411b0ff 100644 --- a/tests/common/setup_drop_relay.rs +++ b/tests/common/setup_drop_relay.rs @@ -34,6 +34,7 @@ impl SetupDropRelay { let (shutdown_tx, mut shutdown_rx) = oneshot::channel(); let handle = tokio::spawn(async move { + let mut connections = tokio::task::JoinSet::new(); loop { tokio::select! { accepted = listener.accept() => { @@ -43,7 +44,7 @@ impl SetupDropRelay { let dropped_tx = dropped_tx.clone(); let dropped_rx = dropped_rx.clone(); let active_websockets = active_websockets.clone(); - tokio::spawn(async move { + connections.spawn(async move { handle_connection( stream, nip11_tx, @@ -55,9 +56,13 @@ impl SetupDropRelay { .await; }); } + result = connections.join_next(), if !connections.is_empty() => { + result.expect("connection task").expect("fixture connection panicked"); + } _ = &mut shutdown_rx => break, } } + connections.shutdown().await; }); Self { @@ -142,6 +147,9 @@ async fn handle_connection( impl Drop for SetupDropRelay { fn drop(&mut self) { + if let Some(handle) = self.handle.take() { + handle.abort(); + } if let Some(tx) = self.shutdown_tx.take() { let _ = tx.send(()); } diff --git a/tests/common/upload_pack_counting_proxy.rs b/tests/common/upload_pack_counting_proxy.rs index 8f38e30..8516c3a 100644 --- a/tests/common/upload_pack_counting_proxy.rs +++ b/tests/common/upload_pack_counting_proxy.rs @@ -31,6 +31,7 @@ use hyper::{Request, Response, StatusCode}; use hyper_util::rt::TokioIo; use tokio::net::TcpListener; use tokio::sync::oneshot; +use tokio::task::JoinSet; /// One observed `POST /git-upload-pack` exchange. #[derive(Debug, Clone)] @@ -77,6 +78,7 @@ impl UploadPackCountingProxy { let accept_info_refs = info_refs.clone(); let handle = tokio::spawn(async move { + let mut connections = JoinSet::new(); // One shared upstream client; keep it plain so request and // response bodies pass through byte-for-byte. let client = reqwest::Client::builder() @@ -92,7 +94,7 @@ impl UploadPackCountingProxy { let client = client.clone(); let exchanges = accept_exchanges.clone(); let info_refs = accept_info_refs.clone(); - tokio::spawn(async move { + connections.spawn(async move { let io = TokioIo::new(stream); let service = service_fn(move |req| { let backend_url = backend_url.clone(); @@ -118,8 +120,10 @@ impl UploadPackCountingProxy { }); } _ = &mut shutdown_rx => break, + _ = connections.join_next(), if !connections.is_empty() => {} } } + connections.shutdown().await; }); Self { @@ -155,7 +159,10 @@ impl UploadPackCountingProxy { let _ = tx.send(()); } if let Some(handle) = self.handle.take() { - let _ = handle.await; + tokio::time::timeout(std::time::Duration::from_secs(5), handle) + .await + .expect("upload-pack proxy tasks should stop") + .expect("upload-pack proxy server task should not panic"); } } } @@ -291,6 +298,41 @@ fn contains(haystack: &[u8], needle: &[u8]) -> bool { mod tests { use super::*; + #[tokio::test] + async fn stop_closes_active_forwarded_requests() { + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + tokio::time::timeout(std::time::Duration::from_secs(5), async { + let backend = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let proxy = UploadPackCountingProxy::start(&format!( + "http://{}", + backend.local_addr().unwrap() + )) + .await; + let address = proxy.url().strip_prefix("http://").unwrap(); + let mut client = tokio::net::TcpStream::connect(address).await.unwrap(); + client + .write_all( + b"GET /info/refs?service=git-upload-pack HTTP/1.1\r\nHost: localhost\r\n\r\n", + ) + .await + .unwrap(); + let (mut upstream, _) = backend.accept().await.unwrap(); + let mut request = [0; 256]; + assert!(upstream.read(&mut request).await.unwrap() > 0); + // The request has reached the backend, which deliberately retains + // the connection without replying while the proxy is stopped. + proxy.stop().await; + let mut byte = [0]; + match client.read(&mut byte).await { + Ok(0) => {} + Err(error) if error.kind() == std::io::ErrorKind::ConnectionReset => {} + result => panic!("proxy retained an active request after stop: {result:?}"), + } + }) + .await + .expect("proxy stop must close active requests"); + } + #[test] fn parse_wants_extracts_oids_from_pkt_lines() { let body = b"0032want aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa\n\ diff --git a/tests/fixture_lifecycle.rs b/tests/fixture_lifecycle.rs new file mode 100644 index 0000000..266ca1e --- /dev/null +++ b/tests/fixture_lifecycle.rs @@ -0,0 +1,58 @@ +//! Accepted connections must not outlive the fixture that owns them. +mod common; + +use std::time::Duration; + +use common::{ + censoring_proxy::CensoringProxy, flapping_relay::FlappingRelay, + neg_limiting_proxy::NegLimitingProxy, req_limiting_proxy::ReqLimitingProxy, + setup_drop_relay::SetupDropRelay, MockRelay, +}; +use futures_util::StreamExt; + +macro_rules! check_shutdown { + ($fixture:expr) => {{ + tokio::time::timeout(Duration::from_secs(10), async { + let fixture = $fixture; + let (mut connection, _) = tokio_tungstenite::connect_async(fixture.url()) + .await + .expect("fixture must finish the WebSocket handshake"); + fixture.stop().await; + while let Some(message) = connection.next().await { + if message.is_err() || message.unwrap().is_close() { + break; + } + } + }) + .await + .expect("fixture stop must close accepted connections"); + }}; +} + +#[tokio::test] +async fn censoring_proxy_closes_connections_on_stop() { + let backend = MockRelay::start().await; + check_shutdown!(CensoringProxy::start(backend.url()).await); +} + +#[tokio::test] +async fn req_limiting_proxy_closes_connections_on_stop() { + let backend = MockRelay::start().await; + check_shutdown!(ReqLimitingProxy::start(backend.url(), 2).await); +} + +#[tokio::test] +async fn neg_limiting_proxy_closes_connections_on_stop() { + let backend = MockRelay::start().await; + check_shutdown!(NegLimitingProxy::start(backend.url(), 2).await); +} + +#[tokio::test] +async fn flapping_relay_closes_connections_on_stop() { + check_shutdown!(FlappingRelay::start().await); +} + +#[tokio::test] +async fn setup_drop_relay_closes_connections_on_stop() { + check_shutdown!(SetupDropRelay::start().await); +}