use std::collections::HashMap;
use std::future::IntoFuture;
use nostr::types::RelayUrl;
use crate::client::Client;
use crate::future::BoxedFuture;
use crate::relay::{Relay, RelayCapabilities};
enum Policy {
All,
WithCapabilities(RelayCapabilities),
}
#[must_use = "Does nothing unless you await!"]
pub struct GetRelays<'client> {
client: &'client Client,
policy: Policy,
}
impl<'client> GetRelays<'client> {
pub(crate) fn new(client: &'client Client) -> Self {
Self {
client,
policy: Policy::WithCapabilities(RelayCapabilities::READ | RelayCapabilities::WRITE),
}
}
#[inline]
pub fn all(mut self) -> Self {
self.policy = Policy::All;
self
}
#[inline]
pub fn with_capabilities(mut self, capabilities: RelayCapabilities) -> Self {
self.policy = Policy::WithCapabilities(capabilities);
self
}
}
impl<'client> IntoFuture for GetRelays<'client> {
type Output = HashMap<RelayUrl, Relay>;
type IntoFuture = BoxedFuture<'client, Self::Output>;
fn into_future(self) -> Self::IntoFuture {
Box::pin(async move {
match self.policy {
Policy::All => self.client.pool().all_relays().await,
Policy::WithCapabilities(capabilities) => {
self.client.pool().relays_with_any_cap(capabilities).await
}
}
})
}
}
#[cfg(test)]
mod tests {
use super::*;
async fn setup_client() -> Client {
let client = Client::default();
client.add_relay("wss://relay1.example.com").await.unwrap();
client.add_relay("wss://relay2.example.com").await.unwrap();
client
.add_relay("wss://relay3.example.com")
.capabilities(RelayCapabilities::DISCOVERY)
.await
.unwrap();
client
.add_relay("wss://relay4.example.com")
.capabilities(RelayCapabilities::GOSSIP)
.await
.unwrap();
client
}
#[tokio::test]
async fn test_get_relays() {
let client = setup_client().await;
let url1 = RelayUrl::parse("wss://relay1.example.com").unwrap();
let url2 = RelayUrl::parse("wss://relay2.example.com").unwrap();
let relays = client.relays().await;
assert_eq!(relays.len(), 2);
assert!(relays.contains_key(&url1));
assert!(relays.contains_key(&url2));
}
#[tokio::test]
async fn test_get_relays_with_capabilities() {
let client = setup_client().await;
let url3 = RelayUrl::parse("wss://relay3.example.com").unwrap();
let url4 = RelayUrl::parse("wss://relay4.example.com").unwrap();
let relays = client
.relays()
.with_capabilities(RelayCapabilities::GOSSIP)
.await;
assert_eq!(relays.len(), 1);
assert!(relays.contains_key(&url4));
let relays = client
.relays()
.with_capabilities(RelayCapabilities::GOSSIP | RelayCapabilities::DISCOVERY)
.await;
assert_eq!(relays.len(), 2);
assert!(relays.contains_key(&url3));
assert!(relays.contains_key(&url4));
}
#[tokio::test]
async fn test_get_all_relays() {
let client = setup_client().await;
let url1 = RelayUrl::parse("wss://relay1.example.com").unwrap();
let url2 = RelayUrl::parse("wss://relay2.example.com").unwrap();
let url3 = RelayUrl::parse("wss://relay3.example.com").unwrap();
let url4 = RelayUrl::parse("wss://relay4.example.com").unwrap();
let relays = client.relays().all().await;
assert_eq!(relays.len(), 4);
assert!(relays.contains_key(&url1));
assert!(relays.contains_key(&url2));
assert!(relays.contains_key(&url3));
assert!(relays.contains_key(&url4));
}
}