#![allow(clippy::mutable_key_type)]
use crate::Dup;
use crate::ack::{AckPolicy, QUORUM_MET_SENTINEL};
use crate::actor::{Actor, ActorContext, Addr};
use crate::message::{BatchPut, Flush, Get, Message, Put};
use crate::types::{Children, NodeData, Value};
use crate::utils::{BoundedHashMap, try_send_or_log};
use async_trait::async_trait;
use log::{debug, error, info};
use rand::{rng, seq::IteratorRandom};
use std::collections::{HashMap, HashSet};
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use web_time::Instant;
static SEEN_MSGS_MAX_SIZE: usize = 10000;
static WRITE_CHANNEL_BOUND: usize = 1024;
struct SeenGetMessage {
from: Addr,
last_reply_checksum: Option<i32>,
}
pub(crate) struct QuorumEntry {
requester: Addr,
required: usize,
received: HashSet<Addr>,
started_at: Instant,
max_timeout: web_time::Duration,
}
impl QuorumEntry {
#[cfg(test)]
fn new(requester: Addr, policy: &AckPolicy) -> Self {
Self {
requester,
required: policy.quorum,
received: HashSet::new(),
started_at: Instant::now(),
max_timeout: policy.timeout,
}
}
fn record_ack(&mut self, from: &Addr) -> Option<usize> {
self.received.insert(from.clone());
if self.received.len() >= self.required {
Some(self.received.len())
} else {
None
}
}
fn is_expired(&self, timeout: web_time::Duration) -> bool {
self.started_at.elapsed() >= timeout
}
}
pub struct Router {
metrics: Arc<crate::metrics::Metrics>,
known_peers: HashSet<Addr>,
peer_addrs: HashMap<String, Addr>,
storage_adapters: HashSet<Addr>,
read_adapters: HashSet<Addr>,
write_adapters: HashSet<Addr>,
network_adapters: HashSet<Addr>,
storage_adapter_actors: Vec<Box<dyn Actor>>,
network_adapter_actors: Vec<Box<dyn Actor>>,
server_peers: HashSet<Addr>,
dup: Dup,
seen_get_messages: BoundedHashMap<String, SeenGetMessage>,
subscribers_by_topic: HashMap<String, HashSet<Addr>>,
msg_counter: AtomicUsize,
quorum_entries: BoundedHashMap<String, QuorumEntry>,
}
#[async_trait]
impl Actor for Router {
async fn pre_start(&mut self, ctx: &ActorContext) {
while let Some(adapter) = self.storage_adapter_actors.pop() {
match adapter.try_clone_storage() {
Some(read_actor) => {
let read_addr = ctx.start_actor(read_actor);
let write_addr = ctx.start_actor_bounded(adapter, WRITE_CHANNEL_BOUND);
self.storage_adapters.insert(read_addr.clone());
self.storage_adapters.insert(write_addr.clone());
self.read_adapters.insert(read_addr);
self.write_adapters.insert(write_addr);
}
None => {
let addr = ctx.start_actor(adapter);
self.storage_adapters.insert(addr.clone());
self.read_adapters.insert(addr.clone());
self.write_adapters.insert(addr);
}
}
}
while let Some(adapter) = self.network_adapter_actors.pop() {
let subscribe_to_everything = adapter.subscribe_to_everything();
let addr = ctx.start_actor(adapter);
self.network_adapters.insert(addr.clone());
if subscribe_to_everything {
self.server_peers.insert(addr);
}
}
#[cfg(not(target_arch = "wasm32"))]
{
let ctx_addr = ctx.addr.clone();
ctx.child_task(async move {
let mut interval = crate::tokio_time::interval(web_time::Duration::from_secs(1));
interval.tick().await; loop {
interval.tick().await;
let _ = ctx_addr.send(Message::CheckQuorumTimeouts);
}
});
}
}
async fn stopping(&mut self, _ctx: &ActorContext) {
info!("Router stopping");
}
async fn handle(&mut self, msg: Message, _ctx: &ActorContext) {
debug!("incoming message {}", msg.clone().to_string());
match msg {
Message::Put(put) => self.handle_put(put),
Message::BatchPut(batch) => {
self.handle_batch_put(batch);
}
Message::Get(get) => self.handle_get(get),
Message::Flush(flush) => self.handle_flush(flush),
Message::Hi { from, peer_id } => {
self.known_peers.insert(from.clone());
if !peer_id.is_empty() {
if let Some(existing) = self.peer_addrs.get(&peer_id) {
if existing != &from {
error!(
"Router peer_id collision: '{}' already mapped to {:?}, rejecting {:?}. Each peer_id must be unique.",
peer_id, existing, from
);
return;
}
}
self.peer_addrs.insert(peer_id, from);
}
}
Message::RtcSignal(rtc) => {
debug!(
"RtcSignal id={} to={:?} known_peers={}",
rtc.id,
rtc.to,
self.known_peers.len()
);
if let Some(to_peer_id) = &rtc.to {
if let Some(addr) = self.peer_addrs.get(to_peer_id) {
debug!(
"RtcSignal delivering to local addr for peer_id={}",
to_peer_id
);
let _ = addr.send(Message::RtcSignal(rtc));
} else {
debug!(
"RtcSignal broadcasting to {} known_peers",
self.known_peers.len()
);
for addr in self.known_peers.iter() {
let _ = addr.send(Message::RtcSignal(rtc.clone()));
}
}
}
}
Message::RegisterQuorum {
put_id,
requester,
policy,
} => {
let _ = self.handle_register_quorum(put_id, requester, policy);
}
Message::CheckQuorumTimeouts => {
self.handle_quorum_timeout_reaper();
}
};
}
}
impl Router {
pub fn new(
storage_adapter_actors: Vec<Box<dyn Actor>>,
network_adapter_actors: Vec<Box<dyn Actor>>,
metrics: Arc<crate::metrics::Metrics>,
) -> Self {
Self {
metrics,
known_peers: HashSet::new(),
peer_addrs: HashMap::new(),
storage_adapters: HashSet::new(),
read_adapters: HashSet::new(),
write_adapters: HashSet::new(),
network_adapters: HashSet::new(),
storage_adapter_actors,
network_adapter_actors,
server_peers: HashSet::new(),
dup: Dup::default_gun(),
seen_get_messages: BoundedHashMap::new(SEEN_MSGS_MAX_SIZE),
subscribers_by_topic: HashMap::new(),
msg_counter: AtomicUsize::new(0),
quorum_entries: BoundedHashMap::new(SEEN_MSGS_MAX_SIZE),
}
}
#[allow(dead_code)]
pub fn metrics(&self) -> Arc<crate::metrics::Metrics> {
self.metrics.clone()
}
fn handle_get(&mut self, get: Get) {
if !get.id.chars().all(char::is_alphanumeric) {
error!("id {}", get.id);
}
if self.is_message_seen(&get.id) {
return;
}
let seen_get_message = SeenGetMessage {
from: get.from.clone(),
last_reply_checksum: get.checksum,
};
self.seen_get_messages
.insert(get.id.clone(), seen_get_message);
let topic = get.node_id.split("/").next().unwrap_or("");
debug!("{} subscribed to {}", get.from, topic);
self.subscribers_by_topic
.entry(topic.to_string())
.or_default()
.insert(get.from.clone());
for addr in self.read_adapters.iter() {
let _ = addr.send(Message::Get(get.clone()));
}
let mut already_sent_to = HashSet::new();
for addr in self.server_peers.iter() {
debug!("send to server peer");
let _ = addr.send(Message::Get(get.clone()));
already_sent_to.insert(addr.clone());
}
let mut errored = HashSet::new();
let mut sent_to = 0;
let mut rng = rng();
if let Some(topic_subscribers) = self.subscribers_by_topic.get(topic) {
let sample = topic_subscribers.iter().choose_multiple(&mut rng, 4);
for addr in sample {
if get.from == *addr {
continue;
}
if already_sent_to.contains(addr) {
continue;
}
already_sent_to.insert(addr.clone());
match addr.send(Message::Get(get.clone())) {
Ok(_) => {
sent_to += 1;
}
_ => {
#[cfg(target_arch = "wasm32")]
web_sys::console::log_1(&format!("router: FAILED to send put to known_peer {}", addr).into());
errored.insert(addr.clone());
}
}
}
}
debug!(
"sent get to a random sample of subscribers of size {}",
sent_to
);
if !errored.is_empty() {
if let Some(topic_subscribers) = self.subscribers_by_topic.get_mut(topic) {
for addr in errored {
topic_subscribers.remove(&addr);
self.known_peers.remove(&addr);
}
}
}
if sent_to < 4 {
let mut errored = HashSet::new();
while let Some(addr) = self.known_peers.iter().choose(&mut rng) {
sent_to += 1;
if sent_to >= 4 {
break;
}
if get.from == *addr {
continue;
}
if already_sent_to.contains(addr) {
continue;
}
already_sent_to.insert(addr.clone());
match addr.send(Message::Get(get.clone())) {
Ok(_) => {}
_ => {
#[cfg(target_arch = "wasm32")]
web_sys::console::log_1(&format!("router: FAILED to send put to known_peer {}", addr).into());
errored.insert(addr.clone());
}
}
}
for addr in errored {
self.known_peers.remove(&addr);
}
}
}
fn handle_put(&mut self, put: Put) {
if self.is_message_seen(&put.id) {
return;
}
if let (Some(ack), Some(hash)) = (&put.in_response_to, put.checksum) {
let checksum_key = format!("{}##{}", ack, hash);
if self.dup.check(&checksum_key) {
debug!("duplicate response checksum: {}", checksum_key);
return;
}
self.dup.track(&checksum_key);
}
match &put.in_response_to {
Some(in_response_to) => {
if let Some(entry) = self.quorum_entries.get_mut(in_response_to) {
let ack_count = entry.record_ack(&put.from);
if let Some(count) = ack_count {
let children: Children = std::collections::BTreeMap::from([(
"_".to_string(),
NodeData {
value: Value::Number(count as f64),
updated_at: 0.0, },
)]);
let mut reply = Put::new_from_kv(
QUORUM_MET_SENTINEL.to_string(),
children,
put.from.clone(),
);
reply.in_response_to = Some(in_response_to.clone());
debug!("quorum met for {} ({} acks)", in_response_to, count);
try_send_or_log(
&entry.requester,
Message::Put(reply),
&self.metrics,
"router:quorum-met",
);
self.quorum_entries.take(in_response_to);
}
return; }
if let Some(seen_get_message) = self.seen_get_messages.get_mut(in_response_to) {
if put.checksum.is_some()
&& put.checksum == seen_get_message.last_reply_checksum
{
debug!("same reply already sent");
return;
}
seen_get_message.last_reply_checksum = put.checksum;
try_send_or_log(
&seen_get_message.from,
Message::Put(put),
&self.metrics,
"router:get-reply",
);
}
}
_ => {
for addr in self.write_adapters.iter() {
if put.from == *addr {
continue;
}
let _ = addr.send(Message::Put(put.clone()));
debug!("sent to write adapter {}", addr);
}
self.handle_put_relay(&put);
}
};
}
fn handle_put_relay(&mut self, put: &Put) {
#[cfg(target_arch = "wasm32")]
{
web_sys::console::log_1(&format!(
"router.handle_put_relay: from={} server_peers={} known_peers={} subscribers={}",
put.from,
self.server_peers.len(),
self.known_peers.len(),
self.subscribers_by_topic.len(),
).into());
for addr in self.known_peers.iter() {
web_sys::console::log_1(&format!(" known_peer: {} (from matches: {})", addr, *addr == put.from).into());
}
}
let mut hops = put.peer_hop_list.clone().unwrap_or_default();
hops.insert(put.from.to_string());
let mut already_sent_to = HashSet::new();
for addr in self.server_peers.iter() {
if put.from == *addr || hops.contains(&addr.to_string()) {
continue;
}
let mut put = put.clone();
put.peer_hop_list = Some(hops.clone());
let _ = addr.send(Message::Put(put));
already_sent_to.insert(addr.clone());
}
let mut sent_to = 0;
for node_id in put.clone().updated_nodes.keys() {
let topic = node_id.split("/").next().unwrap_or("");
if let Some(topic_subscribers) = self.subscribers_by_topic.get_mut(topic) {
topic_subscribers.retain(|addr| {
if put.from == *addr || hops.contains(&addr.to_string()) {
return true;
}
if already_sent_to.contains(addr) {
return true;
}
already_sent_to.insert(addr.clone());
let mut put = put.clone();
put.peer_hop_list = Some(hops.clone());
match addr.send(Message::Put(put)) {
Ok(_) => {
sent_to += 1;
true
}
_ => false,
}
})
}
}
debug!("sent put to {} subscribers", already_sent_to.len());
if already_sent_to.len() < 4 {
#[cfg(target_arch = "wasm32")]
web_sys::console::log_1(&format!("router: entering random sampling, known_peers={}, sent_to={}", self.known_peers.len(), sent_to).into());
let mut rng = rng();
let mut errored = HashSet::new();
while let Some(addr) = self.known_peers.iter().choose(&mut rng) {
sent_to += 1;
if sent_to >= 4 {
break;
}
if already_sent_to.contains(addr) {
continue;
}
already_sent_to.insert(addr.clone());
if put.from == *addr || hops.contains(&addr.to_string()) {
continue;
}
let mut put = put.clone();
put.peer_hop_list = Some(hops.clone());
match addr.send(Message::Put(put)) {
Ok(_) => {
#[cfg(target_arch = "wasm32")]
web_sys::console::log_1(&format!("router: sent put to known_peer {}", addr).into());
debug!("sent put to random peer");
}
_ => {
#[cfg(target_arch = "wasm32")]
web_sys::console::log_1(&format!("router: FAILED to send put to known_peer {}", addr).into());
errored.insert(addr.clone());
}
}
}
for addr in errored {
self.known_peers.remove(&addr);
}
}
}
fn handle_register_quorum(
&mut self,
put_id: String,
requester: Addr,
policy: AckPolicy,
) -> Result<(), String> {
let required = policy.quorum;
let max_timeout = policy.timeout;
let entry = QuorumEntry {
requester,
required,
received: HashSet::new(),
started_at: web_time::Instant::now(),
max_timeout,
};
self.quorum_entries.insert(put_id.clone(), entry);
debug!(
"registered quorum for put_id={} (required: {} peers, timeout: {:?})",
put_id, required, policy.timeout
);
Ok(())
}
fn handle_quorum_timeout_reaper(&mut self) {
let expired_keys: Vec<String> = self
.quorum_entries
.iter()
.filter(|(_k, v)| v.is_expired(v.max_timeout))
.map(|(k, _v)| k.clone())
.collect();
if expired_keys.is_empty() {
return;
}
let mut expired: Vec<(String, QuorumEntry)> = Vec::with_capacity(expired_keys.len());
for key in expired_keys {
if let Some(entry) = self.quorum_entries.take(&key) {
expired.push((key, entry));
}
}
debug!(
"quorum reaper: timing out {} expired entr{}",
expired.len(),
if expired.len() == 1 { "y" } else { "ies" }
);
for (put_id, entry) in expired {
let mut children: Children = std::collections::BTreeMap::new();
children.insert(
"_".to_string(),
NodeData {
value: Value::Bit(true),
updated_at: 0.0,
},
);
let mut reply = Put::new_from_kv(
QUORUM_MET_SENTINEL.to_string(),
children,
entry.requester.clone(),
);
reply.in_response_to = Some(put_id.clone());
try_send_or_log(
&entry.requester,
Message::Put(reply),
&self.metrics,
"router:quorum-timeout",
);
debug!(
"quorum reaper: notified requester of timeout for put_id={}",
put_id
);
}
}
fn handle_batch_put(&mut self, batch: BatchPut) {
for addr in self.write_adapters.iter() {
if batch.from == *addr {
continue;
}
let _ = addr.send(Message::BatchPut(batch.clone()));
}
for put in batch.puts {
if self.is_message_seen(&put.id) {
continue;
}
if let (Some(ack), Some(hash)) = (&put.in_response_to, put.checksum) {
let checksum_key = format!("{}##{}", ack, hash);
if self.dup.check(&checksum_key) {
debug!("batch: duplicate response checksum: {}", checksum_key);
continue;
}
self.dup.track(&checksum_key);
}
if let Some(in_response_to) = &put.in_response_to {
if let Some(seen_get_message) = self.seen_get_messages.get_mut(in_response_to) {
if put.checksum == seen_get_message.last_reply_checksum {
continue;
}
seen_get_message.last_reply_checksum = put.checksum;
try_send_or_log(
&seen_get_message.from,
Message::Put(put),
&self.metrics,
"router:get-reply",
);
}
continue;
}
self.handle_put_relay(&put);
}
}
fn handle_flush(&mut self, flush: Flush) {
let mut sent = HashSet::new();
for addr in self.write_adapters.iter() {
if flush.from == *addr {
continue;
}
if sent.contains(addr) {
continue;
}
sent.insert(addr.clone());
let _ = addr.send(Message::Flush(flush.clone()));
}
debug!("forwarded flush to {} storage write adapters", sent.len());
}
fn is_message_seen(&mut self, id: &String) -> bool {
self.msg_counter.fetch_add(1, Ordering::Relaxed);
if self.dup.check(id) {
debug!("already seen message {}", id);
return true;
}
self.dup.track(id);
false
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::adapters::MemoryStorage;
use crate::metrics::Metrics;
use web_time::Duration;
#[test]
fn test_router_new() {
let storage = vec![Box::new(MemoryStorage::new()) as Box<dyn Actor>];
let metrics = Arc::new(Metrics::new());
let router = Router::new(storage, vec![], metrics);
assert!(router.known_peers.is_empty());
assert!(router.read_adapters.is_empty());
assert!(router.write_adapters.is_empty());
assert!(router.network_adapters.is_empty());
}
#[test]
fn test_router_default_dedup() {
let metrics = Arc::new(Metrics::new());
let router = Router::new(vec![], vec![], metrics);
assert_eq!(router.dup.max(), 999);
assert_eq!(router.dup.age(), web_time::Duration::from_secs(9));
}
#[test]
fn test_router_seen_msg_capacity() {
let metrics = Arc::new(Metrics::new());
let router = Router::new(vec![], vec![], metrics);
assert_eq!(SEEN_MSGS_MAX_SIZE, 10000);
let _ = router; }
#[test]
fn test_router_msg_counter_starts_zero() {
let metrics = Arc::new(Metrics::new());
let router = Router::new(vec![], vec![], metrics);
assert_eq!(router.msg_counter.load(Ordering::Relaxed), 0);
}
fn _make_quorum_entry(required: usize, timeout_ms: u64) -> QuorumEntry {
QuorumEntry::new(
Addr::noop(),
&AckPolicy::any()
.with_quorum(required)
.with_timeout(Duration::from_millis(timeout_ms)),
)
}
#[test]
fn quorum_entry_initial_state() {
let entry = _make_quorum_entry(3, 60_000);
assert_eq!(entry.received.len(), 0);
assert_eq!(entry.required, 3);
assert!(!entry.is_expired(Duration::from_millis(60_000)));
}
#[test]
fn quorum_entry_is_expired_respects_timeout() {
let entry = _make_quorum_entry(1, 60_000);
assert!(
!entry.is_expired(Duration::from_secs(60)),
"fresh entry should not be expired under 60s timeout"
);
assert!(
entry.is_expired(Duration::from_nanos(1)),
"1ns timeout should be exceeded by microsecond-level elapsed"
);
}
#[test]
fn quorum_entry_required_field_from_policy() {
let entry_any = QuorumEntry::new(Addr::noop(), &AckPolicy::any());
assert_eq!(entry_any.required, 1);
let entry_all = QuorumEntry::new(Addr::noop(), &AckPolicy::all());
assert_eq!(entry_all.required, usize::MAX);
assert_eq!(
QuorumEntry::new(Addr::noop(), &AckPolicy::for_peer_count(0)).required,
1
);
assert_eq!(
QuorumEntry::new(Addr::noop(), &AckPolicy::for_peer_count(5)).required,
3
);
assert_eq!(
QuorumEntry::new(Addr::noop(), &AckPolicy::for_peer_count(7)).required,
4
);
}
}