1use 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#[derive(Debug, Clone, PartialEq, Eq)]
28pub struct DiscoveredAddress {
29 pub addr: SocketAddr,
31 pub source: DiscoverySource,
33 pub priority: u32,
35 pub ttl: Option<Duration>,
37}
38
39#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
41pub enum DiscoverySource {
42 LocalInterface,
44 PeerExchange,
46 Observed,
48 Config,
50 Manual,
52 Dns,
54}
55
56impl DiscoverySource {
57 pub fn base_priority(&self) -> u32 {
59 match self {
60 Self::Observed => 100, Self::LocalInterface => 90,
62 Self::PeerExchange => 80,
63 Self::Config => 70,
64 Self::Dns => 60,
65 Self::Manual => 50,
66 }
67 }
68}
69
70pub type DiscoveryResult = Result<DiscoveredAddress, DiscoveryError>;
72
73#[derive(Debug, Clone)]
75pub struct DiscoveryError {
76 pub message: String,
78 pub source: Option<DiscoverySource>,
80 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
92pub trait Discovery: Send + Sync + 'static {
97 fn discover(
102 &self,
103 peer_id: &PeerId,
104 ) -> Pin<Box<dyn Stream<Item = DiscoveryResult> + Send + 'static>>;
105
106 fn name(&self) -> &'static str;
108}
109
110#[derive(Default)]
112pub struct ConcurrentDiscovery {
113 sources: Vec<Arc<dyn Discovery>>,
114}
115
116impl ConcurrentDiscovery {
117 pub fn new() -> Self {
119 Self {
120 sources: Vec::new(),
121 }
122 }
123
124 pub fn add_source<D: Discovery>(&mut self, source: D) {
126 self.sources.push(Arc::new(source));
127 }
128
129 pub fn add_boxed_source(&mut self, source: Arc<dyn Discovery>) {
131 self.sources.push(source);
132 }
133
134 pub fn builder() -> ConcurrentDiscoveryBuilder {
136 ConcurrentDiscoveryBuilder::new()
137 }
138
139 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 pub fn source_count(&self) -> usize {
152 self.sources.len()
153 }
154}
155
156#[derive(Default)]
158pub struct ConcurrentDiscoveryBuilder {
159 sources: Vec<Arc<dyn Discovery>>,
160}
161
162impl ConcurrentDiscoveryBuilder {
163 pub fn new() -> Self {
165 Self {
166 sources: Vec::new(),
167 }
168 }
169
170 pub fn with_source<D: Discovery>(mut self, source: D) -> Self {
172 self.sources.push(Arc::new(source));
173 self
174 }
175
176 pub fn build(self) -> ConcurrentDiscovery {
178 ConcurrentDiscovery {
179 sources: self.sources,
180 }
181 }
182}
183
184pub 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 if this.completed.iter().all(|&c| c) {
205 return Poll::Ready(None);
206 }
207
208 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 if this.completed.iter().all(|&c| c) {
227 Poll::Ready(None)
228 } else {
229 Poll::Pending
230 }
231 }
232}
233
234pub 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 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 pub fn sender(&self) -> mpsc::Sender<DiscoveredAddress> {
254 self.sender.clone()
255 }
256
257 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
287pub struct StaticDiscovery {
289 addresses: Vec<DiscoveredAddress>,
290}
291
292impl StaticDiscovery {
293 pub fn new(addresses: Vec<DiscoveredAddress>) -> Self {
295 Self { addresses }
296 }
297
298 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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 drop(sender1);
658 drop(sender2);
659 }
660}