1use std::{
2 collections::{BTreeSet, VecDeque},
3 sync::{
4 Arc, Mutex,
5 atomic::{AtomicU8, Ordering},
6 },
7 time::{Duration, Instant},
8};
9
10use tokio::sync::oneshot;
11
12use crate::BlobError;
13
14const PRIORITY_CYCLE: [usize; 7] = [0, 0, 0, 0, 1, 1, 2];
15const DEFAULT_GLOBAL_TRANSFERS: usize = 8;
16const DEFAULT_PROVIDER_TRANSFERS: u16 = 4;
17const INITIAL_BACKOFF: Duration = Duration::from_millis(50);
18const MAX_BACKOFF: Duration = Duration::from_secs(2);
19
20#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
22pub enum BlobPriority {
23 High,
24 #[default]
25 Normal,
26 Low,
27}
28
29impl BlobPriority {
30 pub(crate) const fn queue(self) -> usize {
31 match self {
32 Self::High => 0,
33 Self::Normal => 1,
34 Self::Low => 2,
35 }
36 }
37
38 pub(crate) const fn lower(self) -> Self {
39 match self {
40 Self::High => Self::Normal,
41 Self::Normal | Self::Low => Self::Low,
42 }
43 }
44
45 const fn rank(self) -> u8 {
46 match self {
47 Self::Low => 0,
48 Self::Normal => 1,
49 Self::High => 2,
50 }
51 }
52
53 const fn from_rank(rank: u8) -> Self {
54 match rank {
55 2 => Self::High,
56 1 => Self::Normal,
57 _ => Self::Low,
58 }
59 }
60}
61
62#[derive(Clone)]
64pub(crate) struct PrioritySignal {
65 rank: Arc<AtomicU8>,
66}
67
68impl PrioritySignal {
69 pub(crate) fn new(priority: BlobPriority) -> Self {
70 Self {
71 rank: Arc::new(AtomicU8::new(priority.rank())),
72 }
73 }
74
75 pub(crate) fn get(&self) -> BlobPriority {
76 BlobPriority::from_rank(self.rank.load(Ordering::Acquire))
77 }
78
79 pub(crate) fn promote(&self, priority: BlobPriority) {
80 self.rank.fetch_max(priority.rank(), Ordering::AcqRel);
81 }
82}
83
84#[derive(Clone)]
86pub struct FetchScheduler {
87 admission: Arc<Admission>,
88}
89
90impl Default for FetchScheduler {
91 fn default() -> Self {
92 Self::new(DEFAULT_GLOBAL_TRANSFERS).expect("default transfer limit is non-zero")
93 }
94}
95
96impl FetchScheduler {
97 pub fn new(max_concurrent: usize) -> Result<Self, BlobError> {
99 if max_concurrent == 0 {
100 return Err(BlobError::InvalidFetchConcurrency(max_concurrent));
101 }
102 Ok(Self {
103 admission: Arc::new(Admission {
104 max_concurrent,
105 state: Mutex::new(AdmissionState::default()),
106 }),
107 })
108 }
109
110 #[cfg(test)]
111 pub(crate) async fn acquire(
112 &self,
113 priority: BlobPriority,
114 ) -> Result<TransferPermit, BlobError> {
115 self.acquire_promotable(PrioritySignal::new(priority)).await
116 }
117
118 pub(crate) async fn acquire_promotable(
119 &self,
120 priority: PrioritySignal,
121 ) -> Result<TransferPermit, BlobError> {
122 let (sender, receiver) = oneshot::channel();
123 {
124 let mut state = self
125 .admission
126 .state
127 .lock()
128 .map_err(|_| BlobError::SchedulerUnavailable)?;
129 state.queues[priority.get().queue()].push_back(AdmissionWaiter { sender, priority });
130 self.admission.dispatch(&mut state);
131 }
132 receiver.await.map_err(|_| BlobError::SchedulerUnavailable)
133 }
134
135 #[cfg(test)]
136 fn queued(&self) -> [usize; 3] {
137 let state = self.admission.state.lock().unwrap();
138 std::array::from_fn(|index| state.queues[index].len())
139 }
140}
141
142struct Admission {
143 max_concurrent: usize,
144 state: Mutex<AdmissionState>,
145}
146
147#[derive(Default)]
148struct AdmissionState {
149 in_flight: usize,
150 cursor: usize,
151 queues: [VecDeque<AdmissionWaiter>; 3],
152}
153
154struct AdmissionWaiter {
155 sender: oneshot::Sender<TransferPermit>,
156 priority: PrioritySignal,
157}
158
159impl Admission {
160 fn dispatch(self: &Arc<Self>, state: &mut AdmissionState) {
161 while state.in_flight < self.max_concurrent {
162 reclassify_waiters(state);
163 let Some(queue) = next_queue(state) else {
164 break;
165 };
166 let Some(waiter) = state.queues[queue].pop_front() else {
167 continue;
168 };
169 let permit = TransferPermit {
170 admission: self.clone(),
171 active: true,
172 };
173 match waiter.sender.send(permit) {
174 Ok(()) => state.in_flight += 1,
175 Err(mut permit) => permit.active = false,
176 }
177 }
178 }
179
180 fn release(self: &Arc<Self>) {
181 let Ok(mut state) = self.state.lock() else {
182 return;
183 };
184 state.in_flight = state.in_flight.saturating_sub(1);
185 self.dispatch(&mut state);
186 }
187}
188
189fn reclassify_waiters(state: &mut AdmissionState) {
190 let mut waiters = Vec::new();
191 for queue in &mut state.queues {
192 waiters.extend(queue.drain(..));
193 }
194 for waiter in waiters {
195 state.queues[waiter.priority.get().queue()].push_back(waiter);
196 }
197}
198
199fn next_queue(state: &mut AdmissionState) -> Option<usize> {
200 for _ in 0..PRIORITY_CYCLE.len() {
201 let queue = PRIORITY_CYCLE[state.cursor];
202 state.cursor = (state.cursor + 1) % PRIORITY_CYCLE.len();
203 if !state.queues[queue].is_empty() {
204 return Some(queue);
205 }
206 }
207 None
208}
209
210pub(crate) struct TransferPermit {
211 admission: Arc<Admission>,
212 active: bool,
213}
214
215impl Drop for TransferPermit {
216 fn drop(&mut self) {
217 if self.active {
218 self.active = false;
219 self.admission.release();
220 }
221 }
222}
223
224#[derive(Debug, Clone, Copy, PartialEq, Eq)]
226pub struct ProviderMetrics {
227 pub(crate) slot: usize,
228 pub(crate) successes: u64,
229 pub(crate) failures: u64,
230 pub(crate) latency: Duration,
231 pub(crate) bytes_per_second: u64,
232 pub(crate) in_flight: u16,
233 pub(crate) cooling_down: bool,
234}
235
236impl ProviderMetrics {
237 pub const fn slot(&self) -> usize {
238 self.slot
239 }
240
241 pub const fn successes(&self) -> u64 {
242 self.successes
243 }
244
245 pub const fn failures(&self) -> u64 {
246 self.failures
247 }
248
249 pub const fn latency(&self) -> Duration {
250 self.latency
251 }
252
253 pub const fn bytes_per_second(&self) -> u64 {
254 self.bytes_per_second
255 }
256
257 pub const fn in_flight(&self) -> u16 {
258 self.in_flight
259 }
260
261 pub const fn cooling_down(&self) -> bool {
262 self.cooling_down
263 }
264}
265
266#[derive(Debug, Default)]
267struct ProviderState {
268 successes: u64,
269 failures: u64,
270 consecutive_failures: u32,
271 latency_micros: Option<u64>,
272 bytes_per_second: Option<u64>,
273 in_flight: u16,
274 retry_at: Option<Instant>,
275}
276
277#[derive(Debug)]
278pub(crate) struct ProviderBook {
279 states: Vec<ProviderState>,
280 exploration_cursor: usize,
281 per_provider_limit: u16,
282}
283
284impl ProviderBook {
285 pub(crate) fn new(providers: usize) -> Self {
286 Self {
287 states: (0..providers).map(|_| ProviderState::default()).collect(),
288 exploration_cursor: 0,
289 per_provider_limit: DEFAULT_PROVIDER_TRANSFERS,
290 }
291 }
292
293 pub(crate) fn select(
294 &mut self,
295 excluded: &BTreeSet<usize>,
296 expected_bytes: u64,
297 now: Instant,
298 ) -> Option<usize> {
299 let len = self.states.len();
300 for offset in 0..len {
301 let index = (self.exploration_cursor + offset) % len;
302 let state = &self.states[index];
303 if eligible(state, excluded, index, now, self.per_provider_limit)
304 && state.successes == 0
305 && state.failures == 0
306 {
307 self.exploration_cursor = (index + 1) % len;
308 self.states[index].in_flight += 1;
309 return Some(index);
310 }
311 }
312
313 let selected = self
314 .states
315 .iter()
316 .enumerate()
317 .filter(|(index, state)| {
318 eligible(state, excluded, *index, now, self.per_provider_limit)
319 })
320 .min_by_key(|(_, state)| provider_score(state, expected_bytes))
321 .map(|(index, _)| index)?;
322 self.states[selected].in_flight += 1;
323 Some(selected)
324 }
325
326 pub(crate) fn next_ready(&self, excluded: &BTreeSet<usize>, now: Instant) -> Option<Duration> {
327 self.states
328 .iter()
329 .enumerate()
330 .filter(|(index, state)| {
331 !excluded.contains(index) && state.in_flight < self.per_provider_limit
332 })
333 .filter_map(|(_, state)| state.retry_at)
334 .filter_map(|retry_at| retry_at.checked_duration_since(now))
335 .min()
336 }
337
338 pub(crate) fn finish_success(&mut self, provider: usize, bytes: u64, elapsed: Duration) {
339 let state = &mut self.states[provider];
340 state.in_flight = state.in_flight.saturating_sub(1);
341 state.successes += 1;
342 state.consecutive_failures = 0;
343 state.retry_at = None;
344 let elapsed_micros = u64::try_from(elapsed.as_micros())
345 .unwrap_or(u64::MAX)
346 .max(1);
347 update_ewma(&mut state.latency_micros, elapsed_micros);
348 update_ewma(
349 &mut state.bytes_per_second,
350 bytes.saturating_mul(1_000_000) / elapsed_micros,
351 );
352 }
353
354 pub(crate) fn finish_failure(&mut self, provider: usize, now: Instant) {
355 let state = &mut self.states[provider];
356 state.in_flight = state.in_flight.saturating_sub(1);
357 state.failures += 1;
358 state.consecutive_failures = state.consecutive_failures.saturating_add(1);
359 let shift = state.consecutive_failures.saturating_sub(1).min(6);
360 let multiplier = 1_u32 << shift;
361 let base = (INITIAL_BACKOFF * multiplier).min(MAX_BACKOFF);
362 let provider = u32::try_from(provider).unwrap_or(u32::MAX);
363 let jitter_percent = provider
364 .wrapping_mul(17)
365 .wrapping_add(state.consecutive_failures.wrapping_mul(13))
366 % 21;
367 let jitter = base * jitter_percent / 100;
368 state.retry_at = Some(now + (base + jitter).min(MAX_BACKOFF));
369 }
370
371 pub(crate) fn cancel(&mut self, provider: usize) {
372 self.states[provider].in_flight = self.states[provider].in_flight.saturating_sub(1);
373 }
374
375 pub(crate) fn hedge_delay(&self, provider: usize) -> Duration {
376 Duration::from_micros(
377 self.states[provider]
378 .latency_micros
379 .map_or(150_000, |latency| latency.saturating_mul(2))
380 .clamp(50_000, 500_000),
381 )
382 }
383
384 pub(crate) fn snapshots(&self, now: Instant) -> Vec<ProviderMetrics> {
385 self.states
386 .iter()
387 .enumerate()
388 .map(|(slot, state)| ProviderMetrics {
389 slot,
390 successes: state.successes,
391 failures: state.failures,
392 latency: Duration::from_micros(state.latency_micros.unwrap_or_default()),
393 bytes_per_second: state.bytes_per_second.unwrap_or_default(),
394 in_flight: state.in_flight,
395 cooling_down: state.retry_at.is_some_and(|retry_at| retry_at > now),
396 })
397 .collect()
398 }
399}
400
401fn eligible(
402 state: &ProviderState,
403 excluded: &BTreeSet<usize>,
404 index: usize,
405 now: Instant,
406 per_provider_limit: u16,
407) -> bool {
408 !excluded.contains(&index)
409 && state.in_flight < per_provider_limit
410 && state.retry_at.is_none_or(|retry_at| retry_at <= now)
411}
412
413fn provider_score(state: &ProviderState, expected_bytes: u64) -> u128 {
414 let latency = u128::from(state.latency_micros.unwrap_or(100_000));
415 let throughput = u128::from(state.bytes_per_second.unwrap_or(1024 * 1024).max(1));
416 latency
417 + u128::from(expected_bytes).saturating_mul(1_000_000) / throughput
418 + u128::from(state.consecutive_failures) * 250_000
419 + u128::from(state.in_flight) * 50_000
420}
421
422fn update_ewma(value: &mut Option<u64>, sample: u64) {
423 *value = Some(value.map_or(sample, |old| {
424 old.saturating_mul(3).saturating_add(sample) / 4
425 }));
426}
427
428#[cfg(test)]
429mod tests {
430 use tokio::sync::mpsc;
431
432 use super::*;
433
434 #[test]
435 fn unknown_providers_are_explored_before_reuse() {
436 let now = Instant::now();
437 let mut providers = ProviderBook::new(3);
438 let first = providers.select(&BTreeSet::new(), 1024, now).unwrap();
439 providers.cancel(first);
440 let second = providers.select(&BTreeSet::new(), 1024, now).unwrap();
441 providers.cancel(second);
442 let third = providers.select(&BTreeSet::new(), 1024, now).unwrap();
443 assert_eq!([first, second, third], [0, 1, 2]);
444 }
445
446 #[test]
447 fn measured_completion_time_beats_round_robin_order() {
448 let now = Instant::now();
449 let mut providers = ProviderBook::new(2);
450 providers.finish_success(0, 1024, Duration::from_millis(200));
451 providers.finish_success(1, 1024, Duration::from_millis(20));
452 let selected = providers.select(&BTreeSet::new(), 1024, now).unwrap();
453 assert_eq!(selected, 1);
454 }
455
456 #[test]
457 fn failed_provider_cools_down_while_an_alternative_is_ready() {
458 let now = Instant::now();
459 let mut providers = ProviderBook::new(2);
460 providers.finish_failure(0, now);
461 let selected = providers.select(&BTreeSet::new(), 1024, now).unwrap();
462 assert_eq!(selected, 1);
463 assert_eq!(providers.next_ready(&BTreeSet::from([0, 1]), now), None);
464 }
465
466 #[test]
467 fn provider_concurrency_and_backoff_are_bounded() {
468 let now = Instant::now();
469 let mut providers = ProviderBook::new(1);
470 let reservations = (0..DEFAULT_PROVIDER_TRANSFERS)
471 .map(|_| providers.select(&BTreeSet::new(), 1024, now).unwrap())
472 .collect::<Vec<_>>();
473 assert!(reservations.iter().all(|provider| *provider == 0));
474 assert_eq!(providers.select(&BTreeSet::new(), 1024, now), None);
475 for provider in reservations.iter().take(reservations.len() - 1) {
476 providers.cancel(*provider);
477 }
478 providers.finish_failure(*reservations.last().unwrap(), now);
479 let delay = providers.next_ready(&BTreeSet::new(), now).unwrap();
480 assert!(delay >= INITIAL_BACKOFF);
481 assert!(delay <= INITIAL_BACKOFF + INITIAL_BACKOFF / 5);
482 assert_eq!(providers.select(&BTreeSet::new(), 1024, now), None);
483 assert_eq!(
484 providers.select(&BTreeSet::new(), 1024, now + delay),
485 Some(0)
486 );
487 }
488
489 #[test]
490 fn hedge_delay_tracks_latency_with_strict_bounds() {
491 let mut providers = ProviderBook::new(1);
492 assert_eq!(providers.hedge_delay(0), Duration::from_millis(150));
493 providers.finish_success(0, 1024, Duration::from_millis(10));
494 assert_eq!(providers.hedge_delay(0), Duration::from_millis(50));
495 providers.finish_success(0, 1024, Duration::from_secs(5));
496 assert_eq!(providers.hedge_delay(0), Duration::from_millis(500));
497 }
498
499 #[tokio::test]
500 async fn admission_uses_four_two_one_weighting_without_starvation() {
501 let scheduler = FetchScheduler::new(1).unwrap();
502 let blocker = scheduler.acquire(BlobPriority::High).await.unwrap();
503 let (sender, mut receiver) = mpsc::unbounded_channel();
504 let mut tasks = Vec::new();
505 for (priority, count) in [
506 (BlobPriority::High, 4),
507 (BlobPriority::Normal, 2),
508 (BlobPriority::Low, 1),
509 ] {
510 for _ in 0..count {
511 let scheduler = scheduler.clone();
512 let sender = sender.clone();
513 tasks.push(tokio::spawn(async move {
514 let _permit = scheduler.acquire(priority).await.unwrap();
515 sender.send(priority).unwrap();
516 }));
517 }
518 }
519 while scheduler.queued() != [4, 2, 1] {
520 tokio::task::yield_now().await;
521 }
522 drop(blocker);
523 let mut order = Vec::new();
524 for _ in 0..7 {
525 order.push(receiver.recv().await.unwrap());
526 }
527 for task in tasks {
528 task.await.unwrap();
529 }
530 assert_eq!(
531 order,
532 [
533 BlobPriority::High,
534 BlobPriority::High,
535 BlobPriority::High,
536 BlobPriority::Normal,
537 BlobPriority::Normal,
538 BlobPriority::Low,
539 BlobPriority::High,
540 ]
541 );
542 }
543
544 #[tokio::test]
545 async fn cancelled_waiter_does_not_leak_capacity() {
546 let scheduler = FetchScheduler::new(1).unwrap();
547 let blocker = scheduler.acquire(BlobPriority::Normal).await.unwrap();
548 let waiting_scheduler = scheduler.clone();
549 let waiter =
550 tokio::spawn(
551 async move { waiting_scheduler.acquire(BlobPriority::Low).await.unwrap() },
552 );
553 while scheduler.queued() != [0, 0, 1] {
554 tokio::task::yield_now().await;
555 }
556 waiter.abort();
557 drop(blocker);
558 let _permit = tokio::time::timeout(
559 Duration::from_secs(1),
560 scheduler.acquire(BlobPriority::High),
561 )
562 .await
563 .unwrap()
564 .unwrap();
565 }
566
567 #[tokio::test]
568 async fn coalesced_waiter_can_be_promoted_out_of_the_low_queue() {
569 let scheduler = FetchScheduler::new(1).unwrap();
570 let blocker = scheduler.acquire(BlobPriority::Low).await.unwrap();
571 let promoted = PrioritySignal::new(BlobPriority::Low);
572 let promoted_scheduler = scheduler.clone();
573 let promoted_signal = promoted.clone();
574 let (sender, mut receiver) = mpsc::unbounded_channel();
575 let promoted_sender = sender.clone();
576 let promoted_task = tokio::spawn(async move {
577 let _permit = promoted_scheduler
578 .acquire_promotable(promoted_signal)
579 .await
580 .unwrap();
581 promoted_sender.send(BlobPriority::High).unwrap();
582 });
583 while scheduler.queued() != [0, 0, 1] {
584 tokio::task::yield_now().await;
585 }
586 let normal_scheduler = scheduler.clone();
587 let normal_task = tokio::spawn(async move {
588 let _permit = normal_scheduler
589 .acquire(BlobPriority::Normal)
590 .await
591 .unwrap();
592 sender.send(BlobPriority::Normal).unwrap();
593 });
594 while scheduler.queued() != [0, 1, 1] {
595 tokio::task::yield_now().await;
596 }
597 promoted.promote(BlobPriority::High);
598 drop(blocker);
599 assert_eq!(receiver.recv().await, Some(BlobPriority::High));
600 assert_eq!(receiver.recv().await, Some(BlobPriority::Normal));
601 promoted_task.await.unwrap();
602 normal_task.await.unwrap();
603 }
604}