lib

Core libraries for Radroots
git clone https://radroots.dev/git/lib.git
Log | Files | Refs | README

commit 189c49b74b4bafc142b00b76b296477931139e72
parent 8dbf27459be1729709a0ea36bba7470e90479ca3
Author: triesap <tyson@radroots.org>
Date:   Thu,  1 Oct 2026 23:27:37 +0000

transport_nostr: Bound raw fetch ingress before decoding

- Meter decrypted bytes and WebSocket frame work with shared fetch limits.
- Revoke cancelled generations and serialize completion with budget admission.
- Preserve TLS validation, public APIs, and previously collected relay evidence.
- Verify malformed traffic, lifecycle races, and coverage with committed regressions.

Diffstat:
MCargo.lock | 3+++
MCargo.toml | 2++
Mcontracts/architecture/decisions/nostr_fetch_bounds.v1.json | 13++++++++++---
Mcrates/transport_nostr/Cargo.toml | 3+++
Mcrates/transport_nostr/README.md | 34+++++++++++++++++++++++++++++-----
Mcrates/transport_nostr/src/client.rs | 6+++++-
Mcrates/transport_nostr/src/lib.rs | 2++
Mcrates/transport_nostr/src/relay.rs | 129++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++-------
Mcrates/transport_nostr/src/source.rs | 597++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++---
Mcrates/transport_nostr/src/source_budget.rs | 104++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++---------
Acrates/transport_nostr/src/source_ingress.rs | 439+++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
Acrates/transport_nostr/src/source_wire.rs | 231+++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
Acrates/transport_nostr/src/source_wire_tests.rs | 334+++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
Mcrates/transport_nostr/tests/fetch_bounds.rs | 9+++++----
Acrates/transport_nostr/tests/fetch_raw_budget.rs | 290+++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
Mcrates/transport_nostr/tests/network_hardening.rs | 8+++++++-
Mcrates/transport_nostr/tests/package_boundary.rs | 5+++++
17 files changed, 2153 insertions(+), 56 deletions(-)

