use crate::error::{RoutingError, Result};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::{Arc, OnceLock};
use std::time::{Duration, Instant};
use std::fmt;
use tokio::sync::RwLock;
use tracing::{debug, warn};
fn ucb1_c() -> f64 {
static C: OnceLock<f64> = OnceLock::new();
*C.get_or_init(|| {
std::env::var("MURMURATION_UCB1_C")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(2.0)
})
}
const UCB1_MIN_SAMPLES: u64 = 5;
const Q_ALPHA: f64 = 0.15;
const Q_INIT: f64 = 1.0;
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub struct MurmurationAddress {
pub node_id: String,
}
impl MurmurationAddress {
pub fn from_string(addr: &str) -> Result<Self> {
if let Some(stripped) = addr.strip_prefix("mur://") {
Ok(Self {
node_id: stripped.to_string(),
})
} else {
Err(RoutingError::Protocol(format!(
"Invalid Murmuration address format: {}",
addr
)))
}
}
}
impl fmt::Display for MurmurationAddress {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "mur://{}", self.node_id)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MeshMessage {
pub from: String,
pub to: Option<String>, pub data: Vec<u8>,
pub message_id: String,
pub ttl: u8,
pub path: Vec<String>, }
impl MeshMessage {
pub fn new(from: String, to: Option<String>, data: Vec<u8>) -> Self {
Self {
from,
to,
data,
message_id: uuid::Uuid::new_v4().to_string(),
ttl: 10, path: Vec::new(),
}
}
}
pub trait RouterStore: Send + Sync {
fn load(&self) -> Option<Vec<u8>>;
fn save(&self, bytes: &[u8]);
}
pub struct Router {
our_node_id: String,
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>>,
}
#[derive(Debug, Default)]
struct QRoutingState {
q: HashMap<(String, String), f64>,
}
impl QRoutingState {
fn get(&self, dest: &str, peer: &str) -> f64 {
*self
.q
.get(&(dest.to_string(), peer.to_string()))
.unwrap_or(&Q_INIT)
}
fn best_over(&self, dest: &str, neighbours: &[String]) -> f64 {
neighbours
.iter()
.map(|p| self.get(dest, p))
.fold(0.0_f64, f64::max)
}
fn update(&mut self, dest: &str, peer: &str, target: f64) {
let cur = self.get(dest, peer);
self.q.insert(
(dest.to_string(), peer.to_string()),
(1.0 - Q_ALPHA) * cur + Q_ALPHA * target,
);
}
}
#[derive(Debug, Clone)]
pub struct RouteStats {
success_count: u32,
failure_count: u32,
total_latency: Duration,
sample_count: u32,
last_updated: Instant,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
struct UcbPeerStats {
selections: u64,
avg_reward: f64,
}
#[derive(Debug, Default, Serialize, Deserialize)]
struct UcbState {
total_selections: u64,
peers: HashMap<String, UcbPeerStats>,
#[serde(default)]
by_dest: HashMap<String, DestBandit>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
struct DestBandit {
total_selections: u64,
peers: HashMap<String, UcbPeerStats>,
}
impl DestBandit {
fn ucb1_score(&self, peer_id: &str) -> f64 {
match self.peers.get(peer_id) {
None => f64::INFINITY,
Some(s) if s.selections == 0 => f64::INFINITY,
Some(s) => {
let exploration = if self.total_selections > 0 {
(ucb1_c() * (self.total_selections as f64).ln() / s.selections as f64).sqrt()
} else {
0.0
};
s.avg_reward + exploration
}
}
}
fn record_reward(&mut self, peer_id: &str, reward: f64) {
self.total_selections += 1;
let s = self.peers.entry(peer_id.to_string()).or_default();
s.selections += 1;
s.avg_reward += (reward - s.avg_reward) / s.selections as f64;
}
fn selections(&self, peer_id: &str) -> u64 {
self.peers.get(peer_id).map_or(0, |s| s.selections)
}
}
impl UcbState {
fn ucb1_score(&self, peer_id: &str) -> f64 {
match self.peers.get(peer_id) {
None => f64::INFINITY,
Some(s) if s.selections == 0 => f64::INFINITY,
Some(s) => {
let exploration = if self.total_selections > 0 {
(ucb1_c() * (self.total_selections as f64).ln() / s.selections as f64).sqrt()
} else {
0.0
};
s.avg_reward + exploration
}
}
}
fn record_reward(&mut self, peer_id: &str, reward: f64) {
self.total_selections += 1;
let s = self.peers.entry(peer_id.to_string()).or_default();
s.selections += 1;
s.avg_reward += (reward - s.avg_reward) / s.selections as f64;
}
fn selections(&self, peer_id: &str) -> u64 {
self.peers.get(peer_id).map_or(0, |s| s.selections)
}
}
impl Router {
pub fn new(our_node_id: String) -> Self {
Self {
our_node_id,
seen_messages: Arc::new(RwLock::new(HashMap::new())),
message_cache: Arc::new(RwLock::new(HashMap::new())),
route_history: Arc::new(RwLock::new(HashMap::new())),
ucb_state: Arc::new(RwLock::new(UcbState::default())),
q_state: Arc::new(RwLock::new(QRoutingState::default())),
store: None,
}
}
pub fn with_store(our_node_id: String, store: Arc<dyn RouterStore>) -> Self {
let initial_state = store
.load()
.and_then(|bytes| serde_json::from_slice::<UcbState>(&bytes).ok())
.unwrap_or_default();
Self {
our_node_id,
seen_messages: Arc::new(RwLock::new(HashMap::new())),
message_cache: Arc::new(RwLock::new(HashMap::new())),
route_history: Arc::new(RwLock::new(HashMap::new())),
ucb_state: Arc::new(RwLock::new(initial_state)),
q_state: Arc::new(RwLock::new(QRoutingState::default())),
store: Some(store),
}
}
async fn persist_ucb_state(&self) {
if let Some(store) = &self.store {
let state = self.ucb_state.read().await;
match serde_json::to_vec(&*state) {
Ok(bytes) => store.save(&bytes),
Err(e) => warn!("Failed to serialize UCB1 state: {}", e),
}
}
}
pub async fn should_process(&self, message: &MeshMessage) -> bool {
if message.ttl == 0 {
debug!("Message {} dropped: TTL expired", message.message_id);
return false;
}
let seen = self.seen_messages.read().await;
if let Some(timestamp) = seen.get(&message.message_id) {
if timestamp.elapsed() < Duration::from_secs(60) {
debug!("Message {} dropped: already seen", message.message_id);
return false;
}
}
drop(seen);
if message.path.contains(&self.our_node_id) {
debug!("Message {} dropped: loop detected", message.message_id);
return false;
}
true
}
pub async fn mark_seen(&self, message_id: &str) {
let mut seen = self.seen_messages.write().await;
seen.insert(message_id.to_string(), Instant::now());
seen.retain(|_, timestamp| timestamp.elapsed() < Duration::from_secs(300));
}
pub fn is_for_us(&self, message: &MeshMessage) -> bool {
match &message.to {
None => true, Some(to) => to == &self.our_node_id,
}
}
pub fn prepare_for_forwarding(&self, message: &MeshMessage) -> MeshMessage {
let mut forward_msg = message.clone();
forward_msg.ttl = forward_msg.ttl.saturating_sub(1);
forward_msg.path.push(self.our_node_id.clone());
forward_msg
}
pub fn calculate_peer_score(
peer_metrics: &crate::peer::PeerMetrics,
route_stats: Option<&RouteStats>,
) -> f64 {
let latency_score = peer_metrics
.latency
.map(|lat| {
let lat_secs = lat.as_secs_f64();
(1.0 - (lat_secs.min(1.0))).max(0.0)
})
.unwrap_or(0.5);
let uptime_score = (peer_metrics.uptime.as_secs_f64() / 3600.0).min(1.0);
let reliability = peer_metrics.reliability_score() as f64;
let route_success_rate = if let Some(stats) = route_stats {
let total = stats.success_count + stats.failure_count;
if total > 0 {
stats.success_count as f64 / total as f64
} else {
0.5
}
} else {
0.5 };
let base_score = 0.3 * latency_score
+ 0.15 * uptime_score
+ 0.3 * reliability
+ 0.25 * route_success_rate;
if let Some(stats) = route_stats {
if stats.sample_count > 0 {
let avg_latency = if stats.sample_count > 0 {
stats.total_latency.as_secs_f64() / stats.sample_count as f64
} else {
0.0
};
let historical_score = (1.0 - (avg_latency.min(1.0))).max(0.0);
const ALPHA: f64 = 0.7;
const BETA: f64 = 0.3;
return ALPHA * historical_score + BETA * base_score;
}
}
base_score
}
pub fn get_forward_peers(&self, message: &MeshMessage, all_peers: &[String]) -> Vec<String> {
all_peers
.iter()
.filter(|peer_id| {
**peer_id != message.from &&
!message.path.contains(peer_id)
})
.cloned()
.collect()
}
pub async fn get_best_forward_peers(
&self,
message: &MeshMessage,
peer_infos: &[crate::peer::PeerInfo],
max_peers: usize,
) -> Vec<String> {
let route_history = self.route_history.read().await;
let ucb = self.ucb_state.read().await;
let mut scored_peers: Vec<(String, f64)> = peer_infos
.iter()
.filter(|peer| {
peer.node_id != message.from
&& !message.path.contains(&peer.node_id)
&& peer.is_connected()
})
.map(|peer| {
let n_i = ucb.selections(&peer.node_id);
let score = if n_i < UCB1_MIN_SAMPLES {
let heuristic =
Self::calculate_peer_score(&peer.metrics, route_history.get(&peer.node_id));
let bonus = if n_i == 0 { 1.0 } else { 0.5 };
heuristic + bonus
} else {
ucb.ucb1_score(&peer.node_id)
};
(peer.node_id.clone(), score)
})
.collect();
scored_peers.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
scored_peers
.into_iter()
.take(max_peers)
.map(|(peer_id, _)| peer_id)
.collect()
}
pub async fn get_best_forward_peers_toward(
&self,
message: &MeshMessage,
peer_infos: &[crate::peer::PeerInfo],
max_peers: usize,
dest: &str,
) -> Vec<String> {
let route_history = self.route_history.read().await;
let ucb = self.ucb_state.read().await;
let empty = DestBandit::default();
let bandit = ucb.by_dest.get(dest).unwrap_or(&empty);
let mut scored_peers: Vec<(String, f64)> = peer_infos
.iter()
.filter(|peer| {
peer.node_id != message.from
&& !message.path.contains(&peer.node_id)
&& peer.is_connected()
})
.map(|peer| {
let n_i = bandit.selections(&peer.node_id);
let score = if n_i < UCB1_MIN_SAMPLES {
let heuristic =
Self::calculate_peer_score(&peer.metrics, route_history.get(&peer.node_id));
let bonus = if n_i == 0 { 1.0 } else { 0.5 };
heuristic + bonus
} else {
bandit.ucb1_score(&peer.node_id)
};
(peer.node_id.clone(), score)
})
.collect();
scored_peers.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
scored_peers
.into_iter()
.take(max_peers)
.map(|(peer_id, _)| peer_id)
.collect()
}
pub async fn record_route_outcome_toward(
&self,
dest: &str,
peer_id: &str,
success: Option<Duration>,
) {
let reward = match success {
Some(latency) => (1.0 - 2.0 * latency.as_secs_f64()).clamp(0.5, 1.0),
None => 0.0,
};
{
let mut state = self.ucb_state.write().await;
state
.by_dest
.entry(dest.to_string())
.or_default()
.record_reward(peer_id, reward);
}
self.persist_ucb_state().await;
}
pub async fn q_select_toward(
&self,
message: &MeshMessage,
peer_infos: &[crate::peer::PeerInfo],
max_peers: usize,
dest: &str,
) -> Vec<String> {
let q = self.q_state.read().await;
let mut scored: Vec<(String, f64)> = peer_infos
.iter()
.filter(|p| {
p.node_id != message.from
&& !message.path.contains(&p.node_id)
&& p.is_connected()
})
.map(|p| (p.node_id.clone(), q.get(dest, &p.node_id)))
.collect();
scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
scored
.into_iter()
.take(max_peers)
.map(|(id, _)| id)
.collect()
}
pub async fn q_advertised_value(&self, dest: &str, neighbours: &[String]) -> f64 {
self.q_state.read().await.best_over(dest, neighbours)
}
pub async fn q_record(
&self,
dest: &str,
peer: &str,
delivered: bool,
downstream_value: f64,
) {
let target = if delivered { downstream_value } else { 0.0 };
self.q_state.write().await.update(dest, peer, target);
}
pub async fn record_route_success(&self, peer_id: &str, latency: Duration) {
let mut history = self.route_history.write().await;
let stats = history
.entry(peer_id.to_string())
.or_insert_with(|| RouteStats {
success_count: 0,
failure_count: 0,
total_latency: Duration::ZERO,
sample_count: 0,
last_updated: Instant::now(),
});
stats.success_count += 1;
stats.total_latency += latency;
stats.sample_count += 1;
stats.last_updated = Instant::now();
drop(history);
let reward = (1.0 - 2.0 * latency.as_secs_f64()).clamp(0.5, 1.0);
self.ucb_state.write().await.record_reward(peer_id, reward);
self.persist_ucb_state().await;
}
pub async fn record_route_failure(&self, peer_id: &str) {
let mut history = self.route_history.write().await;
let stats = history
.entry(peer_id.to_string())
.or_insert_with(|| RouteStats {
success_count: 0,
failure_count: 0,
total_latency: Duration::ZERO,
sample_count: 0,
last_updated: Instant::now(),
});
stats.failure_count += 1;
stats.last_updated = Instant::now();
drop(history);
self.ucb_state.write().await.record_reward(peer_id, 0.0);
self.persist_ucb_state().await;
}
pub async fn cleanup_cache(&self) {
let mut cache = self.message_cache.write().await;
cache.retain(|_, _msg| {
true });
}
}
impl Clone for Router {
fn clone(&self) -> Self {
Self {
our_node_id: self.our_node_id.clone(),
seen_messages: self.seen_messages.clone(),
message_cache: self.message_cache.clone(),
route_history: self.route_history.clone(),
ucb_state: self.ucb_state.clone(),
q_state: self.q_state.clone(),
store: self.store.clone(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_murmuration_address_parse() {
let addr = MurmurationAddress::from_string("mur://node123").unwrap();
assert_eq!(addr.node_id, "node123");
let invalid = MurmurationAddress::from_string("invalid");
assert!(invalid.is_err());
}
#[test]
fn test_murmuration_address_to_string() {
let addr = MurmurationAddress {
node_id: "node123".to_string(),
};
assert_eq!(addr.to_string(), "mur://node123");
}
#[tokio::test]
async fn test_router_should_process() {
let router = Router::new("our-node".to_string());
let message = MeshMessage::new("peer1".to_string(), None, b"test".to_vec());
assert!(router.should_process(&message).await);
router.mark_seen(&message.message_id).await;
assert!(!router.should_process(&message).await);
}
#[tokio::test]
async fn test_router_is_for_us() {
let router = Router::new("our-node".to_string());
let broadcast = MeshMessage::new("peer1".to_string(), None, b"test".to_vec());
assert!(router.is_for_us(&broadcast));
let directed = MeshMessage::new(
"peer1".to_string(),
Some("our-node".to_string()),
b"test".to_vec(),
);
assert!(router.is_for_us(&directed));
let other = MeshMessage::new(
"peer1".to_string(),
Some("other-node".to_string()),
b"test".to_vec(),
);
assert!(!router.is_for_us(&other));
}
#[tokio::test]
async fn test_router_prepare_for_forwarding() {
let router = Router::new("our-node".to_string());
let message = MeshMessage::new("peer1".to_string(), None, b"test".to_vec());
let original_ttl = message.ttl;
let forwarded = router.prepare_for_forwarding(&message);
assert_eq!(forwarded.ttl, original_ttl - 1);
assert!(forwarded.path.contains(&"our-node".to_string()));
}
#[tokio::test]
async fn test_router_get_forward_peers() {
let router = Router::new("our-node".to_string());
let message = MeshMessage::new("peer1".to_string(), None, b"test".to_vec());
let all_peers = vec![
"peer1".to_string(),
"peer2".to_string(),
"peer3".to_string(),
];
let forward_peers = router.get_forward_peers(&message, &all_peers);
assert!(!forward_peers.contains(&"peer1".to_string()));
assert!(forward_peers.contains(&"peer2".to_string()));
assert!(forward_peers.contains(&"peer3".to_string()));
}
#[tokio::test]
async fn test_router_loop_detection() {
let router = Router::new("our-node".to_string());
let mut message = MeshMessage::new("peer1".to_string(), None, b"test".to_vec());
message.path.push("our-node".to_string());
assert!(!router.should_process(&message).await);
}
use crate::peer::{ConnectionState, PeerInfo};
use std::net::SocketAddr;
fn connected_peer(id: &str) -> PeerInfo {
let addr: SocketAddr = "127.0.0.1:9000".parse().unwrap();
let mut p = PeerInfo::new(id.to_string(), addr);
p.state = ConnectionState::Connected;
p
}
#[tokio::test]
async fn test_q_optimistic_init() {
let router = Router::new("u".to_string());
let v = router.q_advertised_value("dst", &["a".to_string()]).await;
assert_eq!(v, Q_INIT);
assert_eq!(router.q_advertised_value("dst", &[]).await, 0.0);
}
#[tokio::test]
async fn test_q_record_bootstraps_downstream_value() {
let router = Router::new("u".to_string());
router.q_record("dst", "a", true, 0.8).await;
let q_a = router.q_advertised_value("dst", &["a".to_string()]).await;
assert!((q_a - 0.97).abs() < 1e-9, "got {q_a}");
for _ in 0..20 {
router.q_record("dst", "b", false, 0.0).await;
}
let q_b = router
.q_state
.read()
.await
.get("dst", "b");
assert!(q_b < 0.1, "failures should drive Q→0, got {q_b}");
}
#[tokio::test]
async fn test_q_select_prefers_higher_value() {
let router = Router::new("u".to_string());
for _ in 0..30 {
router.q_record("dst", "good", true, 1.0).await;
router.q_record("dst", "bad", false, 0.0).await;
}
let msg = MeshMessage::new("src".to_string(), Some("dst".to_string()), vec![]);
let peers = vec![connected_peer("good"), connected_peer("bad")];
let picked = router.q_select_toward(&msg, &peers, 1, "dst").await;
assert_eq!(picked, vec!["good".to_string()]);
}
#[tokio::test]
async fn test_q_select_excludes_sender_and_path() {
let router = Router::new("u".to_string());
let mut msg = MeshMessage::new("sender".to_string(), Some("dst".to_string()), vec![]);
msg.path.push("visited".to_string());
let peers = vec![
connected_peer("sender"),
connected_peer("visited"),
connected_peer("fresh"),
];
let picked = router.q_select_toward(&msg, &peers, 3, "dst").await;
assert_eq!(picked, vec!["fresh".to_string()]);
}
}