use std::sync::atomic::{AtomicUsize, Ordering};
#[derive(Debug)]
pub struct WeightedRoundRobin {
counter: AtomicUsize,
total_weight: usize,
}
impl WeightedRoundRobin {
#[must_use]
pub const fn new(total_weight: usize) -> Self {
Self {
counter: AtomicUsize::new(0),
total_weight,
}
}
pub const fn set_total_weight(&mut self, total_weight: usize) {
self.total_weight = total_weight;
}
#[must_use]
pub fn select_with_weight(&self, weight: usize) -> Option<usize> {
if weight == 0 {
return None;
}
let counter = self.counter.fetch_add(1, Ordering::Relaxed);
Some(counter % weight)
}
#[must_use]
pub const fn total_weight(&self) -> usize {
self.total_weight
}
}
#[derive(Debug)]
pub struct LeastLoaded {
tie_breaker: AtomicUsize,
}
impl LeastLoaded {
#[must_use]
pub const fn new() -> Self {
Self {
tie_breaker: AtomicUsize::new(0),
}
}
#[must_use]
pub fn should_replace_tie(&self, tie_count: usize) -> bool {
debug_assert!(tie_count > 1);
if tie_count <= 1 {
return false;
}
let counter = self
.tie_breaker
.fetch_add(0x9e37_79b9_7f4a_7c15usize, Ordering::Relaxed);
Self::mix(counter).is_multiple_of(tie_count)
}
#[inline]
fn mix(mut value: usize) -> usize {
value ^= value >> 16;
value = value.wrapping_mul(0x7feb_352d);
value ^= value >> 15;
value = value.wrapping_mul(0x846c_a68b);
value ^ (value >> 16)
}
}
impl Default for LeastLoaded {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
use std::thread;
#[test]
fn test_weighted_round_robin_basic() {
let strategy = WeightedRoundRobin::new(10);
for i in 0..20 {
assert_eq!(
strategy.select_with_weight(strategy.total_weight()),
Some(i % 10)
);
}
}
#[test]
fn test_weighted_round_robin_zero_weight() {
let strategy = WeightedRoundRobin::new(0);
assert_eq!(strategy.select_with_weight(strategy.total_weight()), None);
}
#[test]
fn test_weighted_round_robin_one_weight() {
let strategy = WeightedRoundRobin::new(1);
for _ in 0..100 {
assert_eq!(
strategy.select_with_weight(strategy.total_weight()),
Some(0)
);
}
}
#[test]
fn test_set_total_weight() {
let mut strategy = WeightedRoundRobin::new(10);
assert_eq!(strategy.total_weight(), 10);
strategy.set_total_weight(20);
assert_eq!(strategy.total_weight(), 20);
}
#[test]
fn test_set_total_weight_to_zero() {
let mut strategy = WeightedRoundRobin::new(10);
assert!(
strategy
.select_with_weight(strategy.total_weight())
.is_some()
);
strategy.set_total_weight(0);
assert_eq!(strategy.total_weight(), 0);
assert_eq!(strategy.select_with_weight(strategy.total_weight()), None);
}
#[test]
fn test_weighted_round_robin_large_weight() {
let strategy = WeightedRoundRobin::new(1000);
for i in 0..2000 {
assert_eq!(
strategy.select_with_weight(strategy.total_weight()),
Some(i % 1000)
);
}
}
#[test]
fn test_weighted_round_robin_odd_weight() {
let strategy = WeightedRoundRobin::new(7);
for i in 0..21 {
assert_eq!(
strategy.select_with_weight(strategy.total_weight()),
Some(i % 7)
);
}
}
#[test]
fn test_weighted_round_robin_prime_weight() {
let strategy = WeightedRoundRobin::new(13);
for i in 0..26 {
assert_eq!(
strategy.select_with_weight(strategy.total_weight()),
Some(i % 13)
);
}
}
#[test]
fn test_weighted_round_robin_counter_wraparound() {
let strategy = WeightedRoundRobin::new(10);
strategy.counter.store(usize::MAX - 5, Ordering::Relaxed);
for _ in 0..10 {
let result = strategy.select_with_weight(strategy.total_weight());
assert!(result.is_some());
assert!(result.unwrap() < 10);
}
}
#[test]
fn test_weighted_round_robin_distribution() {
let strategy = WeightedRoundRobin::new(10);
let mut counts = [0usize; 10];
for _ in 0..1000 {
let pos = strategy
.select_with_weight(strategy.total_weight())
.unwrap();
counts[pos] += 1;
}
for &count in &counts {
assert_eq!(count, 100);
}
}
#[test]
fn test_weighted_round_robin_concurrent() {
let strategy = Arc::new(WeightedRoundRobin::new(100));
let mut handles = vec![];
for _ in 0..10 {
let strategy_clone = Arc::clone(&strategy);
handles.push(thread::spawn(move || {
let mut results = vec![];
for _ in 0..100 {
results.push(strategy_clone.select_with_weight(100).unwrap());
}
results
}));
}
let mut all_results = vec![];
for handle in handles {
all_results.extend(handle.join().unwrap());
}
assert_eq!(all_results.len(), 1000);
for result in all_results {
assert!(result < 100);
}
}
#[test]
fn test_weighted_round_robin_concurrent_distribution() {
let strategy = Arc::new(WeightedRoundRobin::new(50));
let mut handles = vec![];
for _ in 0..5 {
let strategy_clone = Arc::clone(&strategy);
handles.push(thread::spawn(move || {
let mut counts = [0usize; 50];
for _ in 0..1000 {
let pos = strategy_clone.select_with_weight(50).unwrap();
counts[pos] += 1;
}
counts
}));
}
let mut total_counts = [0usize; 50];
for handle in handles {
let thread_counts = handle.join().unwrap();
for (i, &count) in thread_counts.iter().enumerate() {
total_counts[i] += count;
}
}
for &count in &total_counts {
assert_eq!(count, 100);
}
}
#[test]
fn test_weighted_round_robin_new_default_state() {
let strategy = WeightedRoundRobin::new(42);
assert_eq!(
strategy.select_with_weight(strategy.total_weight()),
Some(0)
);
assert_eq!(
strategy.select_with_weight(strategy.total_weight()),
Some(1)
);
}
#[test]
fn test_weighted_round_robin_debug_format() {
let strategy = WeightedRoundRobin::new(10);
let debug_str = format!("{strategy:?}");
assert!(debug_str.contains("WeightedRoundRobin"));
}
#[test]
fn test_set_total_weight_multiple_times() {
let mut strategy = WeightedRoundRobin::new(10);
assert_eq!(strategy.total_weight(), 10);
assert_eq!(
strategy.select_with_weight(strategy.total_weight()),
Some(0)
);
strategy.set_total_weight(5);
assert_eq!(strategy.total_weight(), 5);
assert_eq!(
strategy.select_with_weight(strategy.total_weight()),
Some(1)
);
strategy.set_total_weight(20);
assert_eq!(strategy.total_weight(), 20);
assert_eq!(
strategy.select_with_weight(strategy.total_weight()),
Some(2)
);
}
#[test]
fn test_weighted_round_robin_power_of_two() {
let strategy = WeightedRoundRobin::new(64);
for i in 0..128 {
assert_eq!(
strategy.select_with_weight(strategy.total_weight()),
Some(i % 64)
);
}
}
#[test]
fn test_weighted_round_robin_max_usize_weight() {
let strategy = WeightedRoundRobin::new(usize::MAX);
assert_eq!(
strategy.select_with_weight(strategy.total_weight()),
Some(0)
);
assert_eq!(
strategy.select_with_weight(strategy.total_weight()),
Some(1)
);
assert_eq!(
strategy.select_with_weight(strategy.total_weight()),
Some(2)
);
}
}