use std::hash::Hasher;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Mutex, OnceLock};
use std::time::{Duration, Instant};
use dashmap::DashMap;
use tracing::{debug, warn};
use xxhash_rust::xxh3::Xxh3Default;
use arc_swap::ArcSwap;
use crate::error::{Error, Result};
use crate::proto::grpc::block::WorkerInfo;
use crate::proto::grpc::WorkerNetAddress;
const DEFAULT_FAILURE_TTL: Duration = Duration::from_secs(60);
const DEFAULT_WORKER_REFRESH_TTL: Duration = Duration::from_secs(30);
const VIRTUAL_NODES_PER_WORKER: u32 = 100;
fn new_failed_workers_map() -> DashMap<String, Instant> {
DashMap::with_capacity_and_shard_amount(0, 2)
}
pub struct WorkerRouter {
workers: ArcSwap<Vec<WorkerInfo>>,
failed_workers: OnceLock<DashMap<String, Instant>>,
failed_count: AtomicUsize,
failure_ttl: Duration,
last_refresh: Mutex<Instant>,
worker_refresh_ttl: Duration,
local_worker_id: ArcSwap<Option<Option<i64>>>,
hash_ring: ArcSwap<Vec<(u64, usize)>>,
}
impl WorkerRouter {
pub fn new() -> Self {
Self {
workers: ArcSwap::from_pointee(Vec::new()),
failed_workers: OnceLock::new(),
failed_count: AtomicUsize::new(0),
failure_ttl: DEFAULT_FAILURE_TTL,
last_refresh: Mutex::new(Instant::now()),
worker_refresh_ttl: DEFAULT_WORKER_REFRESH_TTL,
local_worker_id: ArcSwap::from_pointee(None),
hash_ring: ArcSwap::from_pointee(Vec::new()),
}
}
pub fn with_failure_ttl(failure_ttl: Duration) -> Self {
Self {
workers: ArcSwap::from_pointee(Vec::new()),
failed_workers: OnceLock::new(),
failed_count: AtomicUsize::new(0),
failure_ttl,
last_refresh: Mutex::new(Instant::now()),
worker_refresh_ttl: DEFAULT_WORKER_REFRESH_TTL,
local_worker_id: ArcSwap::from_pointee(None),
hash_ring: ArcSwap::from_pointee(Vec::new()),
}
}
pub fn with_ttls(failure_ttl: Duration, worker_refresh_ttl: Duration) -> Self {
Self {
workers: ArcSwap::from_pointee(Vec::new()),
failed_workers: OnceLock::new(),
failed_count: AtomicUsize::new(0),
failure_ttl,
last_refresh: Mutex::new(Instant::now()),
worker_refresh_ttl,
local_worker_id: ArcSwap::from_pointee(None),
hash_ring: ArcSwap::from_pointee(Vec::new()),
}
}
pub fn snapshot_from(shared: &WorkerRouter) -> Self {
Self {
workers: ArcSwap::new(shared.workers.load_full()),
hash_ring: ArcSwap::new(shared.hash_ring.load_full()),
failed_workers: OnceLock::new(),
failed_count: AtomicUsize::new(0),
failure_ttl: shared.failure_ttl,
last_refresh: Mutex::new(Instant::now()),
worker_refresh_ttl: shared.worker_refresh_ttl,
local_worker_id: ArcSwap::new(shared.local_worker_id.load_full()),
}
}
pub async fn update_workers(&self, workers: Vec<WorkerInfo>) {
let new_fp = workers_fingerprint(&workers);
let cur_workers = self.workers.load_full();
let cur_fp = workers_fingerprint(&cur_workers);
if new_fp == cur_fp && !cur_workers.is_empty() {
*self
.last_refresh
.lock()
.expect("last_refresh mutex poisoned") = Instant::now();
return;
}
let new_ring = Arc::new(build_hash_ring(&workers));
let new_snapshot = Arc::new(workers);
self.workers.store(new_snapshot.clone());
self.hash_ring.store(new_ring);
*self
.last_refresh
.lock()
.expect("last_refresh mutex poisoned") = Instant::now();
let detected = Self::detect_local_worker(&new_snapshot).await;
let cached = if detected > 0 { Some(detected) } else { None };
self.local_worker_id.store(Arc::new(Some(cached)));
}
pub async fn get_workers(&self) -> Arc<Vec<WorkerInfo>> {
self.workers.load_full()
}
pub fn workers_is_empty(&self) -> bool {
self.workers.load().is_empty()
}
pub async fn needs_refresh(&self) -> bool {
self.last_refresh
.lock()
.expect("last_refresh mutex poisoned")
.elapsed()
>= self.worker_refresh_ttl
}
pub async fn refresh_workers(&self, wm: &crate::client::WorkerManagerClient) -> Result<()> {
match wm.get_worker_info_list().await {
Ok(workers) => {
debug!(count = workers.len(), "worker list refreshed");
self.update_workers(workers).await;
Ok(())
}
Err(e) => {
warn!("worker list refresh failed, keeping stale list: {}", e);
*self
.last_refresh
.lock()
.expect("last_refresh mutex poisoned") = Instant::now();
Ok(())
}
}
}
async fn detect_local_worker(workers: &[WorkerInfo]) -> i64 {
let local_names = Self::local_hostnames();
for w in workers {
if let Some(addr) = &w.address {
let host = addr.host.as_deref().unwrap_or("");
if host.is_empty() {
continue;
}
if local_names.iter().any(|n| n == host) || Self::is_local_address(host) {
let id = w.id.unwrap_or(0);
debug!(host = %host, worker_id = id, "detected local worker");
return id;
}
}
}
0
}
fn local_hostnames() -> Vec<String> {
let mut names = vec![
"localhost".to_string(),
"127.0.0.1".to_string(),
"::1".to_string(),
];
if let Ok(h) = hostname::get() {
if let Ok(s) = h.into_string() {
names.push(s.clone());
if let Some(short) = s.split('.').next() {
names.push(short.to_string());
}
}
}
names
}
fn is_local_address(host: &str) -> bool {
use std::net::UdpSocket;
UdpSocket::bind((host, 0u16)).is_ok()
}
pub async fn select_worker(&self, block_id: i64) -> Result<WorkerInfo> {
let workers = self.workers.load_full();
let ring = self.hash_ring.load_full();
let local = self.local_worker_id.load_full();
if workers.is_empty() {
return Err(Error::NoWorkerAvailable {
message: "no workers registered".to_string(),
});
}
self.cleanup_expired_failures();
let local_id_opt: Option<i64> = match *local {
Some(cached) => cached,
None => {
let detected = Self::detect_local_worker(&workers).await;
let cached_value = if detected > 0 { Some(detected) } else { None };
self.local_worker_id.store(Arc::new(Some(cached_value)));
cached_value
}
};
if let Some(local_id) = local_id_opt {
if let Some(local_w) = workers.iter().find(|w| w.id == Some(local_id)) {
if let Some(addr) = &local_w.address {
if !self.is_failed(&worker_addr_key(addr)) {
return Ok(local_w.clone());
}
}
}
}
if let Some(w) = self.consistent_hash_select_with_ring(block_id, &workers, &ring, true) {
return Ok(w);
}
self.consistent_hash_select_with_ring(block_id, &workers, &ring, false)
.ok_or_else(|| Error::NoWorkerAvailable {
message: format!("no suitable worker for block_id={}", block_id),
})
}
pub fn mark_failed(&self, addr: &WorkerNetAddress) {
let key = worker_addr_key(addr);
let map = self.failed_workers.get_or_init(new_failed_workers_map);
if map.insert(key, Instant::now()).is_none() {
self.failed_count.fetch_add(1, Ordering::Relaxed);
}
}
pub async fn is_block_source_local(&self, block_id: i64) -> bool {
let Ok(selected) = self.select_worker(block_id).await else {
return false;
};
match **self.local_worker_id.load() {
Some(Some(local_id)) => selected.id == Some(local_id),
_ => false,
}
}
pub async fn pick_any_worker(&self) -> Result<WorkerInfo> {
let workers = self.workers.load_full();
if workers.is_empty() {
return Err(Error::NoWorkerAvailable {
message: "no workers registered".to_string(),
});
}
self.cleanup_expired_failures();
let eligible: Vec<WorkerInfo> = workers
.iter()
.filter(|w| {
if let Some(addr) = w.address.as_ref() {
let key = worker_addr_key(addr);
!self.is_failed(&key)
} else {
false
}
})
.cloned()
.collect();
let pool = if eligible.is_empty() {
(*workers).clone()
} else {
eligible
};
if pool.is_empty() {
return Err(Error::NoWorkerAvailable {
message: "no eligible workers".to_string(),
});
}
let idx = rand::Rng::random_range(&mut rand::rng(), 0..pool.len());
Ok(pool[idx].clone())
}
fn is_failed(&self, key: &str) -> bool {
let Some(map) = self.failed_workers.get() else {
return false;
};
if let Some(entry) = map.get(key) {
entry.value().elapsed() < self.failure_ttl
} else {
false
}
}
fn cleanup_expired_failures(&self) {
if self.failed_count.load(Ordering::Relaxed) == 0 {
return;
}
let Some(map) = self.failed_workers.get() else {
return;
};
let ttl = self.failure_ttl;
let mut removed: usize = 0;
map.retain(|_, v| {
if v.elapsed() < ttl {
true
} else {
removed += 1;
false
}
});
if removed > 0 {
self.failed_count.fetch_sub(removed, Ordering::Relaxed);
}
}
#[cfg(test)]
fn failed_workers_is_uninitialised(&self) -> bool {
self.failed_workers.get().is_none()
}
#[cfg(test)]
fn failed_workers_len(&self) -> usize {
self.failed_workers.get().map_or(0, |m| m.len())
}
fn consistent_hash_select_with_ring(
&self,
block_id: i64,
workers: &[WorkerInfo],
ring: &[(u64, usize)],
skip_failed: bool,
) -> Option<WorkerInfo> {
consistent_hash_select_from_ring(block_id, workers, ring, skip_failed, |key| {
self.is_failed(key)
})
}
}
fn consistent_hash_select_from_ring<F>(
block_id: i64,
workers: &[WorkerInfo],
ring: &[(u64, usize)],
skip_failed: bool,
is_failed_fn: F,
) -> Option<WorkerInfo>
where
F: Fn(&str) -> bool,
{
if ring.is_empty() || workers.is_empty() {
return None;
}
let target = hash_block_id(block_id);
let start = ring
.binary_search_by_key(&target, |(h, _)| *h)
.unwrap_or_else(|p| p)
% ring.len();
for offset in 0..ring.len() {
let pos = (start + offset) % ring.len();
let worker_idx = ring[pos].1;
let Some(w) = workers.get(worker_idx) else {
continue;
};
if !skip_failed {
return Some(w.clone());
}
if let Some(addr) = w.address.as_ref() {
if !is_failed_fn(&worker_addr_key(addr)) {
return Some(w.clone());
}
}
}
None
}
pub struct WorkerRouterView {
workers: Arc<Vec<WorkerInfo>>,
hash_ring: Arc<Vec<(u64, usize)>>,
local_worker_id: Option<i64>,
failed_workers: OnceLock<DashMap<String, Instant>>,
failed_count: AtomicUsize,
failure_ttl: Duration,
}
impl WorkerRouterView {
pub fn from_shared(shared: &WorkerRouter) -> Self {
let local_worker_id: Option<i64> = (*shared.local_worker_id.load_full()).flatten();
Self {
workers: shared.workers.load_full(),
hash_ring: shared.hash_ring.load_full(),
local_worker_id,
failed_workers: OnceLock::new(),
failed_count: AtomicUsize::new(0),
failure_ttl: shared.failure_ttl,
}
}
pub fn from_workers(workers: Vec<WorkerInfo>, failure_ttl: Duration) -> Self {
let ring = Arc::new(build_hash_ring(&workers));
Self {
workers: Arc::new(workers),
hash_ring: ring,
local_worker_id: None,
failed_workers: OnceLock::new(),
failed_count: AtomicUsize::new(0),
failure_ttl,
}
}
pub fn empty() -> Self {
Self {
workers: Arc::new(Vec::new()),
hash_ring: Arc::new(Vec::new()),
local_worker_id: None,
failed_workers: OnceLock::new(),
failed_count: AtomicUsize::new(0),
failure_ttl: DEFAULT_FAILURE_TTL,
}
}
pub fn default_failure_ttl() -> Duration {
DEFAULT_FAILURE_TTL
}
pub async fn select_worker(&self, block_id: i64) -> Result<WorkerInfo> {
if self.workers.is_empty() {
return Err(Error::NoWorkerAvailable {
message: "no workers registered".to_string(),
});
}
self.cleanup_expired_failures();
if let Some(local_id) = self.local_worker_id {
if let Some(local_w) = self.workers.iter().find(|w| w.id == Some(local_id)) {
if let Some(addr) = &local_w.address {
if !self.is_failed(&worker_addr_key(addr)) {
return Ok(local_w.clone());
}
}
}
}
if let Some(w) =
consistent_hash_select_from_ring(block_id, &self.workers, &self.hash_ring, true, |k| {
self.is_failed(k)
})
{
return Ok(w);
}
consistent_hash_select_from_ring(block_id, &self.workers, &self.hash_ring, false, |k| {
self.is_failed(k)
})
.ok_or_else(|| Error::NoWorkerAvailable {
message: format!("no suitable worker for block_id={}", block_id),
})
}
pub async fn pick_any_worker(&self) -> Result<WorkerInfo> {
if self.workers.is_empty() {
return Err(Error::NoWorkerAvailable {
message: "no workers registered".to_string(),
});
}
self.cleanup_expired_failures();
let eligible: Vec<WorkerInfo> = self
.workers
.iter()
.filter(|w| {
if let Some(addr) = w.address.as_ref() {
let key = worker_addr_key(addr);
!self.is_failed(&key)
} else {
false
}
})
.cloned()
.collect();
let pool = if eligible.is_empty() {
(*self.workers).clone()
} else {
eligible
};
if pool.is_empty() {
return Err(Error::NoWorkerAvailable {
message: "no eligible workers".to_string(),
});
}
let idx = rand::Rng::random_range(&mut rand::rng(), 0..pool.len());
Ok(pool[idx].clone())
}
pub fn mark_failed(&self, addr: &WorkerNetAddress) {
let key = worker_addr_key(addr);
let map = self.failed_workers.get_or_init(new_failed_workers_map);
if map.insert(key, Instant::now()).is_none() {
self.failed_count.fetch_add(1, Ordering::Relaxed);
}
}
fn is_failed(&self, key: &str) -> bool {
let Some(map) = self.failed_workers.get() else {
return false;
};
if let Some(entry) = map.get(key) {
entry.value().elapsed() < self.failure_ttl
} else {
false
}
}
fn cleanup_expired_failures(&self) {
if self.failed_count.load(Ordering::Relaxed) == 0 {
return;
}
let Some(map) = self.failed_workers.get() else {
return;
};
let ttl = self.failure_ttl;
let mut removed: usize = 0;
map.retain(|_, v| {
if v.elapsed() < ttl {
true
} else {
removed += 1;
false
}
});
if removed > 0 {
self.failed_count.fetch_sub(removed, Ordering::Relaxed);
}
}
#[cfg(test)]
fn local_worker_id(&self) -> Option<i64> {
self.local_worker_id
}
#[cfg(test)]
fn workers_arc(&self) -> &Arc<Vec<WorkerInfo>> {
&self.workers
}
#[cfg(test)]
fn failed_workers_is_uninitialised(&self) -> bool {
self.failed_workers.get().is_none()
}
#[cfg(test)]
fn failed_workers_len(&self) -> usize {
self.failed_workers.get().map_or(0, |m| m.len())
}
#[cfg(test)]
fn hash_ring_arc(&self) -> &Arc<Vec<(u64, usize)>> {
&self.hash_ring
}
}
fn build_hash_ring(workers: &[WorkerInfo]) -> Vec<(u64, usize)> {
let mut ring: Vec<(u64, usize)> =
Vec::with_capacity(workers.len() * VIRTUAL_NODES_PER_WORKER as usize);
for (idx, worker) in workers.iter().enumerate() {
let worker_id = worker.id.unwrap_or(idx as i64);
let virtual_nodes = worker
.virtual_node_num
.unwrap_or(VIRTUAL_NODES_PER_WORKER as i32) as u32;
for vn in 0..virtual_nodes {
let hash = hash_virtual_node(worker_id, vn);
ring.push((hash, idx));
}
}
ring.sort_by_key(|(h, _)| *h);
ring
}
fn workers_fingerprint(workers: &[WorkerInfo]) -> u64 {
if workers.is_empty() {
return 0;
}
let mut tuples: Vec<(i64, &str, i32, i32)> = workers
.iter()
.map(|w| {
let addr = w.address.as_ref();
let host = addr.and_then(|a| a.host.as_deref()).unwrap_or("");
let port = addr.and_then(|a| a.rpc_port).unwrap_or(0);
let vn = w
.virtual_node_num
.unwrap_or(VIRTUAL_NODES_PER_WORKER as i32);
(w.id.unwrap_or(0), host, port, vn)
})
.collect();
tuples.sort_unstable();
let mut h = Xxh3Default::default();
for (id, host, port, vn) in tuples {
h.write(&id.to_le_bytes());
h.write(&(port as i32).to_le_bytes());
h.write(&(vn as i32).to_le_bytes());
h.write(&(host.len() as u32).to_le_bytes());
h.write(host.as_bytes());
h.write(&[0u8]);
}
h.finish()
}
impl Default for WorkerRouter {
fn default() -> Self {
Self::new()
}
}
pub(crate) fn worker_addr_key(addr: &WorkerNetAddress) -> String {
let host = addr.host.as_deref().unwrap_or("unknown");
let port = addr.rpc_port.unwrap_or(0);
let mut s = String::with_capacity(host.len() + 12);
s.push_str(host);
s.push(':');
let mut buf = itoa::Buffer::new();
s.push_str(buf.format(port));
s
}
pub(crate) fn rpc_endpoint(addr: &WorkerNetAddress) -> String {
let host = addr.host.as_deref().unwrap_or("127.0.0.1");
let port = addr.rpc_port.unwrap_or(9203);
let mut s = String::with_capacity(host.len() + 12);
s.push_str(host);
s.push(':');
let mut buf = itoa::Buffer::new();
s.push_str(buf.format(port));
s
}
#[inline]
fn hash_virtual_node(worker_id: i64, vn: u32) -> u64 {
let mut h = Xxh3Default::default();
h.write(&worker_id.to_le_bytes());
h.write(b":");
h.write(&vn.to_le_bytes());
h.finish()
}
#[inline]
fn hash_block_id(block_id: i64) -> u64 {
let mut h = Xxh3Default::default();
h.write(&block_id.to_le_bytes());
h.finish()
}
#[cfg(test)]
mod tests {
use super::*;
fn make_worker(id: i64, host: &str, port: i32) -> WorkerInfo {
WorkerInfo {
id: Some(id),
address: Some(WorkerNetAddress {
host: Some(host.to_string()),
rpc_port: Some(port),
..Default::default()
}),
..Default::default()
}
}
#[tokio::test]
async fn test_select_worker_empty() {
let router = WorkerRouter::new();
assert!(router.select_worker(123).await.is_err());
}
#[tokio::test]
async fn test_select_worker_deterministic() {
let router = WorkerRouter::new();
let workers = vec![
make_worker(1, "w1", 9203),
make_worker(2, "w2", 9203),
make_worker(3, "w3", 9203),
];
router.update_workers(workers).await;
let w1 = router.select_worker(42).await.unwrap();
let w2 = router.select_worker(42).await.unwrap();
assert_eq!(w1.id, w2.id);
}
#[tokio::test]
async fn test_failed_worker_filtered() {
let router = WorkerRouter::with_failure_ttl(Duration::from_secs(3600));
let workers = vec![make_worker(1, "w1", 9203), make_worker(2, "w2", 9203)];
router.update_workers(workers.clone()).await;
router.mark_failed(workers[0].address.as_ref().unwrap());
let selected = router.select_worker(42).await.unwrap();
assert_eq!(selected.id, Some(2));
}
#[tokio::test]
async fn test_pick_any_worker_empty() {
let router = WorkerRouter::new();
assert!(router.pick_any_worker().await.is_err());
}
#[tokio::test]
async fn test_pick_any_worker_returns_eligible() {
let router = WorkerRouter::with_failure_ttl(Duration::from_secs(3600));
let workers = vec![
make_worker(1, "w1", 9203),
make_worker(2, "w2", 9203),
make_worker(3, "w3", 9203),
];
router.update_workers(workers.clone()).await;
router.mark_failed(workers[0].address.as_ref().unwrap());
router.mark_failed(workers[1].address.as_ref().unwrap());
for _ in 0..10 {
let picked = router.pick_any_worker().await.unwrap();
assert_eq!(picked.id, Some(3));
}
}
#[tokio::test]
async fn test_pick_any_worker_fallback_when_all_failed() {
let router = WorkerRouter::with_failure_ttl(Duration::from_secs(3600));
let workers = vec![make_worker(1, "w1", 9203), make_worker(2, "w2", 9203)];
router.update_workers(workers.clone()).await;
router.mark_failed(workers[0].address.as_ref().unwrap());
router.mark_failed(workers[1].address.as_ref().unwrap());
let picked = router.pick_any_worker().await.unwrap();
assert!(picked.id == Some(1) || picked.id == Some(2));
}
#[tokio::test]
async fn test_needs_refresh_false_after_new() {
let router = WorkerRouter::new();
assert!(!router.needs_refresh().await);
}
#[tokio::test]
async fn test_needs_refresh_true_with_zero_ttl() {
let router = WorkerRouter::with_ttls(DEFAULT_FAILURE_TTL, Duration::ZERO);
tokio::time::sleep(Duration::from_millis(1)).await;
assert!(router.needs_refresh().await);
}
#[tokio::test]
async fn test_with_ttls_stores_values() {
let failure = Duration::from_secs(10);
let refresh = Duration::from_secs(5);
let router = WorkerRouter::with_ttls(failure, refresh);
assert_eq!(router.failure_ttl, failure);
assert_eq!(router.worker_refresh_ttl, refresh);
}
#[tokio::test]
async fn test_update_workers_resets_refresh_clock() {
let router = WorkerRouter::with_ttls(DEFAULT_FAILURE_TTL, Duration::ZERO);
tokio::time::sleep(Duration::from_millis(1)).await;
assert!(
router.needs_refresh().await,
"should need refresh before update"
);
let router2 = WorkerRouter::with_ttls(DEFAULT_FAILURE_TTL, Duration::from_secs(60));
router2
.update_workers(vec![make_worker(1, "w1", 9203)])
.await;
assert!(!router2.needs_refresh().await);
}
#[tokio::test]
async fn test_local_worker_preferred() {
let router = WorkerRouter::new();
let workers = vec![
make_worker(1, "remote1", 9203),
make_worker(2, "localhost", 9203), make_worker(3, "remote2", 9203),
];
router.update_workers(workers).await;
for block_id in [1i64, 42, 100, 999, 10_000] {
let selected = router.select_worker(block_id).await.unwrap();
assert_eq!(
selected.id,
Some(2),
"block_id={} should route to local worker",
block_id
);
}
}
#[tokio::test]
async fn test_local_worker_skipped_when_failed() {
let router = WorkerRouter::with_failure_ttl(Duration::from_secs(3600));
let local_worker = make_worker(2, "localhost", 9203);
let workers = vec![
make_worker(1, "remote1", 9203),
local_worker.clone(),
make_worker(3, "remote2", 9203),
];
router.update_workers(workers).await;
router.mark_failed(local_worker.address.as_ref().unwrap());
let selected = router.select_worker(42).await.unwrap();
assert_ne!(
selected.id,
Some(2),
"failed local worker should not be selected"
);
}
#[tokio::test]
async fn test_detect_local_worker_none() {
let workers = vec![
make_worker(1, "remote-host-a.example.com", 9203),
make_worker(2, "remote-host-b.example.com", 9203),
];
let id = WorkerRouter::detect_local_worker(&workers).await;
assert_eq!(id, 0);
}
#[tokio::test]
async fn test_detect_local_worker_loopback() {
let workers = vec![
make_worker(1, "10.0.0.1", 9203),
make_worker(2, "127.0.0.1", 9203),
];
let id = WorkerRouter::detect_local_worker(&workers).await;
assert_eq!(id, 2);
}
#[tokio::test]
async fn test_local_worker_cache_invalidated_on_update() {
let router = WorkerRouter::new();
router
.update_workers(vec![make_worker(1, "remote1", 9203)])
.await;
assert_eq!(**router.local_worker_id.load(), Some(None));
let _ = router.select_worker(1).await;
assert_eq!(**router.local_worker_id.load(), Some(None));
router
.update_workers(vec![
make_worker(1, "remote1", 9203),
make_worker(2, "127.0.0.1", 9203),
])
.await;
assert_eq!(**router.local_worker_id.load(), Some(Some(2)));
let selected = router.select_worker(1).await.unwrap();
assert_eq!(selected.id, Some(2), "new local worker should be preferred");
}
#[tokio::test]
async fn test_update_workers_leaves_local_worker_id_probed() {
let router = WorkerRouter::new();
router
.update_workers(vec![
make_worker(1, "remote-a.example.com", 9203),
make_worker(2, "remote-b.example.com", 9203),
])
.await;
match **router.local_worker_id.load() {
Some(_) => {}
None => panic!("update_workers must leave local_worker_id in the probed state"),
}
assert_eq!(
**router.local_worker_id.load(),
Some(None),
"no local worker → cache must be Some(None), not None (unprobed)"
);
router
.update_workers(vec![
make_worker(1, "remote-a.example.com", 9203),
make_worker(2, "127.0.0.1", 9203),
])
.await;
assert_eq!(
**router.local_worker_id.load(),
Some(Some(2)),
"local worker present → cache must be Some(Some(id))"
);
router
.update_workers(vec![make_worker(3, "remote-c.example.com", 9203)])
.await;
assert!(
(**router.local_worker_id.load()).is_some(),
"post-update state must always be probed"
);
}
#[tokio::test]
async fn test_snapshot_from_shares_hash_ring_arc() {
let shared = WorkerRouter::new();
shared
.update_workers(vec![
make_worker(1, "w1", 9203),
make_worker(2, "w2", 9203),
make_worker(3, "w3", 9203),
])
.await;
let snap = WorkerRouter::snapshot_from(&shared);
let a = shared.hash_ring.load_full();
let b = snap.hash_ring.load_full();
assert!(
Arc::ptr_eq(&a, &b),
"snapshot must reuse the shared hash_ring Arc (no rebuild)"
);
let wa = shared.workers.load_full();
let wb = snap.workers.load_full();
assert!(
Arc::ptr_eq(&wa, &wb),
"snapshot must reuse the shared workers Arc"
);
}
#[tokio::test]
async fn test_snapshot_from_shares_local_worker_id() {
let shared = WorkerRouter::new();
shared
.update_workers(vec![make_worker(1, "w1", 9203), make_worker(2, "w2", 9203)])
.await;
let snap_unprobed = WorkerRouter::snapshot_from(&shared);
let a = shared.local_worker_id.load_full();
let b = snap_unprobed.local_worker_id.load_full();
assert!(
Arc::ptr_eq(&a, &b),
"snapshot must reuse the shared local_worker_id Arc (unprobed case)"
);
shared.local_worker_id.store(Arc::new(Some(None)));
let snap_probed = WorkerRouter::snapshot_from(&shared);
let pa = shared.local_worker_id.load_full();
let pb = snap_probed.local_worker_id.load_full();
assert!(
Arc::ptr_eq(&pa, &pb),
"snapshot must reuse the shared local_worker_id Arc (probed case)"
);
assert_eq!(
*pb,
Some(None),
"snapshot must observe the probed value, not re-probe"
);
}
#[tokio::test]
async fn test_cleanup_expired_failures_empty_is_noop() {
let router = WorkerRouter::new();
router
.update_workers(vec![make_worker(1, "w1", 9203)])
.await;
assert!(router.failed_workers_is_uninitialised());
router.cleanup_expired_failures();
assert!(
router.failed_workers_is_uninitialised(),
"cleanup on empty map must be a no-op (map not allocated)"
);
let worker = make_worker(1, "w1", 9203);
router.mark_failed(worker.address.as_ref().unwrap());
assert_eq!(router.failed_workers_len(), 1);
router.cleanup_expired_failures();
assert_eq!(
router.failed_workers_len(),
1,
"cleanup must not evict non-expired entries"
);
}
#[tokio::test]
async fn test_snapshot_from_select_worker_matches_parent() {
let shared = WorkerRouter::new();
shared
.update_workers(vec![
make_worker(1, "w1", 9203),
make_worker(2, "w2", 9203),
make_worker(3, "w3", 9203),
])
.await;
let snap = WorkerRouter::snapshot_from(&shared);
for block_id in [0i64, 1, 42, 999, 1 << 30, i64::MAX] {
let a = shared.select_worker(block_id).await.unwrap();
let b = snap.select_worker(block_id).await.unwrap();
assert_eq!(
a.id, b.id,
"snapshot and parent must select the same worker for block_id={}",
block_id
);
}
}
#[tokio::test]
async fn test_snapshot_mark_failed_is_isolated_from_parent() {
let shared = WorkerRouter::with_failure_ttl(Duration::from_secs(3600));
let workers = vec![make_worker(1, "w1", 9203), make_worker(2, "w2", 9203)];
shared.update_workers(workers.clone()).await;
let snap = WorkerRouter::snapshot_from(&shared);
snap.mark_failed(workers[0].address.as_ref().unwrap());
assert!(!shared.is_failed(&worker_addr_key(workers[0].address.as_ref().unwrap())));
assert!(snap.is_failed(&worker_addr_key(workers[0].address.as_ref().unwrap())));
}
#[tokio::test]
async fn test_update_workers_fingerprint_skip_rebuild() {
let router = WorkerRouter::new();
let workers = vec![
make_worker(1, "w1", 9203),
make_worker(2, "w2", 9203),
make_worker(3, "w3", 9203),
];
router.update_workers(workers.clone()).await;
let ring_before = router.hash_ring.load_full();
router.update_workers(workers.clone()).await;
let ring_same = router.hash_ring.load_full();
assert!(
Arc::ptr_eq(&ring_before, &ring_same),
"identical worker set must not rebuild the ring"
);
let mut reordered = workers.clone();
reordered.reverse();
router.update_workers(reordered).await;
let ring_reordered = router.hash_ring.load_full();
assert!(
Arc::ptr_eq(&ring_before, &ring_reordered),
"reordered-but-identical worker set must not rebuild the ring"
);
let changed = vec![
make_worker(1, "w1", 9203),
make_worker(2, "w2", 9203),
make_worker(4, "w4", 9203), ];
router.update_workers(changed).await;
let ring_after = router.hash_ring.load_full();
assert!(
!Arc::ptr_eq(&ring_before, &ring_after),
"changed worker set must rebuild the ring"
);
}
#[tokio::test]
async fn test_cleanup_expired_failures_counter_stays_in_sync() {
let ttl = Duration::from_millis(30);
let router = WorkerRouter::with_failure_ttl(ttl);
let a = make_worker(1, "w1", 9203);
let b = make_worker(2, "w2", 9203);
let c = make_worker(3, "w3", 9203);
router.mark_failed(a.address.as_ref().unwrap());
router.mark_failed(b.address.as_ref().unwrap());
router.mark_failed(c.address.as_ref().unwrap());
assert_eq!(router.failed_count.load(Ordering::Relaxed), 3);
assert_eq!(router.failed_workers_len(), 3);
router.mark_failed(a.address.as_ref().unwrap());
router.mark_failed(a.address.as_ref().unwrap());
assert_eq!(
router.failed_count.load(Ordering::Relaxed),
3,
"re-insert must not touch the counter"
);
let healthy = WorkerRouter::new();
healthy.cleanup_expired_failures();
assert_eq!(healthy.failed_count.load(Ordering::Relaxed), 0);
assert!(healthy.failed_workers_is_uninitialised());
tokio::time::sleep(ttl + Duration::from_millis(20)).await;
router.cleanup_expired_failures();
assert_eq!(
router.failed_workers_len(),
0,
"expired entries must be removed"
);
assert_eq!(
router.failed_count.load(Ordering::Relaxed),
0,
"counter must be decremented by exactly the number of removals"
);
router.cleanup_expired_failures();
assert_eq!(router.failed_count.load(Ordering::Relaxed), 0);
}
#[tokio::test]
async fn test_snapshot_failed_count_starts_fresh() {
let shared = WorkerRouter::with_failure_ttl(Duration::from_secs(3600));
shared
.update_workers(vec![make_worker(1, "w1", 9203)])
.await;
shared.mark_failed(&WorkerNetAddress {
host: Some("w1".to_string()),
rpc_port: Some(9203),
..Default::default()
});
assert_eq!(shared.failed_count.load(Ordering::Relaxed), 1);
let snap = WorkerRouter::snapshot_from(&shared);
assert_eq!(
snap.failed_count.load(Ordering::Relaxed),
0,
"snapshot must not inherit parent's failed_count"
);
assert!(snap.failed_workers_is_uninitialised());
}
#[test]
fn test_hash_functions_are_stable() {
assert_eq!(hash_virtual_node(42, 7), hash_virtual_node(42, 7));
assert_eq!(hash_block_id(1234567890), hash_block_id(1234567890));
assert_ne!(hash_virtual_node(42, 0), hash_block_id(42));
}
#[tokio::test]
async fn test_view_failed_workers_is_lazy_initialised() {
let shared = WorkerRouter::new();
shared
.update_workers(vec![make_worker(1, "w1", 9203), make_worker(2, "w2", 9203)])
.await;
let view = WorkerRouterView::from_shared(&shared);
for i in 0..1000 {
let _ = view.select_worker(i).await;
}
assert!(
view.failed_workers_is_uninitialised(),
"OnceLock must stay uninitialised on the happy path (no mark_failed calls)"
);
assert!(
shared.failed_workers_is_uninitialised(),
"shared router must also stay uninitialised on the happy path"
);
}
#[tokio::test]
async fn test_view_mark_failed_init_dashmap_lazily() {
let shared = WorkerRouter::new();
let workers = vec![make_worker(1, "w1", 9203), make_worker(2, "w2", 9203)];
shared.update_workers(workers.clone()).await;
let view = WorkerRouterView::from_shared(&shared);
assert!(view.failed_workers_is_uninitialised());
view.mark_failed(workers[0].address.as_ref().unwrap());
assert!(!view.failed_workers_is_uninitialised());
assert_eq!(view.failed_workers_len(), 1);
assert_eq!(view.failed_count.load(Ordering::Relaxed), 1);
let key = worker_addr_key(workers[0].address.as_ref().unwrap());
assert!(view.is_failed(&key));
view.mark_failed(workers[1].address.as_ref().unwrap());
assert_eq!(view.failed_workers_len(), 2);
assert_eq!(view.failed_count.load(Ordering::Relaxed), 2);
}
#[tokio::test]
async fn test_view_from_shared_shares_hash_ring_arc() {
let shared = WorkerRouter::new();
shared
.update_workers(vec![
make_worker(1, "w1", 9203),
make_worker(2, "w2", 9203),
make_worker(3, "w3", 9203),
])
.await;
let workers_arc_before = shared.workers.load_full();
let ring_arc_before = shared.hash_ring.load_full();
let view = WorkerRouterView::from_shared(&shared);
assert!(
Arc::ptr_eq(&workers_arc_before, view.workers_arc()),
"view must share the shared router's `workers` Arc (no re-allocation)"
);
assert!(
Arc::ptr_eq(&ring_arc_before, view.hash_ring_arc()),
"view must share the shared router's `hash_ring` Arc (no re-build)"
);
}
#[tokio::test]
async fn test_view_from_shared_inherits_local_worker_id() {
let shared_a = WorkerRouter::new();
assert!((**shared_a.local_worker_id.load()).is_none());
let view_a = WorkerRouterView::from_shared(&shared_a);
assert_eq!(
view_a.local_worker_id(),
None,
"unprobed parent → view captures None"
);
let shared_b = WorkerRouter::new();
shared_b
.update_workers(vec![
make_worker(1, "remote-a.example.com", 9203),
make_worker(2, "remote-b.example.com", 9203),
])
.await;
assert_eq!(**shared_b.local_worker_id.load(), Some(None));
let view_b = WorkerRouterView::from_shared(&shared_b);
assert_eq!(
view_b.local_worker_id(),
None,
"probed-no-local parent → view captures None"
);
let shared_c = WorkerRouter::new();
shared_c
.update_workers(vec![
make_worker(1, "remote-a.example.com", 9203),
make_worker(2, "127.0.0.1", 9203),
])
.await;
assert_eq!(**shared_c.local_worker_id.load(), Some(Some(2)));
let view_c = WorkerRouterView::from_shared(&shared_c);
assert_eq!(
view_c.local_worker_id(),
Some(2),
"probed-local-present parent → view captures Some(id)"
);
let shared_d = WorkerRouter::new();
let view_d = WorkerRouterView::from_shared(&shared_d);
shared_d
.update_workers(vec![make_worker(1, "127.0.0.1", 9203)])
.await;
assert!(matches!(
**shared_d.local_worker_id.load(),
Some(Some(1)) | Some(None)
));
assert_eq!(
view_d.local_worker_id(),
None,
"view minted before probe must keep its captured None — parent's later probe does NOT retroactively update the view"
);
}
#[tokio::test]
async fn test_view_from_workers_builds_hash_ring() {
let workers = vec![
make_worker(1, "w1", 9203),
make_worker(2, "w2", 9203),
make_worker(3, "w3", 9203),
];
let view = WorkerRouterView::from_workers(workers, Duration::from_secs(60));
assert_eq!(
view.hash_ring_arc().len(),
3 * VIRTUAL_NODES_PER_WORKER as usize
);
let a = view.select_worker(42).await.unwrap();
let b = view.select_worker(42).await.unwrap();
assert_eq!(a.id, b.id, "same block_id must select the same worker");
assert_eq!(
view.local_worker_id(),
None,
"from_workers must NOT run detect_local_worker (blocking syscall on legacy path)"
);
}
#[tokio::test]
async fn test_view_select_worker_matches_shared_for_all_block_ids() {
let shared = WorkerRouter::new();
shared
.update_workers(vec![
make_worker(1, "w1", 9203),
make_worker(2, "w2", 9203),
make_worker(3, "w3", 9203),
make_worker(4, "w4", 9203),
make_worker(5, "w5", 9203),
])
.await;
let snap = WorkerRouter::snapshot_from(&shared);
let view = WorkerRouterView::from_shared(&shared);
for block_id in [
0i64,
1,
-1,
42,
999,
1 << 20,
1 << 30,
i64::MAX,
i64::MIN,
0x7fff_ffff_ffff_ffff,
] {
let a = shared.select_worker(block_id).await.unwrap();
let b = snap.select_worker(block_id).await.unwrap();
let c = view.select_worker(block_id).await.unwrap();
assert_eq!(
a.id, b.id,
"snapshot must match shared for block_id={}",
block_id
);
assert_eq!(
a.id, c.id,
"view must match shared for block_id={} (A/B parity)",
block_id
);
}
}
#[tokio::test]
async fn test_view_mark_failed_is_isolated_from_shared() {
let shared = WorkerRouter::with_failure_ttl(Duration::from_secs(3600));
let workers = vec![make_worker(1, "w1", 9203), make_worker(2, "w2", 9203)];
shared.update_workers(workers.clone()).await;
let view = WorkerRouterView::from_shared(&shared);
view.mark_failed(workers[0].address.as_ref().unwrap());
assert!(!shared.is_failed(&worker_addr_key(workers[0].address.as_ref().unwrap())));
assert!(view.is_failed(&worker_addr_key(workers[0].address.as_ref().unwrap())));
let view2 = WorkerRouterView::from_shared(&shared);
assert!(!view2.is_failed(&worker_addr_key(workers[0].address.as_ref().unwrap())));
}
#[tokio::test]
async fn test_view_pick_any_worker_semantics() {
let shared_empty = WorkerRouter::new();
let view_empty = WorkerRouterView::from_shared(&shared_empty);
assert!(view_empty.pick_any_worker().await.is_err());
let shared = WorkerRouter::new();
shared
.update_workers(vec![make_worker(1, "w1", 9203), make_worker(2, "w2", 9203)])
.await;
let view = WorkerRouterView::from_shared(&shared);
let picked = view.pick_any_worker().await.unwrap();
assert!(matches!(picked.id, Some(1) | Some(2)));
let w1_addr = shared.workers.load_full()[0].address.clone().unwrap();
let w2_addr = shared.workers.load_full()[1].address.clone().unwrap();
view.mark_failed(&w1_addr);
view.mark_failed(&w2_addr);
let picked_after_fail = view.pick_any_worker().await.unwrap();
assert!(matches!(picked_after_fail.id, Some(1) | Some(2)));
}
#[tokio::test]
async fn test_view_from_workers_no_local_first_when_not_probed() {
let workers = vec![
make_worker(1, "remote-a.example.com", 9203),
make_worker(2, "127.0.0.1", 9203), ];
let view = WorkerRouterView::from_workers(workers, Duration::from_secs(60));
assert_eq!(view.local_worker_id(), None);
let a = view.select_worker(0).await.unwrap();
let b = view.select_worker(0).await.unwrap();
assert_eq!(a.id, b.id);
}
#[tokio::test]
async fn test_view_empty_matches_worker_router_new_semantics() {
let view = WorkerRouterView::empty();
assert!(matches!(
view.select_worker(0).await,
Err(Error::NoWorkerAvailable { .. })
));
assert!(matches!(
view.select_worker(i64::MAX).await,
Err(Error::NoWorkerAvailable { .. })
));
assert!(matches!(
view.pick_any_worker().await,
Err(Error::NoWorkerAvailable { .. })
));
assert_eq!(view.local_worker_id(), None);
view.mark_failed(&WorkerNetAddress {
host: Some("unknown".to_string()),
rpc_port: Some(9203),
..Default::default()
});
assert_eq!(view.failed_count.load(Ordering::Relaxed), 1);
assert_eq!(
WorkerRouterView::default_failure_ttl(),
DEFAULT_FAILURE_TTL,
"public default TTL must equal the shared router's private DEFAULT_FAILURE_TTL"
);
}
}