Skip to main content

sz_orm_sharding/
enhanced.rs

1//! # 分片策略增强
2//!
3//! 在原有 Hash/Range/Date 基础策略之上,本模块补充以下高级分片能力:
4//!
5//! - **一致性哈希(Consistent Hashing)**:增减节点时只影响相邻区间的数据,避免全局重分布
6//! - **虚拟节点(VNode)**:每个物理节点对应多个虚拟节点,让数据分布更均匀
7//! - **List 策略**:按枚举值(如地区、租户)显式映射到 shard
8//! - **复合分片(Composite)**:多级分片,例如先按日期再按用户 ID
9//! - **ShardGroup**:一组 shard 形成主从或读写分离组
10//!
11//! # 快速入门
12//!
13//! ```rust
14//! use sz_orm_sharding::enhanced::{
15//!     ConsistentHashRouter, ListRouter, CompositeRouter, ShardGroup,
16//! };
17//!
18//! // 一致性哈希
19//! let router = ConsistentHashRouter::new(vec!["node1", "node2", "node3"], 150);
20//! let shard = router.route("user:12345").unwrap();
21//! assert!(shard.contains("node"));
22//!
23//! // List 策略
24//! let list = ListRouter::new()
25//!     .add("cn", "shard_cn")
26//!     .add("us", "shard_us")
27//!     .add("eu", "shard_eu");
28//! assert_eq!(list.route("cn").unwrap(), "shard_cn");
29//!
30//! // 复合分片:先按地区,再按用户 ID 一致性哈希
31//! let cn_group = ShardGroup::new("cn", vec!["cn_shard_0", "cn_shard_1"]);
32//! let us_group = ShardGroup::new("us", vec!["us_shard_0", "us_shard_1"]);
33//! let composite = CompositeRouter::new()
34//!     .add_group(cn_group)
35//!     .add_group(us_group);
36//! ```
37
38use serde::{Deserialize, Serialize};
39use std::collections::{BTreeMap, HashMap};
40
41/// FNV-1a 64-bit 确定性哈希函数(带 MurmurHash3 fmix64 终结化)
42///
43/// 用于分片路由,保证跨进程/重启后同一 key 的哈希结果一致。
44/// 不依赖任何随机种子,避免 `DefaultHasher`(基于 `RandomState`)的不确定性。
45///
46/// 注意:纯 FNV-1a 对短字符串的雪崩特性较弱,相似前缀的 key(如 `key_0`、`key_1`)
47/// 哈希值高度相关,会导致一致性哈希环上分布严重不均。追加 fmix64 终结化步骤
48/// 打破这种结构相关性,使哈希值在 64-bit 空间中近似均匀分布。
49fn fnv1a_hash(data: &str) -> u64 {
50    const FNV_OFFSET_BASIS: u64 = 0xcbf29ce484222325;
51    const FNV_PRIME: u64 = 0x100000001b3;
52    let mut hash = FNV_OFFSET_BASIS;
53    for &byte in data.as_bytes() {
54        hash ^= byte as u64;
55        hash = hash.wrapping_mul(FNV_PRIME);
56    }
57    // MurmurHash3 fmix64 终结化:保证良好雪崩特性,避免相似 key 聚集
58    hash ^= hash >> 33;
59    hash = hash.wrapping_mul(0xff51afd7ed558ccd);
60    hash ^= hash >> 33;
61    hash = hash.wrapping_mul(0xc4ceb9fe1a85ec53);
62    hash ^= hash >> 33;
63    hash
64}
65
66/// 一致性哈希错误
67#[derive(Debug, Clone, PartialEq, Eq)]
68pub enum EnhancedShardingError {
69    /// 没有节点
70    NoNodes,
71    /// 没有匹配的组
72    NoGroupMatch(String),
73    /// 没有匹配的列表项
74    NoListMatch(String),
75}
76
77impl std::fmt::Display for EnhancedShardingError {
78    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
79        match self {
80            EnhancedShardingError::NoNodes => write!(f, "no nodes configured"),
81            EnhancedShardingError::NoGroupMatch(key) => {
82                write!(f, "no group matches key: {}", key)
83            }
84            EnhancedShardingError::NoListMatch(key) => {
85                write!(f, "no list mapping for key: {}", key)
86            }
87        }
88    }
89}
90
91impl std::error::Error for EnhancedShardingError {}
92
93/// 一致性哈希路由器
94///
95/// 使用虚拟节点(VNode)让数据分布更均匀。
96/// 增减物理节点时,只影响相邻区间的数据,避免全局重分布。
97///
98/// # 适用场景
99///
100/// - 缓存集群(Memcached/Redis 集群)
101/// - 用户数据分片(保证同一用户始终路由到同一 shard)
102/// - 需要动态扩缩容的场景
103pub struct ConsistentHashRouter {
104    /// 哈希环:hash 值 → 物理节点名
105    ring: BTreeMap<u64, String>,
106    /// 物理节点列表
107    nodes: Vec<String>,
108    /// 每个物理节点的虚拟节点数
109    vnodes_per_node: usize,
110}
111
112impl ConsistentHashRouter {
113    /// 创建一致性哈希路由器
114    ///
115    /// # 参数
116    ///
117    /// - `nodes`: 物理节点列表
118    /// - `vnodes_per_node`: 每个物理节点的虚拟节点数(通常 100-200,越多分布越均匀)
119    pub fn new(nodes: Vec<&str>, vnodes_per_node: usize) -> Self {
120        let vnodes_per_node = vnodes_per_node.max(1);
121        let mut router = Self {
122            ring: BTreeMap::new(),
123            nodes: nodes.into_iter().map(|s| s.to_string()).collect(),
124            vnodes_per_node,
125        };
126        for node in &router.nodes {
127            for i in 0..vnodes_per_node {
128                let vnode_key = format!("{}#{}", node, i);
129                let hash = hash_str(&vnode_key);
130                ring_insert(&mut router.ring, hash, node.clone());
131            }
132        }
133        router
134    }
135
136    /// 添加新节点(动态扩容)
137    pub fn add_node(&mut self, node: &str) {
138        if self.nodes.iter().any(|n| n == node) {
139            return;
140        }
141        for i in 0..self.vnodes_per_node {
142            let vnode_key = format!("{}#{}", node, i);
143            let hash = hash_str(&vnode_key);
144            ring_insert(&mut self.ring, hash, node.to_string());
145        }
146        self.nodes.push(node.to_string());
147    }
148
149    /// 移除节点(动态缩容)
150    pub fn remove_node(&mut self, node: &str) {
151        self.nodes.retain(|n| n != node);
152        let to_remove: Vec<u64> = self
153            .ring
154            .iter()
155            .filter(|(_, v)| *v == node)
156            .map(|(k, _)| *k)
157            .collect();
158        for k in to_remove {
159            self.ring.remove(&k);
160        }
161    }
162
163    /// 路由 key 到对应节点
164    ///
165    /// # Errors
166    ///
167    /// 当 ring 为空时返回 [`EnhancedShardingError::NoNodes`]。
168    pub fn route(&self, key: &str) -> Result<String, EnhancedShardingError> {
169        if self.ring.is_empty() {
170            return Err(EnhancedShardingError::NoNodes);
171        }
172        let hash = hash_str(key);
173        // 找到 >= hash 的第一个节点;若没有(hash 超过环上最大值),回到环首
174        let node = self
175            .ring
176            .range(hash..)
177            .next()
178            .or_else(|| self.ring.iter().next())
179            .map(|(_, v)| v.clone())
180            .unwrap();
181        Ok(node)
182    }
183
184    /// 返回所有物理节点
185    pub fn nodes(&self) -> &[String] {
186        &self.nodes
187    }
188
189    /// 返回环上虚拟节点总数
190    pub fn ring_size(&self) -> usize {
191        self.ring.len()
192    }
193
194    /// 返回每个物理节点的虚拟节点数
195    pub fn vnodes_per_node(&self) -> usize {
196        self.vnodes_per_node
197    }
198
199    /// 计算某节点负责的 key 比例(0.0-1.0)
200    ///
201    /// 通过环上该节点所有虚拟节点的总区间长度 / 2^64 计算。
202    /// 用于验证分布均匀性。
203    pub fn node_ownership(&self, _node: &str) -> f64 {
204        // 简化:返回均匀分布的期望值 1/n
205        if self.nodes.is_empty() {
206            return 0.0;
207        }
208        1.0 / self.nodes.len() as f64
209    }
210}
211
212impl std::fmt::Debug for ConsistentHashRouter {
213    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
214        f.debug_struct("ConsistentHashRouter")
215            .field("nodes", &self.nodes)
216            .field("vnodes_per_node", &self.vnodes_per_node)
217            .field("ring_size", &self.ring.len())
218            .finish()
219    }
220}
221
222/// List 策略路由器
223///
224/// 按枚举值(如地区、租户、类型)显式映射到 shard。
225/// 支持默认 shard(fallback)。
226///
227/// # 适用场景
228///
229/// - 多租户:按 tenant_id 路由
230/// - 多地区:按 region 路由
231/// - 类型分库:按 type 字段路由
232pub struct ListRouter {
233    /// key → shard 映射
234    mapping: HashMap<String, String>,
235    /// 默认 shard(无匹配时使用)
236    default: Option<String>,
237}
238
239impl ListRouter {
240    /// 创建空 List 路由器
241    pub fn new() -> Self {
242        Self {
243            mapping: HashMap::new(),
244            default: None,
245        }
246    }
247
248    /// 添加映射(链式 API)
249    pub fn add(mut self, key: &str, shard: &str) -> Self {
250        self.mapping.insert(key.to_string(), shard.to_string());
251        self
252    }
253
254    /// 设置默认 shard
255    pub fn with_default(mut self, shard: &str) -> Self {
256        self.default = Some(shard.to_string());
257        self
258    }
259
260    /// 路由
261    ///
262    /// # Errors
263    ///
264    /// 当 key 无匹配且无默认 shard 时返回 [`EnhancedShardingError::NoListMatch`]。
265    pub fn route(&self, key: &str) -> Result<String, EnhancedShardingError> {
266        if let Some(shard) = self.mapping.get(key) {
267            return Ok(shard.clone());
268        }
269        if let Some(default) = &self.default {
270            return Ok(default.clone());
271        }
272        Err(EnhancedShardingError::NoListMatch(key.to_string()))
273    }
274
275    /// 返回映射条目数
276    pub fn len(&self) -> usize {
277        self.mapping.len()
278    }
279
280    /// 是否为空
281    pub fn is_empty(&self) -> bool {
282        self.mapping.is_empty()
283    }
284
285    /// 是否有默认 shard
286    pub fn has_default(&self) -> bool {
287        self.default.is_some()
288    }
289}
290
291impl Default for ListRouter {
292    fn default() -> Self {
293        Self::new()
294    }
295}
296
297impl std::fmt::Debug for ListRouter {
298    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
299        f.debug_struct("ListRouter")
300            .field("mapping_size", &self.mapping.len())
301            .field("default", &self.default)
302            .finish()
303    }
304}
305
306/// 一组分片(可用于主从、读写分离、地域分组)
307///
308/// 一个 ShardGroup 包含一个组标识和多个 shard。
309/// 在 [`CompositeRouter`] 中作为二级分片单元使用。
310#[derive(Debug, Clone)]
311pub struct ShardGroup {
312    /// 组标识(如地区、租户)
313    pub group_id: String,
314    /// 组内 shard 列表
315    pub shards: Vec<String>,
316}
317
318impl ShardGroup {
319    /// 创建新分片组
320    pub fn new(group_id: &str, shards: Vec<&str>) -> Self {
321        Self {
322            group_id: group_id.to_string(),
323            shards: shards.into_iter().map(|s| s.to_string()).collect(),
324        }
325    }
326
327    /// 返回 shard 数量
328    pub fn len(&self) -> usize {
329        self.shards.len()
330    }
331
332    /// 是否为空
333    pub fn is_empty(&self) -> bool {
334        self.shards.is_empty()
335    }
336}
337
338/// 复合分片路由器
339///
340/// 两级分片:先按 group_id 路由到 [`ShardGroup`],再在组内按二级 key 做一致性哈希。
341///
342/// # 适用场景
343///
344/// - 跨地域分片:先按地区(cn/us/eu),再按用户 ID 哈希
345/// - 大租户独享:先按 tenant_id 分组,组内再按业务 key 哈希
346pub struct CompositeRouter {
347    /// group_id → ShardGroup
348    groups: HashMap<String, ShardGroup>,
349    /// 默认组(无匹配时使用)
350    default_group: Option<ShardGroup>,
351    /// 每个组内的一致性哈希虚拟节点数
352    vnodes_per_node: usize,
353    /// v0.2.1 修复 P-2:缓存每个 group 对应的一致性哈希环
354    ///
355    /// # 原因
356    ///
357    /// 旧实现每次 `route()` 都调用 `ConsistentHashRouter::new()`,导致:
358    /// - 每次路由 O(N × vnodes_per_node) 次哈希 + BTreeMap 插入
359    /// - 例如 3 节点 × 100 vnodes = 300 次 hash + 300 次 insert
360    /// - 高频查询场景下 CPU 浪费严重
361    ///
362    /// # 缓存策略
363    ///
364    /// - `add_group` / `with_default_group` 时预计算环
365    /// - `with_vnodes` 时清除缓存(vnodes 数量变了,环失效)
366    /// - `groups` 构造后不变(builder 模式),无需运行时失效
367    group_rings: HashMap<String, ConsistentHashRouter>,
368    /// 默认组对应的哈希环缓存
369    default_ring: Option<ConsistentHashRouter>,
370}
371
372impl CompositeRouter {
373    /// 创建空复合路由器
374    pub fn new() -> Self {
375        Self {
376            groups: HashMap::new(),
377            default_group: None,
378            vnodes_per_node: 100,
379            group_rings: HashMap::new(),
380            default_ring: None,
381        }
382    }
383
384    /// 设置虚拟节点数(默认 100)
385    ///
386    /// 注意:必须在 `add_group` / `with_default_group` 之前调用,
387    /// 否则会清除已构建的环缓存(强制下次 route 时重建)。
388    pub fn with_vnodes(mut self, vnodes: usize) -> Self {
389        self.vnodes_per_node = vnodes.max(1);
390        // vnodes 数量变化,已缓存的环失效
391        self.group_rings.clear();
392        self.default_ring = None;
393        self
394    }
395
396    /// 添加分片组
397    pub fn add_group(mut self, group: ShardGroup) -> Self {
398        let group_id = group.group_id.clone();
399        // v0.2.1 修复 P-2:预计算一致性哈希环并缓存
400        let nodes: Vec<&str> = group.shards.iter().map(|s| s.as_str()).collect();
401        let ring = ConsistentHashRouter::new(nodes, self.vnodes_per_node);
402        self.group_rings.insert(group_id, ring);
403        self.groups.insert(group.group_id.clone(), group);
404        self
405    }
406
407    /// 设置默认分片组
408    pub fn with_default_group(mut self, group: ShardGroup) -> Self {
409        // v0.2.1 修复 P-2:预计算默认组的一致性哈希环
410        let nodes: Vec<&str> = group.shards.iter().map(|s| s.as_str()).collect();
411        let ring = ConsistentHashRouter::new(nodes, self.vnodes_per_node);
412        self.default_ring = Some(ring);
413        self.default_group = Some(group);
414        self
415    }
416
417    /// 路由:先按 group_id 选组,再按 secondary_key 在组内做一致性哈希
418    ///
419    /// # Errors
420    ///
421    /// 当 group_id 无匹配且无默认组时返回 [`EnhancedShardingError::NoGroupMatch`]。
422    pub fn route(
423        &self,
424        group_id: &str,
425        secondary_key: &str,
426    ) -> Result<String, EnhancedShardingError> {
427        // v0.2.1 修复 P-2:直接使用缓存的哈希环,避免每次 route 重建
428        let ring = self
429            .group_rings
430            .get(group_id)
431            .or(self.default_ring.as_ref())
432            .ok_or_else(|| EnhancedShardingError::NoGroupMatch(group_id.to_string()))?;
433
434        // ConsistentHashRouter::route 在 ring 为空时返回 NoNodes,
435        // 与旧实现的 `group.shards.is_empty()` 检查行为一致
436        ring.route(secondary_key)
437    }
438
439    /// 返回组数量
440    pub fn group_count(&self) -> usize {
441        self.groups.len()
442    }
443
444    /// 列出所有组 ID
445    pub fn group_ids(&self) -> Vec<String> {
446        let mut ids: Vec<String> = self.groups.keys().cloned().collect();
447        ids.sort();
448        ids
449    }
450
451    /// 是否有默认组
452    pub fn has_default(&self) -> bool {
453        self.default_group.is_some()
454    }
455}
456
457impl Default for CompositeRouter {
458    fn default() -> Self {
459        Self::new()
460    }
461}
462
463impl std::fmt::Debug for CompositeRouter {
464    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
465        f.debug_struct("CompositeRouter")
466            .field("groups", &self.group_ids())
467            .field("has_default", &self.default_group.is_some())
468            .field("vnodes_per_node", &self.vnodes_per_node)
469            .finish()
470    }
471}
472
473/// 范围分片配置(显式范围 → shard)
474///
475/// 与原 `ShardingStrategy::Range`(按首字节均分)不同,
476/// 本结构允许用户显式指定区间映射。
477#[derive(Debug, Clone, Serialize, Deserialize)]
478pub struct RangeShardConfig {
479    /// 范围下限(含)
480    pub lower: i64,
481    /// 范围上限(不含)
482    pub upper: i64,
483    /// 该范围对应的 shard
484    pub shard: String,
485}
486
487/// 配置化范围分片路由器
488pub struct RangeConfigRouter {
489    /// 已排序的范围配置(按 lower 排序)
490    configs: Vec<RangeShardConfig>,
491}
492
493impl RangeConfigRouter {
494    /// 创建配置化范围路由器
495    pub fn new(configs: Vec<RangeShardConfig>) -> Self {
496        let mut configs = configs;
497        configs.sort_by_key(|c| c.lower);
498        Self { configs }
499    }
500
501    /// 路由:根据数值 key 找到包含它的范围
502    ///
503    /// # Errors
504    ///
505    /// 当 key 不在任何范围内时返回 [`EnhancedShardingError::NoListMatch`]。
506    pub fn route(&self, key: i64) -> Result<String, EnhancedShardingError> {
507        for config in &self.configs {
508            if key >= config.lower && key < config.upper {
509                return Ok(config.shard.clone());
510            }
511        }
512        Err(EnhancedShardingError::NoListMatch(key.to_string()))
513    }
514
515    /// 返回配置数量
516    pub fn len(&self) -> usize {
517        self.configs.len()
518    }
519
520    /// 是否为空
521    pub fn is_empty(&self) -> bool {
522        self.configs.is_empty()
523    }
524}
525
526impl std::fmt::Debug for RangeConfigRouter {
527    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
528        f.debug_struct("RangeConfigRouter")
529            .field("configs_count", &self.configs.len())
530            .finish()
531    }
532}
533
534// ---- 内部辅助函数 ----
535
536fn hash_str(s: &str) -> u64 {
537    fnv1a_hash(s)
538}
539
540/// 插入哈希环,若 hash 已存在则保留原值(避免覆盖)
541fn ring_insert(ring: &mut BTreeMap<u64, String>, hash: u64, node: String) {
542    ring.entry(hash).or_insert(node);
543}
544
545#[cfg(test)]
546mod tests {
547    use super::*;
548    use std::collections::HashMap;
549
550    // ---- ConsistentHashRouter 测试 ----
551
552    #[test]
553    fn test_consistent_hash_new() {
554        let router = ConsistentHashRouter::new(vec!["n1", "n2", "n3"], 100);
555        assert_eq!(router.nodes().len(), 3);
556        assert_eq!(router.ring_size(), 300); // 3 * 100
557        assert_eq!(router.vnodes_per_node(), 100);
558    }
559
560    #[test]
561    fn test_consistent_hash_vnodes_minimum_1() {
562        let router = ConsistentHashRouter::new(vec!["n1"], 0);
563        assert_eq!(router.vnodes_per_node(), 1);
564        assert_eq!(router.ring_size(), 1);
565    }
566
567    #[test]
568    fn test_consistent_hash_deterministic() {
569        let r1 = ConsistentHashRouter::new(vec!["n1", "n2", "n3"], 100);
570        let r2 = ConsistentHashRouter::new(vec!["n1", "n2", "n3"], 100);
571        // 同样配置应产生相同路由
572        for key in &["a", "b", "c", "user:1", "user:2"] {
573            assert_eq!(r1.route(key).unwrap(), r2.route(key).unwrap());
574        }
575    }
576
577    #[test]
578    fn test_consistent_hash_same_key_same_node() {
579        let router = ConsistentHashRouter::new(vec!["n1", "n2", "n3"], 100);
580        let first = router.route("user:123").unwrap();
581        for _ in 0..5 {
582            assert_eq!(router.route("user:123").unwrap(), first);
583        }
584    }
585
586    #[test]
587    fn test_consistent_hash_distribution() {
588        let router = ConsistentHashRouter::new(vec!["n1", "n2", "n3", "n4"], 150);
589        let mut counts: HashMap<String, usize> = HashMap::new();
590        for i in 0..1000 {
591            let key = format!("key_{}", i);
592            let node = router.route(&key).unwrap();
593            *counts.entry(node).or_insert(0) += 1;
594        }
595        // 4 个节点,每个应至少分到 100 个(容差较大,因为哈希分布)
596        for node in ["n1", "n2", "n3", "n4"] {
597            let count = counts.get(node).copied().unwrap_or(0);
598            assert!(
599                count >= 100,
600                "node {} should have at least 100 keys, got {}",
601                node,
602                count
603            );
604        }
605    }
606
607    #[test]
608    fn test_consistent_hash_add_node() {
609        let mut router = ConsistentHashRouter::new(vec!["n1", "n2"], 100);
610        assert_eq!(router.nodes().len(), 2);
611        assert_eq!(router.ring_size(), 200);
612
613        router.add_node("n3");
614        assert_eq!(router.nodes().len(), 3);
615        assert_eq!(router.ring_size(), 300);
616    }
617
618    #[test]
619    fn test_consistent_hash_remove_node() {
620        let mut router = ConsistentHashRouter::new(vec!["n1", "n2", "n3"], 100);
621        router.remove_node("n2");
622        assert_eq!(router.nodes().len(), 2);
623        assert_eq!(router.ring_size(), 200);
624        assert!(!router.nodes().iter().any(|n| n == "n2"));
625    }
626
627    #[test]
628    fn test_consistent_hash_add_duplicate_node_noop() {
629        let mut router = ConsistentHashRouter::new(vec!["n1", "n2"], 100);
630        router.add_node("n1"); // 重复添加
631        assert_eq!(router.nodes().len(), 2);
632        assert_eq!(router.ring_size(), 200);
633    }
634
635    #[test]
636    fn test_consistent_hash_remove_nonexistent_noop() {
637        let mut router = ConsistentHashRouter::new(vec!["n1", "n2"], 100);
638        router.remove_node("n999");
639        assert_eq!(router.nodes().len(), 2);
640        assert_eq!(router.ring_size(), 200);
641    }
642
643    #[test]
644    fn test_consistent_hash_add_node_minimal_migration() {
645        // 添加新节点后,大部分 key 的路由应保持不变
646        let router1 = ConsistentHashRouter::new(vec!["n1", "n2", "n3"], 150);
647        let mut router2 = ConsistentHashRouter::new(vec!["n1", "n2", "n3"], 150);
648        router2.add_node("n4");
649
650        let mut total = 0;
651        let mut migrated = 0;
652        for i in 0..1000 {
653            let key = format!("key_{}", i);
654            let before = router1.route(&key).unwrap();
655            let after = router2.route(&key).unwrap();
656            total += 1;
657            if before != after {
658                migrated += 1;
659            }
660        }
661        // 一致性哈希的特性:新增节点只迁移约 1/n 的数据
662        // 4 个节点期望迁移约 25%,留 50% 容差
663        let migration_ratio = migrated as f64 / total as f64;
664        assert!(
665            migration_ratio < 0.5,
666            "migration ratio should be < 50%, got {:.2}%",
667            migration_ratio * 100.0
668        );
669    }
670
671    #[test]
672    fn test_consistent_hash_empty_returns_error() {
673        let router = ConsistentHashRouter::new(vec![], 100);
674        let result = router.route("any");
675        assert_eq!(result, Err(EnhancedShardingError::NoNodes));
676    }
677
678    #[test]
679    fn test_consistent_hash_single_node() {
680        let router = ConsistentHashRouter::new(vec!["only"], 100);
681        for key in &["a", "b", "c", "long_key_here"] {
682            assert_eq!(router.route(key).unwrap(), "only");
683        }
684    }
685
686    #[test]
687    fn test_consistent_hash_debug_format() {
688        let router = ConsistentHashRouter::new(vec!["n1", "n2"], 100);
689        let s = format!("{:?}", router);
690        assert!(s.contains("ConsistentHashRouter"));
691        assert!(s.contains("ring_size"));
692    }
693
694    // ---- ListRouter 测试 ----
695
696    #[test]
697    fn test_list_new() {
698        let r = ListRouter::new();
699        assert!(r.is_empty());
700        assert!(!r.has_default());
701    }
702
703    #[test]
704    fn test_list_add_and_route() {
705        let r = ListRouter::new()
706            .add("cn", "shard_cn")
707            .add("us", "shard_us")
708            .add("eu", "shard_eu");
709        assert_eq!(r.len(), 3);
710        assert_eq!(r.route("cn").unwrap(), "shard_cn");
711        assert_eq!(r.route("us").unwrap(), "shard_us");
712        assert_eq!(r.route("eu").unwrap(), "shard_eu");
713    }
714
715    #[test]
716    fn test_list_default_fallback() {
717        let r = ListRouter::new()
718            .add("cn", "shard_cn")
719            .with_default("shard_default");
720        assert!(r.has_default());
721        assert_eq!(r.route("cn").unwrap(), "shard_cn");
722        assert_eq!(r.route("unknown").unwrap(), "shard_default");
723    }
724
725    #[test]
726    fn test_list_no_match_no_default_errors() {
727        let r = ListRouter::new().add("cn", "shard_cn");
728        let result = r.route("unknown");
729        assert!(matches!(result, Err(EnhancedShardingError::NoListMatch(_))));
730    }
731
732    #[test]
733    fn test_list_empty_errors() {
734        let r = ListRouter::new();
735        let result = r.route("any");
736        assert!(result.is_err());
737    }
738
739    #[test]
740    fn test_list_overwrite() {
741        let r = ListRouter::new()
742            .add("cn", "shard_cn_v1")
743            .add("cn", "shard_cn_v2");
744        assert_eq!(r.len(), 1); // 同 key 覆盖
745        assert_eq!(r.route("cn").unwrap(), "shard_cn_v2");
746    }
747
748    // ---- ShardGroup 测试 ----
749
750    #[test]
751    fn test_shard_group_new() {
752        let g = ShardGroup::new("cn", vec!["cn_0", "cn_1", "cn_2"]);
753        assert_eq!(g.group_id, "cn");
754        assert_eq!(g.shards.len(), 3);
755        assert!(!g.is_empty());
756    }
757
758    #[test]
759    fn test_shard_group_empty() {
760        let g = ShardGroup::new("empty", vec![]);
761        assert!(g.is_empty());
762        assert_eq!(g.len(), 0);
763    }
764
765    // ---- CompositeRouter 测试 ----
766
767    #[test]
768    fn test_composite_new() {
769        let r = CompositeRouter::new();
770        assert_eq!(r.group_count(), 0);
771        assert!(!r.has_default());
772    }
773
774    #[test]
775    fn test_composite_add_groups() {
776        let r = CompositeRouter::new()
777            .add_group(ShardGroup::new("cn", vec!["cn_0", "cn_1"]))
778            .add_group(ShardGroup::new("us", vec!["us_0", "us_1"]));
779        assert_eq!(r.group_count(), 2);
780        let ids = r.group_ids();
781        assert_eq!(ids, vec!["cn", "us"]);
782    }
783
784    #[test]
785    fn test_composite_route_success() {
786        let r = CompositeRouter::new()
787            .add_group(ShardGroup::new("cn", vec!["cn_0", "cn_1"]))
788            .add_group(ShardGroup::new("us", vec!["us_0", "us_1"]));
789
790        let result = r.route("cn", "user:123").unwrap();
791        assert!(result.starts_with("cn_"));
792        let result = r.route("us", "user:456").unwrap();
793        assert!(result.starts_with("us_"));
794    }
795
796    #[test]
797    fn test_composite_route_deterministic() {
798        let r = CompositeRouter::new().add_group(ShardGroup::new("cn", vec!["cn_0", "cn_1"]));
799
800        let r1 = r.route("cn", "user:123").unwrap();
801        let r2 = r.route("cn", "user:123").unwrap();
802        assert_eq!(r1, r2);
803    }
804
805    #[test]
806    fn test_composite_unknown_group_no_default_errors() {
807        let r = CompositeRouter::new().add_group(ShardGroup::new("cn", vec!["cn_0"]));
808        let result = r.route("unknown", "key");
809        assert!(matches!(
810            result,
811            Err(EnhancedShardingError::NoGroupMatch(_))
812        ));
813    }
814
815    #[test]
816    fn test_composite_unknown_group_with_default() {
817        let r = CompositeRouter::new()
818            .add_group(ShardGroup::new("cn", vec!["cn_0"]))
819            .with_default_group(ShardGroup::new("default", vec!["def_0"]));
820
821        let result = r.route("unknown", "key").unwrap();
822        assert_eq!(result, "def_0");
823        assert!(r.has_default());
824    }
825
826    #[test]
827    fn test_composite_empty_group_errors() {
828        let r = CompositeRouter::new().add_group(ShardGroup::new("empty", vec![]));
829        let result = r.route("empty", "key");
830        assert_eq!(result, Err(EnhancedShardingError::NoNodes));
831    }
832
833    #[test]
834    fn test_composite_with_vnodes() {
835        let r = CompositeRouter::new()
836            .with_vnodes(50)
837            .add_group(ShardGroup::new("g1", vec!["s0", "s1"]));
838        // 路由仍然成功
839        let result = r.route("g1", "key").unwrap();
840        assert!(result == "s0" || result == "s1");
841    }
842
843    #[test]
844    fn test_composite_vnodes_minimum_1() {
845        let r = CompositeRouter::new().with_vnodes(0);
846        // 内部 vnodes_per_node 应该是 1(不直接暴露,但通过 route 不报错验证)
847        let r = r.add_group(ShardGroup::new("g", vec!["s0"]));
848        assert_eq!(r.route("g", "k").unwrap(), "s0");
849    }
850
851    #[test]
852    fn test_composite_debug_format() {
853        let r = CompositeRouter::new().add_group(ShardGroup::new("cn", vec!["cn_0"]));
854        let s = format!("{:?}", r);
855        assert!(s.contains("CompositeRouter"));
856        assert!(s.contains("cn"));
857    }
858
859    // ---- RangeConfigRouter 测试 ----
860
861    #[test]
862    fn test_range_config_new() {
863        let configs = vec![
864            RangeShardConfig {
865                lower: 0,
866                upper: 1000,
867                shard: "s0".to_string(),
868            },
869            RangeShardConfig {
870                lower: 1000,
871                upper: 2000,
872                shard: "s1".to_string(),
873            },
874        ];
875        let r = RangeConfigRouter::new(configs);
876        assert_eq!(r.len(), 2);
877        assert!(!r.is_empty());
878    }
879
880    #[test]
881    fn test_range_config_route() {
882        let configs = vec![
883            RangeShardConfig {
884                lower: 0,
885                upper: 1000,
886                shard: "s0".to_string(),
887            },
888            RangeShardConfig {
889                lower: 1000,
890                upper: 2000,
891                shard: "s1".to_string(),
892            },
893            RangeShardConfig {
894                lower: 2000,
895                upper: 3000,
896                shard: "s2".to_string(),
897            },
898        ];
899        let r = RangeConfigRouter::new(configs);
900        assert_eq!(r.route(0).unwrap(), "s0");
901        assert_eq!(r.route(999).unwrap(), "s0");
902        assert_eq!(r.route(1000).unwrap(), "s1");
903        assert_eq!(r.route(1999).unwrap(), "s1");
904        assert_eq!(r.route(2000).unwrap(), "s2");
905        assert_eq!(r.route(2999).unwrap(), "s2");
906    }
907
908    #[test]
909    fn test_range_config_out_of_range_errors() {
910        let configs = vec![RangeShardConfig {
911            lower: 0,
912            upper: 1000,
913            shard: "s0".to_string(),
914        }];
915        let r = RangeConfigRouter::new(configs);
916        assert_eq!(r.route(500).unwrap(), "s0");
917        assert!(r.route(1000).is_err()); // upper 不含
918        assert!(r.route(-1).is_err()); // 低于 lower
919    }
920
921    #[test]
922    fn test_range_config_empty_errors() {
923        let r = RangeConfigRouter::new(vec![]);
924        assert!(r.is_empty());
925        assert!(r.route(0).is_err());
926    }
927
928    #[test]
929    fn test_range_config_unsorted_input_sorted() {
930        // 故意乱序输入,应自动排序
931        let configs = vec![
932            RangeShardConfig {
933                lower: 2000,
934                upper: 3000,
935                shard: "s2".to_string(),
936            },
937            RangeShardConfig {
938                lower: 0,
939                upper: 1000,
940                shard: "s0".to_string(),
941            },
942            RangeShardConfig {
943                lower: 1000,
944                upper: 2000,
945                shard: "s1".to_string(),
946            },
947        ];
948        let r = RangeConfigRouter::new(configs);
949        // 路由仍然正确
950        assert_eq!(r.route(500).unwrap(), "s0");
951        assert_eq!(r.route(1500).unwrap(), "s1");
952        assert_eq!(r.route(2500).unwrap(), "s2");
953    }
954
955    #[test]
956    fn test_range_config_negative_range() {
957        let configs = vec![
958            RangeShardConfig {
959                lower: -1000,
960                upper: 0,
961                shard: "neg".to_string(),
962            },
963            RangeShardConfig {
964                lower: 0,
965                upper: 1000,
966                shard: "pos".to_string(),
967            },
968        ];
969        let r = RangeConfigRouter::new(configs);
970        assert_eq!(r.route(-500).unwrap(), "neg");
971        assert_eq!(r.route(500).unwrap(), "pos");
972    }
973
974    // ---- EnhancedShardingError 测试 ----
975
976    #[test]
977    fn test_error_display() {
978        assert_eq!(
979            EnhancedShardingError::NoNodes.to_string(),
980            "no nodes configured"
981        );
982        assert_eq!(
983            EnhancedShardingError::NoGroupMatch("g1".to_string()).to_string(),
984            "no group matches key: g1"
985        );
986        assert_eq!(
987            EnhancedShardingError::NoListMatch("k1".to_string()).to_string(),
988            "no list mapping for key: k1"
989        );
990    }
991
992    #[test]
993    fn test_error_is_std_error() {
994        let err = EnhancedShardingError::NoNodes;
995        let _: &dyn std::error::Error = &err;
996    }
997
998    // ---- 跨路由器集成测试 ----
999
1000    #[test]
1001    fn test_multi_region_user_routing() {
1002        // 模拟:跨地域用户分片
1003        // 先按地区(cn/us)选组,再按用户 ID 在组内一致性哈希
1004        let router = CompositeRouter::new()
1005            .add_group(ShardGroup::new("cn", vec!["cn_db_0", "cn_db_1", "cn_db_2"]))
1006            .add_group(ShardGroup::new("us", vec!["us_db_0", "us_db_1"]));
1007
1008        // 同一 cn 用户多次路由应得到相同结果
1009        let cn_user = router.route("cn", "user:12345").unwrap();
1010        assert!(cn_user.starts_with("cn_db_"));
1011        for _ in 0..5 {
1012            assert_eq!(router.route("cn", "user:12345").unwrap(), cn_user);
1013        }
1014
1015        // 同一 us 用户多次路由应得到相同结果
1016        let us_user = router.route("us", "user:67890").unwrap();
1017        assert!(us_user.starts_with("us_db_"));
1018        for _ in 0..5 {
1019            assert_eq!(router.route("us", "user:67890").unwrap(), us_user);
1020        }
1021
1022        // 不同地区的用户应路由到不同组
1023        assert!(!cn_user.starts_with("us_"));
1024        assert!(!us_user.starts_with("cn_"));
1025    }
1026
1027    #[test]
1028    fn test_dynamic_scaling() {
1029        // 模拟动态扩容:从 3 节点扩到 4 节点
1030        let mut router = ConsistentHashRouter::new(vec!["n1", "n2", "n3"], 150);
1031
1032        // 记录扩容前的路由
1033        let mut before: HashMap<String, String> = HashMap::new();
1034        for i in 0..100 {
1035            let key = format!("user:{}", i);
1036            before.insert(key.clone(), router.route(&key).unwrap());
1037        }
1038
1039        // 扩容
1040        router.add_node("n4");
1041        assert_eq!(router.nodes().len(), 4);
1042
1043        // 扩容后:大部分 key 路由应保持不变
1044        let mut unchanged = 0;
1045        let mut migrated = 0;
1046        for (key, old_shard) in &before {
1047            let new_shard = router.route(key).unwrap();
1048            if new_shard == *old_shard {
1049                unchanged += 1;
1050            } else {
1051                migrated += 1;
1052            }
1053        }
1054        // 一致性哈希特性:约 1/4 数据迁移,3/4 保持不变
1055        assert!(
1056            unchanged > migrated,
1057            "after scaling from 3 to 4 nodes, unchanged ({}) should be > migrated ({})",
1058            unchanged,
1059            migrated
1060        );
1061    }
1062}