use serde::{Deserialize, Serialize};
use std::collections::{BTreeMap, HashMap};
fn fnv1a_hash(data: &str) -> u64 {
const FNV_OFFSET_BASIS: u64 = 0xcbf29ce484222325;
const FNV_PRIME: u64 = 0x100000001b3;
let mut hash = FNV_OFFSET_BASIS;
for &byte in data.as_bytes() {
hash ^= byte as u64;
hash = hash.wrapping_mul(FNV_PRIME);
}
hash ^= hash >> 33;
hash = hash.wrapping_mul(0xff51afd7ed558ccd);
hash ^= hash >> 33;
hash = hash.wrapping_mul(0xc4ceb9fe1a85ec53);
hash ^= hash >> 33;
hash
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum EnhancedShardingError {
NoNodes,
NoGroupMatch(String),
NoListMatch(String),
}
impl std::fmt::Display for EnhancedShardingError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
EnhancedShardingError::NoNodes => write!(f, "no nodes configured"),
EnhancedShardingError::NoGroupMatch(key) => {
write!(f, "no group matches key: {}", key)
}
EnhancedShardingError::NoListMatch(key) => {
write!(f, "no list mapping for key: {}", key)
}
}
}
}
impl std::error::Error for EnhancedShardingError {}
pub struct ConsistentHashRouter {
ring: BTreeMap<u64, String>,
nodes: Vec<String>,
vnodes_per_node: usize,
}
impl ConsistentHashRouter {
pub fn new(nodes: Vec<&str>, vnodes_per_node: usize) -> Self {
let vnodes_per_node = vnodes_per_node.max(1);
let mut router = Self {
ring: BTreeMap::new(),
nodes: nodes.into_iter().map(|s| s.to_string()).collect(),
vnodes_per_node,
};
for node in &router.nodes {
for i in 0..vnodes_per_node {
let vnode_key = format!("{}#{}", node, i);
let hash = hash_str(&vnode_key);
ring_insert(&mut router.ring, hash, node.clone());
}
}
router
}
pub fn add_node(&mut self, node: &str) {
if self.nodes.iter().any(|n| n == node) {
return;
}
for i in 0..self.vnodes_per_node {
let vnode_key = format!("{}#{}", node, i);
let hash = hash_str(&vnode_key);
ring_insert(&mut self.ring, hash, node.to_string());
}
self.nodes.push(node.to_string());
}
pub fn remove_node(&mut self, node: &str) {
self.nodes.retain(|n| n != node);
let to_remove: Vec<u64> = self
.ring
.iter()
.filter(|(_, v)| *v == node)
.map(|(k, _)| *k)
.collect();
for k in to_remove {
self.ring.remove(&k);
}
}
pub fn route(&self, key: &str) -> Result<String, EnhancedShardingError> {
if self.ring.is_empty() {
return Err(EnhancedShardingError::NoNodes);
}
let hash = hash_str(key);
let node = self
.ring
.range(hash..)
.next()
.or_else(|| self.ring.iter().next())
.map(|(_, v)| v.clone())
.expect("ring is non-empty (checked above)");
Ok(node)
}
pub fn nodes(&self) -> &[String] {
&self.nodes
}
pub fn ring_size(&self) -> usize {
self.ring.len()
}
pub fn vnodes_per_node(&self) -> usize {
self.vnodes_per_node
}
pub fn node_ownership(&self, _node: &str) -> f64 {
if self.nodes.is_empty() {
return 0.0;
}
1.0 / self.nodes.len() as f64
}
}
impl std::fmt::Debug for ConsistentHashRouter {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ConsistentHashRouter")
.field("nodes", &self.nodes)
.field("vnodes_per_node", &self.vnodes_per_node)
.field("ring_size", &self.ring.len())
.finish()
}
}
pub struct ListRouter {
mapping: HashMap<String, String>,
default: Option<String>,
}
impl ListRouter {
pub fn new() -> Self {
Self {
mapping: HashMap::new(),
default: None,
}
}
pub fn add(mut self, key: &str, shard: &str) -> Self {
self.mapping.insert(key.to_string(), shard.to_string());
self
}
pub fn with_default(mut self, shard: &str) -> Self {
self.default = Some(shard.to_string());
self
}
pub fn route(&self, key: &str) -> Result<String, EnhancedShardingError> {
if let Some(shard) = self.mapping.get(key) {
return Ok(shard.clone());
}
if let Some(default) = &self.default {
return Ok(default.clone());
}
Err(EnhancedShardingError::NoListMatch(key.to_string()))
}
pub fn len(&self) -> usize {
self.mapping.len()
}
pub fn is_empty(&self) -> bool {
self.mapping.is_empty()
}
pub fn has_default(&self) -> bool {
self.default.is_some()
}
}
impl Default for ListRouter {
fn default() -> Self {
Self::new()
}
}
impl std::fmt::Debug for ListRouter {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ListRouter")
.field("mapping_size", &self.mapping.len())
.field("default", &self.default)
.finish()
}
}
#[derive(Debug, Clone)]
pub struct ShardGroup {
pub group_id: String,
pub shards: Vec<String>,
}
impl ShardGroup {
pub fn new(group_id: &str, shards: Vec<&str>) -> Self {
Self {
group_id: group_id.to_string(),
shards: shards.into_iter().map(|s| s.to_string()).collect(),
}
}
pub fn len(&self) -> usize {
self.shards.len()
}
pub fn is_empty(&self) -> bool {
self.shards.is_empty()
}
}
pub struct CompositeRouter {
groups: HashMap<String, ShardGroup>,
default_group: Option<ShardGroup>,
vnodes_per_node: usize,
group_rings: HashMap<String, ConsistentHashRouter>,
default_ring: Option<ConsistentHashRouter>,
}
impl CompositeRouter {
pub fn new() -> Self {
Self {
groups: HashMap::new(),
default_group: None,
vnodes_per_node: 100,
group_rings: HashMap::new(),
default_ring: None,
}
}
pub fn with_vnodes(mut self, vnodes: usize) -> Self {
self.vnodes_per_node = vnodes.max(1);
self.group_rings.clear();
self.default_ring = None;
self
}
pub fn add_group(mut self, group: ShardGroup) -> Self {
let group_id = group.group_id.clone();
let nodes: Vec<&str> = group.shards.iter().map(|s| s.as_str()).collect();
let ring = ConsistentHashRouter::new(nodes, self.vnodes_per_node);
self.group_rings.insert(group_id, ring);
self.groups.insert(group.group_id.clone(), group);
self
}
pub fn with_default_group(mut self, group: ShardGroup) -> Self {
let nodes: Vec<&str> = group.shards.iter().map(|s| s.as_str()).collect();
let ring = ConsistentHashRouter::new(nodes, self.vnodes_per_node);
self.default_ring = Some(ring);
self.default_group = Some(group);
self
}
pub fn route(
&self,
group_id: &str,
secondary_key: &str,
) -> Result<String, EnhancedShardingError> {
let ring = self
.group_rings
.get(group_id)
.or(self.default_ring.as_ref())
.ok_or_else(|| EnhancedShardingError::NoGroupMatch(group_id.to_string()))?;
ring.route(secondary_key)
}
pub fn group_count(&self) -> usize {
self.groups.len()
}
pub fn group_ids(&self) -> Vec<String> {
let mut ids: Vec<String> = self.groups.keys().cloned().collect();
ids.sort();
ids
}
pub fn has_default(&self) -> bool {
self.default_group.is_some()
}
}
impl Default for CompositeRouter {
fn default() -> Self {
Self::new()
}
}
impl std::fmt::Debug for CompositeRouter {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("CompositeRouter")
.field("groups", &self.group_ids())
.field("has_default", &self.default_group.is_some())
.field("vnodes_per_node", &self.vnodes_per_node)
.finish()
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RangeShardConfig {
pub lower: i64,
pub upper: i64,
pub shard: String,
}
pub struct RangeConfigRouter {
configs: Vec<RangeShardConfig>,
}
impl RangeConfigRouter {
pub fn new(configs: Vec<RangeShardConfig>) -> Self {
let mut configs = configs;
configs.sort_by_key(|c| c.lower);
Self { configs }
}
pub fn route(&self, key: i64) -> Result<String, EnhancedShardingError> {
for config in &self.configs {
if key >= config.lower && key < config.upper {
return Ok(config.shard.clone());
}
}
Err(EnhancedShardingError::NoListMatch(key.to_string()))
}
pub fn len(&self) -> usize {
self.configs.len()
}
pub fn is_empty(&self) -> bool {
self.configs.is_empty()
}
}
impl std::fmt::Debug for RangeConfigRouter {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RangeConfigRouter")
.field("configs_count", &self.configs.len())
.finish()
}
}
fn hash_str(s: &str) -> u64 {
fnv1a_hash(s)
}
fn ring_insert(ring: &mut BTreeMap<u64, String>, hash: u64, node: String) {
ring.entry(hash).or_insert(node);
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashMap;
#[test]
fn test_consistent_hash_new() {
let router = ConsistentHashRouter::new(vec!["n1", "n2", "n3"], 100);
assert_eq!(router.nodes().len(), 3);
assert_eq!(router.ring_size(), 300); assert_eq!(router.vnodes_per_node(), 100);
}
#[test]
fn test_consistent_hash_vnodes_minimum_1() {
let router = ConsistentHashRouter::new(vec!["n1"], 0);
assert_eq!(router.vnodes_per_node(), 1);
assert_eq!(router.ring_size(), 1);
}
#[test]
fn test_consistent_hash_deterministic() {
let r1 = ConsistentHashRouter::new(vec!["n1", "n2", "n3"], 100);
let r2 = ConsistentHashRouter::new(vec!["n1", "n2", "n3"], 100);
for key in &["a", "b", "c", "user:1", "user:2"] {
assert_eq!(r1.route(key).unwrap(), r2.route(key).unwrap());
}
}
#[test]
fn test_consistent_hash_same_key_same_node() {
let router = ConsistentHashRouter::new(vec!["n1", "n2", "n3"], 100);
let first = router.route("user:123").unwrap();
for _ in 0..5 {
assert_eq!(router.route("user:123").unwrap(), first);
}
}
#[test]
fn test_consistent_hash_distribution() {
let router = ConsistentHashRouter::new(vec!["n1", "n2", "n3", "n4"], 150);
let mut counts: HashMap<String, usize> = HashMap::new();
for i in 0..1000 {
let key = format!("key_{}", i);
let node = router.route(&key).unwrap();
*counts.entry(node).or_insert(0) += 1;
}
for node in ["n1", "n2", "n3", "n4"] {
let count = counts.get(node).copied().unwrap_or(0);
assert!(
count >= 100,
"node {} should have at least 100 keys, got {}",
node,
count
);
}
}
#[test]
fn test_consistent_hash_add_node() {
let mut router = ConsistentHashRouter::new(vec!["n1", "n2"], 100);
assert_eq!(router.nodes().len(), 2);
assert_eq!(router.ring_size(), 200);
router.add_node("n3");
assert_eq!(router.nodes().len(), 3);
assert_eq!(router.ring_size(), 300);
}
#[test]
fn test_consistent_hash_remove_node() {
let mut router = ConsistentHashRouter::new(vec!["n1", "n2", "n3"], 100);
router.remove_node("n2");
assert_eq!(router.nodes().len(), 2);
assert_eq!(router.ring_size(), 200);
assert!(!router.nodes().iter().any(|n| n == "n2"));
}
#[test]
fn test_consistent_hash_add_duplicate_node_noop() {
let mut router = ConsistentHashRouter::new(vec!["n1", "n2"], 100);
router.add_node("n1"); assert_eq!(router.nodes().len(), 2);
assert_eq!(router.ring_size(), 200);
}
#[test]
fn test_consistent_hash_remove_nonexistent_noop() {
let mut router = ConsistentHashRouter::new(vec!["n1", "n2"], 100);
router.remove_node("n999");
assert_eq!(router.nodes().len(), 2);
assert_eq!(router.ring_size(), 200);
}
#[test]
fn test_consistent_hash_add_node_minimal_migration() {
let router1 = ConsistentHashRouter::new(vec!["n1", "n2", "n3"], 150);
let mut router2 = ConsistentHashRouter::new(vec!["n1", "n2", "n3"], 150);
router2.add_node("n4");
let mut total = 0;
let mut migrated = 0;
for i in 0..1000 {
let key = format!("key_{}", i);
let before = router1.route(&key).unwrap();
let after = router2.route(&key).unwrap();
total += 1;
if before != after {
migrated += 1;
}
}
let migration_ratio = migrated as f64 / total as f64;
assert!(
migration_ratio < 0.5,
"migration ratio should be < 50%, got {:.2}%",
migration_ratio * 100.0
);
}
#[test]
fn test_consistent_hash_empty_returns_error() {
let router = ConsistentHashRouter::new(vec![], 100);
let result = router.route("any");
assert_eq!(result, Err(EnhancedShardingError::NoNodes));
}
#[test]
fn test_consistent_hash_single_node() {
let router = ConsistentHashRouter::new(vec!["only"], 100);
for key in &["a", "b", "c", "long_key_here"] {
assert_eq!(router.route(key).unwrap(), "only");
}
}
#[test]
fn test_consistent_hash_debug_format() {
let router = ConsistentHashRouter::new(vec!["n1", "n2"], 100);
let s = format!("{:?}", router);
assert!(s.contains("ConsistentHashRouter"));
assert!(s.contains("ring_size"));
}
#[test]
fn test_list_new() {
let r = ListRouter::new();
assert!(r.is_empty());
assert!(!r.has_default());
}
#[test]
fn test_list_add_and_route() {
let r = ListRouter::new()
.add("cn", "shard_cn")
.add("us", "shard_us")
.add("eu", "shard_eu");
assert_eq!(r.len(), 3);
assert_eq!(r.route("cn").unwrap(), "shard_cn");
assert_eq!(r.route("us").unwrap(), "shard_us");
assert_eq!(r.route("eu").unwrap(), "shard_eu");
}
#[test]
fn test_list_default_fallback() {
let r = ListRouter::new()
.add("cn", "shard_cn")
.with_default("shard_default");
assert!(r.has_default());
assert_eq!(r.route("cn").unwrap(), "shard_cn");
assert_eq!(r.route("unknown").unwrap(), "shard_default");
}
#[test]
fn test_list_no_match_no_default_errors() {
let r = ListRouter::new().add("cn", "shard_cn");
let result = r.route("unknown");
assert!(matches!(result, Err(EnhancedShardingError::NoListMatch(_))));
}
#[test]
fn test_list_empty_errors() {
let r = ListRouter::new();
let result = r.route("any");
assert!(result.is_err());
}
#[test]
fn test_list_overwrite() {
let r = ListRouter::new()
.add("cn", "shard_cn_v1")
.add("cn", "shard_cn_v2");
assert_eq!(r.len(), 1); assert_eq!(r.route("cn").unwrap(), "shard_cn_v2");
}
#[test]
fn test_shard_group_new() {
let g = ShardGroup::new("cn", vec!["cn_0", "cn_1", "cn_2"]);
assert_eq!(g.group_id, "cn");
assert_eq!(g.shards.len(), 3);
assert!(!g.is_empty());
}
#[test]
fn test_shard_group_empty() {
let g = ShardGroup::new("empty", vec![]);
assert!(g.is_empty());
assert_eq!(g.len(), 0);
}
#[test]
fn test_composite_new() {
let r = CompositeRouter::new();
assert_eq!(r.group_count(), 0);
assert!(!r.has_default());
}
#[test]
fn test_composite_add_groups() {
let r = CompositeRouter::new()
.add_group(ShardGroup::new("cn", vec!["cn_0", "cn_1"]))
.add_group(ShardGroup::new("us", vec!["us_0", "us_1"]));
assert_eq!(r.group_count(), 2);
let ids = r.group_ids();
assert_eq!(ids, vec!["cn", "us"]);
}
#[test]
fn test_composite_route_success() {
let r = CompositeRouter::new()
.add_group(ShardGroup::new("cn", vec!["cn_0", "cn_1"]))
.add_group(ShardGroup::new("us", vec!["us_0", "us_1"]));
let result = r.route("cn", "user:123").unwrap();
assert!(result.starts_with("cn_"));
let result = r.route("us", "user:456").unwrap();
assert!(result.starts_with("us_"));
}
#[test]
fn test_composite_route_deterministic() {
let r = CompositeRouter::new().add_group(ShardGroup::new("cn", vec!["cn_0", "cn_1"]));
let r1 = r.route("cn", "user:123").unwrap();
let r2 = r.route("cn", "user:123").unwrap();
assert_eq!(r1, r2);
}
#[test]
fn test_composite_unknown_group_no_default_errors() {
let r = CompositeRouter::new().add_group(ShardGroup::new("cn", vec!["cn_0"]));
let result = r.route("unknown", "key");
assert!(matches!(
result,
Err(EnhancedShardingError::NoGroupMatch(_))
));
}
#[test]
fn test_composite_unknown_group_with_default() {
let r = CompositeRouter::new()
.add_group(ShardGroup::new("cn", vec!["cn_0"]))
.with_default_group(ShardGroup::new("default", vec!["def_0"]));
let result = r.route("unknown", "key").unwrap();
assert_eq!(result, "def_0");
assert!(r.has_default());
}
#[test]
fn test_composite_empty_group_errors() {
let r = CompositeRouter::new().add_group(ShardGroup::new("empty", vec![]));
let result = r.route("empty", "key");
assert_eq!(result, Err(EnhancedShardingError::NoNodes));
}
#[test]
fn test_composite_with_vnodes() {
let r = CompositeRouter::new()
.with_vnodes(50)
.add_group(ShardGroup::new("g1", vec!["s0", "s1"]));
let result = r.route("g1", "key").unwrap();
assert!(result == "s0" || result == "s1");
}
#[test]
fn test_composite_vnodes_minimum_1() {
let r = CompositeRouter::new().with_vnodes(0);
let r = r.add_group(ShardGroup::new("g", vec!["s0"]));
assert_eq!(r.route("g", "k").unwrap(), "s0");
}
#[test]
fn test_composite_debug_format() {
let r = CompositeRouter::new().add_group(ShardGroup::new("cn", vec!["cn_0"]));
let s = format!("{:?}", r);
assert!(s.contains("CompositeRouter"));
assert!(s.contains("cn"));
}
#[test]
fn test_range_config_new() {
let configs = vec![
RangeShardConfig {
lower: 0,
upper: 1000,
shard: "s0".to_string(),
},
RangeShardConfig {
lower: 1000,
upper: 2000,
shard: "s1".to_string(),
},
];
let r = RangeConfigRouter::new(configs);
assert_eq!(r.len(), 2);
assert!(!r.is_empty());
}
#[test]
fn test_range_config_route() {
let configs = vec![
RangeShardConfig {
lower: 0,
upper: 1000,
shard: "s0".to_string(),
},
RangeShardConfig {
lower: 1000,
upper: 2000,
shard: "s1".to_string(),
},
RangeShardConfig {
lower: 2000,
upper: 3000,
shard: "s2".to_string(),
},
];
let r = RangeConfigRouter::new(configs);
assert_eq!(r.route(0).unwrap(), "s0");
assert_eq!(r.route(999).unwrap(), "s0");
assert_eq!(r.route(1000).unwrap(), "s1");
assert_eq!(r.route(1999).unwrap(), "s1");
assert_eq!(r.route(2000).unwrap(), "s2");
assert_eq!(r.route(2999).unwrap(), "s2");
}
#[test]
fn test_range_config_out_of_range_errors() {
let configs = vec![RangeShardConfig {
lower: 0,
upper: 1000,
shard: "s0".to_string(),
}];
let r = RangeConfigRouter::new(configs);
assert_eq!(r.route(500).unwrap(), "s0");
assert!(r.route(1000).is_err()); assert!(r.route(-1).is_err()); }
#[test]
fn test_range_config_empty_errors() {
let r = RangeConfigRouter::new(vec![]);
assert!(r.is_empty());
assert!(r.route(0).is_err());
}
#[test]
fn test_range_config_unsorted_input_sorted() {
let configs = vec![
RangeShardConfig {
lower: 2000,
upper: 3000,
shard: "s2".to_string(),
},
RangeShardConfig {
lower: 0,
upper: 1000,
shard: "s0".to_string(),
},
RangeShardConfig {
lower: 1000,
upper: 2000,
shard: "s1".to_string(),
},
];
let r = RangeConfigRouter::new(configs);
assert_eq!(r.route(500).unwrap(), "s0");
assert_eq!(r.route(1500).unwrap(), "s1");
assert_eq!(r.route(2500).unwrap(), "s2");
}
#[test]
fn test_range_config_negative_range() {
let configs = vec![
RangeShardConfig {
lower: -1000,
upper: 0,
shard: "neg".to_string(),
},
RangeShardConfig {
lower: 0,
upper: 1000,
shard: "pos".to_string(),
},
];
let r = RangeConfigRouter::new(configs);
assert_eq!(r.route(-500).unwrap(), "neg");
assert_eq!(r.route(500).unwrap(), "pos");
}
#[test]
fn test_error_display() {
assert_eq!(
EnhancedShardingError::NoNodes.to_string(),
"no nodes configured"
);
assert_eq!(
EnhancedShardingError::NoGroupMatch("g1".to_string()).to_string(),
"no group matches key: g1"
);
assert_eq!(
EnhancedShardingError::NoListMatch("k1".to_string()).to_string(),
"no list mapping for key: k1"
);
}
#[test]
fn test_error_is_std_error() {
let err = EnhancedShardingError::NoNodes;
let _: &dyn std::error::Error = &err;
}
#[test]
fn test_multi_region_user_routing() {
let router = CompositeRouter::new()
.add_group(ShardGroup::new("cn", vec!["cn_db_0", "cn_db_1", "cn_db_2"]))
.add_group(ShardGroup::new("us", vec!["us_db_0", "us_db_1"]));
let cn_user = router.route("cn", "user:12345").unwrap();
assert!(cn_user.starts_with("cn_db_"));
for _ in 0..5 {
assert_eq!(router.route("cn", "user:12345").unwrap(), cn_user);
}
let us_user = router.route("us", "user:67890").unwrap();
assert!(us_user.starts_with("us_db_"));
for _ in 0..5 {
assert_eq!(router.route("us", "user:67890").unwrap(), us_user);
}
assert!(!cn_user.starts_with("us_"));
assert!(!us_user.starts_with("cn_"));
}
#[test]
fn test_dynamic_scaling() {
let mut router = ConsistentHashRouter::new(vec!["n1", "n2", "n3"], 150);
let mut before: HashMap<String, String> = HashMap::new();
for i in 0..100 {
let key = format!("user:{}", i);
before.insert(key.clone(), router.route(&key).unwrap());
}
router.add_node("n4");
assert_eq!(router.nodes().len(), 4);
let mut unchanged = 0;
let mut migrated = 0;
for (key, old_shard) in &before {
let new_shard = router.route(key).unwrap();
if new_shard == *old_shard {
unchanged += 1;
} else {
migrated += 1;
}
}
assert!(
unchanged > migrated,
"after scaling from 3 to 4 nodes, unchanged ({}) should be > migrated ({})",
unchanged,
migrated
);
}
}