Skip to main content

aisimulate_core/engine/scheduler/
rank.rs

1// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4//! Thin neutral contract adapter over the mechanically moved scheduler cores.
5
6use 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
40/// One preserved vLLM/SGLang scheduler rank behind the neutral contract.
41pub 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        // Preserve the scheduler's command boundary when no model step is in
154        // flight: native G1 mutations produced by the command belong to its
155        // returned effects. Mid-pass mutations remain buffered so neither a
156        // command nor a due physical transfer can expose KV state before the
157        // shared completion boundary.
158        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        // The preserved scheduler retries deferred destination reservations
226        // when a forward pass releases capacity. Keep that wakeup at the
227        // pass-completion boundary: command-time retry is suppressed while a
228        // pass is in flight, and waiting until an unrelated later pass can
229        // leave disaggregated replay permanently asleep.
230        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        // Occupancy is authoritative at the shared completion boundary, but
239        // SGLang cache hit/total are transient observations from this pass.
240        // Preserve those fields while refreshing the rest of the snapshot;
241        // a later live adapter latches the last non-empty observation.
242        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    /// Seed G2 with one prompt, evict it from G1, then queue an H2D owned by a
723    /// still-pending source handoff. Returns `(handoff_id, request_id, due_ms)`.
724    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                    // The observation retains the physical completion timestamp,
948                    // even though it is emitted only at the later model boundary.
949                    && 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}