Skip to main content

sz_orm_rw/
lib.rs

1//! # SZ-ORM RW — 读写分离
2//!
3//! 提供 master/slave 读写分离路由,支持轮询、随机、最少连接三种负载均衡策略,
4//! 写请求路由至 master,读请求在 slave 集群间分配。
5//!
6//! ## 主要类型
7//!
8//! - [`ReadWriteRouter`] — 读写分离路由器
9//! - [`LoadBalanceStrategy`] — 负载均衡策略
10//! - [`HealthChecker`] — 健康检查与故障转移
11//! - [`WeightedSlave`] — 加权 slave 配置
12//! - [`ReadRationing`] — 读写比例控制
13//! - [`LatencyStats`] — 延迟统计
14
15use rand::Rng;
16use serde::{Deserialize, Serialize};
17use std::collections::HashMap;
18use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
19use std::sync::Mutex;
20use std::time::{Duration, Instant};
21
22/// 负载均衡策略
23#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
24pub enum LoadBalanceStrategy {
25    /// 轮询:依次分配请求到各 slave
26    RoundRobin,
27    /// 随机:基于系统时间熵随机选择 slave
28    Random,
29    /// 最少连接:选择当前活跃连接数最少的 slave
30    LeastConnections,
31    /// 加权轮询:根据 slave 权重比例分配请求
32    WeightedRoundRobin,
33}
34
35/// Slave 健康状态
36#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
37pub enum SlaveHealth {
38    /// 健康,可正常服务
39    Healthy,
40    /// 不健康,已被故障转移排除
41    Unhealthy,
42    /// 临时不可用(如维护中)
43    Drained,
44}
45
46/// 加权 slave 配置
47#[derive(Debug, Clone, Serialize, Deserialize)]
48pub struct WeightedSlave {
49    /// slave 地址
50    pub addr: String,
51    /// 权重(>=1),权重越大被选中的概率越高
52    pub weight: u32,
53    /// 当前健康状态
54    pub health: SlaveHealth,
55}
56
57impl WeightedSlave {
58    pub fn new(addr: impl Into<String>, weight: u32) -> Self {
59        Self {
60            addr: addr.into(),
61            weight: weight.max(1),
62            health: SlaveHealth::Healthy,
63        }
64    }
65
66    pub fn with_health(mut self, health: SlaveHealth) -> Self {
67        self.health = health;
68        self
69    }
70}
71
72/// 单个 slave 的延迟统计快照
73#[derive(Debug, Clone, Default, Serialize, Deserialize)]
74pub struct LatencySnapshot {
75    /// 已采样的请求总数
76    pub samples: u64,
77    /// 最小延迟(纳秒)
78    pub min_ns: u64,
79    /// 最大延迟(纳秒)
80    pub max_ns: u64,
81    /// 累计延迟(纳秒),用于计算平均值
82    pub sum_ns: u128,
83}
84
85impl LatencySnapshot {
86    pub fn record(&mut self, latency: Duration) {
87        let ns = latency.as_nanos();
88        self.samples += 1;
89        if self.min_ns == 0 || ns < self.min_ns as u128 {
90            self.min_ns = ns.min(u64::MAX as u128) as u64;
91        }
92        if ns > self.max_ns as u128 {
93            self.max_ns = ns.min(u64::MAX as u128) as u64;
94        }
95        self.sum_ns = self.sum_ns.saturating_add(ns);
96    }
97
98    /// 平均延迟(纳秒),无样本时返回 0
99    pub fn avg_ns(&self) -> u64 {
100        if self.samples == 0 {
101            0
102        } else {
103            (self.sum_ns / self.samples as u128) as u64
104        }
105    }
106
107    pub fn avg(&self) -> Duration {
108        Duration::from_nanos(self.avg_ns())
109    }
110}
111
112/// 延迟统计:为每个 slave 维护一份独立的 [`LatencySnapshot`]
113#[derive(Debug, Default)]
114pub struct LatencyStats {
115    inner: Mutex<HashMap<String, LatencySnapshot>>,
116}
117
118impl LatencyStats {
119    pub fn new() -> Self {
120        Self::default()
121    }
122
123    /// 记录一次 slave 的请求延迟
124    pub fn record(&self, slave: &str, latency: Duration) {
125        if let Ok(mut inner) = self.inner.lock() {
126            inner.entry(slave.to_string()).or_default().record(latency);
127        }
128    }
129
130    /// 返回某个 slave 的快照副本
131    pub fn snapshot(&self, slave: &str) -> LatencySnapshot {
132        match self.inner.lock() {
133            Ok(inner) => inner.get(slave).cloned().unwrap_or_default(),
134            Err(_) => LatencySnapshot::default(),
135        }
136    }
137
138    /// 列出所有 slave 的快照
139    pub fn all(&self) -> Vec<(String, LatencySnapshot)> {
140        match self.inner.lock() {
141            Ok(inner) => inner.iter().map(|(k, v)| (k.clone(), v.clone())).collect(),
142            Err(_) => Vec::new(),
143        }
144    }
145
146    /// 重置某个 slave 的统计
147    pub fn reset(&self, slave: &str) {
148        if let Ok(mut inner) = self.inner.lock() {
149            inner.remove(slave);
150        }
151    }
152}
153
154/// 读写比例控制器:按比例把读请求路由到 master(强一致读)或 slave(弱一致读)
155#[derive(Debug, Serialize, Deserialize)]
156pub struct ReadRationing {
157    /// 0..=100,表示有多少比例的读走 master
158    /// 0 表示全部走 slave,100 表示全部走 master
159    pub master_read_percent: u8,
160    /// 内部计数器,用于轮询决定
161    #[serde(skip)]
162    counter: AtomicU64,
163}
164
165impl Clone for ReadRationing {
166    fn clone(&self) -> Self {
167        Self {
168            master_read_percent: self.master_read_percent,
169            counter: AtomicU64::new(self.counter.load(Ordering::Relaxed)),
170        }
171    }
172}
173
174impl ReadRationing {
175    pub fn new(master_read_percent: u8) -> Self {
176        Self {
177            master_read_percent: master_read_percent.min(100),
178            counter: AtomicU64::new(0),
179        }
180    }
181
182    /// 默认 0% 走 master(全部走 slave)
183    pub fn default_slave_only() -> Self {
184        Self::new(0)
185    }
186
187    /// 默认 100% 走 master(强一致读)
188    pub fn default_master_only() -> Self {
189        Self::new(100)
190    }
191
192    /// 决定本次读请求是否走 master
193    pub fn should_read_master(&self) -> bool {
194        if self.master_read_percent == 0 {
195            return false;
196        }
197        if self.master_read_percent == 100 {
198            return true;
199        }
200        let idx = self.counter.fetch_add(1, Ordering::Relaxed);
201        // 每 100 次循环内,前 master_read_percent 次走 master
202        (idx % 100) < self.master_read_percent as u64
203    }
204
205    /// 修改比例(运行时热更新)
206    pub fn set_percent(&mut self, percent: u8) {
207        self.master_read_percent = percent.min(100);
208        self.counter.store(0, Ordering::Relaxed);
209    }
210}
211
212impl Default for ReadRationing {
213    fn default() -> Self {
214        Self::default_slave_only()
215    }
216}
217
218/// 健康检查器:跟踪每个 slave 的健康状态,并支持故障转移
219pub struct HealthChecker {
220    states: Mutex<HashMap<String, SlaveHealth>>,
221    /// 连续失败次数阈值,达到后标记为 Unhealthy
222    pub failure_threshold: u32,
223    /// 连续失败次数计数
224    failure_counts: Mutex<HashMap<String, u32>>,
225    /// 自动故障转移恢复时间(健康检查恢复后多久重新加入集群)
226    pub recovery_cooldown: Duration,
227}
228
229impl HealthChecker {
230    pub fn new(failure_threshold: u32) -> Self {
231        Self {
232            states: Mutex::new(HashMap::new()),
233            failure_threshold,
234            failure_counts: Mutex::new(HashMap::new()),
235            recovery_cooldown: Duration::from_secs(30),
236        }
237    }
238
239    /// 注册一个 slave,初始状态为 Healthy
240    pub fn register(&self, slave: &str) {
241        if let Ok(mut states) = self.states.lock() {
242            states
243                .entry(slave.to_string())
244                .or_insert(SlaveHealth::Healthy);
245        }
246    }
247
248    /// 标记 slave 为指定健康状态
249    pub fn set_health(&self, slave: &str, health: SlaveHealth) {
250        if let Ok(mut states) = self.states.lock() {
251            states.insert(slave.to_string(), health);
252        }
253        if let Ok(mut counts) = self.failure_counts.lock() {
254            if health == SlaveHealth::Healthy {
255                counts.remove(slave);
256            }
257        }
258    }
259
260    /// 记录一次失败,达到阈值后自动标记为 Unhealthy
261    /// 返回 true 表示触发了故障转移
262    pub fn record_failure(&self, slave: &str) -> bool {
263        let mut triggered = false;
264        if let Ok(mut counts) = self.failure_counts.lock() {
265            let count = counts.entry(slave.to_string()).or_insert(0);
266            *count = count.saturating_add(1);
267            if *count >= self.failure_threshold {
268                triggered = true;
269            }
270        }
271        if triggered {
272            self.set_health(slave, SlaveHealth::Unhealthy);
273        }
274        triggered
275    }
276
277    /// 记录一次成功,重置失败计数
278    pub fn record_success(&self, slave: &str) {
279        if let Ok(mut counts) = self.failure_counts.lock() {
280            counts.remove(slave);
281        }
282    }
283
284    /// 查询 slave 健康状态,未注册返回 None
285    pub fn health(&self, slave: &str) -> Option<SlaveHealth> {
286        self.states.lock().ok().and_then(|s| s.get(slave).copied())
287    }
288
289    /// 列出所有处于指定状态的 slave
290    pub fn list_by_health(&self, health: SlaveHealth) -> Vec<String> {
291        match self.states.lock() {
292            Ok(states) => states
293                .iter()
294                .filter(|(_, h)| **h == health)
295                .map(|(k, _)| k.clone())
296                .collect(),
297            Err(_) => Vec::new(),
298        }
299    }
300
301    /// 返回所有健康 slave 列表
302    pub fn healthy_slaves(&self) -> Vec<String> {
303        self.list_by_health(SlaveHealth::Healthy)
304    }
305
306    /// 返回所有不健康 slave 列表
307    pub fn unhealthy_slaves(&self) -> Vec<String> {
308        self.list_by_health(SlaveHealth::Unhealthy)
309    }
310
311    /// 返回当前连续失败次数
312    pub fn failure_count(&self, slave: &str) -> u32 {
313        self.failure_counts
314            .lock()
315            .ok()
316            .and_then(|c| c.get(slave).copied())
317            .unwrap_or(0)
318    }
319}
320
321impl Default for HealthChecker {
322    fn default() -> Self {
323        Self::new(3)
324    }
325}
326
327/// 读写分离路由器
328///
329/// master 处理写请求,slave 集群处理读请求。
330/// 根据 `LoadBalanceStrategy` 在多个 slave 间分配读请求。
331pub struct ReadWriteRouter {
332    master: String,
333    slaves: Vec<String>,
334    strategy: LoadBalanceStrategy,
335    round_robin_counter: AtomicUsize,
336    connection_counts: Mutex<Vec<usize>>,
337    /// 加权 slave 配置(addr -> weight)
338    weights: Mutex<HashMap<String, u32>>,
339    /// 健康检查器
340    health_checker: HealthChecker,
341    /// 延迟统计
342    latency_stats: LatencyStats,
343    /// 读写比例控制器(决定是否走 master 读)
344    rationing: Mutex<ReadRationing>,
345}
346
347impl ReadWriteRouter {
348    pub fn new(master: &str, slaves: Vec<&str>) -> Self {
349        let slave_count = slaves.len();
350        let mut weights = HashMap::new();
351        let health_checker = HealthChecker::default();
352        for s in &slaves {
353            weights.insert(s.to_string(), 1u32);
354            health_checker.register(s);
355        }
356        Self {
357            master: master.to_string(),
358            slaves: slaves.into_iter().map(|s| s.to_string()).collect(),
359            strategy: LoadBalanceStrategy::RoundRobin,
360            round_robin_counter: AtomicUsize::new(0),
361            connection_counts: Mutex::new(vec![0; slave_count]),
362            weights: Mutex::new(weights),
363            health_checker,
364            latency_stats: LatencyStats::new(),
365            rationing: Mutex::new(ReadRationing::default_slave_only()),
366        }
367    }
368
369    /// 返回 master 节点
370    pub fn master(&self) -> &str {
371        &self.master
372    }
373
374    /// 返回 slave 列表
375    pub fn slaves(&self) -> &[String] {
376        &self.slaves
377    }
378
379    /// 根据当前策略选择一个 slave
380    ///
381    /// 如果 slaves 为空,回退到 master。
382    /// 如果启用健康检查,会跳过 Unhealthy/Drained 的 slave。
383    pub fn slave(&self) -> &str {
384        if self.slaves.is_empty() {
385            return &self.master;
386        }
387        // 若配置了 master 读比例,按比例分流到 master
388        if let Ok(rationing) = self.rationing.lock() {
389            if rationing.should_read_master() {
390                return &self.master;
391            }
392        }
393        // 先尝试选择健康的 slave
394        if let Some(healthy) = self.select_healthy_slave() {
395            return healthy;
396        }
397        // 所有 slave 都不健康时降级到 master(故障转移)
398        &self.master
399    }
400
401    /// 在所有健康 slave 中按策略选择
402    fn select_healthy_slave(&self) -> Option<&str> {
403        let healthy_indices: Vec<usize> = self
404            .slaves
405            .iter()
406            .enumerate()
407            .filter(|(_, s)| {
408                self.health_checker
409                    .health(s)
410                    .unwrap_or(SlaveHealth::Healthy)
411                    == SlaveHealth::Healthy
412            })
413            .map(|(i, _)| i)
414            .collect();
415
416        if healthy_indices.is_empty() {
417            return None;
418        }
419
420        let idx = match self.strategy {
421            LoadBalanceStrategy::RoundRobin => {
422                let counter = self.round_robin_counter.fetch_add(1, Ordering::SeqCst);
423                healthy_indices[counter % healthy_indices.len()]
424            }
425            LoadBalanceStrategy::Random => {
426                // 改进:使用 rand::thread_rng().gen_range 替代基于系统时间的伪随机
427                // 提供 CSPRNG 级别的随机性,避免高并发下时间熵相近导致的分布不均
428                let idx = rand::thread_rng().gen_range(0..healthy_indices.len());
429                healthy_indices[idx]
430            }
431            LoadBalanceStrategy::LeastConnections => {
432                let counts = match self.connection_counts.lock() {
433                    Ok(c) => c,
434                    Err(_) => return Some(&self.slaves[healthy_indices[0]]),
435                };
436                let mut min_idx = healthy_indices[0];
437                let mut min_count = counts[min_idx];
438                for &i in healthy_indices.iter().skip(1) {
439                    if counts[i] < min_count {
440                        min_count = counts[i];
441                        min_idx = i;
442                    }
443                }
444                min_idx
445            }
446            LoadBalanceStrategy::WeightedRoundRobin => self.select_weighted_index(&healthy_indices),
447        };
448        Some(&self.slaves[idx])
449    }
450
451    /// 加权轮询:根据权重在 healthy_indices 中选择
452    fn select_weighted_index(&self, healthy_indices: &[usize]) -> usize {
453        let weights = match self.weights.lock() {
454            Ok(w) => w,
455            Err(_) => return healthy_indices[0],
456        };
457        let total: u64 = healthy_indices
458            .iter()
459            .map(|i| weights.get(&self.slaves[*i]).copied().unwrap_or(1) as u64)
460            .sum();
461        if total == 0 {
462            return healthy_indices[0];
463        }
464        let counter = self.round_robin_counter.fetch_add(1, Ordering::SeqCst);
465        let mut pick = (counter as u64) % total;
466        for &idx in healthy_indices.iter() {
467            let w = weights.get(&self.slaves[idx]).copied().unwrap_or(1) as u64;
468            if pick < w {
469                return idx;
470            }
471            pick -= w;
472        }
473        healthy_indices[healthy_indices.len() - 1]
474    }
475
476    /// 设置负载均衡策略
477    pub fn set_strategy(&mut self, strategy: LoadBalanceStrategy) {
478        self.strategy = strategy;
479    }
480
481    /// 获取当前策略
482    pub fn strategy(&self) -> LoadBalanceStrategy {
483        self.strategy
484    }
485
486    /// 设置 slave 权重
487    pub fn set_weight(&self, slave: &str, weight: u32) -> Result<(), String> {
488        if !self.slaves.iter().any(|s| s == slave) {
489            return Err(format!("unknown slave: {}", slave));
490        }
491        if let Ok(mut weights) = self.weights.lock() {
492            weights.insert(slave.to_string(), weight.max(1));
493        }
494        Ok(())
495    }
496
497    /// 获取 slave 权重
498    pub fn weight(&self, slave: &str) -> Option<u32> {
499        self.weights.lock().ok().and_then(|w| w.get(slave).copied())
500    }
501
502    /// 健康检查器引用
503    pub fn health_checker(&self) -> &HealthChecker {
504        &self.health_checker
505    }
506
507    /// 延迟统计引用
508    pub fn latency_stats(&self) -> &LatencyStats {
509        &self.latency_stats
510    }
511
512    /// 记录一次 slave 请求的延迟
513    pub fn record_latency(&self, slave: &str, latency: Duration) {
514        self.latency_stats.record(slave, latency);
515    }
516
517    /// 测量并记录一次 slave 调用的耗时
518    pub fn measure<F, T>(&self, slave: &str, f: F) -> T
519    where
520        F: FnOnce() -> T,
521    {
522        let start = Instant::now();
523        let result = f();
524        self.record_latency(slave, start.elapsed());
525        result
526    }
527
528    /// 配置读写比例
529    pub fn set_read_rationing(&self, percent: u8) {
530        if let Ok(mut r) = self.rationing.lock() {
531            r.set_percent(percent);
532        }
533    }
534
535    /// 获取当前读写比例
536    pub fn read_rationing_percent(&self) -> u8 {
537        self.rationing
538            .lock()
539            .map(|r| r.master_read_percent)
540            .unwrap_or(0)
541    }
542
543    /// 在某个 slave 上获取连接(增加连接计数,用于 LeastConnections)
544    pub fn acquire(&self, slave: &str) -> Result<(), String> {
545        if let Some(idx) = self.slaves.iter().position(|s| s == slave) {
546            let mut counts = self
547                .connection_counts
548                .lock()
549                .map_err(|e| format!("lock error: {}", e))?;
550            counts[idx] = counts[idx].saturating_add(1);
551            Ok(())
552        } else {
553            Err(format!("unknown slave: {}", slave))
554        }
555    }
556
557    /// 释放某个 slave 的连接(减少连接计数)
558    pub fn release(&self, slave: &str) -> Result<(), String> {
559        if let Some(idx) = self.slaves.iter().position(|s| s == slave) {
560            let mut counts = self
561                .connection_counts
562                .lock()
563                .map_err(|e| format!("lock error: {}", e))?;
564            if counts[idx] > 0 {
565                counts[idx] -= 1;
566            }
567            Ok(())
568        } else {
569            Err(format!("unknown slave: {}", slave))
570        }
571    }
572
573    /// 查询某个 slave 的当前连接数(用于测试和监控)
574    pub fn connection_count(&self, slave: &str) -> Option<usize> {
575        let idx = self.slaves.iter().position(|s| s == slave)?;
576        let counts = self.connection_counts.lock().ok()?;
577        Some(counts[idx])
578    }
579}
580
581#[cfg(test)]
582mod tests {
583    use super::*;
584    use std::collections::HashSet;
585
586    #[test]
587    fn test_router_master() {
588        let router = ReadWriteRouter::new("master:3306", vec!["slave1:3306", "slave2:3306"]);
589        assert_eq!(router.master(), "master:3306");
590    }
591
592    #[test]
593    fn test_router_slaves_list() {
594        let router = ReadWriteRouter::new("m", vec!["s1", "s2", "s3"]);
595        assert_eq!(router.slaves().len(), 3);
596        assert_eq!(router.slaves()[0], "s1");
597        assert_eq!(router.slaves()[2], "s3");
598    }
599
600    #[test]
601    fn test_default_strategy_is_round_robin() {
602        let router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
603        assert_eq!(router.strategy(), LoadBalanceStrategy::RoundRobin);
604    }
605
606    #[test]
607    fn test_set_strategy() {
608        let mut router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
609        router.set_strategy(LoadBalanceStrategy::Random);
610        assert_eq!(router.strategy(), LoadBalanceStrategy::Random);
611        router.set_strategy(LoadBalanceStrategy::LeastConnections);
612        assert_eq!(router.strategy(), LoadBalanceStrategy::LeastConnections);
613    }
614
615    // --- RoundRobin 策略测试 ---
616
617    #[test]
618    fn test_round_robin_cycles_through_slaves() {
619        let mut router = ReadWriteRouter::new("m", vec!["s1", "s2", "s3"]);
620        router.set_strategy(LoadBalanceStrategy::RoundRobin);
621
622        // 第一次调用应该返回 s1,第二次 s2,第三次 s3,第四次回到 s1
623        let first = router.slave().to_string();
624        let second = router.slave().to_string();
625        let third = router.slave().to_string();
626        let fourth = router.slave().to_string();
627
628        assert_eq!(first, "s1");
629        assert_eq!(second, "s2");
630        assert_eq!(third, "s3");
631        assert_eq!(fourth, "s1", "RoundRobin 应在第 4 次回到 s1");
632    }
633
634    #[test]
635    fn test_round_robin_single_slave() {
636        let mut router = ReadWriteRouter::new("m", vec!["only_slave"]);
637        router.set_strategy(LoadBalanceStrategy::RoundRobin);
638        for _ in 0..5 {
639            assert_eq!(router.slave(), "only_slave");
640        }
641    }
642
643    #[test]
644    fn test_round_robin_visits_all_slaves() {
645        let mut router = ReadWriteRouter::new("m", vec!["s1", "s2", "s3", "s4"]);
646        router.set_strategy(LoadBalanceStrategy::RoundRobin);
647
648        let mut visited = HashSet::new();
649        for _ in 0..4 {
650            visited.insert(router.slave().to_string());
651        }
652        assert_eq!(visited.len(), 4, "一轮轮询应该访问所有 4 个 slave");
653    }
654
655    // --- Random 策略测试 ---
656
657    #[test]
658    fn test_random_returns_valid_slave() {
659        let mut router = ReadWriteRouter::new("m", vec!["s1", "s2", "s3"]);
660        router.set_strategy(LoadBalanceStrategy::Random);
661
662        let slaves: HashSet<&str> = ["s1", "s2", "s3"].iter().copied().collect();
663        for _ in 0..20 {
664            let picked = router.slave();
665            assert!(
666                slaves.contains(picked),
667                "随机策略返回了未知 slave: {}",
668                picked
669            );
670        }
671    }
672
673    #[test]
674    fn test_random_eventually_visits_multiple_slaves() {
675        let mut router = ReadWriteRouter::new("m", vec!["s1", "s2", "s3", "s4"]);
676        router.set_strategy(LoadBalanceStrategy::Random);
677
678        let mut visited = HashSet::new();
679        // 大量调用后,应该至少访问 2 个不同的 slave(统计性验证)
680        for _ in 0..200 {
681            visited.insert(router.slave().to_string());
682        }
683        assert!(
684            visited.len() >= 2,
685            "随机策略在 200 次调用后应至少访问 2 个 slave,实际: {}",
686            visited.len()
687        );
688    }
689
690    #[test]
691    fn test_random_single_slave() {
692        let mut router = ReadWriteRouter::new("m", vec!["only"]);
693        router.set_strategy(LoadBalanceStrategy::Random);
694        for _ in 0..10 {
695            assert_eq!(router.slave(), "only");
696        }
697    }
698
699    // --- LeastConnections 策略测试 ---
700
701    #[test]
702    fn test_least_connections_picks_zero_load_slave() {
703        let mut router = ReadWriteRouter::new("m", vec!["s1", "s2", "s3"]);
704        router.set_strategy(LoadBalanceStrategy::LeastConnections);
705
706        // 初始所有 slave 连接数都为 0,应选第一个
707        assert_eq!(router.slave(), "s1");
708    }
709
710    #[test]
711    fn test_least_connections_picks_least_loaded() {
712        let mut router = ReadWriteRouter::new("m", vec!["s1", "s2", "s3"]);
713        router.set_strategy(LoadBalanceStrategy::LeastConnections);
714
715        // 给 s1 和 s2 增加连接
716        router.acquire("s1").unwrap();
717        router.acquire("s1").unwrap();
718        router.acquire("s2").unwrap();
719
720        // s3 连接数为 0,应被选中
721        assert_eq!(router.slave(), "s3");
722
723        // 给 s3 也加一个连接
724        router.acquire("s3").unwrap();
725        // 现在 s1=2, s2=1, s3=1,应选 s2(索引靠前的最少连接)
726        assert_eq!(router.slave(), "s2");
727    }
728
729    #[test]
730    fn test_least_connections_after_release() {
731        let mut router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
732        router.set_strategy(LoadBalanceStrategy::LeastConnections);
733
734        router.acquire("s1").unwrap();
735        router.acquire("s1").unwrap();
736        router.acquire("s2").unwrap();
737
738        // s1=2, s2=1,选 s2
739        assert_eq!(router.slave(), "s2");
740
741        // 释放 s1 的两个连接
742        router.release("s1").unwrap();
743        router.release("s1").unwrap();
744
745        // 现在 s1=0, s2=1,应选 s1
746        assert_eq!(router.slave(), "s1");
747    }
748
749    #[test]
750    fn test_acquire_unknown_slave_returns_error() {
751        let router = ReadWriteRouter::new("m", vec!["s1"]);
752        assert!(router.acquire("nonexistent").is_err());
753    }
754
755    #[test]
756    fn test_release_unknown_slave_returns_error() {
757        let router = ReadWriteRouter::new("m", vec!["s1"]);
758        assert!(router.release("nonexistent").is_err());
759    }
760
761    #[test]
762    fn test_release_below_zero_clamped() {
763        let router = ReadWriteRouter::new("m", vec!["s1"]);
764        // 释放不存在的连接不应导致负数
765        router.release("s1").unwrap();
766        assert_eq!(router.connection_count("s1"), Some(0));
767    }
768
769    #[test]
770    fn test_connection_count_query() {
771        let router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
772        assert_eq!(router.connection_count("s1"), Some(0));
773        router.acquire("s1").unwrap();
774        router.acquire("s1").unwrap();
775        assert_eq!(router.connection_count("s1"), Some(2));
776        assert_eq!(router.connection_count("s2"), Some(0));
777        assert_eq!(router.connection_count("unknown"), None);
778    }
779
780    #[test]
781    fn test_empty_slaves_falls_back_to_master() {
782        let router = ReadWriteRouter::new("only_master", vec![]);
783        // 任何策略下都应回退到 master
784        assert_eq!(router.slave(), "only_master");
785
786        let mut router_rr = router;
787        router_rr.set_strategy(LoadBalanceStrategy::RoundRobin);
788        assert_eq!(router_rr.slave(), "only_master");
789
790        router_rr.set_strategy(LoadBalanceStrategy::Random);
791        assert_eq!(router_rr.slave(), "only_master");
792
793        router_rr.set_strategy(LoadBalanceStrategy::LeastConnections);
794        assert_eq!(router_rr.slave(), "only_master");
795    }
796
797    #[test]
798    fn test_round_robin_concurrent_safe() {
799        // 验证 AtomicUsize 在多线程下不会 panic,且计数正确
800        // 默认策略即为 RoundRobin
801        let router = std::sync::Arc::new(ReadWriteRouter::new("m", vec!["s1", "s2", "s3"]));
802
803        let mut handles = vec![];
804        for _ in 0..4 {
805            let r = std::sync::Arc::clone(&router);
806            handles.push(std::thread::spawn(move || {
807                for _ in 0..10 {
808                    let _ = r.slave();
809                }
810            }));
811        }
812        for h in handles {
813            h.join().unwrap();
814        }
815        // 40 次调用后计数器应该是 40
816        assert_eq!(router.round_robin_counter.load(Ordering::SeqCst), 40);
817    }
818
819    // ====================================================================
820    // 健康检查与故障转移测试
821    // ====================================================================
822
823    #[test]
824    fn test_health_checker_new_slave_is_healthy() {
825        let checker = HealthChecker::new(3);
826        checker.register("s1");
827        assert_eq!(checker.health("s1"), Some(SlaveHealth::Healthy));
828    }
829
830    #[test]
831    fn test_health_checker_unregistered_slave_returns_none() {
832        let checker = HealthChecker::new(3);
833        assert_eq!(checker.health("unknown"), None);
834    }
835
836    #[test]
837    fn test_health_checker_mark_unhealthy() {
838        let checker = HealthChecker::new(3);
839        checker.register("s1");
840        checker.set_health("s1", SlaveHealth::Unhealthy);
841        assert_eq!(checker.health("s1"), Some(SlaveHealth::Unhealthy));
842    }
843
844    #[test]
845    fn test_health_checker_mark_drained() {
846        let checker = HealthChecker::new(3);
847        checker.register("s1");
848        checker.set_health("s1", SlaveHealth::Drained);
849        assert_eq!(checker.health("s1"), Some(SlaveHealth::Drained));
850    }
851
852    #[test]
853    fn test_record_failure_below_threshold_keeps_healthy() {
854        let checker = HealthChecker::new(3);
855        checker.register("s1");
856        // 失败 2 次(< 阈值 3)仍应保持健康
857        assert!(!checker.record_failure("s1"));
858        assert!(!checker.record_failure("s1"));
859        assert_eq!(checker.health("s1"), Some(SlaveHealth::Healthy));
860        assert_eq!(checker.failure_count("s1"), 2);
861    }
862
863    #[test]
864    fn test_record_failure_at_threshold_marks_unhealthy() {
865        let checker = HealthChecker::new(3);
866        checker.register("s1");
867        assert!(!checker.record_failure("s1"));
868        assert!(!checker.record_failure("s1"));
869        // 第 3 次失败应触发故障转移
870        assert!(checker.record_failure("s1"));
871        assert_eq!(checker.health("s1"), Some(SlaveHealth::Unhealthy));
872    }
873
874    #[test]
875    fn test_record_success_resets_failure_count() {
876        let checker = HealthChecker::new(3);
877        checker.register("s1");
878        checker.record_failure("s1");
879        checker.record_failure("s1");
880        assert_eq!(checker.failure_count("s1"), 2);
881        checker.record_success("s1");
882        assert_eq!(checker.failure_count("s1"), 0);
883    }
884
885    #[test]
886    fn test_set_healthy_resets_failure_count() {
887        let checker = HealthChecker::new(2);
888        checker.register("s1");
889        checker.record_failure("s1");
890        checker.record_failure("s1");
891        assert_eq!(checker.health("s1"), Some(SlaveHealth::Unhealthy));
892        // 恢复
893        checker.set_health("s1", SlaveHealth::Healthy);
894        assert_eq!(checker.failure_count("s1"), 0);
895    }
896
897    #[test]
898    fn test_list_by_health() {
899        let checker = HealthChecker::new(3);
900        checker.register("s1");
901        checker.register("s2");
902        checker.register("s3");
903        checker.set_health("s2", SlaveHealth::Unhealthy);
904        let mut healthy = checker.healthy_slaves();
905        healthy.sort();
906        assert_eq!(healthy, vec!["s1".to_string(), "s3".to_string()]);
907        assert_eq!(checker.unhealthy_slaves(), vec!["s2".to_string()]);
908    }
909
910    #[test]
911    fn test_router_failover_to_master_when_all_unhealthy() {
912        let router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
913        router
914            .health_checker()
915            .set_health("s1", SlaveHealth::Unhealthy);
916        router
917            .health_checker()
918            .set_health("s2", SlaveHealth::Unhealthy);
919        // 全部不健康时降级到 master
920        assert_eq!(router.slave(), "m");
921    }
922
923    #[test]
924    fn test_router_skips_unhealthy_slave() {
925        let mut router = ReadWriteRouter::new("m", vec!["s1", "s2", "s3"]);
926        router.set_strategy(LoadBalanceStrategy::RoundRobin);
927        router
928            .health_checker()
929            .set_health("s2", SlaveHealth::Unhealthy);
930
931        // 100 次调用都不应命中 s2
932        for _ in 0..100 {
933            let picked = router.slave().to_string();
934            assert_ne!(picked, "s2", "不应选中不健康的 slave");
935        }
936    }
937
938    #[test]
939    fn test_router_failover_skips_drained_slave() {
940        let router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
941        router
942            .health_checker()
943            .set_health("s1", SlaveHealth::Drained);
944        // 只有 s2 健康
945        for _ in 0..5 {
946            assert_eq!(router.slave(), "s2");
947        }
948    }
949
950    #[test]
951    fn test_router_default_health_checker_threshold_is_3() {
952        let router = ReadWriteRouter::new("m", vec!["s1"]);
953        assert_eq!(router.health_checker().failure_threshold, 3);
954    }
955
956    // ====================================================================
957    // 权重配置测试
958    // ====================================================================
959
960    #[test]
961    fn test_set_weight_for_known_slave() {
962        let router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
963        router.set_weight("s1", 10).unwrap();
964        assert_eq!(router.weight("s1"), Some(10));
965        assert_eq!(router.weight("s2"), Some(1));
966    }
967
968    #[test]
969    fn test_set_weight_for_unknown_slave_errors() {
970        let router = ReadWriteRouter::new("m", vec!["s1"]);
971        assert!(router.set_weight("ghost", 10).is_err());
972    }
973
974    #[test]
975    fn test_set_weight_zero_clamped_to_one() {
976        let router = ReadWriteRouter::new("m", vec!["s1"]);
977        router.set_weight("s1", 0).unwrap();
978        assert_eq!(router.weight("s1"), Some(1));
979    }
980
981    #[test]
982    fn test_weighted_round_robin_respects_weights() {
983        let mut router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
984        router.set_strategy(LoadBalanceStrategy::WeightedRoundRobin);
985        router.set_weight("s1", 9).unwrap();
986        router.set_weight("s2", 1).unwrap();
987
988        // 100 次调用后,s1 应该被选中的次数 >> s2
989        let mut s1_count = 0usize;
990        let mut s2_count = 0usize;
991        for _ in 0..100 {
992            match router.slave() {
993                "s1" => s1_count += 1,
994                "s2" => s2_count += 1,
995                _ => {}
996            }
997        }
998        assert_eq!(s1_count + s2_count, 100);
999        assert!(
1000            s1_count > s2_count * 3,
1001            "权重 9:1 应使 s1 命中次数远多于 s2,实际 s1={}, s2={}",
1002            s1_count,
1003            s2_count
1004        );
1005    }
1006
1007    #[test]
1008    fn test_weighted_round_robin_with_equal_weights_visits_all() {
1009        let mut router = ReadWriteRouter::new("m", vec!["s1", "s2", "s3"]);
1010        router.set_strategy(LoadBalanceStrategy::WeightedRoundRobin);
1011
1012        let mut visited = HashSet::new();
1013        for _ in 0..30 {
1014            visited.insert(router.slave().to_string());
1015        }
1016        assert_eq!(visited.len(), 3);
1017    }
1018
1019    #[test]
1020    fn test_weighted_round_robin_skips_unhealthy() {
1021        let mut router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
1022        router.set_strategy(LoadBalanceStrategy::WeightedRoundRobin);
1023        router.set_weight("s1", 100).unwrap();
1024        router.set_weight("s2", 1).unwrap();
1025        router
1026            .health_checker()
1027            .set_health("s1", SlaveHealth::Unhealthy);
1028
1029        // s1 不健康,所有请求应落到 s2
1030        for _ in 0..10 {
1031            assert_eq!(router.slave(), "s2");
1032        }
1033    }
1034
1035    // ====================================================================
1036    // 读写比例控制测试
1037    // ====================================================================
1038
1039    #[test]
1040    fn test_read_rationing_default_is_zero_percent_master() {
1041        let router = ReadWriteRouter::new("m", vec!["s1"]);
1042        assert_eq!(router.read_rationing_percent(), 0);
1043    }
1044
1045    #[test]
1046    fn test_read_rationing_all_master() {
1047        let router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
1048        router.set_read_rationing(100);
1049        // 100% 走 master
1050        for _ in 0..10 {
1051            assert_eq!(router.slave(), "m");
1052        }
1053    }
1054
1055    #[test]
1056    fn test_read_rationing_all_slave() {
1057        let router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
1058        router.set_read_rationing(0);
1059        // 0% 走 master,全部走 slave
1060        for _ in 0..10 {
1061            let picked = router.slave();
1062            assert!(picked == "s1" || picked == "s2");
1063        }
1064    }
1065
1066    #[test]
1067    fn test_read_rationing_partial_distribution() {
1068        let router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
1069        router.set_read_rationing(30);
1070
1071        let mut master_count = 0usize;
1072        let mut slave_count = 0usize;
1073        for _ in 0..100 {
1074            let picked = router.slave();
1075            if picked == "m" {
1076                master_count += 1;
1077            } else {
1078                slave_count += 1;
1079            }
1080        }
1081        // 30% 应该走 master,允许 ±5 浮动
1082        assert!(
1083            (25..=35).contains(&master_count),
1084            "30% 比例下 master 命中数应在 25-35 之间,实际: {}",
1085            master_count
1086        );
1087        assert_eq!(master_count + slave_count, 100);
1088    }
1089
1090    #[test]
1091    fn test_read_rationing_clamps_above_100() {
1092        let r = ReadRationing::new(150);
1093        assert_eq!(r.master_read_percent, 100);
1094    }
1095
1096    #[test]
1097    fn test_read_rationing_set_percent_resets_counter() {
1098        let mut r = ReadRationing::new(50);
1099        // 调用几次
1100        for _ in 0..5 {
1101            let _ = r.should_read_master();
1102        }
1103        r.set_percent(80);
1104        // counter 应该被重置
1105        assert_eq!(r.counter.load(Ordering::Relaxed), 0);
1106    }
1107
1108    #[test]
1109    fn test_read_rationing_master_only_always_returns_true() {
1110        let r = ReadRationing::default_master_only();
1111        for _ in 0..10 {
1112            assert!(r.should_read_master());
1113        }
1114    }
1115
1116    #[test]
1117    fn test_read_rationing_slave_only_always_returns_false() {
1118        let r = ReadRationing::default_slave_only();
1119        for _ in 0..10 {
1120            assert!(!r.should_read_master());
1121        }
1122    }
1123
1124    // ====================================================================
1125    // 延迟统计测试
1126    // ====================================================================
1127
1128    #[test]
1129    fn test_latency_stats_record_and_snapshot() {
1130        let stats = LatencyStats::new();
1131        stats.record("s1", Duration::from_millis(10));
1132        stats.record("s1", Duration::from_millis(20));
1133        stats.record("s1", Duration::from_millis(30));
1134
1135        let snap = stats.snapshot("s1");
1136        assert_eq!(snap.samples, 3);
1137        assert!(snap.min_ns > 0);
1138        assert!(snap.max_ns >= snap.min_ns);
1139        // 平均应该在 10ms-30ms 之间
1140        let avg = snap.avg();
1141        assert!(avg >= Duration::from_millis(9));
1142        assert!(avg <= Duration::from_millis(31));
1143    }
1144
1145    #[test]
1146    fn test_latency_stats_unknown_slave_returns_default() {
1147        let stats = LatencyStats::new();
1148        let snap = stats.snapshot("ghost");
1149        assert_eq!(snap.samples, 0);
1150        assert_eq!(snap.avg_ns(), 0);
1151    }
1152
1153    #[test]
1154    fn test_latency_stats_reset() {
1155        let stats = LatencyStats::new();
1156        stats.record("s1", Duration::from_millis(10));
1157        assert_eq!(stats.snapshot("s1").samples, 1);
1158        stats.reset("s1");
1159        assert_eq!(stats.snapshot("s1").samples, 0);
1160    }
1161
1162    #[test]
1163    fn test_latency_stats_all_returns_all_slaves() {
1164        let stats = LatencyStats::new();
1165        stats.record("s1", Duration::from_millis(10));
1166        stats.record("s2", Duration::from_millis(20));
1167        let all = stats.all();
1168        assert_eq!(all.len(), 2);
1169    }
1170
1171    #[test]
1172    fn test_router_measure_records_latency() {
1173        let router = ReadWriteRouter::new("m", vec!["s1"]);
1174        let result = router.measure("s1", || 42);
1175        assert_eq!(result, 42);
1176        let snap = router.latency_stats().snapshot("s1");
1177        assert_eq!(snap.samples, 1);
1178        // min_ns 可能在高分辨率计时器上为 0(闭包执行极快),不强制大于 0
1179    }
1180
1181    #[test]
1182    fn test_router_record_latency_increases_sample_count() {
1183        let router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
1184        router.record_latency("s1", Duration::from_micros(100));
1185        router.record_latency("s1", Duration::from_micros(200));
1186        router.record_latency("s2", Duration::from_micros(50));
1187
1188        assert_eq!(router.latency_stats().snapshot("s1").samples, 2);
1189        assert_eq!(router.latency_stats().snapshot("s2").samples, 1);
1190    }
1191
1192    #[test]
1193    fn test_latency_snapshot_avg_with_zero_samples() {
1194        let snap = LatencySnapshot::default();
1195        assert_eq!(snap.avg_ns(), 0);
1196        assert_eq!(snap.avg(), Duration::ZERO);
1197    }
1198
1199    #[test]
1200    fn test_latency_snapshot_min_updates() {
1201        let mut snap = LatencySnapshot::default();
1202        snap.record(Duration::from_millis(50));
1203        assert_eq!(snap.min_ns, 50_000_000);
1204        snap.record(Duration::from_millis(10));
1205        assert_eq!(snap.min_ns, 10_000_000);
1206        snap.record(Duration::from_millis(100));
1207        assert_eq!(snap.min_ns, 10_000_000);
1208    }
1209
1210    #[test]
1211    fn test_latency_snapshot_max_updates() {
1212        let mut snap = LatencySnapshot::default();
1213        snap.record(Duration::from_millis(10));
1214        assert_eq!(snap.max_ns, 10_000_000);
1215        snap.record(Duration::from_millis(50));
1216        assert_eq!(snap.max_ns, 50_000_000);
1217        snap.record(Duration::from_millis(20));
1218        assert_eq!(snap.max_ns, 50_000_000);
1219    }
1220
1221    #[test]
1222    fn test_weighted_slave_default_health_is_healthy() {
1223        let ws = WeightedSlave::new("s1:3306", 5);
1224        assert_eq!(ws.health, SlaveHealth::Healthy);
1225        assert_eq!(ws.weight, 5);
1226    }
1227
1228    #[test]
1229    fn test_weighted_slave_zero_weight_clamped() {
1230        let ws = WeightedSlave::new("s1:3306", 0);
1231        assert_eq!(ws.weight, 1);
1232    }
1233
1234    #[test]
1235    fn test_weighted_slave_with_health_builder() {
1236        let ws = WeightedSlave::new("s1:3306", 5).with_health(SlaveHealth::Unhealthy);
1237        assert_eq!(ws.health, SlaveHealth::Unhealthy);
1238    }
1239}