1use 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
30pub 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
88pub 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 pending_count: Arc<AtomicUsize>,
102 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 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 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 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 pub fn pending_count(&self) -> usize {
399 self.pending_count.load(AtomicOrdering::Relaxed)
400 }
401
402 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 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 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 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 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 !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 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 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 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 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 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 let (queue, slots) = make_queue(num_workers, block_size, isl, Some(0.0));
1693 assert_eq!(queue.pending_count(), 0);
1694
1695 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); 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 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 assert!(
1719 queue.pending_count() < 2,
1720 "pending_count should decrease after free+update, got {}",
1721 queue.pending_count()
1722 );
1723
1724 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 #[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 let (queue, slots, cfg_tx) = make_queue_with_sender(0, block_size, isl, None, None);
2079
2080 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 slots.upsert_worker(WorkerDpRange::new(100, 0, 1)).unwrap();
2094 slots.upsert_worker(WorkerDpRange::new(200, 0, 1)).unwrap();
2095
2096 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 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 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 #[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 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 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 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}