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)
This commit is contained in:
DanConwayDev
2026-09-12 14:48:26 +00:00
parent 3c77e76a1f
commit 59a37b660a
8 changed files with 198 additions and 25 deletions
+51 -18
View File
@@ -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<Mutex<JoinSet<()>>>,
) -> Result<Response<Full<Bytes>>, 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<SplitSink<_, Message>> = None;
let mut backend_task: Option<tokio::task::JoinHandle<()>> = 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<tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>>,
>,
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");
}
}
+9 -1
View File
@@ -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(());
}
+9 -1
View File
@@ -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(());
}
+9 -1
View File
@@ -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(());
}
+9 -1
View File
@@ -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(());
}
+9 -1
View File
@@ -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(());
}
+44 -2
View File
@@ -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\
+58
View File
@@ -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);
}