use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use tokio::sync::broadcast;
use tracing::{debug, warn};
use super::registry::NodeId;
#[derive(Debug, Clone, Serialize, Deserialize)]
#[non_exhaustive]
pub struct RelayMessage {
pub seq: u64,
pub from: NodeId,
pub to: String,
pub topic: String,
pub payload: serde_json::Value,
pub timestamp: DateTime<Utc>,
}
impl RelayMessage {
pub fn new(
seq: u64,
from: impl Into<NodeId>,
to: impl Into<String>,
topic: impl Into<String>,
payload: serde_json::Value,
) -> Self {
Self {
seq,
from: from.into(),
to: to.into(),
topic: topic.into(),
payload,
timestamp: Utc::now(),
}
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct IncomingMessage {
pub message: RelayMessage,
pub is_broadcast: bool,
}
#[derive(Debug, Clone, Default)]
#[non_exhaustive]
pub struct RelayStats {
pub messages_sent: u64,
pub messages_received: u64,
pub duplicates_dropped: u64,
}
pub struct Relay {
node_id: NodeId,
next_seq: AtomicU64,
seen: std::sync::Mutex<HashMap<NodeId, u64>>,
tx: broadcast::Sender<IncomingMessage>,
stats: std::sync::Mutex<RelayStats>,
}
impl Relay {
pub fn new(node_id: impl Into<NodeId>, capacity: usize) -> Self {
let (tx, _) = broadcast::channel(capacity);
Self {
node_id: node_id.into(),
next_seq: AtomicU64::new(1),
seen: std::sync::Mutex::new(HashMap::new()),
tx,
stats: std::sync::Mutex::new(RelayStats::default()),
}
}
pub fn with_defaults(node_id: impl Into<NodeId>) -> Self {
Self::new(node_id, 256)
}
pub fn send(
&self,
to: impl Into<String>,
topic: impl Into<String>,
payload: serde_json::Value,
) -> u64 {
let seq = self.next_seq.fetch_add(1, Ordering::AcqRel);
let msg = RelayMessage {
seq,
from: self.node_id.clone(),
to: to.into(),
topic: topic.into(),
payload,
timestamp: Utc::now(),
};
debug!(seq, to = %msg.to, topic = %msg.topic, "relay: sending message");
let incoming = IncomingMessage {
is_broadcast: msg.to.is_empty(),
message: msg,
};
let _ = self.tx.send(incoming);
if let Ok(mut stats) = self.stats.lock() {
stats.messages_sent += 1;
}
seq
}
pub fn broadcast(&self, topic: impl Into<String>, payload: serde_json::Value) -> u64 {
self.send("", topic, payload)
}
pub fn subscribe(&self) -> broadcast::Receiver<IncomingMessage> {
self.tx.subscribe()
}
pub fn receive(&self, msg: RelayMessage) -> Option<IncomingMessage> {
if msg.from == self.node_id {
return None;
}
if !msg.to.is_empty() && msg.to != self.node_id {
return None;
}
let mut seen = self.seen.lock().unwrap_or_else(|e| {
warn!("relay seen-map mutex was poisoned, resetting");
let mut inner = e.into_inner();
inner.clear(); inner
});
let last_seen = seen.entry(msg.from.clone()).or_insert(0);
if msg.seq <= *last_seen {
debug!(
seq = msg.seq,
from = %msg.from,
"relay: dropping duplicate message"
);
if let Ok(mut stats) = self.stats.lock() {
stats.duplicates_dropped += 1;
}
return None;
}
if msg.seq != *last_seen + 1 {
warn!(
expected = *last_seen + 1,
got = msg.seq,
from = %msg.from,
"relay: sequence gap detected"
);
}
*last_seen = msg.seq;
if let Ok(mut stats) = self.stats.lock() {
stats.messages_received += 1;
}
let incoming = IncomingMessage {
is_broadcast: msg.to.is_empty(),
message: msg,
};
let _ = self.tx.send(incoming.clone());
Some(incoming)
}
pub fn node_id(&self) -> &str {
&self.node_id
}
pub fn stats(&self) -> RelayStats {
self.stats
.lock()
.unwrap_or_else(|e| {
warn!("relay stats mutex was poisoned, recovering");
e.into_inner()
})
.clone()
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn send_increments_seq() {
let relay = Relay::with_defaults("node-a");
let s1 = relay.send("node-b", "task", json!({"id": 1}));
let s2 = relay.send("node-b", "task", json!({"id": 2}));
assert_eq!(s1, 1);
assert_eq!(s2, 2);
}
#[test]
fn broadcast_uses_empty_to() {
let relay = Relay::with_defaults("node-a");
let mut rx = relay.subscribe();
relay.broadcast("heartbeat", json!({}));
let msg = rx.try_recv().expect("should receive broadcast");
assert!(msg.is_broadcast);
assert!(msg.message.to.is_empty());
}
#[test]
fn receive_dedup_drops_duplicate() {
let relay = Relay::with_defaults("node-b");
let msg1 = RelayMessage {
seq: 1,
from: "node-a".into(),
to: "node-b".into(),
topic: "test".into(),
payload: json!({}),
timestamp: Utc::now(),
};
let msg2 = msg1.clone();
assert!(relay.receive(msg1).is_some());
assert!(relay.receive(msg2).is_none());
let stats = relay.stats();
assert_eq!(stats.messages_received, 1);
assert_eq!(stats.duplicates_dropped, 1);
}
#[test]
fn receive_skips_own_messages() {
let relay = Relay::with_defaults("node-a");
let msg = RelayMessage {
seq: 1,
from: "node-a".into(),
to: "".into(),
topic: "test".into(),
payload: json!({}),
timestamp: Utc::now(),
};
assert!(relay.receive(msg).is_none());
}
#[test]
fn receive_skips_messages_for_other_nodes() {
let relay = Relay::with_defaults("node-b");
let msg = RelayMessage {
seq: 1,
from: "node-a".into(),
to: "node-c".into(),
topic: "test".into(),
payload: json!({}),
timestamp: Utc::now(),
};
assert!(relay.receive(msg).is_none());
}
#[test]
fn receive_accepts_broadcast() {
let relay = Relay::with_defaults("node-b");
let msg = RelayMessage {
seq: 1,
from: "node-a".into(),
to: "".into(),
topic: "heartbeat".into(),
payload: json!({"status": "ok"}),
timestamp: Utc::now(),
};
let incoming = relay.receive(msg).expect("should accept broadcast");
assert!(incoming.is_broadcast);
}
#[test]
fn receive_detects_sequence_gap() {
let relay = Relay::with_defaults("node-b");
let msg = RelayMessage {
seq: 2,
from: "node-a".into(),
to: "node-b".into(),
topic: "test".into(),
payload: json!({}),
timestamp: Utc::now(),
};
assert!(relay.receive(msg).is_some());
}
#[test]
fn stats_tracking() {
let relay = Relay::with_defaults("node-a");
relay.send("node-b", "t", json!({}));
relay.send("node-b", "t", json!({}));
let stats = relay.stats();
assert_eq!(stats.messages_sent, 2);
}
#[test]
fn subscriber_receives_sent_messages() {
let relay = Relay::with_defaults("node-a");
let mut rx = relay.subscribe();
relay.send("node-b", "task", json!({"work": true}));
let msg = rx.try_recv().expect("subscriber should receive message");
assert_eq!(msg.message.topic, "task");
}
}