From 50520265599c6d6bc3e30cbfa4e196fe1581ce07 Mon Sep 17 00:00:00 2001 From: Jeff Gardner <202880+erskingardner@users.noreply.github.com> Date: Wed, 30 Sep 2026 09:08:38 +0200 Subject: [PATCH] sdk: require correlated COUNT results and preserve waiter receive errors Pull-Request: https://github.com/nostrdevkit/nostr/pull/1478 Signed-off-by: Yuki Kishimoto --- nostr-sdk/CHANGELOG.md | 5 + nostr-sdk/src/relay/api/send_event.rs | 70 +++++++- nostr-sdk/src/relay/inner.rs | 92 +++++++++- nostr-sdk/src/relay/mod.rs | 232 +++++++++++++++++++++++--- 4 files changed, 364 insertions(+), 35 deletions(-) diff --git a/nostr-sdk/CHANGELOG.md b/nostr-sdk/CHANGELOG.md index 871813789..f502064c6 100644 --- a/nostr-sdk/CHANGELOG.md +++ b/nostr-sdk/CHANGELOG.md @@ -35,6 +35,11 @@ - Add `LocalRelayBuilderNip42::relay_url` to make it possible to configure a custom `relay_url` for the nip42 challenge validation (https://github.com/nostrdevkit/nostr/pull/1476) - Add `LocalRelayBuilder::new_event_channel_size` for customizing the size of the channel used to notify new received events +### Fixed + +- Require a correlated COUNT response before `Relay::count_events` returns a count. +- Preserve the broadcast lag or closure cause in event `OK` and authentication waiters instead of returning a generic premature-exit error. + ### Deprecated - Deprecate `LocalRelayBuilder::max_query_results` and `LocalRelayBuilder::default_filter_limit` (https://github.com/nostrdevkit/nostr/pull/1461) diff --git a/nostr-sdk/src/relay/api/send_event.rs b/nostr-sdk/src/relay/api/send_event.rs index 030da55f5..948620395 100644 --- a/nostr-sdk/src/relay/api/send_event.rs +++ b/nostr-sdk/src/relay/api/send_event.rs @@ -114,7 +114,18 @@ impl EventSendStatus { } } -/// Send event to relay +/// Send an event to one relay. +/// +/// By default, the operation waits for a matching `OK true` acknowledgement +/// from the relay. A matching `OK false` is an explicit rejection. +/// +/// Timeout, disconnection, notification loss, or notification closure while +/// waiting for the acknowledgement leaves publication unconfirmed: the relay +/// may have received the event. Callers should retain the event ID and resolve +/// ambiguity according to their own policy. +/// +/// Use [`SendEvent::wait_for_ok(false)`] to return after send submission without +/// waiting for relay acknowledgement. #[must_use = "Does nothing unless you await!"] pub struct SendEvent<'relay, 'event> { relay: &'relay Relay, @@ -135,7 +146,10 @@ impl<'relay, 'event> SendEvent<'relay, 'event> { } } - /// Wait for OK confirmation by the relay (default: true) + /// Wait for a matching `OK` confirmation by the relay (default: true). + /// + /// If disabled, [`EventSendStatus::Sent`] confirms only SDK send submission, + /// not relay acceptance. #[inline] pub fn wait_for_ok(mut self, enable: bool) -> Self { self.wait_for_ok = enable; @@ -184,8 +198,8 @@ async fn wait_for_authentication( timeout: Duration, ) -> Result<(), Error> { time::timeout(Some(timeout), async { - while let Ok(notification) = notifications.recv().await { - match notification { + loop { + match notifications.recv().await.map_err(Error::from)? { RelayNotification::Authenticated => { return Ok(()); } @@ -198,8 +212,6 @@ async fn wait_for_authentication( _ => (), } } - - Err(Error::state_msg("premature exit")) }) .await .ok_or_else(Error::timeout)? @@ -268,6 +280,7 @@ where #[cfg(test)] mod tests { + use std::error::Error as _; use std::time::Duration; use nostr::prelude::*; @@ -392,4 +405,49 @@ mod tests { // Send as authenticated assert!(relay.send_event(&event).await.is_ok()); } + + #[tokio::test] + async fn authentication_waiter_preserves_receive_failure() { + let (tx, mut rx) = broadcast::channel(2); + + for _ in 0..8 { + tx.send(RelayNotification::Authenticated).unwrap(); + } + + let error = wait_for_authentication(&mut rx, Duration::from_secs(1)) + .await + .unwrap_err(); + + assert!( + error + .source() + .and_then(|source| source.downcast_ref::()) + .is_some_and(|source| matches!(source, broadcast::error::RecvError::Lagged(_))) + ); + + let (tx, mut rx) = broadcast::channel(2); + + drop(tx); + + let error = wait_for_authentication(&mut rx, Duration::from_secs(1)) + .await + .unwrap_err(); + + assert!( + error + .source() + .and_then(|source| source.downcast_ref::()) + .is_some_and(|source| matches!(source, broadcast::error::RecvError::Closed)) + ); + + let (tx, mut rx) = broadcast::channel(2); + + tx.send(RelayNotification::AuthenticationFailed).unwrap(); + + let error = wait_for_authentication(&mut rx, Duration::from_secs(1)) + .await + .unwrap_err(); + + assert_eq!(error.kind(), crate::error::ErrorKind::Rejected); + } } diff --git a/nostr-sdk/src/relay/inner.rs b/nostr-sdk/src/relay/inner.rs index 7a716b0bc..254c823ac 100644 --- a/nostr-sdk/src/relay/inner.rs +++ b/nostr-sdk/src/relay/inner.rs @@ -1470,8 +1470,8 @@ impl InnerRelay { timeout: Duration, ) -> Result<(bool, String), Error> { time::timeout(Some(timeout), async { - while let Ok(notification) = notifications.recv().await { - match notification { + loop { + match notifications.recv().await.map_err(Error::from)? { RelayNotification::Message { message } => { if let RelayMessage::Ok { event_id, @@ -1491,8 +1491,6 @@ impl InnerRelay { _ => (), } } - - Err(Error::state_msg("premature exit")) }) .await .ok_or_else(Error::timeout)? @@ -1875,6 +1873,8 @@ async fn close_ws(tx: &mut WebSocketSink) -> Result<(), Error> { #[cfg(test)] mod tests { + use std::borrow::Cow; + use std::error::Error as _; use std::future::Future; use std::pin::Pin; use std::sync::Arc; @@ -2274,6 +2274,90 @@ mod tests { assert!(!relay.inner.should_resubscribe(&subscription_id).await); } + + #[tokio::test] + async fn ok_waiter_requires_matching_reply_and_preserves_receive_failure() { + let relay = Relay::new(RelayUrl::parse("wss://relay.example.com").unwrap()); + + let wanted = EventId::from_byte_array([0; 32]); + let other = EventId::from_byte_array([1; 32]); + + let (tx, mut rx) = broadcast::channel(2); + + tx.send(RelayNotification::Message { + message: Box::new(RelayMessage::Ok { + event_id: other, + status: true, + message: Cow::Borrowed("unrelated"), + }), + }) + .unwrap(); + + tx.send(RelayNotification::Message { + message: Box::new(RelayMessage::Ok { + event_id: wanted, + status: false, + message: Cow::Borrowed("rejected"), + }), + }) + .unwrap(); + + let (accepted, message) = relay + .inner + .wait_for_ok(&mut rx, &wanted, Duration::from_secs(1)) + .await + .unwrap(); + + assert!(!accepted); + assert_eq!(message, "rejected"); + + let (_silent_tx, mut silent_rx) = broadcast::channel(2); + + let error = relay + .inner + .wait_for_ok(&mut silent_rx, &wanted, Duration::from_millis(20)) + .await + .unwrap_err(); + + assert_eq!(error.kind(), ErrorKind::Timeout); + + let (lag_tx, _) = broadcast::channel(2); + let mut lagged = lag_tx.subscribe(); + + for _ in 0..8 { + lag_tx.send(RelayNotification::Authenticated).unwrap(); + } + + let error = relay + .inner + .wait_for_ok(&mut lagged, &wanted, Duration::from_secs(1)) + .await + .unwrap_err(); + + assert!( + error + .source() + .and_then(|source| source.downcast_ref::()) + .is_some_and(|source| matches!(source, broadcast::error::RecvError::Lagged(_))) + ); + + let (closed_tx, mut closed_rx) = broadcast::channel(2); + + drop(closed_tx); + + let error = relay + .inner + .wait_for_ok(&mut closed_rx, &wanted, Duration::from_secs(1)) + .await + .unwrap_err(); + + assert!( + error + .source() + .and_then(|source| source.downcast_ref::()) + .is_some_and(|source| matches!(source, broadcast::error::RecvError::Closed)) + ); + } } #[cfg(bench)] diff --git a/nostr-sdk/src/relay/mod.rs b/nostr-sdk/src/relay/mod.rs index 920901b74..8ed6f3854 100644 --- a/nostr-sdk/src/relay/mod.rs +++ b/nostr-sdk/src/relay/mod.rs @@ -411,41 +411,34 @@ impl Relay { FetchEvents::new(self, filters.into()) } - /// Count events + /// Count events. + /// + /// A successful zero is returned only for a matching COUNT response. + /// Timeout, notification loss, closure, or relay rejection return an error. pub async fn count_events(&self, filter: Filter, timeout: Duration) -> Result { - let id = SubscriptionId::generate(); + let id: SubscriptionId = SubscriptionId::generate(); + + let mut notifications = self.inner.internal_notification_sender.subscribe(); + let msg = ClientMessage::Count { subscription_id: Cow::Borrowed(&id), filter: Cow::Owned(filter), }; self.send_msg(msg).await?; - let mut count = 0; - - let mut notifications = self.inner.internal_notification_sender.subscribe(); - time::timeout(Some(timeout), async { - while let Ok(notification) = notifications.recv().await { - if let RelayNotification::Message { message } = notification { - if let RelayMessage::Count { - subscription_id, - count: c, - } = *message - { - if subscription_id.as_ref() == &id { - count = c; - break; - } - } - } - } - }) - .await - .ok_or_else(Error::timeout)?; + let fut = time::timeout(Some(timeout), receive_count_reply(&mut notifications, &id)); + let result: Result = fut.await.unwrap_or_else(|| Err(Error::timeout())); // Unsubscribe - self.send_msg(ClientMessage::close(id)).await?; + let close_result: Result<(), Error> = self.send_msg(ClientMessage::close(id)).await; - Ok(count) + match result { + Ok(count) => { + close_result?; + Ok(count) + } + Err(error) => Err(error), + } } /// Sync events with relays (negentropy reconciliation) @@ -455,14 +448,44 @@ impl Relay { } } +async fn receive_count_reply( + notifications: &mut broadcast::Receiver, + id: &SubscriptionId, +) -> Result { + loop { + match notifications.recv().await.map_err(Error::from)? { + RelayNotification::Message { message } => match *message { + RelayMessage::Count { + subscription_id, + count, + } if subscription_id.as_ref() == id => return Ok(count), + RelayMessage::Closed { + subscription_id, + message, + } if subscription_id.as_ref() == id => { + return Err(Error::relay_msg(message.into_owned())); + } + _ => {} + }, + RelayNotification::RelayStatus { status } if status.is_disconnected() => { + return Err(Error::not_connected()); + } + _ => {} + } + } +} + #[cfg(test)] mod tests { + use std::borrow::Cow; use std::collections::HashSet; use std::future::Future; + use std::net::SocketAddr; use std::pin::Pin; use std::sync::Arc; use async_utility::time; + use nostr::message::MachineReadablePrefix; use super::*; use crate::error::{Error, ErrorKind}; @@ -489,6 +512,34 @@ mod tests { } } + #[derive(Debug)] + struct RejectCount; + + impl QueryPolicy for RejectCount { + fn admit_query<'a>( + &'a self, + _query: &'a mut Filter, + _addr: &'a SocketAddr, + ) -> Pin + Send + 'a>> { + Box::pin(async { + QueryPolicyResult::reject(MachineReadablePrefix::Blocked, "count rejected") + }) + } + } + + #[derive(Debug)] + struct SilentCount; + + impl QueryPolicy for SilentCount { + fn admit_query<'a>( + &'a self, + _query: &'a mut Filter, + _addr: &'a SocketAddr, + ) -> Pin + Send + 'a>> { + Box::pin(std::future::pending()) + } + } + fn new_relay(url: RelayUrl, opts: RelayOptions) -> Relay { Relay::builder(url).opts(opts).build() } @@ -1098,4 +1149,135 @@ mod tests { // Must return None, as it's empty assert!(res.is_none()); } + + #[tokio::test] + async fn count_requires_a_matching_response() { + let (tx, mut notifications) = broadcast::channel(4); + let id = SubscriptionId::new("expected-count"); + let other_id = SubscriptionId::new("other-count"); + + tx.send(RelayNotification::Authenticated).unwrap(); + + tx.send(RelayNotification::Message { + message: Box::new(RelayMessage::Count { + subscription_id: Cow::Owned(other_id), + count: 99, + }), + }) + .unwrap(); + tx.send(RelayNotification::Message { + message: Box::new(RelayMessage::Count { + subscription_id: Cow::Owned(id.clone()), + count: 0, + }), + }) + .unwrap(); + + assert_eq!( + receive_count_reply(&mut notifications, &id).await.unwrap(), + 0 + ); + } + + #[tokio::test] + async fn count_receive_loss_and_closure_are_errors() { + let id = SubscriptionId::new("missing-count"); + let (tx, mut notifications) = broadcast::channel(2); + for _ in 0..8 { + tx.send(RelayNotification::Authenticated).unwrap(); + } + let error = receive_count_reply(&mut notifications, &id) + .await + .unwrap_err(); + assert_eq!(error.kind(), ErrorKind::Other); + assert!(error.to_string().contains("lagged")); + + let (tx, mut notifications) = broadcast::channel(2); + drop(tx); + let error = receive_count_reply(&mut notifications, &id) + .await + .unwrap_err(); + assert_eq!(error.kind(), ErrorKind::Other); + assert!(error.to_string().contains("closed")); + } + + #[tokio::test] + async fn count_disconnect_is_reported_before_timeout() { + let id = SubscriptionId::new("interrupted-count"); + let (tx, mut notifications) = broadcast::channel(2); + tx.send(RelayNotification::RelayStatus { + status: RelayStatus::Disconnected, + }) + .unwrap(); + + let error = tokio::time::timeout( + Duration::from_secs(1), + receive_count_reply(&mut notifications, &id), + ) + .await + .unwrap() + .unwrap_err(); + assert_eq!(error.kind(), ErrorKind::State); + assert!(error.to_string().contains("not connected")); + } + + #[tokio::test] + async fn count_with_no_matches_returns_zero() { + let local = LocalRelay::builder().build(); + local.run().await.unwrap(); + let relay = new_relay(local.url().await, RelayOptions::default()); + relay + .try_connect() + .timeout(Duration::from_secs(2)) + .await + .unwrap(); + + assert_eq!( + relay + .count_events(Filter::new(), Duration::from_secs(2)) + .await + .unwrap(), + 0 + ); + } + + #[tokio::test] + async fn count_without_count_reply_reports_rejection() { + let local = LocalRelay::builder().query_policy(RejectCount).build(); + local.run().await.unwrap(); + let relay = new_relay(local.url().await, RelayOptions::default()); + relay + .try_connect() + .timeout(Duration::from_secs(2)) + .await + .unwrap(); + + let error = relay + .count_events(Filter::new(), Duration::from_secs(2)) + .await + .unwrap_err(); + assert_eq!(error.kind(), ErrorKind::Rejected); + assert!(error.to_string().contains("count rejected")); + } + + #[tokio::test] + async fn count_without_a_reply_times_out() { + let local = LocalRelay::builder().query_policy(SilentCount).build(); + local.run().await.unwrap(); + let relay = new_relay(local.url().await, RelayOptions::default()); + relay + .try_connect() + .timeout(Duration::from_secs(2)) + .await + .unwrap(); + + let error = relay + .count_events(Filter::new(), Duration::from_millis(20)) + .await + .unwrap_err(); + assert_eq!(error.kind(), ErrorKind::Timeout); + assert_eq!(relay.status(), RelayStatus::Connected); + relay.shutdown(); + local.shutdown(); + } }