Skip to main content

sz_orm_rw/
lib.rs

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