use rand::Rng;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::Mutex;
use std::time::{Duration, Instant};
#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
pub enum LoadBalanceStrategy {
RoundRobin,
Random,
LeastConnections,
WeightedRoundRobin,
}
#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
pub enum SlaveHealth {
Healthy,
Unhealthy,
Drained,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WeightedSlave {
pub addr: String,
pub weight: u32,
pub health: SlaveHealth,
}
impl WeightedSlave {
pub fn new(addr: impl Into<String>, weight: u32) -> Self {
Self {
addr: addr.into(),
weight: weight.max(1),
health: SlaveHealth::Healthy,
}
}
pub fn with_health(mut self, health: SlaveHealth) -> Self {
self.health = health;
self
}
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct LatencySnapshot {
pub samples: u64,
pub min_ns: u64,
pub max_ns: u64,
pub sum_ns: u128,
}
impl LatencySnapshot {
pub fn record(&mut self, latency: Duration) {
let ns = latency.as_nanos();
self.samples += 1;
if self.min_ns == 0 || ns < self.min_ns as u128 {
self.min_ns = ns.min(u64::MAX as u128) as u64;
}
if ns > self.max_ns as u128 {
self.max_ns = ns.min(u64::MAX as u128) as u64;
}
self.sum_ns = self.sum_ns.saturating_add(ns);
}
pub fn avg_ns(&self) -> u64 {
if self.samples == 0 {
0
} else {
(self.sum_ns / self.samples as u128) as u64
}
}
pub fn avg(&self) -> Duration {
Duration::from_nanos(self.avg_ns())
}
}
#[derive(Debug, Default)]
pub struct LatencyStats {
inner: Mutex<HashMap<String, LatencySnapshot>>,
}
impl LatencyStats {
pub fn new() -> Self {
Self::default()
}
pub fn record(&self, slave: &str, latency: Duration) {
if let Ok(mut inner) = self.inner.lock() {
inner.entry(slave.to_string()).or_default().record(latency);
}
}
pub fn snapshot(&self, slave: &str) -> LatencySnapshot {
match self.inner.lock() {
Ok(inner) => inner.get(slave).cloned().unwrap_or_default(),
Err(_) => LatencySnapshot::default(),
}
}
pub fn all(&self) -> Vec<(String, LatencySnapshot)> {
match self.inner.lock() {
Ok(inner) => inner.iter().map(|(k, v)| (k.clone(), v.clone())).collect(),
Err(_) => Vec::new(),
}
}
pub fn reset(&self, slave: &str) {
if let Ok(mut inner) = self.inner.lock() {
inner.remove(slave);
}
}
}
#[derive(Debug, Serialize, Deserialize)]
pub struct ReadRationing {
pub master_read_percent: u8,
#[serde(skip)]
counter: AtomicU64,
}
impl Clone for ReadRationing {
fn clone(&self) -> Self {
Self {
master_read_percent: self.master_read_percent,
counter: AtomicU64::new(self.counter.load(Ordering::Relaxed)),
}
}
}
impl ReadRationing {
pub fn new(master_read_percent: u8) -> Self {
Self {
master_read_percent: master_read_percent.min(100),
counter: AtomicU64::new(0),
}
}
pub fn default_slave_only() -> Self {
Self::new(0)
}
pub fn default_master_only() -> Self {
Self::new(100)
}
pub fn should_read_master(&self) -> bool {
if self.master_read_percent == 0 {
return false;
}
if self.master_read_percent == 100 {
return true;
}
let idx = self.counter.fetch_add(1, Ordering::Relaxed);
(idx % 100) < self.master_read_percent as u64
}
pub fn set_percent(&mut self, percent: u8) {
self.master_read_percent = percent.min(100);
self.counter.store(0, Ordering::Relaxed);
}
}
impl Default for ReadRationing {
fn default() -> Self {
Self::default_slave_only()
}
}
pub struct HealthChecker {
states: Mutex<HashMap<String, SlaveHealth>>,
pub failure_threshold: u32,
failure_counts: Mutex<HashMap<String, u32>>,
pub recovery_cooldown: Duration,
}
impl HealthChecker {
pub fn new(failure_threshold: u32) -> Self {
Self {
states: Mutex::new(HashMap::new()),
failure_threshold,
failure_counts: Mutex::new(HashMap::new()),
recovery_cooldown: Duration::from_secs(30),
}
}
pub fn register(&self, slave: &str) {
if let Ok(mut states) = self.states.lock() {
states
.entry(slave.to_string())
.or_insert(SlaveHealth::Healthy);
}
}
pub fn set_health(&self, slave: &str, health: SlaveHealth) {
if let Ok(mut states) = self.states.lock() {
states.insert(slave.to_string(), health);
}
if let Ok(mut counts) = self.failure_counts.lock() {
if health == SlaveHealth::Healthy {
counts.remove(slave);
}
}
}
pub fn record_failure(&self, slave: &str) -> bool {
let mut triggered = false;
if let Ok(mut counts) = self.failure_counts.lock() {
let count = counts.entry(slave.to_string()).or_insert(0);
*count = count.saturating_add(1);
if *count >= self.failure_threshold {
triggered = true;
}
}
if triggered {
self.set_health(slave, SlaveHealth::Unhealthy);
}
triggered
}
pub fn record_success(&self, slave: &str) {
if let Ok(mut counts) = self.failure_counts.lock() {
counts.remove(slave);
}
}
pub fn health(&self, slave: &str) -> Option<SlaveHealth> {
self.states.lock().ok().and_then(|s| s.get(slave).copied())
}
pub fn list_by_health(&self, health: SlaveHealth) -> Vec<String> {
match self.states.lock() {
Ok(states) => states
.iter()
.filter(|(_, h)| **h == health)
.map(|(k, _)| k.clone())
.collect(),
Err(_) => Vec::new(),
}
}
pub fn healthy_slaves(&self) -> Vec<String> {
self.list_by_health(SlaveHealth::Healthy)
}
pub fn unhealthy_slaves(&self) -> Vec<String> {
self.list_by_health(SlaveHealth::Unhealthy)
}
pub fn failure_count(&self, slave: &str) -> u32 {
self.failure_counts
.lock()
.ok()
.and_then(|c| c.get(slave).copied())
.unwrap_or(0)
}
}
impl Default for HealthChecker {
fn default() -> Self {
Self::new(3)
}
}
pub struct ReadWriteRouter {
master: String,
slaves: Vec<String>,
strategy: LoadBalanceStrategy,
round_robin_counter: AtomicUsize,
connection_counts: Mutex<Vec<usize>>,
weights: Mutex<HashMap<String, u32>>,
health_checker: HealthChecker,
latency_stats: LatencyStats,
rationing: Mutex<ReadRationing>,
}
impl ReadWriteRouter {
pub fn new(master: &str, slaves: Vec<&str>) -> Self {
let slave_count = slaves.len();
let mut weights = HashMap::new();
let health_checker = HealthChecker::default();
for s in &slaves {
weights.insert(s.to_string(), 1u32);
health_checker.register(s);
}
Self {
master: master.to_string(),
slaves: slaves.into_iter().map(|s| s.to_string()).collect(),
strategy: LoadBalanceStrategy::RoundRobin,
round_robin_counter: AtomicUsize::new(0),
connection_counts: Mutex::new(vec![0; slave_count]),
weights: Mutex::new(weights),
health_checker,
latency_stats: LatencyStats::new(),
rationing: Mutex::new(ReadRationing::default_slave_only()),
}
}
pub fn master(&self) -> &str {
&self.master
}
pub fn slaves(&self) -> &[String] {
&self.slaves
}
pub fn slave(&self) -> &str {
if self.slaves.is_empty() {
return &self.master;
}
if let Ok(rationing) = self.rationing.lock() {
if rationing.should_read_master() {
return &self.master;
}
}
if let Some(healthy) = self.select_healthy_slave() {
return healthy;
}
&self.master
}
fn select_healthy_slave(&self) -> Option<&str> {
let healthy_indices: Vec<usize> = self
.slaves
.iter()
.enumerate()
.filter(|(_, s)| {
self.health_checker
.health(s)
.unwrap_or(SlaveHealth::Healthy)
== SlaveHealth::Healthy
})
.map(|(i, _)| i)
.collect();
if healthy_indices.is_empty() {
return None;
}
let idx = match self.strategy {
LoadBalanceStrategy::RoundRobin => {
let counter = self.round_robin_counter.fetch_add(1, Ordering::SeqCst);
healthy_indices[counter % healthy_indices.len()]
}
LoadBalanceStrategy::Random => {
let idx = rand::thread_rng().gen_range(0..healthy_indices.len());
healthy_indices[idx]
}
LoadBalanceStrategy::LeastConnections => {
let counts = match self.connection_counts.lock() {
Ok(c) => c,
Err(_) => return Some(&self.slaves[healthy_indices[0]]),
};
let mut min_idx = healthy_indices[0];
let mut min_count = counts[min_idx];
for &i in healthy_indices.iter().skip(1) {
if counts[i] < min_count {
min_count = counts[i];
min_idx = i;
}
}
min_idx
}
LoadBalanceStrategy::WeightedRoundRobin => self.select_weighted_index(&healthy_indices),
};
Some(&self.slaves[idx])
}
fn select_weighted_index(&self, healthy_indices: &[usize]) -> usize {
let weights = match self.weights.lock() {
Ok(w) => w,
Err(_) => return healthy_indices[0],
};
let total: u64 = healthy_indices
.iter()
.map(|i| weights.get(&self.slaves[*i]).copied().unwrap_or(1) as u64)
.sum();
if total == 0 {
return healthy_indices[0];
}
let counter = self.round_robin_counter.fetch_add(1, Ordering::SeqCst);
let mut pick = (counter as u64) % total;
for &idx in healthy_indices.iter() {
let w = weights.get(&self.slaves[idx]).copied().unwrap_or(1) as u64;
if pick < w {
return idx;
}
pick -= w;
}
healthy_indices[healthy_indices.len() - 1]
}
pub fn set_strategy(&mut self, strategy: LoadBalanceStrategy) {
self.strategy = strategy;
}
pub fn strategy(&self) -> LoadBalanceStrategy {
self.strategy
}
pub fn set_weight(&self, slave: &str, weight: u32) -> Result<(), String> {
if !self.slaves.iter().any(|s| s == slave) {
return Err(format!("unknown slave: {}", slave));
}
if let Ok(mut weights) = self.weights.lock() {
weights.insert(slave.to_string(), weight.max(1));
}
Ok(())
}
pub fn weight(&self, slave: &str) -> Option<u32> {
self.weights.lock().ok().and_then(|w| w.get(slave).copied())
}
pub fn health_checker(&self) -> &HealthChecker {
&self.health_checker
}
pub fn latency_stats(&self) -> &LatencyStats {
&self.latency_stats
}
pub fn record_latency(&self, slave: &str, latency: Duration) {
self.latency_stats.record(slave, latency);
}
pub fn measure<F, T>(&self, slave: &str, f: F) -> T
where
F: FnOnce() -> T,
{
let start = Instant::now();
let result = f();
self.record_latency(slave, start.elapsed());
result
}
pub fn set_read_rationing(&self, percent: u8) {
if let Ok(mut r) = self.rationing.lock() {
r.set_percent(percent);
}
}
pub fn read_rationing_percent(&self) -> u8 {
self.rationing
.lock()
.map(|r| r.master_read_percent)
.unwrap_or(0)
}
pub fn acquire(&self, slave: &str) -> Result<(), String> {
if let Some(idx) = self.slaves.iter().position(|s| s == slave) {
let mut counts = self
.connection_counts
.lock()
.map_err(|e| format!("lock error: {}", e))?;
counts[idx] = counts[idx].saturating_add(1);
Ok(())
} else {
Err(format!("unknown slave: {}", slave))
}
}
pub fn release(&self, slave: &str) -> Result<(), String> {
if let Some(idx) = self.slaves.iter().position(|s| s == slave) {
let mut counts = self
.connection_counts
.lock()
.map_err(|e| format!("lock error: {}", e))?;
if counts[idx] > 0 {
counts[idx] -= 1;
}
Ok(())
} else {
Err(format!("unknown slave: {}", slave))
}
}
pub fn connection_count(&self, slave: &str) -> Option<usize> {
let idx = self.slaves.iter().position(|s| s == slave)?;
let counts = self.connection_counts.lock().ok()?;
Some(counts[idx])
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashSet;
#[test]
fn test_router_master() {
let router = ReadWriteRouter::new("master:3306", vec!["slave1:3306", "slave2:3306"]);
assert_eq!(router.master(), "master:3306");
}
#[test]
fn test_router_slaves_list() {
let router = ReadWriteRouter::new("m", vec!["s1", "s2", "s3"]);
assert_eq!(router.slaves().len(), 3);
assert_eq!(router.slaves()[0], "s1");
assert_eq!(router.slaves()[2], "s3");
}
#[test]
fn test_default_strategy_is_round_robin() {
let router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
assert_eq!(router.strategy(), LoadBalanceStrategy::RoundRobin);
}
#[test]
fn test_set_strategy() {
let mut router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
router.set_strategy(LoadBalanceStrategy::Random);
assert_eq!(router.strategy(), LoadBalanceStrategy::Random);
router.set_strategy(LoadBalanceStrategy::LeastConnections);
assert_eq!(router.strategy(), LoadBalanceStrategy::LeastConnections);
}
#[test]
fn test_round_robin_cycles_through_slaves() {
let mut router = ReadWriteRouter::new("m", vec!["s1", "s2", "s3"]);
router.set_strategy(LoadBalanceStrategy::RoundRobin);
let first = router.slave().to_string();
let second = router.slave().to_string();
let third = router.slave().to_string();
let fourth = router.slave().to_string();
assert_eq!(first, "s1");
assert_eq!(second, "s2");
assert_eq!(third, "s3");
assert_eq!(fourth, "s1", "RoundRobin 应在第 4 次回到 s1");
}
#[test]
fn test_round_robin_single_slave() {
let mut router = ReadWriteRouter::new("m", vec!["only_slave"]);
router.set_strategy(LoadBalanceStrategy::RoundRobin);
for _ in 0..5 {
assert_eq!(router.slave(), "only_slave");
}
}
#[test]
fn test_round_robin_visits_all_slaves() {
let mut router = ReadWriteRouter::new("m", vec!["s1", "s2", "s3", "s4"]);
router.set_strategy(LoadBalanceStrategy::RoundRobin);
let mut visited = HashSet::new();
for _ in 0..4 {
visited.insert(router.slave().to_string());
}
assert_eq!(visited.len(), 4, "一轮轮询应该访问所有 4 个 slave");
}
#[test]
fn test_random_returns_valid_slave() {
let mut router = ReadWriteRouter::new("m", vec!["s1", "s2", "s3"]);
router.set_strategy(LoadBalanceStrategy::Random);
let slaves: HashSet<&str> = ["s1", "s2", "s3"].iter().copied().collect();
for _ in 0..20 {
let picked = router.slave();
assert!(
slaves.contains(picked),
"随机策略返回了未知 slave: {}",
picked
);
}
}
#[test]
fn test_random_eventually_visits_multiple_slaves() {
let mut router = ReadWriteRouter::new("m", vec!["s1", "s2", "s3", "s4"]);
router.set_strategy(LoadBalanceStrategy::Random);
let mut visited = HashSet::new();
for _ in 0..200 {
visited.insert(router.slave().to_string());
}
assert!(
visited.len() >= 2,
"随机策略在 200 次调用后应至少访问 2 个 slave,实际: {}",
visited.len()
);
}
#[test]
fn test_random_single_slave() {
let mut router = ReadWriteRouter::new("m", vec!["only"]);
router.set_strategy(LoadBalanceStrategy::Random);
for _ in 0..10 {
assert_eq!(router.slave(), "only");
}
}
#[test]
fn test_least_connections_picks_zero_load_slave() {
let mut router = ReadWriteRouter::new("m", vec!["s1", "s2", "s3"]);
router.set_strategy(LoadBalanceStrategy::LeastConnections);
assert_eq!(router.slave(), "s1");
}
#[test]
fn test_least_connections_picks_least_loaded() {
let mut router = ReadWriteRouter::new("m", vec!["s1", "s2", "s3"]);
router.set_strategy(LoadBalanceStrategy::LeastConnections);
router.acquire("s1").unwrap();
router.acquire("s1").unwrap();
router.acquire("s2").unwrap();
assert_eq!(router.slave(), "s3");
router.acquire("s3").unwrap();
assert_eq!(router.slave(), "s2");
}
#[test]
fn test_least_connections_after_release() {
let mut router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
router.set_strategy(LoadBalanceStrategy::LeastConnections);
router.acquire("s1").unwrap();
router.acquire("s1").unwrap();
router.acquire("s2").unwrap();
assert_eq!(router.slave(), "s2");
router.release("s1").unwrap();
router.release("s1").unwrap();
assert_eq!(router.slave(), "s1");
}
#[test]
fn test_acquire_unknown_slave_returns_error() {
let router = ReadWriteRouter::new("m", vec!["s1"]);
assert!(router.acquire("nonexistent").is_err());
}
#[test]
fn test_release_unknown_slave_returns_error() {
let router = ReadWriteRouter::new("m", vec!["s1"]);
assert!(router.release("nonexistent").is_err());
}
#[test]
fn test_release_below_zero_clamped() {
let router = ReadWriteRouter::new("m", vec!["s1"]);
router.release("s1").unwrap();
assert_eq!(router.connection_count("s1"), Some(0));
}
#[test]
fn test_connection_count_query() {
let router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
assert_eq!(router.connection_count("s1"), Some(0));
router.acquire("s1").unwrap();
router.acquire("s1").unwrap();
assert_eq!(router.connection_count("s1"), Some(2));
assert_eq!(router.connection_count("s2"), Some(0));
assert_eq!(router.connection_count("unknown"), None);
}
#[test]
fn test_empty_slaves_falls_back_to_master() {
let router = ReadWriteRouter::new("only_master", vec![]);
assert_eq!(router.slave(), "only_master");
let mut router_rr = router;
router_rr.set_strategy(LoadBalanceStrategy::RoundRobin);
assert_eq!(router_rr.slave(), "only_master");
router_rr.set_strategy(LoadBalanceStrategy::Random);
assert_eq!(router_rr.slave(), "only_master");
router_rr.set_strategy(LoadBalanceStrategy::LeastConnections);
assert_eq!(router_rr.slave(), "only_master");
}
#[test]
fn test_round_robin_concurrent_safe() {
let router = std::sync::Arc::new(ReadWriteRouter::new("m", vec!["s1", "s2", "s3"]));
let mut handles = vec![];
for _ in 0..4 {
let r = std::sync::Arc::clone(&router);
handles.push(std::thread::spawn(move || {
for _ in 0..10 {
let _ = r.slave();
}
}));
}
for h in handles {
h.join().unwrap();
}
assert_eq!(router.round_robin_counter.load(Ordering::SeqCst), 40);
}
#[test]
fn test_health_checker_new_slave_is_healthy() {
let checker = HealthChecker::new(3);
checker.register("s1");
assert_eq!(checker.health("s1"), Some(SlaveHealth::Healthy));
}
#[test]
fn test_health_checker_unregistered_slave_returns_none() {
let checker = HealthChecker::new(3);
assert_eq!(checker.health("unknown"), None);
}
#[test]
fn test_health_checker_mark_unhealthy() {
let checker = HealthChecker::new(3);
checker.register("s1");
checker.set_health("s1", SlaveHealth::Unhealthy);
assert_eq!(checker.health("s1"), Some(SlaveHealth::Unhealthy));
}
#[test]
fn test_health_checker_mark_drained() {
let checker = HealthChecker::new(3);
checker.register("s1");
checker.set_health("s1", SlaveHealth::Drained);
assert_eq!(checker.health("s1"), Some(SlaveHealth::Drained));
}
#[test]
fn test_record_failure_below_threshold_keeps_healthy() {
let checker = HealthChecker::new(3);
checker.register("s1");
assert!(!checker.record_failure("s1"));
assert!(!checker.record_failure("s1"));
assert_eq!(checker.health("s1"), Some(SlaveHealth::Healthy));
assert_eq!(checker.failure_count("s1"), 2);
}
#[test]
fn test_record_failure_at_threshold_marks_unhealthy() {
let checker = HealthChecker::new(3);
checker.register("s1");
assert!(!checker.record_failure("s1"));
assert!(!checker.record_failure("s1"));
assert!(checker.record_failure("s1"));
assert_eq!(checker.health("s1"), Some(SlaveHealth::Unhealthy));
}
#[test]
fn test_record_success_resets_failure_count() {
let checker = HealthChecker::new(3);
checker.register("s1");
checker.record_failure("s1");
checker.record_failure("s1");
assert_eq!(checker.failure_count("s1"), 2);
checker.record_success("s1");
assert_eq!(checker.failure_count("s1"), 0);
}
#[test]
fn test_set_healthy_resets_failure_count() {
let checker = HealthChecker::new(2);
checker.register("s1");
checker.record_failure("s1");
checker.record_failure("s1");
assert_eq!(checker.health("s1"), Some(SlaveHealth::Unhealthy));
checker.set_health("s1", SlaveHealth::Healthy);
assert_eq!(checker.failure_count("s1"), 0);
}
#[test]
fn test_list_by_health() {
let checker = HealthChecker::new(3);
checker.register("s1");
checker.register("s2");
checker.register("s3");
checker.set_health("s2", SlaveHealth::Unhealthy);
let mut healthy = checker.healthy_slaves();
healthy.sort();
assert_eq!(healthy, vec!["s1".to_string(), "s3".to_string()]);
assert_eq!(checker.unhealthy_slaves(), vec!["s2".to_string()]);
}
#[test]
fn test_router_failover_to_master_when_all_unhealthy() {
let router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
router
.health_checker()
.set_health("s1", SlaveHealth::Unhealthy);
router
.health_checker()
.set_health("s2", SlaveHealth::Unhealthy);
assert_eq!(router.slave(), "m");
}
#[test]
fn test_router_skips_unhealthy_slave() {
let mut router = ReadWriteRouter::new("m", vec!["s1", "s2", "s3"]);
router.set_strategy(LoadBalanceStrategy::RoundRobin);
router
.health_checker()
.set_health("s2", SlaveHealth::Unhealthy);
for _ in 0..100 {
let picked = router.slave().to_string();
assert_ne!(picked, "s2", "不应选中不健康的 slave");
}
}
#[test]
fn test_router_failover_skips_drained_slave() {
let router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
router
.health_checker()
.set_health("s1", SlaveHealth::Drained);
for _ in 0..5 {
assert_eq!(router.slave(), "s2");
}
}
#[test]
fn test_router_default_health_checker_threshold_is_3() {
let router = ReadWriteRouter::new("m", vec!["s1"]);
assert_eq!(router.health_checker().failure_threshold, 3);
}
#[test]
fn test_set_weight_for_known_slave() {
let router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
router.set_weight("s1", 10).unwrap();
assert_eq!(router.weight("s1"), Some(10));
assert_eq!(router.weight("s2"), Some(1));
}
#[test]
fn test_set_weight_for_unknown_slave_errors() {
let router = ReadWriteRouter::new("m", vec!["s1"]);
assert!(router.set_weight("ghost", 10).is_err());
}
#[test]
fn test_set_weight_zero_clamped_to_one() {
let router = ReadWriteRouter::new("m", vec!["s1"]);
router.set_weight("s1", 0).unwrap();
assert_eq!(router.weight("s1"), Some(1));
}
#[test]
fn test_weighted_round_robin_respects_weights() {
let mut router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
router.set_strategy(LoadBalanceStrategy::WeightedRoundRobin);
router.set_weight("s1", 9).unwrap();
router.set_weight("s2", 1).unwrap();
let mut s1_count = 0usize;
let mut s2_count = 0usize;
for _ in 0..100 {
match router.slave() {
"s1" => s1_count += 1,
"s2" => s2_count += 1,
_ => {}
}
}
assert_eq!(s1_count + s2_count, 100);
assert!(
s1_count > s2_count * 3,
"权重 9:1 应使 s1 命中次数远多于 s2,实际 s1={}, s2={}",
s1_count,
s2_count
);
}
#[test]
fn test_weighted_round_robin_with_equal_weights_visits_all() {
let mut router = ReadWriteRouter::new("m", vec!["s1", "s2", "s3"]);
router.set_strategy(LoadBalanceStrategy::WeightedRoundRobin);
let mut visited = HashSet::new();
for _ in 0..30 {
visited.insert(router.slave().to_string());
}
assert_eq!(visited.len(), 3);
}
#[test]
fn test_weighted_round_robin_skips_unhealthy() {
let mut router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
router.set_strategy(LoadBalanceStrategy::WeightedRoundRobin);
router.set_weight("s1", 100).unwrap();
router.set_weight("s2", 1).unwrap();
router
.health_checker()
.set_health("s1", SlaveHealth::Unhealthy);
for _ in 0..10 {
assert_eq!(router.slave(), "s2");
}
}
#[test]
fn test_read_rationing_default_is_zero_percent_master() {
let router = ReadWriteRouter::new("m", vec!["s1"]);
assert_eq!(router.read_rationing_percent(), 0);
}
#[test]
fn test_read_rationing_all_master() {
let router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
router.set_read_rationing(100);
for _ in 0..10 {
assert_eq!(router.slave(), "m");
}
}
#[test]
fn test_read_rationing_all_slave() {
let router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
router.set_read_rationing(0);
for _ in 0..10 {
let picked = router.slave();
assert!(picked == "s1" || picked == "s2");
}
}
#[test]
fn test_read_rationing_partial_distribution() {
let router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
router.set_read_rationing(30);
let mut master_count = 0usize;
let mut slave_count = 0usize;
for _ in 0..100 {
let picked = router.slave();
if picked == "m" {
master_count += 1;
} else {
slave_count += 1;
}
}
assert!(
(25..=35).contains(&master_count),
"30% 比例下 master 命中数应在 25-35 之间,实际: {}",
master_count
);
assert_eq!(master_count + slave_count, 100);
}
#[test]
fn test_read_rationing_clamps_above_100() {
let r = ReadRationing::new(150);
assert_eq!(r.master_read_percent, 100);
}
#[test]
fn test_read_rationing_set_percent_resets_counter() {
let mut r = ReadRationing::new(50);
for _ in 0..5 {
let _ = r.should_read_master();
}
r.set_percent(80);
assert_eq!(r.counter.load(Ordering::Relaxed), 0);
}
#[test]
fn test_read_rationing_master_only_always_returns_true() {
let r = ReadRationing::default_master_only();
for _ in 0..10 {
assert!(r.should_read_master());
}
}
#[test]
fn test_read_rationing_slave_only_always_returns_false() {
let r = ReadRationing::default_slave_only();
for _ in 0..10 {
assert!(!r.should_read_master());
}
}
#[test]
fn test_latency_stats_record_and_snapshot() {
let stats = LatencyStats::new();
stats.record("s1", Duration::from_millis(10));
stats.record("s1", Duration::from_millis(20));
stats.record("s1", Duration::from_millis(30));
let snap = stats.snapshot("s1");
assert_eq!(snap.samples, 3);
assert!(snap.min_ns > 0);
assert!(snap.max_ns >= snap.min_ns);
let avg = snap.avg();
assert!(avg >= Duration::from_millis(9));
assert!(avg <= Duration::from_millis(31));
}
#[test]
fn test_latency_stats_unknown_slave_returns_default() {
let stats = LatencyStats::new();
let snap = stats.snapshot("ghost");
assert_eq!(snap.samples, 0);
assert_eq!(snap.avg_ns(), 0);
}
#[test]
fn test_latency_stats_reset() {
let stats = LatencyStats::new();
stats.record("s1", Duration::from_millis(10));
assert_eq!(stats.snapshot("s1").samples, 1);
stats.reset("s1");
assert_eq!(stats.snapshot("s1").samples, 0);
}
#[test]
fn test_latency_stats_all_returns_all_slaves() {
let stats = LatencyStats::new();
stats.record("s1", Duration::from_millis(10));
stats.record("s2", Duration::from_millis(20));
let all = stats.all();
assert_eq!(all.len(), 2);
}
#[test]
fn test_router_measure_records_latency() {
let router = ReadWriteRouter::new("m", vec!["s1"]);
let result = router.measure("s1", || 42);
assert_eq!(result, 42);
let snap = router.latency_stats().snapshot("s1");
assert_eq!(snap.samples, 1);
}
#[test]
fn test_router_record_latency_increases_sample_count() {
let router = ReadWriteRouter::new("m", vec!["s1", "s2"]);
router.record_latency("s1", Duration::from_micros(100));
router.record_latency("s1", Duration::from_micros(200));
router.record_latency("s2", Duration::from_micros(50));
assert_eq!(router.latency_stats().snapshot("s1").samples, 2);
assert_eq!(router.latency_stats().snapshot("s2").samples, 1);
}
#[test]
fn test_latency_snapshot_avg_with_zero_samples() {
let snap = LatencySnapshot::default();
assert_eq!(snap.avg_ns(), 0);
assert_eq!(snap.avg(), Duration::ZERO);
}
#[test]
fn test_latency_snapshot_min_updates() {
let mut snap = LatencySnapshot::default();
snap.record(Duration::from_millis(50));
assert_eq!(snap.min_ns, 50_000_000);
snap.record(Duration::from_millis(10));
assert_eq!(snap.min_ns, 10_000_000);
snap.record(Duration::from_millis(100));
assert_eq!(snap.min_ns, 10_000_000);
}
#[test]
fn test_latency_snapshot_max_updates() {
let mut snap = LatencySnapshot::default();
snap.record(Duration::from_millis(10));
assert_eq!(snap.max_ns, 10_000_000);
snap.record(Duration::from_millis(50));
assert_eq!(snap.max_ns, 50_000_000);
snap.record(Duration::from_millis(20));
assert_eq!(snap.max_ns, 50_000_000);
}
#[test]
fn test_weighted_slave_default_health_is_healthy() {
let ws = WeightedSlave::new("s1:3306", 5);
assert_eq!(ws.health, SlaveHealth::Healthy);
assert_eq!(ws.weight, 5);
}
#[test]
fn test_weighted_slave_zero_weight_clamped() {
let ws = WeightedSlave::new("s1:3306", 0);
assert_eq!(ws.weight, 1);
}
#[test]
fn test_weighted_slave_with_health_builder() {
let ws = WeightedSlave::new("s1:3306", 5).with_health(SlaveHealth::Unhealthy);
assert_eq!(ws.health, SlaveHealth::Unhealthy);
}
}