use std::sync::Arc;
use praxis_core::health::{ClusterHealthState, EndpointHealth};
use super::endpoint::WeightedEndpoint;
const TABLE_SIZE: usize = 65_537;
const SKIP_SEED: u64 = 0x9E37_79B9_7F4A_7C15;
const SENTINEL: u32 = u32::MAX;
pub(crate) struct Maglev {
endpoints: Vec<WeightedEndpoint>,
header: Option<String>,
table: Vec<u32>,
}
impl Maglev {
pub(crate) fn new(endpoints: Vec<WeightedEndpoint>, header: Option<String>) -> Self {
let table = build_table(&endpoints);
Self {
endpoints,
header,
table,
}
}
pub(crate) fn header(&self) -> Option<&str> {
self.header.as_deref()
}
pub(crate) fn select(&self, hash_key: Option<&str>, health: Option<&ClusterHealthState>) -> Option<Arc<str>> {
let len = self.table.len();
if len == 0 {
return None;
}
let key = hash_key.unwrap_or("");
#[expect(clippy::cast_possible_truncation, reason = "modulo fits usize")]
let start = (fnv1a_seeded(key, 0) as usize) % len;
if let Some(state) = health {
for offset in 0..len {
let ep = self.endpoint_at((start + offset) % len);
if state.endpoints().get(ep.index).is_some_and(EndpointHealth::is_healthy) {
return Some(Arc::clone(&ep.address));
}
}
}
Some(Arc::clone(&self.endpoint_at(start).address))
}
#[expect(
clippy::indexing_slicing,
reason = "table slot and owner index are in bounds by construction"
)]
fn endpoint_at(&self, slot: usize) -> &WeightedEndpoint {
&self.endpoints[self.table[slot] as usize]
}
}
struct Replica {
offset: usize,
skip: usize,
owner: u32,
}
fn build_replicas(endpoints: &[WeightedEndpoint]) -> Vec<Replica> {
let mut replicas = Vec::new();
for (idx, ep) in endpoints.iter().enumerate() {
for replica in 0..ep.weight {
let key = format!("{}#{replica}", ep.address);
let offset = usize::try_from(fnv1a_seeded(&key, 0) % TABLE_SIZE as u64).unwrap_or(0);
let skip = usize::try_from(fnv1a_seeded(&key, SKIP_SEED) % (TABLE_SIZE as u64 - 1)).unwrap_or(0) + 1;
#[expect(clippy::cast_possible_truncation, reason = "endpoint count fits u32")]
let owner = idx as u32;
replicas.push(Replica { offset, skip, owner });
}
}
replicas
}
#[expect(clippy::indexing_slicing, reason = "table index is a modulo of its length")]
fn build_table(endpoints: &[WeightedEndpoint]) -> Vec<u32> {
let replicas = build_replicas(endpoints);
if replicas.is_empty() {
return Vec::new();
}
let mut table = vec![SENTINEL; TABLE_SIZE];
let mut cursors = vec![0_usize; replicas.len()];
let mut filled = 0_usize;
loop {
for (r, cursor) in replicas.iter().zip(cursors.iter_mut()) {
let mut c = (r.offset + *cursor * r.skip) % TABLE_SIZE;
while table[c] != SENTINEL {
*cursor += 1;
c = (r.offset + *cursor * r.skip) % TABLE_SIZE;
}
table[c] = r.owner;
*cursor += 1;
filled += 1;
if filled == TABLE_SIZE {
debug_assert!(!table.contains(&SENTINEL), "maglev table must be fully populated");
return table;
}
}
}
}
fn fnv1a_seeded(s: &str, seed: u64) -> u64 {
let mut hash: u64 = 0xCBF2_9CE4_8422_2325 ^ seed;
for byte in s.bytes() {
hash ^= u64::from(byte);
hash = hash.wrapping_mul(0x0000_0100_0000_01B3);
}
hash
}
#[cfg(test)]
#[expect(clippy::allow_attributes, reason = "blanket test suppressions")]
#[allow(
clippy::unwrap_used,
clippy::expect_used,
clippy::indexing_slicing,
clippy::panic,
clippy::too_many_lines,
clippy::cast_precision_loss,
clippy::cast_possible_truncation,
clippy::cast_sign_loss,
clippy::cast_lossless,
reason = "tests"
)]
mod tests {
use std::collections::{HashMap, HashSet};
use praxis_core::health::ClusterHealthEntry;
use super::*;
#[test]
fn same_key_same_endpoint() {
let mg = Maglev::new(endpoints(3), None);
let first = mg.select(Some("/stable"), None).unwrap();
let second = mg.select(Some("/stable"), None).unwrap();
assert_eq!(first, second, "same key should always select same endpoint");
}
#[test]
fn different_keys_reach_all_endpoints() {
let mg = Maglev::new(endpoints(2), None);
let selections: HashSet<Arc<str>> = (0..100)
.map(|i| mg.select(Some(&format!("/k{i}")), None).unwrap())
.collect();
assert_eq!(
selections.len(),
2,
"distinct keys should reach both endpoints across many keys"
);
}
#[test]
fn distribution_is_even() {
let n = 5;
let mg = Maglev::new(endpoints(n), None);
let mut counts: HashMap<Arc<str>, usize> = HashMap::new();
let total = 10_000;
for i in 0..total {
let sel = mg.select(Some(&format!("/key-{i}")), None).unwrap();
*counts.entry(sel).or_default() += 1;
}
let expected = total as f64 / n as f64;
for (addr, count) in &counts {
let ratio = *count as f64 / expected;
assert!(
(0.85..=1.15).contains(&ratio),
"endpoint {addr} share {ratio:.3} should be near 1.0 (count={count}, expected={expected})"
);
}
assert_eq!(counts.len(), n, "all endpoints should receive traffic");
}
#[test]
fn skips_unhealthy() {
let mg = Maglev::new(endpoints(3), None);
let state = health_state(&["10.0.0.1:80", "10.0.0.2:80", "10.0.0.3:80"]);
state.endpoints()[1].mark_unhealthy();
for i in 0..100 {
let sel = mg.select(Some(&format!("/k{i}")), Some(&state)).unwrap();
assert_ne!(&*sel, "10.0.0.2:80", "unhealthy endpoint must never be selected");
}
}
#[test]
fn panic_mode_when_all_unhealthy() {
let mg = Maglev::new(endpoints(2), None);
let state = health_state(&["10.0.0.1:80", "10.0.0.2:80"]);
state.endpoints()[0].mark_unhealthy();
state.endpoints()[1].mark_unhealthy();
let sel = mg.select(Some("/panic"), Some(&state)).unwrap();
assert!(
&*sel == "10.0.0.1:80" || &*sel == "10.0.0.2:80",
"panic mode should still return an endpoint, got: {sel}"
);
}
#[test]
fn select_with_none_hash_key_uses_fallback() {
let mg = Maglev::new(endpoints(3), None);
let first = mg.select(None, None).unwrap();
for _ in 0..10 {
assert_eq!(first, mg.select(None, None).unwrap(), "None key must be deterministic");
}
}
#[test]
fn weight_stability() {
let eps = vec![
WeightedEndpoint {
address: Arc::from("10.0.0.1:80"),
index: 0,
weight: 3,
},
WeightedEndpoint {
address: Arc::from("10.0.0.2:80"),
index: 1,
weight: 1,
},
];
let mg = Maglev::new(eps, None);
let total = 4_000;
let mut ep1 = 0_usize;
for i in 0..total {
let key = format!("/w-{i}");
let sel = mg.select(Some(&key), None).unwrap();
assert_eq!(sel, mg.select(Some(&key), None).unwrap(), "must be deterministic");
if &*sel == "10.0.0.1:80" {
ep1 += 1;
}
}
let ratio = ep1 as f64 / total as f64;
assert!(
(ratio - 0.75).abs() < 0.05,
"weight-3 endpoint share {ratio:.3} should be near 0.75"
);
}
#[test]
fn minimal_disruption_on_backend_removal() {
let four = Maglev::new(endpoints(4), None);
let keys: Vec<String> = (0..10_000).map(|i| format!("/k-{i}")).collect();
let before: Vec<Arc<str>> = keys.iter().map(|k| four.select(Some(k), None).unwrap()).collect();
let three = Maglev::new(endpoints(3), None);
let dropped: Arc<str> = Arc::from("10.0.0.4:80");
let mut survivors = 0_usize;
let mut reassigned = 0_usize;
for (k, prev) in keys.iter().zip(&before) {
if *prev == dropped {
continue; }
survivors += 1;
if four.select(Some(k), None).unwrap() != three.select(Some(k), None).unwrap() {
reassigned += 1;
}
}
let churn = reassigned as f64 / survivors as f64;
assert!(
churn < 0.10,
"Maglev should reassign <10% of surviving keys on removal, got {churn:.3}"
);
}
#[test]
fn single_endpoint_owns_every_key() {
let mg = Maglev::new(endpoints(1), None);
for i in 0..50 {
let sel = mg.select(Some(&format!("/k{i}")), None).unwrap();
assert_eq!(&*sel, "10.0.0.1:80", "a single endpoint must own every key");
}
assert!(
!mg.table.contains(&SENTINEL),
"table must be fully populated with one endpoint"
);
}
#[test]
fn minimal_disruption_on_backend_addition() {
let three = Maglev::new(endpoints(3), None);
let four = Maglev::new(endpoints(4), None);
let added: Arc<str> = Arc::from("10.0.0.4:80");
let mut stayed_existing = 0_usize;
let mut reassigned = 0_usize;
for i in 0..10_000 {
let k = format!("/k-{i}");
let before = three.select(Some(&k), None).unwrap();
let after = four.select(Some(&k), None).unwrap();
if after == added {
continue; }
stayed_existing += 1;
if before != after {
reassigned += 1;
}
}
let churn = reassigned as f64 / stayed_existing as f64;
assert!(
churn < 0.10,
"adding a backend should not reshuffle keys among existing backends, got {churn:.3}"
);
}
#[test]
fn empty_endpoints_returns_none() {
let mg = Maglev::new(Vec::new(), None);
assert!(mg.select(Some("/x"), None).is_none(), "no endpoints should yield None");
}
#[test]
fn table_is_fully_populated() {
let mg = Maglev::new(endpoints(3), None);
assert_eq!(mg.table.len(), TABLE_SIZE, "table must be full size");
assert!(!mg.table.contains(&SENTINEL), "no slot should remain unfilled");
for idx in 0..3_u32 {
assert!(mg.table.contains(&idx), "endpoint {idx} should appear in the table");
}
}
fn endpoints(n: usize) -> Vec<WeightedEndpoint> {
(0..n)
.map(|i| WeightedEndpoint {
address: Arc::from(format!("10.0.0.{}:80", i + 1).as_str()),
index: i,
weight: 1,
})
.collect()
}
fn health_state(addrs: &[&str]) -> ClusterHealthState {
Arc::new(ClusterHealthEntry::new(
addrs.iter().map(|_| EndpointHealth::new()).collect(),
addrs.iter().map(|a| Arc::from(*a)).collect(),
None,
None,
))
}
}