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#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
24pub enum LoadBalanceStrategy {
25 RoundRobin,
27 Random,
29 LeastConnections,
31 WeightedRoundRobin,
33}
34
35#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
37pub enum SlaveHealth {
38 Healthy,
40 Unhealthy,
42 Drained,
44}
45
46#[derive(Debug, Clone, Serialize, Deserialize)]
48pub struct WeightedSlave {
49 pub addr: String,
51 pub weight: u32,
53 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#[derive(Debug, Clone, Default, Serialize, Deserialize)]
74pub struct LatencySnapshot {
75 pub samples: u64,
77 pub min_ns: u64,
79 pub max_ns: u64,
81 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 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#[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 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 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 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 pub fn reset(&self, slave: &str) {
148 if let Ok(mut inner) = self.inner.lock() {
149 inner.remove(slave);
150 }
151 }
152}
153
154#[derive(Debug, Serialize, Deserialize)]
156pub struct ReadRationing {
157 pub master_read_percent: u8,
160 #[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 pub fn default_slave_only() -> Self {
184 Self::new(0)
185 }
186
187 pub fn default_master_only() -> Self {
189 Self::new(100)
190 }
191
192 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 (idx % 100) < self.master_read_percent as u64
203 }
204
205 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
218pub struct HealthChecker {
220 states: Mutex<HashMap<String, SlaveHealth>>,
221 pub failure_threshold: u32,
223 failure_counts: Mutex<HashMap<String, u32>>,
225 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 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 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 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 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 pub fn health(&self, slave: &str) -> Option<SlaveHealth> {
286 self.states.lock().ok().and_then(|s| s.get(slave).copied())
287 }
288
289 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 pub fn healthy_slaves(&self) -> Vec<String> {
303 self.list_by_health(SlaveHealth::Healthy)
304 }
305
306 pub fn unhealthy_slaves(&self) -> Vec<String> {
308 self.list_by_health(SlaveHealth::Unhealthy)
309 }
310
311 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
327pub struct ReadWriteRouter {
332 master: String,
333 slaves: Vec<String>,
334 strategy: LoadBalanceStrategy,
335 round_robin_counter: AtomicUsize,
336 connection_counts: Mutex<Vec<usize>>,
337 weights: Mutex<HashMap<String, u32>>,
339 health_checker: HealthChecker,
341 latency_stats: LatencyStats,
343 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 pub fn master(&self) -> &str {
371 &self.master
372 }
373
374 pub fn slaves(&self) -> &[String] {
376 &self.slaves
377 }
378
379 pub fn slave(&self) -> &str {
384 if self.slaves.is_empty() {
385 return &self.master;
386 }
387 if let Ok(rationing) = self.rationing.lock() {
389 if rationing.should_read_master() {
390 return &self.master;
391 }
392 }
393 if let Some(healthy) = self.select_healthy_slave() {
395 return healthy;
396 }
397 &self.master
399 }
400
401 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 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 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 pub fn set_strategy(&mut self, strategy: LoadBalanceStrategy) {
478 self.strategy = strategy;
479 }
480
481 pub fn strategy(&self) -> LoadBalanceStrategy {
483 self.strategy
484 }
485
486 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 pub fn weight(&self, slave: &str) -> Option<u32> {
499 self.weights.lock().ok().and_then(|w| w.get(slave).copied())
500 }
501
502 pub fn health_checker(&self) -> &HealthChecker {
504 &self.health_checker
505 }
506
507 pub fn latency_stats(&self) -> &LatencyStats {
509 &self.latency_stats
510 }
511
512 pub fn record_latency(&self, slave: &str, latency: Duration) {
514 self.latency_stats.record(slave, latency);
515 }
516
517 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 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 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 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 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 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 #[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 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 #[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 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 #[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 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 router.acquire("s1").unwrap();
717 router.acquire("s1").unwrap();
718 router.acquire("s2").unwrap();
719
720 assert_eq!(router.slave(), "s3");
722
723 router.acquire("s3").unwrap();
725 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 assert_eq!(router.slave(), "s2");
740
741 router.release("s1").unwrap();
743 router.release("s1").unwrap();
744
745 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 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 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 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 assert_eq!(router.round_robin_counter.load(Ordering::SeqCst), 40);
817 }
818
819 #[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 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 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 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 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 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 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 #[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 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 for _ in 0..10 {
1031 assert_eq!(router.slave(), "s2");
1032 }
1033 }
1034
1035 #[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 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 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 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 for _ in 0..5 {
1101 let _ = r.should_read_master();
1102 }
1103 r.set_percent(80);
1104 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 #[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 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 }
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}