Skip to main content

dynamo_kv_router/scheduling/
queue.rs

1// SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4use std::collections::HashMap;
5use std::marker::PhantomData;
6use std::sync::Arc;
7use std::sync::atomic::{AtomicUsize, Ordering as AtomicOrdering};
8use std::time::Duration;
9
10use tokio::sync::{mpsc, oneshot, watch};
11use tokio::time::Instant;
12
13use super::config::RouterQueuePolicy;
14use super::filter::RoutingEligibility;
15use super::overlap_refresh::{
16    NoopOverlapScoresRefresh, OverlapScoresRefresh, read_overlap_refresh_after, refresh_overlap,
17};
18use super::policy_config::{PolicyClassConfig, PolicyProfile};
19use super::policy_queue::{PolicyQueue, QueueSnapshot};
20use super::prefill_load::{PrefillLoadEstimator, effective_prefill_tokens};
21use super::selector::{DefaultWorkerSelector, WorkerSelector};
22use super::types::{
23    KvSchedulerError, OverloadedWorkerProvider, SchedulingContext, SchedulingRequest,
24    SchedulingResponse,
25};
26use crate::protocols::{LocalBlockHash, PrefillLoadHint, WorkerConfigLike, WorkerId};
27use crate::sequences::topology::WorkerDpRange;
28use crate::sequences::{ActiveSequencesMultiWorker, SequencePublisher, SequenceRequest};
29
30/// Large default for max_num_batched_tokens when not configured (effectively disables queueing for that worker)
31pub const DEFAULT_MAX_BATCHED_TOKENS: u64 = 10_000_000;
32
33const ADMISSION_CHANNEL_CAPACITY: usize = 65_536;
34
35struct ClassQueueCounters {
36    pending_count: AtomicUsize,
37    pending_isl_tokens: AtomicUsize,
38    pending_cached_tokens: AtomicUsize,
39}
40
41#[derive(Debug, Clone, Copy, PartialEq, Eq)]
42pub struct ClassQueueStats {
43    pub pending_count: usize,
44    pub pending_isl_tokens: usize,
45    pub pending_cached_tokens: usize,
46}
47
48struct QueuedRequest {
49    request: SchedulingRequest,
50    enqueue_at: Instant,
51    block_hashes: Option<Vec<LocalBlockHash>>,
52}
53
54#[allow(clippy::large_enum_variant)]
55enum AdmissionCommand {
56    Enqueue {
57        request: SchedulingRequest,
58        block_hashes: Option<Vec<LocalBlockHash>>,
59        ack_tx: oneshot::Sender<()>,
60    },
61    Update {
62        ack_tx: oneshot::Sender<()>,
63    },
64}
65
66struct SchedulerQueueActor<
67    P: SequencePublisher,
68    C: WorkerConfigLike,
69    Sel: WorkerSelector<C>,
70    RF: OverlapScoresRefresh,
71> {
72    pending: PolicyQueue<QueuedRequest>,
73    profile: PolicyProfile,
74    pending_count: Arc<AtomicUsize>,
75    pending_isl_tokens: Arc<AtomicUsize>,
76    class_counters: Arc<Vec<ClassQueueCounters>>,
77    slots: Arc<ActiveSequencesMultiWorker<P>>,
78    workers_with_configs: watch::Receiver<HashMap<WorkerId, C>>,
79    start_time: Instant,
80    block_size: u32,
81    selector: Sel,
82    prefill_load_estimator: Option<Arc<dyn PrefillLoadEstimator>>,
83    overlap_scores_refresh: Option<Arc<RF>>,
84    overlap_refresh_after: Option<Duration>,
85    overloaded_worker_provider: Option<OverloadedWorkerProvider>,
86}
87
88/// Queue that gates scheduling requests behind a capacity check.
89/// When all workers exceed `threshold_frac` utilisation the request is parked in `pending`.
90/// When capacity frees up (`update()`), pending requests are scheduled in priority order.
91/// If queueing is disabled (threshold_frac is None), requests are scheduled immediately.
92pub struct SchedulerQueue<
93    P: SequencePublisher,
94    C: WorkerConfigLike,
95    Sel: WorkerSelector<C> = DefaultWorkerSelector,
96    RF: OverlapScoresRefresh = NoopOverlapScoresRefresh,
97> {
98    admission_tx: mpsc::Sender<AdmissionCommand>,
99    /// Number of requests currently parked in the pending queue.
100    /// Incremented after push, decremented after pop. Lock-free reads via `Relaxed` load.
101    pending_count: Arc<AtomicUsize>,
102    /// Sum of `isl_tokens` for requests currently parked in the pending queue.
103    /// Incremented after push, decremented after pop. Lock-free reads via `Relaxed` load.
104    pending_isl_tokens: Arc<AtomicUsize>,
105    class_counters: Arc<Vec<ClassQueueCounters>>,
106    slots: Arc<ActiveSequencesMultiWorker<P>>,
107    workers_with_configs: watch::Receiver<HashMap<WorkerId, C>>,
108    queueing_enabled: bool,
109    supports_overlap_refresh: bool,
110    _marker: PhantomData<(Sel, RF)>,
111}
112
113impl<
114    P: SequencePublisher + 'static,
115    C: WorkerConfigLike + Send + Sync + 'static,
116    Sel: WorkerSelector<C> + Send + 'static,
117    RF: OverlapScoresRefresh + Send + Sync + 'static,
118> SchedulerQueue<P, C, Sel, RF>
119{
120    #[allow(clippy::too_many_arguments)]
121    pub fn new_with_overlap_refresh(
122        slots: Arc<ActiveSequencesMultiWorker<P>>,
123        workers_with_configs: watch::Receiver<HashMap<WorkerId, C>>,
124        threshold_frac: Option<f64>,
125        block_size: u32,
126        selector: Sel,
127        queue_policy: RouterQueuePolicy,
128        prefill_load_estimator: Option<Arc<dyn PrefillLoadEstimator>>,
129        overlap_scores_refresh: Option<Arc<RF>>,
130        overloaded_worker_provider: Option<OverloadedWorkerProvider>,
131    ) -> Self {
132        let profile = PolicyProfile::synthetic(threshold_frac, queue_policy);
133        Self::new_with_policy_profile(
134            slots,
135            workers_with_configs,
136            profile,
137            block_size,
138            selector,
139            prefill_load_estimator,
140            overlap_scores_refresh,
141            overloaded_worker_provider,
142        )
143    }
144
145    #[allow(clippy::too_many_arguments)]
146    pub fn new_with_policy_profile(
147        slots: Arc<ActiveSequencesMultiWorker<P>>,
148        workers_with_configs: watch::Receiver<HashMap<WorkerId, C>>,
149        profile: PolicyProfile,
150        block_size: u32,
151        selector: Sel,
152        prefill_load_estimator: Option<Arc<dyn PrefillLoadEstimator>>,
153        overlap_scores_refresh: Option<Arc<RF>>,
154        overloaded_worker_provider: Option<OverloadedWorkerProvider>,
155    ) -> Self {
156        Self::new_with_policy_profile_and_capacity(
157            slots,
158            workers_with_configs,
159            profile,
160            block_size,
161            selector,
162            prefill_load_estimator,
163            overlap_scores_refresh,
164            overloaded_worker_provider,
165            ADMISSION_CHANNEL_CAPACITY,
166        )
167    }
168
169    #[allow(clippy::too_many_arguments)]
170    fn new_with_policy_profile_and_capacity(
171        slots: Arc<ActiveSequencesMultiWorker<P>>,
172        workers_with_configs: watch::Receiver<HashMap<WorkerId, C>>,
173        profile: PolicyProfile,
174        block_size: u32,
175        selector: Sel,
176        prefill_load_estimator: Option<Arc<dyn PrefillLoadEstimator>>,
177        overlap_scores_refresh: Option<Arc<RF>>,
178        overloaded_worker_provider: Option<OverloadedWorkerProvider>,
179        admission_channel_capacity: usize,
180    ) -> Self {
181        let queueing_enabled = profile
182            .classes()
183            .iter()
184            .any(PolicyClassConfig::queueing_enabled);
185        for class in profile.classes() {
186            tracing::info!(
187                policy_class = class.name,
188                queue_policy = %class.queue_policy,
189                quantum = class.quantum,
190                prefill_busy_threshold = ?class.prefill_busy_threshold,
191                prefill_busy_threshold_frac = ?class.prefill_busy_threshold_frac,
192                "Router policy class configured"
193            );
194        }
195        let overlap_refresh_after = if overlap_scores_refresh.is_some() {
196            let configured = read_overlap_refresh_after();
197            match configured {
198                Some(d) => tracing::info!(
199                    "Router queue overlap-score refresh enabled after {:.1}s wait",
200                    d.as_secs_f64()
201                ),
202                None => tracing::info!(
203                    "Router queue overlap-score refresh disabled via DYN_ROUTER_OVERLAP_REFRESH_AFTER_SECS"
204                ),
205            }
206            configured
207        } else {
208            None
209        };
210        let pending_count = Arc::new(AtomicUsize::new(0));
211        let pending_isl_tokens = Arc::new(AtomicUsize::new(0));
212        let class_counters = Arc::new(
213            profile
214                .classes()
215                .iter()
216                .map(|_| ClassQueueCounters {
217                    pending_count: AtomicUsize::new(0),
218                    pending_isl_tokens: AtomicUsize::new(0),
219                    pending_cached_tokens: AtomicUsize::new(0),
220                })
221                .collect(),
222        );
223        let (admission_tx, admission_rx) = mpsc::channel(admission_channel_capacity);
224        let actor = SchedulerQueueActor {
225            pending: PolicyQueue::new(profile.clone()),
226            profile,
227            pending_count: Arc::clone(&pending_count),
228            pending_isl_tokens: Arc::clone(&pending_isl_tokens),
229            class_counters: Arc::clone(&class_counters),
230            slots: Arc::clone(&slots),
231            workers_with_configs: workers_with_configs.clone(),
232            start_time: Instant::now(),
233            block_size,
234            selector,
235            prefill_load_estimator,
236            overlap_scores_refresh,
237            overlap_refresh_after,
238            overloaded_worker_provider,
239        };
240        tokio::spawn(actor.run(admission_rx));
241        Self {
242            admission_tx,
243            pending_count,
244            pending_isl_tokens,
245            class_counters,
246            slots,
247            workers_with_configs,
248            queueing_enabled,
249            supports_overlap_refresh: overlap_refresh_after.is_some(),
250            _marker: PhantomData,
251        }
252    }
253}
254
255impl<
256    P: SequencePublisher + 'static,
257    C: WorkerConfigLike + Send + Sync + 'static,
258    Sel: WorkerSelector<C> + Send + 'static,
259> SchedulerQueue<P, C, Sel, NoopOverlapScoresRefresh>
260{
261    #[allow(clippy::too_many_arguments)]
262    pub fn new(
263        slots: Arc<ActiveSequencesMultiWorker<P>>,
264        workers_with_configs: watch::Receiver<HashMap<WorkerId, C>>,
265        threshold_frac: Option<f64>,
266        block_size: u32,
267        selector: Sel,
268        queue_policy: RouterQueuePolicy,
269        prefill_load_estimator: Option<Arc<dyn PrefillLoadEstimator>>,
270    ) -> Self {
271        Self::new_with_overlap_refresh(
272            slots,
273            workers_with_configs,
274            threshold_frac,
275            block_size,
276            selector,
277            queue_policy,
278            prefill_load_estimator,
279            None,
280            None,
281        )
282    }
283
284    #[allow(clippy::too_many_arguments)]
285    pub fn new_with_overload_provider(
286        slots: Arc<ActiveSequencesMultiWorker<P>>,
287        workers_with_configs: watch::Receiver<HashMap<WorkerId, C>>,
288        threshold_frac: Option<f64>,
289        block_size: u32,
290        selector: Sel,
291        queue_policy: RouterQueuePolicy,
292        prefill_load_estimator: Option<Arc<dyn PrefillLoadEstimator>>,
293        overloaded_worker_provider: Option<OverloadedWorkerProvider>,
294    ) -> Self {
295        Self::new_with_overlap_refresh(
296            slots,
297            workers_with_configs,
298            threshold_frac,
299            block_size,
300            selector,
301            queue_policy,
302            prefill_load_estimator,
303            None,
304            overloaded_worker_provider,
305        )
306    }
307}
308
309impl<
310    P: SequencePublisher + 'static,
311    C: WorkerConfigLike + Send + Sync + 'static,
312    Sel: WorkerSelector<C> + Send + 'static,
313    RF: OverlapScoresRefresh + Send + Sync + 'static,
314> SchedulerQueue<P, C, Sel, RF>
315{
316    /// Register externally-provided workers in the slot tracker.
317    ///
318    /// Looks up DP rank/size from the discovery watch channel; defaults to
319    /// `(0, 1)` for workers not yet known to discovery.
320    pub fn register_workers(&self, worker_ids: &std::collections::HashSet<u64>) {
321        let discovery_workers = self.workers_with_configs.borrow();
322        for &worker_id in worker_ids {
323            let (dp_start, dp_size) = discovery_workers
324                .get(&worker_id)
325                .map(|runtime_config| {
326                    (
327                        runtime_config.data_parallel_start_rank(),
328                        runtime_config.data_parallel_size(),
329                    )
330                })
331                .unwrap_or((0, 1));
332            let range = WorkerDpRange::new(worker_id, dp_start, dp_size);
333            if let Err(error) = self.slots.upsert_worker(range) {
334                tracing::warn!(worker_id, %error, "Invalid externally-provided worker topology");
335            }
336        }
337    }
338
339    /// Enqueue a new request.
340    /// If queueing is disabled or workers have capacity, schedule immediately.
341    /// Otherwise park in the pending heap.
342    pub async fn enqueue(&self, request: SchedulingRequest) {
343        self.enqueue_with_block_hashes(request, None).await;
344    }
345
346    pub async fn enqueue_with_block_hashes(
347        &self,
348        mut request: SchedulingRequest,
349        block_hashes: Option<Vec<LocalBlockHash>>,
350    ) {
351        let eligibility = request.eligibility();
352
353        if let Err(error) = eligibility.validate_pinned_worker_allowed() {
354            request.respond(Err(error));
355            return;
356        }
357
358        let (ack_tx, ack_rx) = oneshot::channel();
359        let command = AdmissionCommand::Enqueue {
360            request,
361            block_hashes: self.prepare_block_hashes_for_refresh(block_hashes),
362            ack_tx,
363        };
364
365        if let Err(error) = self.admission_tx.send(command).await {
366            let AdmissionCommand::Enqueue { mut request, .. } = error.0 else {
367                return;
368            };
369            request.respond(Err(KvSchedulerError::SubscriberShutdown));
370            return;
371        }
372
373        if ack_rx.await.is_err() {
374            tracing::warn!("scheduler queue actor dropped enqueue acknowledgement");
375        }
376    }
377
378    /// Called on prefill_complete/free. Drains pending requests while workers have capacity.
379    /// Each scheduled request updates active_tokens via add_request, so the prefill-busy check
380    /// sees fresh state on the next iteration.
381    pub async fn update(&self) {
382        if !self.queueing_enabled {
383            return;
384        }
385
386        let (ack_tx, ack_rx) = oneshot::channel();
387        if self
388            .admission_tx
389            .send(AdmissionCommand::Update { ack_tx })
390            .await
391            .is_ok()
392        {
393            let _ = ack_rx.await;
394        }
395    }
396
397    /// Number of requests currently parked in the pending queue (lock-free).
398    pub fn pending_count(&self) -> usize {
399        self.pending_count.load(AtomicOrdering::Relaxed)
400    }
401
402    /// Sum of `isl_tokens` for requests currently parked in the pending queue (lock-free).
403    pub fn pending_isl_tokens(&self) -> usize {
404        self.pending_isl_tokens.load(AtomicOrdering::Relaxed)
405    }
406
407    pub fn class_queue_stats(&self, class_index: usize) -> Option<ClassQueueStats> {
408        let counters = self.class_counters.get(class_index)?;
409        Some(ClassQueueStats {
410            pending_count: counters.pending_count.load(AtomicOrdering::Relaxed),
411            pending_isl_tokens: counters.pending_isl_tokens.load(AtomicOrdering::Relaxed),
412            pending_cached_tokens: counters.pending_cached_tokens.load(AtomicOrdering::Relaxed),
413        })
414    }
415
416    pub fn supports_overlap_refresh(&self) -> bool {
417        self.supports_overlap_refresh
418    }
419
420    fn prepare_block_hashes_for_refresh(
421        &self,
422        block_hashes: Option<Vec<LocalBlockHash>>,
423    ) -> Option<Vec<LocalBlockHash>> {
424        if !self.supports_overlap_refresh {
425            return None;
426        }
427        block_hashes.filter(|hashes| !hashes.is_empty())
428    }
429}
430
431impl<
432    P: SequencePublisher + 'static,
433    C: WorkerConfigLike + Send + Sync + 'static,
434    Sel: WorkerSelector<C> + Send + 'static,
435    RF: OverlapScoresRefresh + Send + Sync + 'static,
436> SchedulerQueueActor<P, C, Sel, RF>
437{
438    async fn run(mut self, mut rx: mpsc::Receiver<AdmissionCommand>) {
439        while let Some(command) = rx.recv().await {
440            match command {
441                AdmissionCommand::Enqueue {
442                    request,
443                    block_hashes,
444                    ack_tx,
445                } => {
446                    self.handle_enqueue(request, block_hashes);
447                    let _ = ack_tx.send(());
448                }
449                AdmissionCommand::Update { ack_tx } => {
450                    self.handle_update().await;
451                    let _ = ack_tx.send(());
452                }
453            }
454        }
455
456        let class_counters = Arc::clone(&self.class_counters);
457        for entry in self.pending.drain() {
458            let class_index = entry.class_index();
459            let snapshot = entry.snapshot();
460            self.pending_count.fetch_sub(1, AtomicOrdering::Relaxed);
461            self.pending_isl_tokens
462                .fetch_sub(snapshot.raw_isl_tokens, AtomicOrdering::Relaxed);
463            let counters = &class_counters[class_index];
464            counters.pending_count.fetch_sub(1, AtomicOrdering::Relaxed);
465            counters
466                .pending_isl_tokens
467                .fetch_sub(snapshot.raw_isl_tokens, AtomicOrdering::Relaxed);
468            counters
469                .pending_cached_tokens
470                .fetch_sub(snapshot.cached_tokens, AtomicOrdering::Relaxed);
471
472            let mut request = entry.into_payload().request;
473            request.respond(Err(KvSchedulerError::SubscriberShutdown));
474        }
475    }
476
477    fn handle_enqueue(
478        &mut self,
479        request: SchedulingRequest,
480        block_hashes: Option<Vec<LocalBlockHash>>,
481    ) {
482        let eligibility = request.eligibility();
483        let decay_now = Instant::now();
484        // Synthetic and explicit selections avoid cache work. Family
485        // classification reuses one worker generation for snapshot and busy checks.
486        let (class_index, snapshot, should_queue) = if let Some(class_index) = self
487            .profile
488            .direct_class_index(request.policy_class.as_deref())
489        {
490            let class = self.profile.class(class_index);
491            let should_queue = self.should_queue(class_index, class, || {
492                self.all_workers_prefill_busy(class, eligibility, decay_now)
493            });
494            (class_index, None, should_queue)
495        } else {
496            let active_tokens = self.slots.active_tokens(decay_now);
497            let workers = self.workers_with_configs.borrow();
498            let snapshot = Self::snapshot_for_with(&request, &workers);
499            let class_index = self
500                .profile
501                .resolve_class_index(request.policy_class.as_deref(), snapshot.uncached_tokens);
502            let class = self.profile.class(class_index);
503            let should_queue = self.should_queue(class_index, class, || {
504                Self::all_workers_prefill_busy_with(&active_tokens, &workers, class, eligibility)
505            });
506            (class_index, Some(snapshot), should_queue)
507        };
508        if !should_queue {
509            self.admit_one(request, decay_now);
510            return;
511        }
512
513        let snapshot = snapshot.unwrap_or_else(|| self.snapshot_for(&request));
514        let class = self.profile.class(class_index);
515        tracing::debug!(policy_class = class.name, "queueing request");
516        let arrival_offset = self.start_time.elapsed().as_secs_f64();
517        let priority_jump = request.priority_jump;
518        let strict_priority = request.strict_priority;
519        let queued = QueuedRequest {
520            request,
521            enqueue_at: decay_now,
522            block_hashes,
523        };
524        let worker_count = self.workers_with_configs.borrow().len();
525        if let Err((rejection, queued)) = self.pending.enqueue(
526            class_index,
527            worker_count,
528            snapshot,
529            arrival_offset,
530            priority_jump,
531            strict_priority,
532            queued,
533        ) {
534            let mut request = queued.request;
535            request.respond(Err(KvSchedulerError::QueueRejected(rejection)));
536            return;
537        }
538        self.pending_count.fetch_add(1, AtomicOrdering::Relaxed);
539        self.pending_isl_tokens
540            .fetch_add(snapshot.raw_isl_tokens, AtomicOrdering::Relaxed);
541        self.add_class_counters(class_index, snapshot);
542    }
543
544    fn should_queue(
545        &self,
546        class_index: usize,
547        class: &PolicyClassConfig,
548        all_workers_busy: impl FnOnce() -> bool,
549    ) -> bool {
550        // Preserve backlog anti-bypass and lazily avoid worker scans when an
551        // earlier condition already decides admission.
552        class.queueing_enabled() && (self.pending.has_backlog(class_index) || all_workers_busy())
553    }
554
555    fn snapshot_for(&self, request: &SchedulingRequest) -> QueueSnapshot {
556        let workers = self.workers_with_configs.borrow();
557        Self::snapshot_for_with(request, &workers)
558    }
559
560    fn snapshot_for_with(
561        request: &SchedulingRequest,
562        workers: &HashMap<WorkerId, C>,
563    ) -> QueueSnapshot {
564        // Cache overlap is sampled once and reused for classification, queue
565        // limits, ordering, DRR cost, and counters.
566        let context = SchedulingContext::new(request, workers);
567        QueueSnapshot::new(request.isl_tokens, context.best_cached_tokens())
568    }
569
570    async fn handle_update(&mut self) {
571        if self.pending.pending_count() == 0 {
572            return;
573        }
574
575        // Continuation draining stays actor-local; never self-send through the
576        // bounded command channel while processing an update.
577        loop {
578            let decay_now = Instant::now();
579            let active_tokens = self.slots.active_tokens(decay_now);
580            let popped = {
581                let configs = self.workers_with_configs.borrow();
582                self.pending.pop_next(|_, class, queued| {
583                    // TODO: This preserves head-of-line blocking within each policy
584                    // class. A blocked constrained head can stall later entries in
585                    // that class until a bounded non-HOL strategy is introduced.
586                    !Self::all_workers_prefill_busy_with(
587                        &active_tokens,
588                        &configs,
589                        class,
590                        queued.request.eligibility(),
591                    )
592                })
593            };
594            let Some(mut popped) = popped else {
595                break;
596            };
597            let snapshot = popped.snapshot();
598            let current_pending_count = self.pending_count.load(AtomicOrdering::Relaxed);
599            debug_assert!(
600                current_pending_count > 0,
601                "pending_count underflow on queue drain"
602            );
603            self.pending_count.fetch_sub(1, AtomicOrdering::Relaxed);
604            let current_pending_isl_tokens = self.pending_isl_tokens.load(AtomicOrdering::Relaxed);
605            debug_assert!(
606                current_pending_isl_tokens >= snapshot.raw_isl_tokens,
607                "pending_isl_tokens underflow: pending={} request_isl_tokens={}",
608                current_pending_isl_tokens,
609                snapshot.raw_isl_tokens
610            );
611            self.pending_isl_tokens
612                .fetch_sub(snapshot.raw_isl_tokens, AtomicOrdering::Relaxed);
613            self.subtract_class_counters(popped.class_index(), snapshot);
614            let queued = popped.payload_mut();
615            // NOTE: Overlap refresh is expected to be very short. We intentionally
616            // accept load crossing the class threshold during this await: busy
617            // thresholds guide admission, not reservation. This differs from main
618            // to avoid reversing counters, heap state, and charged DRR credit.
619            let refreshed = refresh_overlap(
620                self.overlap_scores_refresh.as_deref(),
621                self.overlap_refresh_after,
622                queued.block_hashes.as_deref(),
623                queued.enqueue_at,
624                decay_now,
625            )
626            .await;
627            let wait_ms = queued.enqueue_at.elapsed().as_millis() as u64;
628            if let Some(overlap) = refreshed {
629                tracing::info!(
630                    request_id = queued.request.mode.request_id().unwrap_or("unknown"),
631                    wait_ms,
632                    "refreshed overlap scores after long queue wait"
633                );
634                queued.request.overlap = overlap;
635            }
636            let admit_now = Instant::now();
637            let class_index = popped.class_index();
638            let class = self.profile.class(class_index);
639            let request = popped.into_payload().request;
640            tracing::debug!(
641                policy_class = class.name,
642                "scheduling request from pending queue"
643            );
644            self.admit_one(request, admit_now);
645        }
646    }
647
648    /// Run the full scheduling pipeline for a single request:
649    /// compute projected load -> select worker -> book tracked state -> respond.
650    fn admit_one(&self, mut request: SchedulingRequest, decay_now: Instant) {
651        request.worker_loads = self
652            .slots
653            .project_worker_loads(request.token_seq.as_deref(), decay_now);
654
655        let selection = {
656            let workers = self.workers_with_configs.borrow();
657            let overloaded_worker_ids = self
658                .overloaded_worker_provider
659                .as_ref()
660                .and_then(|provider| provider());
661            let eligibility = request.eligibility_with_overloaded(overloaded_worker_ids.as_ref());
662            self.selector
663                .select_worker(&workers, &request, eligibility, self.block_size)
664                .map(|selection| {
665                    let config = workers
666                        .get(&selection.worker.worker_id)
667                        .expect("selected worker config must exist");
668                    let selected_worker_tiers = request
669                        .overlap
670                        .selected_worker_tiers(selection.worker, config);
671                    (selection, selected_worker_tiers)
672                })
673        };
674
675        let (selection, selected_worker_tiers) = match selection {
676            Ok(s) => s,
677            Err(e) => {
678                tracing::warn!("scheduling failed: {e}");
679                request.respond(Err(e));
680                return;
681            }
682        };
683
684        let response = SchedulingResponse {
685            best_worker: selection.worker,
686            effective_overlap_blocks: selection.effective_overlap_blocks,
687            cached_tokens: selection.cached_tokens,
688            selected_worker_tiers,
689        };
690
691        if !request.mode.is_tracked() {
692            request.respond(Ok(response));
693            return;
694        }
695
696        let request_id = request
697            .mode
698            .tracked_request_id()
699            .expect("tracked mode always has a request ID")
700            .to_string();
701
702        let prefill_load_hint = self.prefill_load_hint_for(
703            request.isl_tokens,
704            selection.cached_tokens,
705            request.track_prefill_tokens,
706        );
707
708        let sequence_request = SequenceRequest {
709            request_id,
710            token_sequence: request.token_seq.take(),
711            track_prefill_tokens: request.track_prefill_tokens,
712            expected_output_tokens: request.expected_output_tokens,
713            prefill_load_hint,
714            worker: selection.worker,
715            lora_name: request.lora_name.take(),
716        };
717        self.book_and_respond(request, sequence_request, response);
718    }
719
720    /// Completes the tracked-admission ownership handoff.
721    ///
722    /// A closed receiver means the actor-owned request was abandoned before
723    /// booking, so there is nothing to install. Otherwise booking precedes the
724    /// response: once delivery succeeds, the response channel no longer tracks
725    /// request lifetime and the caller must install its RAII cleanup owner. If
726    /// delivery loses that race, roll back the booking here.
727    fn book_and_respond(
728        &self,
729        mut request: SchedulingRequest,
730        sequence_request: SequenceRequest,
731        response: SchedulingResponse,
732    ) {
733        if request.response_is_closed() {
734            tracing::debug!(
735                request_id = %sequence_request.request_id,
736                "Skipping scheduler booking for cancelled request"
737            );
738            return;
739        }
740
741        let request_id = sequence_request.request_id.clone();
742        if let Err(error) = self.slots.add_request(sequence_request, Instant::now()) {
743            tracing::warn!(%request_id, %error, "Failed to book scheduler state");
744            request.respond(Err(KvSchedulerError::BookingFailed(error.to_string())));
745            return;
746        }
747
748        if request.respond(Ok(response)) {
749            return;
750        }
751
752        tracing::debug!(%request_id, "Rolling back undelivered scheduler booking");
753        if let Err(error) = self.slots.free(&request_id, Instant::now()) {
754            tracing::error!(%request_id, %error, "Failed to roll back scheduler booking");
755        }
756    }
757
758    fn prefill_load_hint_for(
759        &self,
760        isl_tokens: usize,
761        cached_tokens: usize,
762        track_prefill_tokens: bool,
763    ) -> Option<PrefillLoadHint> {
764        if !track_prefill_tokens {
765            return None;
766        }
767
768        let effective_isl = effective_prefill_tokens(isl_tokens, cached_tokens);
769        if effective_isl == 0 {
770            return None;
771        }
772        let prefix = isl_tokens - effective_isl;
773
774        let expected_prefill_duration = match &self.prefill_load_estimator {
775            Some(estimator) => match estimator.predict_prefill_duration(1, effective_isl, prefix) {
776                Ok(expected_prefill_duration) => Some(expected_prefill_duration),
777                Err(error) => {
778                    tracing::warn!(
779                        effective_isl,
780                        prefix,
781                        "failed to predict prefill duration for active load tracking: {error}"
782                    );
783                    None
784                }
785            },
786            None => None,
787        };
788
789        Some(PrefillLoadHint {
790            initial_effective_prefill_tokens: effective_isl,
791            expected_prefill_duration,
792        })
793    }
794
795    /// Check if all eligible workers are prefill-busy based on threshold.
796    /// When `pinned_worker` is `Some`, only that exact worker/rank is considered.
797    /// Otherwise when `allowed` is `Some`, only those worker IDs are considered;
798    /// otherwise all registered workers are checked.
799    /// Returns false when no eligible workers exist so the request falls
800    /// through to `schedule`, which returns a proper `NoEndpoints` error.
801    fn all_workers_prefill_busy(
802        &self,
803        class: &PolicyClassConfig,
804        eligibility: RoutingEligibility<'_>,
805        decay_now: Instant,
806    ) -> bool {
807        let active_tokens = self.slots.active_tokens(decay_now);
808        let configs = self.workers_with_configs.borrow();
809        Self::all_workers_prefill_busy_with(&active_tokens, &configs, class, eligibility)
810    }
811
812    fn all_workers_prefill_busy_with(
813        active_tokens: &HashMap<crate::protocols::WorkerWithDpRank, usize>,
814        configs: &HashMap<WorkerId, C>,
815        class: &PolicyClassConfig,
816        eligibility: RoutingEligibility<'_>,
817    ) -> bool {
818        if let Some(worker) = eligibility.pinned_worker() {
819            let Ok(config) = eligibility.validate_worker_rank(configs, worker) else {
820                return false;
821            };
822
823            let max_batched = config
824                .max_num_batched_tokens()
825                .unwrap_or(DEFAULT_MAX_BATCHED_TOKENS);
826            let tokens = active_tokens.get(&worker).copied().unwrap_or(0);
827            return class.worker_is_busy(tokens, max_batched);
828        }
829
830        let mut checked_any = false;
831        let has_available = eligibility.any_eligible_worker_rank(configs, |worker, config| {
832            checked_any = true;
833            let max_batched = config
834                .max_num_batched_tokens()
835                .unwrap_or(DEFAULT_MAX_BATCHED_TOKENS);
836            let tokens = active_tokens.get(&worker).copied().unwrap_or(0);
837            !class.worker_is_busy(tokens, max_batched)
838        });
839
840        checked_any && !has_available
841    }
842
843    fn add_class_counters(&self, class_index: usize, snapshot: QueueSnapshot) {
844        let counters = &self.class_counters[class_index];
845        counters.pending_count.fetch_add(1, AtomicOrdering::Relaxed);
846        counters
847            .pending_isl_tokens
848            .fetch_add(snapshot.raw_isl_tokens, AtomicOrdering::Relaxed);
849        counters
850            .pending_cached_tokens
851            .fetch_add(snapshot.cached_tokens, AtomicOrdering::Relaxed);
852    }
853
854    fn subtract_class_counters(&self, class_index: usize, snapshot: QueueSnapshot) {
855        let counters = &self.class_counters[class_index];
856        counters.pending_count.fetch_sub(1, AtomicOrdering::Relaxed);
857        counters
858            .pending_isl_tokens
859            .fetch_sub(snapshot.raw_isl_tokens, AtomicOrdering::Relaxed);
860        counters
861            .pending_cached_tokens
862            .fetch_sub(snapshot.cached_tokens, AtomicOrdering::Relaxed);
863    }
864}
865
866#[cfg(test)]
867mod tests {
868    use std::collections::{HashMap, HashSet};
869    use std::sync::atomic::{AtomicUsize, Ordering};
870    use std::sync::{Arc, Condvar, Mutex as StdMutex};
871    use std::time::Duration;
872
873    use async_trait::async_trait;
874    use rustc_hash::FxHashMap;
875    use tokio::sync::{Barrier, watch};
876
877    use super::*;
878    use crate::protocols::{
879        ActiveLoad, ActiveSequenceEvent, WorkerSelectionResult, WorkerWithDpRank,
880    };
881    use crate::scheduling::OverlapSignals;
882    use crate::scheduling::types::{KvSchedulerError, ScheduleMode};
883    use crate::scheduling::{RefreshedOverlap, RouterPolicyConfig};
884    use crate::sequences::{ActiveSequencesMultiWorker, SequencePublisher};
885    use crate::test_utils::{NoopSequencePublisher, SimpleWorkerConfig};
886    use crate::{DefaultWorkerSelector, WorkerSelector};
887
888    fn decay_now() -> Instant {
889        Instant::now()
890    }
891
892    struct FixedPrefillLoadEstimator {
893        duration: Duration,
894    }
895
896    impl PrefillLoadEstimator for FixedPrefillLoadEstimator {
897        fn predict_prefill_duration(
898            &self,
899            _batch_size: usize,
900            _effective_isl: usize,
901            _prefix: usize,
902        ) -> anyhow::Result<Duration> {
903            Ok(self.duration)
904        }
905    }
906
907    type SchedulingResponseReceiver =
908        tokio::sync::oneshot::Receiver<Result<SchedulingResponse, KvSchedulerError>>;
909
910    struct DropResponseOnLoadPublisher {
911        response_rx: Arc<StdMutex<Option<SchedulingResponseReceiver>>>,
912    }
913
914    impl SequencePublisher for DropResponseOnLoadPublisher {
915        fn publish_event(
916            &self,
917            _event: &ActiveSequenceEvent,
918        ) -> impl std::future::Future<Output = anyhow::Result<()>> + Send {
919            std::future::ready(Ok(()))
920        }
921
922        fn publish_load(&self, _load: ActiveLoad) {
923            self.response_rx.lock().unwrap().take();
924        }
925
926        fn observe_load(&self, _: &WorkerWithDpRank, _: &str, _: usize, _: usize) {}
927    }
928
929    #[derive(Default)]
930    struct SelectorRendezvous {
931        arrivals: StdMutex<usize>,
932        cv: Condvar,
933    }
934
935    impl SelectorRendezvous {
936        fn wait_for_peer(&self) {
937            let mut arrivals = self.arrivals.lock().unwrap();
938            *arrivals += 1;
939
940            if *arrivals == 1 {
941                let _ = self
942                    .cv
943                    .wait_timeout(arrivals, Duration::from_millis(100))
944                    .unwrap();
945                return;
946            }
947
948            self.cv.notify_all();
949        }
950    }
951
952    #[derive(Clone)]
953    struct MinDecodeSelector {
954        rendezvous: Option<Arc<SelectorRendezvous>>,
955    }
956
957    impl WorkerSelector<SimpleWorkerConfig> for MinDecodeSelector {
958        fn select_worker(
959            &self,
960            workers: &HashMap<WorkerId, SimpleWorkerConfig>,
961            request: &SchedulingRequest,
962            eligibility: RoutingEligibility<'_>,
963            block_size: u32,
964        ) -> Result<WorkerSelectionResult, KvSchedulerError> {
965            if let Some(rendezvous) = &self.rendezvous {
966                rendezvous.wait_for_peer();
967            }
968
969            let mut best_worker = None;
970            eligibility.for_each_eligible_worker_rank(workers, |worker, _| {
971                let load = request.worker_load_for(worker);
972                let potential_prefill_tokens = if request.track_prefill_tokens {
973                    load.active_prefill_tokens
974                        .saturating_add(effective_prefill_tokens(
975                            request.isl_tokens,
976                            request.effective_cached_tokens_for(worker),
977                        ))
978                } else {
979                    0
980                };
981                let potential_decode_blocks = load.potential_decode_blocks();
982                let key = (
983                    potential_prefill_tokens,
984                    potential_decode_blocks,
985                    worker.worker_id,
986                    worker.dp_rank,
987                );
988                if best_worker.is_none_or(|(_, best_key)| key < best_key) {
989                    best_worker = Some((worker, key));
990                }
991            });
992
993            let Some((worker, _)) = best_worker else {
994                return Err(KvSchedulerError::NoEndpoints);
995            };
996
997            Ok(WorkerSelectionResult {
998                worker,
999                required_blocks: request.request_blocks(block_size),
1000                effective_overlap_blocks: request.effective_overlap_blocks_for(worker),
1001                cached_tokens: request.effective_cached_tokens_for(worker),
1002            })
1003        }
1004    }
1005
1006    fn make_queue(
1007        num_workers: usize,
1008        block_size: u32,
1009        isl: usize,
1010        threshold_frac: Option<f64>,
1011    ) -> (
1012        Arc<SchedulerQueue<NoopSequencePublisher, SimpleWorkerConfig>>,
1013        Arc<ActiveSequencesMultiWorker<NoopSequencePublisher>>,
1014    ) {
1015        let (queue, slots, _tx) =
1016            make_queue_with_sender(num_workers, block_size, isl, threshold_frac, None);
1017        (queue, slots)
1018    }
1019
1020    #[allow(clippy::type_complexity)]
1021    fn make_queue_with_custom_selector<Sel: WorkerSelector<SimpleWorkerConfig> + Send + 'static>(
1022        num_workers: usize,
1023        block_size: u32,
1024        isl: usize,
1025        threshold_frac: Option<f64>,
1026        selector: Sel,
1027    ) -> (
1028        Arc<SchedulerQueue<NoopSequencePublisher, SimpleWorkerConfig, Sel>>,
1029        Arc<ActiveSequencesMultiWorker<NoopSequencePublisher>>,
1030    ) {
1031        let dp_range: HashMap<u64, (u32, u32)> =
1032            (0..num_workers as u64).map(|id| (id, (0, 1))).collect();
1033        let slots = Arc::new(ActiveSequencesMultiWorker::new(
1034            NoopSequencePublisher,
1035            block_size as usize,
1036            dp_range,
1037            false,
1038            0,
1039            "test",
1040        ));
1041
1042        let mut configs: HashMap<u64, SimpleWorkerConfig> = HashMap::new();
1043        for id in 0..num_workers as u64 {
1044            configs.insert(
1045                id,
1046                SimpleWorkerConfig {
1047                    max_num_batched_tokens: Some(isl as u64),
1048                    ..Default::default()
1049                },
1050            );
1051        }
1052        let (_cfg_tx, cfg_rx) = watch::channel(configs);
1053
1054        let queue = Arc::new(SchedulerQueue::new(
1055            Arc::clone(&slots),
1056            cfg_rx,
1057            threshold_frac,
1058            block_size,
1059            selector,
1060            RouterQueuePolicy::Fcfs,
1061            None,
1062        ));
1063
1064        (queue, slots)
1065    }
1066
1067    #[allow(clippy::type_complexity)]
1068    fn make_queue_with_sender(
1069        num_workers: usize,
1070        block_size: u32,
1071        isl: usize,
1072        threshold_frac: Option<f64>,
1073        prefill_load_estimator: Option<Arc<dyn PrefillLoadEstimator>>,
1074    ) -> (
1075        Arc<SchedulerQueue<NoopSequencePublisher, SimpleWorkerConfig>>,
1076        Arc<ActiveSequencesMultiWorker<NoopSequencePublisher>>,
1077        watch::Sender<HashMap<u64, SimpleWorkerConfig>>,
1078    ) {
1079        let dp_range: HashMap<u64, (u32, u32)> =
1080            (0..num_workers as u64).map(|id| (id, (0, 1))).collect();
1081        let slots = Arc::new(ActiveSequencesMultiWorker::new(
1082            NoopSequencePublisher,
1083            block_size as usize,
1084            dp_range,
1085            false,
1086            0,
1087            "test",
1088        ));
1089
1090        let mut configs: HashMap<u64, SimpleWorkerConfig> = HashMap::new();
1091        for id in 0..num_workers as u64 {
1092            configs.insert(
1093                id,
1094                SimpleWorkerConfig {
1095                    max_num_batched_tokens: Some(isl as u64),
1096                    ..Default::default()
1097                },
1098            );
1099        }
1100        let (cfg_tx, cfg_rx) = watch::channel(configs);
1101
1102        let selector = DefaultWorkerSelector::new(None, "test");
1103        let queue = Arc::new(SchedulerQueue::new(
1104            Arc::clone(&slots),
1105            cfg_rx,
1106            threshold_frac,
1107            block_size,
1108            selector,
1109            RouterQueuePolicy::Fcfs,
1110            prefill_load_estimator,
1111        ));
1112
1113        (queue, slots, cfg_tx)
1114    }
1115
1116    fn policy_profile(yaml: &str) -> PolicyProfile {
1117        RouterPolicyConfig::from_yaml(yaml)
1118            .unwrap()
1119            .resolve_profile(None, None, crate::config::RouterQueuePolicy::Fcfs)
1120    }
1121
1122    #[allow(clippy::type_complexity)]
1123    fn make_queue_with_profile(
1124        num_workers: usize,
1125        block_size: u32,
1126        max_num_batched_tokens: usize,
1127        profile: PolicyProfile,
1128    ) -> (
1129        Arc<SchedulerQueue<NoopSequencePublisher, SimpleWorkerConfig>>,
1130        Arc<ActiveSequencesMultiWorker<NoopSequencePublisher>>,
1131    ) {
1132        let (queue, slots, _cfg_tx) = make_queue_with_profile_and_sender(
1133            num_workers,
1134            block_size,
1135            max_num_batched_tokens,
1136            profile,
1137        );
1138        (queue, slots)
1139    }
1140
1141    #[allow(clippy::type_complexity)]
1142    fn make_queue_with_profile_and_sender(
1143        num_workers: usize,
1144        block_size: u32,
1145        max_num_batched_tokens: usize,
1146        profile: PolicyProfile,
1147    ) -> (
1148        Arc<SchedulerQueue<NoopSequencePublisher, SimpleWorkerConfig>>,
1149        Arc<ActiveSequencesMultiWorker<NoopSequencePublisher>>,
1150        watch::Sender<HashMap<u64, SimpleWorkerConfig>>,
1151    ) {
1152        let dp_range: HashMap<u64, (u32, u32)> =
1153            (0..num_workers as u64).map(|id| (id, (0, 1))).collect();
1154        let slots = Arc::new(ActiveSequencesMultiWorker::new(
1155            NoopSequencePublisher,
1156            block_size as usize,
1157            dp_range,
1158            false,
1159            0,
1160            "test",
1161        ));
1162        let configs = (0..num_workers as u64)
1163            .map(|id| {
1164                (
1165                    id,
1166                    SimpleWorkerConfig {
1167                        max_num_batched_tokens: Some(max_num_batched_tokens as u64),
1168                        ..Default::default()
1169                    },
1170                )
1171            })
1172            .collect();
1173        let (cfg_tx, cfg_rx) = watch::channel(configs);
1174        let queue = Arc::new(SchedulerQueue::new_with_policy_profile(
1175            Arc::clone(&slots),
1176            cfg_rx,
1177            profile,
1178            block_size,
1179            DefaultWorkerSelector::new(None, "test"),
1180            None,
1181            None,
1182            None,
1183        ));
1184        (queue, slots, cfg_tx)
1185    }
1186
1187    fn make_queue_with_overload_provider(
1188        num_workers: usize,
1189        block_size: u32,
1190        isl: usize,
1191        overloaded_worker_provider: OverloadedWorkerProvider,
1192    ) -> (
1193        Arc<SchedulerQueue<NoopSequencePublisher, SimpleWorkerConfig>>,
1194        Arc<ActiveSequencesMultiWorker<NoopSequencePublisher>>,
1195    ) {
1196        let dp_range: HashMap<u64, (u32, u32)> =
1197            (0..num_workers as u64).map(|id| (id, (0, 1))).collect();
1198        let slots = Arc::new(ActiveSequencesMultiWorker::new(
1199            NoopSequencePublisher,
1200            block_size as usize,
1201            dp_range,
1202            false,
1203            0,
1204            "test",
1205        ));
1206
1207        let mut configs: HashMap<u64, SimpleWorkerConfig> = HashMap::new();
1208        for id in 0..num_workers as u64 {
1209            configs.insert(
1210                id,
1211                SimpleWorkerConfig {
1212                    max_num_batched_tokens: Some(isl as u64),
1213                    ..Default::default()
1214                },
1215            );
1216        }
1217        let (_cfg_tx, cfg_rx) = watch::channel(configs);
1218
1219        let selector = DefaultWorkerSelector::new(None, "test");
1220        let queue = Arc::new(SchedulerQueue::new_with_overload_provider(
1221            Arc::clone(&slots),
1222            cfg_rx,
1223            None,
1224            block_size,
1225            selector,
1226            RouterQueuePolicy::Fcfs,
1227            None,
1228            Some(overloaded_worker_provider),
1229        ));
1230
1231        (queue, slots)
1232    }
1233
1234    struct CountingRefresher {
1235        calls: AtomicUsize,
1236        response: RefreshedOverlap,
1237    }
1238
1239    #[async_trait]
1240    impl OverlapScoresRefresh for CountingRefresher {
1241        async fn refresh(&self, _block_hashes: &[LocalBlockHash]) -> Option<RefreshedOverlap> {
1242            self.calls.fetch_add(1, Ordering::Relaxed);
1243            Some(self.response.clone())
1244        }
1245    }
1246
1247    struct BlockingRefresher {
1248        calls: AtomicUsize,
1249        started: tokio::sync::Notify,
1250        release: tokio::sync::Notify,
1251        response: RefreshedOverlap,
1252    }
1253
1254    impl BlockingRefresher {
1255        fn new(response: RefreshedOverlap) -> Self {
1256            Self {
1257                calls: AtomicUsize::new(0),
1258                started: tokio::sync::Notify::new(),
1259                release: tokio::sync::Notify::new(),
1260                response,
1261            }
1262        }
1263
1264        async fn wait_for_calls(&self, target: usize) {
1265            while self.calls.load(Ordering::Relaxed) < target {
1266                self.started.notified().await;
1267            }
1268        }
1269
1270        fn release_one(&self) {
1271            self.release.notify_one();
1272        }
1273    }
1274
1275    #[async_trait]
1276    impl OverlapScoresRefresh for BlockingRefresher {
1277        async fn refresh(&self, _block_hashes: &[LocalBlockHash]) -> Option<RefreshedOverlap> {
1278            self.calls.fetch_add(1, Ordering::Relaxed);
1279            self.started.notify_one();
1280            self.release.notified().await;
1281            Some(self.response.clone())
1282        }
1283    }
1284
1285    #[allow(clippy::type_complexity)]
1286    fn make_queue_with_refresher(
1287        num_workers: usize,
1288        block_size: u32,
1289        isl: usize,
1290        threshold_frac: Option<f64>,
1291        refresher: Arc<CountingRefresher>,
1292    ) -> (
1293        Arc<
1294            SchedulerQueue<
1295                NoopSequencePublisher,
1296                SimpleWorkerConfig,
1297                DefaultWorkerSelector,
1298                CountingRefresher,
1299            >,
1300        >,
1301        Arc<ActiveSequencesMultiWorker<NoopSequencePublisher>>,
1302    ) {
1303        let dp_range: HashMap<u64, (u32, u32)> =
1304            (0..num_workers as u64).map(|id| (id, (0, 1))).collect();
1305        let slots = Arc::new(ActiveSequencesMultiWorker::new(
1306            NoopSequencePublisher,
1307            block_size as usize,
1308            dp_range,
1309            false,
1310            0,
1311            "test",
1312        ));
1313
1314        let mut configs: HashMap<u64, SimpleWorkerConfig> = HashMap::new();
1315        for id in 0..num_workers as u64 {
1316            configs.insert(
1317                id,
1318                SimpleWorkerConfig {
1319                    max_num_batched_tokens: Some(isl as u64),
1320                    ..Default::default()
1321                },
1322            );
1323        }
1324        let (_cfg_tx, cfg_rx) = watch::channel(configs);
1325
1326        let queue = Arc::new(SchedulerQueue::new_with_overlap_refresh(
1327            Arc::clone(&slots),
1328            cfg_rx,
1329            threshold_frac,
1330            block_size,
1331            DefaultWorkerSelector::new(None, "test"),
1332            RouterQueuePolicy::Fcfs,
1333            None,
1334            Some(refresher),
1335            None,
1336        ));
1337
1338        (queue, slots)
1339    }
1340
1341    #[allow(clippy::type_complexity)]
1342    fn make_queue_with_blocking_refresher(
1343        num_workers: usize,
1344        block_size: u32,
1345        isl: usize,
1346        threshold_frac: Option<f64>,
1347        refresher: Arc<BlockingRefresher>,
1348        admission_channel_capacity: usize,
1349    ) -> (
1350        Arc<
1351            SchedulerQueue<
1352                NoopSequencePublisher,
1353                SimpleWorkerConfig,
1354                DefaultWorkerSelector,
1355                BlockingRefresher,
1356            >,
1357        >,
1358        Arc<ActiveSequencesMultiWorker<NoopSequencePublisher>>,
1359    ) {
1360        let dp_range: HashMap<u64, (u32, u32)> =
1361            (0..num_workers as u64).map(|id| (id, (0, 1))).collect();
1362        let slots = Arc::new(ActiveSequencesMultiWorker::new(
1363            NoopSequencePublisher,
1364            block_size as usize,
1365            dp_range,
1366            false,
1367            0,
1368            "test",
1369        ));
1370
1371        let mut configs: HashMap<u64, SimpleWorkerConfig> = HashMap::new();
1372        for id in 0..num_workers as u64 {
1373            configs.insert(
1374                id,
1375                SimpleWorkerConfig {
1376                    max_num_batched_tokens: Some(isl as u64),
1377                    ..Default::default()
1378                },
1379            );
1380        }
1381        let (_cfg_tx, cfg_rx) = watch::channel(configs);
1382
1383        let queue = Arc::new(SchedulerQueue::new_with_policy_profile_and_capacity(
1384            Arc::clone(&slots),
1385            cfg_rx,
1386            PolicyProfile::synthetic(threshold_frac, crate::config::RouterQueuePolicy::Fcfs),
1387            block_size,
1388            DefaultWorkerSelector::new(None, "test"),
1389            None,
1390            Some(refresher),
1391            None,
1392            admission_channel_capacity,
1393        ));
1394
1395        (queue, slots)
1396    }
1397
1398    fn make_request(
1399        request_id: &str,
1400        isl_tokens: usize,
1401    ) -> (
1402        SchedulingRequest,
1403        tokio::sync::oneshot::Receiver<
1404            Result<SchedulingResponse, crate::scheduling::types::KvSchedulerError>,
1405        >,
1406    ) {
1407        let (tx, rx) = tokio::sync::oneshot::channel();
1408        let req = SchedulingRequest {
1409            mode: ScheduleMode::Tracked {
1410                request_id: request_id.to_string(),
1411            },
1412            token_seq: None,
1413            isl_tokens,
1414            overlap: OverlapSignals::default(),
1415            worker_loads: FxHashMap::default(),
1416            track_prefill_tokens: true,
1417            router_config_override: None,
1418            lora_name: None,
1419            priority_jump: 0.0,
1420            strict_priority: 0,
1421            policy_class: None,
1422            expected_output_tokens: None,
1423            pinned_worker: None,
1424            allowed_worker_ids: None,
1425            routing_constraints: crate::protocols::RoutingConstraints::default(),
1426            shared_cache_hits: None,
1427            resp_tx: Some(tx),
1428        };
1429        (req, rx)
1430    }
1431
1432    #[tokio::test(flavor = "multi_thread")]
1433    async fn test_cancelled_pending_request_is_not_booked() {
1434        let isl = 512;
1435        let (queue, slots) = make_queue(1, 16, isl, Some(0.0));
1436
1437        let (first, first_rx) = make_request("first", isl);
1438        queue.enqueue(first).await;
1439        first_rx
1440            .await
1441            .expect("first response sender dropped")
1442            .expect("first request should be scheduled");
1443
1444        let (cancelled, cancelled_rx) = make_request("cancelled", isl);
1445        queue.enqueue(cancelled).await;
1446        assert_eq!(queue.pending_count(), 1);
1447        drop(cancelled_rx);
1448
1449        slots.free(&"first".to_string(), decay_now()).unwrap();
1450        queue.update().await;
1451
1452        assert_eq!(queue.pending_count(), 0);
1453        slots.assert_completely_drained(decay_now());
1454    }
1455
1456    #[tokio::test(flavor = "multi_thread")]
1457    async fn test_strict_priority_drains_before_policy_score() {
1458        let isl = 512;
1459        let (queue, slots) = make_queue(1, 16, isl, Some(0.0));
1460
1461        let (first, first_rx) = make_request("first", isl);
1462        queue.enqueue(first).await;
1463        first_rx.await.unwrap().unwrap();
1464
1465        let (mut low, mut low_rx) = make_request("low", isl);
1466        low.priority_jump = 10_000.0;
1467        queue.enqueue(low).await;
1468
1469        let (mut high, high_rx) = make_request("high", isl);
1470        high.strict_priority = 1;
1471        queue.enqueue(high).await;
1472        assert_eq!(queue.pending_count(), 2);
1473
1474        slots.free(&"first".to_string(), decay_now()).unwrap();
1475        queue.update().await;
1476
1477        let high_response = high_rx.await.unwrap().unwrap();
1478        assert_eq!(high_response.best_worker, WorkerWithDpRank::new(0, 0));
1479        assert!(
1480            low_rx.try_recv().is_err(),
1481            "lower strict priority should remain queued"
1482        );
1483
1484        slots.free(&"high".to_string(), decay_now()).unwrap();
1485        queue.update().await;
1486        low_rx.await.unwrap().unwrap();
1487        assert_eq!(queue.pending_count(), 0);
1488
1489        slots.free(&"low".to_string(), decay_now()).unwrap();
1490        slots.assert_completely_drained(decay_now());
1491    }
1492
1493    #[tokio::test(flavor = "multi_thread")]
1494    async fn test_failed_response_delivery_rolls_back_booking() {
1495        let isl = 512;
1496        let response_rx = Arc::new(StdMutex::new(None));
1497        let publisher = DropResponseOnLoadPublisher {
1498            response_rx: Arc::clone(&response_rx),
1499        };
1500        let slots = Arc::new(ActiveSequencesMultiWorker::new(
1501            publisher,
1502            16,
1503            HashMap::from([(0, (0, 1))]),
1504            false,
1505            0,
1506            "test",
1507        ));
1508        let (_cfg_tx, cfg_rx) = watch::channel(HashMap::from([(
1509            0,
1510            SimpleWorkerConfig {
1511                max_num_batched_tokens: Some(isl as u64),
1512                ..Default::default()
1513            },
1514        )]));
1515        let queue = SchedulerQueue::new(
1516            Arc::clone(&slots),
1517            cfg_rx,
1518            None,
1519            16,
1520            DefaultWorkerSelector::new(None, "test"),
1521            RouterQueuePolicy::Fcfs,
1522            None,
1523        );
1524
1525        let (request, receiver) = make_request("delivery-race", isl);
1526        *response_rx.lock().unwrap() = Some(receiver);
1527        queue.enqueue(request).await;
1528
1529        assert!(response_rx.lock().unwrap().is_none());
1530        slots.assert_completely_drained(decay_now());
1531    }
1532
1533    #[tokio::test(flavor = "multi_thread")]
1534    async fn test_concurrent_flood() {
1535        let block_size = 16;
1536        let isl = 512;
1537        let num_workers = 4;
1538        let num_tasks = 25;
1539
1540        let (queue, slots) = make_queue(num_workers, block_size, isl, None);
1541
1542        let mut handles = Vec::new();
1543        for i in 0..num_tasks {
1544            let queue = Arc::clone(&queue);
1545            let slots = Arc::clone(&slots);
1546            handles.push(tokio::spawn(async move {
1547                let req_id = format!("req-{i}");
1548                let (req, rx) = make_request(&req_id, isl);
1549                queue.enqueue(req).await;
1550                let resp = rx.await.expect("oneshot dropped");
1551                let resp = resp.expect("scheduling failed");
1552                assert!(resp.best_worker.worker_id < num_workers as u64);
1553
1554                slots.mark_prefill_completed(&req_id, decay_now()).unwrap();
1555                slots.free(&req_id, decay_now()).unwrap();
1556                queue.update().await;
1557            }));
1558        }
1559
1560        for h in handles {
1561            h.await.expect("task panicked");
1562        }
1563
1564        let active = slots.active_tokens(decay_now());
1565        for (worker, tokens) in &active {
1566            assert_eq!(
1567                *tokens, 0,
1568                "worker {worker:?} still has {tokens} active tokens"
1569            );
1570        }
1571    }
1572
1573    #[tokio::test(flavor = "multi_thread")]
1574    async fn test_concurrent_immediate_admissions_see_prior_booking() {
1575        let selector = MinDecodeSelector {
1576            rendezvous: Some(Arc::new(SelectorRendezvous::default())),
1577        };
1578        let (queue, slots) = make_queue_with_custom_selector(2, 16, 512, None, selector);
1579        let barrier = Arc::new(Barrier::new(3));
1580
1581        let (req1, rx1) = make_request("req-1", 512);
1582        let queue1 = Arc::clone(&queue);
1583        let barrier1 = Arc::clone(&barrier);
1584        let handle1 = tokio::spawn(async move {
1585            barrier1.wait().await;
1586            queue1.enqueue(req1).await;
1587        });
1588
1589        let (req2, rx2) = make_request("req-2", 512);
1590        let queue2 = Arc::clone(&queue);
1591        let barrier2 = Arc::clone(&barrier);
1592        let handle2 = tokio::spawn(async move {
1593            barrier2.wait().await;
1594            queue2.enqueue(req2).await;
1595        });
1596
1597        barrier.wait().await;
1598        handle1.await.unwrap();
1599        handle2.await.unwrap();
1600
1601        let resp1 = rx1.await.unwrap().unwrap();
1602        let resp2 = rx2.await.unwrap().unwrap();
1603        assert_ne!(
1604            resp1.best_worker, resp2.best_worker,
1605            "second admission should see the first booking and choose the other idle worker"
1606        );
1607
1608        for request_id in ["req-1", "req-2"] {
1609            slots
1610                .mark_prefill_completed(&request_id.to_string(), decay_now())
1611                .unwrap();
1612            slots.free(&request_id.to_string(), decay_now()).unwrap();
1613        }
1614    }
1615
1616    #[tokio::test(flavor = "multi_thread")]
1617    async fn test_queueing_under_pressure() {
1618        let block_size = 16;
1619        let isl = 512;
1620        let num_workers = 2;
1621        let num_requests = 10;
1622
1623        let (queue, slots) = make_queue(num_workers, block_size, isl, Some(0.0));
1624
1625        let mut receivers = Vec::new();
1626        let mut req_ids = Vec::new();
1627
1628        for i in 0..num_requests {
1629            let req_id = format!("pressure-{i}");
1630            let (req, rx) = make_request(&req_id, isl);
1631            queue.enqueue(req).await;
1632            receivers.push(rx);
1633            req_ids.push(req_id);
1634        }
1635
1636        // Drain pending by cycling mark_prefill_completed + free + update
1637        // on already-scheduled requests until all receivers have a response.
1638        for _ in 0..num_requests {
1639            queue.update().await;
1640            for rid in &req_ids {
1641                let _ = slots.mark_prefill_completed(rid, decay_now());
1642                let _ = slots.free(rid, decay_now());
1643            }
1644        }
1645        queue.update().await;
1646
1647        let mut ok_count = 0;
1648        for mut rx in receivers {
1649            if let Ok(result) = rx.try_recv() {
1650                result.expect("scheduling returned error");
1651                ok_count += 1;
1652            }
1653        }
1654        assert_eq!(ok_count, num_requests, "not all requests were scheduled");
1655    }
1656
1657    #[tokio::test(flavor = "multi_thread")]
1658    async fn test_pending_requests_receive_shutdown_on_queue_drop() {
1659        let block_size = 16;
1660        let isl = 512;
1661        let (queue, _slots) = make_queue(1, block_size, isl, Some(0.0));
1662
1663        let (req1, rx1) = make_request("req-1", isl);
1664        queue.enqueue(req1).await;
1665        rx1.await
1666            .expect("first response sender dropped")
1667            .expect("first request should be scheduled");
1668
1669        let (req2, rx2) = make_request("req-2", isl);
1670        queue.enqueue(req2).await;
1671        assert_eq!(queue.pending_count(), 1);
1672
1673        drop(queue);
1674
1675        let response = tokio::time::timeout(Duration::from_secs(1), rx2)
1676            .await
1677            .expect("shutdown response timed out")
1678            .expect("pending response sender dropped");
1679        assert!(matches!(
1680            response,
1681            Err(KvSchedulerError::SubscriberShutdown)
1682        ));
1683    }
1684
1685    #[tokio::test(flavor = "multi_thread")]
1686    async fn test_pending_count() {
1687        let block_size = 16;
1688        let isl = 512;
1689        let num_workers = 1;
1690
1691        // threshold_frac=0.0 means any active tokens trigger queueing
1692        let (queue, slots) = make_queue(num_workers, block_size, isl, Some(0.0));
1693        assert_eq!(queue.pending_count(), 0);
1694
1695        // First request goes through (worker is idle)
1696        let (req1, rx1) = make_request("req-1", isl);
1697        queue.enqueue(req1).await;
1698        let _resp1 = rx1.await.unwrap().unwrap();
1699        assert_eq!(queue.pending_count(), 0); // scheduled immediately
1700
1701        // Second and third requests should be queued (worker is now prefill-busy)
1702        let (req2, _rx2) = make_request("req-2", isl);
1703        queue.enqueue(req2).await;
1704        assert_eq!(queue.pending_count(), 1);
1705
1706        let (req3, _rx3) = make_request("req-3", isl);
1707        queue.enqueue(req3).await;
1708        assert_eq!(queue.pending_count(), 2);
1709
1710        // Free the first request and update — should drain one from pending
1711        slots
1712            .mark_prefill_completed(&"req-1".to_string(), decay_now())
1713            .unwrap();
1714        slots.free(&"req-1".to_string(), decay_now()).unwrap();
1715        queue.update().await;
1716
1717        // After update, one pending request should have been scheduled
1718        assert!(
1719            queue.pending_count() < 2,
1720            "pending_count should decrease after free+update, got {}",
1721            queue.pending_count()
1722        );
1723
1724        // Free req-2 and update to drain remaining
1725        let _ = slots.mark_prefill_completed(&"req-2".to_string(), decay_now());
1726        let _ = slots.free(&"req-2".to_string(), decay_now());
1727        queue.update().await;
1728        let _ = slots.mark_prefill_completed(&"req-3".to_string(), decay_now());
1729        let _ = slots.free(&"req-3".to_string(), decay_now());
1730        queue.update().await;
1731
1732        assert_eq!(queue.pending_count(), 0, "all requests should be drained");
1733    }
1734
1735    #[tokio::test(flavor = "multi_thread")]
1736    async fn policy_classes_apply_independent_thresholds_and_preserve_backlog_order() {
1737        let profile = policy_profile(
1738            r#"
1739default_policy_family: latency
1740uncached_isl_buckets:
1741  - min_tokens: 0
1742    bucket: all
1743policy_classes:
1744  - name: latency
1745    policy_family: latency
1746    cache_bucket: all
1747    quantum: 1
1748    prefill_busy_threshold: 0
1749  - name: bulk
1750    policy_family: bulk
1751    cache_bucket: all
1752    quantum: 1
1753    prefill_busy_threshold: 1024
1754"#,
1755        );
1756        let (queue, slots) = make_queue_with_profile(1, 16, 64, profile);
1757
1758        let (mut active, active_rx) = make_request("active", 64);
1759        active.policy_class = Some("latency".to_string());
1760        queue.enqueue(active).await;
1761        active_rx.await.unwrap().unwrap();
1762
1763        let (mut bulk, bulk_rx) = make_request("bulk", 64);
1764        bulk.policy_class = Some("bulk".to_string());
1765        queue.enqueue(bulk).await;
1766        bulk_rx.await.unwrap().unwrap();
1767
1768        let (mut queued_first, mut queued_first_rx) = make_request("queued-first", 64);
1769        queued_first.policy_class = Some("latency".to_string());
1770        queue.enqueue(queued_first).await;
1771        assert_eq!(queue.pending_count(), 1);
1772
1773        for request_id in ["active", "bulk"] {
1774            slots
1775                .mark_prefill_completed(&request_id.to_string(), decay_now())
1776                .unwrap();
1777            slots.free(&request_id.to_string(), decay_now()).unwrap();
1778        }
1779
1780        let (mut queued_second, mut queued_second_rx) = make_request("queued-second", 64);
1781        queued_second.policy_class = Some("latency".to_string());
1782        queue.enqueue(queued_second).await;
1783        assert_eq!(
1784            queue.pending_count(),
1785            2,
1786            "new arrivals must not bypass backlog"
1787        );
1788        assert!(queued_first_rx.try_recv().is_err());
1789        assert!(queued_second_rx.try_recv().is_err());
1790
1791        queue.update().await;
1792        queued_first_rx
1793            .try_recv()
1794            .expect("first queued request should be admitted")
1795            .expect("first queued request failed");
1796        assert!(
1797            queued_second_rx.try_recv().is_err(),
1798            "second request should remain behind the admitted head"
1799        );
1800
1801        slots
1802            .mark_prefill_completed(&"queued-first".to_string(), decay_now())
1803            .unwrap();
1804        slots
1805            .free(&"queued-first".to_string(), decay_now())
1806            .unwrap();
1807        queue.update().await;
1808        queued_second_rx.await.unwrap().unwrap();
1809    }
1810
1811    #[tokio::test(flavor = "multi_thread")]
1812    async fn policy_families_and_cache_buckets_select_physical_queues() {
1813        let profile = policy_profile(
1814            r#"
1815default_policy_family: standard
1816uncached_isl_buckets:
1817  - min_tokens: 0
1818    bucket: cached
1819  - min_tokens: 32
1820    bucket: uncached
1821policy_classes:
1822  - name: cached
1823    policy_family: standard
1824    cache_bucket: cached
1825    quantum: 1
1826    prefill_busy_threshold: 0
1827  - name: uncached
1828    policy_family: standard
1829    cache_bucket: uncached
1830    quantum: 1
1831    prefill_busy_threshold: 0
1832  - name: latency_cached
1833    policy_family: latency
1834    cache_bucket: cached
1835    quantum: 1
1836    prefill_busy_threshold: 0
1837  - name: latency_uncached
1838    policy_family: latency
1839    cache_bucket: uncached
1840    quantum: 1
1841    prefill_busy_threshold: 0
1842  - name: custom_priority
1843    quantum: 1
1844    prefill_busy_threshold: 0
1845"#,
1846        );
1847        let (queue, _slots) = make_queue_with_profile(1, 16, 64, profile);
1848        let worker = WorkerWithDpRank::new(0, 0);
1849
1850        let (active, active_rx) = make_request("active", 64);
1851        queue.enqueue(active).await;
1852        active_rx.await.unwrap().unwrap();
1853
1854        let (mut latency_cached, _latency_cached_rx) = make_request("latency-cached", 64);
1855        latency_cached.policy_class = Some("latency".to_string());
1856        latency_cached
1857            .overlap
1858            .effective_cached_tokens
1859            .insert(worker, 64);
1860        queue.enqueue(latency_cached).await;
1861
1862        let (mut latency_uncached, _latency_uncached_rx) = make_request("latency-uncached", 64);
1863        latency_uncached.policy_class = Some("latency".to_string());
1864        queue.enqueue(latency_uncached).await;
1865
1866        let (mut unknown_cached, _unknown_cached_rx) = make_request("unknown-cached", 64);
1867        unknown_cached.policy_class = Some("unknown".to_string());
1868        unknown_cached
1869            .overlap
1870            .effective_cached_tokens
1871            .insert(worker, 64);
1872        queue.enqueue(unknown_cached).await;
1873
1874        let (mut ordinary_class_name, _ordinary_class_name_rx) =
1875            make_request("ordinary-class-name", 64);
1876        ordinary_class_name.policy_class = Some("latency_cached".to_string());
1877        queue.enqueue(ordinary_class_name).await;
1878
1879        let (mut custom, _custom_rx) = make_request("custom", 64);
1880        custom.policy_class = Some("custom_priority".to_string());
1881        queue.enqueue(custom).await;
1882
1883        assert_eq!(
1884            queue.class_queue_stats(0),
1885            Some(ClassQueueStats {
1886                pending_count: 1,
1887                pending_isl_tokens: 64,
1888                pending_cached_tokens: 64,
1889            })
1890        );
1891        assert_eq!(
1892            queue.class_queue_stats(1),
1893            Some(ClassQueueStats {
1894                pending_count: 1,
1895                pending_isl_tokens: 64,
1896                pending_cached_tokens: 0,
1897            })
1898        );
1899        assert_eq!(
1900            queue.class_queue_stats(2),
1901            Some(ClassQueueStats {
1902                pending_count: 1,
1903                pending_isl_tokens: 64,
1904                pending_cached_tokens: 64,
1905            })
1906        );
1907        assert_eq!(
1908            queue.class_queue_stats(3),
1909            Some(ClassQueueStats {
1910                pending_count: 1,
1911                pending_isl_tokens: 64,
1912                pending_cached_tokens: 0,
1913            })
1914        );
1915        assert_eq!(
1916            queue.class_queue_stats(4),
1917            Some(ClassQueueStats {
1918                pending_count: 1,
1919                pending_isl_tokens: 64,
1920                pending_cached_tokens: 0,
1921            })
1922        );
1923    }
1924
1925    #[tokio::test(flavor = "multi_thread")]
1926    async fn class_local_limit_rejection_is_typed_and_not_overload() {
1927        let profile = policy_profile(
1928            r#"
1929default_policy_family: capped
1930uncached_isl_buckets:
1931  - min_tokens: 0
1932    bucket: all
1933policy_classes:
1934  - name: capped
1935    policy_family: capped
1936    cache_bucket: all
1937    quantum: 1
1938    prefill_busy_threshold: 0
1939    request_queue_limit_per_worker: 1
1940"#,
1941        );
1942        let (queue, _slots) = make_queue_with_profile(1, 16, 64, profile);
1943
1944        let (active, active_rx) = make_request("active", 64);
1945        queue.enqueue(active).await;
1946        active_rx.await.unwrap().unwrap();
1947
1948        let (queued, _queued_rx) = make_request("queued", 64);
1949        queue.enqueue(queued).await;
1950
1951        let (rejected, rejected_rx) = make_request("rejected", 64);
1952        queue.enqueue(rejected).await;
1953        let error = rejected_rx.await.unwrap().unwrap_err();
1954        let KvSchedulerError::QueueRejected(rejection) = &error else {
1955            panic!("expected queue rejection, got {error:?}");
1956        };
1957        assert_eq!(rejection.policy_class, "capped");
1958        assert_eq!(rejection.limit_kind, super::super::QueueLimitKind::Requests);
1959        assert_eq!(rejection.current, 1);
1960        assert_eq!(rejection.limit, 1);
1961        assert!(!error.is_overload());
1962
1963        assert_eq!(
1964            queue.class_queue_stats(0),
1965            Some(ClassQueueStats {
1966                pending_count: 1,
1967                pending_isl_tokens: 64,
1968                pending_cached_tokens: 0,
1969            })
1970        );
1971    }
1972
1973    #[tokio::test(flavor = "multi_thread")]
1974    async fn per_worker_limit_tracks_discovered_worker_count_without_evicting() {
1975        let profile = policy_profile(
1976            r#"
1977default_policy_family: capped
1978uncached_isl_buckets:
1979  - min_tokens: 0
1980    bucket: all
1981policy_classes:
1982  - name: capped
1983    policy_family: capped
1984    cache_bucket: all
1985    quantum: 1
1986    prefill_busy_threshold: 0
1987    request_queue_limit_per_worker: 1
1988"#,
1989        );
1990        let (queue, _slots, cfg_tx) = make_queue_with_profile_and_sender(1, 16, 64, profile);
1991
1992        let (active, active_rx) = make_request("active", 64);
1993        queue.enqueue(active).await;
1994        active_rx.await.unwrap().unwrap();
1995
1996        let (first, _first_rx) = make_request("first", 64);
1997        queue.enqueue(first).await;
1998
1999        cfg_tx.send_modify(|configs| {
2000            configs.insert(
2001                1,
2002                SimpleWorkerConfig {
2003                    max_num_batched_tokens: Some(64),
2004                    ..Default::default()
2005                },
2006            );
2007        });
2008        let (second, _second_rx) = make_request("second", 64);
2009        queue.enqueue(second).await;
2010        assert_eq!(queue.pending_count(), 2);
2011
2012        cfg_tx.send_modify(|configs| {
2013            configs.remove(&1);
2014        });
2015        let (rejected, rejected_rx) = make_request("rejected", 64);
2016        queue.enqueue(rejected).await;
2017        let error = rejected_rx.await.unwrap().unwrap_err();
2018        let KvSchedulerError::QueueRejected(rejection) = error else {
2019            panic!("expected queue rejection, got {error:?}");
2020        };
2021        assert_eq!(rejection.current, 2);
2022        assert_eq!(rejection.limit, 1);
2023        assert_eq!(queue.pending_count(), 2);
2024    }
2025
2026    #[tokio::test(start_paused = true)]
2027    async fn test_queue_update_uses_decayed_oldest_prefill_load() {
2028        let estimator: Arc<dyn PrefillLoadEstimator> = Arc::new(FixedPrefillLoadEstimator {
2029            duration: Duration::from_secs(10),
2030        });
2031        let (queue, _slots, _cfg_tx) =
2032            make_queue_with_sender(1, 16, 100, Some(0.5), Some(estimator));
2033
2034        let (req1, rx1) = make_request("req-1", 100);
2035        queue.enqueue(req1).await;
2036        let _ = rx1.await.unwrap().unwrap();
2037
2038        let (req2, mut rx2) = make_request("req-2", 100);
2039        queue.enqueue(req2).await;
2040        assert_eq!(queue.pending_count(), 1);
2041
2042        tokio::time::advance(Duration::from_secs(6)).await;
2043        queue.update().await;
2044
2045        let scheduled = rx2
2046            .try_recv()
2047            .expect("queued request should have been scheduled");
2048        let response = scheduled.expect("scheduling returned error");
2049        assert_eq!(response.best_worker.worker_id, 0);
2050        assert_eq!(queue.pending_count(), 0);
2051    }
2052
2053    #[tokio::test(flavor = "multi_thread")]
2054    async fn test_overloaded_provider_filters_at_admission() {
2055        let overloaded_worker_provider: OverloadedWorkerProvider =
2056            Arc::new(|| Some(HashSet::from([0])));
2057        let (queue, _slots) =
2058            make_queue_with_overload_provider(1, 16, 256, overloaded_worker_provider);
2059
2060        let (req, rx) = make_request("overloaded", 256);
2061        queue.enqueue(req).await;
2062
2063        let resp = rx.await.expect("oneshot dropped");
2064        assert!(matches!(
2065            resp,
2066            Err(KvSchedulerError::AllEligibleWorkersOverloaded)
2067        ));
2068    }
2069
2070    /// Simulates the EPP path: router starts with zero workers (skip_initial_worker_wait),
2071    /// then register_workers lazily injects workers before routing.
2072    #[tokio::test(flavor = "multi_thread")]
2073    async fn test_register_workers_lazy_epp_path() {
2074        let block_size = 16;
2075        let isl = 512;
2076
2077        // Start with zero workers (mimics skip_initial_worker_wait=true)
2078        let (queue, slots, cfg_tx) = make_queue_with_sender(0, block_size, isl, None, None);
2079
2080        // Routing with no workers must fail
2081        let (req_fail, rx_fail) = make_request("before-register", isl);
2082        queue.enqueue(req_fail).await;
2083        let resp = rx_fail.await.expect("oneshot dropped");
2084        assert!(
2085            matches!(
2086                resp,
2087                Err(crate::scheduling::types::KvSchedulerError::NoEndpoints)
2088            ),
2089            "expected NoEndpoints before register_workers, got {resp:?}"
2090        );
2091
2092        // Lazily register two workers in the slot tracker (EPP supplies pod list)
2093        slots.upsert_worker(WorkerDpRange::new(100, 0, 1)).unwrap();
2094        slots.upsert_worker(WorkerDpRange::new(200, 0, 1)).unwrap();
2095
2096        // Also update the config watch so the selector can see these workers
2097        let mut configs = HashMap::new();
2098        for &id in &[100_u64, 200_u64] {
2099            configs.insert(
2100                id,
2101                SimpleWorkerConfig {
2102                    max_num_batched_tokens: Some(isl as u64),
2103                    ..Default::default()
2104                },
2105            );
2106        }
2107        cfg_tx.send(configs).unwrap();
2108
2109        // Routing after registration must succeed and pick one of the registered workers
2110        let (req_ok, rx_ok) = make_request("after-register", isl);
2111        queue.enqueue(req_ok).await;
2112        let resp = rx_ok
2113            .await
2114            .expect("oneshot dropped")
2115            .expect("scheduling failed");
2116        assert!(
2117            resp.best_worker.worker_id == 100 || resp.best_worker.worker_id == 200,
2118            "expected worker 100 or 200, got {}",
2119            resp.best_worker.worker_id
2120        );
2121
2122        // Clean up
2123        slots
2124            .mark_prefill_completed(&"after-register".to_string(), decay_now())
2125            .unwrap();
2126        slots
2127            .free(&"after-register".to_string(), decay_now())
2128            .unwrap();
2129    }
2130
2131    /// Register_workers is additive: calling with a new set does NOT remove old workers.
2132    #[tokio::test(flavor = "multi_thread")]
2133    async fn test_register_workers_additive() {
2134        let block_size = 16;
2135        let isl = 256;
2136
2137        let (queue, slots, cfg_tx) = make_queue_with_sender(0, block_size, isl, None, None);
2138
2139        // Register worker 10 in slots and config
2140        slots.upsert_worker(WorkerDpRange::new(10, 0, 1)).unwrap();
2141
2142        let mut configs = HashMap::new();
2143        configs.insert(
2144            10_u64,
2145            SimpleWorkerConfig {
2146                max_num_batched_tokens: Some(isl as u64),
2147                ..Default::default()
2148            },
2149        );
2150        cfg_tx.send(configs.clone()).unwrap();
2151
2152        // Register worker 20 (worker 10 must NOT be evicted)
2153        slots.upsert_worker(WorkerDpRange::new(20, 0, 1)).unwrap();
2154
2155        configs.insert(
2156            20_u64,
2157            SimpleWorkerConfig {
2158                max_num_batched_tokens: Some(isl as u64),
2159                ..Default::default()
2160            },
2161        );
2162        cfg_tx.send(configs).unwrap();
2163
2164        // Send enough requests to statistically prove both workers are available
2165        let mut seen = std::collections::HashSet::new();
2166        for i in 0..20 {
2167            let req_id = format!("add-{i}");
2168            let (req, rx) = make_request(&req_id, isl);
2169            queue.enqueue(req).await;
2170            let resp = rx
2171                .await
2172                .expect("oneshot dropped")
2173                .expect("scheduling failed");
2174            seen.insert(resp.best_worker.worker_id);
2175            slots.mark_prefill_completed(&req_id, decay_now()).unwrap();
2176            slots.free(&req_id, decay_now()).unwrap();
2177        }
2178
2179        assert!(
2180            seen.contains(&10) && seen.contains(&20),
2181            "both workers should be reachable after additive registration, saw: {seen:?}"
2182        );
2183    }
2184
2185    #[tokio::test(flavor = "multi_thread")]
2186    async fn allowed_worker_request_joins_backlog_and_dispatches_within_allow_list() {
2187        let block_size = 16;
2188        let isl = 256;
2189        let (queue, slots) = make_queue(2, block_size, isl, Some(0.0));
2190
2191        let (active_a, active_a_rx) = make_request("active-a", isl);
2192        queue.enqueue(active_a).await;
2193        let active_a_worker = active_a_rx.await.unwrap().unwrap().best_worker.worker_id;
2194
2195        let (active_b, active_b_rx) = make_request("active-b", isl);
2196        queue.enqueue(active_b).await;
2197        active_b_rx.await.unwrap().unwrap();
2198
2199        let (backlog_head, backlog_head_rx) = make_request("backlog-head", isl);
2200        queue.enqueue(backlog_head).await;
2201        assert_eq!(queue.pending_count(), 1);
2202
2203        slots
2204            .mark_prefill_completed(&"active-a".to_string(), decay_now())
2205            .unwrap();
2206        slots.free(&"active-a".to_string(), decay_now()).unwrap();
2207
2208        let (mut allowed, mut allowed_rx) = make_request("allowed", isl);
2209        allowed.allowed_worker_ids = Some(HashSet::from([active_a_worker]));
2210        queue.enqueue(allowed).await;
2211        assert_eq!(
2212            queue.pending_count(),
2213            2,
2214            "allow-list request must not bypass the existing class backlog"
2215        );
2216        assert!(allowed_rx.try_recv().is_err());
2217
2218        queue.update().await;
2219        let backlog_head_worker = backlog_head_rx
2220            .await
2221            .unwrap()
2222            .unwrap()
2223            .best_worker
2224            .worker_id;
2225        assert!(allowed_rx.try_recv().is_err());
2226
2227        slots
2228            .mark_prefill_completed(&"backlog-head".to_string(), decay_now())
2229            .unwrap();
2230        slots
2231            .free(&"backlog-head".to_string(), decay_now())
2232            .unwrap();
2233        queue.update().await;
2234
2235        let allowed_worker = allowed_rx.await.unwrap().unwrap().best_worker.worker_id;
2236        assert_eq!(allowed_worker, active_a_worker);
2237
2238        for request_id in ["active-b", "allowed"] {
2239            slots
2240                .mark_prefill_completed(&request_id.to_string(), decay_now())
2241                .unwrap();
2242            slots.free(&request_id.to_string(), decay_now()).unwrap();
2243        }
2244        assert_eq!(backlog_head_worker, active_a_worker);
2245    }
2246
2247    #[tokio::test(flavor = "multi_thread")]
2248    async fn test_pinned_worker_conflict_with_allowed_ids_fails_early() {
2249        let (queue, _slots) = make_queue(1, 16, 256, Some(0.0));
2250        let (mut req, rx) = make_request("conflict", 256);
2251        req.pinned_worker = Some(WorkerWithDpRank::new(0, 0));
2252        req.allowed_worker_ids = Some(HashSet::from([1]));
2253
2254        queue.enqueue(req).await;
2255
2256        let resp = rx.await.expect("oneshot dropped");
2257        assert!(matches!(
2258            resp,
2259            Err(KvSchedulerError::PinnedWorkerNotAllowed { worker_id: 0 })
2260        ));
2261    }
2262
2263    #[tokio::test(flavor = "multi_thread")]
2264    async fn test_disallowed_worker_ids_fail_without_queueing() {
2265        let (queue, _slots) = make_queue(1, 16, 256, Some(0.0));
2266        let (mut req, rx) = make_request("disallowed", 256);
2267        req.allowed_worker_ids = Some(HashSet::from([999]));
2268
2269        queue.enqueue(req).await;
2270
2271        let resp = rx.await.expect("oneshot dropped");
2272        assert!(matches!(resp, Err(KvSchedulerError::NoEndpoints)));
2273        assert_eq!(queue.pending_count(), 0);
2274    }
2275
2276    #[tokio::test(flavor = "multi_thread")]
2277    async fn test_incompatible_required_taints_fail_without_queueing() {
2278        let (queue, _slots, cfg_tx) = make_queue_with_sender(1, 16, 256, Some(0.0), None);
2279        let mut configs = HashMap::new();
2280        configs.insert(
2281            0_u64,
2282            SimpleWorkerConfig {
2283                max_num_batched_tokens: Some(256),
2284                taints: HashSet::from(["mdc-a".to_string()]),
2285                ..Default::default()
2286            },
2287        );
2288        cfg_tx.send(configs).unwrap();
2289
2290        let (mut req, rx) = make_request("tainted", 256);
2291        req.routing_constraints = crate::protocols::RoutingConstraints {
2292            required_taints: HashSet::from(["mdc-b".to_string()]),
2293            preferred_taints: HashMap::new(),
2294        };
2295
2296        queue.enqueue(req).await;
2297
2298        let resp = rx.await.expect("oneshot dropped");
2299        assert!(matches!(resp, Err(KvSchedulerError::NoEndpoints)));
2300        assert_eq!(queue.pending_count(), 0);
2301    }
2302
2303    #[tokio::test(flavor = "multi_thread")]
2304    async fn test_pinned_head_blocks_class_backlog_despite_other_worker_capacity() {
2305        let (queue, slots) = make_queue(2, 16, 256, Some(0.0));
2306
2307        let (mut first, first_rx) = make_request("pinned-1", 256);
2308        first.pinned_worker = Some(WorkerWithDpRank::new(1, 0));
2309        queue.enqueue(first).await;
2310        let first_resp = first_rx.await.unwrap().unwrap();
2311        assert_eq!(first_resp.best_worker, WorkerWithDpRank::new(1, 0));
2312
2313        let (mut second, mut second_rx) = make_request("pinned-2", 256);
2314        second.pinned_worker = Some(WorkerWithDpRank::new(1, 0));
2315        queue.enqueue(second).await;
2316        assert_eq!(queue.pending_count(), 1);
2317        assert!(
2318            second_rx.try_recv().is_err(),
2319            "request should remain queued"
2320        );
2321
2322        let (unpinned, mut unpinned_rx) = make_request("unpinned", 256);
2323        queue.enqueue(unpinned).await;
2324        assert_eq!(queue.pending_count(), 2);
2325
2326        queue.update().await;
2327
2328        assert_eq!(queue.pending_count(), 2);
2329        assert!(
2330            unpinned_rx.try_recv().is_err(),
2331            "unpinned request should remain queued behind the pinned head"
2332        );
2333        assert!(
2334            second_rx.try_recv().is_err(),
2335            "pinned request should still be queued"
2336        );
2337
2338        slots
2339            .mark_prefill_completed(&"pinned-1".to_string(), decay_now())
2340            .unwrap();
2341        slots.free(&"pinned-1".to_string(), decay_now()).unwrap();
2342        queue.update().await;
2343
2344        let second_resp = second_rx
2345            .try_recv()
2346            .expect("pinned request should have been scheduled");
2347        let second_resp = second_resp.expect("scheduling returned error");
2348        assert_eq!(second_resp.best_worker, WorkerWithDpRank::new(1, 0));
2349
2350        let unpinned_resp = unpinned_rx
2351            .try_recv()
2352            .expect("unpinned request should have been scheduled");
2353        let unpinned_resp = unpinned_resp.expect("scheduling returned error");
2354        assert_eq!(unpinned_resp.best_worker, WorkerWithDpRank::new(0, 0));
2355        assert_eq!(queue.pending_count(), 0);
2356    }
2357
2358    #[tokio::test(flavor = "multi_thread")]
2359    async fn test_queue_prefill_busy_check_ignores_untracked_prefill_tokens() {
2360        let (queue, slots) = make_queue(1, 16, 256, Some(0.0));
2361
2362        let (mut req1, rx1) = make_request("req-1", 256);
2363        req1.track_prefill_tokens = false;
2364        queue.enqueue(req1).await;
2365        let _resp1 = rx1.await.unwrap().unwrap();
2366        assert_eq!(
2367            slots
2368                .active_tokens(decay_now())
2369                .get(&WorkerWithDpRank::new(0, 0))
2370                .copied(),
2371            Some(0)
2372        );
2373
2374        let (req2, rx2) = make_request("req-2", 256);
2375        queue.enqueue(req2).await;
2376        let _resp2 = rx2.await.unwrap().unwrap();
2377        assert_eq!(queue.pending_count(), 0);
2378
2379        let _ = slots.mark_prefill_completed(&"req-1".to_string(), decay_now());
2380        let _ = slots.free(&"req-1".to_string(), decay_now());
2381        let _ = slots.mark_prefill_completed(&"req-2".to_string(), decay_now());
2382        let _ = slots.free(&"req-2".to_string(), decay_now());
2383    }
2384
2385    #[tokio::test(flavor = "current_thread", start_paused = true)]
2386    async fn update_refresh_can_change_selected_worker_after_queue_wait() {
2387        let block_size = 16u32;
2388        let isl = 64usize;
2389        let refresher = Arc::new(CountingRefresher {
2390            calls: AtomicUsize::new(0),
2391            response: RefreshedOverlap {
2392                tier_overlap_blocks: Default::default(),
2393                effective_overlap_blocks: HashMap::from([
2394                    (WorkerWithDpRank::new(0, 0), 1.0),
2395                    (WorkerWithDpRank::new(1, 0), 9.0),
2396                ]),
2397                effective_cached_tokens: HashMap::from([
2398                    (WorkerWithDpRank::new(0, 0), 16),
2399                    (WorkerWithDpRank::new(1, 0), 144),
2400                ]),
2401            },
2402        });
2403        let (queue, slots) =
2404            make_queue_with_refresher(2, block_size, isl, Some(0.0), refresher.clone());
2405
2406        let (mut req1, rx1) = make_request("req-1", isl);
2407        req1.overlap
2408            .effective_overlap_blocks
2409            .insert(WorkerWithDpRank::new(0, 0), 3.0);
2410        req1.overlap
2411            .effective_cached_tokens
2412            .insert(WorkerWithDpRank::new(0, 0), 48);
2413        queue.enqueue(req1).await;
2414        let resp1 = rx1.await.expect("rx1 dropped").expect("req-1 failed");
2415        assert_eq!(resp1.best_worker, WorkerWithDpRank::new(0, 0));
2416
2417        let (mut req2, rx2) = make_request("req-2", isl);
2418        req2.overlap
2419            .effective_overlap_blocks
2420            .insert(WorkerWithDpRank::new(1, 0), 3.0);
2421        req2.overlap
2422            .effective_cached_tokens
2423            .insert(WorkerWithDpRank::new(1, 0), 48);
2424        queue.enqueue(req2).await;
2425        let resp2 = rx2.await.expect("rx2 dropped").expect("req-2 failed");
2426        assert_eq!(resp2.best_worker, WorkerWithDpRank::new(1, 0));
2427
2428        let (mut req3, rx3) = make_request("req-3", isl);
2429        req3.overlap
2430            .effective_overlap_blocks
2431            .insert(WorkerWithDpRank::new(0, 0), 8.0);
2432        req3.overlap
2433            .effective_overlap_blocks
2434            .insert(WorkerWithDpRank::new(1, 0), 2.0);
2435        req3.overlap
2436            .effective_cached_tokens
2437            .insert(WorkerWithDpRank::new(0, 0), 128);
2438        req3.overlap
2439            .effective_cached_tokens
2440            .insert(WorkerWithDpRank::new(1, 0), 32);
2441        queue
2442            .enqueue_with_block_hashes(req3, Some(vec![LocalBlockHash(42)]))
2443            .await;
2444        assert_eq!(queue.pending_count(), 1);
2445        assert_eq!(refresher.calls.load(Ordering::Relaxed), 0);
2446
2447        tokio::time::advance(Duration::from_secs(11)).await;
2448
2449        slots.free(&"req-1".to_string(), decay_now()).unwrap();
2450        slots.free(&"req-2".to_string(), decay_now()).unwrap();
2451        queue.update().await;
2452
2453        let resp3 = rx3.await.expect("rx3 dropped").expect("req-3 failed");
2454        assert_eq!(refresher.calls.load(Ordering::Relaxed), 1);
2455        assert_eq!(resp3.best_worker, WorkerWithDpRank::new(1, 0));
2456        assert_eq!(resp3.effective_overlap_blocks, 9.0);
2457        assert_eq!(resp3.cached_tokens, 144);
2458        assert_eq!(queue.pending_count(), 0);
2459    }
2460
2461    #[tokio::test(flavor = "current_thread", start_paused = true)]
2462    async fn selected_request_dispatches_after_refresh_if_worker_becomes_busy() {
2463        let block_size = 16u32;
2464        let isl = 64usize;
2465        let worker = WorkerWithDpRank::new(0, 0);
2466        let refresher = Arc::new(BlockingRefresher::new(RefreshedOverlap {
2467            tier_overlap_blocks: Default::default(),
2468            effective_overlap_blocks: HashMap::from([(worker, 7.0)]),
2469            effective_cached_tokens: HashMap::from([(worker, 56)]),
2470        }));
2471        let (queue, slots) = make_queue_with_blocking_refresher(
2472            1,
2473            block_size,
2474            isl,
2475            Some(0.0),
2476            refresher.clone(),
2477            ADMISSION_CHANNEL_CAPACITY,
2478        );
2479
2480        let (req1, rx1) = make_request("req-1", isl);
2481        queue.enqueue(req1).await;
2482        let _ = rx1.await.expect("rx1 dropped").expect("req-1 failed");
2483
2484        let (mut req2, rx2) = make_request("req-2", isl);
2485        req2.overlap
2486            .effective_overlap_blocks
2487            .insert(WorkerWithDpRank::new(0, 0), 4.0);
2488        req2.overlap
2489            .effective_cached_tokens
2490            .insert(WorkerWithDpRank::new(0, 0), 64);
2491        queue
2492            .enqueue_with_block_hashes(req2, Some(vec![LocalBlockHash(42)]))
2493            .await;
2494        assert_eq!(queue.pending_count(), 1);
2495        assert_eq!(
2496            queue.class_queue_stats(0).unwrap().pending_cached_tokens,
2497            64
2498        );
2499
2500        slots
2501            .mark_prefill_completed(&"req-1".to_string(), decay_now())
2502            .unwrap();
2503        slots.free(&"req-1".to_string(), decay_now()).unwrap();
2504
2505        tokio::time::advance(Duration::from_secs(11)).await;
2506
2507        let update = {
2508            let queue = Arc::clone(&queue);
2509            tokio::spawn(async move {
2510                queue.update().await;
2511            })
2512        };
2513        refresher.wait_for_calls(1).await;
2514        assert_eq!(
2515            queue.pending_count(),
2516            0,
2517            "DRR-selected request must be removed before refresh"
2518        );
2519        assert_eq!(
2520            queue.class_queue_stats(0).unwrap().pending_cached_tokens,
2521            0,
2522            "queue counters must reflect the irrevocable dequeue"
2523        );
2524
2525        slots
2526            .add_request(
2527                SequenceRequest {
2528                    request_id: "occupy-during-refresh".to_string(),
2529                    token_sequence: None,
2530                    track_prefill_tokens: true,
2531                    expected_output_tokens: None,
2532                    prefill_load_hint: Some(PrefillLoadHint {
2533                        initial_effective_prefill_tokens: isl,
2534                        expected_prefill_duration: None,
2535                    }),
2536                    worker,
2537                    lora_name: None,
2538                },
2539                decay_now(),
2540            )
2541            .unwrap();
2542
2543        refresher.release_one();
2544        update.await.unwrap();
2545
2546        let resp2 = rx2.await.expect("rx2 dropped").expect("req-2 failed");
2547        assert_eq!(refresher.calls.load(Ordering::Relaxed), 1);
2548        assert_eq!(resp2.best_worker, worker);
2549        assert_eq!(resp2.effective_overlap_blocks, 7.0);
2550        assert_eq!(resp2.cached_tokens, 56);
2551        assert_eq!(queue.pending_count(), 0);
2552
2553        for request_id in ["occupy-during-refresh", "req-2"] {
2554            slots
2555                .mark_prefill_completed(&request_id.to_string(), decay_now())
2556                .unwrap();
2557            slots.free(&request_id.to_string(), decay_now()).unwrap();
2558        }
2559    }
2560
2561    #[tokio::test(flavor = "current_thread", start_paused = true)]
2562    async fn continuation_drain_does_not_self_send_into_saturated_actor_channel() {
2563        let block_size = 16u32;
2564        let isl = 64usize;
2565        let refresher = Arc::new(BlockingRefresher::new(RefreshedOverlap::default()));
2566        let (queue, slots) =
2567            make_queue_with_blocking_refresher(1, block_size, isl, Some(0.0), refresher.clone(), 1);
2568
2569        let (active, active_rx) = make_request("active", isl);
2570        queue.enqueue(active).await;
2571        active_rx.await.unwrap().unwrap();
2572
2573        let (queued, queued_rx) = make_request("queued", isl);
2574        queue
2575            .enqueue_with_block_hashes(queued, Some(vec![LocalBlockHash(42)]))
2576            .await;
2577        slots
2578            .mark_prefill_completed(&"active".to_string(), decay_now())
2579            .unwrap();
2580        slots.free(&"active".to_string(), decay_now()).unwrap();
2581        tokio::time::advance(Duration::from_secs(11)).await;
2582
2583        let update = {
2584            let queue = Arc::clone(&queue);
2585            tokio::spawn(async move { queue.update().await })
2586        };
2587        refresher.wait_for_calls(1).await;
2588
2589        let (following, following_rx) = make_request("following", isl);
2590        let enqueue = {
2591            let queue = Arc::clone(&queue);
2592            tokio::spawn(async move { queue.enqueue(following).await })
2593        };
2594        tokio::task::yield_now().await;
2595        assert_eq!(
2596            queue.admission_tx.capacity(),
2597            0,
2598            "test must saturate the actor command channel"
2599        );
2600
2601        refresher.release_one();
2602        tokio::time::timeout(Duration::from_secs(1), update)
2603            .await
2604            .expect("update deadlocked with a full actor command channel")
2605            .unwrap();
2606        queued_rx.await.unwrap().unwrap();
2607
2608        slots
2609            .mark_prefill_completed(&"queued".to_string(), decay_now())
2610            .unwrap();
2611        slots.free(&"queued".to_string(), decay_now()).unwrap();
2612        queue.update().await;
2613        following_rx.await.unwrap().unwrap();
2614        enqueue.await.unwrap();
2615    }
2616}