1use super::entry::CachedPeer;
11use crate::reachability::{ReachabilityScope, socket_addr_scope};
12use rand::Rng;
13use std::collections::HashSet;
14
15#[derive(Debug, Clone, Copy)]
17pub enum SelectionStrategy {
18 BestFirst,
20 EpsilonGreedy {
22 epsilon: f64,
24 },
25 Random,
27}
28
29impl Default for SelectionStrategy {
30 fn default() -> Self {
31 Self::EpsilonGreedy { epsilon: 0.1 }
32 }
33}
34
35const fn scope_rank(scope: Option<ReachabilityScope>) -> u8 {
36 match scope {
37 Some(ReachabilityScope::Loopback) => 1,
38 Some(ReachabilityScope::LocalNetwork) => 2,
39 Some(ReachabilityScope::Global) => 3,
40 None => 0,
41 }
42}
43
44fn helper_preference_score(
45 peer: &CachedPeer,
46 require_relay: bool,
47 require_coordination: bool,
48) -> u8 {
49 if !require_relay && !require_coordination {
50 return 0;
51 }
52
53 let scope_score = scope_rank(peer.capabilities.direct_reachability_scope);
54 let global_bonus = u8::from(
55 (require_relay && peer.capabilities.supports_relay)
56 || (require_coordination && peer.capabilities.supports_coordination),
57 );
58
59 scope_score.saturating_mul(2).saturating_add(global_bonus)
60}
61
62fn scope_match_score(observed: ReachabilityScope, target_scope: Option<ReachabilityScope>) -> u8 {
63 match target_scope {
64 Some(ReachabilityScope::Global) => u8::from(observed == ReachabilityScope::Global) * 3,
65 Some(ReachabilityScope::LocalNetwork) => match observed {
66 ReachabilityScope::LocalNetwork => 3,
67 ReachabilityScope::Global => 2,
68 ReachabilityScope::Loopback => 0,
69 },
70 Some(ReachabilityScope::Loopback) => u8::from(observed == ReachabilityScope::Loopback) * 3,
71 None => scope_rank(Some(observed)),
72 }
73}
74
75fn best_relay_score(peer: &CachedPeer, target: std::net::SocketAddr) -> u8 {
76 let target_scope = socket_addr_scope(target);
77 let target_is_ipv4 = target.is_ipv4();
78
79 let direct_score = peer
80 .capabilities
81 .reachable_addresses
82 .iter()
83 .filter(|entry| entry.address.is_ipv4() == target_is_ipv4)
84 .filter_map(|entry| {
85 let scope_score = scope_match_score(entry.scope, target_scope);
86 (scope_score > 0).then_some(scope_score.saturating_add(4))
87 })
88 .max()
89 .unwrap_or(0);
90
91 let observed_score = peer
92 .capabilities
93 .external_addresses
94 .iter()
95 .filter(|addr| addr.is_ipv4() == target_is_ipv4)
96 .filter_map(|addr| {
97 let scope_score = socket_addr_scope(*addr)
98 .map(|scope| scope_match_score(scope, target_scope))
99 .unwrap_or(0);
100 (scope_score > 0).then_some(scope_score.saturating_add(2))
101 })
102 .max()
103 .unwrap_or(0);
104
105 let stored_score = peer
106 .addresses
107 .iter()
108 .filter(|addr| addr.is_ipv4() == target_is_ipv4)
109 .filter_map(|addr| {
110 let scope_score = socket_addr_scope(*addr)
111 .map(|scope| scope_match_score(scope, target_scope))
112 .unwrap_or(0);
113 (scope_score > 0).then_some(scope_score)
114 })
115 .max()
116 .unwrap_or(0);
117
118 direct_score.max(observed_score).max(stored_score)
119}
120
121pub fn select_epsilon_greedy(peers: &[CachedPeer], count: usize, epsilon: f64) -> Vec<&CachedPeer> {
134 if peers.is_empty() || count == 0 {
135 return Vec::new();
136 }
137
138 let mut rng = rand::thread_rng();
139 let mut selected = Vec::with_capacity(count.min(peers.len()));
140 let mut used_indices = HashSet::new();
141
142 let mut sorted_indices: Vec<usize> = (0..peers.len()).collect();
144 sorted_indices.sort_by(|&a, &b| {
145 peers[b]
146 .quality_score
147 .partial_cmp(&peers[a].quality_score)
148 .unwrap_or(std::cmp::Ordering::Equal)
149 });
150
151 let target_count = count.min(peers.len());
153 let explore_count = ((target_count as f64) * epsilon).ceil() as usize;
154 let exploit_count = target_count.saturating_sub(explore_count);
155
156 for &idx in sorted_indices.iter().take(exploit_count) {
158 if used_indices.insert(idx) && selected.len() < target_count {
159 selected.push(&peers[idx]);
160 }
161 }
162
163 let remaining: Vec<usize> = (0..peers.len())
166 .filter(|idx| !used_indices.contains(idx))
167 .collect();
168
169 if !remaining.is_empty() && selected.len() < target_count {
170 let (untested, tested): (Vec<_>, Vec<_>) = remaining.iter().partition(|&&idx| {
172 peers[idx].stats.success_count + peers[idx].stats.failure_count == 0
173 });
174
175 let explore_pool = if !untested.is_empty() {
177 untested
178 } else {
179 tested
180 };
181
182 let mut explore_indices: Vec<usize> = explore_pool.into_iter().copied().collect();
184 for i in (1..explore_indices.len()).rev() {
186 let j = rng.gen_range(0..=i);
187 explore_indices.swap(i, j);
188 }
189
190 for &idx in explore_indices.iter() {
191 if selected.len() >= target_count {
192 break;
193 }
194 if used_indices.insert(idx) {
195 selected.push(&peers[idx]);
196 }
197 }
198 }
199
200 for &idx in &sorted_indices {
202 if selected.len() >= target_count {
203 break;
204 }
205 if used_indices.insert(idx) {
206 selected.push(&peers[idx]);
207 }
208 }
209
210 selected
211}
212
213#[allow(dead_code)]
219pub fn select_with_capabilities(
220 peers: &[CachedPeer],
221 count: usize,
222 require_relay: bool,
223 require_coordination: bool,
224) -> Vec<&CachedPeer> {
225 if peers.is_empty() || count == 0 {
226 return Vec::new();
227 }
228
229 let mut candidates: Vec<&CachedPeer> = peers.iter().collect();
230
231 candidates.sort_by(|a, b| {
234 let a_pref = helper_preference_score(a, require_relay, require_coordination);
235 let b_pref = helper_preference_score(b, require_relay, require_coordination);
236 b_pref
237 .cmp(&a_pref)
238 .then_with(|| {
239 b.capabilities
240 .direct_reachability_scope
241 .cmp(&a.capabilities.direct_reachability_scope)
242 })
243 .then_with(|| {
244 b.quality_score
245 .partial_cmp(&a.quality_score)
246 .unwrap_or(std::cmp::Ordering::Equal)
247 })
248 });
249
250 candidates.into_iter().take(count).collect()
251}
252
253pub fn select_relays_for_target(
265 peers: &[CachedPeer],
266 count: usize,
267 target: std::net::SocketAddr,
268 prefer_dual_stack: bool,
269) -> Vec<&CachedPeer> {
270 if peers.is_empty() || count == 0 {
271 return Vec::new();
272 }
273
274 let target_is_ipv4 = target.is_ipv4();
275
276 let mut candidates: Vec<&CachedPeer> = peers
277 .iter()
278 .filter(|p| {
279 let preferred = p.preferred_addresses();
280 preferred.is_empty()
281 || preferred
282 .iter()
283 .any(|addr| addr.is_ipv4() == target_is_ipv4)
284 })
285 .collect();
286
287 if candidates.is_empty() {
288 return Vec::new();
289 }
290
291 candidates.sort_by(|a, b| {
292 let a_pref = best_relay_score(a, target);
293 let b_pref = best_relay_score(b, target);
294
295 b_pref
296 .cmp(&a_pref)
297 .then_with(|| {
298 if prefer_dual_stack {
299 b.capabilities
300 .supports_dual_stack()
301 .cmp(&a.capabilities.supports_dual_stack())
302 } else {
303 std::cmp::Ordering::Equal
304 }
305 })
306 .then_with(|| {
307 b.capabilities
308 .direct_reachability_scope
309 .cmp(&a.capabilities.direct_reachability_scope)
310 })
311 .then_with(|| {
312 b.quality_score
313 .partial_cmp(&a.quality_score)
314 .unwrap_or(std::cmp::Ordering::Equal)
315 })
316 });
317
318 candidates.into_iter().take(count).collect()
319}
320
321pub fn select_dual_stack_relays(peers: &[CachedPeer], count: usize) -> Vec<&CachedPeer> {
325 let mut filtered: Vec<&CachedPeer> = peers
326 .iter()
327 .filter(|p| p.capabilities.supports_dual_stack())
328 .collect();
329
330 if filtered.is_empty() {
331 return Vec::new();
332 }
333
334 filtered.sort_by(|a, b| {
335 let a_pref = helper_preference_score(a, true, false);
336 let b_pref = helper_preference_score(b, true, false);
337 b_pref
338 .cmp(&a_pref)
339 .then_with(|| {
340 b.capabilities
341 .direct_reachability_scope
342 .cmp(&a.capabilities.direct_reachability_scope)
343 })
344 .then_with(|| {
345 b.quality_score
346 .partial_cmp(&a.quality_score)
347 .unwrap_or(std::cmp::Ordering::Equal)
348 })
349 });
350
351 filtered.into_iter().take(count).collect()
352}
353
354#[allow(dead_code)]
356pub fn select_by_strategy(
357 peers: &[CachedPeer],
358 count: usize,
359 strategy: SelectionStrategy,
360) -> Vec<&CachedPeer> {
361 match strategy {
362 SelectionStrategy::BestFirst => {
363 let mut sorted: Vec<&CachedPeer> = peers.iter().collect();
364 sorted.sort_by(|a, b| {
365 b.quality_score
366 .partial_cmp(&a.quality_score)
367 .unwrap_or(std::cmp::Ordering::Equal)
368 });
369 sorted.into_iter().take(count).collect()
370 }
371 SelectionStrategy::EpsilonGreedy { epsilon } => {
372 select_epsilon_greedy(peers, count, epsilon)
373 }
374 SelectionStrategy::Random => {
375 let mut rng = rand::thread_rng();
376 let mut indices: Vec<usize> = (0..peers.len()).collect();
377 for i in (1..indices.len()).rev() {
379 let j = rng.gen_range(0..=i);
380 indices.swap(i, j);
381 }
382 indices.into_iter().take(count).map(|i| &peers[i]).collect()
383 }
384 }
385}
386
387#[cfg(test)]
388mod tests {
389 use super::*;
390 use crate::bootstrap_cache::entry::PeerSource;
391 use crate::nat_traversal_api::PeerId;
392
393 fn create_test_peers(count: usize) -> Vec<CachedPeer> {
394 (0..count)
395 .map(|i| {
396 let mut peer = CachedPeer::new(
397 PeerId([i as u8; 32]),
398 vec![format!("127.0.0.1:{}", 9000 + i).parse().unwrap()],
399 PeerSource::Seed,
400 );
401 peer.quality_score = i as f64 / count as f64;
403 peer
404 })
405 .collect()
406 }
407
408 #[test]
409 fn test_select_empty() {
410 let peers: Vec<CachedPeer> = vec![];
411 let selected = select_epsilon_greedy(&peers, 5, 0.1);
412 assert!(selected.is_empty());
413 }
414
415 #[test]
416 fn test_select_pure_exploitation() {
417 let peers = create_test_peers(10);
418 let selected = select_epsilon_greedy(&peers, 5, 0.0);
420
421 assert_eq!(selected.len(), 5);
422 for i in 0..4 {
424 assert!(selected[i].quality_score >= selected[i + 1].quality_score);
425 }
426 assert!((selected[0].quality_score - 0.9).abs() < 0.01);
428 }
429
430 #[test]
431 fn test_select_with_exploration() {
432 let peers = create_test_peers(20);
433 let mut has_variation = false;
436 let first_selection = select_epsilon_greedy(&peers, 10, 0.5);
437
438 for _ in 0..10 {
439 let selection = select_epsilon_greedy(&peers, 10, 0.5);
440 if selection.iter().map(|p| p.peer_id).collect::<Vec<_>>()
441 != first_selection
442 .iter()
443 .map(|p| p.peer_id)
444 .collect::<Vec<_>>()
445 {
446 has_variation = true;
447 break;
448 }
449 }
450 assert!(has_variation, "Expected variation with epsilon=0.5");
452 }
453
454 #[test]
455 fn test_select_more_than_available() {
456 let peers = create_test_peers(3);
457 let selected = select_epsilon_greedy(&peers, 10, 0.1);
458 assert_eq!(selected.len(), 3); }
460
461 #[test]
462 fn test_select_with_capabilities_prefers_broader_scope() {
463 let mut peers = create_test_peers(3);
464
465 peers[0].capabilities.direct_reachability_scope = Some(ReachabilityScope::LocalNetwork);
466 peers[1].capabilities.direct_reachability_scope = Some(ReachabilityScope::Global);
467 peers[1].capabilities.supports_relay = true;
468 peers[1].capabilities.supports_coordination = true;
469
470 let relays = select_with_capabilities(&peers, 3, true, false);
471 assert_eq!(relays.len(), 3);
472 assert_eq!(
473 relays[0].peer_id, peers[1].peer_id,
474 "global evidence should rank first"
475 );
476 assert_eq!(
477 relays[1].peer_id, peers[0].peer_id,
478 "local evidence should outrank unknown peers"
479 );
480 }
481
482 #[test]
483 fn test_best_first_strategy() {
484 let peers = create_test_peers(10);
485 let selected = select_by_strategy(&peers, 5, SelectionStrategy::BestFirst);
486
487 assert_eq!(selected.len(), 5);
488 for i in 0..4 {
490 assert!(selected[i].quality_score >= selected[i + 1].quality_score);
491 }
492 }
493
494 #[test]
495 fn test_random_strategy() {
496 let peers = create_test_peers(20);
497 let mut has_variation = false;
499 let first_selection = select_by_strategy(&peers, 10, SelectionStrategy::Random);
500
501 for _ in 0..10 {
502 let selection = select_by_strategy(&peers, 10, SelectionStrategy::Random);
503 if selection.iter().map(|p| p.peer_id).collect::<Vec<_>>()
504 != first_selection
505 .iter()
506 .map(|p| p.peer_id)
507 .collect::<Vec<_>>()
508 {
509 has_variation = true;
510 break;
511 }
512 }
513 assert!(has_variation, "Random selection should vary");
514 }
515
516 fn create_relay_peer_with_addresses(
517 id: u8,
518 quality: f64,
519 ipv4_addrs: Vec<&str>,
520 ipv6_addrs: Vec<&str>,
521 ) -> CachedPeer {
522 let mut peer = CachedPeer::new(PeerId([id; 32]), vec![], PeerSource::Seed);
523 peer.quality_score = quality;
524
525 for addr in ipv4_addrs {
526 peer.capabilities
527 .external_addresses
528 .push(addr.parse().unwrap());
529 }
530 for addr in ipv6_addrs {
531 peer.capabilities
532 .external_addresses
533 .push(addr.parse().unwrap());
534 }
535
536 peer.capabilities.direct_reachability_scope = peer
537 .capabilities
538 .external_addresses
539 .iter()
540 .filter_map(|addr| socket_addr_scope(*addr))
541 .max();
542 let globally_reachable = peer
543 .capabilities
544 .external_addresses
545 .iter()
546 .filter_map(|addr| socket_addr_scope(*addr))
547 .any(|scope| scope == ReachabilityScope::Global);
548 peer.capabilities.supports_relay = globally_reachable;
549 peer.capabilities.supports_coordination = globally_reachable;
550
551 peer
552 }
553
554 #[test]
555 fn test_select_relays_for_ipv4_target() {
556 let peers = vec![
557 create_relay_peer_with_addresses(
559 1,
560 0.9,
561 vec!["1.2.3.4:9000"],
562 vec!["[2001:db8::10]:9000"],
563 ),
564 create_relay_peer_with_addresses(2, 0.7, vec!["5.6.7.8:9001"], vec![]),
566 create_relay_peer_with_addresses(3, 0.95, vec![], vec!["[2001:db8::1]:9002"]),
568 ];
569
570 let selected = select_relays_for_target(&peers, 10, "8.8.8.8:443".parse().unwrap(), false);
571 assert_eq!(selected.len(), 2);
572
573 let ids: Vec<u8> = selected.iter().map(|p| p.peer_id.0[0]).collect();
575 assert!(ids.contains(&1)); assert!(ids.contains(&2)); assert!(!ids.contains(&3)); }
579
580 #[test]
581 fn test_select_relays_for_ipv6_target() {
582 let peers = vec![
583 create_relay_peer_with_addresses(
585 1,
586 0.9,
587 vec!["1.2.3.4:9000"],
588 vec!["[2001:db8::10]:9000"],
589 ),
590 create_relay_peer_with_addresses(2, 0.95, vec!["5.6.7.8:9001"], vec![]),
592 create_relay_peer_with_addresses(3, 0.7, vec![], vec!["[2001:db8::1]:9002"]),
594 ];
595
596 let selected = select_relays_for_target(
597 &peers,
598 10,
599 "[2001:4860:4860::8888]:443".parse().unwrap(),
600 false,
601 );
602 assert_eq!(selected.len(), 2);
603
604 let ids: Vec<u8> = selected.iter().map(|p| p.peer_id.0[0]).collect();
606 assert!(ids.contains(&1)); assert!(!ids.contains(&2)); assert!(ids.contains(&3)); }
610
611 #[test]
612 fn test_select_relays_prefer_dual_stack() {
613 let peers = vec![
614 create_relay_peer_with_addresses(
616 1,
617 0.5,
618 vec!["1.2.3.4:9000"],
619 vec!["[2001:db8::10]:9000"],
620 ),
621 create_relay_peer_with_addresses(2, 0.9, vec!["5.6.7.8:9001"], vec![]),
623 ];
624
625 let selected = select_relays_for_target(&peers, 10, "8.8.8.8:443".parse().unwrap(), false);
627 assert_eq!(selected[0].peer_id.0[0], 2); let selected = select_relays_for_target(&peers, 10, "8.8.8.8:443".parse().unwrap(), true);
631 assert_eq!(selected[0].peer_id.0[0], 1); }
633
634 #[test]
635 fn test_select_dual_stack_relays() {
636 let peers = vec![
637 create_relay_peer_with_addresses(
639 1,
640 0.9,
641 vec!["1.2.3.4:9000"],
642 vec!["[2001:db8::10]:9000"],
643 ),
644 create_relay_peer_with_addresses(2, 0.8, vec!["5.6.7.8:9001"], vec![]),
646 create_relay_peer_with_addresses(3, 0.7, vec![], vec!["[2001:db8::1]:9002"]),
648 create_relay_peer_with_addresses(
650 4,
651 0.6,
652 vec!["9.9.9.9:9003"],
653 vec!["[2001:db8::2]:9003"],
654 ),
655 ];
656
657 let selected = select_dual_stack_relays(&peers, 10);
658 assert_eq!(selected.len(), 2);
659
660 for peer in &selected {
662 assert!(peer.capabilities.supports_dual_stack());
663 }
664
665 assert!(selected[0].quality_score >= selected[1].quality_score);
667 }
668
669 #[test]
670 fn test_select_relays_excludes_non_relays() {
671 let mut peers = vec![create_relay_peer_with_addresses(
672 1,
673 0.9,
674 vec!["1.2.3.4:9000"],
675 vec![],
676 )];
677
678 let mut non_relay = CachedPeer::new(PeerId([2; 32]), vec![], PeerSource::Seed);
680 non_relay.quality_score = 0.99;
681 non_relay.capabilities.supports_relay = false;
682 non_relay
683 .capabilities
684 .external_addresses
685 .push("5.6.7.8:9001".parse().unwrap());
686 non_relay.capabilities.direct_reachability_scope = Some(ReachabilityScope::LocalNetwork);
687 peers.push(non_relay);
688
689 let selected = select_relays_for_target(&peers, 10, "8.8.8.8:443".parse().unwrap(), false);
690 assert_eq!(selected.len(), 2);
691 assert_eq!(selected[0].peer_id.0[0], 1);
693 }
694
695 #[test]
696 fn test_select_relays_empty_when_no_match() {
697 let peers = vec![
698 create_relay_peer_with_addresses(1, 0.9, vec![], vec!["[2001:db8::1]:9000"]),
700 ];
701
702 let selected = select_relays_for_target(&peers, 10, "8.8.8.8:443".parse().unwrap(), false);
704 assert!(selected.is_empty());
705 }
706}