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