diff --git a/Cargo.lock b/Cargo.lock @@ -3557,11 +3557,14 @@ dependencies = [ "radroots_nostr", "radroots_protocol", "radroots_transport", + "rustls", "serde_json", "sha2", "tokio", + "tokio-rustls", "tokio-tungstenite", "url", + "webpki-roots 0.26.11", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml @@ -253,6 +253,7 @@ tempfile = { version = "3" } tar = { version = "0.4" } thiserror = { version = "1" } tokio = { version = "1" } +tokio-rustls = { version = "0.26.4", default-features = false } tokio-util = { version = "0.7", features = ["rt"] } tokio-tungstenite = { version = "0.26.2", default-features = false, features = [ "connect", @@ -268,6 +269,7 @@ uuid = { version = "1.22.0", features = ["v4", "v7"] } uniffi = { version = "0.29.4" } wasm-bindgen = { version = "0.2" } wasm-bindgen-test = { version = "0.3" } +webpki-roots = { version = "0.26.11" } walkdir = { version = "2" } x509-parser = { version = "0.17", default-features = false } zstd = { version = "0.13", default-features = false } diff --git a/contracts/architecture/decisions/nostr_fetch_bounds.v1.json b/contracts/architecture/decisions/nostr_fetch_bounds.v1.json @@ -13,12 +13,19 @@ "max_events_per_fetch": 4096, "max_event_json_bytes_per_fetch": 8388608, "max_notifications_per_fetch": 8192, + "max_raw_incoming_bytes_per_fetch": 8388608, + "max_raw_frames_per_fetch": 8192, + "max_raw_data_message_attempts_per_fetch": 4096, + "max_http_upgrade_bytes": 16384, + "max_active_fetch_registrations_per_transport": 64, + "predecode_scratch_bytes": 4096, + "physical_read_reserve": "The meter admits at most 8 MiB to WebSocket framing and Nostr decoding. One fixed 4096-byte plaintext scratch read may observe bytes beyond the remaining allowance and discard the entire chunk before forwarding. Kernel and TLS implementation buffers remain outside this plaintext admission claim.", "max_returned_events": 1000, "deadline": "The minimum of the caller absolute deadline and configured request timeout is frozen once before relay scheduling. Connect, REQ, collection and queued starts consume its remaining duration; no later relay receives a fresh request timeout.", "collection": "Charge event inventory and JSON bytes across all relays before retaining raw events; duplicates and malformed events consume work. Defensively enforce the same limits before shared event decoding and candidate collection.", - "parse_boundary": "The WebSocket message/frame bound precedes upstream message decoding. The shared event-JSON bound and aggregate inventory precede canonical decoding. Notification processing is finite; upstream socket tasks remain subject to bounded messages and request auto-close.", + "parse_boundary": "Meter decrypted incoming bytes before WebSocket assembly and Nostr decoding, including bounded HTTP upgrade bytes and every frame header and payload. Count every frame start, including control and zero-length continuation frames; conservatively count every text/binary message start as a data attempt without parsing JSON. Share one monotonic byte/frame/attempt budget across all selected relays, queued batches and connection generations. Malformed, discarded, duplicate and unsolicited inputs consume this budget. Retain independent SDK notification and canonical event-JSON limits.", "completion": "EOSE confirms only the selected relay subscription within the request budget. Exhausted resources remain Partial, elapsed deadlines remain Cancelled, and neither implies global-history completeness.", "partial_finalization": "After bounded network collection, bounded local normalization retains previously collected admissible events and each relay outcome. One slow relay must not discard another relay's earlier evidence. No further network request is started during finalization.", - "cancellation": "Unpolled futures perform no I/O. Dropping polled work cannot await cleanup; published subscriptions retain the original bounded auto-close deadline. No durable admission or publication rollback is claimed.", - "compatibility": "No new public Rust type, feature, dependency, transport or product policy is required. Callers retain bounded FetchRequest and existing typed outcomes; large or continuous results can now terminate earlier as partial evidence." + "cancellation": "Unpolled futures perform no I/O. Cancellation, deadline or resource failure invalidates selected receive generations and signals SDK disconnect with reconnect disabled, including previously completed batches. SDK task teardown remains asynchronous; no task-join claim is made. Normal completion removes the registration and preserves shared authenticated sockets. Overlapping subscription, delivery and authentication operations may truthfully observe disconnection. Fresh explicit operations may establish new generations, while old generations remain closed. No durable admission or publication rollback is claimed.", + "compatibility": "No new public Rust type, feature, transport or product policy is introduced. The private connector directly uses already-locked rustls, tokio-rustls and WebPKI roots to meter decrypted data while preserving destination pinning, original-host SNI and certificate verification. Callers retain bounded FetchRequest and existing typed outcomes; resource exhaustion is sticky Partial evidence and cannot become Complete after EOSE." } diff --git a/crates/transport_nostr/Cargo.toml b/crates/transport_nostr/Cargo.toml @@ -46,11 +46,14 @@ async-wsocket = { workspace = true } futures = { workspace = true } nostr-sdk = { workspace = true } nostr-relay-pool = { workspace = true } +rustls = { workspace = true } serde_json = { workspace = true, features = ["std"] } sha2 = { workspace = true, default-features = false } tokio = { workspace = true, features = ["net", "time"] } +tokio-rustls = { workspace = true } tokio-tungstenite = { workspace = true } url = { workspace = true } +webpki-roots = { workspace = true } [dev-dependencies] tokio = { workspace = true, features = [ diff --git a/crates/transport_nostr/README.md b/crates/transport_nostr/README.md @@ -161,11 +161,24 @@ strictly before its absolute deadline. Deadline expiry is `Cancelled`, and a relay result that exceeds the bounded inventory is `Partial`; neither state is rewritten as completion even when it carries admissible events. -One fetch shares an 8 MiB raw event-JSON budget, a 4,096-event inventory and an -8,192-notification work limit across all relay batches. Each relay retains at -most 1,000 events; duplicate and malformed observations still consume the -budget. An event is at most 256 KiB, and the WebSocket connector explicitly -limits both frames and complete messages to 512 KiB before upstream decoding. +One fetch shares an 8 MiB incoming-byte budget, a 4,096-data-message attempt +limit and an 8,192-frame work limit across all selected relays and connection +generations. Accounting precedes WebSocket message assembly and Nostr decoding: +it includes decrypted HTTP upgrade bytes, frame headers and payloads, malformed +and discarded messages, duplicate events, unsolicited traffic, control frames, +and empty continuation frames. Text and binary message starts conservatively +consume data attempts without inspecting their JSON. The HTTP upgrade response +is also limited to 16 KiB. Each relay retains at most 1,000 events. An event is +at most 256 KiB, and the WebSocket connector explicitly limits both frames and +complete messages to 512 KiB. SDK notification and canonical event collection +retain independent limits of 8,192 notifications, 4,096 events and 8 MiB. +Exhaustion is sticky and wakes all affected fetch collectors; later EOSE cannot +turn it into completion. At most 64 fetch calls may hold ingress registrations +on one transport; additional calls fail closed as partial evidence. +The meter reads into a fixed 4 KiB plaintext scratch buffer before admission; +at most that scratch capacity can be read beyond the remaining allowance and +discarded without reaching WebSocket assembly or Nostr parsing. This bound +does not claim control over kernel or TLS implementation buffers. Shared canonical decoding rechecks the event, aggregate-byte and inventory bounds before candidate collection. A fetch returns at most the caller's validated 1,000-event page limit. EOSE describes only the requested relay @@ -240,6 +253,17 @@ response is therefore reported as unavailable or unknown evidence, never as proof that no publication occurred. Fetch and live observation have no local durable commit point. +A cancelled, timed-out or resource-limited fetch invalidates its selected +connection generations and signals SDK disconnection with reconnect disabled. +This includes connections whose earlier relay batch already completed; their +collected evidence remains intact. Receive admission stops at the generation +boundary, while ordinary SDK task teardown is asynchronous. These authenticated +sockets are shared with subscriptions, delivery and authentication, so those +overlapping operations can also observe disconnection. A successful fetch +removes its registration and preserves the shared connections. A later explicit +operation may establish a fresh connection and receives a fresh fetch budget; +closed generations never regain receive authority. + Dropping a pending subscription `next` or `cancel` future records a cancellation request in the retained capability; its next operation awaits relay unsubscription and returns the stable cancelled terminal result. diff --git a/crates/transport_nostr/src/client.rs b/crates/transport_nostr/src/client.rs @@ -230,6 +230,7 @@ impl NostrTransport { pub fn new(config: Config) -> Self { let connector = crate::relay::HardenedWebsocketTransport::new(config.endpoints()); let writers = connector.writers.clone(); + let ingress = connector.ingress.clone(); let client = nostr_sdk::Client::builder() .websocket_transport(connector) .build(); @@ -238,7 +239,10 @@ impl NostrTransport { Self { config, client: Arc::new(crate::sink::LiveRelayClient::new(client.clone(), writers)), - source_client: Arc::new(crate::source::LiveRelaySourceClient::new(client.clone())), + source_client: Arc::new(crate::source::LiveRelaySourceClient::new( + client.clone(), + ingress, + )), subscription_client: Arc::new(crate::subscription::LiveRelaySubscriptionClient::new( client.clone(), )), diff --git a/crates/transport_nostr/src/lib.rs b/crates/transport_nostr/src/lib.rs @@ -12,6 +12,8 @@ mod relay; mod sink; mod socket_write; mod source; +mod source_ingress; +mod source_wire; mod status; mod subscription; diff --git a/crates/transport_nostr/src/relay.rs b/crates/transport_nostr/src/relay.rs @@ -1,9 +1,9 @@ //! Nostr relay identifiers and network policy. use crate::{Error, RelayEndpoint}; +use async_wsocket::Message; use async_wsocket::futures_util::stream::SplitSink; -use async_wsocket::futures_util::{Sink, SinkExt, StreamExt, TryStreamExt}; -use async_wsocket::{Message, WebSocket}; +use async_wsocket::futures_util::{Sink, SinkExt, StreamExt}; use core::fmt; use core::pin::Pin; use nostr_relay_pool::ConnectionMode; @@ -16,6 +16,8 @@ use std::sync::Arc; use std::task::{Context, Poll}; use std::time::Duration; use tokio::net::TcpStream; +use tokio_tungstenite::tungstenite::Message as WireMessage; +use tokio_tungstenite::{MaybeTlsStream, WebSocketStream as TokioWebSocketStream}; use url::Url; const MAX_RESOLVED_ADDRESSES: usize = 32; @@ -98,11 +100,17 @@ impl fmt::Display for RelayUrl { pub(crate) struct HardenedWebsocketTransport { policies: Arc<BTreeMap<String, RelayUrlPolicy>>, pub(crate) writers: crate::socket_write::WriterRegistry, + pub(crate) ingress: crate::source_ingress::IngressRegistry, } impl HardenedWebsocketTransport { pub(crate) fn new(endpoints: &[RelayEndpoint]) -> Self { Self { + ingress: crate::source_ingress::IngressRegistry::new( + endpoints + .iter() + .map(|endpoint| endpoint.url().as_str().to_owned()), + ), writers: crate::socket_write::WriterRegistry::new( endpoints .iter() @@ -155,31 +163,53 @@ impl WebSocketTransport for HardenedWebsocketTransport { .ok_or_else(|| policy_error("relay URL port is missing"))?; let connect = async { + let ingress = self + .ingress + .connection(relay.as_str()) + .ok_or_else(|| policy_error("relay ingress unavailable"))?; let addresses = resolve_bounded(host, port).await?; relay .validate_resolved_addresses(policy, addresses.iter().map(SocketAddr::ip)) .map_err(|_| policy_error("relay DNS result is denied by network policy"))?; let tcp = connect_pinned(addresses.as_slice()).await?; + let plaintext = if parsed.scheme() == "wss" { + let mut roots = rustls::RootCertStore::empty(); + roots.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned()); + let tls_config = rustls::ClientConfig::builder() + .with_root_certificates(roots) + .with_no_client_auth(); + let server_name = rustls::pki_types::ServerName::try_from( + host.trim_matches(['[', ']']).to_owned(), + ) + .map_err(|_| policy_error("relay TLS server name is invalid"))?; + let tls = tokio_rustls::TlsConnector::from(Arc::new(tls_config)) + .connect(server_name, tcp) + .await + .map_err(TransportError::backend)?; + MaybeTlsStream::Rustls(tls) + } else { + MaybeTlsStream::Plain(tcp) + }; + let metered = crate::source_wire::MeteredIo::new(plaintext, ingress); let config = tokio_tungstenite::tungstenite::protocol::WebSocketConfig::default() .max_message_size(Some(MAX_WIRE_MESSAGE_BYTES)) .max_frame_size(Some(MAX_WIRE_MESSAGE_BYTES)); - let (stream, _) = tokio_tungstenite::client_async_tls_with_config( + let (stream, _) = tokio_tungstenite::client_async_with_config( relay.as_str(), - tcp, + metered, Some(config), - None, ) .await .map_err(TransportError::backend)?; - let socket = WebSocket::Tokio(stream); - let (tx, rx) = socket.split(); + let (tx, rx) = stream.split(); let writer = crate::socket_write::SocketWriter::new(Box::new(HardenedTransportSink(tx))); self.writers.install(relay.as_str(), &writer)?; let sink: WebSocketSink = Box::new(crate::socket_write::SharedSocketSink::new(writer)); - let stream: WebSocketStream = - Box::pin(rx.map_err(TransportError::backend)) as WebSocketStream; + let stream: WebSocketStream = Box::pin( + rx.map(|result| result.map_err(TransportError::backend).and_then(from_wire)), + ); Ok((sink, stream)) }; @@ -243,7 +273,20 @@ pub(crate) fn policy_error(message: &'static str) -> TransportError { TransportError::backend(NetworkPolicyError(message)) } -struct HardenedTransportSink(SplitSink<WebSocket, Message>); +type MeteredSocket = TokioWebSocketStream<crate::source_wire::MeteredIo<MaybeTlsStream<TcpStream>>>; + +struct HardenedTransportSink(SplitSink<MeteredSocket, WireMessage>); + +fn from_wire(message: WireMessage) -> Result<Message, TransportError> { + Ok(match message { + WireMessage::Text(text) => Message::Text(text.to_string()), + WireMessage::Binary(bytes) => Message::Binary(bytes.to_vec()), + WireMessage::Ping(bytes) => Message::Ping(bytes.to_vec()), + WireMessage::Pong(bytes) => Message::Pong(bytes.to_vec()), + WireMessage::Close(frame) => Message::Close(frame.map(Into::into)), + WireMessage::Frame(_) => return Err(policy_error("unexpected raw WebSocket frame")), + }) +} impl Sink<Message> for HardenedTransportSink { type Error = TransportError; @@ -261,7 +304,7 @@ impl Sink<Message> for HardenedTransportSink { #[cfg_attr(coverage_nightly, coverage(off))] fn start_send(mut self: Pin<&mut Self>, item: Message) -> Result<(), Self::Error> { Pin::new(&mut self.0) - .start_send_unpin(item) + .start_send_unpin(item.into()) .map_err(TransportError::backend) } @@ -393,6 +436,70 @@ mod tests { use super::*; #[test] + fn metered_socket_conversion_preserves_data_control_and_close_messages() { + for message in [ + WireMessage::Text("text".into()), + WireMessage::Binary(vec![0, 255].into()), + WireMessage::Ping(vec![1, 2].into()), + WireMessage::Pong(vec![3].into()), + WireMessage::Close(None), + WireMessage::Close(Some(tokio_tungstenite::tungstenite::protocol::CloseFrame { + code: tokio_tungstenite::tungstenite::protocol::frame::coding::CloseCode::Normal, + reason: "finished".into(), + })), + ] { + let restored: WireMessage = from_wire(message.clone()).expect("conversion").into(); + assert_eq!(restored, message); + } + } + + #[tokio::test] + async fn secure_relay_requires_tls_before_any_http_upgrade() { + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("listener"); + let endpoint = format!("wss://{}", listener.local_addr().expect("address")); + let server = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.expect("accept"); + let mut prefix = [0; 5]; + socket + .read_exact(&mut prefix) + .await + .expect("TLS record prefix"); + socket + .write_all(b"HTTP/1.1 101 Switching Protocols\r\n\r\n") + .await + .expect("plaintext response"); + prefix + }); + let profile = crate::profile::test_profile( + crate::RelayProfileKind::Simulator, + RelayUrlPolicy::Local, + [endpoint.as_str()], + ) + .expect("local TLS profile"); + let transport = HardenedWebsocketTransport::new(profile.endpoints()); + let url = Url::parse(&endpoint).expect("URL"); + let result = transport + .connect(&url, &ConnectionMode::Direct, Duration::from_secs(5)) + .await; + assert!( + result.is_err(), + "plaintext peer must not complete a secure relay connection" + ); + let prefix = tokio::time::timeout(Duration::from_secs(5), server) + .await + .expect("server deadline") + .expect("server task"); + assert_eq!( + prefix[0], 0x16, + "client must send a TLS handshake, never an HTTP request" + ); + assert_eq!(prefix[1], 0x03, "TLS record version family"); + } + + #[test] fn policies_classify_literal_and_named_destinations() { assert!(RelayUrl::parse("wss://relay.example.com", RelayUrlPolicy::Public).is_ok()); assert!(RelayUrl::parse("wss://10.0.0.1", RelayUrlPolicy::Public).is_err()); diff --git a/crates/transport_nostr/src/source.rs b/crates/transport_nostr/src/source.rs @@ -19,7 +19,7 @@ use std::sync::Arc; use std::time::{SystemTime, UNIX_EPOCH}; #[path = "source_budget.rs"] -mod budget; +pub(crate) mod budget; use budget::FetchBudget; #[path = "source_window.rs"] @@ -64,24 +64,28 @@ pub(crate) trait RelaySourceClient: Send + Sync { #[derive(Clone, Debug)] pub(crate) struct LiveRelaySourceClient { client: nostr_sdk::Client, + ingress: crate::source_ingress::IngressRegistry, } impl LiveRelaySourceClient { - pub(crate) const fn new(client: nostr_sdk::Client) -> Self { - Self { client } + pub(crate) const fn new( + client: nostr_sdk::Client, + ingress: crate::source_ingress::IngressRegistry, + ) -> Self { + Self { client, ingress } } #[cfg(test)] pub(crate) fn isolated() -> Self { let client = nostr_sdk::Client::default(); client.automatic_authentication(false); - Self::new(client) + Self::new(client, crate::source_ingress::IngressRegistry::default()) } } impl RelaySourceClient for LiveRelaySourceClient { - // The live SDK loop requires external relays. Selection and normalized - // result handling are covered through the injected source boundary. + // The live SDK loop is exercised by bounded loopback regressions; injected + // source tests independently cover selection and normalized results. #[cfg_attr(coverage_nightly, coverage(off))] fn fetch<'a>(&'a self, query: SourceQuery) -> BoxFuture<'a, Vec<RelayFetchBatch>> { Box::pin(async move { @@ -94,12 +98,58 @@ impl RelaySourceClient for LiveRelaySourceClient { max_connections, } = query; let budget = Arc::new(FetchBudget::default()); - stream::iter(relays.into_iter().map(|relay| { + // An unencodable selector has no network work or ingress to admit. + if !selector.kinds().is_empty() + && selector + .kinds() + .iter() + .all(|kind| u16::try_from(*kind).is_err()) + { + return relays + .into_iter() + .map(|relay| RelayFetchBatch { + relay, + result: RelayFetchResult::Complete(Vec::new()), + }) + .collect(); + } + let connections = match admit_fetch_connections(&self.client, &relays, deadline).await { + Ok(connections) => connections, + Err(result) => { + return relays + .into_iter() + .map(|relay| RelayFetchBatch { + relay, + result: result.clone(), + }) + .collect(); + } + }; + let registration = self.ingress.register( + relays.iter().map(|relay| relay.as_str().to_owned()), + Arc::clone(&budget), + ); + let Some(registration) = registration else { + // No ingress or query I/O was admitted, so existing shared + // connections retain their previous operation ownership. + connections.finish(); + return relays + .into_iter() + .map(|relay| RelayFetchBatch { + relay, + result: RelayFetchResult::ResourceLimit(Vec::new()), + }) + .collect(); + }; + let mut batches: Vec<RelayFetchBatch> = stream::iter(relays.into_iter().map(|relay| { let selector = selector.clone(); let budget = Arc::clone(&budget); async move { let url = relay.as_str().to_owned(); let result = async { + if budget.exhausted() { + return Ok(RelayFetchResult::ResourceLimit(Vec::new())); + } if tokio::time::Instant::now() >= deadline { return Ok(RelayFetchResult::Timeout(Vec::new())); } @@ -109,9 +159,6 @@ impl RelaySourceClient for LiveRelaySourceClient { .filter_map(|kind| u16::try_from(*kind).ok()) .map(Kind::from) .collect::<Vec<_>>(); - if !selector.kinds().is_empty() && kinds.is_empty() { - return Ok(RelayFetchResult::Complete(Vec::new())); - } let authors = selector .authors() .iter() @@ -140,10 +187,6 @@ impl RelaySourceClient for LiveRelaySourceClient { } let connected = tokio::time::timeout_at(deadline, async { self.client - .add_relay(url.as_str()) - .await - .map_err(|error| error.to_string())?; - self.client .try_connect_relay( url.as_str(), connect_timeout.min( @@ -187,28 +230,138 @@ impl RelaySourceClient for LiveRelaySourceClient { } Err(_) => return Ok(RelayFetchResult::Timeout(Vec::new())), } - Ok(collect_until_eose( + let result = collect_until_eose( &mut notifications, &subscription_id, deadline, &budget, ) - .await) + .await; + Ok(result) } .await; RelayFetchBatch { relay, - result: result.unwrap_or_else(RelayFetchResult::Failed), + result: match result { + Err(_) if budget.exhausted() => { + RelayFetchResult::ResourceLimit(Vec::new()) + } + Ok(RelayFetchResult::Timeout(events)) if budget.exhausted() => { + RelayFetchResult::ResourceLimit(events) + } + result => result.unwrap_or_else(RelayFetchResult::Failed), + }, } } })) .buffered(max_connections) .collect() - .await + .await; + if finish_fetch(&mut batches, registration) { + connections.finish(); + } + batches }) } } +async fn admit_fetch_connections( + client: &nostr_sdk::Client, + relays: &[RelayUrl], + deadline: tokio::time::Instant, +) -> Result<FetchConnections, RelayFetchResult> { + let connections = FetchConnections::default(); + let admitted = tokio::time::timeout_at(deadline, async { + for relay in relays { + if tokio::time::Instant::now() >= deadline { + return Err(RelayFetchResult::Timeout(Vec::new())); + } + // Adding a relay is inert SDK metadata admission, not connection + // work. Own even preexisting queued sockets before ingress can be + // revoked by this fetch's registration. + client + .add_relay(relay.as_str()) + .await + .map_err(|error| RelayFetchResult::Failed(error.to_string()))?; + connections + .retain( + client + .relay(relay.as_str()) + .await + .map_err(|error| RelayFetchResult::Failed(error.to_string()))?, + ) + .map_err(RelayFetchResult::Failed)?; + } + if tokio::time::Instant::now() >= deadline { + Err(RelayFetchResult::Timeout(Vec::new())) + } else { + Ok(()) + } + }) + .await; + match admitted { + Ok(Ok(())) => Ok(connections), + Ok(Err(result)) => Err(result), + Err(_) => Err(RelayFetchResult::Timeout(Vec::new())), + } +} + +fn finish_fetch( + batches: &mut [RelayFetchBatch], + registration: crate::source_ingress::FetchRegistration, +) -> bool { + if !batches + .iter() + .all(|batch| matches!(batch.result, RelayFetchResult::Complete(_))) + { + return false; + } + if registration.finish() { + return true; + } + // Earlier EOSE evidence remains useful, but finalization cannot declare + // an entirely complete fetch after ingress exhausted its shared budget. + if let Some(last) = batches.last_mut() + && let RelayFetchResult::Complete(events) = &mut last.result + { + last.result = RelayFetchResult::ResourceLimit(std::mem::take(events)); + } + false +} + +// The SDK owns its socket tasks. Its synchronous disconnect signal disables +// reconnect and tears down those tasks, including when the fetch future drops. +#[derive(Default)] +struct FetchConnections(std::sync::Mutex<Vec<nostr_sdk::Relay>>); + +impl FetchConnections { + fn retain(&self, relay: nostr_sdk::Relay) -> Result<(), String> { + self.0 + .lock() + .map_err(|_| String::from("fetch connection state unavailable"))? + .push(relay); + Ok(()) + } + + fn finish(&self) { + if let Ok(mut relays) = self.0.lock() { + relays.clear(); + } + } +} + +impl Drop for FetchConnections { + fn drop(&mut self) { + let relays = self + .0 + .get_mut() + .unwrap_or_else(std::sync::PoisonError::into_inner); + for relay in relays { + relay.disconnect(); + } + } +} + async fn collect_until_eose( notifications: &mut tokio::sync::broadcast::Receiver<RelayNotification>, subscription_id: &SubscriptionId, @@ -217,10 +370,18 @@ async fn collect_until_eose( ) -> RelayFetchResult { let mut events = Vec::new(); loop { + if budget.exhausted() { + return RelayFetchResult::ResourceLimit(events); + } if tokio::time::Instant::now() >= deadline { return RelayFetchResult::Timeout(events); } - let notification = match tokio::time::timeout_at(deadline, notifications.recv()).await { + let received = tokio::select! { + biased; + _ = budget.wait_exhausted() => return RelayFetchResult::ResourceLimit(events), + received = tokio::time::timeout_at(deadline, notifications.recv()) => received, + }; + let notification = match received { Ok(Ok(notification)) => notification, Ok(Err(error)) => return RelayFetchResult::Failed(error.to_string()), Err(_) => return RelayFetchResult::Timeout(events), @@ -248,7 +409,9 @@ async fn collect_until_eose( RelayNotification::Message { message: RelayMessage::EndOfStoredEvents(observed_subscription), } if observed_subscription.as_ref() == subscription_id => { - return if tokio::time::Instant::now() < deadline { + return if budget.exhausted() { + RelayFetchResult::ResourceLimit(events) + } else if tokio::time::Instant::now() < deadline { RelayFetchResult::Complete(events) } else { RelayFetchResult::Timeout(events) @@ -978,6 +1141,400 @@ mod tests { ); } + #[tokio::test] + async fn exhausted_ingress_wins_over_queued_eose() { + let subscription_id = SubscriptionId::generate(); + let (sender, mut receiver) = tokio::sync::broadcast::channel(1); + sender + .send(RelayNotification::Message { + message: RelayMessage::EndOfStoredEvents(Cow::Owned(subscription_id.clone())), + }) + .unwrap(); + let budget = FetchBudget::default(); + assert!(!budget.wire(budget::MAX_FETCH_BYTES + 1, 0, 0)); + assert_eq!( + collect_until_eose( + &mut receiver, + &subscription_id, + tokio::time::Instant::now() + Duration::from_secs(1), + &budget + ) + .await, + RelayFetchResult::ResourceLimit(Vec::new()) + ); + } + + #[tokio::test] + async fn exhaustion_after_eose_before_finalization_retains_events_and_earlier_evidence() { + let relays = ["wss://one.example", "wss://two.example"]; + let registry = + crate::source_ingress::IngressRegistry::new(relays.into_iter().map(str::to_owned)); + let budget = Arc::new(FetchBudget::default()); + let registration = registry + .register(relays.into_iter().map(str::to_owned), Arc::clone(&budget)) + .unwrap(); + let connection = registry.connection(relays[1]).unwrap(); + let mut batches = Vec::new(); + for url in relays { + let id = SubscriptionId::generate(); + let (sender, mut receiver) = tokio::sync::broadcast::channel(2); + sender + .send(RelayNotification::Message { + message: RelayMessage::Event { + subscription_id: Cow::Owned(id.clone()), + event: Cow::Owned(nostr_sdk::prelude::Event::from_json(FIRST).unwrap()), + }, + }) + .unwrap(); + sender + .send(RelayNotification::Message { + message: RelayMessage::EndOfStoredEvents(Cow::Owned(id.clone())), + }) + .unwrap(); + batches.push(RelayFetchBatch { + relay: RelayUrl::parse(url, RelayUrlPolicy::Public).unwrap(), + result: collect_until_eose( + &mut receiver, + &id, + tokio::time::Instant::now() + Duration::from_secs(1), + &budget, + ) + .await, + }); + } + let before = batches.clone(); + assert!(before.iter().all( + |batch| matches!(&batch.result, RelayFetchResult::Complete(events) if events.len() == 1) + )); + assert!(!connection.charge(budget::MAX_FETCH_BYTES + 1, 0, 0)); + assert!(!finish_fetch(&mut batches, registration)); + assert_eq!(batches[0].result, before[0].result); + let RelayFetchResult::Complete(expected) = &before[1].result else { + panic!("completed collection"); + }; + assert_eq!( + batches[1].result, + RelayFetchResult::ResourceLimit(expected.clone()) + ); + } + + #[test] + fn finalization_preserves_success_partial_and_empty_results() { + for (result, exhausted, expected_complete) in [ + ( + Some(RelayFetchResult::Complete(vec![FIRST.to_owned()])), + false, + true, + ), + ( + Some(RelayFetchResult::Timeout(vec![FIRST.to_owned()])), + false, + false, + ), + (None, false, true), + (None, true, false), + ] { + let registry = crate::source_ingress::IngressRegistry::new( + ["wss://one.example".to_owned()].into_iter(), + ); + let budget = Arc::new(FetchBudget::default()); + let registration = registry + .register( + ["wss://one.example".to_owned()].into_iter(), + Arc::clone(&budget), + ) + .unwrap(); + let connection = registry.connection("wss://one.example").unwrap(); + let mut batches = result + .map(|result| RelayFetchBatch { + relay: RelayUrl::parse("wss://one.example", RelayUrlPolicy::Public).unwrap(), + result, + }) + .into_iter() + .collect::<Vec<_>>(); + let before = batches + .iter() + .map(|batch| batch.result.clone()) + .collect::<Vec<_>>(); + if exhausted { + assert!(!connection.charge(budget::MAX_FETCH_BYTES + 1, 0, 0)); + } + assert_eq!(finish_fetch(&mut batches, registration), expected_complete); + assert_eq!( + batches + .iter() + .map(|batch| batch.result.clone()) + .collect::<Vec<_>>(), + before + ); + assert_eq!( + connection.admitted(&std::task::Context::from_waker( + futures::task::noop_waker_ref() + )), + expected_complete + ); + } + } + + #[tokio::test] + async fn completed_connection_ownership_is_retained_until_the_whole_query_finishes() { + let client = nostr_sdk::Client::default(); + client.add_relay("wss://one.example").await.unwrap(); + let relay = client.relay("wss://one.example").await.unwrap(); + let connections = FetchConnections::default(); + connections.retain(relay.clone()).unwrap(); + assert_eq!(connections.0.lock().unwrap().len(), 1); + connections.finish(); + assert!(connections.0.lock().unwrap().is_empty()); + connections.retain(relay).unwrap(); + drop(connections); + } + + #[tokio::test] + async fn complete_connection_inventory_is_inert_and_preserves_shared_sockets_on_success() { + let client = nostr_sdk::Client::default(); + let relays = ["wss://one.example", "wss://two.example"] + .map(|url| RelayUrl::parse(url, RelayUrlPolicy::Public).unwrap()); + let connections = admit_fetch_connections( + &client, + &relays, + tokio::time::Instant::now() + Duration::from_secs(1), + ) + .await + .unwrap(); + assert_eq!(connections.0.lock().unwrap().len(), 2); + for relay in &relays { + assert_eq!( + client.relay(relay.as_str()).await.unwrap().status(), + nostr_sdk::prelude::RelayStatus::Initialized + ); + } + connections.finish(); + drop(connections); + for relay in &relays { + assert_eq!( + client.relay(relay.as_str()).await.unwrap().status(), + nostr_sdk::prelude::RelayStatus::Initialized + ); + } + } + + #[tokio::test] + async fn expired_inventory_performs_no_sdk_admission() { + let client = nostr_sdk::Client::default(); + let relay = RelayUrl::parse("wss://one.example", RelayUrlPolicy::Public).unwrap(); + for relays in [vec![relay], Vec::new()] { + assert!(matches!( + admit_fetch_connections(&client, &relays, tokio::time::Instant::now()).await, + Err(RelayFetchResult::Timeout(_)) + )); + assert!(client.relays().await.is_empty()); + } + } + + #[tokio::test] + async fn failed_inventory_releases_every_handle_already_admitted() { + let client = nostr_sdk::Client::builder() + .opts( + nostr_sdk::ClientOptions::default() + .pool(nostr_sdk::prelude::RelayPoolOptions::default().max_relays(Some(1))), + ) + .build(); + let relays = ["wss://one.example", "wss://two.example"] + .map(|url| RelayUrl::parse(url, RelayUrlPolicy::Public).unwrap()); + assert!(matches!( + admit_fetch_connections( + &client, + &relays, + tokio::time::Instant::now() + Duration::from_secs(1) + ) + .await, + Err(RelayFetchResult::Failed(_)) + )); + assert_eq!(client.relays().await.len(), 1); + assert_eq!( + client.relay(relays[0].as_str()).await.unwrap().status(), + nostr_sdk::prelude::RelayStatus::Terminated + ); + } + + #[derive(Clone, Copy)] + enum QueuedCleanup { + Cancel, + Deadline, + Resource, + } + + async fn queued_connection_cleanup(mode: QueuedCleanup) { + use futures::SinkExt; + use tokio_tungstenite::{accept_async, tungstenite::Message}; + let first = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let second = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let urls = [ + format!("ws://{}", first.local_addr().unwrap()), + format!("ws://{}", second.local_addr().unwrap()), + ]; + let endpoints = urls + .iter() + .map(|url| { + crate::RelayEndpoint::new(url, RelayUrlPolicy::Local, crate::RelayAccess::ReadOnly) + .unwrap() + }) + .collect::<Vec<_>>(); + let connector = crate::relay::HardenedWebsocketTransport::new(&endpoints); + let ingress = connector.ingress.clone(); + let client = nostr_sdk::Client::builder() + .websocket_transport(connector) + .build(); + client.automatic_authentication(false); + let (started, observed) = tokio::sync::oneshot::channel(); + let first_server = tokio::spawn(async move { + let (tcp, _) = first.accept().await.unwrap(); + let mut socket = accept_async(tcp).await.unwrap(); + let mut started = Some(started); + while let Some(message) = socket.next().await { + match message { + Ok(Message::Text(text)) => { + let value: serde_json::Value = serde_json::from_str(&text).unwrap(); + if value[0] == "REQ" { + if let Some(started) = started.take() { + let _ = started.send(()); + } + if matches!(mode, QueuedCleanup::Resource) { + let malformed = serde_json::to_string(&( + "EVENT", + &value[1], + serde_json::json!({"invalid": "x".repeat(384 * 1024)}), + )) + .unwrap(); + for _ in 0..24 { + if socket + .send(Message::Text(malformed.clone().into())) + .await + .is_err() + { + return; + } + } + } + } + } + Ok(Message::Close(_)) | Err(_) => return, + _ => { + if socket.flush().await.is_err() { + return; + } + } + } + } + }); + let queued_requests = Arc::new(AtomicUsize::new(0)); + let requests = Arc::clone(&queued_requests); + let (closed, closure) = tokio::sync::oneshot::channel(); + let second_server = tokio::spawn(async move { + let (tcp, _) = second.accept().await.unwrap(); + let mut socket = accept_async(tcp).await.unwrap(); + while let Some(message) = socket.next().await { + match message { + Ok(Message::Text(text)) => { + let value: serde_json::Value = serde_json::from_str(&text).unwrap(); + if value[0] == "REQ" { + requests.fetch_add(1, AtomicOrdering::SeqCst); + } + } + Ok(Message::Close(_)) | Err(_) => break, + _ => { + if socket.flush().await.is_err() { + break; + } + } + } + } + let _ = closed.send(()); + }); + client.add_relay(urls[1].as_str()).await.unwrap(); + client + .try_connect_relay(urls[1].as_str(), Duration::from_secs(1)) + .await + .unwrap(); + let queued = client.relay(urls[1].as_str()).await.unwrap(); + let source = LiveRelaySourceClient::new(client, ingress); + let query = SourceQuery { + relays: urls + .iter() + .map(|url| RelayUrl::parse(url, RelayUrlPolicy::Local).unwrap()) + .collect(), + selector: radroots_transport::source::FetchSelector::all(), + until_unix_seconds: None, + connect_timeout: Duration::from_secs(1), + deadline: tokio::time::Instant::now() + + if matches!(mode, QueuedCleanup::Deadline) { + Duration::from_millis(500) + } else { + Duration::from_secs(5) + }, + max_connections: 1, + }; + let mut fetch = source.fetch(query); + let started = tokio::time::timeout(Duration::from_secs(10), async { + tokio::select! { + biased; + result = observed => result.is_ok(), + _ = &mut fetch => false, + } + }) + .await; + let batches = if matches!(mode, QueuedCleanup::Cancel) { + None + } else { + Some(tokio::time::timeout(Duration::from_secs(10), &mut fetch).await) + }; + drop(fetch); + let status = queued.status(); + let closed = tokio::time::timeout(Duration::from_secs(2), closure).await; + for server in [first_server, second_server] { + server.abort(); + let result = server.await; + assert!(result.is_ok() || result.unwrap_err().is_cancelled()); + } + assert!( + started.unwrap(), + "first relay must own the only scheduled batch" + ); + assert_eq!(queued_requests.load(AtomicOrdering::SeqCst), 0); + assert_eq!( + status, + nostr_sdk::prelude::RelayStatus::Terminated, + "queued preexisting socket must lose automatic reconnect authority" + ); + closed.unwrap().unwrap(); + if let Some(batches) = batches { + let batches = batches.unwrap(); + assert_eq!(batches.len(), 2); + assert!(batches.iter().all(|batch| match mode { + QueuedCleanup::Deadline => matches!(batch.result, RelayFetchResult::Timeout(_)), + QueuedCleanup::Resource => + matches!(batch.result, RelayFetchResult::ResourceLimit(_)), + QueuedCleanup::Cancel => false, + })); + } + } + + #[tokio::test(flavor = "multi_thread")] + async fn cancellation_terminates_preexisting_queued_relay_reconnect_authority() { + queued_connection_cleanup(QueuedCleanup::Cancel).await; + } + + #[tokio::test(flavor = "multi_thread")] + async fn deadline_terminates_preexisting_queued_relay_reconnect_authority() { + queued_connection_cleanup(QueuedCleanup::Deadline).await; + } + + #[tokio::test(flavor = "multi_thread")] + async fn exhaustion_terminates_preexisting_queued_relay_reconnect_authority() { + queued_connection_cleanup(QueuedCleanup::Resource).await; + } + #[test] fn source_rejects_unconfigured_targets_and_expired_deadlines() { let unconfigured_targets = TargetSet::new(vec![ diff --git a/crates/transport_nostr/src/source_budget.rs b/crates/transport_nostr/src/source_budget.rs @@ -1,28 +1,69 @@ use std::sync::Mutex; -pub(super) const MAX_EVENT_BYTES: usize = radroots_event_codec::decode::MAX_EVENT_JSON_BYTES; -pub(super) const MAX_FETCH_BYTES: usize = 8 * 1024 * 1024; -pub(super) const MAX_FETCH_EVENTS: usize = 4096; -pub(super) const MAX_FETCH_NOTIFICATIONS: usize = 8192; +pub(crate) const MAX_EVENT_BYTES: usize = radroots_event_codec::decode::MAX_EVENT_JSON_BYTES; +pub(crate) const MAX_FETCH_BYTES: usize = 8 * 1024 * 1024; +pub(crate) const MAX_FETCH_EVENTS: usize = 4096; +pub(crate) const MAX_FETCH_NOTIFICATIONS: usize = 8192; #[derive(Debug, Default)] struct Usage { bytes: usize, events: usize, notifications: usize, + wire_bytes: usize, + wire_messages: usize, + wire_data: usize, + exhausted: bool, } /// One monotonic inventory shared by all relay batches. Reservations are not /// refunded when a duplicate, malformed event or completed batch is discarded. #[derive(Debug, Default)] -pub(super) struct FetchBudget(Mutex<Usage>); +pub(crate) struct FetchBudget(Mutex<Usage>, tokio::sync::Notify); impl FetchBudget { + pub(crate) fn wire(&self, bytes: usize, frames: usize, data: usize) -> bool { + let Ok(mut usage) = self.0.lock() else { + return false; + }; + if usage.exhausted + || bytes > MAX_FETCH_BYTES - usage.wire_bytes + || frames > MAX_FETCH_NOTIFICATIONS - usage.wire_messages + || data > MAX_FETCH_EVENTS - usage.wire_data + { + usage.exhausted = true; + self.1.notify_waiters(); + return false; + } + usage.wire_bytes += bytes; + usage.wire_messages += frames; + usage.wire_data += data; + true + } + + pub(crate) fn exhausted(&self) -> bool { + self.0.lock().map_or(true, |usage| usage.exhausted) + } + + pub(super) async fn wait_exhausted(&self) { + loop { + let notified = self.1.notified(); + tokio::pin!(notified); + notified.as_mut().enable(); + if self.exhausted() { + return; + } + notified.await; + } + } + pub(super) fn notification(&self) -> bool { let Ok(mut usage) = self.0.lock() else { return false; }; - if usage.notifications == MAX_FETCH_NOTIFICATIONS { + if usage.exhausted || usage.notifications == MAX_FETCH_NOTIFICATIONS { + usage.exhausted = true; + self.1.notify_waiters(); return false; } usage.notifications += 1; @@ -30,13 +71,16 @@ impl FetchBudget { } pub(super) fn event(&self, bytes: usize) -> bool { - if bytes > MAX_EVENT_BYTES { - return false; - } let Ok(mut usage) = self.0.lock() else { return false; }; - if usage.events == MAX_FETCH_EVENTS || bytes > MAX_FETCH_BYTES - usage.bytes { + if usage.exhausted + || bytes > MAX_EVENT_BYTES + || usage.events == MAX_FETCH_EVENTS + || bytes > MAX_FETCH_BYTES - usage.bytes + { + usage.exhausted = true; + self.1.notify_waiters(); return false; } usage.events += 1; @@ -52,7 +96,7 @@ mod tests { #[test] fn each_maximum_is_accepted_and_the_next_unit_is_rejected() { let bytes = FetchBudget::default(); - assert!(!bytes.event(MAX_EVENT_BYTES + 1)); + assert!(!FetchBudget::default().event(MAX_EVENT_BYTES + 1)); for _ in 0..MAX_FETCH_BYTES / MAX_EVENT_BYTES { assert!(bytes.event(MAX_EVENT_BYTES)); } @@ -84,4 +128,42 @@ mod tests { assert_eq!(accepted * MAX_EVENT_BYTES, MAX_FETCH_BYTES); assert!(!budget.event(1)); } + + #[tokio::test] + async fn raw_bytes_frames_and_data_attempts_have_sticky_independent_limits() { + for (bytes, frames, data) in [ + (MAX_FETCH_BYTES, 0, 0), + (0, MAX_FETCH_NOTIFICATIONS, 0), + (0, 0, MAX_FETCH_EVENTS), + ] { + let budget = FetchBudget::default(); + assert!(budget.wire(bytes, frames, data)); + assert!(!budget.exhausted()); + assert!(!budget.wire( + usize::from(bytes > 0), + usize::from(frames > 0), + usize::from(data > 0) + )); + assert!(budget.exhausted()); + assert!(!budget.wire(0, 0, 0)); + assert!(!budget.event(0)); + assert!(!budget.notification()); + tokio::time::timeout(std::time::Duration::from_secs(1), budget.wait_exhausted()) + .await + .unwrap(); + } + } + + #[tokio::test] + async fn exhaustion_wakes_every_waiter_without_lost_notifications() { + let budget = FetchBudget::default(); + let first = budget.wait_exhausted(); + let second = budget.wait_exhausted(); + tokio::pin!(first, second); + assert!(futures::poll!(&mut first).is_pending()); + assert!(futures::poll!(&mut second).is_pending()); + assert!(!budget.wire(MAX_FETCH_BYTES + 1, 0, 0)); + assert!(futures::poll!(&mut first).is_ready()); + assert!(futures::poll!(&mut second).is_ready()); + } } diff --git a/crates/transport_nostr/src/source_ingress.rs b/crates/transport_nostr/src/source_ingress.rs @@ -0,0 +1,439 @@ +//! Bounded fetch admission at the shared socket's pre-Nostr decoding boundary. + +use crate::source::budget::FetchBudget; +use futures::task::AtomicWaker; +use std::{ + collections::{BTreeMap, BTreeSet}, + sync::{Arc, Mutex}, + task::Context, +}; + +const MAX_ACTIVE_FETCHES: usize = 64; + +#[derive(Debug, Default)] +struct RelayState { + generation: u64, + invalidated: u64, + waker: AtomicWaker, + write_waker: AtomicWaker, +} + +#[derive(Debug)] +struct Registration { + targets: BTreeSet<String>, + budget: Arc<FetchBudget>, +} + +#[derive(Debug, Default)] +struct State { + sequence: u64, + relays: BTreeMap<String, RelayState>, + active: BTreeMap<u64, Registration>, +} + +#[derive(Clone, Debug, Default)] +pub(crate) struct IngressRegistry(Arc<Mutex<State>>); + +impl IngressRegistry { + pub(crate) fn new(targets: impl Iterator<Item = String>) -> Self { + Self(Arc::new(Mutex::new(State { + relays: targets + .map(|target| (target, RelayState::default())) + .collect(), + ..State::default() + }))) + } + + pub(crate) fn register( + &self, + targets: impl Iterator<Item = String>, + budget: Arc<FetchBudget>, + ) -> Option<FetchRegistration> { + let mut state = self.0.lock().ok()?; + if state.active.len() == MAX_ACTIVE_FETCHES { + return None; + } + let targets: BTreeSet<_> = targets.collect(); + if targets + .iter() + .any(|target| !state.relays.contains_key(target)) + { + return None; + } + let id = state.sequence.checked_add(1)?; + state.sequence = id; + state.active.insert(id, Registration { targets, budget }); + Some(FetchRegistration { + registry: self.clone(), + id, + finished: false, + }) + } + + pub(crate) fn connection(&self, relay: &str) -> Option<IngressConnection> { + let mut state = self.0.lock().ok()?; + let row = state.relays.get_mut(relay)?; + row.generation = row.generation.checked_add(1)?; + row.waker.wake(); + row.write_waker.wake(); + Some(IngressConnection { + registry: self.clone(), + relay: relay.to_owned(), + generation: row.generation, + }) + } + + fn admitted(&self, relay: &str, generation: u64, context: &Context<'_>, write: bool) -> bool { + let Ok(state) = self.0.lock() else { + return false; + }; + let Some(row) = state.relays.get(relay) else { + return false; + }; + if generation <= row.invalidated || generation != row.generation { + return false; + } + if write { + row.write_waker.register(context.waker()); + } else { + row.waker.register(context.waker()); + } + true + } + + fn charge( + &self, + relay: &str, + generation: u64, + bytes: usize, + frames: usize, + data: usize, + ) -> bool { + let Ok(mut state) = self.0.lock() else { + return false; + }; + let Some(row) = state.relays.get(relay) else { + return false; + }; + if generation <= row.invalidated || generation != row.generation { + return false; + } + let mut denied = BTreeSet::new(); + for entry in state.active.values() { + if entry.targets.contains(relay) && !entry.budget.wire(bytes, frames, data) { + denied.extend(entry.targets.iter().cloned()); + } + } + for target in &denied { + if let Some(row) = state.relays.get_mut(target) { + row.invalidated = row.generation; + row.waker.wake(); + row.write_waker.wake(); + } + } + denied.is_empty() + } + + fn remove(&self, id: u64, cancel: bool) -> bool { + let Ok(mut state) = self.0.lock() else { + return false; + }; + let Some(entry) = state.active.remove(&id) else { + return false; + }; + // Charge and completion use this same mutex. No ingress can exhaust + // this registration between the budget check and its removal. + let complete = !cancel && !entry.budget.exhausted(); + if !complete { + for target in entry.targets { + if let Some(row) = state.relays.get_mut(&target) { + row.invalidated = row.generation; + row.waker.wake(); + row.write_waker.wake(); + } + } + } + complete + } +} + +pub(crate) struct FetchRegistration { + registry: IngressRegistry, + id: u64, + finished: bool, +} + +impl FetchRegistration { + pub(crate) fn finish(mut self) -> bool { + let complete = self.registry.remove(self.id, false); + self.finished = true; + complete + } +} + +impl Drop for FetchRegistration { + fn drop(&mut self) { + if !self.finished { + self.registry.remove(self.id, true); + } + } +} + +#[derive(Clone, Debug)] +pub(crate) struct IngressConnection { + registry: IngressRegistry, + relay: String, + generation: u64, +} + +impl IngressConnection { + pub(crate) fn admitted(&self, context: &Context<'_>) -> bool { + self.registry + .admitted(&self.relay, self.generation, context, false) + } + + pub(crate) fn admitted_write(&self, context: &Context<'_>) -> bool { + self.registry + .admitted(&self.relay, self.generation, context, true) + } + + pub(crate) fn charge(&self, bytes: usize, frames: usize, data: usize) -> bool { + self.registry + .charge(&self.relay, self.generation, bytes, frames, data) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::source::budget::{MAX_FETCH_BYTES, MAX_FETCH_EVENTS, MAX_FETCH_NOTIFICATIONS}; + + fn registry() -> IngressRegistry { + IngressRegistry::new(["one".to_owned(), "two".to_owned()].into_iter()) + } + + fn register( + registry: &IngressRegistry, + targets: &[&str], + budget: &Arc<FetchBudget>, + ) -> FetchRegistration { + registry + .register( + targets.iter().map(|target| (*target).to_owned()), + Arc::clone(budget), + ) + .unwrap() + } + + fn admitted(connection: &IngressConnection) -> bool { + connection.admitted(&Context::from_waker(futures::task::noop_waker_ref())) + } + + #[test] + fn cancellation_closes_the_exact_generation_and_a_new_fetch_can_reconnect() { + let registry = registry(); + let budget = Arc::new(FetchBudget::default()); + let registration = register(&registry, &["one"], &budget); + let first = registry.connection("one").unwrap(); + assert!(admitted(&first)); + drop(registration); + assert!(!admitted(&first)); + assert!(!first.charge(1, 1, 1)); + let fresh = register(&registry, &["one"], &Arc::new(FetchBudget::default())); + let second = registry.connection("one").unwrap(); + assert!(admitted(&second)); + assert!(second.charge(1, 1, 1)); + assert!(!admitted(&first)); + fresh.finish(); + assert!(admitted(&second)); + assert!(second.charge(MAX_FETCH_BYTES + 1, 0, 0)); + } + + #[test] + fn all_batches_and_reconnections_share_one_monotonic_budget() { + let registry = registry(); + let budget = Arc::new(FetchBudget::default()); + let registration = register(&registry, &["one", "two"], &budget); + let first = registry.connection("one").unwrap(); + assert!(first.charge(MAX_FETCH_BYTES / 2, 1, 1)); + let retried = registry.connection("one").unwrap(); + assert!(!admitted(&first)); + assert!(!first.charge(1, 0, 0)); + assert!(retried.charge(MAX_FETCH_BYTES / 2, 1, 1)); + let second = registry.connection("two").unwrap(); + assert!(!second.charge(1, 0, 0)); + assert!(budget.exhausted()); + assert!(!admitted(&retried)); + assert!(!admitted(&second)); + drop(registration); + assert!(!admitted(&retried)); + } + + #[test] + fn overlapping_fetches_each_charge_shared_traffic_and_unselected_relays_do_not() { + let registry = registry(); + let first_budget = Arc::new(FetchBudget::default()); + let second_budget = Arc::new(FetchBudget::default()); + let first = register(&registry, &["one"], &first_budget); + let second = register(&registry, &["one"], &second_budget); + let other = registry.connection("two").unwrap(); + assert!(other.charge( + MAX_FETCH_BYTES + 1, + MAX_FETCH_NOTIFICATIONS + 1, + MAX_FETCH_EVENTS + 1 + )); + let connection = registry.connection("one").unwrap(); + assert!(connection.charge(MAX_FETCH_BYTES, 0, 0)); + first.finish(); + assert!(!connection.charge(1, 0, 0)); + assert!(!first_budget.exhausted()); + assert!(second_budget.exhausted()); + assert!(admitted(&other)); + drop(second); + } + + #[test] + fn registration_capacity_unknown_targets_and_counter_overflow_fail_closed() { + let registry = registry(); + let budget = Arc::new(FetchBudget::default()); + assert!( + registry + .register(["unknown".to_owned()].into_iter(), Arc::clone(&budget)) + .is_none() + ); + assert!(registry.connection("unknown").is_none()); + let mut registrations = (0..MAX_ACTIVE_FETCHES) + .map(|_| register(&registry, &["one"], &budget)) + .collect::<Vec<_>>(); + assert!( + registry + .register(["one".to_owned()].into_iter(), Arc::clone(&budget)) + .is_none() + ); + registrations.pop().unwrap().finish(); + register(&registry, &["one"], &budget).finish(); + for registration in registrations { + registration.finish(); + } + registry.0.lock().unwrap().sequence = u64::MAX; + assert!( + registry + .register(["one".to_owned()].into_iter(), budget) + .is_none() + ); + registry + .0 + .lock() + .unwrap() + .relays + .get_mut("one") + .unwrap() + .generation = u64::MAX; + assert!(registry.connection("one").is_none()); + } + + #[test] + fn cancel_notifies_the_pending_reader() { + struct Wake(std::sync::atomic::AtomicBool); + impl futures::task::ArcWake for Wake { + fn wake_by_ref(arc_self: &Arc<Self>) { + arc_self.0.store(true, std::sync::atomic::Ordering::SeqCst); + } + } + let registry = registry(); + let registration = register(&registry, &["one"], &Arc::new(FetchBudget::default())); + let connection = registry.connection("one").unwrap(); + let wake = Arc::new(Wake(std::sync::atomic::AtomicBool::new(false))); + let waker = futures::task::waker(Arc::clone(&wake)); + assert!(connection.admitted(&Context::from_waker(&waker))); + drop(registration); + assert!(wake.0.load(std::sync::atomic::Ordering::SeqCst)); + assert!(!admitted(&connection)); + } + + #[test] + fn a_stale_generation_cannot_steal_live_reader_or_writer_wakeups() { + struct Wake(std::sync::atomic::AtomicBool); + impl futures::task::ArcWake for Wake { + fn wake_by_ref(arc_self: &Arc<Self>) { + arc_self.0.store(true, std::sync::atomic::Ordering::SeqCst); + } + } + let registry = registry(); + let registration = register(&registry, &["one"], &Arc::new(FetchBudget::default())); + let stale = registry.connection("one").unwrap(); + let current = registry.connection("one").unwrap(); + let read = Arc::new(Wake(std::sync::atomic::AtomicBool::new(false))); + let write = Arc::new(Wake(std::sync::atomic::AtomicBool::new(false))); + let read_waker = futures::task::waker(Arc::clone(&read)); + let write_waker = futures::task::waker(Arc::clone(&write)); + assert!(current.admitted(&Context::from_waker(&read_waker))); + assert!(current.admitted_write(&Context::from_waker(&write_waker))); + assert!(!stale.admitted(&Context::from_waker(futures::task::noop_waker_ref()))); + assert!(!stale.admitted_write(&Context::from_waker(futures::task::noop_waker_ref()))); + drop(registration); + assert!(read.0.load(std::sync::atomic::Ordering::SeqCst)); + assert!(write.0.load(std::sync::atomic::Ordering::SeqCst)); + } + + #[test] + fn finalization_and_ingress_have_one_ordered_completion_boundary() { + for ingress_first in [false, true] { + let registry = registry(); + let budget = Arc::new(FetchBudget::default()); + let registration = register(&registry, &["one"], &budget); + let connection = registry.connection("one").unwrap(); + assert!(connection.charge(MAX_FETCH_BYTES, 0, 0)); + if ingress_first { + assert!(!connection.charge(1, 0, 0)); + assert!(!registration.finish()); + assert!(budget.exhausted()); + assert!(!admitted(&connection)); + } else { + assert!(registration.finish()); + assert!(connection.charge(1, 0, 0)); + assert!(!budget.exhausted()); + assert!(admitted(&connection)); + } + assert!(registry.0.lock().unwrap().active.is_empty()); + } + } + + #[test] + fn racing_completion_and_last_byte_cannot_both_claim_admission() { + for _ in 0..32 { + let registry = registry(); + let budget = Arc::new(FetchBudget::default()); + let registration = register(&registry, &["one"], &budget); + let connection = registry.connection("one").unwrap(); + assert!(connection.charge(MAX_FETCH_BYTES, 0, 0)); + let barrier = std::sync::Barrier::new(2); + let (complete, admitted_after_limit) = std::thread::scope(|scope| { + let finish = scope.spawn(|| { + barrier.wait(); + registration.finish() + }); + barrier.wait(); + let admitted_after_limit = connection.charge(1, 0, 0); + (finish.join().unwrap(), admitted_after_limit) + }); + assert_eq!(complete, admitted_after_limit); + assert_eq!(budget.exhausted(), !complete); + } + } + + #[test] + fn missing_or_poisoned_registration_cannot_finalize_successfully() { + let registry = registry(); + let budget = Arc::new(FetchBudget::default()); + let missing = register(&registry, &["one"], &budget); + registry.0.lock().unwrap().active.remove(&missing.id); + assert!(!missing.finish()); + let poisoned = register(&registry, &["one"], &budget); + let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| { + let _guard = registry.0.lock().unwrap(); + panic!("poison the registration lock"); + })); + assert!(!poisoned.finish()); + } +} diff --git a/crates/transport_nostr/src/source_wire.rs b/crates/transport_nostr/src/source_wire.rs @@ -0,0 +1,231 @@ +//! Fixed-memory admission of decrypted ingress before WebSocket decoding. + +use crate::source_ingress::IngressConnection; +use std::io; +use std::pin::Pin; +use std::task::{Context, Poll}; +use tokio::io::{AsyncRead, AsyncWrite, ReadBuf}; + +const MAX_UPGRADE_BYTES: usize = 16 * 1024; +const READ_CHUNK_BYTES: usize = 4096; + +/// Counts HTTP response and frame overhead conservatively in the same raw-byte +/// allowance as payloads. No message or payload is retained by this layer. +pub(crate) struct MeteredIo<S> { + inner: S, + ingress: IngressConnection, + parser: WireParser, + failed: bool, +} + +impl<S> MeteredIo<S> { + pub(crate) fn new(inner: S, ingress: IngressConnection) -> Self { + Self { + inner, + ingress, + parser: WireParser::default(), + failed: false, + } + } + + fn admitted(&self, context: &Context<'_>) -> io::Result<()> { + if self.failed || !self.ingress.admitted(context) { + Err(denied()) + } else { + Ok(()) + } + } +} + +impl<S: AsyncRead + Unpin> AsyncRead for MeteredIo<S> { + fn poll_read( + self: Pin<&mut Self>, + context: &mut Context<'_>, + output: &mut ReadBuf<'_>, + ) -> Poll<io::Result<()>> { + let this = self.get_mut(); + if let Err(error) = this.admitted(context) { + return Poll::Ready(Err(error)); + } + if output.remaining() == 0 { + return Poll::Ready(Ok(())); + } + let mut scratch = [0; READ_CHUNK_BYTES]; + let length = output.remaining().min(scratch.len()); + let mut input = ReadBuf::new(&mut scratch[..length]); + match Pin::new(&mut this.inner).poll_read(context, &mut input) { + Poll::Pending => Poll::Pending, + Poll::Ready(Err(error)) => { + this.failed = true; + Poll::Ready(Err(error)) + } + Poll::Ready(Ok(())) => { + let bytes = input.filled(); + let result = if !this.ingress.charge(bytes.len(), 0, 0) { + Err(denied()) + } else if bytes.is_empty() { + this.parser.eof() + } else { + this.parser + .consume(bytes, |data| this.ingress.charge(0, 1, usize::from(data))) + }; + if let Err(error) = result { + this.failed = true; + return Poll::Ready(Err(error)); + } + // Cancellation may race a ready read. Revalidate before any + // plaintext becomes visible to the WebSocket implementation. + if let Err(error) = this.admitted(context) { + this.failed = true; + return Poll::Ready(Err(error)); + } + output.put_slice(bytes); + Poll::Ready(Ok(())) + } + } + } +} + +impl<S: AsyncWrite + Unpin> AsyncWrite for MeteredIo<S> { + fn poll_write( + self: Pin<&mut Self>, + context: &mut Context<'_>, + bytes: &[u8], + ) -> Poll<io::Result<usize>> { + let this = self.get_mut(); + if this.failed || !this.ingress.admitted_write(context) { + return Poll::Ready(Err(denied())); + } + Pin::new(&mut this.inner).poll_write(context, bytes) + } + + fn poll_flush(self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<io::Result<()>> { + let this = self.get_mut(); + if this.failed || !this.ingress.admitted_write(context) { + return Poll::Ready(Err(denied())); + } + Pin::new(&mut this.inner).poll_flush(context) + } + + fn poll_shutdown(self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<io::Result<()>> { + // Closing the underlying transport remains possible after revocation. + Pin::new(&mut self.get_mut().inner).poll_shutdown(context) + } +} + +#[derive(Default)] +struct WireParser { + upgrade_bytes: usize, + delimiter: u32, + upgraded: bool, + header: [u8; 10], + header_len: usize, + header_needed: usize, + payload_remaining: usize, +} + +impl WireParser { + fn consume( + &mut self, + mut bytes: &[u8], + mut admit_frame: impl FnMut(bool) -> bool, + ) -> io::Result<()> { + while let Some((&byte, rest)) = bytes.split_first() { + if !self.upgraded { + self.upgrade_bytes += 1; + if self.upgrade_bytes > MAX_UPGRADE_BYTES { + return Err(invalid("WebSocket upgrade response exceeds its limit")); + } + self.delimiter = (self.delimiter << 8) | u32::from(byte); + self.upgraded = self.delimiter == u32::from_be_bytes(*b"\r\n\r\n"); + bytes = rest; + } else if self.payload_remaining != 0 { + let consumed = self.payload_remaining.min(bytes.len()); + self.payload_remaining -= consumed; + bytes = &bytes[consumed..]; + } else { + if self.header_len == 0 { + let opcode = byte & 0x0f; + if !admit_frame(matches!(opcode, 1 | 2)) { + return Err(denied()); + } + if byte & 0x70 != 0 || !matches!(opcode, 0 | 1 | 2 | 8 | 9 | 10) { + return Err(invalid("unsupported WebSocket frame flags or opcode")); + } + if opcode >= 8 && byte & 0x80 == 0 { + return Err(invalid("fragmented WebSocket control frame")); + } + self.header_needed = 2; + } + self.header[self.header_len] = byte; + self.header_len += 1; + bytes = rest; + if self.header_len == 2 { + if byte & 0x80 != 0 { + return Err(invalid("masked server WebSocket frame")); + } + self.header_needed = match byte { + 126 => 4, + 127 => 10, + _ => 2, + }; + } + if self.header_len == self.header_needed { + self.payload_remaining = self.payload_length()?; + self.header_len = 0; + } + } + } + Ok(()) + } + + fn payload_length(&self) -> io::Result<usize> { + let (length, minimum) = match self.header[1] { + 126 => ( + u64::from(u16::from_be_bytes([self.header[2], self.header[3]])), + 126, + ), + 127 => { + let mut extended = [0; 8]; + extended.copy_from_slice(&self.header[2..10]); + (u64::from_be_bytes(extended), 65536) + } + length => (u64::from(length), 0), + }; + if length < minimum || length > crate::relay::MAX_WIRE_MESSAGE_BYTES as u64 { + return Err(invalid( + "WebSocket frame length is noncanonical or exceeds its limit", + )); + } + if self.header[0] & 0x0f >= 8 && length > 125 { + return Err(invalid("WebSocket control frame exceeds its limit")); + } + usize::try_from(length).map_err(|_| invalid("WebSocket frame length cannot be represented")) + } + + fn eof(&self) -> io::Result<()> { + if !self.upgraded || self.header_len != 0 || self.payload_remaining != 0 { + Err(io::Error::new( + io::ErrorKind::UnexpectedEof, + "truncated WebSocket ingress", + )) + } else { + Ok(()) + } + } +} + +fn invalid(message: &'static str) -> io::Error { + io::Error::new(io::ErrorKind::InvalidData, message) +} + +fn denied() -> io::Error { + io::Error::new( + io::ErrorKind::ConnectionAborted, + "relay ingress budget or generation revoked", + ) +} + +#[cfg(test)] +#[path = "source_wire_tests.rs"] +mod tests; diff --git a/crates/transport_nostr/src/source_wire_tests.rs b/crates/transport_nostr/src/source_wire_tests.rs @@ -0,0 +1,334 @@ +use super::*; +use crate::source::budget::{FetchBudget, MAX_FETCH_BYTES}; +use crate::source_ingress::{FetchRegistration, IngressRegistry}; +use futures::task::{ArcWake, waker}; +use std::sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, +}; + +const UPGRADE: &[u8] = b"HTTP/1.1 101 Switching Protocols\r\n\r\n"; +const RELAY: &str = "ws://127.0.0.1:1234"; + +fn upgraded() -> WireParser { + let mut parser = WireParser::default(); + parser + .consume(UPGRADE, |_| panic!("HTTP is not a frame")) + .expect("upgrade"); + parser +} + +fn header(length: u64) -> Vec<u8> { + let mut bytes = vec![0x82, 127]; + bytes.extend(length.to_be_bytes()); + bytes +} + +#[test] +fn maximum_frame_length_is_admitted_and_next_byte_or_overflow_is_denied() { + let maximum = crate::relay::MAX_WIRE_MESSAGE_BYTES as u64; + let mut parser = upgraded(); + parser + .consume(&header(maximum), |_| true) + .expect("exact limit"); + assert_eq!(parser.payload_remaining, maximum as usize); + for length in [maximum + 1, 1 << 63, u64::MAX] { + assert!(upgraded().consume(&header(length), |_| true).is_err()); + } +} + +#[test] +fn headers_payloads_and_upgrade_can_arrive_one_byte_at_a_time() { + let mut parser = WireParser::default(); + let mut bytes = UPGRADE.to_vec(); + bytes.extend([0x01, 126, 0, 126]); + bytes.extend([b'x'; 126]); + bytes.extend([0x89, 0, 0x80, 2, b'y', b'z', 0x82, 0]); + let mut frames = 0; + let mut data = 0; + for byte in bytes { + parser + .consume(&[byte], |is_data| { + frames += 1; + data += usize::from(is_data); + true + }) + .expect("chopped ingress"); + } + assert_eq!((frames, data), (4, 2)); + parser.eof().expect("complete frames"); +} + +#[test] +fn zero_length_continuations_and_controls_each_require_admission() { + let mut parser = upgraded(); + let mut admissions = 0; + let error = parser + .consume(&[0x01, 0, 0x00, 0, 0x89, 0, 0x80, 0], |is_data| { + admissions += 1; + assert_eq!(is_data, admissions == 1); + admissions <= 3 + }) + .expect_err("fourth frame is denied"); + assert_eq!(error.kind(), io::ErrorKind::ConnectionAborted); + assert_eq!(admissions, 4); +} + +#[test] +fn invalid_flags_masks_lengths_and_controls_fail_closed() { + for bytes in [ + vec![0xc1, 0], + vec![0x83, 0], + vec![0x09, 0], + vec![0x81, 0x80], + vec![0x82, 126, 0, 125], + header(65535), + vec![0x89, 126, 0, 126], + ] { + assert!( + upgraded().consume(&bytes, |_| true).is_err(), + "accepted {bytes:?}" + ); + } +} + +#[test] +fn upgrade_prefix_has_an_exact_fixed_bound() { + let mut exact = vec![b'x'; MAX_UPGRADE_BYTES - 4]; + exact.extend(b"\r\n\r\n"); + let mut parser = WireParser::default(); + parser.consume(&exact, |_| true).expect("exact maximum"); + parser.eof().expect("upgrade boundary"); + let mut excessive = vec![b'x']; + excessive.extend(exact); + assert!(WireParser::default().consume(&excessive, |_| true).is_err()); +} + +#[test] +fn eof_rejects_partial_upgrade_header_and_payload() { + assert!(WireParser::default().eof().is_err()); + for bytes in [vec![0x81], vec![0x82, 126, 0], vec![0x81, 2, b'x']] { + let mut parser = upgraded(); + parser.consume(&bytes, |_| true).expect("partial input"); + assert_eq!( + parser.eof().expect_err("truncation").kind(), + io::ErrorKind::UnexpectedEof + ); + } + upgraded().eof().expect("frame boundary"); +} + +#[derive(Default)] +struct Calls { + reads: AtomicUsize, + writes: AtomicUsize, + wakes: AtomicUsize, +} + +impl ArcWake for Calls { + fn wake_by_ref(arc_self: &Arc<Self>) { + arc_self.wakes.fetch_add(1, Ordering::SeqCst); + } +} + +struct TestIo { + bytes: Vec<u8>, + offset: usize, + chunk: usize, + pending: bool, + calls: Arc<Calls>, +} + +impl AsyncRead for TestIo { + fn poll_read( + mut self: Pin<&mut Self>, + _: &mut Context<'_>, + output: &mut ReadBuf<'_>, + ) -> Poll<io::Result<()>> { + self.calls.reads.fetch_add(1, Ordering::SeqCst); + if self.pending { + return Poll::Pending; + } + let count = (self.bytes.len() - self.offset) + .min(output.remaining()) + .min(self.chunk); + output.put_slice(&self.bytes[self.offset..self.offset + count]); + self.offset += count; + Poll::Ready(Ok(())) + } +} + +impl AsyncWrite for TestIo { + fn poll_write( + self: Pin<&mut Self>, + _: &mut Context<'_>, + bytes: &[u8], + ) -> Poll<io::Result<usize>> { + self.calls.writes.fetch_add(1, Ordering::SeqCst); + Poll::Ready(Ok(bytes.len())) + } + fn poll_flush(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<io::Result<()>> { + self.calls.writes.fetch_add(1, Ordering::SeqCst); + Poll::Ready(Ok(())) + } + fn poll_shutdown(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<io::Result<()>> { + Poll::Ready(Ok(())) + } +} + +fn fixture( + bytes: Vec<u8>, +) -> ( + MeteredIo<TestIo>, + Arc<FetchBudget>, + FetchRegistration, + Arc<Calls>, +) { + let registry = IngressRegistry::new([RELAY.to_owned()].into_iter()); + let budget = Arc::new(FetchBudget::default()); + let registration = registry + .register([RELAY.to_owned()].into_iter(), budget.clone()) + .expect("registration"); + let connection = registry.connection(RELAY).expect("connection"); + let calls = Arc::new(Calls::default()); + let inner = TestIo { + bytes, + offset: 0, + chunk: READ_CHUNK_BYTES, + pending: false, + calls: calls.clone(), + }; + ( + MeteredIo::new(inner, connection), + budget, + registration, + calls, + ) +} + +fn read(meter: &mut MeteredIo<TestIo>, calls: &Arc<Calls>) -> (Poll<io::Result<()>>, usize) { + let waker = waker(calls.clone()); + let mut context = Context::from_waker(&waker); + let mut output = [0; READ_CHUNK_BYTES]; + let mut buffer = ReadBuf::new(&mut output); + let result = Pin::new(meter).poll_read(&mut context, &mut buffer); + (result, buffer.filled().len()) +} + +#[test] +fn raw_upgrade_and_frame_overhead_count_before_decode() { + let mut bytes = UPGRADE.to_vec(); + bytes.extend([0x81, 0]); + let length = bytes.len(); + let (mut meter, budget, _registration, calls) = fixture(bytes); + assert!(budget.wire(MAX_FETCH_BYTES - length, 0, 0)); + assert!(matches!(read(&mut meter, &calls), (Poll::Ready(Ok(())), count) if count == length)); + assert!(!budget.wire(1, 0, 0)); +} + +#[test] +fn excessive_read_is_not_exposed_and_subsequent_io_is_rejected() { + let (mut meter, budget, _registration, calls) = fixture(UPGRADE.to_vec()); + assert!(budget.wire(MAX_FETCH_BYTES, 0, 0)); + assert!(matches!(read(&mut meter, &calls), (Poll::Ready(Err(_)), 0))); + assert!(budget.exhausted()); + let reads = calls.reads.load(Ordering::SeqCst); + assert!(matches!(read(&mut meter, &calls), (Poll::Ready(Err(_)), 0))); + assert_eq!(calls.reads.load(Ordering::SeqCst), reads); + let waker = waker(calls.clone()); + let mut context = Context::from_waker(&waker); + assert!(matches!( + Pin::new(&mut meter).poll_write(&mut context, b"REQ"), + Poll::Ready(Err(_)) + )); + assert!(matches!( + Pin::new(&mut meter).poll_flush(&mut context), + Poll::Ready(Err(_)) + )); + assert_eq!(calls.writes.load(Ordering::SeqCst), 0); +} + +#[test] +fn cancellation_wakes_pending_read_and_revokes_both_io_directions() { + let (mut meter, _budget, registration, calls) = fixture(Vec::new()); + meter.inner.pending = true; + assert!(matches!(read(&mut meter, &calls), (Poll::Pending, 0))); + let writer_calls = Arc::new(Calls::default()); + let writer_waker = waker(writer_calls.clone()); + let mut writer_context = Context::from_waker(&writer_waker); + assert!(matches!( + Pin::new(&mut meter).poll_write(&mut writer_context, b"REQ"), + Poll::Ready(Ok(3)) + )); + assert!(matches!( + Pin::new(&mut meter).poll_flush(&mut writer_context), + Poll::Ready(Ok(())) + )); + drop(registration); + assert!(calls.wakes.load(Ordering::SeqCst) > 0); + assert!(writer_calls.wakes.load(Ordering::SeqCst) > 0); + assert!(matches!(read(&mut meter, &calls), (Poll::Ready(Err(_)), 0))); + let waker = waker(calls.clone()); + let mut context = Context::from_waker(&waker); + assert!(matches!( + Pin::new(&mut meter).poll_write(&mut context, b"REQ"), + Poll::Ready(Err(_)) + )); + assert!(matches!( + Pin::new(&mut meter).poll_shutdown(&mut context), + Poll::Ready(Ok(())) + )); + assert_eq!(calls.reads.load(Ordering::SeqCst), 1); + assert_eq!(calls.writes.load(Ordering::SeqCst), 2); +} + +#[test] +fn continuation_flood_exhausts_frame_budget_without_assembled_message() { + let mut bytes = UPGRADE.to_vec(); + bytes.extend([0x01, 0]); + for _ in 0..8192 { + bytes.extend([0x00, 0]); + } + let (mut meter, budget, _registration, calls) = fixture(bytes); + loop { + match read(&mut meter, &calls) { + (Poll::Ready(Ok(())), count) => assert!(count > 0, "unexpected EOF"), + (Poll::Ready(Err(_)), 0) => break, + other => panic!("unexpected result {other:?}"), + } + } + assert!(budget.exhausted()); +} + +#[test] +fn empty_data_frames_exhaust_attempt_budget() { + let mut bytes = UPGRADE.to_vec(); + for _ in 0..4097 { + bytes.extend([0x81, 0]); + } + let (mut meter, budget, _registration, calls) = fixture(bytes); + loop { + match read(&mut meter, &calls) { + (Poll::Ready(Ok(())), count) => assert!(count > 0, "unexpected EOF"), + (Poll::Ready(Err(_)), 0) => break, + other => panic!("unexpected result {other:?}"), + } + } + assert!(budget.exhausted()); +} + +#[test] +fn read_chunks_remain_bounded_and_truncation_is_terminal() { + let mut bytes = UPGRADE.to_vec(); + bytes.extend([0x81, 126, 0x20, 0]); + bytes.extend([b'x'; 5000]); + let (mut meter, _budget, _registration, calls) = fixture(bytes); + assert!(matches!( + read(&mut meter, &calls), + (Poll::Ready(Ok(())), READ_CHUNK_BYTES) + )); + assert!(matches!(read(&mut meter, &calls), (Poll::Ready(Ok(())), count) if count > 0)); + assert!( + matches!(read(&mut meter, &calls), (Poll::Ready(Err(error)), 0) if error.kind() == io::ErrorKind::UnexpectedEof) + ); +} diff --git a/crates/transport_nostr/tests/fetch_bounds.rs b/crates/transport_nostr/tests/fetch_bounds.rs @@ -218,7 +218,7 @@ async fn oversized_wire_messages_never_become_completed_fetch_evidence() { } #[tokio::test(flavor = "multi_thread")] -async fn dropping_a_polled_fetch_retains_the_original_remote_auto_close_bound() { +async fn dropping_a_polled_fetch_closes_the_subscription_or_connection() { let (listener, url) = listener().await; let (started, observed) = oneshot::channel(); let server = tokio::spawn(async move { @@ -226,8 +226,10 @@ async fn dropping_a_polled_fetch_retains_the_original_remote_auto_close_bound() let mut socket = accept_async(stream).await.unwrap(); let mut started = Some(started); while let Some(message) = socket.next().await { - let Ok(Message::Text(message)) = message else { - continue; + let message = match message { + Ok(Message::Text(message)) => message, + Ok(Message::Close(_)) | Err(_) => return, + _ => continue, }; let values: Value = serde_json::from_str(&message).unwrap(); if values[0] == "REQ" { @@ -236,7 +238,6 @@ async fn dropping_a_polled_fetch_retains_the_original_remote_auto_close_bound() return; } } - panic!("the published subscription must receive CLOSE"); }); let (transport, request) = transport(&[url], 1); let mut fetch = Box::pin(transport.fetch(request)); diff --git a/crates/transport_nostr/tests/fetch_raw_budget.rs b/crates/transport_nostr/tests/fetch_raw_budget.rs @@ -0,0 +1,290 @@ +use futures::{SinkExt, StreamExt}; +use radroots_transport::{ + EventSource, FetchRequest, TargetSet, outcome::FetchTargetState, source::FetchBounds, +}; +use radroots_transport_nostr::{ + Config, NostrTransport, RelayAccess, RelayEndpoint, RelayProfile, RelayProfileKind, + RelayUrlPolicy, +}; +use serde_json::Value; +use std::time::{Duration, SystemTime, UNIX_EPOCH}; +use tokio::net::{TcpListener, TcpStream}; +use tokio_tungstenite::{WebSocketStream, accept_async, tungstenite::Message}; + +const WATCHDOG: Duration = Duration::from_secs(10); + +fn transport(urls: &[String]) -> (NostrTransport, FetchRequest) { + let profile = RelayProfile::explicit( + RelayProfileKind::Simulator, + urls.iter().map(|url| { + RelayEndpoint::new(url, RelayUrlPolicy::Local, RelayAccess::ReadOnly).unwrap() + }), + ) + .unwrap(); + let config = Config::from_profile(profile) + .with_timeouts(1000, 5000, 500) + .unwrap() + .with_max_connections(1) + .unwrap(); + let targets = TargetSet::new( + config + .read_relays() + .map(|relay| relay.to_target().unwrap()) + .collect(), + ) + .unwrap(); + let now = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap() + .as_millis() as u64; + let request = FetchRequest::new( + "raw-budget-loopback", + targets, + FetchBounds::new(10, now + 5000).unwrap(), + ) + .unwrap(); + (NostrTransport::new(config), request) +} + +async fn listener() -> (TcpListener, String) { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let url = format!("ws://{}", listener.local_addr().unwrap()); + (listener, url) +} + +// A WebSocket pong is produced only after the client's receive stream has +// consumed all preceding messages. No scheduling sleep is an admission barrier. +async fn barrier(socket: &mut WebSocketStream<TcpStream>, sequence: usize) -> bool { + let marker = sequence.to_be_bytes().to_vec(); + if socket + .send(Message::Ping(marker.clone().into())) + .await + .is_err() + { + return false; + } + while let Some(message) = socket.next().await { + match message { + Ok(Message::Pong(value)) if value.as_ref() == marker.as_slice() => return true, + Ok(Message::Close(_)) | Err(_) => return false, + _ => { + if socket.flush().await.is_err() { + return false; + } + } + } + } + false +} + +async fn serve_malformed(listener: TcpListener, frames: usize) { + let (stream, _) = listener.accept().await.unwrap(); + let mut socket = accept_async(stream).await.unwrap(); + while let Some(message) = socket.next().await { + let text = match message { + Ok(Message::Text(text)) => text, + Ok(Message::Close(_)) | Err(_) => return, + _ => { + let _ = socket.flush().await; + continue; + } + }; + let values: Value = serde_json::from_str(&text).unwrap(); + if values[0] != "REQ" { + continue; + } + let malformed = serde_json::to_string(&( + "EVENT", + &values[1], + serde_json::json!({ "invalid_event": "x".repeat(384 * 1024) }), + )) + .unwrap(); + assert!(malformed.len() < 512 * 1024); + for sequence in 0..frames { + if socket + .send(Message::Text(malformed.clone().into())) + .await + .is_err() + || !barrier(&mut socket, sequence).await + { + return; + } + } + if socket + .send(Message::Text( + serde_json::to_string(&("EOSE", &values[1])).unwrap().into(), + )) + .await + .is_err() + { + return; + } + } +} + +async fn fetch_malformed(frames: usize) -> FetchTargetState { + let (listener, url) = listener().await; + let server = tokio::spawn(serve_malformed(listener, frames)); + let (transport, request) = transport(&[url]); + let result = tokio::time::timeout(WATCHDOG, transport.fetch(request)).await; + server.abort(); + let joined = server.await; + assert!(joined.is_ok() || joined.unwrap_err().is_cancelled()); + let page = result.unwrap().unwrap(); + assert!(page.events().is_empty()); + assert_eq!(page.target_outcomes().len(), 1); + page.target_outcomes()[0].state() +} + +#[tokio::test(flavor = "multi_thread")] +async fn empty_eose_remains_complete_with_raw_ingress_accounting() { + assert_eq!(fetch_malformed(0).await, FetchTargetState::Complete); +} + +#[tokio::test(flavor = "multi_thread")] +async fn malformed_event_traffic_exhausts_raw_budget_before_eose() { + assert_eq!(fetch_malformed(24).await, FetchTargetState::Partial); +} + +async fn finish_server(mut server: tokio::task::JoinHandle<()>) { + let result = tokio::time::timeout(WATCHDOG, &mut server).await; + if result.is_err() { + server.abort(); + let _ = server.await; + } + result + .expect("client must close the bounded relay connection") + .unwrap(); +} + +#[tokio::test(flavor = "multi_thread")] +async fn completed_batches_do_not_refund_raw_bytes_for_later_relays() { + let (one, one_url) = listener().await; + let (two, two_url) = listener().await; + let servers = [ + tokio::spawn(serve_malformed(one, 12)), + tokio::spawn(serve_malformed(two, 12)), + ]; + let (transport, request) = transport(&[one_url, two_url]); + let result = tokio::time::timeout(WATCHDOG, transport.fetch(request)).await; + for server in servers { + finish_server(server).await; + } + let page = result.unwrap().unwrap(); + assert!(page.events().is_empty()); + assert_eq!( + page.target_outcomes() + .iter() + .filter(|outcome| outcome.state() == FetchTargetState::Complete) + .count(), + 1 + ); + assert_eq!( + page.target_outcomes() + .iter() + .filter(|outcome| outcome.state() == FetchTargetState::Partial) + .count(), + 1 + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn unsolicited_malformed_traffic_is_bounded_before_the_request() { + let (listener, url) = listener().await; + let server = tokio::spawn(async move { + let (stream, _) = listener.accept().await.unwrap(); + let mut socket = accept_async(stream).await.unwrap(); + for sequence in 0..24 { + let malformed = format!( + "[\"EVENT\",\"unsolicited\",{{\"invalid\":\"{}\"}}]", + "x".repeat(384 * 1024) + ); + if socket.send(Message::Text(malformed.into())).await.is_err() + || !barrier(&mut socket, sequence).await + { + return; + } + } + panic!("oversized unsolicited inventory must close before admission"); + }); + let (transport, request) = transport(&[url]); + let result = tokio::time::timeout(WATCHDOG, transport.fetch(request)).await; + finish_server(server).await; + assert_eq!( + result.unwrap().unwrap().target_outcomes()[0].state(), + FetchTargetState::Partial + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn zero_payload_continuation_frames_consume_work_before_message_decoding() { + use tokio_tungstenite::tungstenite::protocol::frame::{ + Frame, + coding::{Data, OpCode}, + }; + let (listener, url) = listener().await; + let server = tokio::spawn(async move { + let (stream, _) = listener.accept().await.unwrap(); + let mut socket = accept_async(stream).await.unwrap(); + while let Some(message) = socket.next().await { + let text = match message { + Ok(Message::Text(text)) => text, + Ok(Message::Close(_)) | Err(_) => return, + _ => { + let _ = socket.flush().await; + continue; + } + }; + let values: Value = serde_json::from_str(&text).unwrap(); + if values[0] != "REQ" { + continue; + } + socket + .send(Message::Frame(Frame::message( + Vec::new(), + OpCode::Data(Data::Text), + false, + ))) + .await + .unwrap(); + for sequence in 0..8193 { + if socket + .send(Message::Frame(Frame::message( + Vec::new(), + OpCode::Data(Data::Continue), + false, + ))) + .await + .is_err() + { + return; + } + if sequence % 128 == 0 && !barrier(&mut socket, sequence).await { + return; + } + } + if socket + .send(Message::Frame(Frame::message( + b"[]".to_vec(), + OpCode::Data(Data::Continue), + true, + ))) + .await + .is_err() + { + return; + } + let _ = socket + .send(Message::Text( + serde_json::to_string(&("EOSE", &values[1])).unwrap().into(), + )) + .await; + } + }); + let (transport, request) = transport(&[url]); + let result = tokio::time::timeout(WATCHDOG, transport.fetch(request)).await; + finish_server(server).await; + let page = result.unwrap().unwrap(); + assert!(page.events().is_empty()); + assert_eq!(page.target_outcomes()[0].state(), FetchTargetState::Partial); +} diff --git a/crates/transport_nostr/tests/network_hardening.rs b/crates/transport_nostr/tests/network_hardening.rs @@ -11,7 +11,13 @@ const CLIENT_SOURCE: &str = include_str!("../src/client.rs"); fn tls_verification_and_pinned_dns_are_non_configurable_live_defaults() { assert!(WORKSPACE_MANIFEST.contains("rustls-tls-webpki-roots")); for required in [ - "client_async_tls_with_config(", + "client_async_with_config(", + "tokio_rustls::TlsConnector::from(Arc::new(tls_config))", + "roots.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned())", + ".with_no_client_auth()", + "rustls::pki_types::ServerName::try_from(", + "host.trim_matches(['[', ']']).to_owned()", + "crate::source_wire::MeteredIo::new(plaintext, ingress)", ".max_message_size(Some(MAX_WIRE_MESSAGE_BYTES))", ".max_frame_size(Some(MAX_WIRE_MESSAGE_BYTES))", "validate_resolved_addresses(", diff --git a/crates/transport_nostr/tests/package_boundary.rs b/crates/transport_nostr/tests/package_boundary.rs @@ -49,6 +49,8 @@ fn manifest_and_root_match_the_governed_transport_boundary() { "sink", "socket_write", "source", + "source_ingress", + "source_wire", "status", "subscription" ]) @@ -326,8 +328,11 @@ fn adapter_owns_no_storage_outbox_or_orchestration_surface() { "socket_write_tests.rs".to_owned(), "source.rs".to_owned(), "source_budget.rs".to_owned(), + "source_ingress.rs".to_owned(), "source_paging_tests.rs".to_owned(), "source_window.rs".to_owned(), + "source_wire.rs".to_owned(), + "source_wire_tests.rs".to_owned(), "status.rs".to_owned(), "subscription.rs".to_owned(), ])