Skip to main content

ant_quic/
discovery_trait.rs

1// Copyright 2024 Saorsa Labs Ltd.
2//
3// This Saorsa Network Software is licensed under the General Public License (GPL), version 3.
4// Please see the file LICENSE-GPL, or visit <http://www.gnu.org/licenses/> for the full text.
5//
6// Full details available at https://saorsalabs.com/licenses
7
8//! Discovery trait for stream composition
9//!
10//! Provides a trait-based abstraction for address discovery that allows
11//! composing multiple discovery sources into a unified stream.
12//!
13//! This is inspired by iroh's `Discovery` trait and `ConcurrentDiscovery`.
14
15use std::net::SocketAddr;
16use std::pin::Pin;
17use std::sync::Arc;
18use std::task::{Context, Poll};
19use std::time::Duration;
20
21use futures_util::stream::Stream;
22use tokio::sync::mpsc;
23
24use crate::nat_traversal_api::PeerId;
25
26/// Information about a discovered address
27#[derive(Debug, Clone, PartialEq, Eq)]
28pub struct DiscoveredAddress {
29    /// The discovered socket address
30    pub addr: SocketAddr,
31    /// Source of the discovery
32    pub source: DiscoverySource,
33    /// Priority of this address (higher = better)
34    pub priority: u32,
35    /// Time-to-live for this discovery
36    pub ttl: Option<Duration>,
37}
38
39/// Source of address discovery
40#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
41pub enum DiscoverySource {
42    /// Discovered from local network interfaces
43    LocalInterface,
44    /// Discovered via peer exchange
45    PeerExchange,
46    /// Observed by a remote peer
47    Observed,
48    /// From configuration or known peers
49    Config,
50    /// Manual/explicit discovery
51    Manual,
52    /// From DNS resolution
53    Dns,
54}
55
56impl DiscoverySource {
57    /// Get base priority for this source
58    pub fn base_priority(&self) -> u32 {
59        match self {
60            Self::Observed => 100, // Highest - verified by peer
61            Self::LocalInterface => 90,
62            Self::PeerExchange => 80,
63            Self::Config => 70,
64            Self::Dns => 60,
65            Self::Manual => 50,
66        }
67    }
68}
69
70/// Result of a discovery operation
71pub type DiscoveryResult = Result<DiscoveredAddress, DiscoveryError>;
72
73/// Error from discovery operations
74#[derive(Debug, Clone)]
75pub struct DiscoveryError {
76    /// Error message
77    pub message: String,
78    /// Source that failed
79    pub source: Option<DiscoverySource>,
80    /// Whether this error is retryable
81    pub retryable: bool,
82}
83
84impl std::fmt::Display for DiscoveryError {
85    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
86        write!(f, "Discovery error: {}", self.message)
87    }
88}
89
90impl std::error::Error for DiscoveryError {}
91
92/// Trait for address discovery sources
93///
94/// Implementations provide a stream of discovered addresses
95/// that can be composed with other discovery sources.
96pub trait Discovery: Send + Sync + 'static {
97    /// Discover addresses for a given peer
98    ///
99    /// Returns a stream of discovered addresses. The stream may
100    /// continue indefinitely or terminate when discovery is complete.
101    fn discover(
102        &self,
103        peer_id: &PeerId,
104    ) -> Pin<Box<dyn Stream<Item = DiscoveryResult> + Send + 'static>>;
105
106    /// Get the name of this discovery source (for logging)
107    fn name(&self) -> &'static str;
108}
109
110/// Combines multiple discovery sources into a concurrent stream
111#[derive(Default)]
112pub struct ConcurrentDiscovery {
113    sources: Vec<Arc<dyn Discovery>>,
114}
115
116impl ConcurrentDiscovery {
117    /// Create a new concurrent discovery with no sources
118    pub fn new() -> Self {
119        Self {
120            sources: Vec::new(),
121        }
122    }
123
124    /// Add a discovery source
125    pub fn add_source<D: Discovery>(&mut self, source: D) {
126        self.sources.push(Arc::new(source));
127    }
128
129    /// Add a boxed discovery source
130    pub fn add_boxed_source(&mut self, source: Arc<dyn Discovery>) {
131        self.sources.push(source);
132    }
133
134    /// Create a builder for fluent construction
135    pub fn builder() -> ConcurrentDiscoveryBuilder {
136        ConcurrentDiscoveryBuilder::new()
137    }
138
139    /// Discover addresses from all sources concurrently
140    pub fn discover(&self, peer_id: &PeerId) -> ConcurrentDiscoveryStream {
141        let mut streams = Vec::new();
142
143        for source in &self.sources {
144            streams.push(source.discover(peer_id));
145        }
146
147        ConcurrentDiscoveryStream::new(streams)
148    }
149
150    /// Number of discovery sources
151    pub fn source_count(&self) -> usize {
152        self.sources.len()
153    }
154}
155
156/// Builder for ConcurrentDiscovery
157#[derive(Default)]
158pub struct ConcurrentDiscoveryBuilder {
159    sources: Vec<Arc<dyn Discovery>>,
160}
161
162impl ConcurrentDiscoveryBuilder {
163    /// Create a new builder
164    pub fn new() -> Self {
165        Self {
166            sources: Vec::new(),
167        }
168    }
169
170    /// Add a discovery source
171    pub fn with_source<D: Discovery>(mut self, source: D) -> Self {
172        self.sources.push(Arc::new(source));
173        self
174    }
175
176    /// Build the concurrent discovery
177    pub fn build(self) -> ConcurrentDiscovery {
178        ConcurrentDiscovery {
179            sources: self.sources,
180        }
181    }
182}
183
184/// Stream that polls multiple discovery sources concurrently
185pub struct ConcurrentDiscoveryStream {
186    streams: Vec<Pin<Box<dyn Stream<Item = DiscoveryResult> + Send + 'static>>>,
187    completed: Vec<bool>,
188}
189
190impl ConcurrentDiscoveryStream {
191    fn new(streams: Vec<Pin<Box<dyn Stream<Item = DiscoveryResult> + Send + 'static>>>) -> Self {
192        let completed = vec![false; streams.len()];
193        Self { streams, completed }
194    }
195}
196
197impl Stream for ConcurrentDiscoveryStream {
198    type Item = DiscoveryResult;
199
200    fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
201        let this = &mut *self;
202
203        // Check if all streams are done
204        if this.completed.iter().all(|&c| c) {
205            return Poll::Ready(None);
206        }
207
208        // Poll each stream, returning the first ready result
209        for i in 0..this.streams.len() {
210            if this.completed[i] {
211                continue;
212            }
213
214            match this.streams[i].as_mut().poll_next(cx) {
215                Poll::Ready(Some(result)) => {
216                    return Poll::Ready(Some(result));
217                }
218                Poll::Ready(None) => {
219                    this.completed[i] = true;
220                }
221                Poll::Pending => {}
222            }
223        }
224
225        // Check again if all completed during this poll
226        if this.completed.iter().all(|&c| c) {
227            Poll::Ready(None)
228        } else {
229            Poll::Pending
230        }
231    }
232}
233
234/// A simple discovery source that yields addresses from a channel
235pub struct ChannelDiscovery {
236    name: &'static str,
237    sender: mpsc::Sender<DiscoveredAddress>,
238    receiver: Arc<tokio::sync::Mutex<mpsc::Receiver<DiscoveredAddress>>>,
239}
240
241impl ChannelDiscovery {
242    /// Create a new channel-based discovery
243    pub fn new(name: &'static str, buffer_size: usize) -> Self {
244        let (sender, receiver) = mpsc::channel(buffer_size);
245        Self {
246            name,
247            sender,
248            receiver: Arc::new(tokio::sync::Mutex::new(receiver)),
249        }
250    }
251
252    /// Get a sender to push discovered addresses
253    pub fn sender(&self) -> mpsc::Sender<DiscoveredAddress> {
254        self.sender.clone()
255    }
256
257    /// Push a discovered address
258    pub async fn push(
259        &self,
260        addr: DiscoveredAddress,
261    ) -> Result<(), mpsc::error::SendError<DiscoveredAddress>> {
262        self.sender.send(addr).await
263    }
264}
265
266impl Discovery for ChannelDiscovery {
267    fn discover(
268        &self,
269        _peer_id: &PeerId,
270    ) -> Pin<Box<dyn Stream<Item = DiscoveryResult> + Send + 'static>> {
271        let receiver = self.receiver.clone();
272
273        Box::pin(futures_util::stream::unfold(
274            receiver,
275            |receiver| async move {
276                let mut guard = receiver.lock().await;
277                guard.recv().await.map(|addr| (Ok(addr), receiver.clone()))
278            },
279        ))
280    }
281
282    fn name(&self) -> &'static str {
283        self.name
284    }
285}
286
287/// Discovery source from static/configured addresses
288pub struct StaticDiscovery {
289    addresses: Vec<DiscoveredAddress>,
290}
291
292impl StaticDiscovery {
293    /// Create a new static discovery with the given addresses
294    pub fn new(addresses: Vec<DiscoveredAddress>) -> Self {
295        Self { addresses }
296    }
297
298    /// Create from socket addresses with default settings
299    pub fn from_addrs(addrs: Vec<SocketAddr>) -> Self {
300        let addresses = addrs
301            .into_iter()
302            .map(|addr| DiscoveredAddress {
303                addr,
304                source: DiscoverySource::Config,
305                priority: DiscoverySource::Config.base_priority(),
306                ttl: None,
307            })
308            .collect();
309        Self { addresses }
310    }
311}
312
313impl Discovery for StaticDiscovery {
314    fn discover(
315        &self,
316        _peer_id: &PeerId,
317    ) -> Pin<Box<dyn Stream<Item = DiscoveryResult> + Send + 'static>> {
318        let addresses = self.addresses.clone();
319        Box::pin(futures_util::stream::iter(addresses.into_iter().map(Ok)))
320    }
321
322    fn name(&self) -> &'static str {
323        "static"
324    }
325}
326
327#[cfg(test)]
328mod tests {
329    use super::*;
330    use futures_util::StreamExt;
331
332    fn test_addr(port: u16) -> SocketAddr {
333        format!("192.168.1.1:{}", port).parse().unwrap()
334    }
335
336    fn test_peer_id() -> PeerId {
337        PeerId([0u8; 32])
338    }
339
340    fn test_discovered_addr(
341        port: u16,
342        source: DiscoverySource,
343        priority: u32,
344    ) -> DiscoveredAddress {
345        DiscoveredAddress {
346            addr: test_addr(port),
347            source,
348            priority,
349            ttl: None,
350        }
351    }
352
353    // DiscoverySource tests
354
355    #[test]
356    fn test_discovery_source_priority_values() {
357        assert_eq!(DiscoverySource::Observed.base_priority(), 100);
358        assert_eq!(DiscoverySource::LocalInterface.base_priority(), 90);
359        assert_eq!(DiscoverySource::PeerExchange.base_priority(), 80);
360        assert_eq!(DiscoverySource::Config.base_priority(), 70);
361        assert_eq!(DiscoverySource::Dns.base_priority(), 60);
362        assert_eq!(DiscoverySource::Manual.base_priority(), 50);
363    }
364
365    #[test]
366    fn test_discovery_source_order() {
367        assert!(
368            DiscoverySource::Observed.base_priority()
369                > DiscoverySource::LocalInterface.base_priority()
370        );
371        assert!(
372            DiscoverySource::LocalInterface.base_priority()
373                > DiscoverySource::PeerExchange.base_priority()
374        );
375        assert!(
376            DiscoverySource::PeerExchange.base_priority() > DiscoverySource::Config.base_priority()
377        );
378        assert!(DiscoverySource::Config.base_priority() > DiscoverySource::Dns.base_priority());
379        assert!(DiscoverySource::Dns.base_priority() > DiscoverySource::Manual.base_priority());
380    }
381
382    #[test]
383    fn test_discovery_source_equality() {
384        assert_eq!(DiscoverySource::Observed, DiscoverySource::Observed);
385        assert_ne!(DiscoverySource::Observed, DiscoverySource::Config);
386    }
387
388    #[test]
389    fn test_discovery_source_clone_copy() {
390        let a = DiscoverySource::Observed;
391        let b = a;
392        assert_eq!(a, b);
393    }
394
395    #[test]
396    fn test_discovery_source_debug() {
397        assert_eq!(format!("{:?}", DiscoverySource::Observed), "Observed");
398        assert_eq!(format!("{:?}", DiscoverySource::Dns), "Dns");
399    }
400
401    // DiscoveredAddress tests
402
403    #[test]
404    fn test_discovered_address_clone() {
405        let addr = test_discovered_addr(5000, DiscoverySource::Observed, 100);
406        let cloned = addr.clone();
407        assert_eq!(addr, cloned);
408    }
409
410    #[test]
411    fn test_discovered_address_debug() {
412        let addr = test_discovered_addr(8080, DiscoverySource::Config, 70);
413        let debug = format!("{addr:?}");
414        assert!(debug.contains("8080"));
415        assert!(debug.contains("Config"));
416    }
417
418    #[test]
419    fn test_discovered_address_equality() {
420        let a = test_discovered_addr(5000, DiscoverySource::Config, 70);
421        let b = test_discovered_addr(5000, DiscoverySource::Config, 70);
422        let c = test_discovered_addr(5001, DiscoverySource::Config, 70);
423        assert_eq!(a, b);
424        assert_ne!(a, c);
425    }
426
427    #[test]
428    fn test_discovered_address_different_sources_not_equal() {
429        let a = test_discovered_addr(5000, DiscoverySource::Config, 70);
430        let b = test_discovered_addr(5000, DiscoverySource::Observed, 100);
431        assert_ne!(a, b);
432    }
433
434    // DiscoveryError tests
435
436    #[test]
437    fn test_discovery_error_display() {
438        let err = DiscoveryError {
439            message: "test error".to_string(),
440            source: Some(DiscoverySource::Dns),
441            retryable: true,
442        };
443        let display = err.to_string();
444        assert!(display.contains("test error"));
445        assert!(display.contains("Discovery error"));
446    }
447
448    #[test]
449    fn test_discovery_error_clone() {
450        let err = DiscoveryError {
451            message: "err".to_string(),
452            source: None,
453            retryable: false,
454        };
455        let cloned = err.clone();
456        assert_eq!(err.message, cloned.message);
457        assert_eq!(err.source, cloned.source);
458        assert_eq!(err.retryable, cloned.retryable);
459    }
460
461    #[test]
462    fn test_discovery_error_debug() {
463        let err = DiscoveryError {
464            message: "debug me".to_string(),
465            source: Some(DiscoverySource::Config),
466            retryable: true,
467        };
468        let debug = format!("{err:?}");
469        assert!(debug.contains("debug me"));
470        assert!(debug.contains("Config"));
471    }
472
473    #[test]
474    fn test_discovery_error_retryable_flag() {
475        let err_retryable = DiscoveryError {
476            message: "retry".to_string(),
477            source: None,
478            retryable: true,
479        };
480        let err_not = DiscoveryError {
481            message: "fatal".to_string(),
482            source: None,
483            retryable: false,
484        };
485        assert!(err_retryable.retryable);
486        assert!(!err_not.retryable);
487    }
488
489    #[test]
490    fn test_discovery_error_with_source() {
491        let err = DiscoveryError {
492            message: "dns failed".to_string(),
493            source: Some(DiscoverySource::Dns),
494            retryable: true,
495        };
496        assert_eq!(err.source, Some(DiscoverySource::Dns));
497    }
498
499    // StaticDiscovery tests
500
501    #[test]
502    fn test_static_discovery_name() {
503        let discovery = StaticDiscovery::from_addrs(vec![]);
504        assert_eq!(discovery.name(), "static");
505    }
506
507    #[tokio::test]
508    async fn test_static_discovery() {
509        let addrs = vec![test_addr(5000), test_addr(5001)];
510        let discovery = StaticDiscovery::from_addrs(addrs.clone());
511        let mut stream = discovery.discover(&test_peer_id());
512        let first = stream.next().await.unwrap().unwrap();
513        assert_eq!(first.addr, addrs[0]);
514        let second = stream.next().await.unwrap().unwrap();
515        assert_eq!(second.addr, addrs[1]);
516        assert!(stream.next().await.is_none());
517    }
518
519    #[tokio::test]
520    async fn test_static_discovery_empty() {
521        let discovery = StaticDiscovery::from_addrs(vec![]);
522        let mut stream = discovery.discover(&test_peer_id());
523        assert!(stream.next().await.is_none());
524    }
525
526    #[tokio::test]
527    async fn test_static_discovery_new() {
528        let addr = test_discovered_addr(9000, DiscoverySource::Config, 70);
529        let discovery = StaticDiscovery::new(vec![addr.clone()]);
530        assert_eq!(discovery.name(), "static");
531        let mut stream = discovery.discover(&test_peer_id());
532        let result = stream.next().await.unwrap().unwrap();
533        assert_eq!(result.addr, addr.addr);
534    }
535
536    // ConcurrentDiscovery tests
537
538    #[test]
539    fn test_concurrent_discovery_new_empty() {
540        let discovery = ConcurrentDiscovery::new();
541        assert_eq!(discovery.source_count(), 0);
542    }
543
544    #[test]
545    fn test_concurrent_discovery_default() {
546        let discovery = ConcurrentDiscovery::default();
547        assert_eq!(discovery.source_count(), 0);
548    }
549
550    #[test]
551    fn test_concurrent_discovery_add_source() {
552        let mut discovery = ConcurrentDiscovery::new();
553        assert_eq!(discovery.source_count(), 0);
554        discovery.add_source(StaticDiscovery::from_addrs(vec![test_addr(5000)]));
555        assert_eq!(discovery.source_count(), 1);
556    }
557
558    #[test]
559    fn test_concurrent_discovery_multiple_sources() {
560        let mut discovery = ConcurrentDiscovery::new();
561        discovery.add_source(StaticDiscovery::from_addrs(vec![test_addr(5000)]));
562        discovery.add_source(StaticDiscovery::from_addrs(vec![test_addr(6000)]));
563        assert_eq!(discovery.source_count(), 2);
564    }
565
566    #[test]
567    fn test_concurrent_discovery_add_boxed_source() {
568        let mut discovery = ConcurrentDiscovery::new();
569        let source: Arc<dyn Discovery> = Arc::new(StaticDiscovery::from_addrs(vec![]));
570        discovery.add_boxed_source(source);
571        assert_eq!(discovery.source_count(), 1);
572    }
573
574    #[tokio::test]
575    async fn test_concurrent_discovery_with_two_sources() {
576        let addrs1 = vec![test_addr(5000)];
577        let addrs2 = vec![test_addr(6000)];
578
579        let discovery = ConcurrentDiscovery::builder()
580            .with_source(StaticDiscovery::from_addrs(addrs1))
581            .with_source(StaticDiscovery::from_addrs(addrs2))
582            .build();
583
584        assert_eq!(discovery.source_count(), 2);
585        let mut stream = discovery.discover(&test_peer_id());
586        let mut found_ports = vec![];
587        while let Some(result) = stream.next().await {
588            found_ports.push(result.unwrap().addr.port());
589        }
590        assert!(found_ports.contains(&5000));
591        assert!(found_ports.contains(&6000));
592    }
593
594    #[tokio::test]
595    async fn test_empty_concurrent_discovery() {
596        let discovery = ConcurrentDiscovery::new();
597        let mut stream = discovery.discover(&test_peer_id());
598        assert!(stream.next().await.is_none());
599    }
600
601    // ConcurrentDiscoveryBuilder tests
602
603    #[test]
604    fn test_builder_empty() {
605        let discovery = ConcurrentDiscoveryBuilder::new().build();
606        assert_eq!(discovery.source_count(), 0);
607    }
608
609    #[test]
610    fn test_builder_pattern() {
611        let discovery = ConcurrentDiscoveryBuilder::new()
612            .with_source(StaticDiscovery::from_addrs(vec![test_addr(5000)]))
613            .with_source(StaticDiscovery::from_addrs(vec![test_addr(6000)]))
614            .build();
615        assert_eq!(discovery.source_count(), 2);
616    }
617
618    #[test]
619    fn test_builder_single_source() {
620        let discovery = ConcurrentDiscoveryBuilder::new()
621            .with_source(StaticDiscovery::from_addrs(vec![test_addr(5000)]))
622            .build();
623        assert_eq!(discovery.source_count(), 1);
624    }
625
626    // ChannelDiscovery tests
627
628    #[tokio::test]
629    async fn test_channel_discovery() {
630        let discovery = ChannelDiscovery::new("test", 10);
631        let sender = discovery.sender();
632        tokio::spawn(async move {
633            sender
634                .send(test_discovered_addr(7000, DiscoverySource::Observed, 100))
635                .await
636                .unwrap();
637        });
638        let mut stream = discovery.discover(&test_peer_id());
639        let result = tokio::time::timeout(Duration::from_millis(100), stream.next()).await;
640        assert!(result.is_ok());
641        let addr = result.unwrap().unwrap().unwrap();
642        assert_eq!(addr.addr.port(), 7000);
643    }
644
645    #[test]
646    fn test_channel_discovery_name() {
647        let discovery = ChannelDiscovery::new("custom-name", 5);
648        assert_eq!(discovery.name(), "custom-name");
649    }
650
651    #[test]
652    fn test_channel_discovery_sender_is_cloneable() {
653        let discovery = ChannelDiscovery::new("clone-test", 5);
654        let sender1 = discovery.sender();
655        let sender2 = discovery.sender();
656        // Both senders should be usable
657        drop(sender1);
658        drop(sender2);
659    }
660}