1use 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#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
30pub enum LoadBalanceStrategy {
31 RoundRobin,
33 Random,
35 LeastConnections,
37 WeightedRoundRobin,
39}
40
41#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
43pub enum SlaveHealth {
44 Healthy,
46 Unhealthy,
48 Drained,
50}
51
52#[derive(Debug, Clone, Serialize, Deserialize)]
54pub struct WeightedSlave {
55 pub addr: String,
57 pub weight: u32,
59 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#[derive(Debug, Clone, Default, Serialize, Deserialize)]
80pub struct LatencySnapshot {
81 pub samples: u64,
83 pub min_ns: u64,
85 pub max_ns: u64,
87 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 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#[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 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 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 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 pub fn reset(&self, slave: &str) {
154 if let Ok(mut inner) = self.inner.lock() {
155 inner.remove(slave);
156 }
157 }
158}
159
160#[derive(Debug, Serialize, Deserialize)]
162pub struct ReadRationing {
163 pub master_read_percent: u8,
166 #[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 pub fn default_slave_only() -> Self {
190 Self::new(0)
191 }
192
193 pub fn default_master_only() -> Self {
195 Self::new(100)
196 }
197
198 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 (idx % 100) < self.master_read_percent as u64
209 }
210
211 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
224pub struct HealthChecker {
226 states: Mutex<HashMap<String, SlaveHealth>>,
227 pub failure_threshold: u32,
229 failure_counts: Mutex<HashMap<String, u32>>,
231 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 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 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 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 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 pub fn health(&self, slave: &str) -> Option<SlaveHealth> {
292 self.states.lock().ok().and_then(|s| s.get(slave).copied())
293 }
294
295 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 pub fn healthy_slaves(&self) -> Vec<String> {
309 self.list_by_health(SlaveHealth::Healthy)
310 }
311
312 pub fn unhealthy_slaves(&self) -> Vec<String> {
314 self.list_by_health(SlaveHealth::Unhealthy)
315 }
316
317 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
333pub struct ReadWriteRouter {
338 master: String,
339 slaves: Vec<String>,
340 strategy: LoadBalanceStrategy,
341 round_robin_counter: AtomicUsize,
342 connection_counts: Mutex<Vec<usize>>,
343 weights: Mutex<HashMap<String, u32>>,
345 health_checker: HealthChecker,
347 latency_stats: LatencyStats,
349 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 pub fn master(&self) -> &str {
377 &self.master
378 }
379
380 pub fn slaves(&self) -> &[String] {
382 &self.slaves
383 }
384
385 pub fn slave(&self) -> &str {
390 if self.slaves.is_empty() {
391 return &self.master;
392 }
393 if let Ok(rationing) = self.rationing.lock() {
395 if rationing.should_read_master() {
396 return &self.master;
397 }
398 }
399 if let Some(healthy) = self.select_healthy_slave() {
401 return healthy;
402 }
403 &self.master
405 }
406
407 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 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 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 pub fn set_strategy(&mut self, strategy: LoadBalanceStrategy) {
484 self.strategy = strategy;
485 }
486
487 pub fn strategy(&self) -> LoadBalanceStrategy {
489 self.strategy
490 }
491
492 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 pub fn weight(&self, slave: &str) -> Option<u32> {
505 self.weights.lock().ok().and_then(|w| w.get(slave).copied())
506 }
507
508 pub fn health_checker(&self) -> &HealthChecker {
510 &self.health_checker
511 }
512
513 pub fn latency_stats(&self) -> &LatencyStats {
515 &self.latency_stats
516 }
517
518 pub fn record_latency(&self, slave: &str, latency: Duration) {
520 self.latency_stats.record(slave, latency);
521 }
522
523 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 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 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 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 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 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 #[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 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 #[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 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 #[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 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 router.acquire("s1").unwrap();
723 router.acquire("s1").unwrap();
724 router.acquire("s2").unwrap();
725
726 assert_eq!(router.slave(), "s3");
728
729 router.acquire("s3").unwrap();
731 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 assert_eq!(router.slave(), "s2");
746
747 router.release("s1").unwrap();
749 router.release("s1").unwrap();
750
751 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 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 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 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 assert_eq!(router.round_robin_counter.load(Ordering::SeqCst), 40);
823 }
824
825 #[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 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 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 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 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 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 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 #[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 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 for _ in 0..10 {
1037 assert_eq!(router.slave(), "s2");
1038 }
1039 }
1040
1041 #[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 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 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 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 for _ in 0..5 {
1107 let _ = r.should_read_master();
1108 }
1109 r.set_percent(80);
1110 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 #[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 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 }
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}