1use std::collections::HashMap;
7use std::sync::Arc;
8
9use anyhow::{Result, anyhow, ensure};
10use uuid::Uuid;
11
12use crate::engine::common::perf_model::PerfModel;
13use crate::engine::common::protocols::{
14 DirectRequest, EngineType, KvTransferTimingMode, MockEngineArgs,
15 PreemptionMode as CorePreemptionMode, SglangArgs, WorkerType as CoreWorkerType,
16};
17use crate::engine::generalized::{CommandContext, RankEngine, RankIdentity, RankPass};
18use crate::engine::{
19 Admission, Backend, Command, CommandEffects, CommandResult, EngineConfig, ForwardPassMetrics,
20 HandoffId, HostOffloadObserver, LifecycleEvent, Metrics, Output, PassCompletionEffects,
21 PassStartEffects, PendingPass, PreemptionMode, Request, TimingModel, TransferTimingMode,
22 WorkerType,
23};
24
25use super::{
26 EngineCore, EnginePassResult, KvEventVisibility, MockerMetrics,
27 SchedulerCommand as CoreCommand, SchedulerCommandEffects as CoreCommandEffects,
28 SchedulerCommandResult as CoreCommandResult, SchedulerLifecycleEvent as CoreLifecycle,
29 SglangCore, VllmCore,
30};
31
32pub fn engine_seed_offset(identity: RankIdentity) -> Result<u64> {
33 identity
34 .worker_id
35 .checked_mul(u64::from(identity.dp_size.get()))
36 .and_then(|base| base.checked_add(u64::from(identity.dp_rank)))
37 .ok_or_else(|| anyhow!("native mock-engine seed offset overflow"))
38}
39
40pub struct SchedulerRank {
42 core: EngineCore,
43 handoff_requests: HashMap<HandoffId, Uuid>,
44}
45
46impl SchedulerRank {
47 pub(crate) fn set_host_offload_observer(&mut self, observer: Arc<dyn HostOffloadObserver>) {
48 self.core.set_host_offload_observer(observer);
49 }
50
51 pub fn new_with_timing_model(
52 identity: RankIdentity,
53 config: &EngineConfig,
54 timing: Arc<dyn TimingModel>,
55 seed_offset: u64,
56 ) -> Result<Self> {
57 config.validate()?;
58 ensure!(
59 config.native_host_offload.is_none() || identity.dp_size.get() == 1,
60 "native_host_offload supports only dp_size=1 in the initial implementation"
61 );
62 let args = core_args(config, timing);
63 let capture_kv_events = config.emit_kv_events;
64 let core = match config.backend {
65 Backend::Vllm | Backend::Trtllm => EngineCore::Vllm(VllmCore::new_with_worker_rank(
66 args,
67 identity.worker_id,
68 identity.dp_rank,
69 seed_offset,
70 capture_kv_events,
71 )),
72 Backend::Sglang => EngineCore::Sglang(SglangCore::new_with_worker_rank(
73 args,
74 identity.worker_id,
75 identity.dp_rank,
76 seed_offset,
77 capture_kv_events,
78 )),
79 };
80 Ok(Self {
81 core,
82 handoff_requests: HashMap::new(),
83 })
84 }
85
86 fn core_command(command: Command) -> CoreCommand {
87 match command {
88 Command::Submit(request) => CoreCommand::Submit(core_request(request)),
89 Command::CancelRequest { request_id, .. } => CoreCommand::CancelRequest { request_id },
90 Command::SubmitHandoffPrefill {
91 handoff_id,
92 request,
93 } => CoreCommand::SubmitHandoffPrefill {
94 handoff_id,
95 request: core_request(request),
96 },
97 Command::ReserveDestination {
98 handoff_id,
99 request,
100 } => CoreCommand::ReserveDestination {
101 handoff_id,
102 request: core_request(request),
103 },
104 Command::ActivateDestination { handoff_id } => {
105 CoreCommand::ActivateDestination { handoff_id }
106 }
107 Command::ReleaseSource { handoff_id } => CoreCommand::ReleaseSource { handoff_id },
108 Command::CancelSource { handoff_id } => CoreCommand::CancelSource { handoff_id },
109 Command::CancelDestination { handoff_id } => {
110 CoreCommand::CancelDestination { handoff_id }
111 }
112 }
113 }
114
115 fn metrics(&self) -> Metrics {
116 let metrics = match &self.core {
117 EngineCore::Vllm(core) => core.mocker_metrics(),
118 EngineCore::Sglang(core) => core.mocker_metrics(),
119 };
120 map_metrics(metrics)
121 }
122}
123
124impl RankEngine for SchedulerRank {
125 type Config = EngineConfig;
126 type Command = Command;
127 type CommandEffects = CommandEffects;
128 type PassStartEffects = PassStartEffects;
129 type PendingPass = PendingPass;
130 type PassCompletionEffects = PassCompletionEffects;
131 type InternalEffects = PassStartEffects;
132
133 fn new(identity: RankIdentity, config: &Self::Config) -> Result<Self> {
134 let timing = config.built_in_timing_model()?;
135 let seed_offset = engine_seed_offset(identity)?;
136 Self::new_with_timing_model(identity, config, timing, seed_offset)
137 }
138
139 fn apply_command_effects(
140 &mut self,
141 command: Self::Command,
142 context: CommandContext,
143 pending_pass: Option<&mut Self::PendingPass>,
144 ) -> Result<Self::CommandEffects> {
145 let pending_suppression = pending_output_suppression(&command, &self.handoff_requests);
146 let handoff_update = handoff_tracking_update(&command);
147 let core_command = Self::core_command(command);
148 let mut effects = self.core.apply_command_effects_at(
149 core_command,
150 context.allow_immediate_admission(),
151 context.now_ms,
152 )?;
153 if !context.pass_in_flight {
159 effects.kv_events.extend(self.core.drain_kv_events());
160 }
161 let suppressed_pending_output = if let (Some((request_id, discard_on_noop)), Some(pending)) =
162 (pending_suppression, pending_pass)
163 && (effects.result != CoreCommandResult::Noop || discard_on_noop)
164 {
165 let before = pending.effects.outputs.len();
166 pending
167 .effects
168 .outputs
169 .retain(|output| output.request_id != request_id);
170 pending
171 .effects
172 .lifecycle_events
173 .retain(|event| match *event {
174 LifecycleEvent::SourceHeld { request_id: id, .. }
175 | LifecycleEvent::DestinationReserved { request_id: id, .. } => {
176 id != request_id
177 }
178 });
179 before != pending.effects.outputs.len()
180 } else {
181 false
182 };
183 if effects.result != CoreCommandResult::Noop || suppressed_pending_output {
184 self.apply_handoff_tracking_update(handoff_update);
185 }
186 for request_id in &effects.retired_requests {
187 self.handoff_requests
188 .retain(|_, tracked_request| tracked_request != request_id);
189 }
190 map_command_effects(effects, self.metrics(), suppressed_pending_output)
191 }
192
193 fn is_ready(&self) -> bool {
194 self.core.is_ready()
195 }
196
197 fn waiting_for_external_command(&self) -> bool {
198 self.core.waiting_for_external_command()
199 }
200
201 fn execute_pass(
202 &mut self,
203 now_ms: f64,
204 ) -> Result<RankPass<Self::PassStartEffects, Self::PendingPass>> {
205 let pass = self.core.try_execute_hidden_pass(now_ms)?;
206 let end_ms = pass.end_ms;
207 let (same_timestamp_retry, start_effects, completion_effects) = split_pass(pass)?;
208 Ok(RankPass {
209 end_ms,
210 same_timestamp_retry,
211 start_effects,
212 pending: PendingPass {
213 started_at_ms: now_ms,
214 effects: completion_effects,
215 },
216 })
217 }
218
219 fn complete_pass(
220 &mut self,
221 mut pending: Self::PendingPass,
222 end_ms: f64,
223 ) -> Result<Self::PassCompletionEffects> {
224 self.core.complete_engine_boundary(end_ms);
225 pending.effects.lifecycle_events.extend(
231 self.core
232 .retry_pending_destinations()
233 .into_iter()
234 .map(map_lifecycle),
235 );
236 let completion_kv_events = self.core.drain_kv_events();
237 pending.effects.kv_events.extend(completion_kv_events);
238 let sglang_cache_hit_tokens = pending.effects.metrics.sglang_cache_hit_tokens;
243 let sglang_cache_total_tokens = pending.effects.metrics.sglang_cache_total_tokens;
244 pending.effects.metrics = self.metrics();
245 pending.effects.metrics.sglang_cache_hit_tokens = sglang_cache_hit_tokens;
246 pending.effects.metrics.sglang_cache_total_tokens = sglang_cache_total_tokens;
247 pending.effects.forward_pass_metrics.duration_ms =
248 (end_ms - pending.started_at_ms).max(0.0);
249 for output in &pending.effects.outputs {
250 if output.completed {
251 self.handoff_requests
252 .retain(|_, request_id| *request_id != output.request_id);
253 }
254 }
255 Ok(pending.effects)
256 }
257
258 fn complete_idle_group_pass(
259 &mut self,
260 started_at_ms: f64,
261 end_ms: f64,
262 ) -> Result<Option<Self::PassCompletionEffects>> {
263 self.core.complete_engine_boundary(end_ms);
264 let lifecycle_events = self
265 .core
266 .retry_pending_destinations()
267 .into_iter()
268 .map(map_lifecycle)
269 .collect::<Vec<_>>();
270 let kv_events = self.core.drain_kv_events();
271 Ok(Some(PassCompletionEffects {
272 lifecycle_events,
273 kv_events,
274 metrics: self.metrics(),
275 forward_pass_metrics: ForwardPassMetrics {
276 duration_ms: (end_ms - started_at_ms).max(0.0),
277 ..Default::default()
278 },
279 ..PassCompletionEffects::default()
280 }))
281 }
282
283 fn next_internal_deadline_ms(&self) -> Option<f64> {
284 self.core.next_internal_deadline_ms()
285 }
286
287 fn process_internal_work(
288 &mut self,
289 now_ms: f64,
290 pass_in_flight: bool,
291 ) -> Result<Self::InternalEffects> {
292 if pass_in_flight {
293 return Ok(PassStartEffects::default());
294 }
295 self.core.process_internal_work(now_ms);
296 Ok(PassStartEffects {
297 kv_events: self.core.drain_kv_events(),
298 ..PassStartEffects::default()
299 })
300 }
301
302 fn is_drained(&self) -> bool {
303 self.core.is_drained()
304 }
305}
306
307#[derive(Clone, Copy)]
308enum HandoffTrackingUpdate {
309 None,
310 Insert(HandoffId, Uuid),
311 RemoveHandoff(HandoffId),
312 RemoveRequest(Uuid),
313}
314
315impl SchedulerRank {
316 fn apply_handoff_tracking_update(&mut self, update: HandoffTrackingUpdate) {
317 match update {
318 HandoffTrackingUpdate::None => {}
319 HandoffTrackingUpdate::Insert(handoff_id, request_id) => {
320 self.handoff_requests.insert(handoff_id, request_id);
321 }
322 HandoffTrackingUpdate::RemoveHandoff(handoff_id) => {
323 self.handoff_requests.remove(&handoff_id);
324 }
325 HandoffTrackingUpdate::RemoveRequest(request_id) => self
326 .handoff_requests
327 .retain(|_, tracked_request| *tracked_request != request_id),
328 }
329 }
330}
331
332fn handoff_tracking_update(command: &Command) -> HandoffTrackingUpdate {
333 match command {
334 Command::SubmitHandoffPrefill {
335 handoff_id,
336 request,
337 }
338 | Command::ReserveDestination {
339 handoff_id,
340 request,
341 } => HandoffTrackingUpdate::Insert(*handoff_id, request.request_id),
342 Command::ReleaseSource { handoff_id }
343 | Command::CancelSource { handoff_id }
344 | Command::CancelDestination { handoff_id } => {
345 HandoffTrackingUpdate::RemoveHandoff(*handoff_id)
346 }
347 Command::CancelRequest { request_id, .. } => {
348 HandoffTrackingUpdate::RemoveRequest(*request_id)
349 }
350 Command::Submit(_) | Command::ActivateDestination { .. } => HandoffTrackingUpdate::None,
351 }
352}
353
354fn core_args(config: &EngineConfig, timing: Arc<dyn TimingModel>) -> MockEngineArgs {
355 MockEngineArgs {
356 engine_type: match config.backend {
357 Backend::Vllm => EngineType::Vllm,
358 Backend::Sglang => EngineType::Sglang,
359 Backend::Trtllm => EngineType::Trtllm,
360 },
361 num_gpu_blocks: config.num_gpu_blocks,
362 block_size: config.block_size,
363 max_model_len: config.max_model_len,
364 max_num_seqs: Some(config.max_num_seqs),
365 max_num_batched_tokens: Some(config.max_num_batched_tokens),
366 enable_prefix_caching: config.enable_prefix_caching,
367 enable_chunked_prefill: config.enable_chunked_prefill,
368 speedup_ratio: config.speedup_ratio,
369 decode_speedup_ratio: config.decode_speedup_ratio,
370 worker_type: match config.worker_type {
371 WorkerType::Aggregated => CoreWorkerType::Aggregated,
372 WorkerType::Prefill => CoreWorkerType::Prefill,
373 WorkerType::Decode => CoreWorkerType::Decode,
374 },
375 perf_model: Arc::new(PerfModel::External { timing }),
376 aic_nextn: config.aic_nextn,
377 aic_nextn_accept_rates: config.aic_nextn_accept_rates.clone(),
378 aic_mtp_seed: config.aic_mtp_seed,
379 kv_transfer_bytes_per_token: config.kv_transfer_bytes_per_token,
380 kv_cache_bytes_per_token: config.kv_cache_bytes_per_token,
381 native_host_offload: config.native_host_offload,
382 kv_transfer_bandwidth: config.kv_transfer_bandwidth,
383 kv_transfer_timing_mode: match config.kv_transfer_timing_mode {
384 TransferTimingMode::FullPrompt => KvTransferTimingMode::FullPrompt,
385 TransferTimingMode::DestinationMissing => KvTransferTimingMode::DestinationMissing,
386 },
387 preemption_mode: match config.preemption_mode {
388 PreemptionMode::Lifo => CorePreemptionMode::Lifo,
389 PreemptionMode::Fifo => CorePreemptionMode::Fifo,
390 },
391 sglang: Some(SglangArgs {
392 schedule_policy: Some(
393 match config.sglang.schedule_policy {
394 crate::engine::SglangSchedulePolicy::Fifo => "fifo",
395 crate::engine::SglangSchedulePolicy::Lpm => "lpm",
396 }
397 .to_string(),
398 ),
399 page_size: Some(config.block_size),
400 max_prefill_tokens: Some(config.sglang.max_prefill_tokens),
401 chunked_prefill_size: Some(config.sglang.chunked_prefill_size),
402 clip_max_new_tokens: Some(config.sglang.clip_max_new_tokens),
403 schedule_conservativeness: Some(config.sglang.schedule_conservativeness),
404 }),
405 emit_kv_events: config.emit_kv_events,
406 emit_kv_token_ids: config.emit_kv_token_ids,
407 }
408}
409
410fn core_request(request: Request) -> DirectRequest {
411 DirectRequest {
412 tokens: request.tokens,
413 max_output_tokens: request.max_output_tokens,
414 output_token_ids: request.output_token_ids,
415 uuid: Some(request.request_id),
416 arrival_timestamp_ms: None,
417 }
418}
419
420fn pending_output_suppression(
421 command: &Command,
422 handoffs: &HashMap<HandoffId, Uuid>,
423) -> Option<(Uuid, bool)> {
424 match *command {
425 Command::CancelRequest {
426 request_id,
427 discard_pending_output,
428 } => Some((request_id, discard_pending_output)),
429 Command::CancelSource { handoff_id } | Command::CancelDestination { handoff_id } => {
430 handoffs
431 .get(&handoff_id)
432 .copied()
433 .map(|request_id| (request_id, false))
434 }
435 _ => None,
436 }
437}
438
439fn map_command_effects(
440 effects: CoreCommandEffects,
441 metrics: Metrics,
442 suppressed_pending_output: bool,
443) -> Result<CommandEffects> {
444 let result = match effects.result {
445 CoreCommandResult::Submitted(id) => CommandResult::Submitted(id),
446 CoreCommandResult::DestinationAccepted { request_id } => {
447 CommandResult::DestinationAccepted { request_id }
448 }
449 CoreCommandResult::Applied => CommandResult::Applied,
450 CoreCommandResult::Noop => CommandResult::Noop,
451 };
452 Ok(CommandEffects {
453 result,
454 lifecycle_events: effects
455 .lifecycle_events
456 .into_iter()
457 .map(map_lifecycle)
458 .collect(),
459 kv_events: effects.kv_events,
460 retired_requests: effects.retired_requests,
461 metrics,
462 suppressed_pending_output,
463 })
464}
465
466fn map_lifecycle(event: CoreLifecycle) -> LifecycleEvent {
467 match event {
468 CoreLifecycle::SourceHeld {
469 handoff_id,
470 request_id,
471 transfer_timing,
472 } => LifecycleEvent::SourceHeld {
473 handoff_id,
474 request_id,
475 transfer_timing,
476 },
477 CoreLifecycle::DestinationReserved {
478 handoff_id,
479 request_id,
480 transferable_prompt_tokens,
481 } => LifecycleEvent::DestinationReserved {
482 handoff_id,
483 request_id,
484 transferable_prompt_tokens,
485 },
486 }
487}
488
489fn map_metrics(metrics: MockerMetrics) -> Metrics {
490 Metrics {
491 dp_rank: metrics.dp_rank,
492 active_blocks: metrics.active_decode_blocks,
493 inactive_blocks: metrics.inactive_decode_blocks,
494 total_blocks: metrics.total_blocks,
495 cache_usage: metrics.gpu_cache_usage_perc,
496 physical_cache_usage: metrics.physical_gpu_cache_usage_perc,
497 running_requests: metrics.running_requests,
498 waiting_requests: metrics.waiting_requests,
499 preemptions_total: metrics.vllm_preemptions_total,
500 sglang_cache_hit_tokens: metrics.sglang_cache_hit_tokens,
501 sglang_cache_total_tokens: metrics.sglang_cache_total_tokens,
502 }
503}
504
505fn split_pass(
506 pass: EnginePassResult,
507) -> Result<(
508 crate::engine::generalized::SameTimestampRetry,
509 PassStartEffects,
510 PassCompletionEffects,
511)> {
512 let EnginePassResult {
513 same_timestamp_retry,
514 output_signals,
515 admissions,
516 pressure_events,
517 lifecycle_events,
518 mocker_metrics,
519 kv_event_visibility,
520 kv_events,
521 fpm,
522 ..
523 } = pass;
524 let (start_kv, completion_kv) = match kv_event_visibility {
525 KvEventVisibility::PassEnd => (Vec::new(), kv_events),
526 };
527 let start = PassStartEffects {
528 admissions: admissions
529 .into_iter()
530 .map(|admission| Admission {
531 request_id: admission.uuid,
532 reused_input_tokens: admission.reused_input_tokens,
533 cache_tier_attribution: admission.cache_tier_attribution,
534 })
535 .collect(),
536 pressure_events,
537 kv_events: start_kv,
538 };
539 let completion = PassCompletionEffects {
540 outputs: output_signals
541 .into_iter()
542 .map(|output| Output {
543 request_id: output.uuid,
544 token_id: output.token_id,
545 completed: output.completed,
546 rejected: output.rejected,
547 cached_tokens: output.cached_tokens,
548 })
549 .collect(),
550 lifecycle_events: lifecycle_events.into_iter().map(map_lifecycle).collect(),
551 kv_events: completion_kv,
552 metrics: map_metrics(mocker_metrics),
553 forward_pass_metrics: fpm.map(map_fpm).unwrap_or_default(),
554 };
555 Ok((same_timestamp_retry, start, completion))
556}
557
558fn map_fpm(fpm: crate::engine::common::protocols::ForwardPassSnapshot) -> ForwardPassMetrics {
559 ForwardPassMetrics {
560 num_prefill_requests: fpm.num_prefill_requests,
561 sum_prefill_tokens: fpm.sum_prefill_tokens,
562 var_prefill_length: fpm.var_prefill_length,
563 sum_prefill_kv_tokens: fpm.sum_prefill_kv_tokens,
564 num_decode_requests: fpm.num_decode_requests,
565 sum_decode_kv_tokens: fpm.sum_decode_kv_tokens,
566 var_decode_kv_tokens: fpm.var_decode_kv_tokens,
567 num_queued_prefill: fpm.num_queued_prefill,
568 sum_queued_prefill_tokens: fpm.sum_queued_prefill_tokens,
569 var_queued_prefill_length: fpm.var_queued_prefill_length,
570 num_queued_decode: fpm.num_queued_decode,
571 sum_queued_decode_kv_tokens: fpm.sum_queued_decode_kv_tokens,
572 var_queued_decode_kv_tokens: fpm.var_queued_decode_kv_tokens,
573 duration_ms: fpm.wall_time_secs * 1_000.0,
574 }
575}
576
577#[cfg(test)]
578mod tests {
579 use std::num::NonZeroU32;
580 use std::sync::Mutex;
581
582 use super::*;
583 use crate::engine::{
584 HostOffloadObservation, HostOffloadObservationData, NativeHostOffloadConfig, PressureKind,
585 TimingModelConfig,
586 };
587
588 #[derive(Debug, Clone, Copy, PartialEq)]
589 struct CapturedHostEvent {
590 request_id: Uuid,
591 kind: &'static str,
592 at_ms: f64,
593 }
594
595 #[derive(Default)]
596 struct HostEventCapture(Mutex<Vec<CapturedHostEvent>>);
597
598 impl HostEventCapture {
599 fn snapshot(&self) -> Vec<CapturedHostEvent> {
600 self.0.lock().expect("host event capture poisoned").clone()
601 }
602 }
603
604 impl HostOffloadObserver for HostEventCapture {
605 fn record(&self, observation: HostOffloadObservation<'_>) {
606 let (kind, at_ms) = match observation.event {
607 HostOffloadObservationData::LoadQueued { at_ms, .. } => ("load_queued", at_ms),
608 HostOffloadObservationData::LoadCompleted { at_ms, .. } => {
609 ("load_completed", at_ms)
610 }
611 HostOffloadObservationData::LoadCancelled { at_ms, .. } => {
612 ("load_cancelled", at_ms)
613 }
614 _ => return,
615 };
616 self.0
617 .lock()
618 .expect("host event capture poisoned")
619 .push(CapturedHostEvent {
620 request_id: observation.request_id,
621 kind,
622 at_ms,
623 });
624 }
625 }
626
627 fn rank() -> SchedulerRank {
628 rank_for_worker(WorkerType::Aggregated)
629 }
630
631 fn rank_for_worker(worker_type: WorkerType) -> SchedulerRank {
632 let config = EngineConfig {
633 worker_type,
634 num_gpu_blocks: 8,
635 block_size: 4,
636 max_num_seqs: 2,
637 max_num_batched_tokens: 16,
638 speedup_ratio: 0.0,
639 timing_model: TimingModelConfig::Fixed {
640 prefill_ms: 10.0,
641 decode_ms: 10.0,
642 },
643 ..EngineConfig::default()
644 };
645 SchedulerRank::new(
646 RankIdentity {
647 worker_id: 1,
648 dp_rank: 0,
649 dp_size: NonZeroU32::MIN,
650 },
651 &config,
652 )
653 .unwrap()
654 }
655
656 fn host_rank(observer: Arc<HostEventCapture>) -> SchedulerRank {
657 host_rank_with_capacity(observer, 2, 4)
658 }
659
660 fn host_rank_with_capacity(
661 observer: Arc<HostEventCapture>,
662 g1_blocks: usize,
663 host_blocks: usize,
664 ) -> SchedulerRank {
665 let config = EngineConfig {
666 num_gpu_blocks: g1_blocks,
667 block_size: 4,
668 max_num_seqs: 4,
669 max_num_batched_tokens: 16,
670 kv_cache_bytes_per_token: Some(250_000),
671 native_host_offload: Some(
672 NativeHostOffloadConfig::new(host_blocks).with_bandwidths(1.0, 1.0),
673 ),
674 speedup_ratio: 0.0,
675 timing_model: TimingModelConfig::Fixed {
676 prefill_ms: 10.0,
677 decode_ms: 10.0,
678 },
679 ..EngineConfig::default()
680 };
681 let mut rank = SchedulerRank::new(
682 RankIdentity {
683 worker_id: 2,
684 dp_rank: 0,
685 dp_size: NonZeroU32::MIN,
686 },
687 &config,
688 )
689 .unwrap();
690 rank.set_host_offload_observer(observer);
691 rank
692 }
693
694 fn submit_completed_prompt(
695 rank: &mut SchedulerRank,
696 request_id: Uuid,
697 tokens: Vec<u32>,
698 now_ms: f64,
699 ) -> f64 {
700 let effects = rank
701 .apply_command_effects(
702 Command::Submit(Request {
703 request_id,
704 tokens,
705 max_output_tokens: 0,
706 output_token_ids: None,
707 }),
708 CommandContext {
709 now_ms,
710 pass_in_flight: false,
711 },
712 None,
713 )
714 .unwrap();
715 assert_eq!(effects.result, CommandResult::Submitted(request_id));
716 let pass = rank.execute_pass(now_ms).unwrap();
717 let end_ms = pass.end_ms;
718 rank.complete_pass(pass.pending, end_ms).unwrap();
719 end_ms
720 }
721
722 fn queue_source_h2d(rank: &mut SchedulerRank) -> (HandoffId, Uuid, f64) {
725 let seed_id = Uuid::from_u128(93_001);
726 let seed_end = submit_completed_prompt(rank, seed_id, vec![1, 2, 3, 4], 0.0);
727 let store_due = rank.next_internal_deadline_ms().unwrap();
728 assert!(store_due >= seed_end);
729 rank.process_internal_work(store_due, false).unwrap();
730
731 let evict_id = Uuid::from_u128(93_002);
732 let evict_end =
733 submit_completed_prompt(rank, evict_id, vec![5, 6, 7, 8, 9, 10, 11, 12], store_due);
734 let mut restore_at_ms = evict_end;
735 while let Some(deadline) = rank.next_internal_deadline_ms() {
736 restore_at_ms = restore_at_ms.max(deadline);
737 rank.process_internal_work(restore_at_ms, false).unwrap();
738 }
739
740 let handoff_id = HandoffId::from(Uuid::from_u128(93_003));
741 let restore_id = Uuid::from_u128(93_004);
742 let effects = rank
743 .apply_command_effects(
744 Command::SubmitHandoffPrefill {
745 handoff_id,
746 request: Request {
747 request_id: restore_id,
748 tokens: vec![1, 2, 3, 4],
749 max_output_tokens: 0,
750 output_token_ids: None,
751 },
752 },
753 CommandContext {
754 now_ms: restore_at_ms,
755 pass_in_flight: false,
756 },
757 None,
758 )
759 .unwrap();
760 assert_eq!(effects.result, CommandResult::Submitted(restore_id));
761 let pass = rank.execute_pass(restore_at_ms).unwrap();
762 rank.complete_pass(pass.pending, pass.end_ms).unwrap();
763 let due_ms = rank.next_internal_deadline_ms().unwrap();
764 (handoff_id, restore_id, due_ms)
765 }
766
767 fn start_request_pass(
768 rank: &mut SchedulerRank,
769 request_id: Uuid,
770 output_token_ids: Vec<u32>,
771 ) -> PendingPass {
772 let effects = rank
773 .apply_command_effects(
774 Command::Submit(Request {
775 request_id,
776 tokens: vec![1, 2, 3, 4],
777 max_output_tokens: output_token_ids.len(),
778 output_token_ids: Some(output_token_ids),
779 }),
780 CommandContext {
781 now_ms: 0.0,
782 pass_in_flight: false,
783 },
784 None,
785 )
786 .unwrap();
787 assert_eq!(effects.result, CommandResult::Submitted(request_id));
788 let pass = rank.execute_pass(0.0).unwrap();
789 assert!(
790 pass.pending
791 .effects
792 .outputs
793 .iter()
794 .any(|output| output.request_id == request_id)
795 );
796 pass.pending
797 }
798
799 #[test]
800 fn ordinary_cancel_suppresses_pending_output_when_scheduler_state_is_removed() {
801 let request_id = Uuid::from_u128(90_001);
802 let mut rank = rank();
803 let mut pending = start_request_pass(&mut rank, request_id, vec![5, 6]);
804
805 let effects = rank
806 .apply_command_effects(
807 Command::CancelRequest {
808 request_id,
809 discard_pending_output: false,
810 },
811 CommandContext {
812 now_ms: 1.0,
813 pass_in_flight: true,
814 },
815 Some(&mut pending),
816 )
817 .unwrap();
818
819 assert_eq!(effects.result, CommandResult::Applied);
820 assert!(effects.suppressed_pending_output);
821 assert!(pending.effects.outputs.is_empty());
822 }
823
824 #[test]
825 fn ordinary_noop_cancel_preserves_pending_output() {
826 let request_id = Uuid::from_u128(90_002);
827 let mut rank = rank();
828 let mut pending = start_request_pass(&mut rank, request_id, vec![5]);
829
830 let effects = rank
831 .apply_command_effects(
832 Command::CancelRequest {
833 request_id,
834 discard_pending_output: false,
835 },
836 CommandContext {
837 now_ms: 1.0,
838 pass_in_flight: true,
839 },
840 Some(&mut pending),
841 )
842 .unwrap();
843
844 assert_eq!(effects.result, CommandResult::Noop);
845 assert!(!effects.suppressed_pending_output);
846 assert_eq!(pending.effects.outputs.len(), 1);
847 }
848
849 #[test]
850 fn explicit_discard_suppresses_pending_output_after_noop_cancellation() {
851 let request_id = Uuid::from_u128(90_003);
852 let mut rank = rank();
853 let mut pending = start_request_pass(&mut rank, request_id, vec![5]);
854
855 let effects = rank
856 .apply_command_effects(
857 Command::CancelRequest {
858 request_id,
859 discard_pending_output: true,
860 },
861 CommandContext {
862 now_ms: 1.0,
863 pass_in_flight: true,
864 },
865 Some(&mut pending),
866 )
867 .unwrap();
868
869 assert_eq!(effects.result, CommandResult::Noop);
870 assert!(effects.suppressed_pending_output);
871 assert!(pending.effects.outputs.is_empty());
872 }
873
874 #[test]
875 fn mid_pass_command_and_internal_call_keep_due_h2d_hidden() {
876 let observer = Arc::new(HostEventCapture::default());
877 let mut rank = host_rank(Arc::clone(&observer));
878 let (_handoff_id, restore_id, h2d_due_ms) = queue_source_h2d(&mut rank);
879 assert!(
880 observer
881 .snapshot()
882 .iter()
883 .any(|event| { event.request_id == restore_id && event.kind == "load_queued" })
884 );
885
886 let busy_id = Uuid::from_u128(93_005);
887 let submission = rank
888 .apply_command_effects(
889 Command::Submit(Request {
890 request_id: busy_id,
891 tokens: vec![21, 22, 23, 24],
892 max_output_tokens: 0,
893 output_token_ids: None,
894 }),
895 CommandContext {
896 now_ms: h2d_due_ms - 0.5,
897 pass_in_flight: false,
898 },
899 None,
900 )
901 .unwrap();
902 assert_eq!(submission.result, CommandResult::Submitted(busy_id));
903 let pass = rank.execute_pass(h2d_due_ms - 0.5).unwrap();
904 assert!(pass.end_ms > h2d_due_ms);
905 let mut pending = pass.pending;
906
907 let arrival_id = Uuid::from_u128(93_006);
908 let arrival = rank
909 .apply_command_effects(
910 Command::Submit(Request {
911 request_id: arrival_id,
912 tokens: vec![31, 32, 33, 34],
913 max_output_tokens: 0,
914 output_token_ids: None,
915 }),
916 CommandContext {
917 now_ms: h2d_due_ms + 0.5,
918 pass_in_flight: true,
919 },
920 Some(&mut pending),
921 )
922 .unwrap();
923 assert_eq!(arrival.result, CommandResult::Submitted(arrival_id));
924 assert!(arrival.kv_events.is_empty());
925 assert!(
926 !observer
927 .snapshot()
928 .iter()
929 .any(|event| { event.request_id == restore_id && event.kind == "load_completed" })
930 );
931
932 let internal = rank.process_internal_work(h2d_due_ms + 0.5, true).unwrap();
933 assert_eq!(internal, PassStartEffects::default());
934 assert!(
935 !observer
936 .snapshot()
937 .iter()
938 .any(|event| { event.request_id == restore_id && event.kind == "load_completed" })
939 );
940
941 rank.complete_pass(pending, pass.end_ms).unwrap();
942 let events = observer.snapshot();
943 assert!(
944 events.iter().any(|event| {
945 event.request_id == restore_id
946 && event.kind == "load_completed"
947 && event.at_ms == h2d_due_ms
950 }),
951 "events: {events:?}, pass end: {}",
952 pass.end_ms
953 );
954 }
955
956 #[test]
957 fn rejected_submit_at_due_deadline_does_not_settle_host_work() {
958 let observer = Arc::new(HostEventCapture::default());
959 let mut rank = host_rank(Arc::clone(&observer));
960 let (_handoff_id, restore_id, h2d_due_ms) = queue_source_h2d(&mut rank);
961 let events_before = observer.snapshot();
962 let metrics_before = rank.metrics();
963 let deadline_before = rank.next_internal_deadline_ms();
964
965 let error = rank
966 .apply_command_effects(
967 Command::Submit(Request {
968 request_id: restore_id,
969 tokens: vec![41, 42, 43, 44],
970 max_output_tokens: 0,
971 output_token_ids: None,
972 }),
973 CommandContext {
974 now_ms: h2d_due_ms + 1.0,
975 pass_in_flight: false,
976 },
977 None,
978 )
979 .unwrap_err();
980 assert!(format!("{error:#}").contains("already active"));
981 assert_eq!(rank.metrics(), metrics_before);
982 assert_eq!(rank.next_internal_deadline_ms(), deadline_before);
983 assert_eq!(observer.snapshot(), events_before);
984 assert!(rank.core.drain_kv_events().is_empty());
985 }
986
987 #[test]
988 fn same_pass_queues_multiple_host_loads_with_checked_headroom_updates() {
989 let observer = Arc::new(HostEventCapture::default());
990 let mut rank = host_rank_with_capacity(Arc::clone(&observer), 4, 8);
991 let mut now_ms = 0.0;
992 for (request_id, tokens) in [
993 (Uuid::from_u128(95_001), vec![1, 2, 3, 4]),
994 (Uuid::from_u128(95_002), vec![11, 12, 13, 14]),
995 ] {
996 now_ms = submit_completed_prompt(&mut rank, request_id, tokens, now_ms);
997 while let Some(deadline) = rank.next_internal_deadline_ms() {
998 now_ms = now_ms.max(deadline);
999 rank.process_internal_work(now_ms, false).unwrap();
1000 }
1001 }
1002 now_ms = submit_completed_prompt(
1003 &mut rank,
1004 Uuid::from_u128(95_003),
1005 (100..116).collect(),
1006 now_ms,
1007 );
1008 while let Some(deadline) = rank.next_internal_deadline_ms() {
1009 now_ms = now_ms.max(deadline);
1010 rank.process_internal_work(now_ms, false).unwrap();
1011 }
1012
1013 let loads = [
1014 (Uuid::from_u128(95_004), vec![1, 2, 3, 4, 21, 22, 23, 24]),
1015 (
1016 Uuid::from_u128(95_005),
1017 vec![11, 12, 13, 14, 31, 32, 33, 34],
1018 ),
1019 ];
1020 for (request_id, tokens) in &loads {
1021 let effects = rank
1022 .apply_command_effects(
1023 Command::Submit(Request {
1024 request_id: *request_id,
1025 tokens: tokens.clone(),
1026 max_output_tokens: 0,
1027 output_token_ids: None,
1028 }),
1029 CommandContext {
1030 now_ms,
1031 pass_in_flight: false,
1032 },
1033 None,
1034 )
1035 .unwrap();
1036 assert_eq!(effects.result, CommandResult::Submitted(*request_id));
1037 }
1038
1039 let pass = rank.execute_pass(now_ms).unwrap();
1040 let events = observer.snapshot();
1041 for (request_id, _) in loads {
1042 assert!(events.iter().any(|event| {
1043 event.request_id == request_id
1044 && event.kind == "load_queued"
1045 && event.at_ms == now_ms
1046 }));
1047 }
1048 rank.complete_pass(pass.pending, pass.end_ms).unwrap();
1049 }
1050
1051 #[test]
1052 fn mid_pass_request_and_source_cancel_observe_command_time_without_completing_h2d() {
1053 for cancel_source in [false, true] {
1054 let observer = Arc::new(HostEventCapture::default());
1055 let mut rank = host_rank(Arc::clone(&observer));
1056 let (handoff_id, restore_id, h2d_due_ms) = queue_source_h2d(&mut rank);
1057
1058 let busy_id = Uuid::from_u128(94_000 + u128::from(cancel_source));
1059 rank.apply_command_effects(
1060 Command::Submit(Request {
1061 request_id: busy_id,
1062 tokens: vec![21, 22, 23, 24],
1063 max_output_tokens: 0,
1064 output_token_ids: None,
1065 }),
1066 CommandContext {
1067 now_ms: h2d_due_ms - 0.5,
1068 pass_in_flight: false,
1069 },
1070 None,
1071 )
1072 .unwrap();
1073 let pass = rank.execute_pass(h2d_due_ms - 0.5).unwrap();
1074 assert!(pass.end_ms > h2d_due_ms);
1075 let mut pending = pass.pending;
1076 let cancel_at_ms = h2d_due_ms + 0.5;
1077 let command = if cancel_source {
1078 Command::CancelSource { handoff_id }
1079 } else {
1080 Command::CancelRequest {
1081 request_id: restore_id,
1082 discard_pending_output: false,
1083 }
1084 };
1085 let effects = rank
1086 .apply_command_effects(
1087 command,
1088 CommandContext {
1089 now_ms: cancel_at_ms,
1090 pass_in_flight: true,
1091 },
1092 Some(&mut pending),
1093 )
1094 .unwrap();
1095 assert_eq!(effects.result, CommandResult::Applied);
1096 let events = observer.snapshot();
1097 assert!(events.iter().any(|event| {
1098 event.request_id == restore_id
1099 && event.kind == "load_cancelled"
1100 && event.at_ms == cancel_at_ms
1101 }));
1102 assert!(
1103 !events.iter().any(|event| {
1104 event.request_id == restore_id && event.kind == "load_completed"
1105 })
1106 );
1107
1108 rank.complete_pass(pending, pass.end_ms).unwrap();
1109 assert!(
1110 !observer.snapshot().iter().any(|event| {
1111 event.request_id == restore_id && event.kind == "load_completed"
1112 })
1113 );
1114 }
1115 }
1116
1117 #[test]
1118 fn cancel_source_uses_command_time_outside_a_pass() {
1119 let observer = Arc::new(HostEventCapture::default());
1120 let mut rank = host_rank(Arc::clone(&observer));
1121 let (handoff_id, restore_id, h2d_due_ms) = queue_source_h2d(&mut rank);
1122 let cancel_at_ms = h2d_due_ms - 0.25;
1123
1124 let effects = rank
1125 .apply_command_effects(
1126 Command::CancelSource { handoff_id },
1127 CommandContext {
1128 now_ms: cancel_at_ms,
1129 pass_in_flight: false,
1130 },
1131 None,
1132 )
1133 .unwrap();
1134 assert_eq!(effects.result, CommandResult::Applied);
1135 assert!(observer.snapshot().iter().any(|event| {
1136 event.request_id == restore_id
1137 && event.kind == "load_cancelled"
1138 && event.at_ms == cancel_at_ms
1139 }));
1140 }
1141
1142 #[test]
1143 fn handoff_tracking_is_inserted_on_success_and_cleared_by_cancel() {
1144 let mut rank = rank_for_worker(WorkerType::Decode);
1145 let handoff_id = HandoffId::from(Uuid::from_u128(91_001));
1146 let request_id = Uuid::from_u128(91_002);
1147 let reservation = rank
1148 .apply_command_effects(
1149 Command::ReserveDestination {
1150 handoff_id,
1151 request: Request {
1152 request_id,
1153 tokens: vec![1, 2, 3, 4],
1154 max_output_tokens: 1,
1155 output_token_ids: Some(vec![5]),
1156 },
1157 },
1158 CommandContext {
1159 now_ms: 0.0,
1160 pass_in_flight: false,
1161 },
1162 None,
1163 )
1164 .unwrap();
1165 assert!(matches!(
1166 reservation.result,
1167 CommandResult::DestinationAccepted { .. }
1168 ));
1169 assert_eq!(rank.handoff_requests.get(&handoff_id), Some(&request_id));
1170
1171 let cancellation = rank
1172 .apply_command_effects(
1173 Command::CancelDestination { handoff_id },
1174 CommandContext {
1175 now_ms: 0.0,
1176 pass_in_flight: false,
1177 },
1178 None,
1179 )
1180 .unwrap();
1181 assert_eq!(cancellation.result, CommandResult::Applied);
1182 assert!(!rank.handoff_requests.contains_key(&handoff_id));
1183 }
1184
1185 #[test]
1186 fn pass_start_exposes_vllm_preemption_pressure_event() {
1187 let config = EngineConfig {
1188 num_gpu_blocks: 6,
1189 block_size: 4,
1190 max_num_seqs: 2,
1191 max_num_batched_tokens: 16,
1192 enable_prefix_caching: false,
1193 enable_chunked_prefill: true,
1194 speedup_ratio: 0.0,
1195 preemption_mode: PreemptionMode::Lifo,
1196 timing_model: TimingModelConfig::Fixed {
1197 prefill_ms: 10.0,
1198 decode_ms: 10.0,
1199 },
1200 ..EngineConfig::default()
1201 };
1202 let mut rank = SchedulerRank::new(
1203 RankIdentity {
1204 worker_id: 7,
1205 dp_rank: 0,
1206 dp_size: NonZeroU32::MIN,
1207 },
1208 &config,
1209 )
1210 .unwrap();
1211 let first = Uuid::from_u128(92_001);
1212 let second = Uuid::from_u128(92_002);
1213 for (request_id, tokens) in [
1214 (first, (0..8).collect::<Vec<_>>()),
1215 (second, (100..108).collect::<Vec<_>>()),
1216 ] {
1217 let effects = rank
1218 .apply_command_effects(
1219 Command::Submit(Request {
1220 request_id,
1221 tokens,
1222 max_output_tokens: 8,
1223 output_token_ids: None,
1224 }),
1225 CommandContext {
1226 now_ms: 0.0,
1227 pass_in_flight: false,
1228 },
1229 None,
1230 )
1231 .unwrap();
1232 assert_eq!(effects.result, CommandResult::Submitted(request_id));
1233 }
1234
1235 let mut now_ms = 0.0;
1236 let mut observed = None;
1237 for _ in 0..16 {
1238 let pass = rank.execute_pass(now_ms).unwrap();
1239 if let Some(event) = pass.start_effects.pressure_events.first() {
1240 assert_eq!(pass.start_effects.pressure_events.len(), 1);
1241 observed = Some(event.clone());
1242 }
1243 let end_ms = pass.end_ms;
1244 rank.complete_pass(pass.pending, end_ms).unwrap();
1245 if observed.is_some() {
1246 break;
1247 }
1248 now_ms = end_ms.max(now_ms + 1.0);
1249 }
1250
1251 let event = observed.expect("tight native G1 capacity should preempt one vLLM request");
1252 assert_eq!(event.at_ms, now_ms);
1253 assert_eq!(event.kind, PressureKind::VllmPreemption);
1254 assert_eq!(event.request_id, second);
1255 assert_eq!(event.state_before.running_requests, 2);
1256 assert_eq!(event.state_before.waiting_requests, Some(0));
1257 assert_eq!(event.state_after.running_requests, 1);
1258 assert_eq!(event.state_after.waiting_requests, Some(1));
1259 assert!(event.request_active_blocks_before > 0);
1260 assert!(event.state_after.active_blocks < event.state_before.active_blocks);
1261 assert_eq!(event.logical_available_blocks_before, None);
1262 assert_eq!(event.required_blocks_before, None);
1263 }
1264
1265 #[test]
1266 fn terminal_handoff_output_clears_tracking() {
1267 let mut rank = rank_for_worker(WorkerType::Prefill);
1268 let handoff_id = HandoffId::from(Uuid::from_u128(92_001));
1269 let request_id = Uuid::from_u128(92_002);
1270 let submission = rank
1271 .apply_command_effects(
1272 Command::SubmitHandoffPrefill {
1273 handoff_id,
1274 request: Request {
1275 request_id,
1276 tokens: vec![1, 2, 3, 4],
1277 max_output_tokens: 1,
1278 output_token_ids: Some(vec![5]),
1279 },
1280 },
1281 CommandContext {
1282 now_ms: 0.0,
1283 pass_in_flight: false,
1284 },
1285 None,
1286 )
1287 .unwrap();
1288 assert_eq!(submission.result, CommandResult::Submitted(request_id));
1289 assert_eq!(rank.handoff_requests.get(&handoff_id), Some(&request_id));
1290
1291 let pass = rank.execute_pass(0.0).unwrap();
1292 let completion = rank.complete_pass(pass.pending, pass.end_ms).unwrap();
1293 assert!(
1294 completion
1295 .outputs
1296 .iter()
1297 .any(|output| output.request_id == request_id && output.completed)
1298 );
1299 assert!(!rank.handoff_requests.contains_key(&handoff_id));
1300 }
1301}