1use std::{
4 collections::HashMap,
5 num::NonZeroU64,
6 sync::Arc,
7 time::{Duration, Instant},
8};
9
10use rand::seq::SliceRandom;
11
12#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
14pub enum SelectionStrategy {
15 #[default]
17 RoundRobin,
18 Random,
20 ShuffledRoundRobin,
25}
26
27#[derive(Debug, Clone, PartialEq, Eq)]
38pub struct RotationPolicy {
39 pub requests_per_proxy: Option<NonZeroU64>,
41 pub max_age: Option<Duration>,
43 pub strategy: SelectionStrategy,
45}
46
47impl Default for RotationPolicy {
48 fn default() -> Self {
49 Self {
50 requests_per_proxy: NonZeroU64::new(1),
51 max_age: None,
52 strategy: SelectionStrategy::RoundRobin,
53 }
54 }
55}
56
57impl RotationPolicy {
58 pub fn every(requests: NonZeroU64) -> Self {
60 Self {
61 requests_per_proxy: Some(requests),
62 ..Self::default()
63 }
64 }
65
66 pub fn sticky() -> Self {
68 Self {
69 requests_per_proxy: None,
70 ..Self::default()
71 }
72 }
73
74 pub fn for_duration(duration: Duration) -> crate::Result<Self> {
77 let policy = Self {
78 max_age: Some(duration),
79 ..Self::sticky()
80 };
81 policy.validate()?;
82 Ok(policy)
83 }
84
85 pub fn validate(&self) -> crate::Result<()> {
87 if self.max_age.is_some_and(|age| age.is_zero()) {
88 return Err(crate::Error::Config(
89 "rotation max_age must be greater than zero",
90 ));
91 }
92 Ok(())
93 }
94}
95
96#[derive(Debug, Default)]
98pub(crate) struct RotationState {
99 current: Option<String>,
100 requests: u64,
101 selected_at: Option<Instant>,
102 rotate_next: bool,
103 candidates: Option<Arc<[String]>>,
106 positions: HashMap<String, usize>,
107 shuffled_remaining: Vec<usize>,
108 shuffled_started: bool,
109 #[cfg(test)]
110 snapshot_rebuilds: usize,
111}
112
113impl RotationState {
114 #[cfg(test)]
115 fn select(
116 &mut self,
117 candidates: &[String],
118 policy: &RotationPolicy,
119 now: Instant,
120 ) -> Option<String> {
121 let candidates = match &self.candidates {
124 Some(previous) if previous.as_ref() == candidates => Arc::clone(previous),
125 _ => Arc::from(candidates),
126 };
127 self.select_shared(&candidates, policy, now)
128 }
129
130 pub(crate) fn select_shared(
131 &mut self,
132 candidates: &Arc<[String]>,
133 policy: &RotationPolicy,
134 now: Instant,
135 ) -> Option<String> {
136 if candidates.is_empty() {
137 self.rotate_next = true;
139 return None;
140 }
141
142 let removed_successor = self.update_candidates(candidates);
143
144 let quota_reached = policy
145 .requests_per_proxy
146 .is_some_and(|limit| self.requests >= limit.get());
147 let age_reached = policy.max_age.is_some_and(|limit| {
148 self.selected_at
149 .is_some_and(|started| now.saturating_duration_since(started) >= limit)
150 });
151 let current_index = self
152 .current
153 .as_ref()
154 .and_then(|current| self.positions.get(current))
155 .copied();
156
157 if current_index.is_some() && !self.rotate_next && !quota_reached && !age_reached {
158 if policy.strategy == SelectionStrategy::ShuffledRoundRobin && !self.shuffled_started {
159 self.start_shuffled_round(candidates.len(), current_index);
162 }
163 self.requests = self.requests.saturating_add(1);
164 return self.current.clone();
165 }
166
167 let index = match policy.strategy {
168 SelectionStrategy::RoundRobin => current_index
169 .map(|index| (index + 1) % candidates.len())
170 .or(removed_successor)
171 .unwrap_or(0),
172 SelectionStrategy::Random => Self::next_random(candidates.len(), current_index),
173 SelectionStrategy::ShuffledRoundRobin => {
174 self.next_shuffled(candidates.len(), current_index)
175 }
176 };
177 let selected = candidates[index].clone();
178 self.current = Some(selected.clone());
179 self.requests = 1;
180 self.selected_at = Some(now);
181 self.rotate_next = false;
182 Some(selected)
183 }
184
185 pub(crate) fn has_current(&self) -> bool {
186 self.current.is_some()
187 }
188
189 pub(crate) fn seed(&mut self, id: &str, now: Instant) {
191 self.current = Some(id.to_owned());
192 self.requests = 0;
193 self.selected_at = Some(now);
194 self.rotate_next = false;
195 self.shuffled_remaining.clear();
196 self.shuffled_started = false;
197 }
198
199 pub(crate) fn invalidate(&mut self, id: &str) {
200 if self.current.as_deref() == Some(id) {
201 self.force_rotate();
202 }
203 }
204
205 pub(crate) fn force_rotate(&mut self) {
206 self.rotate_next = true;
207 }
208
209 fn update_candidates(&mut self, candidates: &Arc<[String]>) -> Option<usize> {
212 if self
213 .candidates
214 .as_ref()
215 .is_some_and(|previous| Arc::ptr_eq(previous, candidates))
216 {
217 return None;
218 }
219 let positions: HashMap<_, _> = candidates
220 .iter()
221 .enumerate()
222 .map(|(index, id)| (id.clone(), index))
223 .collect();
224 let successor = self.current.as_ref().and_then(|current| {
225 if positions.contains_key(current) {
226 return None;
227 }
228 let index = *self.positions.get(current)?;
229 let previous = self.candidates.as_ref()?;
230 previous[index + 1..]
231 .iter()
232 .chain(&previous[..index])
233 .find_map(|id| positions.get(id).copied())
234 });
235 if let Some(previous) = &self.candidates {
236 self.shuffled_remaining = self
237 .shuffled_remaining
238 .iter()
239 .filter_map(|index| positions.get(&previous[*index]).copied())
240 .collect();
241 }
242 self.positions = positions;
243 self.candidates = Some(Arc::clone(candidates));
244 #[cfg(test)]
245 {
246 self.snapshot_rebuilds += 1;
247 }
248 successor
249 }
250
251 fn next_random(len: usize, current: Option<usize>) -> usize {
254 if len == 1 {
255 return 0;
256 }
257 if let Some(current) = current {
258 let index = rand::random_range(0..len - 1);
259 index + usize::from(index >= current)
260 } else {
261 rand::random_range(0..len)
262 }
263 }
264
265 fn start_shuffled_round(&mut self, len: usize, already_visited: Option<usize>) {
266 self.shuffled_remaining = (0..len)
267 .filter(|index| Some(*index) != already_visited)
268 .collect();
269 self.shuffled_remaining.shuffle(&mut rand::rng());
270 self.shuffled_started = true;
271 }
272
273 fn next_shuffled(&mut self, len: usize, current: Option<usize>) -> usize {
274 if self.shuffled_remaining.is_empty() {
275 self.start_shuffled_round(len, None);
276 let last = self.shuffled_remaining.len() - 1;
277 if len > 1 && self.shuffled_remaining.last().copied() == current {
278 self.shuffled_remaining.swap(0, last);
279 }
280 }
281 self.shuffled_remaining
282 .pop()
283 .expect("a shuffled round is nonempty for a nonempty candidate list")
284 }
285}
286
287#[cfg(test)]
288mod tests {
289 use super::*;
290 use std::{
291 collections::HashMap,
292 sync::{Arc, Mutex},
293 thread,
294 };
295
296 fn nodes(ids: &[&str]) -> Vec<String> {
297 ids.iter().map(|id| (*id).to_owned()).collect()
298 }
299
300 fn every(requests: u64) -> RotationPolicy {
301 RotationPolicy {
302 requests_per_proxy: NonZeroU64::new(requests),
303 ..RotationPolicy::default()
304 }
305 }
306
307 fn select(
308 state: &mut RotationState,
309 nodes: &[String],
310 policy: &RotationPolicy,
311 now: Instant,
312 ) -> String {
313 state.select(nodes, policy, now).expect("eligible nodes")
314 }
315
316 #[test]
317 fn an_explicit_batch_rotates_on_twenty_first_acquisition() {
318 let mut state = RotationState::default();
319 let nodes = nodes(&["a", "b", "c"]);
320 let now = Instant::now();
321 for expected in ["a", "b", "c", "a"] {
322 for _ in 0..20 {
323 assert_eq!(select(&mut state, &nodes, &every(20), now), expected);
324 }
325 }
326 }
327
328 #[test]
329 fn default_visits_every_node_in_order_without_concentrating_a_batch() {
330 let mut state = RotationState::default();
331 let nodes = nodes(&["a", "b", "c"]);
332 let now = Instant::now();
333 let actual: Vec<_> = (0..7)
334 .map(|_| select(&mut state, &nodes, &RotationPolicy::default(), now))
335 .collect();
336 assert_eq!(actual, ["a", "b", "c", "a", "b", "c", "a"]);
337 }
338
339 #[test]
340 fn elapsed_time_rotates_at_exact_boundary_and_resets_the_clock() {
341 let mut state = RotationState::default();
342 let nodes = nodes(&["a", "b"]);
343 let start = Instant::now();
344 let policy = RotationPolicy {
345 requests_per_proxy: None,
346 max_age: Some(Duration::from_secs(10)),
347 ..RotationPolicy::default()
348 };
349 assert_eq!(select(&mut state, &nodes, &policy, start), "a");
350 assert_eq!(
351 select(
352 &mut state,
353 &nodes,
354 &policy,
355 start + Duration::from_millis(9_999)
356 ),
357 "a"
358 );
359 assert_eq!(
360 select(&mut state, &nodes, &policy, start + Duration::from_secs(10)),
361 "b"
362 );
363 assert_eq!(
364 select(&mut state, &nodes, &policy, start + Duration::from_secs(19)),
365 "b"
366 );
367 assert_eq!(
368 select(&mut state, &nodes, &policy, start + Duration::from_secs(20)),
369 "a"
370 );
371 }
372
373 #[test]
374 fn either_limit_triggers_rotation() {
375 let mut state = RotationState::default();
376 let nodes = nodes(&["a", "b", "c"]);
377 let start = Instant::now();
378 let policy = RotationPolicy {
379 max_age: Some(Duration::from_secs(10)),
380 ..every(2)
381 };
382 assert_eq!(select(&mut state, &nodes, &policy, start), "a");
383 assert_eq!(select(&mut state, &nodes, &policy, start), "a");
384 assert_eq!(select(&mut state, &nodes, &policy, start), "b");
385 assert_eq!(
386 select(&mut state, &nodes, &policy, start + Duration::from_secs(10)),
387 "c"
388 );
389 assert_eq!(
390 select(&mut state, &nodes, &policy, start + Duration::from_secs(10)),
391 "c"
392 );
393 assert_eq!(
394 select(&mut state, &nodes, &policy, start + Duration::from_secs(10)),
395 "a"
396 );
397 }
398
399 #[test]
400 fn random_rotation_never_repeats_when_an_alternative_exists() {
401 let mut state = RotationState::default();
402 let nodes = nodes(&["a", "b", "c", "d"]);
403 let now = Instant::now();
404 let policy = RotationPolicy {
405 strategy: SelectionStrategy::Random,
406 ..every(1)
407 };
408 let mut previous = None;
409 for _ in 0..1_000 {
410 let selected = select(&mut state, &nodes, &policy, now);
411 assert!(nodes.contains(&selected));
412 assert_ne!(previous.as_ref(), Some(&selected));
413 previous = Some(selected);
414 }
415 }
416
417 #[test]
418 fn shuffled_rounds_cover_every_node_without_repeating_at_the_boundary() {
419 let mut state = RotationState::default();
420 let candidates: Arc<[String]> = (0..16).map(|index| index.to_string()).collect();
421 let policy = RotationPolicy {
422 strategy: SelectionStrategy::ShuffledRoundRobin,
423 ..RotationPolicy::default()
424 };
425 let now = Instant::now();
426 let mut previous = None;
427 for _ in 0..20 {
428 let mut visited = std::collections::HashSet::new();
429 for _ in 0..candidates.len() {
430 let selected = state.select_shared(&candidates, &policy, now).unwrap();
431 assert_ne!(previous.as_ref(), Some(&selected));
432 assert!(visited.insert(selected.clone()));
433 previous = Some(selected);
434 }
435 assert_eq!(visited.len(), candidates.len());
436 }
437 }
438
439 #[test]
440 fn shuffled_rounds_apply_the_request_budget_to_each_visited_node() {
441 let mut state = RotationState::default();
442 let candidates: Arc<[String]> = nodes(&["a", "b", "c", "d"]).into();
443 let policy = RotationPolicy {
444 strategy: SelectionStrategy::ShuffledRoundRobin,
445 ..every(3)
446 };
447 let now = Instant::now();
448 for _ in 0..5 {
449 let mut visited = std::collections::HashSet::new();
450 for _ in 0..candidates.len() {
451 let selected = state.select_shared(&candidates, &policy, now).unwrap();
452 assert!(visited.insert(selected.clone()));
453 for _ in 0..2 {
454 assert_eq!(
455 state.select_shared(&candidates, &policy, now),
456 Some(selected.clone())
457 );
458 }
459 }
460 }
461 }
462
463 #[test]
464 fn shuffled_updates_finish_surviving_nodes_before_adding_new_nodes() {
465 let mut state = RotationState::default();
466 let candidates: Arc<[String]> = nodes(&["a", "b", "c", "d", "e"]).into();
467 let policy = RotationPolicy {
468 strategy: SelectionStrategy::ShuffledRoundRobin,
469 ..RotationPolicy::default()
470 };
471 let now = Instant::now();
472 let first = state.select_shared(&candidates, &policy, now).unwrap();
473 let removed = candidates[*state.shuffled_remaining.last().unwrap()].clone();
474 let expected: std::collections::HashSet<_> = candidates
475 .iter()
476 .filter(|id| **id != first && **id != removed)
477 .cloned()
478 .collect();
479 let updated: Arc<[String]> = candidates
482 .iter()
483 .rev()
484 .filter(|id| **id != removed)
485 .cloned()
486 .chain(["new".to_owned()])
487 .collect();
488 let actual: std::collections::HashSet<_> = (0..expected.len())
489 .map(|_| state.select_shared(&updated, &policy, now).unwrap())
490 .collect();
491 assert_eq!(actual, expected);
492
493 let next_round: std::collections::HashSet<_> = (0..updated.len())
494 .map(|_| state.select_shared(&updated, &policy, now).unwrap())
495 .collect();
496 assert_eq!(next_round, updated.iter().cloned().collect());
497 assert!(next_round.contains("new"));
498 assert!(!next_round.contains(&removed));
499 }
500
501 #[test]
502 fn seeded_sessions_start_at_the_assigned_node_without_spending_quota() {
503 let candidates: Arc<[String]> = nodes(&["a", "b", "c", "d"]).into();
504 let now = Instant::now();
505 let mut state = RotationState::default();
506 assert!(!state.has_current());
507 state.seed("c", now);
508 assert!(state.has_current());
509 for _ in 0..3 {
510 assert_eq!(
511 state.select_shared(&candidates, &every(3), now).as_deref(),
512 Some("c")
513 );
514 }
515 assert_eq!(
516 state.select_shared(&candidates, &every(3), now).as_deref(),
517 Some("d")
518 );
519
520 let policy = RotationPolicy {
521 strategy: SelectionStrategy::ShuffledRoundRobin,
522 ..RotationPolicy::default()
523 };
524 state.seed("b", now);
525 assert_eq!(
526 state.select_shared(&candidates, &policy, now).as_deref(),
527 Some("b")
528 );
529 let remaining: std::collections::HashSet<_> = (0..3)
530 .map(|_| state.select_shared(&candidates, &policy, now).unwrap())
531 .collect();
532 assert_eq!(remaining, nodes(&["a", "c", "d"]).into_iter().collect());
533 }
534
535 #[test]
536 fn large_shared_snapshots_are_indexed_only_when_the_allocation_changes() {
537 let candidates: Arc<[String]> = (0..4_096).map(|index| index.to_string()).collect();
538 let now = Instant::now();
539 for strategy in [
540 SelectionStrategy::RoundRobin,
541 SelectionStrategy::Random,
542 SelectionStrategy::ShuffledRoundRobin,
543 ] {
544 let mut state = RotationState::default();
545 let policy = RotationPolicy {
546 strategy,
547 ..every(7)
548 };
549 for _ in 0..20_000 {
550 state.select_shared(&candidates, &policy, now).unwrap();
551 }
552 assert_eq!(state.snapshot_rebuilds, 1);
553 assert!(Arc::ptr_eq(state.candidates.as_ref().unwrap(), &candidates));
554
555 let new_snapshot: Arc<[String]> = candidates.to_vec().into();
556 state.select_shared(&new_snapshot, &policy, now).unwrap();
557 assert_eq!(state.snapshot_rebuilds, 2);
558 assert!(Arc::ptr_eq(
559 state.candidates.as_ref().unwrap(),
560 &new_snapshot
561 ));
562 }
563 }
564
565 #[test]
566 fn manual_rotation_preserves_the_round_robin_position() {
567 let mut state = RotationState::default();
568 let nodes = nodes(&["a", "b", "c"]);
569 let now = Instant::now();
570 let policy = every(0); for expected in ["a", "b", "c", "a"] {
572 for _ in 0..30 {
573 assert_eq!(select(&mut state, &nodes, &policy, now), expected);
574 }
575 state.force_rotate();
576 }
577 }
578
579 #[test]
580 fn invalidation_only_forces_rotation_for_the_current_node() {
581 let mut state = RotationState::default();
582 let nodes = nodes(&["a", "b", "c"]);
583 let now = Instant::now();
584 let policy = every(20);
585 assert_eq!(select(&mut state, &nodes, &policy, now), "a");
586 state.invalidate("b");
587 assert_eq!(select(&mut state, &nodes, &policy, now), "a");
588 state.invalidate("a");
589 assert_eq!(select(&mut state, &nodes, &policy, now), "b");
590 }
591
592 #[test]
593 fn removal_continues_from_the_removed_nodes_successor_and_recovery_rejoins() {
594 let mut state = RotationState::default();
595 let all = nodes(&["a", "b", "c", "d"]);
596 let now = Instant::now();
597 let policy = every(1);
598 assert_eq!(select(&mut state, &all, &policy, now), "a");
599 assert_eq!(select(&mut state, &all, &policy, now), "b");
600 state.invalidate("b");
601 assert_eq!(
602 select(&mut state, &nodes(&["a", "c", "d"]), &policy, now),
603 "c"
604 );
605 assert_eq!(select(&mut state, &all, &policy, now), "d");
606 assert_eq!(select(&mut state, &all, &policy, now), "a");
607 assert_eq!(select(&mut state, &all, &policy, now), "b");
608 }
609
610 #[test]
611 fn removal_skips_missing_successors_and_wraps() {
612 let mut state = RotationState::default();
613 let all = nodes(&["a", "b", "c", "d"]);
614 let now = Instant::now();
615 let policy = every(1);
616 assert_eq!(select(&mut state, &all, &policy, now), "a");
617 assert_eq!(select(&mut state, &all, &policy, now), "b");
618 assert_eq!(select(&mut state, &nodes(&["a", "d"]), &policy, now), "d");
619 assert_eq!(select(&mut state, &nodes(&["a", "new"]), &policy, now), "a");
620 assert_eq!(select(&mut state, &nodes(&["new"]), &policy, now), "new");
621 }
622
623 #[test]
624 fn list_changes_do_not_reset_a_valid_nodes_request_allowance() {
625 let mut state = RotationState::default();
626 let now = Instant::now();
627 let policy = every(3);
628 assert_eq!(select(&mut state, &nodes(&["a", "b"]), &policy, now), "a");
629 assert_eq!(
630 select(&mut state, &nodes(&["b", "a", "c"]), &policy, now),
631 "a"
632 );
633 assert_eq!(select(&mut state, &nodes(&["a", "c"]), &policy, now), "a");
634 assert_eq!(select(&mut state, &nodes(&["a", "c"]), &policy, now), "c");
635 }
636
637 #[test]
638 fn empty_list_preserves_position_until_nodes_return() {
639 let mut state = RotationState::default();
640 let all = nodes(&["a", "b", "c"]);
641 let now = Instant::now();
642 let policy = every(20);
643 assert_eq!(state.select(&[], &policy, now), None);
644 assert_eq!(select(&mut state, &all, &policy, now), "a");
645 state.force_rotate();
646 assert_eq!(select(&mut state, &all, &policy, now), "b");
647 assert_eq!(state.select(&[], &policy, now), None);
648 assert_eq!(select(&mut state, &all, &policy, now), "c");
649 }
650
651 #[test]
652 fn a_single_node_remains_usable_with_every_strategy() {
653 for strategy in [
654 SelectionStrategy::RoundRobin,
655 SelectionStrategy::Random,
656 SelectionStrategy::ShuffledRoundRobin,
657 ] {
658 let mut state = RotationState::default();
659 let nodes = nodes(&["only"]);
660 let now = Instant::now();
661 let policy = RotationPolicy {
662 strategy,
663 ..every(1)
664 };
665 for _ in 0..50 {
666 assert_eq!(select(&mut state, &nodes, &policy, now), "only");
667 }
668 state.invalidate("only");
669 assert_eq!(select(&mut state, &nodes, &policy, now), "only");
670 state.force_rotate();
671 assert_eq!(select(&mut state, &nodes, &policy, now), "only");
672 }
673 }
674
675 #[test]
676 fn count_saturates_without_limits_and_rotates_at_the_maximum_limit() {
677 let now = Instant::now();
678 let nodes = nodes(&["a", "b"]);
679 let mut state = RotationState::default();
680 assert_eq!(select(&mut state, &nodes, &every(0), now), "a");
681 state.requests = u64::MAX - 1;
682 assert_eq!(select(&mut state, &nodes, &every(u64::MAX), now), "a");
683 assert_eq!(select(&mut state, &nodes, &every(u64::MAX), now), "b");
684 state.requests = u64::MAX;
685 assert_eq!(select(&mut state, &nodes, &every(0), now), "b");
686 assert_eq!(state.requests, u64::MAX);
687 }
688
689 #[test]
690 fn configuration_rejects_zero_time_but_allows_manual_only_rotation() {
691 assert!(
692 RotationPolicy {
693 max_age: Some(Duration::ZERO),
694 ..every(1)
695 }
696 .validate()
697 .is_err()
698 );
699 assert!(every(0).validate().is_ok());
700 assert!(RotationPolicy::default().validate().is_ok());
701 }
702
703 #[test]
704 fn constructors_select_independent_count_time_and_sticky_policies() {
705 let count = NonZeroU64::new(20).unwrap();
706 let counted = RotationPolicy::every(count);
707 assert_eq!(counted.requests_per_proxy, Some(count));
708 assert_eq!(counted.max_age, None);
709
710 let sticky = RotationPolicy::sticky();
711 assert_eq!(sticky.requests_per_proxy, None);
712 assert_eq!(sticky.max_age, None);
713
714 let duration = Duration::from_secs(30);
715 let timed = RotationPolicy::for_duration(duration).unwrap();
716 assert_eq!(timed.requests_per_proxy, None);
717 assert_eq!(timed.max_age, Some(duration));
718 assert!(RotationPolicy::for_duration(Duration::ZERO).is_err());
719 }
720
721 #[test]
722 fn concurrent_acquisitions_receive_exact_serialized_batches() {
723 let state = Arc::new(Mutex::new((RotationState::default(), Vec::new())));
724 let nodes = Arc::new(nodes(&["a", "b", "c", "d"]));
725 let mut threads = Vec::new();
726 for _ in 0..8 {
727 let state = Arc::clone(&state);
728 let nodes = Arc::clone(&nodes);
729 threads.push(thread::spawn(move || {
730 for _ in 0..100 {
731 let mut locked = state.lock().unwrap();
732 let selected = select(&mut locked.0, &nodes, &every(20), Instant::now());
733 locked.1.push(selected);
734 }
735 }));
736 }
737 for thread in threads {
738 thread.join().unwrap();
739 }
740 let locked = state.lock().unwrap();
741 let counts = locked.1.iter().fold(HashMap::new(), |mut counts, id| {
742 *counts.entry(id.as_str()).or_insert(0) += 1;
743 counts
744 });
745 for id in nodes.iter() {
746 assert_eq!(counts[id.as_str()], 200);
747 }
748 for (index, batch) in locked.1.chunks_exact(20).enumerate() {
749 assert!(batch.iter().all(|id| id == &nodes[index % nodes.len()]));
750 }
751 }
752}