1use crate::error::{RoutingError, Result};
8use serde::{Deserialize, Serialize};
9use std::collections::HashMap;
10use std::sync::{Arc, OnceLock};
11use std::time::{Duration, Instant};
12use std::fmt;
13use tokio::sync::RwLock;
14use tracing::{debug, warn};
15
16fn ucb1_c() -> f64 {
23 static C: OnceLock<f64> = OnceLock::new();
24 *C.get_or_init(|| {
25 std::env::var("MURMURATION_UCB1_C")
26 .ok()
27 .and_then(|v| v.parse().ok())
28 .unwrap_or(2.0)
29 })
30}
31const UCB1_MIN_SAMPLES: u64 = 5;
33
34const Q_ALPHA: f64 = 0.15;
36const Q_INIT: f64 = 1.0;
39
40#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
42pub struct MurmurationAddress {
43 pub node_id: String,
44}
45
46impl MurmurationAddress {
47 pub fn from_string(addr: &str) -> Result<Self> {
49 if let Some(stripped) = addr.strip_prefix("mur://") {
50 Ok(Self {
51 node_id: stripped.to_string(),
52 })
53 } else {
54 Err(RoutingError::Protocol(format!(
55 "Invalid Murmuration address format: {}",
56 addr
57 )))
58 }
59 }
60}
61
62impl fmt::Display for MurmurationAddress {
63 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
64 write!(f, "mur://{}", self.node_id)
65 }
66}
67
68#[derive(Debug, Clone, Serialize, Deserialize)]
70pub struct MeshMessage {
71 pub from: String,
72 pub to: Option<String>, pub data: Vec<u8>,
74 pub message_id: String,
75 pub ttl: u8,
76 pub path: Vec<String>, }
78
79impl MeshMessage {
80 pub fn new(from: String, to: Option<String>, data: Vec<u8>) -> Self {
82 Self {
83 from,
84 to,
85 data,
86 message_id: uuid::Uuid::new_v4().to_string(),
87 ttl: 10, path: Vec::new(),
89 }
90 }
91
92 }
95
96pub trait RouterStore: Send + Sync {
102 fn load(&self) -> Option<Vec<u8>>;
104 fn save(&self, bytes: &[u8]);
106}
107
108pub struct Router {
114 our_node_id: String,
115 seen_messages: Arc<RwLock<HashMap<String, Instant>>>, message_cache: Arc<RwLock<HashMap<String, MeshMessage>>>, route_history: Arc<RwLock<HashMap<String, RouteStats>>>, ucb_state: Arc<RwLock<UcbState>>, q_state: Arc<RwLock<QRoutingState>>, store: Option<Arc<dyn RouterStore>>,
122}
123
124#[derive(Debug, Default)]
133struct QRoutingState {
134 q: HashMap<(String, String), f64>,
135}
136
137impl QRoutingState {
138 fn get(&self, dest: &str, peer: &str) -> f64 {
139 *self
140 .q
141 .get(&(dest.to_string(), peer.to_string()))
142 .unwrap_or(&Q_INIT)
143 }
144
145 fn best_over(&self, dest: &str, neighbours: &[String]) -> f64 {
148 neighbours
149 .iter()
150 .map(|p| self.get(dest, p))
151 .fold(0.0_f64, f64::max)
152 }
153
154 fn update(&mut self, dest: &str, peer: &str, target: f64) {
155 let cur = self.get(dest, peer);
156 self.q.insert(
157 (dest.to_string(), peer.to_string()),
158 (1.0 - Q_ALPHA) * cur + Q_ALPHA * target,
159 );
160 }
161}
162
163#[derive(Debug, Clone)]
165pub struct RouteStats {
166 success_count: u32,
167 failure_count: u32,
168 total_latency: Duration,
169 sample_count: u32,
170 last_updated: Instant,
171}
172
173#[derive(Debug, Clone, Default, Serialize, Deserialize)]
175struct UcbPeerStats {
176 selections: u64,
178 avg_reward: f64,
180}
181
182#[derive(Debug, Default, Serialize, Deserialize)]
198struct UcbState {
199 total_selections: u64,
201 peers: HashMap<String, UcbPeerStats>,
202 #[serde(default)]
204 by_dest: HashMap<String, DestBandit>,
205}
206
207#[derive(Debug, Clone, Default, Serialize, Deserialize)]
209struct DestBandit {
210 total_selections: u64,
211 peers: HashMap<String, UcbPeerStats>,
212}
213
214impl DestBandit {
215 fn ucb1_score(&self, peer_id: &str) -> f64 {
216 match self.peers.get(peer_id) {
217 None => f64::INFINITY,
218 Some(s) if s.selections == 0 => f64::INFINITY,
219 Some(s) => {
220 let exploration = if self.total_selections > 0 {
221 (ucb1_c() * (self.total_selections as f64).ln() / s.selections as f64).sqrt()
222 } else {
223 0.0
224 };
225 s.avg_reward + exploration
226 }
227 }
228 }
229
230 fn record_reward(&mut self, peer_id: &str, reward: f64) {
231 self.total_selections += 1;
232 let s = self.peers.entry(peer_id.to_string()).or_default();
233 s.selections += 1;
234 s.avg_reward += (reward - s.avg_reward) / s.selections as f64;
235 }
236
237 fn selections(&self, peer_id: &str) -> u64 {
238 self.peers.get(peer_id).map_or(0, |s| s.selections)
239 }
240}
241
242impl UcbState {
243 fn ucb1_score(&self, peer_id: &str) -> f64 {
246 match self.peers.get(peer_id) {
247 None => f64::INFINITY,
248 Some(s) if s.selections == 0 => f64::INFINITY,
249 Some(s) => {
250 let exploration = if self.total_selections > 0 {
251 (ucb1_c() * (self.total_selections as f64).ln() / s.selections as f64).sqrt()
252 } else {
253 0.0
254 };
255 s.avg_reward + exploration
256 }
257 }
258 }
259
260 fn record_reward(&mut self, peer_id: &str, reward: f64) {
262 self.total_selections += 1;
263 let s = self.peers.entry(peer_id.to_string()).or_default();
264 s.selections += 1;
265 s.avg_reward += (reward - s.avg_reward) / s.selections as f64;
267 }
268
269 fn selections(&self, peer_id: &str) -> u64 {
271 self.peers.get(peer_id).map_or(0, |s| s.selections)
272 }
273}
274
275impl Router {
276 pub fn new(our_node_id: String) -> Self {
278 Self {
279 our_node_id,
280 seen_messages: Arc::new(RwLock::new(HashMap::new())),
281 message_cache: Arc::new(RwLock::new(HashMap::new())),
282 route_history: Arc::new(RwLock::new(HashMap::new())),
283 ucb_state: Arc::new(RwLock::new(UcbState::default())),
284 q_state: Arc::new(RwLock::new(QRoutingState::default())),
285 store: None,
286 }
287 }
288
289 pub fn with_store(our_node_id: String, store: Arc<dyn RouterStore>) -> Self {
292 let initial_state = store
293 .load()
294 .and_then(|bytes| serde_json::from_slice::<UcbState>(&bytes).ok())
295 .unwrap_or_default();
296 Self {
297 our_node_id,
298 seen_messages: Arc::new(RwLock::new(HashMap::new())),
299 message_cache: Arc::new(RwLock::new(HashMap::new())),
300 route_history: Arc::new(RwLock::new(HashMap::new())),
301 ucb_state: Arc::new(RwLock::new(initial_state)),
302 q_state: Arc::new(RwLock::new(QRoutingState::default())),
303 store: Some(store),
304 }
305 }
306
307 async fn persist_ucb_state(&self) {
309 if let Some(store) = &self.store {
310 let state = self.ucb_state.read().await;
311 match serde_json::to_vec(&*state) {
312 Ok(bytes) => store.save(&bytes),
313 Err(e) => warn!("Failed to serialize UCB1 state: {}", e),
314 }
315 }
316 }
317
318 pub async fn should_process(&self, message: &MeshMessage) -> bool {
320 if message.ttl == 0 {
322 debug!("Message {} dropped: TTL expired", message.message_id);
323 return false;
324 }
325
326 let seen = self.seen_messages.read().await;
328 if let Some(timestamp) = seen.get(&message.message_id) {
329 if timestamp.elapsed() < Duration::from_secs(60) {
330 debug!("Message {} dropped: already seen", message.message_id);
331 return false;
332 }
333 }
334 drop(seen);
335
336 if message.path.contains(&self.our_node_id) {
338 debug!("Message {} dropped: loop detected", message.message_id);
339 return false;
340 }
341
342 true
343 }
344
345 pub async fn mark_seen(&self, message_id: &str) {
347 let mut seen = self.seen_messages.write().await;
348 seen.insert(message_id.to_string(), Instant::now());
349
350 seen.retain(|_, timestamp| timestamp.elapsed() < Duration::from_secs(300));
352 }
353
354 pub fn is_for_us(&self, message: &MeshMessage) -> bool {
356 match &message.to {
357 None => true, Some(to) => to == &self.our_node_id,
359 }
360 }
361
362 pub fn prepare_for_forwarding(&self, message: &MeshMessage) -> MeshMessage {
364 let mut forward_msg = message.clone();
365 forward_msg.ttl = forward_msg.ttl.saturating_sub(1);
366 forward_msg.path.push(self.our_node_id.clone());
367 forward_msg
368 }
369
370 pub fn calculate_peer_score(
373 peer_metrics: &crate::peer::PeerMetrics,
374 route_stats: Option<&RouteStats>,
375 ) -> f64 {
376 let latency_score = peer_metrics
378 .latency
379 .map(|lat| {
380 let lat_secs = lat.as_secs_f64();
381 (1.0 - (lat_secs.min(1.0))).max(0.0)
382 })
383 .unwrap_or(0.5); let uptime_score = (peer_metrics.uptime.as_secs_f64() / 3600.0).min(1.0);
387
388 let reliability = peer_metrics.reliability_score() as f64;
390
391 let route_success_rate = if let Some(stats) = route_stats {
393 let total = stats.success_count + stats.failure_count;
394 if total > 0 {
395 stats.success_count as f64 / total as f64
396 } else {
397 0.5
398 }
399 } else {
400 0.5 };
402
403 let base_score = 0.3 * latency_score
405 + 0.15 * uptime_score
406 + 0.3 * reliability
407 + 0.25 * route_success_rate;
408
409 if let Some(stats) = route_stats {
411 if stats.sample_count > 0 {
412 let avg_latency = if stats.sample_count > 0 {
413 stats.total_latency.as_secs_f64() / stats.sample_count as f64
414 } else {
415 0.0
416 };
417 let historical_score = (1.0 - (avg_latency.min(1.0))).max(0.0);
418
419 const ALPHA: f64 = 0.7;
421 const BETA: f64 = 0.3;
422 return ALPHA * historical_score + BETA * base_score;
423 }
424 }
425
426 base_score
427 }
428
429 pub fn get_forward_peers(&self, message: &MeshMessage, all_peers: &[String]) -> Vec<String> {
431 all_peers
432 .iter()
433 .filter(|peer_id| {
434 **peer_id != message.from &&
436 !message.path.contains(peer_id)
438 })
439 .cloned()
440 .collect()
441 }
442
443 pub async fn get_best_forward_peers(
452 &self,
453 message: &MeshMessage,
454 peer_infos: &[crate::peer::PeerInfo],
455 max_peers: usize,
456 ) -> Vec<String> {
457 let route_history = self.route_history.read().await;
458 let ucb = self.ucb_state.read().await;
459
460 let mut scored_peers: Vec<(String, f64)> = peer_infos
461 .iter()
462 .filter(|peer| {
463 peer.node_id != message.from
464 && !message.path.contains(&peer.node_id)
465 && peer.is_connected()
466 })
467 .map(|peer| {
468 let n_i = ucb.selections(&peer.node_id);
469 let score = if n_i < UCB1_MIN_SAMPLES {
470 let heuristic =
473 Self::calculate_peer_score(&peer.metrics, route_history.get(&peer.node_id));
474 let bonus = if n_i == 0 { 1.0 } else { 0.5 };
475 heuristic + bonus
476 } else {
477 ucb.ucb1_score(&peer.node_id)
478 };
479 (peer.node_id.clone(), score)
480 })
481 .collect();
482
483 scored_peers.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
484 scored_peers
485 .into_iter()
486 .take(max_peers)
487 .map(|(peer_id, _)| peer_id)
488 .collect()
489 }
490
491 pub async fn get_best_forward_peers_toward(
503 &self,
504 message: &MeshMessage,
505 peer_infos: &[crate::peer::PeerInfo],
506 max_peers: usize,
507 dest: &str,
508 ) -> Vec<String> {
509 let route_history = self.route_history.read().await;
510 let ucb = self.ucb_state.read().await;
511 let empty = DestBandit::default();
512 let bandit = ucb.by_dest.get(dest).unwrap_or(&empty);
513
514 let mut scored_peers: Vec<(String, f64)> = peer_infos
515 .iter()
516 .filter(|peer| {
517 peer.node_id != message.from
518 && !message.path.contains(&peer.node_id)
519 && peer.is_connected()
520 })
521 .map(|peer| {
522 let n_i = bandit.selections(&peer.node_id);
523 let score = if n_i < UCB1_MIN_SAMPLES {
524 let heuristic =
525 Self::calculate_peer_score(&peer.metrics, route_history.get(&peer.node_id));
526 let bonus = if n_i == 0 { 1.0 } else { 0.5 };
527 heuristic + bonus
528 } else {
529 bandit.ucb1_score(&peer.node_id)
530 };
531 (peer.node_id.clone(), score)
532 })
533 .collect();
534
535 scored_peers.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
536 scored_peers
537 .into_iter()
538 .take(max_peers)
539 .map(|(peer_id, _)| peer_id)
540 .collect()
541 }
542
543 pub async fn record_route_outcome_toward(
549 &self,
550 dest: &str,
551 peer_id: &str,
552 success: Option<Duration>,
553 ) {
554 let reward = match success {
555 Some(latency) => (1.0 - 2.0 * latency.as_secs_f64()).clamp(0.5, 1.0),
556 None => 0.0,
557 };
558 {
559 let mut state = self.ucb_state.write().await;
560 state
561 .by_dest
562 .entry(dest.to_string())
563 .or_default()
564 .record_reward(peer_id, reward);
565 }
566 self.persist_ucb_state().await;
567 }
568
569 pub async fn q_select_toward(
585 &self,
586 message: &MeshMessage,
587 peer_infos: &[crate::peer::PeerInfo],
588 max_peers: usize,
589 dest: &str,
590 ) -> Vec<String> {
591 let q = self.q_state.read().await;
592 let mut scored: Vec<(String, f64)> = peer_infos
593 .iter()
594 .filter(|p| {
595 p.node_id != message.from
596 && !message.path.contains(&p.node_id)
597 && p.is_connected()
598 })
599 .map(|p| (p.node_id.clone(), q.get(dest, &p.node_id)))
600 .collect();
601 scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
602 scored
603 .into_iter()
604 .take(max_peers)
605 .map(|(id, _)| id)
606 .collect()
607 }
608
609 pub async fn q_advertised_value(&self, dest: &str, neighbours: &[String]) -> f64 {
614 self.q_state.read().await.best_over(dest, neighbours)
615 }
616
617 pub async fn q_record(
624 &self,
625 dest: &str,
626 peer: &str,
627 delivered: bool,
628 downstream_value: f64,
629 ) {
630 let target = if delivered { downstream_value } else { 0.0 };
631 self.q_state.write().await.update(dest, peer, target);
632 }
633
634 pub async fn record_route_success(&self, peer_id: &str, latency: Duration) {
638 let mut history = self.route_history.write().await;
639 let stats = history
640 .entry(peer_id.to_string())
641 .or_insert_with(|| RouteStats {
642 success_count: 0,
643 failure_count: 0,
644 total_latency: Duration::ZERO,
645 sample_count: 0,
646 last_updated: Instant::now(),
647 });
648
649 stats.success_count += 1;
650 stats.total_latency += latency;
651 stats.sample_count += 1;
652 stats.last_updated = Instant::now();
653 drop(history);
654
655 let reward = (1.0 - 2.0 * latency.as_secs_f64()).clamp(0.5, 1.0);
657 self.ucb_state.write().await.record_reward(peer_id, reward);
658 self.persist_ucb_state().await;
659 }
660
661 pub async fn record_route_failure(&self, peer_id: &str) {
664 let mut history = self.route_history.write().await;
665 let stats = history
666 .entry(peer_id.to_string())
667 .or_insert_with(|| RouteStats {
668 success_count: 0,
669 failure_count: 0,
670 total_latency: Duration::ZERO,
671 sample_count: 0,
672 last_updated: Instant::now(),
673 });
674
675 stats.failure_count += 1;
676 stats.last_updated = Instant::now();
677 drop(history);
678
679 self.ucb_state.write().await.record_reward(peer_id, 0.0);
681 self.persist_ucb_state().await;
682 }
683
684 pub async fn cleanup_cache(&self) {
686 let mut cache = self.message_cache.write().await;
687 cache.retain(|_, _msg| {
688 true });
691 }
692}
693
694impl Clone for Router {
695 fn clone(&self) -> Self {
696 Self {
697 our_node_id: self.our_node_id.clone(),
698 seen_messages: self.seen_messages.clone(),
699 message_cache: self.message_cache.clone(),
700 route_history: self.route_history.clone(),
701 ucb_state: self.ucb_state.clone(),
702 q_state: self.q_state.clone(),
703 store: self.store.clone(),
704 }
705 }
706}
707
708#[cfg(test)]
709mod tests {
710 use super::*;
711
712 #[test]
713 fn test_murmuration_address_parse() {
714 let addr = MurmurationAddress::from_string("mur://node123").unwrap();
715 assert_eq!(addr.node_id, "node123");
716
717 let invalid = MurmurationAddress::from_string("invalid");
718 assert!(invalid.is_err());
719 }
720
721 #[test]
722 fn test_murmuration_address_to_string() {
723 let addr = MurmurationAddress {
724 node_id: "node123".to_string(),
725 };
726 assert_eq!(addr.to_string(), "mur://node123");
727 }
728
729 #[tokio::test]
730 async fn test_router_should_process() {
731 let router = Router::new("our-node".to_string());
732 let message = MeshMessage::new("peer1".to_string(), None, b"test".to_vec());
733
734 assert!(router.should_process(&message).await);
736
737 router.mark_seen(&message.message_id).await;
739
740 assert!(!router.should_process(&message).await);
742 }
743
744 #[tokio::test]
745 async fn test_router_is_for_us() {
746 let router = Router::new("our-node".to_string());
747
748 let broadcast = MeshMessage::new("peer1".to_string(), None, b"test".to_vec());
750 assert!(router.is_for_us(&broadcast));
751
752 let directed = MeshMessage::new(
754 "peer1".to_string(),
755 Some("our-node".to_string()),
756 b"test".to_vec(),
757 );
758 assert!(router.is_for_us(&directed));
759
760 let other = MeshMessage::new(
762 "peer1".to_string(),
763 Some("other-node".to_string()),
764 b"test".to_vec(),
765 );
766 assert!(!router.is_for_us(&other));
767 }
768
769 #[tokio::test]
770 async fn test_router_prepare_for_forwarding() {
771 let router = Router::new("our-node".to_string());
772 let message = MeshMessage::new("peer1".to_string(), None, b"test".to_vec());
773 let original_ttl = message.ttl;
774
775 let forwarded = router.prepare_for_forwarding(&message);
776
777 assert_eq!(forwarded.ttl, original_ttl - 1);
778 assert!(forwarded.path.contains(&"our-node".to_string()));
779 }
780
781 #[tokio::test]
782 async fn test_router_get_forward_peers() {
783 let router = Router::new("our-node".to_string());
784 let message = MeshMessage::new("peer1".to_string(), None, b"test".to_vec());
785 let all_peers = vec![
786 "peer1".to_string(),
787 "peer2".to_string(),
788 "peer3".to_string(),
789 ];
790
791 let forward_peers = router.get_forward_peers(&message, &all_peers);
792
793 assert!(!forward_peers.contains(&"peer1".to_string()));
795 assert!(forward_peers.contains(&"peer2".to_string()));
796 assert!(forward_peers.contains(&"peer3".to_string()));
797 }
798
799 #[tokio::test]
800 async fn test_router_loop_detection() {
801 let router = Router::new("our-node".to_string());
802 let mut message = MeshMessage::new("peer1".to_string(), None, b"test".to_vec());
803 message.path.push("our-node".to_string());
804
805 assert!(!router.should_process(&message).await);
807 }
808
809 use crate::peer::{ConnectionState, PeerInfo};
812 use std::net::SocketAddr;
813
814 fn connected_peer(id: &str) -> PeerInfo {
815 let addr: SocketAddr = "127.0.0.1:9000".parse().unwrap();
816 let mut p = PeerInfo::new(id.to_string(), addr);
817 p.state = ConnectionState::Connected;
818 p
819 }
820
821 #[tokio::test]
822 async fn test_q_optimistic_init() {
823 let router = Router::new("u".to_string());
824 let v = router.q_advertised_value("dst", &["a".to_string()]).await;
827 assert_eq!(v, Q_INIT);
828 assert_eq!(router.q_advertised_value("dst", &[]).await, 0.0);
830 }
831
832 #[tokio::test]
833 async fn test_q_record_bootstraps_downstream_value() {
834 let router = Router::new("u".to_string());
835 router.q_record("dst", "a", true, 0.8).await;
837 let q_a = router.q_advertised_value("dst", &["a".to_string()]).await;
839 assert!((q_a - 0.97).abs() < 1e-9, "got {q_a}");
840
841 for _ in 0..20 {
843 router.q_record("dst", "b", false, 0.0).await;
844 }
845 let q_b = router
846 .q_state
847 .read()
848 .await
849 .get("dst", "b");
850 assert!(q_b < 0.1, "failures should drive Q→0, got {q_b}");
851 }
852
853 #[tokio::test]
854 async fn test_q_select_prefers_higher_value() {
855 let router = Router::new("u".to_string());
856 for _ in 0..30 {
858 router.q_record("dst", "good", true, 1.0).await;
859 router.q_record("dst", "bad", false, 0.0).await;
860 }
861 let msg = MeshMessage::new("src".to_string(), Some("dst".to_string()), vec![]);
862 let peers = vec![connected_peer("good"), connected_peer("bad")];
863 let picked = router.q_select_toward(&msg, &peers, 1, "dst").await;
864 assert_eq!(picked, vec!["good".to_string()]);
865 }
866
867 #[tokio::test]
868 async fn test_q_select_excludes_sender_and_path() {
869 let router = Router::new("u".to_string());
870 let mut msg = MeshMessage::new("sender".to_string(), Some("dst".to_string()), vec![]);
871 msg.path.push("visited".to_string());
872 let peers = vec![
873 connected_peer("sender"),
874 connected_peer("visited"),
875 connected_peer("fresh"),
876 ];
877 let picked = router.q_select_toward(&msg, &peers, 3, "dst").await;
878 assert_eq!(picked, vec!["fresh".to_string()]);
879 }
880}