diff --git a/src/sync/relay_connection.rs b/src/sync/relay_connection.rs index 429fa94..75e175b 100644 --- a/src/sync/relay_connection.rs +++ b/src/sync/relay_connection.rs @@ -2833,8 +2833,12 @@ impl RelayConnection { } } +#[cfg(test)] +mod test_relay; + #[cfg(test)] mod tests { + use super::test_relay::TestRelay; use super::*; #[test] @@ -3111,10 +3115,8 @@ mod tests { #[tokio::test] async fn fetch_events_targets_the_connections_exact_relay() { - let configured = LocalRelayBuilder::default().build(); - configured.run().await.expect("start configured relay"); - let other = LocalRelayBuilder::default().build(); - other.run().await.expect("start other relay"); + let configured = TestRelay::start(LocalRelayBuilder::default()).await; + let other = TestRelay::start(LocalRelayBuilder::default()).await; let expected = EventBuilder::new(Kind::TextNote, "only on the other relay") .finalize(&Keys::generate()) @@ -3167,8 +3169,7 @@ mod tests { const AUTHORS: usize = 3; const EVENTS_PER_AUTHOR: usize = 400; - let relay = LocalRelayBuilder::default().build(); - relay.run().await.expect("start burst relay"); + let relay = TestRelay::start(LocalRelayBuilder::default()).await; let mut authors = Vec::with_capacity(AUTHORS); for author_index in 0..AUTHORS { let keys = Keys::generate(); @@ -3229,8 +3230,7 @@ mod tests { #[tokio::test] async fn immediate_empty_eose_cannot_arrive_before_permit_registration() { - let relay = LocalRelayBuilder::default().build(); - relay.run().await.expect("start empty relay"); + let relay = TestRelay::start(LocalRelayBuilder::default()).await; let connection = RelayConnection::new( relay.url().await.to_string(), Some(Keys::generate()), @@ -3276,8 +3276,7 @@ mod tests { #[tokio::test] async fn peer_closed_subscription_is_removed_from_sdk_registry() { - let relay = LocalRelayBuilder::default().build(); - relay.run().await.expect("start registry relay"); + let relay = TestRelay::start(LocalRelayBuilder::default()).await; let connection = RelayConnection::new( relay.url().await.to_string(), Some(Keys::generate()), @@ -3311,8 +3310,7 @@ mod tests { #[tokio::test] async fn query_rate_closed_paces_later_wire_requests() { - let relay = LocalRelayBuilder::default().queries_per_minute(1).build(); - relay.run().await.expect("start query-limited relay"); + let relay = TestRelay::start(LocalRelayBuilder::default().queries_per_minute(1)).await; let connection = RelayConnection::new( relay.url().await.to_string(), Some(Keys::generate()), @@ -3816,8 +3814,7 @@ mod tests { #[tokio::test] async fn minimum_churn_extension_keeps_full_and_auxiliary_subscriptions_open() { - let relay = LocalRelayBuilder::default().build(); - relay.run().await.expect("start local relay"); + let relay = TestRelay::start(LocalRelayBuilder::default()).await; let connection = RelayConnection::new( relay.url().await.to_string(), Some(Keys::generate()), @@ -3879,8 +3876,7 @@ mod tests { #[tokio::test] async fn failed_minimum_churn_extension_restores_only_the_retired_tail() { - let relay = LocalRelayBuilder::default().build(); - relay.run().await.expect("start local relay"); + let relay = TestRelay::start(LocalRelayBuilder::default()).await; let connection = RelayConnection::new( relay.url().await.to_string(), Some(Keys::generate()), @@ -3971,8 +3967,7 @@ mod tests { #[tokio::test] async fn partial_close_failure_restores_the_already_closed_tail_group() { - let relay = LocalRelayBuilder::default().build(); - relay.run().await.expect("start local relay"); + let relay = TestRelay::start(LocalRelayBuilder::default()).await; let connection = RelayConnection::new( relay.url().await.to_string(), Some(Keys::generate()), @@ -4414,8 +4409,7 @@ mod tests { #[tokio::test] async fn immediate_transient_capacity_rejects_a_saturated_session() { - let relay = LocalRelayBuilder::default().build(); - relay.run().await.expect("start local relay"); + let relay = TestRelay::start(LocalRelayBuilder::default()).await; let relay_url = relay.url().await.to_string(); let connection = permissive_connection(&relay_url, Keys::generate()); connection.connect(3).await.expect("connect local relay"); diff --git a/src/sync/relay_connection/test_relay.rs b/src/sync/relay_connection/test_relay.rs new file mode 100644 index 0000000..c7cde68 --- /dev/null +++ b/src/sync/relay_connection/test_relay.rs @@ -0,0 +1,136 @@ +//! Local relay fixture that owns its listener from allocation through shutdown. + +use std::convert::Infallible; +use std::ops::Deref; +use std::sync::{Arc, Mutex}; +use std::time::Duration; + +use http_body_util::Full; +use hyper::body::Bytes; +use hyper::service::service_fn; +use hyper::{Request, Response}; +use hyper_util::rt::TokioIo; +use nostr_sdk::prelude::{LocalRelay, LocalRelayBuilder}; +use tokio::net::TcpListener; +use tokio::task::{JoinHandle, JoinSet}; +use tokio_tungstenite::tungstenite::handshake::derive_accept_key; + +pub(super) struct TestRelay { + relay: LocalRelay, + url: String, + task: JoinHandle<()>, +} + +impl TestRelay { + pub(super) async fn start(builder: LocalRelayBuilder) -> Self { + // The SDK's run() probes and releases a port before binding it again. + // Keep this listener bound and give upgraded streams to the SDK instead. + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("bind test relay"); + let url = format!("ws://{}", listener.local_addr().expect("relay address")); + let relay = builder.build(); + let server = relay.clone(); + let task = tokio::spawn(async move { + let mut connections = JoinSet::new(); + loop { + tokio::select! { + accepted = listener.accept() => { + let (stream, addr) = accepted.expect("accept test relay connection"); + let relay = server.clone(); + connections.spawn(async move { + let upgrade = Arc::new(Mutex::new(None)); + let requested_upgrade = upgrade.clone(); + let service = service_fn(move |req: Request| { + let key = req.headers().get("sec-websocket-key").cloned(); + let response = if let Some(key) = key { + let accept = derive_accept_key(key.as_bytes()); + *requested_upgrade.lock().expect("upgrade lock") = + Some(hyper::upgrade::on(req)); + Response::builder() + .status(101) + .header("connection", "upgrade") + .header("upgrade", "websocket") + .header("sec-websocket-accept", accept) + } else { + Response::builder().status(400) + }; + async move { + Ok::<_, Infallible>(response.body(Full::new(Bytes::new())).unwrap()) + } + }); + let http = hyper::server::conn::http1::Builder::new() + .serve_connection(TokioIo::new(stream), service) + .with_upgrades(); + if !matches!(tokio::time::timeout(Duration::from_secs(5), http).await, Ok(Ok(()))) { + return; + } + let upgrade = upgrade.lock().expect("upgrade lock").take(); + if let Some(upgrade) = upgrade { + if let Ok(Ok(stream)) = tokio::time::timeout(Duration::from_secs(5), upgrade).await { + let _ = relay.take_connection(TokioIo::new(stream), addr).await; + } + } + }); + } + result = connections.join_next(), if !connections.is_empty() => { + result.expect("connection task").expect("connection task panicked"); + } + } + } + }); + Self { relay, url, task } + } + + pub(super) async fn url(&self) -> String { + self.url.clone() + } + + pub(super) fn shutdown(&self) { + self.task.abort(); + self.relay.shutdown(); + } +} + +impl Deref for TestRelay { + type Target = LocalRelay; + + fn deref(&self) -> &Self::Target { + &self.relay + } +} + +impl Drop for TestRelay { + fn drop(&mut self) { + // Aborting the listener also drops its JoinSet and aborts connections. + self.shutdown(); + } +} + +#[tokio::test] +async fn concurrent_relays_keep_distinct_reachable_listeners() { + let mut starts = JoinSet::new(); + for _ in 0..32 { + starts.spawn(TestRelay::start(LocalRelayBuilder::default())); + } + let mut relays = Vec::new(); + let mut urls = std::collections::HashSet::new(); + while let Some(result) = starts.join_next().await { + let relay = result.expect("start task"); + assert!(urls.insert(relay.url().await.to_owned())); + relays.push(relay); + } + for relay in &relays { + let (mut socket, _) = tokio::time::timeout( + Duration::from_secs(5), + tokio_tungstenite::connect_async(relay.url().await), + ) + .await + .expect("handshake deadline") + .expect("connect test relay"); + tokio::time::timeout(Duration::from_secs(5), socket.close(None)) + .await + .expect("close deadline") + .expect("close test connection"); + } +}