use std::borrow::Cow;
use std::cmp;
use std::collections::HashMap;
#[cfg(not(target_arch = "wasm32"))]
use std::net::SocketAddr;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
use async_utility::time;
use futures::{Stream, StreamExt};
use nostr_database::prelude::*;
use tokio::sync::{broadcast, oneshot};
mod api;
mod builder;
mod capabilities;
mod constants;
mod inner;
mod limits;
mod notification;
mod options;
mod ping;
mod stats;
mod status;
pub use self::api::*;
pub use self::builder::*;
pub use self::capabilities::*;
pub(crate) use self::constants::DEFAULT_FETCH_EVENTS_LIMIT;
use self::inner::InnerRelay;
pub use self::limits::*;
pub use self::notification::*;
pub use self::options::*;
pub use self::stats::*;
pub use self::status::*;
use crate::client::ClientNotification;
use crate::error::Error;
use crate::shared::SharedState;
use crate::stream::NotificationStream;
#[derive(Debug, Clone, PartialEq, Eq)]
enum SubscriptionAutoClosedReason {
AuthenticationFailed,
Closed(String),
Completed,
}
#[derive(Debug)]
enum SubscriptionActivity {
ReceivedEvent(Event),
Closed(SubscriptionAutoClosedReason),
}
#[derive(Debug)]
pub struct Relay {
pub(crate) inner: InnerRelay,
atomic_counter: Arc<AtomicUsize>,
}
impl Clone for Relay {
fn clone(&self) -> Self {
self.atomic_counter.fetch_add(1, Ordering::SeqCst);
Self {
inner: self.inner.clone(),
atomic_counter: self.atomic_counter.clone(),
}
}
}
impl Drop for Relay {
fn drop(&mut self) {
if self.atomic_counter.fetch_sub(1, Ordering::SeqCst) == 1 {
self.shutdown();
}
}
}
impl PartialEq for Relay {
fn eq(&self, other: &Self) -> bool {
self.inner.url == other.inner.url
}
}
impl Eq for Relay {}
impl PartialOrd for Relay {
fn partial_cmp(&self, other: &Self) -> Option<cmp::Ordering> {
Some(self.cmp(other))
}
}
impl Ord for Relay {
fn cmp(&self, other: &Self) -> cmp::Ordering {
self.inner.url.cmp(&other.inner.url)
}
}
impl Relay {
#[inline]
pub(crate) fn new_shared(
url: RelayUrl,
state: SharedState,
capabilities: RelayCapabilities,
opts: RelayOptions,
) -> Self {
Self {
inner: InnerRelay::new(url, state, capabilities, opts),
atomic_counter: Arc::new(AtomicUsize::new(1)),
}
}
#[inline]
pub fn new(url: RelayUrl) -> Self {
Self::builder(url).build()
}
#[inline]
pub fn builder(url: RelayUrl) -> RelayBuilder {
RelayBuilder::new(url)
}
fn from_builder(builder: RelayBuilder) -> Self {
let state: SharedState = SharedState::new(
builder.database,
builder.websocket_transport,
None,
builder.admit_policy,
builder.authenticator,
None,
);
Self::new_shared(builder.url, state, builder.capabilities, builder.opts)
}
#[inline]
pub fn url(&self) -> &RelayUrl {
&self.inner.url
}
#[inline]
#[cfg(not(target_arch = "wasm32"))]
pub fn proxy(&self) -> Option<SocketAddr> {
self.inner.proxy()
}
#[inline]
pub fn status(&self) -> RelayStatus {
self.inner.status()
}
#[inline]
pub fn capabilities(&self) -> &Arc<AtomicRelayCapabilities> {
&self.inner.capabilities
}
#[inline]
pub async fn subscriptions(&self) -> HashMap<SubscriptionId, Vec<Filter>> {
self.inner.subscriptions().await
}
#[inline]
pub async fn subscription(&self, id: &SubscriptionId) -> Option<Vec<Filter>> {
self.inner.subscription(id).await
}
#[inline]
pub fn opts(&self) -> &RelayOptions {
&self.inner.opts
}
#[inline]
pub fn stats(&self) -> &RelayConnectionStats {
&self.inner.stats
}
#[inline]
pub(super) fn set_notification_sender(
&mut self,
notification_sender: broadcast::Sender<ClientNotification>,
) {
self.inner.set_notification_sender(notification_sender);
}
#[inline]
pub fn notifications(&self) -> Pin<Box<dyn Stream<Item = RelayNotification> + Send>> {
let status: RelayStatus = self.status();
if status.is_banned() || status.is_shutdown() {
return Box::pin(futures::stream::empty());
}
let rx = self.inner.internal_notification_sender.subscribe();
let (tx, rx_done) = oneshot::channel();
let mut tx: Option<oneshot::Sender<()>> = Some(tx);
Box::pin(
NotificationStream::new(rx)
.inspect(move |notification| {
if let RelayNotification::RelayStatus { status } = ¬ification {
if status.is_banned() || status.is_shutdown() {
if let Some(tx) = tx.take() {
let _ = tx.send(());
}
}
}
})
.take_until(rx_done),
)
}
pub fn connect(&self) {
if !self.status().can_connect() {
return;
}
self.inner.set_status(RelayStatus::Pending, false);
self.inner.spawn_connection_task(None);
}
pub async fn wait_for_connection(&self, timeout: Duration) {
let status: RelayStatus = self.status();
if status.is_connected()
|| status.is_terminated()
|| status.is_banned()
|| status.is_shutdown()
{
return;
}
let mut notifications = self.inner.internal_notification_sender.subscribe();
time::timeout(Some(timeout), async {
while let Ok(notification) = notifications.recv().await {
if let RelayNotification::RelayStatus { status } = notification {
match status {
RelayStatus::Initialized
| RelayStatus::Pending
| RelayStatus::Connecting
| RelayStatus::Disconnected => {}
RelayStatus::Connected
| RelayStatus::Terminated
| RelayStatus::Banned
| RelayStatus::Sleeping
| RelayStatus::Shutdown => break,
}
}
}
})
.await;
}
pub fn try_connect(&self) -> TryConnect<'_> {
TryConnect::new(self)
}
#[inline]
pub fn disconnect(&self) {
self.inner.disconnect()
}
#[inline]
pub fn ban(&self) {
self.inner.ban()
}
#[inline]
pub fn shutdown(&self) {
self.inner.shutdown()
}
#[inline]
pub fn send_msg<'msg>(&self, msg: ClientMessage<'msg>) -> SendMessage<'_, 'msg> {
SendMessage::new(self, msg)
}
#[inline]
pub fn send_event<'event>(&self, event: &'event Event) -> SendEvent<'_, 'event> {
SendEvent::new(self, event)
}
#[inline]
pub fn subscribe<F>(&self, filters: F) -> Subscribe<'_>
where
F: Into<Vec<Filter>>,
{
Subscribe::new(self, filters.into())
}
#[inline]
pub fn unsubscribe<'id>(&self, id: &'id SubscriptionId) -> Unsubscribe<'_, 'id> {
Unsubscribe::new(self, id)
}
#[inline]
pub fn unsubscribe_all(&self) -> UnsubscribeAll<'_> {
UnsubscribeAll::new(self)
}
#[inline]
pub fn stream_events<F>(&self, filters: F) -> StreamEvents<'_>
where
F: Into<Vec<Filter>>,
{
StreamEvents::new(self, filters.into())
}
#[inline]
pub fn fetch_events<F>(&self, filters: F) -> FetchEvents<'_>
where
F: Into<Vec<Filter>>,
{
FetchEvents::new(self, filters.into())
}
pub async fn count_events(&self, filter: Filter, timeout: Duration) -> Result<usize, Error> {
let id = SubscriptionId::generate();
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)?;
self.send_msg(ClientMessage::close(id)).await?;
Ok(count)
}
#[inline]
pub fn sync(&self, filter: Filter) -> SyncEvents<'_> {
SyncEvents::new(self, filter)
}
}
#[cfg(test)]
mod tests {
use std::collections::HashSet;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use async_utility::time;
use super::*;
use crate::error::{Error, ErrorKind};
use crate::local_relay::*;
use crate::policy::{AdmitPolicy, AdmitStatus};
#[derive(Debug)]
struct CustomTestPolicy {
banned_relays: HashSet<RelayUrl>,
}
impl AdmitPolicy for CustomTestPolicy {
fn admit_connection<'a>(
&'a self,
relay_url: &'a RelayUrl,
) -> Pin<Box<dyn Future<Output = Result<AdmitStatus, Error>> + Send + 'a>> {
Box::pin(async move {
if self.banned_relays.contains(relay_url) {
Ok(AdmitStatus::rejected("banned"))
} else {
Ok(AdmitStatus::Success)
}
})
}
}
fn new_relay(url: RelayUrl, opts: RelayOptions) -> Relay {
Relay::builder(url).opts(opts).build()
}
async fn setup_subscription_relay() -> (SubscriptionId, Relay, MockRelay) {
let mock = MockRelay::run().await.unwrap();
let url = mock.url().await;
let relay: Relay = new_relay(url.clone(), RelayOptions::default());
relay.connect();
let filter = Filter::new().kind(Kind::TextNote);
let id = relay.subscribe(filter).await.unwrap();
(id, relay, mock)
}
fn check_relay_is_sleeping(relay: &Relay) {
assert_eq!(relay.status(), RelayStatus::Sleeping);
assert!(relay.status().can_connect());
assert!(!relay.inner.is_running());
}
#[tokio::test]
async fn test_status_with_reconnection_enabled() {
let mock = MockRelay::run().await.unwrap();
let url = mock.url().await;
let relay: Relay = new_relay(url, RelayOptions::default());
assert_eq!(relay.status(), RelayStatus::Initialized);
relay
.try_connect()
.timeout(Duration::from_secs(3))
.await
.unwrap();
assert_eq!(relay.status(), RelayStatus::Connected);
mock.shutdown();
time::sleep(Duration::from_millis(100)).await;
assert_eq!(relay.status(), RelayStatus::Disconnected);
assert!(relay.inner.is_running());
}
#[tokio::test]
async fn test_status_with_reconnection_disabled() {
let mock = MockRelay::run().await.unwrap();
let url = mock.url().await;
let relay: Relay = new_relay(url, RelayOptions::default().reconnect(false));
assert_eq!(relay.status(), RelayStatus::Initialized);
relay
.try_connect()
.timeout(Duration::from_secs(3))
.await
.unwrap();
assert_eq!(relay.status(), RelayStatus::Connected);
mock.shutdown();
time::sleep(Duration::from_millis(100)).await;
assert_eq!(relay.status(), RelayStatus::Terminated);
assert!(!relay.inner.is_running());
}
#[tokio::test]
async fn test_disconnect() {
let mock = MockRelay::run().await.unwrap();
let url = mock.url().await;
let relay: Relay = new_relay(url, RelayOptions::default());
assert_eq!(relay.status(), RelayStatus::Initialized);
relay
.try_connect()
.timeout(Duration::from_secs(3))
.await
.unwrap();
assert_eq!(relay.status(), RelayStatus::Connected);
relay.disconnect();
time::sleep(Duration::from_millis(100)).await;
assert_eq!(relay.status(), RelayStatus::Terminated);
assert!(!relay.inner.is_running());
}
#[tokio::test]
async fn test_disconnect_non_connected_relay() {
let url = RelayUrl::parse("wss://127.0.0.1:666").unwrap();
let opts = RelayOptions::default()
.adjust_retry_interval(false)
.retry_interval(Duration::from_secs(1));
let relay: Relay = new_relay(url, opts);
assert_eq!(relay.status(), RelayStatus::Initialized);
relay.connect();
time::sleep(Duration::from_secs(1)).await;
assert!(relay.inner.is_running());
assert_eq!(relay.status(), RelayStatus::Disconnected);
time::sleep(Duration::from_secs(3)).await;
relay.disconnect();
time::sleep(Duration::from_millis(100)).await;
assert_eq!(relay.status(), RelayStatus::Terminated);
assert!(!relay.inner.is_running());
}
#[tokio::test]
async fn test_connect() {
let mock = MockRelay::run().await.unwrap();
let url = mock.url().await;
let relay: Relay = new_relay(url, RelayOptions::default());
assert_eq!(relay.status(), RelayStatus::Initialized);
relay.connect();
relay.wait_for_connection(Duration::from_secs(1)).await;
assert_eq!(relay.status(), RelayStatus::Connected);
assert!(relay.inner.is_running());
}
#[tokio::test]
async fn test_connect_to_unreachable_relay() {
let url = RelayUrl::parse("wss://127.0.0.1:666").unwrap();
let relay: Relay = new_relay(url, RelayOptions::default());
assert_eq!(relay.status(), RelayStatus::Initialized);
relay.connect();
time::sleep(Duration::from_secs(1)).await;
assert_eq!(relay.status(), RelayStatus::Disconnected);
assert!(relay.inner.is_running());
}
#[tokio::test]
async fn test_disconnect_unresponsive_relay_that_connect() {
let opts = LocalRelayTestOptions {
unresponsive_connection: Some(Duration::from_secs(2)),
..Default::default()
};
let mock = MockRelay::run_with_opts(opts).await.unwrap();
let url = mock.url().await;
let relay: Relay = new_relay(url, RelayOptions::default());
assert_eq!(relay.status(), RelayStatus::Initialized);
relay.connect();
time::sleep(Duration::from_secs(1)).await;
assert_eq!(relay.status(), RelayStatus::Connecting);
time::sleep(Duration::from_secs(2)).await;
assert_eq!(relay.status(), RelayStatus::Connected);
relay.disconnect();
time::sleep(Duration::from_millis(100)).await;
assert_eq!(relay.status(), RelayStatus::Terminated);
assert!(!relay.inner.is_running());
}
#[tokio::test]
async fn test_disconnect_unresponsive_relay_that_not_connect() {
let opts = LocalRelayTestOptions {
unresponsive_connection: Some(Duration::from_secs(10)),
..Default::default()
};
let mock = MockRelay::run_with_opts(opts).await.unwrap();
let url = mock.url().await;
let relay: Relay = new_relay(url, RelayOptions::default());
assert_eq!(relay.status(), RelayStatus::Initialized);
relay.connect();
time::sleep(Duration::from_secs(1)).await;
assert_eq!(relay.status(), RelayStatus::Connecting);
relay.disconnect();
time::sleep(Duration::from_millis(100)).await;
assert_eq!(relay.status(), RelayStatus::Terminated);
assert!(!relay.inner.is_running());
}
#[tokio::test]
async fn test_disconnect_unresponsive_during_try_connect() {
let opts = LocalRelayTestOptions {
unresponsive_connection: Some(Duration::from_secs(10)),
..Default::default()
};
let mock = MockRelay::run_with_opts(opts).await.unwrap();
let url = mock.url().await;
let relay: Relay = new_relay(url, RelayOptions::default());
assert_eq!(relay.status(), RelayStatus::Initialized);
let r = relay.clone();
tokio::spawn(async move {
time::sleep(Duration::from_secs(3)).await;
r.disconnect();
});
let res = relay.try_connect().timeout(Duration::from_secs(7)).await;
let err = res.unwrap_err();
assert_eq!(err.kind(), ErrorKind::Rejected);
assert_eq!(err.to_string(), "received termination request");
assert_eq!(relay.status(), RelayStatus::Terminated);
assert!(!relay.inner.is_running());
}
#[tokio::test]
async fn test_ban_relay() {
let mock = MockRelay::run().await.unwrap();
let url = mock.url().await;
let relay = new_relay(url, RelayOptions::default());
assert_eq!(relay.status(), RelayStatus::Initialized);
relay
.try_connect()
.timeout(Duration::from_secs(2))
.await
.unwrap();
assert_eq!(relay.status(), RelayStatus::Connected);
relay.ban();
assert_eq!(relay.status(), RelayStatus::Banned);
assert!(!relay.inner.is_running());
let res = relay.try_connect().timeout(Duration::from_secs(2)).await;
let err = res.unwrap_err();
assert_eq!(err.kind(), ErrorKind::State);
assert_eq!(err.to_string(), "relay banned");
assert_eq!(relay.status(), RelayStatus::Banned);
relay.disconnect();
assert_eq!(relay.status(), RelayStatus::Banned);
let res = relay.inner.ensure_operational();
let err = res.unwrap_err();
assert_eq!(err.kind(), ErrorKind::State);
assert_eq!(err.to_string(), "relay banned");
}
#[tokio::test]
async fn test_shutdown() {
let mock = MockRelay::run().await.unwrap();
let url = mock.url().await;
let relay: Relay = new_relay(url, RelayOptions::default());
assert_eq!(relay.status(), RelayStatus::Initialized);
relay
.try_connect()
.timeout(Duration::from_secs(3))
.await
.unwrap();
assert_eq!(relay.status(), RelayStatus::Connected);
relay.shutdown();
time::sleep(Duration::from_millis(100)).await;
assert_eq!(relay.status(), RelayStatus::Shutdown);
assert!(!relay.inner.is_running());
let res = relay.try_connect().timeout(Duration::from_secs(3)).await;
let err = res.unwrap_err();
assert_eq!(err.kind(), ErrorKind::State);
assert_eq!(err.to_string(), "shutdown");
}
#[tokio::test]
async fn test_shutdown_on_drop() {
let mock = MockRelay::run().await.unwrap();
let url = mock.url().await;
let inner: InnerRelay = {
let relay: Relay = Relay::new(url);
relay
.try_connect()
.timeout(Duration::from_secs(3))
.await
.unwrap();
assert_eq!(relay.status(), RelayStatus::Connected);
let inner: InnerRelay = relay.inner.clone();
{
let r2 = relay.clone();
tokio::spawn(async move {
assert_eq!(r2.atomic_counter.load(Ordering::SeqCst), 2);
time::sleep(Duration::from_secs(1)).await;
});
}
time::sleep(Duration::from_secs(3)).await;
assert_eq!(relay.atomic_counter.load(Ordering::SeqCst), 1);
inner
};
time::sleep(Duration::from_secs(1)).await;
assert_eq!(inner.status(), RelayStatus::Shutdown);
assert!(!inner.is_running());
}
#[tokio::test]
async fn test_wait_for_connection() {
let opts = LocalRelayTestOptions {
unresponsive_connection: Some(Duration::from_secs(2)),
..Default::default()
};
let mock = MockRelay::run_with_opts(opts).await.unwrap();
let url = mock.url().await;
let relay: Relay = new_relay(url, RelayOptions::default());
assert_eq!(relay.status(), RelayStatus::Initialized);
relay.connect();
relay.wait_for_connection(Duration::from_millis(500)).await;
assert_eq!(relay.status(), RelayStatus::Connecting);
relay.wait_for_connection(Duration::from_secs(3)).await;
assert_eq!(relay.status(), RelayStatus::Connected);
}
#[tokio::test]
async fn test_unsubscribe() {
let (id, relay, _mock) = setup_subscription_relay().await;
time::sleep(Duration::from_secs(1)).await;
assert!(relay.subscription(&id).await.is_some());
relay.unsubscribe(&id).await.unwrap();
assert!(relay.subscription(&id).await.is_none());
}
#[tokio::test]
async fn test_unsubscribe_all() {
let (_id, relay, _mock) = setup_subscription_relay().await;
time::sleep(Duration::from_secs(1)).await;
relay.unsubscribe_all().await.unwrap();
relay.subscriptions().await.is_empty();
}
#[tokio::test]
async fn test_admit_connection() {
let mock = MockRelay::run().await.unwrap();
let url = mock.url().await;
let mut relay = new_relay(url.clone(), RelayOptions::default());
relay.inner.state.admit_policy = Some(Arc::new(CustomTestPolicy {
banned_relays: HashSet::from([url]),
}));
assert_eq!(relay.status(), RelayStatus::Initialized);
relay.connect();
time::sleep(Duration::from_secs(2)).await;
assert_eq!(relay.status(), RelayStatus::Terminated);
assert!(!relay.inner.is_running());
let res = relay.try_connect().timeout(Duration::from_secs(2)).await;
let err = res.unwrap_err();
assert_eq!(err.kind(), ErrorKind::Rejected);
assert_eq!(err.to_string(), "connection rejected: reason=banned");
assert_eq!(relay.status(), RelayStatus::Terminated);
assert!(!relay.inner.is_running());
}
#[tokio::test]
async fn test_sleep_when_idle() {
let mock = MockRelay::run().await.unwrap();
let url = mock.url().await;
let opts = RelayOptions::default().sleep_when_idle(SleepWhenIdle::Enabled {
timeout: Duration::from_secs(2),
});
let relay = new_relay(url, opts);
relay
.try_connect()
.timeout(Duration::from_secs(2))
.await
.unwrap();
assert_eq!(relay.status(), RelayStatus::Connected);
time::sleep(Duration::from_secs(3)).await;
check_relay_is_sleeping(&relay);
let event = EventBuilder::new(Kind::TextNote, "text wake-up")
.finalize(&Keys::generate())
.unwrap();
relay.send_event(&event).await.unwrap();
assert_eq!(relay.status(), RelayStatus::Connected);
time::sleep(Duration::from_secs(3)).await;
check_relay_is_sleeping(&relay);
let filter = Filter::new().kind(Kind::TextNote);
let _ = relay
.fetch_events(filter)
.timeout(Duration::from_secs(10))
.await
.unwrap();
assert_eq!(relay.status(), RelayStatus::Connected);
time::sleep(Duration::from_secs(3)).await;
check_relay_is_sleeping(&relay);
let filter = Filter::new().kind(Kind::TextNote);
let _ = relay.sync(filter).await.unwrap();
assert_eq!(relay.status(), RelayStatus::Connected);
time::sleep(Duration::from_secs(3)).await;
check_relay_is_sleeping(&relay);
}
#[tokio::test]
async fn test_sleep_when_idle_with_long_lived_subscription() {
let mock = MockRelay::run().await.unwrap();
let url = mock.url().await;
let opts = RelayOptions::default().sleep_when_idle(SleepWhenIdle::Enabled {
timeout: Duration::from_secs(2),
});
let relay = new_relay(url, opts);
relay
.try_connect()
.timeout(Duration::from_secs(2))
.await
.unwrap();
assert_eq!(relay.status(), RelayStatus::Connected);
let filter = Filter::new().kind(Kind::TextNote);
relay.subscribe(filter).await.unwrap();
time::sleep(Duration::from_secs(5)).await;
assert_eq!(relay.status(), RelayStatus::Connected);
}
#[tokio::test]
async fn test_terminate_notification_stream_on_shutdown() {
let mock = MockRelay::run().await.unwrap();
let url = mock.url().await;
let relay = Relay::new(url);
relay
.try_connect()
.timeout(Duration::from_secs(2))
.await
.unwrap();
assert_eq!(relay.status(), RelayStatus::Connected);
let r = relay.clone();
tokio::spawn(async move {
time::sleep(Duration::from_secs(3)).await;
r.shutdown()
});
let fut = async {
let mut notifications = relay.notifications();
let mut received = false;
while let Some(n) = notifications.next().await {
if let RelayNotification::RelayStatus {
status: RelayStatus::Shutdown,
} = n
{
received = true;
}
}
assert!(received);
};
tokio::time::timeout(Duration::from_secs(5), fut)
.await
.unwrap();
let mut notifications = relay.notifications();
let res = tokio::time::timeout(Duration::from_secs(1), notifications.next())
.await
.unwrap();
assert!(res.is_none());
}
}