1use std::collections::{HashMap, HashSet, VecDeque};
5use std::marker::PhantomData;
6use std::sync::Arc;
7use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering as AtomicOrdering};
8use std::time::Duration;
9
10use crossbeam_queue::SegQueue;
11use rustc_hash::FxHashSet;
12use tokio::sync::{mpsc, oneshot, watch};
13use tokio::time::Instant;
14
15use super::config::RouterQueuePolicy;
16use super::filter::RoutingEligibility;
17use super::overlap_refresh::{
18 NoopOverlapScoresRefresh, OverlapScoresRefresh, read_overlap_refresh_after, refresh_overlap,
19};
20use super::policy_config::{PolicyClassConfig, PolicyProfile};
21use super::policy_queue::{PolicyQueue, QueueSnapshot};
22use super::prefill_load::{PrefillLoadEstimator, effective_prefill_tokens};
23use super::queue_admission::{
24 AdmissionAction, AdmissionDecision, AdmissionTicket, ClassAdmissionAction,
25 PolicyClassAdmissionPolicies, RequestProgressUpdater, WorkerEligibility,
26 WorkerEligibilitySnapshot, WorkerPlacement,
27};
28use super::selector::{DefaultWorkerSelector, WorkerSelector};
29use super::types::{
30 KvSchedulerError, OverloadedWorkerProvider, SchedulingContext, SchedulingRequest,
31 SchedulingResponse,
32};
33use crate::protocols::{
34 LocalBlockHash, PrefillLoadHint, WorkerConfigLike, WorkerId, WorkerWithDpRank,
35};
36use crate::sequences::topology::WorkerDpRange;
37use crate::sequences::{ActiveSequencesMultiWorker, SequencePublisher, SequenceRequest};
38
39pub const DEFAULT_MAX_BATCHED_TOKENS: u64 = 10_000_000;
41
42const ADMISSION_CHANNEL_CAPACITY: usize = 65_536;
43
44struct ClassQueueCounters {
45 pending_count: AtomicUsize,
46 pending_isl_tokens: AtomicUsize,
47 pending_cached_tokens: AtomicUsize,
48}
49
50#[derive(Debug, Clone, Copy, PartialEq, Eq)]
51pub struct ClassQueueStats {
52 pub pending_count: usize,
53 pub pending_isl_tokens: usize,
54 pub pending_cached_tokens: usize,
55}
56
57struct QueuedRequest {
58 request: SchedulingRequest,
59 enqueue_at: Instant,
60 block_hashes: Option<Vec<LocalBlockHash>>,
61 admission: Option<RequestAdmission>,
62}
63
64struct RequestAdmission {
65 ticket: AdmissionTicket,
66 progress: RequestProgressUpdater,
67}
68
69#[allow(clippy::large_enum_variant)]
70enum AdmissionCommand {
71 Enqueue {
72 request: SchedulingRequest,
73 block_hashes: Option<Vec<LocalBlockHash>>,
74 lease: Option<Box<RequestLifecycleLease>>,
75 ack_tx: oneshot::Sender<Option<Box<RequestLifecycleLease>>>,
76 },
77 Update {
78 worker: Option<WorkerWithDpRank>,
79 ack_tx: oneshot::Sender<()>,
80 },
81 Reconcile {
82 force: bool,
83 ack_tx: oneshot::Sender<()>,
84 },
85 Dispatched {
86 request_id: String,
87 ticket: AdmissionTicket,
88 },
89 Cleanup,
90}
91
92#[derive(Debug, Clone, Copy)]
93struct TrackedAdmission {
94 ticket: AdmissionTicket,
95 queue_class_index: Option<usize>,
96 worker: Option<WorkerWithDpRank>,
97 dispatched: bool,
98}
99
100#[derive(Debug, PartialEq, Eq)]
101struct AdmissionCleanupEntry {
102 ticket: Option<AdmissionTicket>,
103 request_id: String,
104 context_tokens: Option<usize>,
105 dispatched: bool,
106}
107
108#[derive(Default)]
109struct AdmissionCleanup {
110 dirty: SegQueue<AdmissionCleanupEntry>,
111 pending: AtomicBool,
112}
113
114impl AdmissionCleanup {
115 fn enqueue(&self, cleanup: AdmissionCleanupEntry) -> bool {
116 self.dirty.push(cleanup);
117 !self.pending.swap(true, AtomicOrdering::AcqRel)
118 }
119
120 fn drain(&self) -> Vec<AdmissionCleanupEntry> {
121 if !self.pending.load(AtomicOrdering::Acquire) {
122 return Vec::new();
123 }
124
125 let mut dirty = Vec::new();
129 loop {
130 while let Some(cleanup) = self.dirty.pop() {
131 dirty.push(cleanup);
132 }
133 self.pending.store(false, AtomicOrdering::Release);
134 if self.dirty.is_empty() {
135 return dirty;
136 }
137 self.pending.store(true, AtomicOrdering::Release);
138 }
139 }
140}
141
142#[must_use = "dropping the lease reports the request outcome to the scheduler actor"]
147pub struct RequestLifecycleLease {
148 cleanup: Arc<AdmissionCleanup>,
149 actor_tx: mpsc::Sender<AdmissionCommand>,
150 ticket: Option<AdmissionTicket>,
151 request_id: Option<String>,
152 context_tokens: Option<usize>,
153 dispatched: bool,
154}
155
156impl std::fmt::Debug for RequestLifecycleLease {
157 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
158 formatter
159 .debug_struct("RequestLifecycleLease")
160 .field("ticket", &self.ticket)
161 .field("request_id", &self.request_id)
162 .field("context_tokens", &self.context_tokens)
163 .field("dispatched", &self.dispatched)
164 .finish_non_exhaustive()
165 }
166}
167
168impl RequestLifecycleLease {
169 pub fn mark_completed(&mut self, context_tokens: usize) {
170 self.context_tokens = Some(context_tokens);
171 }
172
173 pub async fn mark_dispatched(&mut self) {
174 self.dispatched = true;
175 let (Some(request_id), Some(ticket)) = (self.request_id.clone(), self.ticket) else {
176 return;
177 };
178 let _ = self
179 .actor_tx
180 .send(AdmissionCommand::Dispatched { request_id, ticket })
181 .await;
182 }
183
184 pub(crate) fn disarm(&mut self) {
185 self.request_id = None;
186 }
187}
188
189impl Drop for RequestLifecycleLease {
190 fn drop(&mut self) {
191 let Some(request_id) = self.request_id.take() else {
192 return;
193 };
194 if self.cleanup.enqueue(AdmissionCleanupEntry {
195 ticket: self.ticket,
196 request_id,
197 context_tokens: self.context_tokens,
198 dispatched: self.dispatched,
199 }) {
200 let _ = self.actor_tx.try_send(AdmissionCommand::Cleanup);
201 }
202 }
203}
204
205struct SchedulerQueueActor<
206 P: SequencePublisher,
207 C: WorkerConfigLike,
208 Sel: WorkerSelector<C>,
209 RF: OverlapScoresRefresh,
210> {
211 pending: PolicyQueue<QueuedRequest>,
212 tracked_admissions: HashMap<String, TrackedAdmission>,
213 cleanup: Arc<AdmissionCleanup>,
214 queueing_enabled: bool,
215 profile: PolicyProfile,
216 queue_recheck_interval: Duration,
217 next_queue_recheck: Instant,
218 pending_count: Arc<AtomicUsize>,
219 pending_isl_tokens: Arc<AtomicUsize>,
220 class_counters: Arc<Vec<ClassQueueCounters>>,
221 slots: Arc<ActiveSequencesMultiWorker<P>>,
222 workers_with_configs: watch::Receiver<HashMap<WorkerId, C>>,
223 start_time: Instant,
224 block_size: u32,
225 selector: Sel,
226 prefill_load_estimator: Option<Arc<dyn PrefillLoadEstimator>>,
227 overlap_scores_refresh: Option<Arc<RF>>,
228 overlap_refresh_after: Option<Duration>,
229 overloaded_worker_provider: Option<OverloadedWorkerProvider>,
230}
231
232pub struct SchedulerQueue<
237 P: SequencePublisher,
238 C: WorkerConfigLike,
239 Sel: WorkerSelector<C> = DefaultWorkerSelector,
240 RF: OverlapScoresRefresh = NoopOverlapScoresRefresh,
241> {
242 admission_tx: mpsc::Sender<AdmissionCommand>,
243 cleanup: Arc<AdmissionCleanup>,
244 pending_count: Arc<AtomicUsize>,
247 pending_isl_tokens: Arc<AtomicUsize>,
250 class_counters: Arc<Vec<ClassQueueCounters>>,
251 slots: Arc<ActiveSequencesMultiWorker<P>>,
252 workers_with_configs: watch::Receiver<HashMap<WorkerId, C>>,
253 queueing_enabled: bool,
254 admission_enabled: bool,
255 supports_overlap_refresh: bool,
256 _marker: PhantomData<(Sel, RF)>,
257}
258
259impl<
260 P: SequencePublisher + 'static,
261 C: WorkerConfigLike + Send + Sync + 'static,
262 Sel: WorkerSelector<C> + Send + 'static,
263 RF: OverlapScoresRefresh + Send + Sync + 'static,
264> SchedulerQueue<P, C, Sel, RF>
265{
266 #[allow(clippy::too_many_arguments)]
267 pub fn new_with_overlap_refresh(
268 slots: Arc<ActiveSequencesMultiWorker<P>>,
269 workers_with_configs: watch::Receiver<HashMap<WorkerId, C>>,
270 threshold_frac: Option<f64>,
271 block_size: u32,
272 selector: Sel,
273 queue_policy: RouterQueuePolicy,
274 prefill_load_estimator: Option<Arc<dyn PrefillLoadEstimator>>,
275 overlap_scores_refresh: Option<Arc<RF>>,
276 overloaded_worker_provider: Option<OverloadedWorkerProvider>,
277 ) -> Self {
278 let profile = PolicyProfile::synthetic(threshold_frac, queue_policy);
279 Self::new_with_policy_profile(
280 slots,
281 workers_with_configs,
282 profile,
283 block_size,
284 selector,
285 prefill_load_estimator,
286 overlap_scores_refresh,
287 overloaded_worker_provider,
288 Duration::from_secs(60),
289 PolicyClassAdmissionPolicies::new(),
290 )
291 .expect("synthetic policy profile does not require admission policies")
292 }
293
294 #[allow(clippy::too_many_arguments)]
295 pub fn new_with_policy_profile(
296 slots: Arc<ActiveSequencesMultiWorker<P>>,
297 workers_with_configs: watch::Receiver<HashMap<WorkerId, C>>,
298 profile: PolicyProfile,
299 block_size: u32,
300 selector: Sel,
301 prefill_load_estimator: Option<Arc<dyn PrefillLoadEstimator>>,
302 overlap_scores_refresh: Option<Arc<RF>>,
303 overloaded_worker_provider: Option<OverloadedWorkerProvider>,
304 queue_recheck_interval: Duration,
305 admission_policies: PolicyClassAdmissionPolicies,
306 ) -> Result<Self, KvSchedulerError> {
307 Self::new_with_policy_profile_and_capacity(
308 slots,
309 workers_with_configs,
310 profile,
311 block_size,
312 selector,
313 prefill_load_estimator,
314 overlap_scores_refresh,
315 overloaded_worker_provider,
316 queue_recheck_interval,
317 admission_policies,
318 ADMISSION_CHANNEL_CAPACITY,
319 )
320 }
321
322 #[allow(clippy::too_many_arguments)]
323 fn new_with_policy_profile_and_capacity(
324 slots: Arc<ActiveSequencesMultiWorker<P>>,
325 workers_with_configs: watch::Receiver<HashMap<WorkerId, C>>,
326 profile: PolicyProfile,
327 block_size: u32,
328 selector: Sel,
329 prefill_load_estimator: Option<Arc<dyn PrefillLoadEstimator>>,
330 overlap_scores_refresh: Option<Arc<RF>>,
331 overloaded_worker_provider: Option<OverloadedWorkerProvider>,
332 queue_recheck_interval: Duration,
333 admission_policies: PolicyClassAdmissionPolicies,
334 admission_channel_capacity: usize,
335 ) -> Result<Self, KvSchedulerError> {
336 let admission_enabled = !admission_policies.is_empty();
337 let pending = PolicyQueue::new_with_admission_policies(
338 profile.clone(),
339 queue_recheck_interval,
340 admission_policies,
341 )?;
342 let queueing_enabled = profile
343 .classes()
344 .iter()
345 .any(PolicyClassConfig::queueing_enabled)
346 || admission_enabled;
347 for class in profile.classes() {
348 tracing::info!(
349 policy_class = class.name,
350 queue_policy = %class.queue_policy,
351 quantum = class.quantum,
352 prefill_busy_threshold = ?class.prefill_busy_threshold,
353 prefill_busy_threshold_frac = ?class.prefill_busy_threshold_frac,
354 "Router policy class configured"
355 );
356 }
357 let overlap_refresh_after = if overlap_scores_refresh.is_some() {
358 let configured = read_overlap_refresh_after();
359 match configured {
360 Some(d) => tracing::info!(
361 "Router queue overlap-score refresh enabled after {:.1}s wait",
362 d.as_secs_f64()
363 ),
364 None => tracing::info!(
365 "Router queue overlap-score refresh disabled via DYN_ROUTER_OVERLAP_REFRESH_AFTER_SECS"
366 ),
367 }
368 configured
369 } else {
370 None
371 };
372 let pending_count = Arc::new(AtomicUsize::new(0));
373 let pending_isl_tokens = Arc::new(AtomicUsize::new(0));
374 let class_counters = Arc::new(
375 profile
376 .classes()
377 .iter()
378 .map(|_| ClassQueueCounters {
379 pending_count: AtomicUsize::new(0),
380 pending_isl_tokens: AtomicUsize::new(0),
381 pending_cached_tokens: AtomicUsize::new(0),
382 })
383 .collect(),
384 );
385 let (admission_tx, admission_rx) = mpsc::channel(admission_channel_capacity);
386 let cleanup = Arc::new(AdmissionCleanup::default());
387 let now = Instant::now();
388 let actor = SchedulerQueueActor {
389 pending,
390 tracked_admissions: HashMap::new(),
391 cleanup: Arc::clone(&cleanup),
392 queueing_enabled,
393 profile,
394 queue_recheck_interval,
395 next_queue_recheck: now + queue_recheck_interval,
396 pending_count: Arc::clone(&pending_count),
397 pending_isl_tokens: Arc::clone(&pending_isl_tokens),
398 class_counters: Arc::clone(&class_counters),
399 slots: Arc::clone(&slots),
400 workers_with_configs: workers_with_configs.clone(),
401 start_time: Instant::now(),
402 block_size,
403 selector,
404 prefill_load_estimator,
405 overlap_scores_refresh,
406 overlap_refresh_after,
407 overloaded_worker_provider,
408 };
409 tokio::spawn(actor.run(admission_rx));
410 Ok(Self {
411 admission_tx,
412 cleanup,
413 pending_count,
414 pending_isl_tokens,
415 class_counters,
416 slots,
417 workers_with_configs,
418 queueing_enabled,
419 admission_enabled,
420 supports_overlap_refresh: overlap_refresh_after.is_some(),
421 _marker: PhantomData,
422 })
423 }
424}
425
426impl<
427 P: SequencePublisher + 'static,
428 C: WorkerConfigLike + Send + Sync + 'static,
429 Sel: WorkerSelector<C> + Send + 'static,
430> SchedulerQueue<P, C, Sel, NoopOverlapScoresRefresh>
431{
432 #[allow(clippy::too_many_arguments)]
433 pub fn new(
434 slots: Arc<ActiveSequencesMultiWorker<P>>,
435 workers_with_configs: watch::Receiver<HashMap<WorkerId, C>>,
436 threshold_frac: Option<f64>,
437 block_size: u32,
438 selector: Sel,
439 queue_policy: RouterQueuePolicy,
440 prefill_load_estimator: Option<Arc<dyn PrefillLoadEstimator>>,
441 ) -> Self {
442 Self::new_with_overlap_refresh(
443 slots,
444 workers_with_configs,
445 threshold_frac,
446 block_size,
447 selector,
448 queue_policy,
449 prefill_load_estimator,
450 None,
451 None,
452 )
453 }
454
455 #[allow(clippy::too_many_arguments)]
456 pub fn new_with_overload_provider(
457 slots: Arc<ActiveSequencesMultiWorker<P>>,
458 workers_with_configs: watch::Receiver<HashMap<WorkerId, C>>,
459 threshold_frac: Option<f64>,
460 block_size: u32,
461 selector: Sel,
462 queue_policy: RouterQueuePolicy,
463 prefill_load_estimator: Option<Arc<dyn PrefillLoadEstimator>>,
464 overloaded_worker_provider: Option<OverloadedWorkerProvider>,
465 ) -> Self {
466 Self::new_with_overlap_refresh(
467 slots,
468 workers_with_configs,
469 threshold_frac,
470 block_size,
471 selector,
472 queue_policy,
473 prefill_load_estimator,
474 None,
475 overloaded_worker_provider,
476 )
477 }
478}
479
480impl<
481 P: SequencePublisher + 'static,
482 C: WorkerConfigLike + Send + Sync + 'static,
483 Sel: WorkerSelector<C> + Send + 'static,
484 RF: OverlapScoresRefresh + Send + Sync + 'static,
485> SchedulerQueue<P, C, Sel, RF>
486{
487 pub fn register_workers(&self, worker_ids: &std::collections::HashSet<u64>) {
492 let discovery_workers = self.workers_with_configs.borrow();
493 for &worker_id in worker_ids {
494 let (dp_start, dp_size) = discovery_workers
495 .get(&worker_id)
496 .map(|runtime_config| {
497 (
498 runtime_config.data_parallel_start_rank(),
499 runtime_config.data_parallel_size(),
500 )
501 })
502 .unwrap_or((0, 1));
503 let range = WorkerDpRange::new(worker_id, dp_start, dp_size);
504 if let Err(error) = self.slots.upsert_worker(range) {
505 tracing::warn!(worker_id, %error, "Invalid externally-provided worker topology");
506 }
507 }
508 }
509
510 pub async fn enqueue(&self, request: SchedulingRequest) {
514 self.enqueue_with_block_hashes(request, None).await;
515 }
516
517 pub async fn enqueue_with_block_hashes(
518 &self,
519 request: SchedulingRequest,
520 block_hashes: Option<Vec<LocalBlockHash>>,
521 ) {
522 let _ = self
523 .enqueue_with_block_hashes_and_lease(request, block_hashes, None)
524 .await;
525 }
526
527 pub(crate) async fn enqueue_with_block_hashes_and_lease(
528 &self,
529 mut request: SchedulingRequest,
530 block_hashes: Option<Vec<LocalBlockHash>>,
531 lease: Option<Box<RequestLifecycleLease>>,
532 ) -> Option<Box<RequestLifecycleLease>> {
533 if self.queueing_enabled && lease.is_none() && request.mode.lifecycle_request_id().is_some()
534 {
535 request.respond(Err(KvSchedulerError::BookingFailed(
536 "admission-managed requests must be scheduled through LocalScheduler".to_string(),
537 )));
538 return None;
539 }
540
541 let eligibility = request.eligibility();
542
543 if let Err(error) = eligibility.validate_pinned_worker_allowed() {
544 request.respond(Err(error));
545 return None;
546 }
547
548 let (ack_tx, ack_rx) = oneshot::channel();
549 let command = AdmissionCommand::Enqueue {
550 request,
551 block_hashes: self.prepare_block_hashes_for_refresh(block_hashes),
552 lease,
553 ack_tx,
554 };
555
556 if let Err(error) = self.admission_tx.send(command).await {
557 let AdmissionCommand::Enqueue { mut request, .. } = error.0 else {
558 return None;
559 };
560 request.respond(Err(KvSchedulerError::SubscriberShutdown));
561 return None;
562 }
563
564 match ack_rx.await {
565 Ok(lease) => lease,
566 Err(_) => {
567 tracing::warn!("scheduler queue actor dropped enqueue acknowledgement");
568 None
569 }
570 }
571 }
572
573 pub(crate) fn new_request_lifecycle_lease(
574 &self,
575 request_id: Option<&str>,
576 ) -> Option<Box<RequestLifecycleLease>> {
577 if !self.queueing_enabled {
578 return None;
579 }
580 request_id?;
581 Some(Box::new(RequestLifecycleLease {
582 cleanup: Arc::clone(&self.cleanup),
583 actor_tx: self.admission_tx.clone(),
584 ticket: None,
585 request_id: None,
586 context_tokens: None,
587 dispatched: false,
588 }))
589 }
590
591 pub async fn update(&self) {
595 self.update_after(None).await;
596 }
597
598 pub(crate) async fn update_worker(&self, worker: WorkerWithDpRank) {
599 self.update_after(Some(worker)).await;
600 }
601
602 async fn update_after(&self, worker: Option<WorkerWithDpRank>) {
603 if !self.queueing_enabled {
604 return;
605 }
606
607 let (ack_tx, ack_rx) = oneshot::channel();
608 if self
609 .admission_tx
610 .send(AdmissionCommand::Update { worker, ack_tx })
611 .await
612 .is_ok()
613 {
614 let _ = ack_rx.await;
615 }
616 }
617
618 #[cfg(test)]
619 pub(crate) async fn reconcile(&self) {
620 self.send_reconcile(true).await;
621 }
622
623 pub(crate) async fn periodic_reconcile(&self) {
624 self.send_reconcile(false).await;
625 }
626
627 async fn send_reconcile(&self, force: bool) {
628 if !self.admission_enabled {
629 self.update().await;
630 return;
631 }
632
633 let (ack_tx, ack_rx) = oneshot::channel();
634 if self
635 .admission_tx
636 .send(AdmissionCommand::Reconcile { force, ack_tx })
637 .await
638 .is_ok()
639 {
640 let _ = ack_rx.await;
641 }
642 }
643
644 pub fn pending_count(&self) -> usize {
646 self.pending_count.load(AtomicOrdering::Relaxed)
647 }
648
649 pub fn pending_isl_tokens(&self) -> usize {
651 self.pending_isl_tokens.load(AtomicOrdering::Relaxed)
652 }
653
654 pub fn class_queue_stats(&self, class_index: usize) -> Option<ClassQueueStats> {
655 let counters = self.class_counters.get(class_index)?;
656 Some(ClassQueueStats {
657 pending_count: counters.pending_count.load(AtomicOrdering::Relaxed),
658 pending_isl_tokens: counters.pending_isl_tokens.load(AtomicOrdering::Relaxed),
659 pending_cached_tokens: counters.pending_cached_tokens.load(AtomicOrdering::Relaxed),
660 })
661 }
662
663 pub fn supports_overlap_refresh(&self) -> bool {
664 self.supports_overlap_refresh
665 }
666
667 fn prepare_block_hashes_for_refresh(
668 &self,
669 block_hashes: Option<Vec<LocalBlockHash>>,
670 ) -> Option<Vec<LocalBlockHash>> {
671 if !self.supports_overlap_refresh {
672 return None;
673 }
674 block_hashes.filter(|hashes| !hashes.is_empty())
675 }
676}
677
678impl<
679 P: SequencePublisher + 'static,
680 C: WorkerConfigLike + Send + Sync + 'static,
681 Sel: WorkerSelector<C> + Send + 'static,
682 RF: OverlapScoresRefresh + Send + Sync + 'static,
683> SchedulerQueueActor<P, C, Sel, RF>
684{
685 async fn run(mut self, mut rx: mpsc::Receiver<AdmissionCommand>) {
686 let mut commands_since_cleanup = 0usize;
687 while let Some(command) = rx.recv().await {
688 let drain_cleanup = self.queueing_enabled && {
689 commands_since_cleanup += 1;
690 let drain_cleanup = rx.is_empty() || commands_since_cleanup == 256;
691 if drain_cleanup {
692 commands_since_cleanup = 0;
693 }
694 drain_cleanup
695 };
696 match command {
697 AdmissionCommand::Enqueue {
698 request,
699 block_hashes,
700 mut lease,
701 ack_tx,
702 } => {
703 let request_id = lease
704 .as_ref()
705 .and_then(|_| request.mode.tracked_request_id().map(str::to_owned));
706 let (enqueue_ready, owns_lifecycle) =
707 self.handle_enqueue(request, block_hashes);
708 if let Some(lease) = lease.as_mut()
709 && owns_lifecycle
710 {
711 lease.ticket = request_id.as_deref().and_then(|request_id| {
712 self.tracked_admissions
713 .get(request_id)
714 .map(|tracked| tracked.ticket)
715 });
716 lease.request_id = request_id;
717 }
718 let made_ready = enqueue_ready | (drain_cleanup && self.drain_cleanup());
719 if made_ready {
720 self.handle_update(None).await;
721 }
722 let _ = ack_tx.send(lease);
723 }
724 AdmissionCommand::Update { worker, ack_tx } => {
725 self.handle_update(worker).await;
726 if drain_cleanup && self.drain_cleanup() {
727 self.handle_update(None).await;
728 }
729 let _ = ack_tx.send(());
730 }
731 AdmissionCommand::Reconcile { force, ack_tx } => {
732 self.handle_reconcile(force).await;
733 if drain_cleanup && self.drain_cleanup() {
734 self.handle_update(None).await;
735 }
736 let _ = ack_tx.send(());
737 }
738 AdmissionCommand::Dispatched { request_id, ticket } => {
739 if self.handle_dispatched(&request_id, ticket)
740 | (drain_cleanup && self.drain_cleanup())
741 {
742 self.handle_update(None).await;
743 }
744 }
745 AdmissionCommand::Cleanup => {
746 if self.drain_cleanup() {
747 self.handle_update(None).await;
748 }
749 }
750 }
751 }
752 self.drain_cleanup();
753
754 let class_counters = Arc::clone(&self.class_counters);
755 for entry in self.pending.drain() {
756 let class_index = entry.class_index();
757 let snapshot = entry.snapshot();
758 self.pending_count.fetch_sub(1, AtomicOrdering::Relaxed);
759 self.pending_isl_tokens
760 .fetch_sub(snapshot.raw_isl_tokens, AtomicOrdering::Relaxed);
761 let counters = &class_counters[class_index];
762 counters.pending_count.fetch_sub(1, AtomicOrdering::Relaxed);
763 counters
764 .pending_isl_tokens
765 .fetch_sub(snapshot.raw_isl_tokens, AtomicOrdering::Relaxed);
766 counters
767 .pending_cached_tokens
768 .fetch_sub(snapshot.cached_tokens, AtomicOrdering::Relaxed);
769
770 let mut request = entry.into_payload().request;
771 request.respond(Err(KvSchedulerError::SubscriberShutdown));
772 }
773 }
774
775 fn handle_enqueue(
776 &mut self,
777 mut request: SchedulingRequest,
778 block_hashes: Option<Vec<LocalBlockHash>>,
779 ) -> (bool, bool) {
780 let decay_now = Instant::now();
781 let (admission_class_index, mut snapshot) = if let Some(class_index) = self
784 .profile
785 .direct_class_index(request.policy_class.as_deref())
786 {
787 (class_index, None)
788 } else {
789 let workers = self.workers_with_configs.borrow();
790 let snapshot = Self::snapshot_for_with(&request, &workers);
791 let class_index = self
792 .profile
793 .resolve_class_index(request.policy_class.as_deref(), snapshot.uncached_tokens);
794 (class_index, Some(snapshot))
795 };
796 let mut queue_class_index = admission_class_index;
797 if let Some(request_id) = request.mode.lifecycle_request_id()
798 && self.tracked_admissions.contains_key(request_id)
799 {
800 request.respond(Err(KvSchedulerError::BookingFailed(format!(
801 "request {request_id} already has an active admission"
802 ))));
803 return (false, false);
804 }
805
806 let has_admission_policy = self.pending.has_admission_policy(admission_class_index);
807 if request.mode.is_tracked()
808 && request.mode.lifecycle_request_id().is_none()
809 && has_admission_policy
810 {
811 request.respond(Err(KvSchedulerError::BookingFailed(format!(
812 "policy class {:?} requires lifecycle-tracked scheduling",
813 self.profile.class(admission_class_index).name
814 ))));
815 return (false, false);
816 }
817
818 let mut admission = if request.mode.lifecycle_request_id().is_some() && has_admission_policy
819 {
820 let allowed_worker_ids = request.allowed_worker_ids.clone();
821 let pinned_worker = request.pinned_worker;
822 let routing_constraints = request.routing_constraints.clone();
823 let workers = self.workers_with_configs.clone();
824 let overloaded_worker_provider = self.overloaded_worker_provider.clone();
825 let worker_eligibility = WorkerEligibility::new(move || {
826 let workers = workers.borrow();
827 let overloaded_worker_ids = overloaded_worker_provider
828 .as_ref()
829 .and_then(|provider| provider());
830 let structural_eligibility = RoutingEligibility::new(
831 allowed_worker_ids.as_ref(),
832 None,
833 pinned_worker,
834 &routing_constraints,
835 );
836 let mut structural_workers = FxHashSet::default();
837 structural_eligibility.for_each_eligible_worker_rank(&workers, |worker, _| {
838 structural_workers.insert(worker);
839 });
840 let Some(overloaded_worker_ids) = overloaded_worker_ids.as_ref() else {
841 return WorkerEligibilitySnapshot::new(structural_workers);
842 };
843 let mut available_workers = structural_workers.clone();
844 available_workers
845 .retain(|worker| !overloaded_worker_ids.contains(&worker.worker_id));
846 WorkerEligibilitySnapshot::with_availability(structural_workers, available_workers)
847 });
848 self.pending
849 .admit(
850 admission_class_index,
851 request.session_id.as_deref(),
852 request.isl_tokens,
853 worker_eligibility,
854 )
855 .map(|(ticket, progress, decision)| {
856 (RequestAdmission { ticket, progress }, decision)
857 })
858 } else {
859 None
860 };
861 let mut deferred = false;
862 if let Some((request_admission, decision)) = admission.as_ref() {
863 match decision {
864 AdmissionDecision::Bypass => admission = None,
865 AdmissionDecision::Ready(placement) => {
866 if let Err(error) = apply_admission_placement(&mut request, *placement) {
867 request.respond(Err(error));
868 return (self.abort_admission(request_admission.ticket), false);
869 }
870 if matches!(placement, WorkerPlacement::Exact(_)) {
871 let exact_snapshot = self.snapshot_for(&request);
872 queue_class_index = self.profile.resolve_class_index(
873 request.policy_class.as_deref(),
874 exact_snapshot.uncached_tokens,
875 );
876 snapshot = Some(exact_snapshot);
877 }
878 }
879 AdmissionDecision::Defer => deferred = true,
880 }
881 }
882
883 let class = self.profile.class(queue_class_index);
884 let should_queue = deferred
885 || self.should_queue(queue_class_index, class, || {
886 self.all_workers_prefill_busy(class, request.eligibility(), decay_now)
887 });
888 if !should_queue {
889 return self.admit_one(
890 request,
891 decay_now,
892 admission.map(|(admission, _)| admission),
893 );
894 }
895
896 let snapshot = snapshot.unwrap_or_else(|| self.snapshot_for(&request));
897 tracing::debug!(policy_class = class.name, deferred, "queueing request");
898 let arrival_offset = self.start_time.elapsed().as_secs_f64();
899 let priority_jump = request.priority_jump;
900 let strict_priority = request.strict_priority;
901 let placement = request
902 .pinned_worker
903 .map_or(WorkerPlacement::Any, WorkerPlacement::Exact);
904 let deferred_id = deferred
905 .then(|| admission.as_ref().map(|(admission, _)| admission.ticket.id))
906 .flatten();
907 let tracked_admission = admission.as_ref().map(|(admission, _)| {
908 (
909 request
910 .mode
911 .tracked_request_id()
912 .expect("admitted request is tracked")
913 .to_owned(),
914 admission.ticket,
915 )
916 });
917 let queued = QueuedRequest {
918 request,
919 enqueue_at: decay_now,
920 block_hashes,
921 admission: admission.map(|(admission, _)| admission),
922 };
923 let worker_count = self.workers_with_configs.borrow().len();
924 let enqueue = match deferred_id {
925 Some(admission_id) => self.pending.enqueue_deferred(
926 queue_class_index,
927 worker_count,
928 snapshot,
929 arrival_offset,
930 priority_jump,
931 strict_priority,
932 admission_id,
933 queued,
934 ),
935 None => self.pending.enqueue(
936 queue_class_index,
937 worker_count,
938 snapshot,
939 arrival_offset,
940 priority_jump,
941 strict_priority,
942 placement,
943 queued,
944 ),
945 };
946 if let Err((rejection, queued)) = enqueue {
947 let made_ready = queued
948 .admission
949 .as_ref()
950 .is_some_and(|admission| self.abort_admission(admission.ticket));
951 let mut request = queued.request;
952 request.respond(Err(KvSchedulerError::QueueRejected(rejection)));
953 return (made_ready, false);
954 }
955 if let Some((request_id, ticket)) = tracked_admission {
956 self.tracked_admissions.insert(
957 request_id,
958 TrackedAdmission {
959 ticket,
960 queue_class_index: Some(queue_class_index),
961 worker: None,
962 dispatched: false,
963 },
964 );
965 }
966 self.pending_count.fetch_add(1, AtomicOrdering::Relaxed);
967 self.pending_isl_tokens
968 .fetch_add(snapshot.raw_isl_tokens, AtomicOrdering::Relaxed);
969 self.add_class_counters(queue_class_index, snapshot);
970 (false, true)
971 }
972
973 fn should_queue(
974 &self,
975 class_index: usize,
976 class: &PolicyClassConfig,
977 all_workers_busy: impl FnOnce() -> bool,
978 ) -> bool {
979 class.queueing_enabled() && (self.pending.has_backlog(class_index) || all_workers_busy())
982 }
983
984 fn snapshot_for(&self, request: &SchedulingRequest) -> QueueSnapshot {
985 let workers = self.workers_with_configs.borrow();
986 Self::snapshot_for_with(request, &workers)
987 }
988
989 fn snapshot_for_with(
990 request: &SchedulingRequest,
991 workers: &HashMap<WorkerId, C>,
992 ) -> QueueSnapshot {
993 let context = SchedulingContext::new(request, workers);
996 QueueSnapshot::new(request.isl_tokens, context.best_cached_tokens())
997 }
998
999 fn handle_dispatched(&mut self, request_id: &str, ticket: AdmissionTicket) -> bool {
1000 let Some(tracked) = self.tracked_admissions.get_mut(request_id) else {
1001 return false;
1002 };
1003 if tracked.ticket != ticket {
1004 return false;
1005 }
1006 if tracked.dispatched {
1007 return false;
1008 }
1009 let Some(worker) = tracked.worker else {
1010 tracing::debug!(%request_id, "Ignoring dispatch before queue admission");
1011 return false;
1012 };
1013 tracked.dispatched = true;
1014 let actions = self.pending.dispatched(tracked.ticket, worker);
1015 self.apply_admission_actions(actions)
1016 }
1017
1018 fn drain_cleanup(&mut self) -> bool {
1019 let dirty = self.cleanup.drain();
1020 if dirty.is_empty() {
1021 return false;
1022 }
1023
1024 let mut made_ready = false;
1025 let mut removed_ready_head = false;
1026 let mut ready_by_class: HashMap<usize, FxHashSet<_>> = HashMap::new();
1027 let mut unmanaged_request_ids = HashSet::new();
1028 for cleanup in dirty {
1029 let request_id = &cleanup.request_id;
1030 let tracked = self.tracked_admissions.get(request_id).copied();
1031 if tracked.is_some_and(|tracked| Some(tracked.ticket) != cleanup.ticket) {
1032 continue;
1033 }
1034 if cleanup.dispatched
1035 && let Some(ticket) = cleanup.ticket
1036 {
1037 made_ready |= self.handle_dispatched(request_id, ticket);
1038 }
1039
1040 if tracked.is_none_or(|tracked| tracked.worker.is_some())
1041 && self.slots.request_worker(request_id).is_some()
1042 {
1043 if let Err(error) = self.slots.free(request_id, Instant::now()) {
1044 tracing::error!(%request_id, %error, "Failed to release dropped scheduler booking");
1045 }
1046 made_ready = true;
1047 }
1048
1049 let Some(tracked) = tracked else {
1050 unmanaged_request_ids.insert(cleanup.request_id);
1051 continue;
1052 };
1053 if tracked.worker.is_some() {
1054 self.tracked_admissions.remove(request_id);
1055 made_ready |= self.finish_admission(tracked.ticket, cleanup.context_tokens);
1056 continue;
1057 }
1058
1059 self.tracked_admissions.remove(request_id);
1060 let queue_class_index = tracked
1061 .queue_class_index
1062 .expect("queued admission must retain its physical class");
1063 if let Some(entry) = self
1064 .pending
1065 .remove_deferred(queue_class_index, tracked.ticket.id)
1066 {
1067 self.subtract_pending_counters(entry.class_index(), entry.snapshot());
1068 } else {
1069 ready_by_class
1070 .entry(queue_class_index)
1071 .or_default()
1072 .insert(tracked.ticket.id);
1073 }
1074 made_ready |= self.finish_admission(tracked.ticket, cleanup.context_tokens);
1075 }
1076
1077 for (class_index, tickets) in ready_by_class {
1080 let (removed, class_head_removed) =
1081 self.pending.take_if_in_class(class_index, |queued| {
1082 queued
1083 .admission
1084 .as_ref()
1085 .is_some_and(|admission| tickets.contains(&admission.ticket.id))
1086 });
1087 debug_assert_eq!(removed.len(), tickets.len());
1088 removed_ready_head |= class_head_removed;
1089 for entry in removed {
1090 self.subtract_pending_counters(class_index, entry.snapshot());
1091 }
1092 }
1093 if !unmanaged_request_ids.is_empty() {
1094 for class_index in 0..self.profile.classes().len() {
1095 let (removed, class_head_removed) =
1096 self.pending.take_if_in_class(class_index, |queued| {
1097 queued
1098 .request
1099 .mode
1100 .tracked_request_id()
1101 .is_some_and(|request_id| unmanaged_request_ids.contains(request_id))
1102 });
1103 removed_ready_head |= class_head_removed;
1104 for entry in removed {
1105 self.subtract_pending_counters(class_index, entry.snapshot());
1106 }
1107 }
1108 }
1109 made_ready || (removed_ready_head && self.has_dispatchable_ready_head())
1110 }
1111
1112 fn has_dispatchable_ready_head(&self) -> bool {
1113 let active_tokens = self.slots.active_tokens(Instant::now());
1114 let configs = self.workers_with_configs.borrow();
1115 self.pending.any_ready_head(|_, class, queued| {
1116 !Self::all_workers_prefill_busy_with(
1117 &active_tokens,
1118 &configs,
1119 class,
1120 queued.request.eligibility(),
1121 )
1122 })
1123 }
1124
1125 fn finish_admission(&mut self, ticket: AdmissionTicket, context_tokens: Option<usize>) -> bool {
1126 match context_tokens {
1127 Some(context_tokens) => self.complete_admission(ticket, context_tokens),
1128 None => self.abort_admission(ticket),
1129 }
1130 }
1131
1132 fn complete_admission(&mut self, ticket: AdmissionTicket, context_tokens: usize) -> bool {
1133 let actions = self.pending.completed(ticket, context_tokens);
1134 self.apply_admission_actions(actions)
1135 }
1136
1137 fn abort_admission(&mut self, ticket: AdmissionTicket) -> bool {
1138 let actions = self.pending.aborted(ticket);
1139 self.apply_admission_actions(actions)
1140 }
1141
1142 fn apply_admission_actions(
1143 &mut self,
1144 actions: impl IntoIterator<Item = ClassAdmissionAction>,
1145 ) -> bool {
1146 let mut made_ready = false;
1147 let mut actions: VecDeque<_> = actions.into_iter().collect();
1148 while let Some(class_action) = actions.pop_front() {
1149 let AdmissionAction::MakeReady { id, placement } = class_action.action;
1150
1151 let prepared = {
1152 let Some(queued) = self
1153 .pending
1154 .deferred_payload_mut(class_action.class_index, id)
1155 else {
1156 tracing::debug!(
1157 admission_id = id.get(),
1158 "Ignoring unknown make-ready action"
1159 );
1160 continue;
1161 };
1162 match apply_admission_placement(&mut queued.request, placement) {
1163 Err(error) => Err(error),
1164 Ok(()) => {
1165 let effective_placement = queued
1166 .request
1167 .pinned_worker
1168 .map_or(placement, WorkerPlacement::Exact);
1169 let replacement =
1170 if matches!(effective_placement, WorkerPlacement::Exact(_)) {
1171 let workers = self.workers_with_configs.borrow();
1172 Some((
1173 Self::snapshot_for_with(&queued.request, &workers),
1174 queued
1175 .enqueue_at
1176 .duration_since(self.start_time)
1177 .as_secs_f64(),
1178 queued.request.priority_jump,
1179 ))
1180 } else {
1181 None
1182 };
1183 let target_class_index = replacement.as_ref().map_or(
1184 class_action.class_index,
1185 |(snapshot, _, _)| {
1186 self.profile.resolve_class_index(
1187 queued.request.policy_class.as_deref(),
1188 snapshot.uncached_tokens,
1189 )
1190 },
1191 );
1192 let request_id =
1193 (target_class_index != class_action.class_index).then(|| {
1194 queued
1195 .request
1196 .mode
1197 .tracked_request_id()
1198 .expect("admitted request is tracked")
1199 .to_owned()
1200 });
1201 Ok((
1202 effective_placement,
1203 target_class_index,
1204 replacement,
1205 request_id,
1206 ))
1207 }
1208 }
1209 };
1210
1211 let (effective_placement, target_class_index, replacement, request_id) = match prepared
1212 {
1213 Ok(prepared) => prepared,
1214 Err(error) => {
1215 let Some(entry) = self.pending.remove_deferred(class_action.class_index, id)
1216 else {
1217 continue;
1218 };
1219 let snapshot = entry.snapshot();
1220 self.subtract_pending_counters(class_action.class_index, snapshot);
1221 let mut queued = entry.into_payload();
1222 if let Some(admission) = queued.admission {
1223 if let Some(request_id) = queued.request.mode.tracked_request_id() {
1224 self.tracked_admissions.remove(request_id);
1225 }
1226 actions.extend(self.pending.aborted(admission.ticket));
1227 }
1228 queued.request.respond(Err(error));
1229 continue;
1230 }
1231 };
1232
1233 let new_snapshot = replacement.map(|(snapshot, _, _)| snapshot);
1234 if let Some(old_snapshot) = self.pending.make_ready(
1235 class_action.class_index,
1236 target_class_index,
1237 id,
1238 effective_placement,
1239 replacement,
1240 ) {
1241 if let Some(new_snapshot) = new_snapshot {
1242 if class_action.class_index == target_class_index {
1243 self.replace_pending_snapshot_counters(
1244 class_action.class_index,
1245 old_snapshot,
1246 new_snapshot,
1247 );
1248 } else {
1249 self.subtract_class_counters(class_action.class_index, old_snapshot);
1250 self.add_class_counters(target_class_index, new_snapshot);
1251 let tracked = self
1252 .tracked_admissions
1253 .get_mut(request_id.as_deref().expect("reclassified request ID"))
1254 .expect("reclassified admission must be tracked");
1255 tracked.queue_class_index = Some(target_class_index);
1256 }
1257 }
1258 made_ready = true;
1259 } else {
1260 tracing::debug!(
1261 admission_id = id.get(),
1262 "Ignoring duplicate make-ready action"
1263 );
1264 }
1265 }
1266 made_ready
1267 }
1268
1269 fn subtract_pending_counters(&self, class_index: usize, snapshot: QueueSnapshot) {
1270 self.pending_count.fetch_sub(1, AtomicOrdering::Relaxed);
1271 self.pending_isl_tokens
1272 .fetch_sub(snapshot.raw_isl_tokens, AtomicOrdering::Relaxed);
1273 self.subtract_class_counters(class_index, snapshot);
1274 }
1275
1276 fn replace_pending_snapshot_counters(
1277 &self,
1278 class_index: usize,
1279 old: QueueSnapshot,
1280 new: QueueSnapshot,
1281 ) {
1282 debug_assert_eq!(old.raw_isl_tokens, new.raw_isl_tokens);
1283 let counter = &self.class_counters[class_index].pending_cached_tokens;
1284 if new.cached_tokens >= old.cached_tokens {
1285 counter.fetch_add(
1286 new.cached_tokens - old.cached_tokens,
1287 AtomicOrdering::Relaxed,
1288 );
1289 } else {
1290 counter.fetch_sub(
1291 old.cached_tokens - new.cached_tokens,
1292 AtomicOrdering::Relaxed,
1293 );
1294 }
1295 }
1296
1297 async fn handle_reconcile(&mut self, force: bool) {
1298 let now = Instant::now();
1299 let queue_due = force || now >= self.next_queue_recheck;
1300 if !force && queue_due {
1301 self.next_queue_recheck = now + self.queue_recheck_interval;
1302 }
1303 let actions = self.pending.reconcile_admission(now, force);
1304 let made_ready = self.apply_admission_actions(actions);
1305 if queue_due || made_ready {
1306 self.handle_update(None).await;
1307 }
1308 }
1309
1310 async fn handle_update(&mut self, worker: Option<WorkerWithDpRank>) {
1311 if !self.pending.has_ready() {
1312 return;
1313 }
1314
1315 if let Some(worker) = worker {
1316 self.pending.recheck_worker(worker);
1317 } else {
1318 self.pending.recheck_all_workers();
1321 }
1322
1323 loop {
1326 let decay_now = Instant::now();
1327 let active_tokens = self.slots.active_tokens(decay_now);
1328 let popped = {
1329 let configs = self.workers_with_configs.borrow();
1330 self.pending.pop_next(|_, class, queued| {
1331 !Self::all_workers_prefill_busy_with(
1335 &active_tokens,
1336 &configs,
1337 class,
1338 queued.request.eligibility(),
1339 )
1340 })
1341 };
1342 let Some(mut popped) = popped else {
1343 break;
1344 };
1345 let snapshot = popped.snapshot();
1346 let current_pending_count = self.pending_count.load(AtomicOrdering::Relaxed);
1347 debug_assert!(
1348 current_pending_count > 0,
1349 "pending_count underflow on queue drain"
1350 );
1351 self.pending_count.fetch_sub(1, AtomicOrdering::Relaxed);
1352 let current_pending_isl_tokens = self.pending_isl_tokens.load(AtomicOrdering::Relaxed);
1353 debug_assert!(
1354 current_pending_isl_tokens >= snapshot.raw_isl_tokens,
1355 "pending_isl_tokens underflow: pending={} request_isl_tokens={}",
1356 current_pending_isl_tokens,
1357 snapshot.raw_isl_tokens
1358 );
1359 self.pending_isl_tokens
1360 .fetch_sub(snapshot.raw_isl_tokens, AtomicOrdering::Relaxed);
1361 self.subtract_class_counters(popped.class_index(), snapshot);
1362 let queued = popped.payload_mut();
1363 let refreshed = refresh_overlap(
1368 self.overlap_scores_refresh.as_deref(),
1369 self.overlap_refresh_after,
1370 queued.block_hashes.as_deref(),
1371 queued.enqueue_at,
1372 decay_now,
1373 )
1374 .await;
1375 let wait_ms = queued.enqueue_at.elapsed().as_millis() as u64;
1376 if let Some(overlap) = refreshed {
1377 tracing::info!(
1378 request_id = queued.request.mode.request_id().unwrap_or("unknown"),
1379 wait_ms,
1380 "refreshed overlap scores after long queue wait"
1381 );
1382 queued.request.overlap = overlap;
1383 }
1384 let admit_now = Instant::now();
1385 let class_index = popped.class_index();
1386 let class = self.profile.class(class_index);
1387 let queued = popped.into_payload();
1388 let admission = queued.admission;
1389 let request = queued.request;
1390 tracing::debug!(
1391 policy_class = class.name,
1392 "scheduling request from pending queue"
1393 );
1394 let _ = self.admit_one(request, admit_now, admission);
1395 }
1396 }
1397
1398 fn admit_one(
1401 &mut self,
1402 mut request: SchedulingRequest,
1403 decay_now: Instant,
1404 admission: Option<RequestAdmission>,
1405 ) -> (bool, bool) {
1406 let admission_key = admission.as_ref().map(|_| {
1407 request
1408 .mode
1409 .tracked_request_id()
1410 .expect("admitted request is tracked")
1411 .to_owned()
1412 });
1413 request.worker_loads = self
1414 .slots
1415 .project_worker_loads(request.token_seq.as_deref(), decay_now);
1416
1417 let selection = {
1418 let workers = self.workers_with_configs.borrow();
1419 let overloaded_worker_ids = self
1420 .overloaded_worker_provider
1421 .as_ref()
1422 .and_then(|provider| provider());
1423 let eligibility = request.eligibility_with_overloaded(overloaded_worker_ids.as_ref());
1424 self.selector
1425 .select_worker(&workers, &request, eligibility, self.block_size)
1426 .map(|selection| {
1427 let config = workers
1428 .get(&selection.worker.worker_id)
1429 .expect("selected worker config must exist");
1430 let selected_worker_tiers = request
1431 .overlap
1432 .selected_worker_tiers(selection.worker, config);
1433 (selection, selected_worker_tiers)
1434 })
1435 };
1436
1437 let (selection, selected_worker_tiers) = match selection {
1438 Ok(s) => s,
1439 Err(e) => {
1440 tracing::warn!("scheduling failed: {e}");
1441 request.respond(Err(e));
1442 if let Some(request_id) = admission_key.as_deref() {
1443 self.tracked_admissions.remove(request_id);
1444 }
1445 return (
1446 admission
1447 .as_ref()
1448 .is_some_and(|admission| self.abort_admission(admission.ticket)),
1449 false,
1450 );
1451 }
1452 };
1453
1454 let (admission, request_progress) = match admission {
1455 Some(RequestAdmission { ticket, progress }) => (Some(ticket), Some(progress)),
1456 None => (None, None),
1457 };
1458
1459 let response = SchedulingResponse {
1460 best_worker: selection.worker,
1461 effective_overlap_blocks: selection.effective_overlap_blocks,
1462 cached_tokens: selection.cached_tokens,
1463 selected_worker_tiers,
1464 request_progress,
1465 lifecycle_lease: None,
1466 potential_decode_blocks: selection.potential_decode_blocks,
1467 };
1468
1469 if !request.mode.is_tracked() {
1470 request.respond(Ok(response));
1471 debug_assert!(
1472 admission.is_none(),
1473 "query-only selection bypasses admission"
1474 );
1475 return (false, false);
1476 }
1477
1478 let request_id = request
1479 .mode
1480 .tracked_request_id()
1481 .expect("tracked mode always has a request ID")
1482 .to_string();
1483
1484 let prefill_load_hint = self.prefill_load_hint_for(
1485 request.isl_tokens,
1486 selection.cached_tokens,
1487 request.track_prefill_tokens,
1488 );
1489
1490 let sequence_request = SequenceRequest {
1491 request_id,
1492 token_sequence: request.token_seq.take(),
1493 track_prefill_tokens: request.track_prefill_tokens,
1494 expected_output_tokens: request.expected_output_tokens,
1495 prefill_load_hint,
1496 worker: selection.worker,
1497 lora_name: request.lora_name.take(),
1498 };
1499 let delivered = self.book_and_respond(request, sequence_request, response);
1500 if let Some(ticket) = admission {
1501 if delivered {
1502 let request_id = admission_key.expect("admitted request has a lifecycle key");
1503 if let Some(tracked) = self.tracked_admissions.get_mut(&request_id) {
1504 debug_assert_eq!(tracked.ticket, ticket);
1505 tracked.queue_class_index = None;
1506 tracked.worker = Some(selection.worker);
1507 } else {
1508 self.tracked_admissions.insert(
1509 request_id,
1510 TrackedAdmission {
1511 ticket,
1512 queue_class_index: None,
1513 worker: Some(selection.worker),
1514 dispatched: false,
1515 },
1516 );
1517 }
1518 } else {
1519 if let Some(request_id) = admission_key.as_deref() {
1520 self.tracked_admissions.remove(request_id);
1521 }
1522 return (self.abort_admission(ticket), false);
1523 }
1524 }
1525 (false, delivered)
1526 }
1527
1528 fn book_and_respond(
1536 &self,
1537 mut request: SchedulingRequest,
1538 sequence_request: SequenceRequest,
1539 response: SchedulingResponse,
1540 ) -> bool {
1541 if request.response_is_closed() {
1542 tracing::debug!(
1543 request_id = %sequence_request.request_id,
1544 "Skipping scheduler booking for cancelled request"
1545 );
1546 return false;
1547 }
1548
1549 let request_id = sequence_request.request_id.clone();
1550 if let Err(error) = self.slots.add_request(sequence_request, Instant::now()) {
1551 tracing::warn!(%request_id, %error, "Failed to book scheduler state");
1552 request.respond(Err(KvSchedulerError::BookingFailed(error.to_string())));
1553 return false;
1554 }
1555
1556 if request.respond(Ok(response)) {
1557 return true;
1558 }
1559
1560 tracing::debug!(%request_id, "Rolling back undelivered scheduler booking");
1561 if let Err(error) = self.slots.free(&request_id, Instant::now()) {
1562 tracing::error!(%request_id, %error, "Failed to roll back scheduler booking");
1563 }
1564 false
1565 }
1566
1567 fn prefill_load_hint_for(
1568 &self,
1569 isl_tokens: usize,
1570 cached_tokens: usize,
1571 track_prefill_tokens: bool,
1572 ) -> Option<PrefillLoadHint> {
1573 if !track_prefill_tokens {
1574 return None;
1575 }
1576
1577 let effective_isl = effective_prefill_tokens(isl_tokens, cached_tokens);
1578 if effective_isl == 0 {
1579 return None;
1580 }
1581 let prefix = isl_tokens - effective_isl;
1582
1583 let expected_prefill_duration = match &self.prefill_load_estimator {
1584 Some(estimator) => match estimator.predict_prefill_duration(1, effective_isl, prefix) {
1585 Ok(expected_prefill_duration) => Some(expected_prefill_duration),
1586 Err(error) => {
1587 tracing::warn!(
1588 effective_isl,
1589 prefix,
1590 "failed to predict prefill duration for active load tracking: {error}"
1591 );
1592 None
1593 }
1594 },
1595 None => None,
1596 };
1597
1598 Some(PrefillLoadHint {
1599 initial_effective_prefill_tokens: effective_isl,
1600 expected_prefill_duration,
1601 })
1602 }
1603
1604 fn all_workers_prefill_busy(
1611 &self,
1612 class: &PolicyClassConfig,
1613 eligibility: RoutingEligibility<'_>,
1614 decay_now: Instant,
1615 ) -> bool {
1616 let active_tokens = self.slots.active_tokens(decay_now);
1617 let configs = self.workers_with_configs.borrow();
1618 Self::all_workers_prefill_busy_with(&active_tokens, &configs, class, eligibility)
1619 }
1620
1621 fn all_workers_prefill_busy_with(
1622 active_tokens: &HashMap<crate::protocols::WorkerWithDpRank, usize>,
1623 configs: &HashMap<WorkerId, C>,
1624 class: &PolicyClassConfig,
1625 eligibility: RoutingEligibility<'_>,
1626 ) -> bool {
1627 if let Some(worker) = eligibility.pinned_worker() {
1628 let Ok(config) = eligibility.validate_worker_rank(configs, worker) else {
1629 return false;
1630 };
1631
1632 let max_batched = config
1633 .max_num_batched_tokens()
1634 .unwrap_or(DEFAULT_MAX_BATCHED_TOKENS);
1635 let tokens = active_tokens.get(&worker).copied().unwrap_or(0);
1636 return class.worker_is_busy(tokens, max_batched);
1637 }
1638
1639 let mut checked_any = false;
1640 let has_available = eligibility.any_eligible_worker_rank(configs, |worker, config| {
1641 checked_any = true;
1642 let max_batched = config
1643 .max_num_batched_tokens()
1644 .unwrap_or(DEFAULT_MAX_BATCHED_TOKENS);
1645 let tokens = active_tokens.get(&worker).copied().unwrap_or(0);
1646 !class.worker_is_busy(tokens, max_batched)
1647 });
1648
1649 checked_any && !has_available
1650 }
1651
1652 fn add_class_counters(&self, class_index: usize, snapshot: QueueSnapshot) {
1653 let counters = &self.class_counters[class_index];
1654 counters.pending_count.fetch_add(1, AtomicOrdering::Relaxed);
1655 counters
1656 .pending_isl_tokens
1657 .fetch_add(snapshot.raw_isl_tokens, AtomicOrdering::Relaxed);
1658 counters
1659 .pending_cached_tokens
1660 .fetch_add(snapshot.cached_tokens, AtomicOrdering::Relaxed);
1661 }
1662
1663 fn subtract_class_counters(&self, class_index: usize, snapshot: QueueSnapshot) {
1664 let counters = &self.class_counters[class_index];
1665 counters.pending_count.fetch_sub(1, AtomicOrdering::Relaxed);
1666 counters
1667 .pending_isl_tokens
1668 .fetch_sub(snapshot.raw_isl_tokens, AtomicOrdering::Relaxed);
1669 counters
1670 .pending_cached_tokens
1671 .fetch_sub(snapshot.cached_tokens, AtomicOrdering::Relaxed);
1672 }
1673}
1674
1675fn apply_admission_placement(
1676 request: &mut SchedulingRequest,
1677 placement: WorkerPlacement,
1678) -> Result<(), KvSchedulerError> {
1679 let WorkerPlacement::Exact(worker) = placement else {
1680 return Ok(());
1681 };
1682 if request.pinned_worker.is_some_and(|pinned| pinned != worker) {
1683 return Err(KvSchedulerError::BookingFailed(format!(
1684 "admission placement {worker:?} conflicts with the request's pinned worker"
1685 )));
1686 }
1687 request.pinned_worker = Some(worker);
1688 request.eligibility().validate_pinned_worker_allowed()
1689}
1690
1691#[cfg(test)]
1692mod tests {
1693 use std::collections::{HashMap, HashSet};
1694 use std::sync::atomic::{AtomicUsize, Ordering};
1695 use std::sync::{Arc, Condvar, Mutex as StdMutex};
1696 use std::time::Duration;
1697
1698 use async_trait::async_trait;
1699 use rustc_hash::FxHashMap;
1700 use tokio::sync::{Barrier, watch};
1701
1702 use super::*;
1703 use crate::protocols::{
1704 ActiveLoad, ActiveSequenceEvent, WorkerSelectionResult, WorkerWithDpRank,
1705 };
1706 use crate::scheduling::OverlapSignals;
1707 use crate::scheduling::types::{KvSchedulerError, ScheduleMode};
1708 use crate::scheduling::{
1709 AdmissionEvent, AdmissionId, AdmissionRequest, PolicyClassAdmissionPolicy,
1710 RefreshedOverlap, RequestProgress, RouterPolicyConfig,
1711 };
1712 use crate::sequences::{ActiveSequencesMultiWorker, SequencePublisher};
1713 use crate::test_utils::{NoopSequencePublisher, SimpleWorkerConfig};
1714 use crate::{DefaultWorkerSelector, WorkerSelector};
1715
1716 fn decay_now() -> Instant {
1717 Instant::now()
1718 }
1719
1720 struct FixedPrefillLoadEstimator {
1721 duration: Duration,
1722 }
1723
1724 impl PrefillLoadEstimator for FixedPrefillLoadEstimator {
1725 fn predict_prefill_duration(
1726 &self,
1727 _batch_size: usize,
1728 _effective_isl: usize,
1729 _prefix: usize,
1730 ) -> anyhow::Result<Duration> {
1731 Ok(self.duration)
1732 }
1733 }
1734
1735 type SchedulingResponseReceiver =
1736 tokio::sync::oneshot::Receiver<Result<SchedulingResponse, KvSchedulerError>>;
1737
1738 struct DropResponseOnLoadPublisher {
1739 response_rx: Arc<StdMutex<Option<SchedulingResponseReceiver>>>,
1740 }
1741
1742 impl SequencePublisher for DropResponseOnLoadPublisher {
1743 fn enqueue_event(&self, _event: ActiveSequenceEvent) -> anyhow::Result<()> {
1744 Ok(())
1745 }
1746
1747 fn publish_load(&self, _load: ActiveLoad) {
1748 self.response_rx.lock().unwrap().take();
1749 }
1750
1751 fn observe_load(&self, _: &WorkerWithDpRank, _: &str, _: usize, _: usize) {}
1752 }
1753
1754 #[derive(Default)]
1755 struct SelectorRendezvous {
1756 arrivals: StdMutex<usize>,
1757 cv: Condvar,
1758 }
1759
1760 impl SelectorRendezvous {
1761 fn wait_for_peer(&self) {
1762 let mut arrivals = self.arrivals.lock().unwrap();
1763 *arrivals += 1;
1764
1765 if *arrivals == 1 {
1766 let _ = self
1767 .cv
1768 .wait_timeout(arrivals, Duration::from_millis(100))
1769 .unwrap();
1770 return;
1771 }
1772
1773 self.cv.notify_all();
1774 }
1775 }
1776
1777 #[derive(Clone)]
1778 struct MinDecodeSelector {
1779 rendezvous: Option<Arc<SelectorRendezvous>>,
1780 }
1781
1782 impl WorkerSelector<SimpleWorkerConfig> for MinDecodeSelector {
1783 fn select_worker(
1784 &self,
1785 workers: &HashMap<WorkerId, SimpleWorkerConfig>,
1786 request: &SchedulingRequest,
1787 eligibility: RoutingEligibility<'_>,
1788 block_size: u32,
1789 ) -> Result<WorkerSelectionResult, KvSchedulerError> {
1790 if let Some(rendezvous) = &self.rendezvous {
1791 rendezvous.wait_for_peer();
1792 }
1793
1794 let mut best_worker = None;
1795 eligibility.for_each_eligible_worker_rank(workers, |worker, _| {
1796 let load = request.worker_load_for(worker);
1797 let potential_prefill_tokens = if request.track_prefill_tokens {
1798 load.active_prefill_tokens
1799 .saturating_add(effective_prefill_tokens(
1800 request.isl_tokens,
1801 request.effective_cached_tokens_for(worker),
1802 ))
1803 } else {
1804 0
1805 };
1806 let potential_decode_blocks = load.potential_decode_blocks();
1807 let key = (
1808 potential_prefill_tokens,
1809 potential_decode_blocks,
1810 worker.worker_id,
1811 worker.dp_rank,
1812 );
1813 if best_worker.is_none_or(|(_, best_key)| key < best_key) {
1814 best_worker = Some((worker, key));
1815 }
1816 });
1817
1818 let Some((worker, _)) = best_worker else {
1819 return Err(KvSchedulerError::NoEndpoints);
1820 };
1821
1822 Ok(WorkerSelectionResult {
1823 worker,
1824 required_blocks: request.request_blocks(block_size),
1825 effective_overlap_blocks: request.effective_overlap_blocks_for(worker),
1826 cached_tokens: request.effective_cached_tokens_for(worker),
1827 potential_decode_blocks: request
1828 .potential_decode_blocks_after_admission(worker, block_size),
1829 })
1830 }
1831 }
1832
1833 fn make_queue(
1834 num_workers: usize,
1835 block_size: u32,
1836 isl: usize,
1837 threshold_frac: Option<f64>,
1838 ) -> (
1839 Arc<SchedulerQueue<NoopSequencePublisher, SimpleWorkerConfig>>,
1840 Arc<ActiveSequencesMultiWorker<NoopSequencePublisher>>,
1841 ) {
1842 let (queue, slots, _tx) =
1843 make_queue_with_sender(num_workers, block_size, isl, threshold_frac, None);
1844 (queue, slots)
1845 }
1846
1847 #[allow(clippy::type_complexity)]
1848 fn make_queue_with_custom_selector<Sel: WorkerSelector<SimpleWorkerConfig> + Send + 'static>(
1849 num_workers: usize,
1850 block_size: u32,
1851 isl: usize,
1852 threshold_frac: Option<f64>,
1853 selector: Sel,
1854 ) -> (
1855 Arc<SchedulerQueue<NoopSequencePublisher, SimpleWorkerConfig, Sel>>,
1856 Arc<ActiveSequencesMultiWorker<NoopSequencePublisher>>,
1857 ) {
1858 let dp_range: HashMap<u64, (u32, u32)> =
1859 (0..num_workers as u64).map(|id| (id, (0, 1))).collect();
1860 let slots = Arc::new(ActiveSequencesMultiWorker::new(
1861 NoopSequencePublisher,
1862 block_size as usize,
1863 dp_range,
1864 false,
1865 0,
1866 "test",
1867 ));
1868
1869 let mut configs: HashMap<u64, SimpleWorkerConfig> = HashMap::new();
1870 for id in 0..num_workers as u64 {
1871 configs.insert(
1872 id,
1873 SimpleWorkerConfig {
1874 max_num_batched_tokens: Some(isl as u64),
1875 ..Default::default()
1876 },
1877 );
1878 }
1879 let (_cfg_tx, cfg_rx) = watch::channel(configs);
1880
1881 let queue = Arc::new(SchedulerQueue::new(
1882 Arc::clone(&slots),
1883 cfg_rx,
1884 threshold_frac,
1885 block_size,
1886 selector,
1887 RouterQueuePolicy::Fcfs,
1888 None,
1889 ));
1890
1891 (queue, slots)
1892 }
1893
1894 #[allow(clippy::type_complexity)]
1895 fn make_queue_with_sender(
1896 num_workers: usize,
1897 block_size: u32,
1898 isl: usize,
1899 threshold_frac: Option<f64>,
1900 prefill_load_estimator: Option<Arc<dyn PrefillLoadEstimator>>,
1901 ) -> (
1902 Arc<SchedulerQueue<NoopSequencePublisher, SimpleWorkerConfig>>,
1903 Arc<ActiveSequencesMultiWorker<NoopSequencePublisher>>,
1904 watch::Sender<HashMap<u64, SimpleWorkerConfig>>,
1905 ) {
1906 let dp_range: HashMap<u64, (u32, u32)> =
1907 (0..num_workers as u64).map(|id| (id, (0, 1))).collect();
1908 let slots = Arc::new(ActiveSequencesMultiWorker::new(
1909 NoopSequencePublisher,
1910 block_size as usize,
1911 dp_range,
1912 false,
1913 0,
1914 "test",
1915 ));
1916
1917 let mut configs: HashMap<u64, SimpleWorkerConfig> = HashMap::new();
1918 for id in 0..num_workers as u64 {
1919 configs.insert(
1920 id,
1921 SimpleWorkerConfig {
1922 max_num_batched_tokens: Some(isl as u64),
1923 ..Default::default()
1924 },
1925 );
1926 }
1927 let (cfg_tx, cfg_rx) = watch::channel(configs);
1928
1929 let selector = DefaultWorkerSelector::new(None, "test");
1930 let queue = Arc::new(SchedulerQueue::new(
1931 Arc::clone(&slots),
1932 cfg_rx,
1933 threshold_frac,
1934 block_size,
1935 selector,
1936 RouterQueuePolicy::Fcfs,
1937 prefill_load_estimator,
1938 ));
1939
1940 (queue, slots, cfg_tx)
1941 }
1942
1943 fn policy_profile(yaml: &str) -> PolicyProfile {
1944 RouterPolicyConfig::from_yaml(yaml)
1945 .unwrap()
1946 .resolve_profile(None, None, crate::config::RouterQueuePolicy::Fcfs)
1947 }
1948
1949 #[allow(clippy::type_complexity)]
1950 fn make_queue_with_profile(
1951 num_workers: usize,
1952 block_size: u32,
1953 max_num_batched_tokens: usize,
1954 profile: PolicyProfile,
1955 ) -> (
1956 Arc<SchedulerQueue<NoopSequencePublisher, SimpleWorkerConfig>>,
1957 Arc<ActiveSequencesMultiWorker<NoopSequencePublisher>>,
1958 ) {
1959 let (queue, slots, _cfg_tx) = make_queue_with_profile_and_sender(
1960 num_workers,
1961 block_size,
1962 max_num_batched_tokens,
1963 profile,
1964 );
1965 (queue, slots)
1966 }
1967
1968 #[allow(clippy::type_complexity)]
1969 fn make_queue_with_profile_and_sender(
1970 num_workers: usize,
1971 block_size: u32,
1972 max_num_batched_tokens: usize,
1973 profile: PolicyProfile,
1974 ) -> (
1975 Arc<SchedulerQueue<NoopSequencePublisher, SimpleWorkerConfig>>,
1976 Arc<ActiveSequencesMultiWorker<NoopSequencePublisher>>,
1977 watch::Sender<HashMap<u64, SimpleWorkerConfig>>,
1978 ) {
1979 let dp_range: HashMap<u64, (u32, u32)> =
1980 (0..num_workers as u64).map(|id| (id, (0, 1))).collect();
1981 let slots = Arc::new(ActiveSequencesMultiWorker::new(
1982 NoopSequencePublisher,
1983 block_size as usize,
1984 dp_range,
1985 false,
1986 0,
1987 "test",
1988 ));
1989 let configs = (0..num_workers as u64)
1990 .map(|id| {
1991 (
1992 id,
1993 SimpleWorkerConfig {
1994 max_num_batched_tokens: Some(max_num_batched_tokens as u64),
1995 ..Default::default()
1996 },
1997 )
1998 })
1999 .collect();
2000 let (cfg_tx, cfg_rx) = watch::channel(configs);
2001 let queue = Arc::new(
2002 SchedulerQueue::new_with_policy_profile(
2003 Arc::clone(&slots),
2004 cfg_rx,
2005 profile,
2006 block_size,
2007 DefaultWorkerSelector::new(None, "test"),
2008 None,
2009 None,
2010 None,
2011 Duration::from_secs(60),
2012 PolicyClassAdmissionPolicies::new(),
2013 )
2014 .unwrap(),
2015 );
2016 (queue, slots, cfg_tx)
2017 }
2018
2019 fn make_queue_with_overload_provider(
2020 num_workers: usize,
2021 block_size: u32,
2022 isl: usize,
2023 overloaded_worker_provider: OverloadedWorkerProvider,
2024 ) -> (
2025 Arc<SchedulerQueue<NoopSequencePublisher, SimpleWorkerConfig>>,
2026 Arc<ActiveSequencesMultiWorker<NoopSequencePublisher>>,
2027 ) {
2028 let dp_range: HashMap<u64, (u32, u32)> =
2029 (0..num_workers as u64).map(|id| (id, (0, 1))).collect();
2030 let slots = Arc::new(ActiveSequencesMultiWorker::new(
2031 NoopSequencePublisher,
2032 block_size as usize,
2033 dp_range,
2034 false,
2035 0,
2036 "test",
2037 ));
2038
2039 let mut configs: HashMap<u64, SimpleWorkerConfig> = HashMap::new();
2040 for id in 0..num_workers as u64 {
2041 configs.insert(
2042 id,
2043 SimpleWorkerConfig {
2044 max_num_batched_tokens: Some(isl as u64),
2045 ..Default::default()
2046 },
2047 );
2048 }
2049 let (_cfg_tx, cfg_rx) = watch::channel(configs);
2050
2051 let selector = DefaultWorkerSelector::new(None, "test");
2052 let queue = Arc::new(SchedulerQueue::new_with_overload_provider(
2053 Arc::clone(&slots),
2054 cfg_rx,
2055 None,
2056 block_size,
2057 selector,
2058 RouterQueuePolicy::Fcfs,
2059 None,
2060 Some(overloaded_worker_provider),
2061 ));
2062
2063 (queue, slots)
2064 }
2065
2066 struct CountingRefresher {
2067 calls: AtomicUsize,
2068 response: RefreshedOverlap,
2069 }
2070
2071 #[async_trait]
2072 impl OverlapScoresRefresh for CountingRefresher {
2073 async fn refresh(&self, _block_hashes: &[LocalBlockHash]) -> Option<RefreshedOverlap> {
2074 self.calls.fetch_add(1, Ordering::Relaxed);
2075 Some(self.response.clone())
2076 }
2077 }
2078
2079 struct BlockingRefresher {
2080 calls: AtomicUsize,
2081 started: tokio::sync::Notify,
2082 release: tokio::sync::Notify,
2083 response: RefreshedOverlap,
2084 }
2085
2086 impl BlockingRefresher {
2087 fn new(response: RefreshedOverlap) -> Self {
2088 Self {
2089 calls: AtomicUsize::new(0),
2090 started: tokio::sync::Notify::new(),
2091 release: tokio::sync::Notify::new(),
2092 response,
2093 }
2094 }
2095
2096 async fn wait_for_calls(&self, target: usize) {
2097 while self.calls.load(Ordering::Relaxed) < target {
2098 self.started.notified().await;
2099 }
2100 }
2101
2102 fn release_one(&self) {
2103 self.release.notify_one();
2104 }
2105 }
2106
2107 #[async_trait]
2108 impl OverlapScoresRefresh for BlockingRefresher {
2109 async fn refresh(&self, _block_hashes: &[LocalBlockHash]) -> Option<RefreshedOverlap> {
2110 self.calls.fetch_add(1, Ordering::Relaxed);
2111 self.started.notify_one();
2112 self.release.notified().await;
2113 Some(self.response.clone())
2114 }
2115 }
2116
2117 #[allow(clippy::type_complexity)]
2118 fn make_queue_with_refresher(
2119 num_workers: usize,
2120 block_size: u32,
2121 isl: usize,
2122 threshold_frac: Option<f64>,
2123 refresher: Arc<CountingRefresher>,
2124 ) -> (
2125 Arc<
2126 SchedulerQueue<
2127 NoopSequencePublisher,
2128 SimpleWorkerConfig,
2129 DefaultWorkerSelector,
2130 CountingRefresher,
2131 >,
2132 >,
2133 Arc<ActiveSequencesMultiWorker<NoopSequencePublisher>>,
2134 ) {
2135 let dp_range: HashMap<u64, (u32, u32)> =
2136 (0..num_workers as u64).map(|id| (id, (0, 1))).collect();
2137 let slots = Arc::new(ActiveSequencesMultiWorker::new(
2138 NoopSequencePublisher,
2139 block_size as usize,
2140 dp_range,
2141 false,
2142 0,
2143 "test",
2144 ));
2145
2146 let mut configs: HashMap<u64, SimpleWorkerConfig> = HashMap::new();
2147 for id in 0..num_workers as u64 {
2148 configs.insert(
2149 id,
2150 SimpleWorkerConfig {
2151 max_num_batched_tokens: Some(isl as u64),
2152 ..Default::default()
2153 },
2154 );
2155 }
2156 let (_cfg_tx, cfg_rx) = watch::channel(configs);
2157
2158 let queue = Arc::new(SchedulerQueue::new_with_overlap_refresh(
2159 Arc::clone(&slots),
2160 cfg_rx,
2161 threshold_frac,
2162 block_size,
2163 DefaultWorkerSelector::new(None, "test"),
2164 RouterQueuePolicy::Fcfs,
2165 None,
2166 Some(refresher),
2167 None,
2168 ));
2169
2170 (queue, slots)
2171 }
2172
2173 #[allow(clippy::type_complexity)]
2174 fn make_queue_with_blocking_refresher(
2175 num_workers: usize,
2176 block_size: u32,
2177 isl: usize,
2178 threshold_frac: Option<f64>,
2179 refresher: Arc<BlockingRefresher>,
2180 admission_channel_capacity: usize,
2181 ) -> (
2182 Arc<
2183 SchedulerQueue<
2184 NoopSequencePublisher,
2185 SimpleWorkerConfig,
2186 DefaultWorkerSelector,
2187 BlockingRefresher,
2188 >,
2189 >,
2190 Arc<ActiveSequencesMultiWorker<NoopSequencePublisher>>,
2191 ) {
2192 let dp_range: HashMap<u64, (u32, u32)> =
2193 (0..num_workers as u64).map(|id| (id, (0, 1))).collect();
2194 let slots = Arc::new(ActiveSequencesMultiWorker::new(
2195 NoopSequencePublisher,
2196 block_size as usize,
2197 dp_range,
2198 false,
2199 0,
2200 "test",
2201 ));
2202
2203 let mut configs: HashMap<u64, SimpleWorkerConfig> = HashMap::new();
2204 for id in 0..num_workers as u64 {
2205 configs.insert(
2206 id,
2207 SimpleWorkerConfig {
2208 max_num_batched_tokens: Some(isl as u64),
2209 ..Default::default()
2210 },
2211 );
2212 }
2213 let (_cfg_tx, cfg_rx) = watch::channel(configs);
2214
2215 let queue = Arc::new(
2216 SchedulerQueue::new_with_policy_profile_and_capacity(
2217 Arc::clone(&slots),
2218 cfg_rx,
2219 PolicyProfile::synthetic(threshold_frac, crate::config::RouterQueuePolicy::Fcfs),
2220 block_size,
2221 DefaultWorkerSelector::new(None, "test"),
2222 None,
2223 Some(refresher),
2224 None,
2225 Duration::from_secs(60),
2226 PolicyClassAdmissionPolicies::new(),
2227 admission_channel_capacity,
2228 )
2229 .unwrap(),
2230 );
2231
2232 (queue, slots)
2233 }
2234
2235 fn make_request(
2236 request_id: &str,
2237 isl_tokens: usize,
2238 ) -> (
2239 SchedulingRequest,
2240 tokio::sync::oneshot::Receiver<
2241 Result<SchedulingResponse, crate::scheduling::types::KvSchedulerError>,
2242 >,
2243 ) {
2244 let (tx, rx) = tokio::sync::oneshot::channel();
2245 let req = SchedulingRequest {
2246 mode: ScheduleMode::Tracked {
2247 request_id: request_id.to_string(),
2248 },
2249 token_seq: None,
2250 isl_tokens,
2251 overlap: OverlapSignals::default(),
2252 worker_loads: FxHashMap::default(),
2253 track_prefill_tokens: true,
2254 router_config_override: None,
2255 lora_name: None,
2256 priority_jump: 0.0,
2257 strict_priority: 0,
2258 policy_class: None,
2259 session_id: None,
2260 expected_output_tokens: None,
2261 pinned_worker: None,
2262 allowed_worker_ids: None,
2263 routing_constraints: crate::protocols::RoutingConstraints::default(),
2264 shared_cache_hits: None,
2265 resp_tx: Some(tx),
2266 };
2267 (req, rx)
2268 }
2269
2270 fn make_admission_request(
2271 request_id: &str,
2272 isl_tokens: usize,
2273 ) -> (
2274 SchedulingRequest,
2275 tokio::sync::oneshot::Receiver<
2276 Result<SchedulingResponse, crate::scheduling::types::KvSchedulerError>,
2277 >,
2278 ) {
2279 let (mut request, response) = make_request(request_id, isl_tokens);
2280 request.mode = ScheduleMode::TrackedWithLifecycle {
2281 request_id: request_id.to_owned(),
2282 };
2283 (request, response)
2284 }
2285
2286 async fn enqueue_with_lease(
2287 queue: &SchedulerQueue<NoopSequencePublisher, SimpleWorkerConfig>,
2288 request: SchedulingRequest,
2289 ) -> Box<RequestLifecycleLease> {
2290 let request_id = request
2291 .mode
2292 .lifecycle_request_id()
2293 .expect("admission test request must be tracked");
2294 let lease = queue.new_request_lifecycle_lease(Some(request_id)).unwrap();
2295 queue
2296 .enqueue_with_block_hashes_and_lease(request, None, Some(lease))
2297 .await
2298 .expect("actor must return the accepted admission lease")
2299 }
2300
2301 #[derive(Default)]
2302 struct GateState {
2303 deferred: Option<AdmissionId>,
2304 session_id: Option<String>,
2305 context_tokens: usize,
2306 progress: Option<RequestProgress>,
2307 dispatched: Vec<WorkerWithDpRank>,
2308 completed_context_tokens: Vec<usize>,
2309 aborted: Vec<AdmissionId>,
2310 }
2311
2312 struct ReconcileGate {
2313 state: Arc<StdMutex<GateState>>,
2314 }
2315
2316 impl PolicyClassAdmissionPolicy for ReconcileGate {
2317 fn admit(&mut self, request: AdmissionRequest<'_>) -> AdmissionDecision {
2318 let mut state = self.state.lock().unwrap();
2319 state.deferred = Some(request.id());
2320 state.session_id = request.session_id().map(str::to_owned);
2321 state.context_tokens = request.context_tokens();
2322 state.progress = Some(request.progress().clone());
2323 AdmissionDecision::Defer
2324 }
2325
2326 fn on_event(&mut self, event: AdmissionEvent) -> Vec<AdmissionAction> {
2327 let mut state = self.state.lock().unwrap();
2328 match event {
2329 AdmissionEvent::Dispatched { worker, .. } => {
2330 state.dispatched.push(worker);
2331 Vec::new()
2332 }
2333 AdmissionEvent::Completed { id, context_tokens } => {
2334 if state.deferred == Some(id) {
2335 state.deferred = None;
2336 }
2337 state.completed_context_tokens.push(context_tokens);
2338 Vec::new()
2339 }
2340 AdmissionEvent::Aborted { id } => {
2341 if state.deferred == Some(id) {
2342 state.deferred = None;
2343 }
2344 state.aborted.push(id);
2345 Vec::new()
2346 }
2347 AdmissionEvent::Reconcile => state
2348 .deferred
2349 .take()
2350 .map(|id| {
2351 vec![AdmissionAction::MakeReady {
2352 id,
2353 placement: WorkerPlacement::Exact(WorkerWithDpRank::new(0, 0)),
2354 }]
2355 })
2356 .unwrap_or_default(),
2357 }
2358 }
2359 }
2360
2361 fn make_queue_with_admission_policy(
2362 policy: Box<dyn PolicyClassAdmissionPolicy>,
2363 ) -> (
2364 Arc<SchedulerQueue<NoopSequencePublisher, SimpleWorkerConfig>>,
2365 Arc<ActiveSequencesMultiWorker<NoopSequencePublisher>>,
2366 ) {
2367 make_queue_with_admission_policy_and_workers(policy, 1)
2368 }
2369
2370 fn make_queue_with_admission_policy_and_workers(
2371 policy: Box<dyn PolicyClassAdmissionPolicy>,
2372 worker_count: u64,
2373 ) -> (
2374 Arc<SchedulerQueue<NoopSequencePublisher, SimpleWorkerConfig>>,
2375 Arc<ActiveSequencesMultiWorker<NoopSequencePublisher>>,
2376 ) {
2377 let profile = policy_profile(
2378 r#"
2379default_policy_family: standard
2380uncached_isl_buckets:
2381 - min_tokens: 0
2382 bucket: all
2383policy_classes:
2384 - name: standard
2385 policy_family: standard
2386 cache_bucket: all
2387 quantum: 1
2388 - name: agents
2389 prefill_busy_threshold: 0
2390 quantum: 1
2391"#,
2392 );
2393 make_queue_with_profile_and_admission_policy(profile, "agents", policy, worker_count)
2394 }
2395
2396 fn make_queue_with_profile_and_admission_policy(
2397 profile: PolicyProfile,
2398 class_name: &str,
2399 policy: Box<dyn PolicyClassAdmissionPolicy>,
2400 worker_count: u64,
2401 ) -> (
2402 Arc<SchedulerQueue<NoopSequencePublisher, SimpleWorkerConfig>>,
2403 Arc<ActiveSequencesMultiWorker<NoopSequencePublisher>>,
2404 ) {
2405 let slots = Arc::new(ActiveSequencesMultiWorker::new(
2406 NoopSequencePublisher,
2407 16,
2408 (0..worker_count).map(|worker| (worker, (0, 1))).collect(),
2409 false,
2410 0,
2411 "test",
2412 ));
2413 let (_cfg_tx, cfg_rx) = watch::channel(
2414 (0..worker_count)
2415 .map(|worker| {
2416 (
2417 worker,
2418 SimpleWorkerConfig {
2419 max_num_batched_tokens: Some(1_000),
2420 ..Default::default()
2421 },
2422 )
2423 })
2424 .collect(),
2425 );
2426 let mut policies = PolicyClassAdmissionPolicies::new();
2427 policies.insert(class_name.to_owned(), policy);
2428 let queue = Arc::new(
2429 SchedulerQueue::new_with_policy_profile(
2430 Arc::clone(&slots),
2431 cfg_rx,
2432 profile,
2433 16,
2434 DefaultWorkerSelector::new(None, "test"),
2435 None,
2436 None,
2437 None,
2438 Duration::from_secs(60),
2439 policies,
2440 )
2441 .unwrap(),
2442 );
2443 (queue, slots)
2444 }
2445
2446 #[tokio::test]
2447 async fn admission_policy_defers_releases_and_observes_lifecycle() {
2448 let state = Arc::new(StdMutex::new(GateState::default()));
2449 let (queue, slots) = make_queue_with_admission_policy(Box::new(ReconcileGate {
2450 state: Arc::clone(&state),
2451 }));
2452 let (mut request, response) = make_admission_request("deferred", 64);
2453 request.policy_class = Some("agents".to_owned());
2454 request.session_id = Some("session-a".to_owned());
2455
2456 let mut lease = enqueue_with_lease(&queue, request).await;
2457 assert_eq!(queue.pending_count(), 1);
2458 {
2459 let state = state.lock().unwrap();
2460 assert_eq!(state.session_id.as_deref(), Some("session-a"));
2461 assert_eq!(state.context_tokens, 64);
2462 assert!(state.dispatched.is_empty());
2463 }
2464
2465 queue.reconcile().await;
2466 let selected = response.await.unwrap().unwrap();
2467 assert_eq!(selected.best_worker, WorkerWithDpRank::new(0, 0));
2468 let progress = selected.request_progress.unwrap();
2469 progress.update_context_tokens(80);
2470 assert_eq!(
2471 state
2472 .lock()
2473 .unwrap()
2474 .progress
2475 .as_ref()
2476 .unwrap()
2477 .context_tokens(),
2478 80
2479 );
2480 assert_eq!(queue.pending_count(), 0);
2481 lease.mark_dispatched().await;
2482 lease.mark_dispatched().await;
2483 lease.mark_completed(81);
2484 drop(lease);
2485 queue.update().await;
2486 {
2487 let state = state.lock().unwrap();
2488 assert_eq!(state.dispatched, vec![WorkerWithDpRank::new(0, 0)]);
2489 assert_eq!(state.completed_context_tokens, vec![81]);
2490 }
2491 slots.assert_completely_drained(decay_now());
2492 }
2493
2494 #[tokio::test]
2495 async fn cancelled_deferred_request_receives_one_terminal_event() {
2496 let state = Arc::new(StdMutex::new(GateState::default()));
2497 let (queue, _slots) = make_queue_with_admission_policy(Box::new(ReconcileGate {
2498 state: Arc::clone(&state),
2499 }));
2500 let (mut request, response) = make_admission_request("cancelled", 64);
2501 request.policy_class = Some("agents".to_owned());
2502 request.session_id = Some("session-a".to_owned());
2503
2504 let cancellation = enqueue_with_lease(&queue, request).await;
2505 drop(response);
2506 drop(cancellation);
2507 queue.update().await;
2508
2509 assert_eq!(queue.pending_count(), 0);
2510 let state = state.lock().unwrap();
2511 assert!(state.dispatched.is_empty());
2512 assert_eq!(state.aborted, vec![AdmissionId::new(0)]);
2513 }
2514
2515 struct ReadyGate {
2516 state: Arc<StdMutex<GateState>>,
2517 }
2518
2519 impl PolicyClassAdmissionPolicy for ReadyGate {
2520 fn admit(&mut self, _request: AdmissionRequest<'_>) -> AdmissionDecision {
2521 AdmissionDecision::Ready(WorkerPlacement::Any)
2522 }
2523
2524 fn on_event(&mut self, event: AdmissionEvent) -> Vec<AdmissionAction> {
2525 let mut state = self.state.lock().unwrap();
2526 match event {
2527 AdmissionEvent::Dispatched { worker, .. } => state.dispatched.push(worker),
2528 AdmissionEvent::Completed { context_tokens, .. } => {
2529 state.completed_context_tokens.push(context_tokens);
2530 }
2531 AdmissionEvent::Aborted { id } => state.aborted.push(id),
2532 AdmissionEvent::Reconcile => {}
2533 }
2534 Vec::new()
2535 }
2536 }
2537
2538 struct ExactReadyGate(WorkerWithDpRank);
2539
2540 impl PolicyClassAdmissionPolicy for ExactReadyGate {
2541 fn admit(&mut self, _request: AdmissionRequest<'_>) -> AdmissionDecision {
2542 AdmissionDecision::Ready(WorkerPlacement::Exact(self.0))
2543 }
2544 }
2545
2546 struct OrderedLifecycleGate(Arc<StdMutex<Vec<&'static str>>>);
2547
2548 impl PolicyClassAdmissionPolicy for OrderedLifecycleGate {
2549 fn admit(&mut self, _request: AdmissionRequest<'_>) -> AdmissionDecision {
2550 AdmissionDecision::Ready(WorkerPlacement::Any)
2551 }
2552
2553 fn on_event(&mut self, event: AdmissionEvent) -> Vec<AdmissionAction> {
2554 let event = match event {
2555 AdmissionEvent::Dispatched { .. } => "dispatched",
2556 AdmissionEvent::Completed { .. } => "completed",
2557 AdmissionEvent::Aborted { .. } => "aborted",
2558 AdmissionEvent::Reconcile => return Vec::new(),
2559 };
2560 self.0.lock().unwrap().push(event);
2561 Vec::new()
2562 }
2563 }
2564
2565 #[tokio::test]
2566 async fn lease_cleanup_preserves_dispatch_before_terminal_event() {
2567 let events = Arc::new(StdMutex::new(Vec::new()));
2568 let (queue, slots) =
2569 make_queue_with_admission_policy(Box::new(OrderedLifecycleGate(Arc::clone(&events))));
2570 let (mut request, response) = make_admission_request("ordered-cleanup", 64);
2571 request.policy_class = Some("agents".to_owned());
2572 let mut lease = enqueue_with_lease(&queue, request).await;
2573 response.await.unwrap().unwrap();
2574
2575 let (full_tx, _full_rx) = mpsc::channel(1);
2577 full_tx.send(AdmissionCommand::Cleanup).await.unwrap();
2578 lease.actor_tx = full_tx;
2579 {
2580 let mut report = std::pin::pin!(lease.mark_dispatched());
2581 let result = std::future::poll_fn(|cx| {
2582 std::task::Poll::Ready(std::future::Future::poll(report.as_mut(), cx))
2583 })
2584 .await;
2585 assert!(result.is_pending());
2586 }
2587 lease.mark_completed(96);
2588 drop(lease);
2589 queue.update().await;
2590
2591 assert_eq!(*events.lock().unwrap(), ["dispatched", "completed"]);
2592 slots.assert_completely_drained(decay_now());
2593 }
2594
2595 #[tokio::test]
2596 async fn cancellation_after_response_send_rolls_back_booking() {
2597 let state = Arc::new(StdMutex::new(GateState::default()));
2598 let (queue, slots) = make_queue_with_admission_policy(Box::new(ReadyGate {
2599 state: Arc::clone(&state),
2600 }));
2601 let (mut request, response) = make_admission_request("cancelled-after-handoff", 64);
2602 request.policy_class = Some("agents".to_owned());
2603 let cancellation = enqueue_with_lease(&queue, request).await;
2604 assert_eq!(
2605 slots
2606 .active_request_counts()
2607 .get(&WorkerWithDpRank::new(0, 0))
2608 .copied(),
2609 Some(1)
2610 );
2611
2612 drop(cancellation);
2613 tokio::time::timeout(Duration::from_secs(1), async {
2614 while state.lock().unwrap().aborted.len() != 1 {
2615 tokio::task::yield_now().await;
2616 }
2617 })
2618 .await
2619 .expect("cancellation did not abort admission");
2620
2621 slots.assert_completely_drained(decay_now());
2622 assert_eq!(state.lock().unwrap().aborted, vec![AdmissionId::new(0)]);
2623 drop(response);
2624 }
2625
2626 #[tokio::test]
2627 async fn dropped_completed_lease_commits_authoritative_context() {
2628 let state = Arc::new(StdMutex::new(GateState::default()));
2629 let (queue, slots) = make_queue_with_admission_policy(Box::new(ReadyGate {
2630 state: Arc::clone(&state),
2631 }));
2632 let (mut request, response) = make_admission_request("completed-after-terminal", 64);
2633 request.policy_class = Some("agents".to_owned());
2634 let mut lease = enqueue_with_lease(&queue, request).await;
2635 response.await.unwrap().unwrap();
2636 drop(queue);
2637 lease.mark_completed(96);
2638 drop(lease);
2639
2640 tokio::time::timeout(Duration::from_secs(1), async {
2641 while state.lock().unwrap().completed_context_tokens != [96] {
2642 tokio::task::yield_now().await;
2643 }
2644 })
2645 .await
2646 .expect("completed lease did not commit admission context");
2647 slots.assert_completely_drained(decay_now());
2648 assert!(state.lock().unwrap().aborted.is_empty());
2649 }
2650
2651 #[test]
2652 fn lease_drop_coalesces_actor_wakes_and_preserves_cleanup() {
2653 let cleanup = Arc::new(AdmissionCleanup::default());
2654 let (actor_tx, mut actor_rx) = mpsc::channel(1);
2655 let lease = |id: u64, request_id: &str| RequestLifecycleLease {
2656 cleanup: Arc::clone(&cleanup),
2657 actor_tx: actor_tx.clone(),
2658 ticket: Some(AdmissionTicket {
2659 class_index: 0,
2660 id: AdmissionId::new(id),
2661 }),
2662 request_id: Some(request_id.to_owned()),
2663 context_tokens: None,
2664 dispatched: false,
2665 };
2666
2667 drop(lease(1, "fast-path"));
2668 assert!(matches!(actor_rx.try_recv(), Ok(AdmissionCommand::Cleanup)));
2669 assert_eq!(
2670 cleanup.drain(),
2671 [AdmissionCleanupEntry {
2672 ticket: Some(AdmissionTicket {
2673 class_index: 0,
2674 id: AdmissionId::new(1),
2675 }),
2676 request_id: "fast-path".to_owned(),
2677 context_tokens: None,
2678 dispatched: false,
2679 }]
2680 );
2681
2682 drop(lease(2, "first"));
2683 drop(lease(3, "coalesced"));
2684 assert!(matches!(actor_rx.try_recv(), Ok(AdmissionCommand::Cleanup)));
2685 assert!(actor_rx.try_recv().is_err());
2686 assert_eq!(
2687 cleanup.drain(),
2688 [
2689 AdmissionCleanupEntry {
2690 ticket: Some(AdmissionTicket {
2691 class_index: 0,
2692 id: AdmissionId::new(2),
2693 }),
2694 request_id: "first".to_owned(),
2695 context_tokens: None,
2696 dispatched: false,
2697 },
2698 AdmissionCleanupEntry {
2699 ticket: Some(AdmissionTicket {
2700 class_index: 0,
2701 id: AdmissionId::new(3),
2702 }),
2703 request_id: "coalesced".to_owned(),
2704 context_tokens: None,
2705 dispatched: false,
2706 },
2707 ]
2708 );
2709 }
2710
2711 #[tokio::test]
2712 async fn duplicate_request_id_does_not_replace_active_admission() {
2713 let state = Arc::new(StdMutex::new(GateState::default()));
2714 let (queue, slots) = make_queue_with_admission_policy(Box::new(ReadyGate {
2715 state: Arc::clone(&state),
2716 }));
2717 let (mut first, first_response) = make_admission_request("duplicate", 64);
2718 first.policy_class = Some("agents".to_owned());
2719 let first_lease = enqueue_with_lease(&queue, first).await;
2720 first_response.await.unwrap().unwrap();
2721
2722 let (mut duplicate, duplicate_response) = make_admission_request("duplicate", 64);
2723 duplicate.policy_class = Some("agents".to_owned());
2724 drop(enqueue_with_lease(&queue, duplicate).await);
2725 assert!(matches!(
2726 duplicate_response.await.unwrap(),
2727 Err(KvSchedulerError::BookingFailed(message))
2728 if message == "request duplicate already has an active admission"
2729 ));
2730 assert_eq!(
2731 slots
2732 .active_request_counts()
2733 .get(&WorkerWithDpRank::new(0, 0))
2734 .copied(),
2735 Some(1)
2736 );
2737 assert!(state.lock().unwrap().aborted.is_empty());
2738
2739 drop(first_lease);
2740 queue.update().await;
2741 assert_eq!(state.lock().unwrap().aborted, vec![AdmissionId::new(0)]);
2742 slots.free(&"duplicate".to_owned(), decay_now()).unwrap();
2743 }
2744
2745 struct BypassGate {
2746 events: Arc<AtomicUsize>,
2747 }
2748
2749 impl PolicyClassAdmissionPolicy for BypassGate {
2750 fn admit(&mut self, _request: AdmissionRequest<'_>) -> AdmissionDecision {
2751 AdmissionDecision::Bypass
2752 }
2753
2754 fn on_event(&mut self, _event: AdmissionEvent) -> Vec<AdmissionAction> {
2755 self.events.fetch_add(1, Ordering::Relaxed);
2756 Vec::new()
2757 }
2758 }
2759
2760 #[tokio::test]
2761 async fn disabled_queueing_has_no_cancellation_lease() {
2762 let (queue, _slots) = make_queue(1, 16, 64, None);
2763
2764 assert!(
2765 queue
2766 .new_request_lifecycle_lease(Some("default-path"))
2767 .is_none()
2768 );
2769 }
2770
2771 #[tokio::test]
2772 async fn raw_queue_rejects_admission_mode_without_lease() {
2773 let state = Arc::new(StdMutex::new(GateState::default()));
2774 let (queue, _slots) = make_queue_with_admission_policy(Box::new(ReconcileGate {
2775 state: Arc::clone(&state),
2776 }));
2777 let (mut request, response) = make_admission_request("raw-admission", 64);
2778 request.policy_class = Some("agents".to_owned());
2779
2780 queue.enqueue(request).await;
2781
2782 let error = response.await.unwrap().unwrap_err();
2783 assert!(matches!(
2784 error,
2785 KvSchedulerError::BookingFailed(message)
2786 if message.contains("must be scheduled through LocalScheduler")
2787 ));
2788 assert_eq!(queue.pending_count(), 0);
2789 assert!(state.lock().unwrap().deferred.is_none());
2790 }
2791
2792 #[tokio::test]
2793 async fn legacy_tracked_request_cannot_bypass_admission_policy() {
2794 let state = Arc::new(StdMutex::new(GateState::default()));
2795 let (queue, _slots) = make_queue_with_admission_policy(Box::new(ReadyGate { state }));
2796 let (mut request, response) = make_request("legacy", 64);
2797 request.policy_class = Some("agents".to_owned());
2798
2799 queue.enqueue(request).await;
2800 let error = response.await.unwrap().unwrap_err();
2801
2802 assert!(matches!(error, KvSchedulerError::BookingFailed(message)
2803 if message.contains("requires lifecycle-tracked scheduling")));
2804 assert_eq!(queue.pending_count(), 0);
2805 }
2806
2807 #[tokio::test]
2808 async fn bypassed_request_has_no_admission_lifecycle() {
2809 let events = Arc::new(AtomicUsize::new(0));
2810 let (queue, slots) = make_queue_with_admission_policy(Box::new(BypassGate {
2811 events: Arc::clone(&events),
2812 }));
2813 let (mut request, response) = make_admission_request("bypassed", 64);
2814 request.policy_class = Some("agents".to_owned());
2815 {
2816 let mut lease = enqueue_with_lease(&queue, request).await;
2817 let selected = response.await.unwrap().unwrap();
2818 assert!(selected.request_progress.is_none());
2819 lease.disarm();
2820 }
2821
2822 assert_eq!(events.load(Ordering::Relaxed), 0);
2823 slots.free(&"bypassed".to_owned(), decay_now()).unwrap();
2824
2825 let (request, response) = make_request("unmanaged", 64);
2826 queue.enqueue(request).await;
2827 let selected = response.await.unwrap().unwrap();
2828 assert!(selected.request_progress.is_none());
2829 slots.free(&"unmanaged".to_owned(), decay_now()).unwrap();
2830 }
2831
2832 #[tokio::test]
2833 async fn cancellation_after_bypassed_handoff_rolls_back_booking() {
2834 let events = Arc::new(AtomicUsize::new(0));
2835 let (queue, slots) = make_queue_with_admission_policy(Box::new(BypassGate {
2836 events: Arc::clone(&events),
2837 }));
2838 let (mut request, response) = make_admission_request("bypassed-handoff", 64);
2839 request.policy_class = Some("agents".to_owned());
2840 let cancellation = enqueue_with_lease(&queue, request).await;
2841 assert_eq!(
2842 slots
2843 .active_request_counts()
2844 .get(&WorkerWithDpRank::new(0, 0))
2845 .copied(),
2846 Some(1)
2847 );
2848
2849 drop(cancellation);
2850 queue.update().await;
2851
2852 slots.assert_completely_drained(decay_now());
2853 assert_eq!(events.load(Ordering::Relaxed), 0);
2854 drop(response);
2855 }
2856
2857 #[derive(Default)]
2858 struct FinishReleaseGate {
2859 first: Option<AdmissionId>,
2860 deferred: Option<AdmissionId>,
2861 }
2862
2863 impl PolicyClassAdmissionPolicy for FinishReleaseGate {
2864 fn admit(&mut self, request: AdmissionRequest<'_>) -> AdmissionDecision {
2865 if self.first.is_none() {
2866 self.first = Some(request.id());
2867 AdmissionDecision::Ready(WorkerPlacement::Any)
2868 } else {
2869 self.deferred = Some(request.id());
2870 AdmissionDecision::Defer
2871 }
2872 }
2873
2874 fn on_event(&mut self, event: AdmissionEvent) -> Vec<AdmissionAction> {
2875 let id = match event {
2876 AdmissionEvent::Completed { id, .. } | AdmissionEvent::Aborted { id } => id,
2877 _ => return Vec::new(),
2878 };
2879 if self.first != Some(id) {
2880 return Vec::new();
2881 }
2882 self.deferred
2883 .take()
2884 .map(|id| {
2885 vec![AdmissionAction::MakeReady {
2886 id,
2887 placement: WorkerPlacement::Any,
2888 }]
2889 })
2890 .unwrap_or_default()
2891 }
2892 }
2893
2894 #[tokio::test]
2895 async fn lifecycle_action_drains_without_an_unrelated_update() {
2896 let (queue, slots) = make_queue_with_admission_policy(Box::<FinishReleaseGate>::default());
2897 let (mut first, first_response) = make_admission_request("first-admitted", 64);
2898 first.policy_class = Some("agents".to_owned());
2899 let mut first_lease = enqueue_with_lease(&queue, first).await;
2900 first_response.await.unwrap().unwrap();
2901
2902 let (mut second, second_response) = make_admission_request("second-deferred", 64);
2903 second.policy_class = Some("agents".to_owned());
2904 let second_lease = enqueue_with_lease(&queue, second).await;
2905 assert_eq!(queue.pending_count(), 1);
2906
2907 first_lease.mark_completed(64);
2908 drop(first_lease);
2909 tokio::time::timeout(Duration::from_secs(1), second_response)
2910 .await
2911 .expect("finish action did not drain the queue")
2912 .unwrap()
2913 .unwrap();
2914 assert_eq!(queue.pending_count(), 0);
2915 drop(second_lease);
2916 queue.update().await;
2917 slots.assert_completely_drained(decay_now());
2918 }
2919
2920 #[derive(Default)]
2921 struct PreservePinGate {
2922 deferred: Option<AdmissionId>,
2923 }
2924
2925 impl PolicyClassAdmissionPolicy for PreservePinGate {
2926 fn admit(&mut self, request: AdmissionRequest<'_>) -> AdmissionDecision {
2927 if self.deferred.is_none() {
2928 self.deferred = Some(request.id());
2929 AdmissionDecision::Defer
2930 } else {
2931 AdmissionDecision::Ready(WorkerPlacement::Any)
2932 }
2933 }
2934
2935 fn on_event(&mut self, event: AdmissionEvent) -> Vec<AdmissionAction> {
2936 if !matches!(event, AdmissionEvent::Reconcile) {
2937 return Vec::new();
2938 }
2939 self.deferred
2940 .take()
2941 .map(|id| {
2942 vec![AdmissionAction::MakeReady {
2943 id,
2944 placement: WorkerPlacement::Any,
2945 }]
2946 })
2947 .unwrap_or_default()
2948 }
2949 }
2950
2951 #[tokio::test]
2952 async fn make_ready_any_preserves_existing_exact_worker_lane() {
2953 let (queue, slots) =
2954 make_queue_with_admission_policy_and_workers(Box::<PreservePinGate>::default(), 2);
2955 for worker_id in 0..2 {
2956 let (mut blocker, response) = make_request(&format!("blocker-{worker_id}"), 64);
2957 blocker.pinned_worker = Some(WorkerWithDpRank::new(worker_id, 0));
2958 queue.enqueue(blocker).await;
2959 assert_eq!(
2960 response.await.unwrap().unwrap().best_worker,
2961 WorkerWithDpRank::new(worker_id, 0)
2962 );
2963 }
2964
2965 let (mut pinned, pinned_response) = make_admission_request("pinned-deferred", 64);
2966 pinned.policy_class = Some("agents".to_owned());
2967 pinned.pinned_worker = Some(WorkerWithDpRank::new(0, 0));
2968 let pinned_lease = enqueue_with_lease(&queue, pinned).await;
2969
2970 let (mut runnable, runnable_response) = make_admission_request("runnable-shared", 64);
2971 runnable.policy_class = Some("agents".to_owned());
2972 let runnable_lease = enqueue_with_lease(&queue, runnable).await;
2973 assert_eq!(queue.pending_count(), 2);
2974
2975 slots.free(&"blocker-1".to_owned(), decay_now()).unwrap();
2976 queue.reconcile().await;
2977 assert_eq!(
2978 tokio::time::timeout(Duration::from_secs(1), runnable_response)
2979 .await
2980 .expect("pinned lane blocked shared work")
2981 .unwrap()
2982 .unwrap()
2983 .best_worker,
2984 WorkerWithDpRank::new(1, 0)
2985 );
2986
2987 slots.free(&"blocker-0".to_owned(), decay_now()).unwrap();
2988 queue.update().await;
2989 assert_eq!(
2990 pinned_response.await.unwrap().unwrap().best_worker,
2991 WorkerWithDpRank::new(0, 0)
2992 );
2993 drop(runnable_lease);
2994 drop(pinned_lease);
2995 queue.update().await;
2996 slots.assert_completely_drained(decay_now());
2997 }
2998
2999 #[tokio::test]
3000 async fn family_ready_exact_uses_pinned_worker_queue_cost() {
3001 let profile = policy_profile(
3002 r#"
3003default_policy_family: agents
3004uncached_isl_buckets:
3005 - min_tokens: 0
3006 bucket: cached
3007 - min_tokens: 32
3008 bucket: uncached
3009policy_classes:
3010 - name: agents_cached
3011 policy_family: agents
3012 cache_bucket: cached
3013 prefill_busy_threshold: 0
3014 quantum: 1
3015 - name: agents_uncached
3016 policy_family: agents
3017 cache_bucket: uncached
3018 prefill_busy_threshold: 0
3019 quantum: 1
3020"#,
3021 );
3022 let worker = WorkerWithDpRank::new(0, 0);
3023 let (queue, slots) = make_queue_with_profile_and_admission_policy(
3024 profile,
3025 "agents_cached",
3026 Box::new(ExactReadyGate(worker)),
3027 2,
3028 );
3029 let (mut blocker, blocker_response) = make_request("family-exact-blocker", 64);
3030 blocker.pinned_worker = Some(worker);
3031 queue.enqueue(blocker).await;
3032 blocker_response.await.unwrap().unwrap();
3033
3034 let (mut request, response) = make_admission_request("family-exact", 64);
3035 request
3036 .overlap
3037 .effective_cached_tokens
3038 .insert(WorkerWithDpRank::new(1, 0), 64);
3039 let lease = enqueue_with_lease(&queue, request).await;
3040
3041 assert_eq!(queue.class_queue_stats(0).unwrap().pending_cached_tokens, 0);
3042 assert_eq!(queue.class_queue_stats(0).unwrap().pending_count, 0);
3043 assert_eq!(queue.class_queue_stats(1).unwrap().pending_count, 1);
3044 slots
3045 .free(&"family-exact-blocker".to_owned(), decay_now())
3046 .unwrap();
3047 queue.update().await;
3048 assert_eq!(response.await.unwrap().unwrap().best_worker, worker);
3049 drop(lease);
3050 queue.update().await;
3051 slots.assert_completely_drained(decay_now());
3052 }
3053
3054 #[tokio::test]
3055 async fn make_ready_exact_recomputes_queue_cost_for_pinned_worker() {
3056 let state = Arc::new(StdMutex::new(GateState::default()));
3057 let (queue, slots) = make_queue_with_admission_policy_and_workers(
3058 Box::new(ReconcileGate {
3059 state: Arc::clone(&state),
3060 }),
3061 2,
3062 );
3063 let worker = WorkerWithDpRank::new(0, 0);
3064 let (mut blocker, blocker_response) = make_request("exact-cost-blocker", 64);
3065 blocker.pinned_worker = Some(worker);
3066 queue.enqueue(blocker).await;
3067 blocker_response.await.unwrap().unwrap();
3068
3069 let (mut request, response) = make_admission_request("exact-cost", 64);
3070 request.policy_class = Some("agents".to_owned());
3071 request
3072 .overlap
3073 .effective_cached_tokens
3074 .insert(WorkerWithDpRank::new(1, 0), 64);
3075 let lease = enqueue_with_lease(&queue, request).await;
3076 assert_eq!(
3077 queue.class_queue_stats(1).unwrap().pending_cached_tokens,
3078 64
3079 );
3080
3081 queue.reconcile().await;
3082 assert_eq!(queue.class_queue_stats(1).unwrap().pending_cached_tokens, 0);
3083
3084 slots
3085 .free(&"exact-cost-blocker".to_owned(), decay_now())
3086 .unwrap();
3087 queue.update().await;
3088 assert_eq!(response.await.unwrap().unwrap().best_worker, worker);
3089 drop(lease);
3090 queue.update().await;
3091 slots.assert_completely_drained(decay_now());
3092 }
3093
3094 #[tokio::test]
3095 async fn deferred_exact_make_ready_reclassifies_physical_queue() {
3096 let profile = policy_profile(
3097 r#"
3098default_policy_family: agents
3099uncached_isl_buckets:
3100 - min_tokens: 0
3101 bucket: cached
3102 - min_tokens: 32
3103 bucket: uncached
3104policy_classes:
3105 - name: agents_cached
3106 policy_family: agents
3107 cache_bucket: cached
3108 prefill_busy_threshold: 0
3109 quantum: 1
3110 - name: agents_uncached
3111 policy_family: agents
3112 cache_bucket: uncached
3113 prefill_busy_threshold: 0
3114 quantum: 1
3115"#,
3116 );
3117 let state = Arc::new(StdMutex::new(GateState::default()));
3118 let worker = WorkerWithDpRank::new(0, 0);
3119 let (queue, slots) = make_queue_with_profile_and_admission_policy(
3120 profile,
3121 "agents_cached",
3122 Box::new(ReconcileGate {
3123 state: Arc::clone(&state),
3124 }),
3125 2,
3126 );
3127 let (mut blocker, blocker_response) = make_request("reclassify-blocker", 64);
3128 blocker.pinned_worker = Some(worker);
3129 queue.enqueue(blocker).await;
3130 blocker_response.await.unwrap().unwrap();
3131
3132 let (mut request, response) = make_admission_request("reclassify-deferred", 64);
3133 request
3134 .overlap
3135 .effective_cached_tokens
3136 .insert(WorkerWithDpRank::new(1, 0), 64);
3137 let lease = enqueue_with_lease(&queue, request).await;
3138 assert_eq!(queue.class_queue_stats(0).unwrap().pending_count, 1);
3139 assert_eq!(
3140 queue.class_queue_stats(0).unwrap().pending_cached_tokens,
3141 64
3142 );
3143
3144 queue.reconcile().await;
3145 assert_eq!(queue.class_queue_stats(0).unwrap().pending_count, 0);
3146 assert_eq!(queue.class_queue_stats(1).unwrap().pending_count, 1);
3147 assert_eq!(queue.class_queue_stats(1).unwrap().pending_cached_tokens, 0);
3148
3149 slots
3150 .free(&"reclassify-blocker".to_owned(), decay_now())
3151 .unwrap();
3152 queue.update().await;
3153 assert_eq!(response.await.unwrap().unwrap().best_worker, worker);
3154 drop(lease);
3155 queue.update().await;
3156 slots.assert_completely_drained(decay_now());
3157 }
3158
3159 #[tokio::test]
3160 async fn cancelled_ready_requests_release_accounting_immediately() {
3161 let state = Arc::new(StdMutex::new(GateState::default()));
3162 let (queue, slots) = make_queue_with_admission_policy(Box::new(ReadyGate {
3163 state: Arc::clone(&state),
3164 }));
3165
3166 let (blocker, blocker_response) = make_request("blocker", 64);
3167 queue.enqueue(blocker).await;
3168 blocker_response.await.unwrap().unwrap();
3169
3170 for request_id in [
3171 "cancelled-ready-0",
3172 "cancelled-ready-1",
3173 "cancelled-ready-2",
3174 ] {
3175 let (mut cancelled, cancelled_response) = make_admission_request(request_id, 64);
3176 cancelled.policy_class = Some("agents".to_owned());
3177 let cancellation = enqueue_with_lease(&queue, cancelled).await;
3178 drop(cancelled_response);
3179 drop(cancellation);
3180 }
3181 queue.update().await;
3182 assert_eq!(state.lock().unwrap().aborted.len(), 3);
3183 assert_eq!(queue.pending_count(), 0);
3184 assert_eq!(queue.pending_isl_tokens(), 0);
3185 assert_eq!(
3186 queue.class_queue_stats(1),
3187 Some(ClassQueueStats {
3188 pending_count: 0,
3189 pending_isl_tokens: 0,
3190 pending_cached_tokens: 0,
3191 })
3192 );
3193 let state = state.lock().unwrap();
3194 assert!(state.dispatched.is_empty());
3195 assert_eq!(state.aborted.len(), 3);
3196 assert_eq!(
3197 state.aborted.iter().copied().collect::<HashSet<_>>(),
3198 [
3199 AdmissionId::new(0),
3200 AdmissionId::new(1),
3201 AdmissionId::new(2)
3202 ]
3203 .into_iter()
3204 .collect()
3205 );
3206 slots.free(&"blocker".to_owned(), decay_now()).unwrap();
3207 }
3208
3209 #[tokio::test]
3210 async fn cancelled_ready_head_redrives_newly_exposed_request() {
3211 let state = Arc::new(StdMutex::new(GateState::default()));
3212 let (queue, slots) =
3213 make_queue_with_admission_policy_and_workers(Box::new(ReadyGate { state }), 2);
3214 let worker_0 = WorkerWithDpRank::new(0, 0);
3215 let (mut blocker, blocker_response) = make_request("redrive-blocker", 64);
3216 blocker.pinned_worker = Some(worker_0);
3217 queue.enqueue(blocker).await;
3218 blocker_response.await.unwrap().unwrap();
3219
3220 let (mut cancelled, cancelled_response) = make_admission_request("redrive-cancelled", 64);
3221 cancelled.policy_class = Some("agents".to_owned());
3222 cancelled.allowed_worker_ids = Some(HashSet::from([0]));
3223 let cancelled_lease = enqueue_with_lease(&queue, cancelled).await;
3224
3225 let (mut exposed, exposed_response) = make_admission_request("redrive-exposed", 64);
3226 exposed.policy_class = Some("agents".to_owned());
3227 exposed.allowed_worker_ids = Some(HashSet::from([1]));
3228 let exposed_lease = enqueue_with_lease(&queue, exposed).await;
3229 assert_eq!(queue.pending_count(), 2);
3230
3231 drop(cancelled_response);
3232 drop(cancelled_lease);
3233 let selected = tokio::time::timeout(Duration::from_secs(1), exposed_response)
3234 .await
3235 .expect("cancellation did not redrive the exposed request")
3236 .unwrap()
3237 .unwrap();
3238 assert_eq!(selected.best_worker, WorkerWithDpRank::new(1, 0));
3239
3240 drop(exposed_lease);
3241 queue.update().await;
3242 slots
3243 .free(&"redrive-blocker".to_owned(), decay_now())
3244 .unwrap();
3245 slots.assert_completely_drained(decay_now());
3246 }
3247
3248 #[tokio::test]
3249 async fn cancelled_ready_cleanup_does_not_remove_reused_request_id() {
3250 let state = Arc::new(StdMutex::new(GateState::default()));
3251 let (queue, slots) = make_queue_with_admission_policy(Box::new(ReadyGate {
3252 state: Arc::clone(&state),
3253 }));
3254
3255 let (blocker, blocker_response) = make_request("blocker", 64);
3256 queue.enqueue(blocker).await;
3257 blocker_response.await.unwrap().unwrap();
3258
3259 let (mut cancelled, cancelled_response) = make_admission_request("reused", 64);
3260 cancelled.policy_class = Some("agents".to_owned());
3261 let cancellation = enqueue_with_lease(&queue, cancelled).await;
3262 let stale_ticket = cancellation.ticket.unwrap();
3263 drop(cancelled_response);
3264 drop(cancellation);
3265 queue.update().await;
3266 assert_eq!(state.lock().unwrap().aborted, vec![AdmissionId::new(0)]);
3267
3268 let (mut replacement, replacement_response) = make_admission_request("reused", 64);
3269 replacement.policy_class = Some("agents".to_owned());
3270 let mut replacement_lease = enqueue_with_lease(&queue, replacement).await;
3271 assert_eq!(queue.pending_count(), 1);
3272 assert_eq!(state.lock().unwrap().aborted, vec![AdmissionId::new(0)]);
3273
3274 slots.free(&"blocker".to_owned(), decay_now()).unwrap();
3275 queue.update().await;
3276 let selected = replacement_response.await.unwrap().unwrap();
3277 assert!(selected.request_progress.is_some());
3278
3279 queue
3280 .admission_tx
3281 .send(AdmissionCommand::Dispatched {
3282 request_id: "reused".to_owned(),
3283 ticket: stale_ticket,
3284 })
3285 .await
3286 .unwrap();
3287 queue.update().await;
3288 assert!(state.lock().unwrap().dispatched.is_empty());
3289
3290 replacement_lease.mark_dispatched().await;
3291 replacement_lease.mark_completed(64);
3292 drop(replacement_lease);
3293 queue.update().await;
3294 assert_eq!(state.lock().unwrap().dispatched.len(), 1);
3295 assert_eq!(state.lock().unwrap().completed_context_tokens, vec![64]);
3296 slots.assert_completely_drained(decay_now());
3297 }
3298
3299 #[tokio::test]
3300 async fn backend_abort_finishes_without_dispatching_admission() {
3301 let state = Arc::new(StdMutex::new(GateState::default()));
3302 let (queue, slots) = make_queue_with_admission_policy(Box::new(ReconcileGate {
3303 state: Arc::clone(&state),
3304 }));
3305 let (mut request, response) = make_admission_request("backend-abort", 64);
3306 request.policy_class = Some("agents".to_owned());
3307
3308 let lease = enqueue_with_lease(&queue, request).await;
3309 queue.reconcile().await;
3310 response.await.unwrap().unwrap();
3311 drop(lease);
3312 queue.update().await;
3313
3314 let state = state.lock().unwrap();
3315 assert!(state.dispatched.is_empty());
3316 assert_eq!(state.aborted, vec![AdmissionId::new(0)]);
3317 slots.assert_completely_drained(decay_now());
3318 }
3319
3320 #[tokio::test(flavor = "multi_thread")]
3321 async fn test_cancelled_pending_request_is_not_booked() {
3322 let isl = 512;
3323 let (queue, slots) = make_queue(1, 16, isl, Some(0.0));
3324
3325 let (first, first_rx) = make_request("first", isl);
3326 queue.enqueue(first).await;
3327 first_rx
3328 .await
3329 .expect("first response sender dropped")
3330 .expect("first request should be scheduled");
3331
3332 let (cancelled, cancelled_rx) = make_request("cancelled", isl);
3333 queue.enqueue(cancelled).await;
3334 assert_eq!(queue.pending_count(), 1);
3335 drop(cancelled_rx);
3336
3337 slots.free(&"first".to_string(), decay_now()).unwrap();
3338 queue.update().await;
3339
3340 assert_eq!(queue.pending_count(), 0);
3341 slots.assert_completely_drained(decay_now());
3342 }
3343
3344 #[tokio::test]
3345 async fn dropped_legacy_lease_retracts_pending_request_immediately() {
3346 let isl = 512;
3347 let (queue, slots) = make_queue(1, 16, isl, Some(0.0));
3348
3349 let (first, first_rx) = make_request("legacy-first", isl);
3350 queue.enqueue(first).await;
3351 first_rx.await.unwrap().unwrap();
3352
3353 let (cancelled, cancelled_rx) = make_request("legacy-cancelled", isl);
3354 let lease = queue
3355 .new_request_lifecycle_lease(Some("legacy-cancelled"))
3356 .unwrap();
3357 let lease = queue
3358 .enqueue_with_block_hashes_and_lease(cancelled, None, Some(lease))
3359 .await
3360 .unwrap();
3361 assert_eq!(queue.pending_count(), 1);
3362
3363 drop(cancelled_rx);
3364 drop(lease);
3365 tokio::time::timeout(Duration::from_secs(1), async {
3366 while queue.pending_count() != 0 {
3367 tokio::task::yield_now().await;
3368 }
3369 })
3370 .await
3371 .expect("legacy lease did not retract pending request");
3372
3373 slots.free(&"legacy-first".to_owned(), decay_now()).unwrap();
3374 slots.assert_completely_drained(decay_now());
3375 }
3376
3377 #[tokio::test(flavor = "multi_thread")]
3378 async fn test_strict_priority_drains_before_policy_score() {
3379 let isl = 512;
3380 let (queue, slots) = make_queue(1, 16, isl, Some(0.0));
3381
3382 let (first, first_rx) = make_request("first", isl);
3383 queue.enqueue(first).await;
3384 first_rx.await.unwrap().unwrap();
3385
3386 let (mut low, mut low_rx) = make_request("low", isl);
3387 low.priority_jump = 10_000.0;
3388 queue.enqueue(low).await;
3389
3390 let (mut high, high_rx) = make_request("high", isl);
3391 high.strict_priority = 1;
3392 queue.enqueue(high).await;
3393 assert_eq!(queue.pending_count(), 2);
3394
3395 slots.free(&"first".to_string(), decay_now()).unwrap();
3396 queue.update().await;
3397
3398 let high_response = high_rx.await.unwrap().unwrap();
3399 assert_eq!(high_response.best_worker, WorkerWithDpRank::new(0, 0));
3400 assert!(
3401 low_rx.try_recv().is_err(),
3402 "lower strict priority should remain queued"
3403 );
3404
3405 slots.free(&"high".to_string(), decay_now()).unwrap();
3406 queue.update().await;
3407 low_rx.await.unwrap().unwrap();
3408 assert_eq!(queue.pending_count(), 0);
3409
3410 slots.free(&"low".to_string(), decay_now()).unwrap();
3411 slots.assert_completely_drained(decay_now());
3412 }
3413
3414 #[tokio::test(flavor = "multi_thread")]
3415 async fn test_failed_response_delivery_rolls_back_booking() {
3416 let isl = 512;
3417 let response_rx = Arc::new(StdMutex::new(None));
3418 let publisher = DropResponseOnLoadPublisher {
3419 response_rx: Arc::clone(&response_rx),
3420 };
3421 let slots = Arc::new(ActiveSequencesMultiWorker::new(
3422 publisher,
3423 16,
3424 HashMap::from([(0, (0, 1))]),
3425 false,
3426 0,
3427 "test",
3428 ));
3429 let (_cfg_tx, cfg_rx) = watch::channel(HashMap::from([(
3430 0,
3431 SimpleWorkerConfig {
3432 max_num_batched_tokens: Some(isl as u64),
3433 ..Default::default()
3434 },
3435 )]));
3436 let queue = SchedulerQueue::new(
3437 Arc::clone(&slots),
3438 cfg_rx,
3439 None,
3440 16,
3441 DefaultWorkerSelector::new(None, "test"),
3442 RouterQueuePolicy::Fcfs,
3443 None,
3444 );
3445
3446 let (request, receiver) = make_request("delivery-race", isl);
3447 *response_rx.lock().unwrap() = Some(receiver);
3448 queue.enqueue(request).await;
3449
3450 assert!(response_rx.lock().unwrap().is_none());
3451 slots.assert_completely_drained(decay_now());
3452 }
3453
3454 #[tokio::test(flavor = "multi_thread")]
3455 async fn test_concurrent_flood() {
3456 let block_size = 16;
3457 let isl = 512;
3458 let num_workers = 4;
3459 let num_tasks = 25;
3460
3461 let (queue, slots) = make_queue(num_workers, block_size, isl, None);
3462
3463 let mut handles = Vec::new();
3464 for i in 0..num_tasks {
3465 let queue = Arc::clone(&queue);
3466 let slots = Arc::clone(&slots);
3467 handles.push(tokio::spawn(async move {
3468 let req_id = format!("req-{i}");
3469 let (req, rx) = make_request(&req_id, isl);
3470 queue.enqueue(req).await;
3471 let resp = rx.await.expect("oneshot dropped");
3472 let resp = resp.expect("scheduling failed");
3473 assert!(resp.best_worker.worker_id < num_workers as u64);
3474
3475 slots.mark_prefill_completed(&req_id, decay_now()).unwrap();
3476 slots.free(&req_id, decay_now()).unwrap();
3477 queue.update().await;
3478 }));
3479 }
3480
3481 for h in handles {
3482 h.await.expect("task panicked");
3483 }
3484
3485 let active = slots.active_tokens(decay_now());
3486 for (worker, tokens) in &active {
3487 assert_eq!(
3488 *tokens, 0,
3489 "worker {worker:?} still has {tokens} active tokens"
3490 );
3491 }
3492 }
3493
3494 #[tokio::test(flavor = "multi_thread")]
3495 async fn test_concurrent_immediate_admissions_see_prior_booking() {
3496 let selector = MinDecodeSelector {
3497 rendezvous: Some(Arc::new(SelectorRendezvous::default())),
3498 };
3499 let (queue, slots) = make_queue_with_custom_selector(2, 16, 512, None, selector);
3500 let barrier = Arc::new(Barrier::new(3));
3501
3502 let (req1, rx1) = make_request("req-1", 512);
3503 let queue1 = Arc::clone(&queue);
3504 let barrier1 = Arc::clone(&barrier);
3505 let handle1 = tokio::spawn(async move {
3506 barrier1.wait().await;
3507 queue1.enqueue(req1).await;
3508 });
3509
3510 let (req2, rx2) = make_request("req-2", 512);
3511 let queue2 = Arc::clone(&queue);
3512 let barrier2 = Arc::clone(&barrier);
3513 let handle2 = tokio::spawn(async move {
3514 barrier2.wait().await;
3515 queue2.enqueue(req2).await;
3516 });
3517
3518 barrier.wait().await;
3519 handle1.await.unwrap();
3520 handle2.await.unwrap();
3521
3522 let resp1 = rx1.await.unwrap().unwrap();
3523 let resp2 = rx2.await.unwrap().unwrap();
3524 assert_ne!(
3525 resp1.best_worker, resp2.best_worker,
3526 "second admission should see the first booking and choose the other idle worker"
3527 );
3528
3529 for request_id in ["req-1", "req-2"] {
3530 slots
3531 .mark_prefill_completed(&request_id.to_string(), decay_now())
3532 .unwrap();
3533 slots.free(&request_id.to_string(), decay_now()).unwrap();
3534 }
3535 }
3536
3537 #[tokio::test(flavor = "multi_thread")]
3538 async fn test_queueing_under_pressure() {
3539 let block_size = 16;
3540 let isl = 512;
3541 let num_workers = 2;
3542 let num_requests = 10;
3543
3544 let (queue, slots) = make_queue(num_workers, block_size, isl, Some(0.0));
3545
3546 let mut receivers = Vec::new();
3547 let mut req_ids = Vec::new();
3548
3549 for i in 0..num_requests {
3550 let req_id = format!("pressure-{i}");
3551 let (req, rx) = make_request(&req_id, isl);
3552 queue.enqueue(req).await;
3553 receivers.push(rx);
3554 req_ids.push(req_id);
3555 }
3556
3557 for _ in 0..num_requests {
3560 queue.update().await;
3561 for rid in &req_ids {
3562 let _ = slots.mark_prefill_completed(rid, decay_now());
3563 let _ = slots.free(rid, decay_now());
3564 }
3565 }
3566 queue.update().await;
3567
3568 let mut ok_count = 0;
3569 for mut rx in receivers {
3570 if let Ok(result) = rx.try_recv() {
3571 result.expect("scheduling returned error");
3572 ok_count += 1;
3573 }
3574 }
3575 assert_eq!(ok_count, num_requests, "not all requests were scheduled");
3576 }
3577
3578 #[tokio::test(flavor = "multi_thread")]
3579 async fn test_pending_requests_receive_shutdown_on_queue_drop() {
3580 let block_size = 16;
3581 let isl = 512;
3582 let (queue, _slots) = make_queue(1, block_size, isl, Some(0.0));
3583
3584 let (req1, rx1) = make_request("req-1", isl);
3585 queue.enqueue(req1).await;
3586 rx1.await
3587 .expect("first response sender dropped")
3588 .expect("first request should be scheduled");
3589
3590 let (req2, rx2) = make_request("req-2", isl);
3591 queue.enqueue(req2).await;
3592 assert_eq!(queue.pending_count(), 1);
3593
3594 drop(queue);
3595
3596 let response = tokio::time::timeout(Duration::from_secs(1), rx2)
3597 .await
3598 .expect("shutdown response timed out")
3599 .expect("pending response sender dropped");
3600 assert!(matches!(
3601 response,
3602 Err(KvSchedulerError::SubscriberShutdown)
3603 ));
3604 }
3605
3606 #[tokio::test(flavor = "multi_thread")]
3607 async fn test_pending_count() {
3608 let block_size = 16;
3609 let isl = 512;
3610 let num_workers = 1;
3611
3612 let (queue, slots) = make_queue(num_workers, block_size, isl, Some(0.0));
3614 assert_eq!(queue.pending_count(), 0);
3615
3616 let (req1, rx1) = make_request("req-1", isl);
3618 queue.enqueue(req1).await;
3619 let _resp1 = rx1.await.unwrap().unwrap();
3620 assert_eq!(queue.pending_count(), 0); let (req2, _rx2) = make_request("req-2", isl);
3624 queue.enqueue(req2).await;
3625 assert_eq!(queue.pending_count(), 1);
3626
3627 let (req3, _rx3) = make_request("req-3", isl);
3628 queue.enqueue(req3).await;
3629 assert_eq!(queue.pending_count(), 2);
3630
3631 slots
3633 .mark_prefill_completed(&"req-1".to_string(), decay_now())
3634 .unwrap();
3635 slots.free(&"req-1".to_string(), decay_now()).unwrap();
3636 queue.update().await;
3637
3638 assert!(
3640 queue.pending_count() < 2,
3641 "pending_count should decrease after free+update, got {}",
3642 queue.pending_count()
3643 );
3644
3645 let _ = slots.mark_prefill_completed(&"req-2".to_string(), decay_now());
3647 let _ = slots.free(&"req-2".to_string(), decay_now());
3648 queue.update().await;
3649 let _ = slots.mark_prefill_completed(&"req-3".to_string(), decay_now());
3650 let _ = slots.free(&"req-3".to_string(), decay_now());
3651 queue.update().await;
3652
3653 assert_eq!(queue.pending_count(), 0, "all requests should be drained");
3654 }
3655
3656 #[tokio::test(flavor = "multi_thread")]
3657 async fn policy_classes_apply_independent_thresholds_and_preserve_backlog_order() {
3658 let profile = policy_profile(
3659 r#"
3660default_policy_family: latency
3661uncached_isl_buckets:
3662 - min_tokens: 0
3663 bucket: all
3664policy_classes:
3665 - name: latency
3666 policy_family: latency
3667 cache_bucket: all
3668 quantum: 1
3669 prefill_busy_threshold: 0
3670 - name: bulk
3671 policy_family: bulk
3672 cache_bucket: all
3673 quantum: 1
3674 prefill_busy_threshold: 1024
3675"#,
3676 );
3677 let (queue, slots) = make_queue_with_profile(1, 16, 64, profile);
3678
3679 let (mut active, active_rx) = make_request("active", 64);
3680 active.policy_class = Some("latency".to_string());
3681 queue.enqueue(active).await;
3682 active_rx.await.unwrap().unwrap();
3683
3684 let (mut bulk, bulk_rx) = make_request("bulk", 64);
3685 bulk.policy_class = Some("bulk".to_string());
3686 queue.enqueue(bulk).await;
3687 bulk_rx.await.unwrap().unwrap();
3688
3689 let (mut queued_first, mut queued_first_rx) = make_request("queued-first", 64);
3690 queued_first.policy_class = Some("latency".to_string());
3691 queue.enqueue(queued_first).await;
3692 assert_eq!(queue.pending_count(), 1);
3693
3694 for request_id in ["active", "bulk"] {
3695 slots
3696 .mark_prefill_completed(&request_id.to_string(), decay_now())
3697 .unwrap();
3698 slots.free(&request_id.to_string(), decay_now()).unwrap();
3699 }
3700
3701 let (mut queued_second, mut queued_second_rx) = make_request("queued-second", 64);
3702 queued_second.policy_class = Some("latency".to_string());
3703 queue.enqueue(queued_second).await;
3704 assert_eq!(
3705 queue.pending_count(),
3706 2,
3707 "new arrivals must not bypass backlog"
3708 );
3709 assert!(queued_first_rx.try_recv().is_err());
3710 assert!(queued_second_rx.try_recv().is_err());
3711
3712 queue.update().await;
3713 queued_first_rx
3714 .try_recv()
3715 .expect("first queued request should be admitted")
3716 .expect("first queued request failed");
3717 assert!(
3718 queued_second_rx.try_recv().is_err(),
3719 "second request should remain behind the admitted head"
3720 );
3721
3722 slots
3723 .mark_prefill_completed(&"queued-first".to_string(), decay_now())
3724 .unwrap();
3725 slots
3726 .free(&"queued-first".to_string(), decay_now())
3727 .unwrap();
3728 queue.update().await;
3729 queued_second_rx.await.unwrap().unwrap();
3730 }
3731
3732 #[tokio::test(flavor = "multi_thread")]
3733 async fn policy_families_and_cache_buckets_select_physical_queues() {
3734 let profile = policy_profile(
3735 r#"
3736default_policy_family: standard
3737uncached_isl_buckets:
3738 - min_tokens: 0
3739 bucket: cached
3740 - min_tokens: 32
3741 bucket: uncached
3742policy_classes:
3743 - name: cached
3744 policy_family: standard
3745 cache_bucket: cached
3746 quantum: 1
3747 prefill_busy_threshold: 0
3748 - name: uncached
3749 policy_family: standard
3750 cache_bucket: uncached
3751 quantum: 1
3752 prefill_busy_threshold: 0
3753 - name: latency_cached
3754 policy_family: latency
3755 cache_bucket: cached
3756 quantum: 1
3757 prefill_busy_threshold: 0
3758 - name: latency_uncached
3759 policy_family: latency
3760 cache_bucket: uncached
3761 quantum: 1
3762 prefill_busy_threshold: 0
3763 - name: custom_priority
3764 quantum: 1
3765 prefill_busy_threshold: 0
3766"#,
3767 );
3768 let (queue, _slots) = make_queue_with_profile(1, 16, 64, profile);
3769 let worker = WorkerWithDpRank::new(0, 0);
3770
3771 let (active, active_rx) = make_request("active", 64);
3772 queue.enqueue(active).await;
3773 active_rx.await.unwrap().unwrap();
3774
3775 let (mut latency_cached, _latency_cached_rx) = make_request("latency-cached", 64);
3776 latency_cached.policy_class = Some("latency".to_string());
3777 latency_cached
3778 .overlap
3779 .effective_cached_tokens
3780 .insert(worker, 64);
3781 queue.enqueue(latency_cached).await;
3782
3783 let (mut latency_uncached, _latency_uncached_rx) = make_request("latency-uncached", 64);
3784 latency_uncached.policy_class = Some("latency".to_string());
3785 queue.enqueue(latency_uncached).await;
3786
3787 let (mut unknown_cached, _unknown_cached_rx) = make_request("unknown-cached", 64);
3788 unknown_cached.policy_class = Some("unknown".to_string());
3789 unknown_cached
3790 .overlap
3791 .effective_cached_tokens
3792 .insert(worker, 64);
3793 queue.enqueue(unknown_cached).await;
3794
3795 let (mut ordinary_class_name, _ordinary_class_name_rx) =
3796 make_request("ordinary-class-name", 64);
3797 ordinary_class_name.policy_class = Some("latency_cached".to_string());
3798 queue.enqueue(ordinary_class_name).await;
3799
3800 let (mut custom, _custom_rx) = make_request("custom", 64);
3801 custom.policy_class = Some("custom_priority".to_string());
3802 queue.enqueue(custom).await;
3803
3804 assert_eq!(
3805 queue.class_queue_stats(0),
3806 Some(ClassQueueStats {
3807 pending_count: 1,
3808 pending_isl_tokens: 64,
3809 pending_cached_tokens: 64,
3810 })
3811 );
3812 assert_eq!(
3813 queue.class_queue_stats(1),
3814 Some(ClassQueueStats {
3815 pending_count: 1,
3816 pending_isl_tokens: 64,
3817 pending_cached_tokens: 0,
3818 })
3819 );
3820 assert_eq!(
3821 queue.class_queue_stats(2),
3822 Some(ClassQueueStats {
3823 pending_count: 1,
3824 pending_isl_tokens: 64,
3825 pending_cached_tokens: 64,
3826 })
3827 );
3828 assert_eq!(
3829 queue.class_queue_stats(3),
3830 Some(ClassQueueStats {
3831 pending_count: 1,
3832 pending_isl_tokens: 64,
3833 pending_cached_tokens: 0,
3834 })
3835 );
3836 assert_eq!(
3837 queue.class_queue_stats(4),
3838 Some(ClassQueueStats {
3839 pending_count: 1,
3840 pending_isl_tokens: 64,
3841 pending_cached_tokens: 0,
3842 })
3843 );
3844 }
3845
3846 #[tokio::test(flavor = "multi_thread")]
3847 async fn class_local_limit_rejection_is_typed_and_not_overload() {
3848 let profile = policy_profile(
3849 r#"
3850default_policy_family: capped
3851uncached_isl_buckets:
3852 - min_tokens: 0
3853 bucket: all
3854policy_classes:
3855 - name: capped
3856 policy_family: capped
3857 cache_bucket: all
3858 quantum: 1
3859 prefill_busy_threshold: 0
3860 request_queue_limit_per_worker: 1
3861"#,
3862 );
3863 let (queue, _slots) = make_queue_with_profile(1, 16, 64, profile);
3864
3865 let (active, active_rx) = make_request("active", 64);
3866 queue.enqueue(active).await;
3867 active_rx.await.unwrap().unwrap();
3868
3869 let (queued, _queued_rx) = make_request("queued", 64);
3870 queue.enqueue(queued).await;
3871
3872 let (rejected, rejected_rx) = make_request("rejected", 64);
3873 queue.enqueue(rejected).await;
3874 let error = rejected_rx.await.unwrap().unwrap_err();
3875 let KvSchedulerError::QueueRejected(rejection) = &error else {
3876 panic!("expected queue rejection, got {error:?}");
3877 };
3878 assert_eq!(rejection.policy_class, "capped");
3879 assert_eq!(rejection.limit_kind, super::super::QueueLimitKind::Requests);
3880 assert_eq!(rejection.current, 1);
3881 assert_eq!(rejection.limit, 1);
3882 assert!(!error.is_overload());
3883
3884 assert_eq!(
3885 queue.class_queue_stats(0),
3886 Some(ClassQueueStats {
3887 pending_count: 1,
3888 pending_isl_tokens: 64,
3889 pending_cached_tokens: 0,
3890 })
3891 );
3892 }
3893
3894 #[tokio::test(flavor = "multi_thread")]
3895 async fn per_worker_limit_tracks_discovered_worker_count_without_evicting() {
3896 let profile = policy_profile(
3897 r#"
3898default_policy_family: capped
3899uncached_isl_buckets:
3900 - min_tokens: 0
3901 bucket: all
3902policy_classes:
3903 - name: capped
3904 policy_family: capped
3905 cache_bucket: all
3906 quantum: 1
3907 prefill_busy_threshold: 0
3908 request_queue_limit_per_worker: 1
3909"#,
3910 );
3911 let (queue, _slots, cfg_tx) = make_queue_with_profile_and_sender(1, 16, 64, profile);
3912
3913 let (active, active_rx) = make_request("active", 64);
3914 queue.enqueue(active).await;
3915 active_rx.await.unwrap().unwrap();
3916
3917 let (first, _first_rx) = make_request("first", 64);
3918 queue.enqueue(first).await;
3919
3920 cfg_tx.send_modify(|configs| {
3921 configs.insert(
3922 1,
3923 SimpleWorkerConfig {
3924 max_num_batched_tokens: Some(64),
3925 ..Default::default()
3926 },
3927 );
3928 });
3929 let (second, _second_rx) = make_request("second", 64);
3930 queue.enqueue(second).await;
3931 assert_eq!(queue.pending_count(), 2);
3932
3933 cfg_tx.send_modify(|configs| {
3934 configs.remove(&1);
3935 });
3936 let (rejected, rejected_rx) = make_request("rejected", 64);
3937 queue.enqueue(rejected).await;
3938 let error = rejected_rx.await.unwrap().unwrap_err();
3939 let KvSchedulerError::QueueRejected(rejection) = error else {
3940 panic!("expected queue rejection, got {error:?}");
3941 };
3942 assert_eq!(rejection.current, 2);
3943 assert_eq!(rejection.limit, 1);
3944 assert_eq!(queue.pending_count(), 2);
3945 }
3946
3947 #[tokio::test(start_paused = true)]
3948 async fn test_queue_update_uses_decayed_oldest_prefill_load() {
3949 let estimator: Arc<dyn PrefillLoadEstimator> = Arc::new(FixedPrefillLoadEstimator {
3950 duration: Duration::from_secs(10),
3951 });
3952 let (queue, _slots, _cfg_tx) =
3953 make_queue_with_sender(1, 16, 100, Some(0.5), Some(estimator));
3954
3955 let (req1, rx1) = make_request("req-1", 100);
3956 queue.enqueue(req1).await;
3957 let _ = rx1.await.unwrap().unwrap();
3958
3959 let (req2, mut rx2) = make_request("req-2", 100);
3960 queue.enqueue(req2).await;
3961 assert_eq!(queue.pending_count(), 1);
3962
3963 tokio::time::advance(Duration::from_secs(6)).await;
3964 queue.update().await;
3965
3966 let scheduled = rx2
3967 .try_recv()
3968 .expect("queued request should have been scheduled");
3969 let response = scheduled.expect("scheduling returned error");
3970 assert_eq!(response.best_worker.worker_id, 0);
3971 assert_eq!(queue.pending_count(), 0);
3972 }
3973
3974 #[tokio::test(flavor = "multi_thread")]
3975 async fn test_overloaded_provider_filters_at_admission() {
3976 let overloaded_worker_provider: OverloadedWorkerProvider =
3977 Arc::new(|| Some(HashSet::from([0])));
3978 let (queue, _slots) =
3979 make_queue_with_overload_provider(1, 16, 256, overloaded_worker_provider);
3980
3981 let (req, rx) = make_request("overloaded", 256);
3982 queue.enqueue(req).await;
3983
3984 let resp = rx.await.expect("oneshot dropped");
3985 assert!(matches!(
3986 resp,
3987 Err(KvSchedulerError::AllEligibleWorkersOverloaded)
3988 ));
3989 }
3990
3991 #[tokio::test(flavor = "multi_thread")]
3994 async fn test_register_workers_lazy_epp_path() {
3995 let block_size = 16;
3996 let isl = 512;
3997
3998 let (queue, slots, cfg_tx) = make_queue_with_sender(0, block_size, isl, None, None);
4000
4001 let (req_fail, rx_fail) = make_request("before-register", isl);
4003 queue.enqueue(req_fail).await;
4004 let resp = rx_fail.await.expect("oneshot dropped");
4005 assert!(
4006 matches!(
4007 resp,
4008 Err(crate::scheduling::types::KvSchedulerError::NoEndpoints)
4009 ),
4010 "expected NoEndpoints before register_workers, got {resp:?}"
4011 );
4012
4013 slots.upsert_worker(WorkerDpRange::new(100, 0, 1)).unwrap();
4015 slots.upsert_worker(WorkerDpRange::new(200, 0, 1)).unwrap();
4016
4017 let mut configs = HashMap::new();
4019 for &id in &[100_u64, 200_u64] {
4020 configs.insert(
4021 id,
4022 SimpleWorkerConfig {
4023 max_num_batched_tokens: Some(isl as u64),
4024 ..Default::default()
4025 },
4026 );
4027 }
4028 cfg_tx.send(configs).unwrap();
4029
4030 let (req_ok, rx_ok) = make_request("after-register", isl);
4032 queue.enqueue(req_ok).await;
4033 let resp = rx_ok
4034 .await
4035 .expect("oneshot dropped")
4036 .expect("scheduling failed");
4037 assert!(
4038 resp.best_worker.worker_id == 100 || resp.best_worker.worker_id == 200,
4039 "expected worker 100 or 200, got {}",
4040 resp.best_worker.worker_id
4041 );
4042
4043 slots
4045 .mark_prefill_completed(&"after-register".to_string(), decay_now())
4046 .unwrap();
4047 slots
4048 .free(&"after-register".to_string(), decay_now())
4049 .unwrap();
4050 }
4051
4052 #[tokio::test(flavor = "multi_thread")]
4054 async fn test_register_workers_additive() {
4055 let block_size = 16;
4056 let isl = 256;
4057
4058 let (queue, slots, cfg_tx) = make_queue_with_sender(0, block_size, isl, None, None);
4059
4060 slots.upsert_worker(WorkerDpRange::new(10, 0, 1)).unwrap();
4062
4063 let mut configs = HashMap::new();
4064 configs.insert(
4065 10_u64,
4066 SimpleWorkerConfig {
4067 max_num_batched_tokens: Some(isl as u64),
4068 ..Default::default()
4069 },
4070 );
4071 cfg_tx.send(configs.clone()).unwrap();
4072
4073 slots.upsert_worker(WorkerDpRange::new(20, 0, 1)).unwrap();
4075
4076 configs.insert(
4077 20_u64,
4078 SimpleWorkerConfig {
4079 max_num_batched_tokens: Some(isl as u64),
4080 ..Default::default()
4081 },
4082 );
4083 cfg_tx.send(configs).unwrap();
4084
4085 let mut seen = std::collections::HashSet::new();
4087 for i in 0..20 {
4088 let req_id = format!("add-{i}");
4089 let (req, rx) = make_request(&req_id, isl);
4090 queue.enqueue(req).await;
4091 let resp = rx
4092 .await
4093 .expect("oneshot dropped")
4094 .expect("scheduling failed");
4095 seen.insert(resp.best_worker.worker_id);
4096 slots.mark_prefill_completed(&req_id, decay_now()).unwrap();
4097 slots.free(&req_id, decay_now()).unwrap();
4098 }
4099
4100 assert!(
4101 seen.contains(&10) && seen.contains(&20),
4102 "both workers should be reachable after additive registration, saw: {seen:?}"
4103 );
4104 }
4105
4106 #[tokio::test(flavor = "multi_thread")]
4107 async fn allowed_worker_request_joins_backlog_and_dispatches_within_allow_list() {
4108 let block_size = 16;
4109 let isl = 256;
4110 let (queue, slots) = make_queue(2, block_size, isl, Some(0.0));
4111
4112 let (active_a, active_a_rx) = make_request("active-a", isl);
4113 queue.enqueue(active_a).await;
4114 let active_a_worker = active_a_rx.await.unwrap().unwrap().best_worker.worker_id;
4115
4116 let (active_b, active_b_rx) = make_request("active-b", isl);
4117 queue.enqueue(active_b).await;
4118 active_b_rx.await.unwrap().unwrap();
4119
4120 let (backlog_head, backlog_head_rx) = make_request("backlog-head", isl);
4121 queue.enqueue(backlog_head).await;
4122 assert_eq!(queue.pending_count(), 1);
4123
4124 slots
4125 .mark_prefill_completed(&"active-a".to_string(), decay_now())
4126 .unwrap();
4127 slots.free(&"active-a".to_string(), decay_now()).unwrap();
4128
4129 let (mut allowed, mut allowed_rx) = make_request("allowed", isl);
4130 allowed.allowed_worker_ids = Some(HashSet::from([active_a_worker]));
4131 queue.enqueue(allowed).await;
4132 assert_eq!(
4133 queue.pending_count(),
4134 2,
4135 "allow-list request must not bypass the existing class backlog"
4136 );
4137 assert!(allowed_rx.try_recv().is_err());
4138
4139 queue.update().await;
4140 let backlog_head_worker = backlog_head_rx
4141 .await
4142 .unwrap()
4143 .unwrap()
4144 .best_worker
4145 .worker_id;
4146 assert!(allowed_rx.try_recv().is_err());
4147
4148 slots
4149 .mark_prefill_completed(&"backlog-head".to_string(), decay_now())
4150 .unwrap();
4151 slots
4152 .free(&"backlog-head".to_string(), decay_now())
4153 .unwrap();
4154 queue.update().await;
4155
4156 let allowed_worker = allowed_rx.await.unwrap().unwrap().best_worker.worker_id;
4157 assert_eq!(allowed_worker, active_a_worker);
4158
4159 for request_id in ["active-b", "allowed"] {
4160 slots
4161 .mark_prefill_completed(&request_id.to_string(), decay_now())
4162 .unwrap();
4163 slots.free(&request_id.to_string(), decay_now()).unwrap();
4164 }
4165 assert_eq!(backlog_head_worker, active_a_worker);
4166 }
4167
4168 #[tokio::test(flavor = "multi_thread")]
4169 async fn test_pinned_worker_conflict_with_allowed_ids_fails_early() {
4170 let (queue, _slots) = make_queue(1, 16, 256, Some(0.0));
4171 let (mut req, rx) = make_request("conflict", 256);
4172 req.pinned_worker = Some(WorkerWithDpRank::new(0, 0));
4173 req.allowed_worker_ids = Some(HashSet::from([1]));
4174
4175 queue.enqueue(req).await;
4176
4177 let resp = rx.await.expect("oneshot dropped");
4178 assert!(matches!(
4179 resp,
4180 Err(KvSchedulerError::PinnedWorkerNotAllowed { worker_id: 0 })
4181 ));
4182 }
4183
4184 #[tokio::test(flavor = "multi_thread")]
4185 async fn test_disallowed_worker_ids_fail_without_queueing() {
4186 let (queue, _slots) = make_queue(1, 16, 256, Some(0.0));
4187 let (mut req, rx) = make_request("disallowed", 256);
4188 req.allowed_worker_ids = Some(HashSet::from([999]));
4189
4190 queue.enqueue(req).await;
4191
4192 let resp = rx.await.expect("oneshot dropped");
4193 assert!(matches!(resp, Err(KvSchedulerError::NoEndpoints)));
4194 assert_eq!(queue.pending_count(), 0);
4195 }
4196
4197 #[tokio::test(flavor = "multi_thread")]
4198 async fn test_incompatible_required_taints_fail_without_queueing() {
4199 let (queue, _slots, cfg_tx) = make_queue_with_sender(1, 16, 256, Some(0.0), None);
4200 let mut configs = HashMap::new();
4201 configs.insert(
4202 0_u64,
4203 SimpleWorkerConfig {
4204 max_num_batched_tokens: Some(256),
4205 taints: HashSet::from(["mdc-a".to_string()]),
4206 ..Default::default()
4207 },
4208 );
4209 cfg_tx.send(configs).unwrap();
4210
4211 let (mut req, rx) = make_request("tainted", 256);
4212 req.routing_constraints = crate::protocols::RoutingConstraints {
4213 required_taints: HashSet::from(["mdc-b".to_string()]),
4214 preferred_taints: HashMap::new(),
4215 };
4216
4217 queue.enqueue(req).await;
4218
4219 let resp = rx.await.expect("oneshot dropped");
4220 assert!(matches!(resp, Err(KvSchedulerError::NoEndpoints)));
4221 assert_eq!(queue.pending_count(), 0);
4222 }
4223
4224 #[tokio::test(flavor = "multi_thread")]
4225 async fn test_blocked_pinned_lane_does_not_block_other_worker() {
4226 let (queue, slots) = make_queue(2, 16, 256, Some(0.0));
4227
4228 let (mut first, first_rx) = make_request("pinned-1", 256);
4229 first.pinned_worker = Some(WorkerWithDpRank::new(1, 0));
4230 queue.enqueue(first).await;
4231 let first_resp = first_rx.await.unwrap().unwrap();
4232 assert_eq!(first_resp.best_worker, WorkerWithDpRank::new(1, 0));
4233
4234 let (mut second, mut second_rx) = make_request("pinned-2", 256);
4235 second.pinned_worker = Some(WorkerWithDpRank::new(1, 0));
4236 queue.enqueue(second).await;
4237 assert_eq!(queue.pending_count(), 1);
4238 assert!(
4239 second_rx.try_recv().is_err(),
4240 "request should remain queued"
4241 );
4242
4243 let (mut other_worker, mut other_worker_rx) = make_request("pinned-0", 256);
4244 other_worker.pinned_worker = Some(WorkerWithDpRank::new(0, 0));
4245 queue.enqueue(other_worker).await;
4246 assert_eq!(queue.pending_count(), 2);
4247
4248 queue.update().await;
4249
4250 assert_eq!(queue.pending_count(), 1);
4251 let other_worker_resp = other_worker_rx
4252 .try_recv()
4253 .expect("other worker request should have been scheduled")
4254 .expect("scheduling returned error");
4255 assert_eq!(other_worker_resp.best_worker, WorkerWithDpRank::new(0, 0));
4256 assert!(
4257 second_rx.try_recv().is_err(),
4258 "pinned request should still be queued"
4259 );
4260
4261 slots
4262 .mark_prefill_completed(&"pinned-1".to_string(), decay_now())
4263 .unwrap();
4264 slots.free(&"pinned-1".to_string(), decay_now()).unwrap();
4265 queue.update_worker(WorkerWithDpRank::new(1, 0)).await;
4266
4267 let second_resp = second_rx
4268 .try_recv()
4269 .expect("pinned request should have been scheduled");
4270 let second_resp = second_resp.expect("scheduling returned error");
4271 assert_eq!(second_resp.best_worker, WorkerWithDpRank::new(1, 0));
4272 assert_eq!(queue.pending_count(), 0);
4273 }
4274
4275 #[tokio::test(flavor = "multi_thread")]
4276 async fn test_queue_prefill_busy_check_ignores_untracked_prefill_tokens() {
4277 let (queue, slots) = make_queue(1, 16, 256, Some(0.0));
4278
4279 let (mut req1, rx1) = make_request("req-1", 256);
4280 req1.track_prefill_tokens = false;
4281 queue.enqueue(req1).await;
4282 let _resp1 = rx1.await.unwrap().unwrap();
4283 assert_eq!(
4284 slots
4285 .active_tokens(decay_now())
4286 .get(&WorkerWithDpRank::new(0, 0))
4287 .copied(),
4288 Some(0)
4289 );
4290
4291 let (req2, rx2) = make_request("req-2", 256);
4292 queue.enqueue(req2).await;
4293 let _resp2 = rx2.await.unwrap().unwrap();
4294 assert_eq!(queue.pending_count(), 0);
4295
4296 let _ = slots.mark_prefill_completed(&"req-1".to_string(), decay_now());
4297 let _ = slots.free(&"req-1".to_string(), decay_now());
4298 let _ = slots.mark_prefill_completed(&"req-2".to_string(), decay_now());
4299 let _ = slots.free(&"req-2".to_string(), decay_now());
4300 }
4301
4302 #[tokio::test(flavor = "current_thread", start_paused = true)]
4303 async fn update_refresh_can_change_selected_worker_after_queue_wait() {
4304 let block_size = 16u32;
4305 let isl = 64usize;
4306 let refresher = Arc::new(CountingRefresher {
4307 calls: AtomicUsize::new(0),
4308 response: RefreshedOverlap {
4309 tier_overlap_blocks: Default::default(),
4310 effective_overlap_blocks: HashMap::from([
4311 (WorkerWithDpRank::new(0, 0), 1.0),
4312 (WorkerWithDpRank::new(1, 0), 9.0),
4313 ]),
4314 effective_cached_tokens: HashMap::from([
4315 (WorkerWithDpRank::new(0, 0), 16),
4316 (WorkerWithDpRank::new(1, 0), 144),
4317 ]),
4318 },
4319 });
4320 let (queue, slots) =
4321 make_queue_with_refresher(2, block_size, isl, Some(0.0), refresher.clone());
4322
4323 let (mut req1, rx1) = make_request("req-1", isl);
4324 req1.overlap
4325 .effective_overlap_blocks
4326 .insert(WorkerWithDpRank::new(0, 0), 3.0);
4327 req1.overlap
4328 .effective_cached_tokens
4329 .insert(WorkerWithDpRank::new(0, 0), 48);
4330 queue.enqueue(req1).await;
4331 let resp1 = rx1.await.expect("rx1 dropped").expect("req-1 failed");
4332 assert_eq!(resp1.best_worker, WorkerWithDpRank::new(0, 0));
4333
4334 let (mut req2, rx2) = make_request("req-2", isl);
4335 req2.overlap
4336 .effective_overlap_blocks
4337 .insert(WorkerWithDpRank::new(1, 0), 3.0);
4338 req2.overlap
4339 .effective_cached_tokens
4340 .insert(WorkerWithDpRank::new(1, 0), 48);
4341 queue.enqueue(req2).await;
4342 let resp2 = rx2.await.expect("rx2 dropped").expect("req-2 failed");
4343 assert_eq!(resp2.best_worker, WorkerWithDpRank::new(1, 0));
4344
4345 let (mut req3, rx3) = make_request("req-3", isl);
4346 req3.overlap
4347 .effective_overlap_blocks
4348 .insert(WorkerWithDpRank::new(0, 0), 8.0);
4349 req3.overlap
4350 .effective_overlap_blocks
4351 .insert(WorkerWithDpRank::new(1, 0), 2.0);
4352 req3.overlap
4353 .effective_cached_tokens
4354 .insert(WorkerWithDpRank::new(0, 0), 128);
4355 req3.overlap
4356 .effective_cached_tokens
4357 .insert(WorkerWithDpRank::new(1, 0), 32);
4358 queue
4359 .enqueue_with_block_hashes(req3, Some(vec![LocalBlockHash(42)]))
4360 .await;
4361 assert_eq!(queue.pending_count(), 1);
4362 assert_eq!(refresher.calls.load(Ordering::Relaxed), 0);
4363
4364 tokio::time::advance(Duration::from_secs(11)).await;
4365
4366 slots.free(&"req-1".to_string(), decay_now()).unwrap();
4367 slots.free(&"req-2".to_string(), decay_now()).unwrap();
4368 queue.update().await;
4369
4370 let resp3 = rx3.await.expect("rx3 dropped").expect("req-3 failed");
4371 assert_eq!(refresher.calls.load(Ordering::Relaxed), 1);
4372 assert_eq!(resp3.best_worker, WorkerWithDpRank::new(1, 0));
4373 assert_eq!(resp3.effective_overlap_blocks, 9.0);
4374 assert_eq!(resp3.cached_tokens, 144);
4375 assert_eq!(queue.pending_count(), 0);
4376 }
4377
4378 #[tokio::test(flavor = "current_thread", start_paused = true)]
4379 async fn selected_request_dispatches_after_refresh_if_worker_becomes_busy() {
4380 let block_size = 16u32;
4381 let isl = 64usize;
4382 let worker = WorkerWithDpRank::new(0, 0);
4383 let refresher = Arc::new(BlockingRefresher::new(RefreshedOverlap {
4384 tier_overlap_blocks: Default::default(),
4385 effective_overlap_blocks: HashMap::from([(worker, 7.0)]),
4386 effective_cached_tokens: HashMap::from([(worker, 56)]),
4387 }));
4388 let (queue, slots) = make_queue_with_blocking_refresher(
4389 1,
4390 block_size,
4391 isl,
4392 Some(0.0),
4393 refresher.clone(),
4394 ADMISSION_CHANNEL_CAPACITY,
4395 );
4396
4397 let (req1, rx1) = make_request("req-1", isl);
4398 queue.enqueue(req1).await;
4399 let _ = rx1.await.expect("rx1 dropped").expect("req-1 failed");
4400
4401 let (mut req2, rx2) = make_request("req-2", isl);
4402 req2.overlap
4403 .effective_overlap_blocks
4404 .insert(WorkerWithDpRank::new(0, 0), 4.0);
4405 req2.overlap
4406 .effective_cached_tokens
4407 .insert(WorkerWithDpRank::new(0, 0), 64);
4408 queue
4409 .enqueue_with_block_hashes(req2, Some(vec![LocalBlockHash(42)]))
4410 .await;
4411 assert_eq!(queue.pending_count(), 1);
4412 assert_eq!(
4413 queue.class_queue_stats(0).unwrap().pending_cached_tokens,
4414 64
4415 );
4416
4417 slots
4418 .mark_prefill_completed(&"req-1".to_string(), decay_now())
4419 .unwrap();
4420 slots.free(&"req-1".to_string(), decay_now()).unwrap();
4421
4422 tokio::time::advance(Duration::from_secs(11)).await;
4423
4424 let update = {
4425 let queue = Arc::clone(&queue);
4426 tokio::spawn(async move {
4427 queue.update().await;
4428 })
4429 };
4430 refresher.wait_for_calls(1).await;
4431 assert_eq!(
4432 queue.pending_count(),
4433 0,
4434 "DRR-selected request must be removed before refresh"
4435 );
4436 assert_eq!(
4437 queue.class_queue_stats(0).unwrap().pending_cached_tokens,
4438 0,
4439 "queue counters must reflect the irrevocable dequeue"
4440 );
4441
4442 slots
4443 .add_request(
4444 SequenceRequest {
4445 request_id: "occupy-during-refresh".to_string(),
4446 token_sequence: None,
4447 track_prefill_tokens: true,
4448 expected_output_tokens: None,
4449 prefill_load_hint: Some(PrefillLoadHint {
4450 initial_effective_prefill_tokens: isl,
4451 expected_prefill_duration: None,
4452 }),
4453 worker,
4454 lora_name: None,
4455 },
4456 decay_now(),
4457 )
4458 .unwrap();
4459
4460 refresher.release_one();
4461 update.await.unwrap();
4462
4463 let resp2 = rx2.await.expect("rx2 dropped").expect("req-2 failed");
4464 assert_eq!(refresher.calls.load(Ordering::Relaxed), 1);
4465 assert_eq!(resp2.best_worker, worker);
4466 assert_eq!(resp2.effective_overlap_blocks, 7.0);
4467 assert_eq!(resp2.cached_tokens, 56);
4468 assert_eq!(queue.pending_count(), 0);
4469
4470 for request_id in ["occupy-during-refresh", "req-2"] {
4471 slots
4472 .mark_prefill_completed(&request_id.to_string(), decay_now())
4473 .unwrap();
4474 slots.free(&request_id.to_string(), decay_now()).unwrap();
4475 }
4476 }
4477
4478 #[tokio::test(flavor = "current_thread", start_paused = true)]
4479 async fn cancelled_enqueue_wait_keeps_cleanup_behind_command() {
4480 let block_size = 16u32;
4481 let isl = 64usize;
4482 let refresher = Arc::new(BlockingRefresher::new(RefreshedOverlap::default()));
4483 let (queue, slots) =
4484 make_queue_with_blocking_refresher(1, block_size, isl, Some(0.0), refresher.clone(), 1);
4485
4486 let (active, active_rx) = make_request("active", isl);
4487 queue.enqueue(active).await;
4488 active_rx.await.unwrap().unwrap();
4489
4490 let (queued, queued_rx) = make_request("queued", isl);
4491 queue
4492 .enqueue_with_block_hashes(queued, Some(vec![LocalBlockHash(42)]))
4493 .await;
4494 slots.free(&"active".to_owned(), decay_now()).unwrap();
4495 tokio::time::advance(Duration::from_secs(11)).await;
4496
4497 let update = {
4498 let queue = Arc::clone(&queue);
4499 tokio::spawn(async move { queue.update().await })
4500 };
4501 refresher.wait_for_calls(1).await;
4502
4503 let (cancelled, cancelled_rx) = make_request("cancelled", isl);
4504 let lease = queue
4505 .new_request_lifecycle_lease(Some("cancelled"))
4506 .unwrap();
4507 let enqueue = {
4508 let queue = Arc::clone(&queue);
4509 tokio::spawn(async move {
4510 queue
4511 .enqueue_with_block_hashes_and_lease(cancelled, None, Some(lease))
4512 .await
4513 })
4514 };
4515 tokio::task::yield_now().await;
4516 assert_eq!(queue.admission_tx.capacity(), 0);
4517 drop(cancelled_rx);
4518 enqueue.abort();
4519 assert!(enqueue.await.unwrap_err().is_cancelled());
4520
4521 refresher.release_one();
4524 update.await.unwrap();
4525 queue.update().await;
4526 assert_eq!(queue.pending_count(), 0);
4527
4528 queued_rx.await.unwrap().unwrap();
4529 slots.free(&"queued".to_owned(), decay_now()).unwrap();
4530 slots.assert_completely_drained(decay_now());
4531 }
4532
4533 #[tokio::test(flavor = "current_thread", start_paused = true)]
4534 async fn continuation_drain_does_not_self_send_into_saturated_actor_channel() {
4535 let block_size = 16u32;
4536 let isl = 64usize;
4537 let refresher = Arc::new(BlockingRefresher::new(RefreshedOverlap::default()));
4538 let (queue, slots) =
4539 make_queue_with_blocking_refresher(1, block_size, isl, Some(0.0), refresher.clone(), 1);
4540
4541 let (active, active_rx) = make_request("active", isl);
4542 queue.enqueue(active).await;
4543 active_rx.await.unwrap().unwrap();
4544
4545 let (queued, queued_rx) = make_request("queued", isl);
4546 queue
4547 .enqueue_with_block_hashes(queued, Some(vec![LocalBlockHash(42)]))
4548 .await;
4549 slots
4550 .mark_prefill_completed(&"active".to_string(), decay_now())
4551 .unwrap();
4552 slots.free(&"active".to_string(), decay_now()).unwrap();
4553 tokio::time::advance(Duration::from_secs(11)).await;
4554
4555 let update = {
4556 let queue = Arc::clone(&queue);
4557 tokio::spawn(async move { queue.update().await })
4558 };
4559 refresher.wait_for_calls(1).await;
4560
4561 let (following, following_rx) = make_request("following", isl);
4562 let enqueue = {
4563 let queue = Arc::clone(&queue);
4564 tokio::spawn(async move { queue.enqueue(following).await })
4565 };
4566 tokio::task::yield_now().await;
4567 assert_eq!(
4568 queue.admission_tx.capacity(),
4569 0,
4570 "test must saturate the actor command channel"
4571 );
4572
4573 refresher.release_one();
4574 tokio::time::timeout(Duration::from_secs(1), update)
4575 .await
4576 .expect("update deadlocked with a full actor command channel")
4577 .unwrap();
4578 queued_rx.await.unwrap().unwrap();
4579
4580 slots
4581 .mark_prefill_completed(&"queued".to_string(), decay_now())
4582 .unwrap();
4583 slots.free(&"queued".to_string(), decay_now()).unwrap();
4584 queue.update().await;
4585 following_rx.await.unwrap().unwrap();
4586 enqueue.await.unwrap();
4587 }
4588}