Skip to main content

dynamo_mocker/common/
protocols.rs

1// SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4use derive_builder::Builder;
5use dynamo_kv_router::config::RouterQueuePolicy;
6use serde::{Deserialize, Serialize};
7use std::path::{Path, PathBuf};
8use std::str::FromStr;
9use std::sync::Arc;
10use uuid::Uuid;
11use validator::{Validate, ValidationError};
12
13use crate::common::perf_model::PerfModel;
14use dynamo_kv_router::protocols::{KvCacheEvent, StorageTier};
15use dynamo_tokens::Token;
16
17/// Trait for publishing KV cache events.
18/// This abstracts the runtime dependency so mocker components can remain generic.
19pub trait KvCacheEventSink: Send + Sync {
20    fn publish(&self, event: KvCacheEvent) -> anyhow::Result<()>;
21
22    fn publish_with_storage_tier(
23        &self,
24        event: KvCacheEvent,
25        _storage_tier: StorageTier,
26    ) -> anyhow::Result<()> {
27        self.publish(event)
28    }
29
30    /// Publishes events that share one source visibility boundary.
31    ///
32    /// Implementations that do not have a native batch representation retain
33    /// singleton delivery semantics by default.
34    fn publish_batch_with_storage_tiers(
35        &self,
36        events: Vec<(KvCacheEvent, StorageTier)>,
37    ) -> anyhow::Result<()> {
38        let mut first_error = None;
39        for (event, storage_tier) in events {
40            if let Err(error) = self.publish_with_storage_tier(event, storage_tier) {
41                first_error.get_or_insert(error);
42            }
43        }
44        first_error.map_or(Ok(()), Err)
45    }
46}
47
48/// Raw KV event payload used by transport-specific publishers such as the
49/// vLLM-native ZMQ event stream.
50#[derive(Debug, Clone)]
51pub struct RawKvEvent {
52    pub event: KvCacheEvent,
53    pub block_token_ids: Option<Vec<Vec<u32>>>,
54    pub storage_tier: StorageTier,
55}
56
57/// Trait for publishing transport-specific raw KV event payloads.
58pub trait RawKvEventSink: Send + Sync {
59    fn publish(&self, event: RawKvEvent) -> anyhow::Result<()>;
60
61    /// Publishes raw events that share one source visibility boundary.
62    ///
63    /// Implementations that do not have a native batch representation retain
64    /// singleton delivery semantics by default.
65    fn publish_batch(&self, events: Vec<RawKvEvent>) -> anyhow::Result<()> {
66        let mut first_error = None;
67        for event in events {
68            if let Err(error) = self.publish(event) {
69                first_error.get_or_insert(error);
70            }
71        }
72        first_error.map_or(Ok(()), Err)
73    }
74}
75
76/// Shared KV event publisher bundle used by schedulers and KV managers.
77#[derive(Clone, Default)]
78pub struct KvEventPublishers {
79    event_sink: Option<Arc<dyn KvCacheEventSink>>,
80    raw_sink: Option<Arc<dyn RawKvEventSink>>,
81}
82
83impl KvEventPublishers {
84    pub fn new(
85        event_sink: Option<Arc<dyn KvCacheEventSink>>,
86        raw_sink: Option<Arc<dyn RawKvEventSink>>,
87    ) -> Self {
88        Self {
89            event_sink,
90            raw_sink,
91        }
92    }
93
94    pub fn raw_enabled(&self) -> bool {
95        self.raw_sink.is_some()
96    }
97
98    pub fn is_empty(&self) -> bool {
99        self.event_sink.is_none() && self.raw_sink.is_none()
100    }
101
102    pub fn publish(
103        &self,
104        event: KvCacheEvent,
105        block_token_ids: Option<&[Vec<u32>]>,
106    ) -> anyhow::Result<()> {
107        self.publish_with_storage_tier(event, block_token_ids, StorageTier::Device)
108    }
109
110    pub fn publish_with_storage_tier(
111        &self,
112        event: KvCacheEvent,
113        block_token_ids: Option<&[Vec<u32>]>,
114        storage_tier: StorageTier,
115    ) -> anyhow::Result<()> {
116        if let Some(sink) = self.event_sink.as_ref() {
117            sink.publish_with_storage_tier(event.clone(), storage_tier)?;
118        }
119
120        if let Some(sink) = self.raw_sink.as_ref() {
121            sink.publish(RawKvEvent {
122                event,
123                block_token_ids: block_token_ids.map(|token_ids| token_ids.to_vec()),
124                storage_tier,
125            })?;
126        }
127
128        Ok(())
129    }
130
131    /// Publishes normal KV events without also forwarding them to a raw sink.
132    ///
133    /// Deferred live-scheduler forwarding uses this to preserve its source
134    /// visibility boundary for normal and raw sinks independently.
135    pub(crate) fn publish_event_sink_batch_only(
136        &self,
137        events: Vec<(KvCacheEvent, StorageTier)>,
138    ) -> anyhow::Result<()> {
139        if let Some(sink) = self.event_sink.as_ref() {
140            sink.publish_batch_with_storage_tiers(events)?;
141        }
142        Ok(())
143    }
144
145    /// Publishes raw events as one source visibility boundary.
146    pub(crate) fn publish_raw_batch(&self, events: Vec<RawKvEvent>) -> anyhow::Result<()> {
147        if let Some(sink) = self.raw_sink.as_ref() {
148            sink.publish_batch(events)?;
149        }
150        Ok(())
151    }
152}
153
154/// Replay-neutral per-pass metrics shared by offline and Live Mocker drivers.
155pub use aisimulate_core::replay::ForwardPassSnapshot;
156
157/// Trait for publishing forward pass metrics snapshots.
158/// This abstracts the FPM publishing pipeline so mocker schedulers remain generic.
159pub trait FpmSink: Send + Sync {
160    fn publish(&self, snapshot: ForwardPassSnapshot) -> anyhow::Result<()>;
161}
162
163/// Optional FPM sink used by schedulers.
164/// Wraps `Option<Arc<dyn FpmSink>>` for ergonomic passing and no-op default behavior.
165#[derive(Clone, Default)]
166pub struct FpmPublisher {
167    sink: Option<Arc<dyn FpmSink>>,
168}
169
170impl FpmPublisher {
171    pub fn new(sink: Option<Arc<dyn FpmSink>>) -> Self {
172        Self { sink }
173    }
174
175    pub fn publish(&self, snapshot: ForwardPassSnapshot) -> anyhow::Result<()> {
176        if let Some(sink) = &self.sink {
177            sink.publish(snapshot)?;
178        }
179        Ok(())
180    }
181}
182
183/// Replay-owned request DTO shared by Dynamo's compatibility and Live Mocker
184/// drivers. The type remains provider-neutral; Dynamo-specific metadata is
185/// interpreted only by Dynamo adapters.
186pub use aisimulate_core::replay::DirectRequest;
187
188/// Signal for output token generation with completion status
189#[derive(Debug, Clone, Serialize, Deserialize)]
190pub struct OutputSignal {
191    pub uuid: Uuid,
192    #[serde(default, skip_serializing_if = "Option::is_none")]
193    pub token_id: Option<Token>,
194    /// Terminal flag: the request's lifecycle has ended. Replay drivers free
195    /// resources and advance/notify on this.
196    pub completed: bool,
197    /// Set with `completed` when the request was rejected without ever running
198    /// (its footprint exceeds the whole KV pool); drivers free/advance but
199    /// exclude it from token/latency/throughput stats.
200    #[serde(default)]
201    pub rejected: bool,
202    #[serde(default, skip_serializing_if = "Option::is_none")]
203    pub handoff_delay_ms: Option<f64>,
204    /// Prompt tokens served from KV cache at admission (scheduler truth,
205    /// post-eviction). Set once, on the request's first output signal.
206    #[serde(default, skip_serializing_if = "Option::is_none")]
207    pub cached_tokens: Option<usize>,
208}
209
210/// Preemption policy for evicting decode requests under memory pressure
211#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
212#[serde(rename_all = "lowercase")]
213pub enum PreemptionMode {
214    /// Evict the newest request (matches vLLM v1 default)
215    #[default]
216    Lifo,
217    /// Evict the oldest request
218    Fifo,
219}
220
221impl FromStr for PreemptionMode {
222    type Err = String;
223
224    fn from_str(value: &str) -> Result<Self, Self::Err> {
225        match value.to_ascii_lowercase().as_str() {
226            "lifo" => Ok(Self::Lifo),
227            "fifo" => Ok(Self::Fifo),
228            _ => Err(format!(
229                "Invalid preemption_mode: '{value}'. Must be 'lifo' or 'fifo'."
230            )),
231        }
232    }
233}
234
235/// Engine type for selecting scheduling and KV cache simulation behavior
236#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
237#[serde(rename_all = "lowercase")]
238pub enum EngineType {
239    /// vLLM-style scheduling with hash-based block KV cache
240    #[default]
241    Vllm,
242    /// SGLang-style scheduling with radix-tree KV cache
243    Sglang,
244    /// TensorRT-LLM-style scheduling. Reuses the vLLM scheduler
245    /// core with a TensorRT-LLM-style admission policy.
246    Trtllm,
247}
248
249impl FromStr for EngineType {
250    type Err = String;
251
252    fn from_str(value: &str) -> Result<Self, Self::Err> {
253        match value.to_ascii_lowercase().as_str() {
254            "vllm" => Ok(Self::Vllm),
255            "sglang" => Ok(Self::Sglang),
256            "trtllm" => Ok(Self::Trtllm),
257            _ => Err(format!(
258                "Invalid engine_type '{value}'. Must be 'vllm', 'sglang', or 'trtllm'."
259            )),
260        }
261    }
262}
263
264/// Worker type for disaggregated serving configurations
265#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
266#[serde(rename_all = "lowercase")]
267pub enum WorkerType {
268    /// Standard aggregated worker handling both prefill and decode
269    #[default]
270    Aggregated,
271    /// Dedicated prefill worker in disaggregated mode
272    Prefill,
273    /// Dedicated decode worker in disaggregated mode
274    Decode,
275}
276
277impl FromStr for WorkerType {
278    type Err = String;
279
280    fn from_str(value: &str) -> Result<Self, Self::Err> {
281        match value.to_ascii_lowercase().as_str() {
282            "aggregated" => Ok(Self::Aggregated),
283            "prefill" => Ok(Self::Prefill),
284            "decode" => Ok(Self::Decode),
285            _ => Err(format!(
286                "Invalid worker_type '{value}'. Must be 'aggregated', 'prefill', or 'decode'."
287            )),
288        }
289    }
290}
291
292/// Physical KV footprint used to model a coordinated disaggregated transfer.
293#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
294#[serde(rename_all = "snake_case")]
295pub enum KvTransferTimingMode {
296    /// Charge the source request's full logical prompt length.
297    #[default]
298    FullPrompt,
299    /// Charge only the physical prompt footprint missing at the destination.
300    DestinationMissing,
301}
302
303impl FromStr for KvTransferTimingMode {
304    type Err = String;
305
306    fn from_str(value: &str) -> Result<Self, Self::Err> {
307        match value.to_ascii_lowercase().as_str() {
308            "full_prompt" => Ok(Self::FullPrompt),
309            "destination_missing" => Ok(Self::DestinationMissing),
310            _ => Err(format!(
311                "Invalid kv_transfer_timing_mode '{value}'. Must be 'full_prompt' or 'destination_missing'."
312            )),
313        }
314    }
315}
316
317/// Configuration for reasoning/thinking token output in the mocker.
318///
319/// When set, the mocker wraps the first portion of each response in thinking
320/// boundary tokens: `[start_token, random..., end_token, random...]`.
321#[derive(Debug, Clone, Serialize, Deserialize, Validate)]
322pub struct ReasoningConfig {
323    pub start_thinking_token_id: u32,
324    pub end_thinking_token_id: u32,
325    #[validate(range(min = 0.0, max = 1.0))]
326    pub thinking_ratio: f64,
327}
328
329impl ReasoningConfig {
330    /// Number of thinking tokens (including start/end boundaries) for a given osl.
331    /// Returns 0 if osl < 2 (thinking disabled). Otherwise clamps to [2, osl].
332    pub fn num_thinking_tokens(&self, max_output_tokens: usize) -> usize {
333        if max_output_tokens < 2 {
334            return 0;
335        }
336        let raw = (max_output_tokens as f64 * self.thinking_ratio).floor() as usize;
337        if raw == 0 {
338            return 0;
339        }
340        raw.max(2).min(max_output_tokens)
341    }
342
343    /// Number of response tokens after the thinking block.
344    pub fn num_response_tokens(&self, max_output_tokens: usize) -> usize {
345        max_output_tokens.saturating_sub(self.num_thinking_tokens(max_output_tokens))
346    }
347}
348
349/// SGLang-specific configuration parameters.
350///
351/// Grouped into a nested struct to keep the `MockEngineArgs` namespace clean,
352/// following the same pattern as [`ReasoningConfig`].
353#[derive(Debug, Clone, Serialize, Deserialize, Validate, Default)]
354pub struct SglangArgs {
355    /// Scheduling policy: "fifo"/"fcfs" or "lpm". Default: "fifo".
356    pub schedule_policy: Option<String>,
357    /// Radix cache page size in tokens. Default: 1.
358    #[validate(range(min = 1))]
359    pub page_size: Option<usize>,
360    /// Maximum prefill tokens budget per batch. Default: 16384.
361    #[validate(range(min = 1))]
362    pub max_prefill_tokens: Option<usize>,
363    /// Chunked prefill size (max tokens per chunk). Default: 8192.
364    #[validate(range(min = 1))]
365    pub chunked_prefill_size: Option<usize>,
366    /// Clip max new tokens for admission budget. Default: 4096.
367    #[validate(range(min = 1))]
368    pub clip_max_new_tokens: Option<usize>,
369    /// Schedule conservativeness factor (0.0–1.0). Default: 1.0.
370    #[validate(range(min = 0.0, max = 1.0))]
371    pub schedule_conservativeness: Option<f64>,
372}
373
374/// TensorRT-LLM-specific configuration parameters.
375///
376/// Grouped into a nested struct to keep the `MockEngineArgs` namespace clean,
377/// following the same pattern as [`SglangArgs`].
378#[derive(Debug, Clone, Serialize, Deserialize, Validate, Default)]
379pub struct TrtllmArgs {
380    /// Capacity scheduler policy, supported only `"guaranteed_no_evict"`
381    /// (TensorRT-LLM's default). Default: `"guaranteed_no_evict"`.
382    pub capacity_scheduler_policy: Option<String>,
383}
384
385/// Keeps omitted JSON fields distinct from explicit `null` so serde can replace
386/// the old hand-written parser without losing input-config semantics.
387#[derive(Debug, Clone, Default)]
388enum OptionalConfigValue<T> {
389    #[default]
390    Missing,
391    Present(Option<T>),
392}
393
394impl<'de, T> Deserialize<'de> for OptionalConfigValue<T>
395where
396    T: Deserialize<'de>,
397{
398    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
399    where
400        D: serde::Deserializer<'de>,
401    {
402        Option::<T>::deserialize(deserializer).map(Self::Present)
403    }
404}
405
406impl<T> OptionalConfigValue<T> {
407    fn into_nullable(self) -> Option<Option<T>> {
408        match self {
409            Self::Missing => None,
410            Self::Present(value) => Some(value),
411        }
412    }
413
414    fn into_non_null(self, field: &str) -> Result<Option<T>, String> {
415        match self {
416            Self::Missing => Ok(None),
417            Self::Present(Some(value)) => Ok(Some(value)),
418            Self::Present(None) => Err(format!("{field} must not be null")),
419        }
420    }
421}
422
423#[derive(Debug, Clone, Default, Deserialize)]
424#[serde(default, deny_unknown_fields)]
425struct MockEngineArgsSerde {
426    engine_type: OptionalConfigValue<String>,
427    num_gpu_blocks: OptionalConfigValue<usize>,
428    block_size: OptionalConfigValue<usize>,
429    max_model_len: OptionalConfigValue<usize>,
430    max_num_seqs: OptionalConfigValue<usize>,
431    max_num_batched_tokens: OptionalConfigValue<usize>,
432    enable_prefix_caching: OptionalConfigValue<bool>,
433    enable_chunked_prefill: OptionalConfigValue<bool>,
434    speedup_ratio: OptionalConfigValue<f64>,
435    decode_speedup_ratio: OptionalConfigValue<f64>,
436    dp_size: OptionalConfigValue<u32>,
437    startup_time: OptionalConfigValue<f64>,
438    worker_type: OptionalConfigValue<String>,
439    is_prefill: OptionalConfigValue<bool>,
440    is_decode: OptionalConfigValue<bool>,
441    planner_profile_data: OptionalConfigValue<PathBuf>,
442    aic_backend: OptionalConfigValue<String>,
443    aic_system: OptionalConfigValue<String>,
444    aic_backend_version: OptionalConfigValue<String>,
445    aic_tp_size: OptionalConfigValue<usize>,
446    aic_model_path: OptionalConfigValue<String>,
447    aic_moe_tp_size: OptionalConfigValue<usize>,
448    aic_moe_ep_size: OptionalConfigValue<usize>,
449    aic_attention_dp_size: OptionalConfigValue<usize>,
450    aic_gemm_dtype: OptionalConfigValue<String>,
451    aic_moe_dtype: OptionalConfigValue<String>,
452    aic_fmha_dtype: OptionalConfigValue<String>,
453    aic_kv_cache_dtype: OptionalConfigValue<String>,
454    aic_comm_dtype: OptionalConfigValue<String>,
455    aic_nextn: OptionalConfigValue<usize>,
456    aic_nextn_accept_rates: OptionalConfigValue<String>,
457    aic_mtp_seed: OptionalConfigValue<u64>,
458    gpu_memory_utilization: OptionalConfigValue<f64>,
459    mem_fraction_static: OptionalConfigValue<f64>,
460    free_gpu_memory_fraction: OptionalConfigValue<f64>,
461    enable_local_indexer: OptionalConfigValue<bool>,
462    bootstrap_port: OptionalConfigValue<u16>,
463    handoff_session_timeout_ms: OptionalConfigValue<u64>,
464    #[serde(alias = "kv_transfer_bytes_per_token")]
465    kv_bytes_per_token: OptionalConfigValue<usize>,
466    kv_cache_bytes_per_token: OptionalConfigValue<usize>,
467    kv_transfer_bandwidth: OptionalConfigValue<f64>,
468    kv_transfer_timing_mode: OptionalConfigValue<String>,
469    reasoning: OptionalConfigValue<ReasoningConfig>,
470    response_replay_trace_path: OptionalConfigValue<PathBuf>,
471    zmq_kv_events_port: OptionalConfigValue<u16>,
472    zmq_replay_port: OptionalConfigValue<u16>,
473    preemption_mode: OptionalConfigValue<String>,
474    router_queue_policy: OptionalConfigValue<String>,
475    sglang: OptionalConfigValue<SglangArgs>,
476    trtllm: OptionalConfigValue<TrtllmArgs>,
477    timing_model: OptionalConfigValue<TimingModelSerde>,
478    #[serde(rename = "has_perf_model")]
479    _has_perf_model: OptionalConfigValue<serde_json::Value>,
480}
481
482#[derive(Debug, Clone, Deserialize)]
483#[serde(tag = "type", rename_all = "lowercase", deny_unknown_fields)]
484enum TimingModelSerde {
485    Default,
486    Polynomial,
487    Fixed { prefill_ms: f64, decode_ms: f64 },
488}
489
490fn load_perf_model(path: &Path) -> Arc<PerfModel> {
491    match PerfModel::from_npz(path) {
492        Ok(model) => {
493            tracing::info!("Successfully loaded performance model from: {:?}", path);
494            Arc::new(model)
495        }
496        Err(e) => {
497            tracing::error!(
498                "Failed to load performance model from {:?}: {}. Falling back to polynomial model.",
499                path,
500                e
501            );
502            Arc::new(PerfModel::default())
503        }
504    }
505}
506
507/// Configuration arguments for MockEngine
508#[derive(Debug, Clone, Serialize, Deserialize, Builder, Validate)]
509#[serde(try_from = "MockEngineArgsSerde")]
510#[validate(schema(function = "validate_mock_engine_args"))]
511#[builder(pattern = "owned", build_fn(public))]
512pub struct MockEngineArgs {
513    /// Engine type: vLLM, SGLang, or TensorRT-LLM simulation
514    #[builder(default = "EngineType::Vllm")]
515    pub engine_type: EngineType,
516
517    /// Usable simulated G1 capacity. This preserves the mocker's historical
518    /// convention across backends. A raw vLLM `num_gpu_blocks` value also
519    /// includes its reserved null block, so parity runs configure real vLLM
520    /// with one additional total block.
521    #[builder(default = "16384")]
522    #[validate(range(min = 1))]
523    pub num_gpu_blocks: usize,
524
525    #[builder(default = "0")]
526    pub block_size: usize,
527
528    /// Optional vLLM sequence-length limit, including prompt and generated
529    /// tokens. Requests with no room to generate are rejected before admission.
530    #[builder(default = "None")]
531    #[validate(range(min = 1))]
532    pub max_model_len: Option<usize>,
533
534    // This was 1024 in the past but reverted back to 256
535    #[builder(default = Some(256))]
536    #[validate(range(min = 1))]
537    pub max_num_seqs: Option<usize>,
538
539    // default for open api server, for llm class it's 16384
540    #[builder(default = Some(8192))]
541    #[validate(range(min = 1))]
542    pub max_num_batched_tokens: Option<usize>,
543
544    #[builder(default = true)]
545    pub enable_prefix_caching: bool,
546
547    #[builder(default = true)]
548    pub enable_chunked_prefill: bool,
549
550    #[builder(default = "1.0")]
551    #[validate(range(min = 0.0))]
552    pub speedup_ratio: f64,
553
554    /// Additional speedup multiplier applied only to decode steps.
555    /// Models speculative decoding (e.g. Eagle) where decode throughput improves
556    /// without affecting prefill latency. The effective decode speedup is
557    /// `speedup_ratio * decode_speedup_ratio`.
558    #[builder(default = "1.0")]
559    #[validate(range(min = 0.0))]
560    pub decode_speedup_ratio: f64,
561
562    #[builder(default = "1")]
563    #[validate(range(min = 1))]
564    pub dp_size: u32,
565
566    /// Optional startup time in seconds to simulate engine initialization delay
567    #[builder(default = "None")]
568    #[validate(range(min = 0.0))]
569    pub startup_time: Option<f64>,
570
571    /// Worker type for disaggregated serving (Aggregated, Prefill, or Decode)
572    #[builder(default = "WorkerType::Aggregated")]
573    pub worker_type: WorkerType,
574
575    /// Original planner profile NPZ path used to materialize `perf_model`.
576    #[builder(default = "None")]
577    pub planner_profile_data: Option<PathBuf>,
578
579    /// Performance model for timing predictions (not serialized, loaded from planner_profile_data)
580    #[serde(skip)]
581    #[builder(default = "Arc::new(PerfModel::default())")]
582    pub perf_model: Arc<PerfModel>,
583
584    /// If set, indicates direct AIC SDK calls should be used.
585    /// The value is the backend name (e.g., "sglang", "vllm").
586    /// The Python layer reads this and overrides perf_model with an Aiconfigurator callback.
587    #[serde(skip)]
588    #[builder(default = "None")]
589    pub aic_backend: Option<String>,
590
591    /// AIC GPU system name (e.g., "h200_sxm"). Required when aic_backend is set.
592    #[serde(skip)]
593    #[builder(default = "None")]
594    pub aic_system: Option<String>,
595
596    /// AIC performance-database slot ("current", "previous", or "next" when available),
597    /// or a version assigned to one of those slots.
598    /// If None, uses the release database's "current" slot.
599    #[serde(skip)]
600    #[builder(default = "None")]
601    pub aic_backend_version: Option<String>,
602
603    /// Tensor parallel size for AIC latency prediction.
604    /// Only affects AIC performance model lookups, not mocker scheduling.
605    #[serde(skip)]
606    #[builder(default = "None")]
607    pub aic_tp_size: Option<usize>,
608
609    /// HuggingFace model path for AIC latency prediction (e.g., "nvidia/Llama-3.1-8B-Instruct-FP8").
610    #[serde(skip)]
611    #[builder(default = "None")]
612    pub aic_model_path: Option<String>,
613
614    /// MoE tensor-parallel size for AIC latency prediction (e.g., 4 for pure MoE-TP).
615    /// Required for MoE models; must satisfy: aic_tp_size * aic_attention_dp_size == aic_moe_tp_size * aic_moe_ep_size.
616    #[serde(skip)]
617    #[builder(default = "None")]
618    pub aic_moe_tp_size: Option<usize>,
619
620    /// MoE expert-parallel size for AIC latency prediction (e.g., 4 for pure EP).
621    /// Required for MoE models; must satisfy: aic_tp_size * aic_attention_dp_size == aic_moe_tp_size * aic_moe_ep_size.
622    #[serde(skip)]
623    #[builder(default = "None")]
624    pub aic_moe_ep_size: Option<usize>,
625
626    /// Attention data-parallel size for AIC latency prediction (default: 1).
627    /// Corresponds to the `dp` dimension in AIC CLI output.
628    /// Must satisfy: aic_tp_size * aic_attention_dp_size == aic_moe_tp_size * aic_moe_ep_size.
629    #[serde(skip)]
630    #[builder(default = "None")]
631    pub aic_attention_dp_size: Option<usize>,
632
633    /// Weight dtype override for AIC latency prediction.
634    #[serde(skip)]
635    #[builder(default = "None")]
636    pub aic_gemm_dtype: Option<String>,
637
638    /// MoE kernel dtype override for AIC latency prediction.
639    #[serde(skip)]
640    #[builder(default = "None")]
641    pub aic_moe_dtype: Option<String>,
642
643    /// Activation dtype override for AIC latency prediction.
644    #[serde(skip)]
645    #[builder(default = "None")]
646    pub aic_fmha_dtype: Option<String>,
647
648    /// KV-cache dtype override for AIC latency prediction.
649    #[serde(skip)]
650    #[builder(default = "None")]
651    pub aic_kv_cache_dtype: Option<String>,
652
653    /// Communication (collective) dtype override for AIC latency prediction.
654    #[serde(skip)]
655    #[builder(default = "None")]
656    pub aic_comm_dtype: Option<String>,
657
658    /// MTP/Eagle speculative-decoding draft-token count (1..=5).
659    /// The mocker samples accepted drafts while AIC supplies undiscounted
660    /// verification-round latency.
661    #[builder(default = "None")]
662    #[validate(range(min = 1, max = 5))]
663    pub aic_nextn: Option<usize>,
664
665    /// Conditional acceptance rates for draft tokens, comma-separated.
666    /// Entry i is P(draft i accepted | every earlier draft was accepted).
667    #[builder(default = "None")]
668    pub aic_nextn_accept_rates: Option<String>,
669
670    /// Base RNG seed for MTP burst sampling. Worker rank is added with
671    /// wrapping arithmetic before constructing each worker-local sampler.
672    #[builder(default = "42")]
673    pub aic_mtp_seed: u64,
674
675    /// GPU memory fraction for AIC KV capacity estimation with vLLM.
676    #[builder(default = "None")]
677    #[validate(range(min = 0.0, max = 1.0))]
678    pub gpu_memory_utilization: Option<f64>,
679
680    /// Static memory fraction for AIC KV capacity estimation with SGLang.
681    #[builder(default = "None")]
682    #[validate(range(min = 0.0, max = 1.0))]
683    pub mem_fraction_static: Option<f64>,
684
685    /// Fraction of *free* GPU memory (after weights/buffers) allocated to the KV
686    /// cache, for AIC KV capacity estimation with TRT-LLM. Mirrors TRT-LLM's
687    /// `KvCacheConfig.free_gpu_memory_fraction`. Unlike vLLM's
688    /// `gpu_memory_utilization` (a fraction of *total* memory), this is a
689    /// fraction of what remains after the model is loaded.
690    #[builder(default = "None")]
691    #[validate(range(min = 0.0, max = 1.0))]
692    pub free_gpu_memory_fraction: Option<f64>,
693
694    /// Enable worker-local KV indexer for tracking this worker's own KV cache state
695    #[builder(default = "false")]
696    pub enable_local_indexer: bool,
697
698    /// Bootstrap port for disaggregated serving rendezvous.
699    /// Prefill workers listen on this port; decode workers connect to it.
700    /// If None, bootstrap rendezvous is disabled.
701    #[builder(default = "None")]
702    pub bootstrap_port: Option<u16>,
703
704    /// Absolute live handoff session timeout, excluding modeled transfer delay.
705    #[builder(default = "300_000")]
706    #[validate(range(min = 1))]
707    pub handoff_session_timeout_ms: u64,
708
709    /// Bytes transferred per token, auto-computed from model config by Python CLI.
710    /// Formula: num_layers * 2 * num_kv_heads * head_dim * dtype_bytes
711    #[builder(default = "None")]
712    pub kv_bytes_per_token: Option<usize>,
713
714    /// Physical KV-cache bytes occupied by one token, independent of transfer geometry.
715    #[builder(default = "None")]
716    #[serde(skip_serializing_if = "Option::is_none")]
717    pub kv_cache_bytes_per_token: Option<usize>,
718
719    /// KV cache transfer bandwidth in GB/s for disaggregated serving latency simulation.
720    /// Default: 64.0 (inter-node InfiniBand). Set to 0 to disable KV transfer delay.
721    /// For intra-node NVLink, typical value is ~450.
722    #[builder(default = "None")]
723    #[validate(range(min = 0.0))]
724    pub kv_transfer_bandwidth: Option<f64>,
725
726    /// Selects whether disaggregated transfer timing charges the full prompt
727    /// or only the physical prompt footprint missing at the destination.
728    #[builder(default = "KvTransferTimingMode::FullPrompt")]
729    pub kv_transfer_timing_mode: KvTransferTimingMode,
730
731    /// Reasoning/thinking token configuration.
732    /// When set, the mocker wraps output in thinking boundary tokens.
733    #[builder(default = "None")]
734    pub reasoning: Option<ReasoningConfig>,
735
736    /// Optional Mooncake trace with exact output token IDs keyed by
737    /// `output_replay_id` annotations. Direct replay paths carry the same token
738    /// IDs on `DirectRequest` and do not need this lookup.
739    #[builder(default = "None")]
740    pub response_replay_trace_path: Option<PathBuf>,
741
742    /// ZMQ port for publishing KV events in vLLM's native wire format.
743    /// When set, the scheduler publishes to a ZMQ PUB socket instead of directly to NATS.
744    /// A KvEventPublisher relay subscribes to this socket and forwards events to NATS.
745    #[builder(default = "None")]
746    pub zmq_kv_events_port: Option<u16>,
747
748    /// ZMQ ROUTER port for replay of buffered KV event batches.
749    /// When set alongside `zmq_kv_events_port`, the mocker binds a ROUTER socket
750    /// that streams back buffered batches by sequence number on request.
751    /// Port is offset by dp_rank (replay_port + dp_rank).
752    #[builder(default = "None")]
753    pub zmq_replay_port: Option<u16>,
754
755    /// Preemption mode for decode eviction under memory pressure.
756    /// Lifo (default) evicts the newest request; Fifo evicts the oldest.
757    #[builder(default)]
758    pub preemption_mode: PreemptionMode,
759
760    /// Optional replay-only override for the router queue policy.
761    #[builder(default = "None")]
762    pub router_queue_policy: Option<RouterQueuePolicy>,
763
764    /// SGLang-specific configuration. Only used when `engine_type == Sglang`.
765    #[builder(default = "None")]
766    pub sglang: Option<SglangArgs>,
767
768    /// TensorRT-LLM-specific configuration. Only used when `engine_type == Trtllm`.
769    #[builder(default = "None")]
770    pub trtllm: Option<TrtllmArgs>,
771}
772
773fn mock_engine_args_validation_error(code: &'static str, message: String) -> ValidationError {
774    let mut error = ValidationError::new(code);
775    error.message = Some(message.into());
776    error
777}
778
779fn validate_mock_engine_args(args: &MockEngineArgs) -> Result<(), ValidationError> {
780    if args.block_size == 0 {
781        return Err(mock_engine_args_validation_error(
782            "block_size_zero",
783            "block_size must be greater than 0".to_string(),
784        ));
785    }
786
787    if matches!(args.engine_type, EngineType::Vllm | EngineType::Trtllm) && args.block_size < 2 {
788        return Err(mock_engine_args_validation_error(
789            "shared_scheduler_block_size_too_small",
790            format!(
791                "the vLLM/TRT-LLM scheduler requires block_size to be at least 2 for engine_type={:?}, got block_size={}",
792                args.engine_type, args.block_size,
793            ),
794        ));
795    }
796
797    if args.max_model_len.is_some() && args.engine_type != EngineType::Vllm {
798        return Err(mock_engine_args_validation_error(
799            "max_model_len_requires_vllm",
800            format!(
801                "max_model_len is supported only for engine_type=vllm, got engine_type={:?}",
802                args.engine_type
803            ),
804        ));
805    }
806    if args.aic_nextn.is_some() && args.decode_speedup_ratio != 1.0 {
807        return Err(mock_engine_args_validation_error(
808            "mtp_decode_speedup_conflict",
809            format!(
810                "aic_nextn requires decode_speedup_ratio=1.0 because MTP output acceleration is modeled by burst sampling, got {}",
811                args.decode_speedup_ratio
812            ),
813        ));
814    }
815
816    if args.aic_nextn.is_none() && args.aic_nextn_accept_rates.is_some() {
817        return Err(mock_engine_args_validation_error(
818            "mtp_rates_without_nextn",
819            "aic_nextn_accept_rates requires aic_nextn".to_string(),
820        ));
821    }
822
823    if let Some(policy) = args
824        .trtllm
825        .as_ref()
826        .and_then(|trtllm| trtllm.capacity_scheduler_policy.as_deref())
827        && policy != "guaranteed_no_evict"
828    {
829        return Err(mock_engine_args_validation_error(
830            "trtllm_unsupported_capacity_scheduler_policy",
831            format!(
832                "engine_type=trtllm v1 supports only capacity_scheduler_policy='guaranteed_no_evict', got '{policy}'",
833            ),
834        ));
835    }
836
837    if args.engine_type != EngineType::Sglang {
838        return Ok(());
839    }
840
841    if let Some(page_size) = args.sglang.as_ref().and_then(|sglang| sglang.page_size)
842        && args.block_size != page_size
843    {
844        return Err(mock_engine_args_validation_error(
845            "sglang_block_size_page_size_mismatch",
846            format!(
847                "engine_type=sglang requires block_size and sglang.page_size to match when both are set, got block_size={} and sglang.page_size={page_size}",
848                args.block_size,
849            ),
850        ));
851    }
852
853    if let Some(chunked_prefill_size) = args
854        .sglang
855        .as_ref()
856        .and_then(|sglang| sglang.chunked_prefill_size)
857        && chunked_prefill_size % args.block_size != 0
858    {
859        return Err(mock_engine_args_validation_error(
860            "sglang_chunked_prefill_size_not_divisible_by_block_size",
861            format!(
862                "engine_type=sglang requires sglang.chunked_prefill_size to be divisible by block_size, got chunked_prefill_size={} and block_size={}",
863                chunked_prefill_size, args.block_size,
864            ),
865        ));
866    }
867
868    Ok(())
869}
870
871impl TryFrom<MockEngineArgsSerde> for MockEngineArgs {
872    type Error = String;
873
874    fn try_from(compat: MockEngineArgsSerde) -> Result<Self, Self::Error> {
875        let mut builder = Self::builder();
876
877        if let Some(engine_type) = compat.engine_type.into_non_null("engine_type")? {
878            builder = builder.engine_type(engine_type.parse()?);
879        }
880        if let Some(Some(num_gpu_blocks)) = compat.num_gpu_blocks.into_nullable() {
881            builder = builder.num_gpu_blocks(num_gpu_blocks);
882        }
883        if let Some(block_size) = compat.block_size.into_non_null("block_size")? {
884            builder = builder.block_size(block_size);
885        }
886        if let Some(max_model_len) = compat.max_model_len.into_nullable() {
887            builder = builder.max_model_len(max_model_len);
888        }
889        if let Some(max_num_seqs) = compat.max_num_seqs.into_nullable() {
890            builder = builder.max_num_seqs(max_num_seqs);
891        }
892        if let Some(max_num_batched_tokens) = compat.max_num_batched_tokens.into_nullable() {
893            builder = builder.max_num_batched_tokens(max_num_batched_tokens);
894        }
895        if let Some(enable_prefix_caching) = compat
896            .enable_prefix_caching
897            .into_non_null("enable_prefix_caching")?
898        {
899            builder = builder.enable_prefix_caching(enable_prefix_caching);
900        }
901        if let Some(enable_chunked_prefill) = compat
902            .enable_chunked_prefill
903            .into_non_null("enable_chunked_prefill")?
904        {
905            builder = builder.enable_chunked_prefill(enable_chunked_prefill);
906        }
907        if let Some(speedup_ratio) = compat.speedup_ratio.into_non_null("speedup_ratio")? {
908            builder = builder.speedup_ratio(speedup_ratio);
909        }
910        if let Some(decode_speedup_ratio) = compat
911            .decode_speedup_ratio
912            .into_non_null("decode_speedup_ratio")?
913        {
914            builder = builder.decode_speedup_ratio(decode_speedup_ratio);
915        }
916        if let Some(dp_size) = compat.dp_size.into_non_null("dp_size")? {
917            builder = builder.dp_size(dp_size);
918        }
919        if let Some(startup_time) = compat.startup_time.into_nullable() {
920            builder = builder.startup_time(startup_time);
921        }
922
923        let worker_type = if let Some(worker_type) =
924            compat.worker_type.into_non_null("worker_type")?
925        {
926            worker_type.parse()?
927        } else {
928            let is_prefill = compat
929                .is_prefill
930                .into_non_null("is_prefill")?
931                .unwrap_or(false);
932            let is_decode = compat
933                .is_decode
934                .into_non_null("is_decode")?
935                .unwrap_or(false);
936
937            match (is_prefill, is_decode) {
938                (false, false) => WorkerType::Aggregated,
939                (true, false) => WorkerType::Prefill,
940                (false, true) => WorkerType::Decode,
941                (true, true) => {
942                    return Err(
943                        "Invalid worker configuration: is_prefill and is_decode cannot both be true."
944                            .to_string(),
945                    );
946                }
947            }
948        };
949        builder = builder.worker_type(worker_type);
950
951        if let Some(planner_profile_data) = compat.planner_profile_data.into_nullable() {
952            builder = builder.planner_profile_data(planner_profile_data.clone());
953            if let Some(path) = planner_profile_data {
954                builder = builder.perf_model(load_perf_model(&path));
955            }
956        }
957        if let Some(timing_model) = compat.timing_model.into_nullable().flatten() {
958            let perf_model = match timing_model {
959                TimingModelSerde::Default | TimingModelSerde::Polynomial => PerfModel::Polynomial,
960                TimingModelSerde::Fixed {
961                    prefill_ms,
962                    decode_ms,
963                } => {
964                    if !prefill_ms.is_finite()
965                        || prefill_ms < 0.0
966                        || !decode_ms.is_finite()
967                        || decode_ms < 0.0
968                    {
969                        return Err(
970                            "fixed timing prefill_ms and decode_ms must be finite and nonnegative"
971                                .to_string(),
972                        );
973                    }
974                    PerfModel::Fixed {
975                        prefill_ms,
976                        decode_ms,
977                    }
978                }
979            };
980            builder = builder.perf_model(Arc::new(perf_model));
981        }
982
983        if let Some(aic_backend) = compat.aic_backend.into_nullable() {
984            builder = builder.aic_backend(aic_backend);
985        }
986        if let Some(aic_system) = compat.aic_system.into_nullable() {
987            builder = builder.aic_system(aic_system);
988        }
989        if let Some(aic_backend_version) = compat.aic_backend_version.into_nullable() {
990            builder = builder.aic_backend_version(aic_backend_version);
991        }
992        if let Some(aic_tp_size) = compat.aic_tp_size.into_nullable() {
993            builder = builder.aic_tp_size(aic_tp_size);
994        }
995        if let Some(aic_model_path) = compat.aic_model_path.into_nullable() {
996            builder = builder.aic_model_path(aic_model_path);
997        }
998        if let Some(aic_moe_tp_size) = compat.aic_moe_tp_size.into_nullable() {
999            builder = builder.aic_moe_tp_size(aic_moe_tp_size);
1000        }
1001        if let Some(aic_moe_ep_size) = compat.aic_moe_ep_size.into_nullable() {
1002            builder = builder.aic_moe_ep_size(aic_moe_ep_size);
1003        }
1004        if let Some(aic_attention_dp_size) = compat.aic_attention_dp_size.into_nullable() {
1005            builder = builder.aic_attention_dp_size(aic_attention_dp_size);
1006        }
1007        if let Some(aic_gemm_dtype) = compat.aic_gemm_dtype.into_nullable() {
1008            builder = builder.aic_gemm_dtype(aic_gemm_dtype);
1009        }
1010        if let Some(aic_moe_dtype) = compat.aic_moe_dtype.into_nullable() {
1011            builder = builder.aic_moe_dtype(aic_moe_dtype);
1012        }
1013        if let Some(aic_fmha_dtype) = compat.aic_fmha_dtype.into_nullable() {
1014            builder = builder.aic_fmha_dtype(aic_fmha_dtype);
1015        }
1016        if let Some(aic_kv_cache_dtype) = compat.aic_kv_cache_dtype.into_nullable() {
1017            builder = builder.aic_kv_cache_dtype(aic_kv_cache_dtype);
1018        }
1019        if let Some(aic_comm_dtype) = compat.aic_comm_dtype.into_nullable() {
1020            builder = builder.aic_comm_dtype(aic_comm_dtype);
1021        }
1022        if let Some(aic_nextn) = compat.aic_nextn.into_nullable() {
1023            builder = builder.aic_nextn(aic_nextn);
1024        }
1025        if let Some(aic_nextn_accept_rates) = compat.aic_nextn_accept_rates.into_nullable() {
1026            builder = builder.aic_nextn_accept_rates(aic_nextn_accept_rates);
1027        }
1028        if let Some(aic_mtp_seed) = compat.aic_mtp_seed.into_non_null("aic_mtp_seed")? {
1029            builder = builder.aic_mtp_seed(aic_mtp_seed);
1030        }
1031        if let Some(gpu_memory_utilization) = compat.gpu_memory_utilization.into_nullable() {
1032            builder = builder.gpu_memory_utilization(gpu_memory_utilization);
1033        }
1034        if let Some(mem_fraction_static) = compat.mem_fraction_static.into_nullable() {
1035            builder = builder.mem_fraction_static(mem_fraction_static);
1036        }
1037        if let Some(free_gpu_memory_fraction) = compat.free_gpu_memory_fraction.into_nullable() {
1038            builder = builder.free_gpu_memory_fraction(free_gpu_memory_fraction);
1039        }
1040        if let Some(enable_local_indexer) = compat
1041            .enable_local_indexer
1042            .into_non_null("enable_local_indexer")?
1043        {
1044            builder = builder.enable_local_indexer(enable_local_indexer);
1045        }
1046        if let Some(bootstrap_port) = compat.bootstrap_port.into_nullable() {
1047            builder = builder.bootstrap_port(bootstrap_port);
1048        }
1049        if let Some(timeout_ms) = compat
1050            .handoff_session_timeout_ms
1051            .into_non_null("handoff_session_timeout_ms")?
1052        {
1053            builder = builder.handoff_session_timeout_ms(timeout_ms);
1054        }
1055        if let Some(kv_bytes_per_token) = compat.kv_bytes_per_token.into_nullable() {
1056            builder = builder.kv_bytes_per_token(kv_bytes_per_token);
1057        }
1058        if let Some(kv_cache_bytes_per_token) = compat.kv_cache_bytes_per_token.into_nullable() {
1059            builder = builder.kv_cache_bytes_per_token(kv_cache_bytes_per_token);
1060        }
1061        if let Some(kv_transfer_bandwidth) = compat.kv_transfer_bandwidth.into_nullable() {
1062            builder = builder.kv_transfer_bandwidth(kv_transfer_bandwidth);
1063        }
1064        if let Some(mode) = compat
1065            .kv_transfer_timing_mode
1066            .into_non_null("kv_transfer_timing_mode")?
1067        {
1068            builder = builder.kv_transfer_timing_mode(mode.parse()?);
1069        }
1070        if let Some(reasoning) = compat.reasoning.into_nullable() {
1071            builder = builder.reasoning(reasoning);
1072        }
1073        if let Some(response_replay_trace_path) = compat.response_replay_trace_path.into_nullable()
1074        {
1075            builder = builder.response_replay_trace_path(response_replay_trace_path);
1076        }
1077        if let Some(zmq_kv_events_port) = compat.zmq_kv_events_port.into_nullable() {
1078            builder = builder.zmq_kv_events_port(zmq_kv_events_port);
1079        }
1080        if let Some(zmq_replay_port) = compat.zmq_replay_port.into_nullable() {
1081            builder = builder.zmq_replay_port(zmq_replay_port);
1082        }
1083        if let Some(preemption_mode) = compat.preemption_mode.into_non_null("preemption_mode")? {
1084            builder = builder.preemption_mode(preemption_mode.parse()?);
1085        }
1086        if let Some(router_queue_policy) = compat.router_queue_policy.into_nullable() {
1087            let router_queue_policy = router_queue_policy
1088                .map(|policy| policy.parse().map_err(|e: String| e))
1089                .transpose()?;
1090            builder = builder.router_queue_policy(router_queue_policy);
1091        }
1092        if let Some(sglang) = compat.sglang.into_nullable() {
1093            builder = builder.sglang(sglang);
1094        }
1095        if let Some(trtllm) = compat.trtllm.into_nullable() {
1096            builder = builder.trtllm(trtllm);
1097        }
1098
1099        builder
1100            .build()
1101            .map_err(|e| format!("Failed to build MockEngineArgs: {e}"))?
1102            .normalized()
1103            .map_err(|e| e.to_string())
1104    }
1105}
1106
1107impl Default for MockEngineArgs {
1108    fn default() -> MockEngineArgs {
1109        MockEngineArgsBuilder::default()
1110            .build()
1111            .expect("Failed to build default MockEngineArgs")
1112            .normalized()
1113            .expect("Failed to normalize default MockEngineArgs")
1114    }
1115}
1116
1117impl MockEngineArgs {
1118    const DEFAULT_VLLM_BLOCK_SIZE: usize = 64;
1119    const DEFAULT_SGLANG_BLOCK_SIZE: usize = 1;
1120    const DEFAULT_TRTLLM_BLOCK_SIZE: usize = 32;
1121
1122    pub fn builder() -> MockEngineArgsBuilder {
1123        MockEngineArgsBuilder::default()
1124    }
1125
1126    /// GPUs occupied by one worker (engine), derived from tensor parallelism
1127    /// and the materialized DP topology. AIC-backed replay uses
1128    /// `aic_tp_size × aic_attention_dp_size`; non-AIC replay still counts one
1129    /// GPU for every independently modeled `dp_size` rank. Used to turn
1130    /// provisioned worker-seconds into GPU-hours.
1131    pub fn aic_gpus_per_worker(&self) -> usize {
1132        self.aic_tp_size.unwrap_or(1) * self.dp_size.max(1) as usize
1133    }
1134
1135    /// Finite ownership bound for live handoff queues and sessions.
1136    ///
1137    /// An unset runnable-sequence limit is semantically unbounded, so use the
1138    /// physical KV block count as the conservative process-local bound.
1139    pub fn effective_handoff_capacity(&self) -> usize {
1140        self.max_num_seqs.unwrap_or(self.num_gpu_blocks).max(1)
1141    }
1142
1143    pub fn normalized(mut self) -> anyhow::Result<Self> {
1144        self.materialize_defaults();
1145        self.validate_config()?;
1146        Ok(self)
1147    }
1148
1149    fn materialize_defaults(&mut self) {
1150        match self.engine_type {
1151            EngineType::Vllm => {
1152                if self.block_size == 0 {
1153                    self.block_size = Self::DEFAULT_VLLM_BLOCK_SIZE;
1154                }
1155            }
1156            EngineType::Sglang => {
1157                let page_size = self.sglang.as_ref().and_then(|sglang| sglang.page_size);
1158                match (self.block_size, page_size) {
1159                    (0, None) => {
1160                        self.block_size = Self::DEFAULT_SGLANG_BLOCK_SIZE;
1161                    }
1162                    (0, Some(page_size)) => {
1163                        self.block_size = page_size;
1164                    }
1165                    (_, Some(_)) => {}
1166                    (_, None) => {}
1167                }
1168            }
1169            EngineType::Trtllm => {
1170                if self.block_size == 0 {
1171                    self.block_size = Self::DEFAULT_TRTLLM_BLOCK_SIZE;
1172                }
1173            }
1174        }
1175    }
1176
1177    fn validate_config(&mut self) -> anyhow::Result<()> {
1178        self.validate()
1179            .map_err(|error| anyhow::anyhow!("Failed to validate MockEngineArgs: {error}"))?;
1180        if let Some(nextn) = self.aic_nextn {
1181            let rates = crate::common::speculative::normalize_conditional_accept_rates(
1182                nextn,
1183                self.aic_nextn_accept_rates.as_deref(),
1184            )?;
1185            self.aic_nextn_accept_rates =
1186                Some(crate::common::speculative::format_accept_rates(&rates));
1187        }
1188        Ok(())
1189    }
1190
1191    pub fn is_prefill(&self) -> bool {
1192        self.worker_type == WorkerType::Prefill
1193    }
1194
1195    pub fn is_decode(&self) -> bool {
1196        self.worker_type == WorkerType::Decode
1197    }
1198
1199    pub fn needs_kv_publisher(&self) -> bool {
1200        self.enable_prefix_caching && !self.is_decode()
1201    }
1202
1203    pub fn undiscounted_aic_accept_rates(&self) -> Option<String> {
1204        crate::common::speculative::undiscounted_aic_accept_rates(self.aic_nextn)
1205    }
1206
1207    /// Create MockEngineArgs from a JSON file containing extra engine arguments
1208    pub fn from_json_file(path: &Path) -> anyhow::Result<Self> {
1209        let file_content = std::fs::read_to_string(path)?;
1210        Self::from_json_str(&file_content)
1211    }
1212
1213    pub fn from_json_str(content: &str) -> anyhow::Result<Self> {
1214        let mut deserializer = serde_json::Deserializer::from_str(content);
1215        let args = serde_path_to_error::deserialize(&mut deserializer)
1216            .map_err(|error| anyhow::anyhow!("{error}"))?;
1217        deserializer
1218            .end()
1219            .map_err(|error| anyhow::anyhow!("{error}"))?;
1220        Ok(args)
1221    }
1222}
1223
1224#[cfg(test)]
1225mod tests {
1226    use std::sync::Mutex;
1227
1228    use super::*;
1229    use serde_json::json;
1230
1231    #[derive(Default)]
1232    struct FailingRawSink {
1233        attempts: Mutex<Vec<u64>>,
1234    }
1235
1236    impl RawKvEventSink for FailingRawSink {
1237        fn publish(&self, event: RawKvEvent) -> anyhow::Result<()> {
1238            self.attempts.lock().unwrap().push(event.event.event_id);
1239            if event.event.event_id == 2 {
1240                anyhow::bail!("injected raw sink failure");
1241            }
1242            Ok(())
1243        }
1244    }
1245
1246    #[test]
1247    fn raw_sink_batch_fallback_attempts_later_events_after_failure() {
1248        let sink = FailingRawSink::default();
1249        let error = sink
1250            .publish_batch(
1251                (1..=3)
1252                    .map(|event_id| RawKvEvent {
1253                        event: KvCacheEvent {
1254                            event_id,
1255                            data: dynamo_kv_router::protocols::KvCacheEventData::Cleared,
1256                            dp_rank: 0,
1257                        },
1258                        block_token_ids: None,
1259                        storage_tier: StorageTier::Device,
1260                    })
1261                    .collect(),
1262            )
1263            .unwrap_err();
1264
1265        assert_eq!(error.to_string(), "injected raw sink failure");
1266        assert_eq!(*sink.attempts.lock().unwrap(), vec![1, 2, 3]);
1267    }
1268
1269    #[test]
1270    fn direct_request_priorities_are_backward_compatible() {
1271        let legacy = json!({
1272            "tokens": [1, 2],
1273            "max_output_tokens": 3,
1274            "uuid": null,
1275            "dp_rank": 0,
1276            "arrival_timestamp_ms": null
1277        });
1278        let request: DirectRequest = serde_json::from_value(legacy).unwrap();
1279        assert_eq!(request.priority, 0);
1280        assert_eq!(request.strict_priority, 0);
1281        assert_eq!(request.router_priorities(), (0.0, 0));
1282
1283        let rendered = serde_json::to_value(&request).unwrap();
1284        assert!(rendered.get("priority").is_none());
1285        assert!(rendered.get("strict_priority").is_none());
1286    }
1287
1288    #[test]
1289    fn direct_request_derives_router_priorities() {
1290        let negative: DirectRequest = serde_json::from_value(json!({
1291            "tokens": [1],
1292            "max_output_tokens": 1,
1293            "uuid": null,
1294            "dp_rank": 0,
1295            "arrival_timestamp_ms": null,
1296            "priority": -7,
1297            "strict_priority": 4
1298        }))
1299        .unwrap();
1300        assert_eq!(negative.router_priorities(), (0.0, 4));
1301
1302        let positive = DirectRequest {
1303            priority: 9,
1304            strict_priority: 5,
1305            ..negative
1306        };
1307        assert_eq!(positive.router_priorities(), (9.0, 5));
1308        let rendered = serde_json::to_value(&positive).unwrap();
1309        assert_eq!(rendered["priority"], 9);
1310        assert_eq!(rendered["strict_priority"], 5);
1311    }
1312
1313    #[test]
1314    fn test_mock_engine_args_json_round_trip_preserves_worker_type_and_nulls() {
1315        let args = MockEngineArgs::builder()
1316            .worker_type(WorkerType::Decode)
1317            .max_model_len(Some(32768))
1318            .max_num_seqs(None)
1319            .max_num_batched_tokens(None)
1320            .reasoning(None)
1321            .sglang(None)
1322            .build()
1323            .unwrap()
1324            .normalized()
1325            .unwrap();
1326
1327        let mut payload = serde_json::json!({
1328            "engine_type": "vllm",
1329            "num_gpu_blocks": args.num_gpu_blocks,
1330            "block_size": args.block_size,
1331            "max_num_seqs": args.max_num_seqs,
1332            "max_num_batched_tokens": args.max_num_batched_tokens,
1333            "enable_prefix_caching": args.enable_prefix_caching,
1334            "enable_chunked_prefill": args.enable_chunked_prefill,
1335            "speedup_ratio": args.speedup_ratio,
1336            "decode_speedup_ratio": args.decode_speedup_ratio,
1337            "dp_size": args.dp_size,
1338            "startup_time": args.startup_time,
1339            "worker_type": "decode",
1340            "planner_profile_data": args.planner_profile_data,
1341            "aic_backend": args.aic_backend,
1342            "aic_system": args.aic_system,
1343            "aic_backend_version": args.aic_backend_version,
1344            "aic_tp_size": args.aic_tp_size,
1345            "aic_model_path": args.aic_model_path,
1346            "enable_local_indexer": args.enable_local_indexer,
1347            "bootstrap_port": args.bootstrap_port,
1348            "handoff_session_timeout_ms": args.handoff_session_timeout_ms,
1349            "kv_bytes_per_token": args.kv_bytes_per_token,
1350            "kv_transfer_bandwidth": args.kv_transfer_bandwidth,
1351            "kv_transfer_timing_mode": "full_prompt",
1352            "reasoning": args.reasoning,
1353            "zmq_kv_events_port": args.zmq_kv_events_port,
1354            "zmq_replay_port": args.zmq_replay_port,
1355            "preemption_mode": "lifo",
1356            "router_queue_policy": args.router_queue_policy.map(|policy| policy.to_string()),
1357            "sglang": args.sglang,
1358            "has_perf_model": true,
1359        });
1360        payload["max_model_len"] = serde_json::json!(args.max_model_len);
1361
1362        let restored = MockEngineArgs::from_json_str(&payload.to_string()).unwrap();
1363
1364        assert_eq!(restored.worker_type, WorkerType::Decode);
1365        assert_eq!(restored.max_model_len, Some(32768));
1366        assert_eq!(restored.max_num_seqs, None);
1367        assert_eq!(restored.max_num_batched_tokens, None);
1368        assert_eq!(
1369            restored.kv_transfer_timing_mode,
1370            KvTransferTimingMode::FullPrompt
1371        );
1372    }
1373
1374    #[test]
1375    fn test_mock_engine_args_accepts_legacy_enum_case_and_writes_lowercase() {
1376        let args = MockEngineArgs::from_json_str(
1377            &json!({
1378                "engine_type": "VLLM",
1379                "worker_type": "Aggregated",
1380                "preemption_mode": "Lifo",
1381            })
1382            .to_string(),
1383        )
1384        .unwrap();
1385
1386        let serialized = serde_json::to_value(args).unwrap();
1387        assert_eq!(serialized["engine_type"], "vllm");
1388        assert_eq!(serialized["worker_type"], "aggregated");
1389        assert_eq!(serialized["preemption_mode"], "lifo");
1390    }
1391
1392    #[test]
1393    fn test_mock_engine_args_json_accepts_aic_quant_dtypes() {
1394        let args = MockEngineArgs::from_json_str(
1395            &json!({
1396                "aic_gemm_dtype": "fp8_block",
1397                "aic_moe_dtype": "w4a16_mxfp4",
1398                "aic_fmha_dtype": "bfloat16",
1399                "aic_kv_cache_dtype": "fp8",
1400                "aic_comm_dtype": "fp8",
1401            })
1402            .to_string(),
1403        )
1404        .unwrap();
1405
1406        assert_eq!(args.aic_gemm_dtype.as_deref(), Some("fp8_block"));
1407        assert_eq!(args.aic_moe_dtype.as_deref(), Some("w4a16_mxfp4"));
1408        assert_eq!(args.aic_fmha_dtype.as_deref(), Some("bfloat16"));
1409        assert_eq!(args.aic_kv_cache_dtype.as_deref(), Some("fp8"));
1410        assert_eq!(args.aic_comm_dtype.as_deref(), Some("fp8"));
1411    }
1412
1413    #[test]
1414    fn test_mock_engine_args_json_rejects_unknown_and_invalid_types() {
1415        let unknown = MockEngineArgs::from_json_str(&json!({"unknown": true}).to_string())
1416            .expect_err("unknown fields should be rejected");
1417        assert!(
1418            unknown.to_string().contains("unknown field"),
1419            "unexpected error: {unknown}",
1420        );
1421
1422        let invalid =
1423            MockEngineArgs::from_json_str(&json!({"gpu_memory_utilization": "bad"}).to_string())
1424                .expect_err("wrongly typed fields should be rejected");
1425        assert!(
1426            invalid.to_string().contains("gpu_memory_utilization"),
1427            "unexpected error: {invalid}",
1428        );
1429
1430        let trailing = MockEngineArgs::from_json_str(r#"{"block_size": 16} true"#)
1431            .expect_err("trailing JSON should be rejected");
1432        assert!(
1433            trailing.to_string().contains("trailing characters"),
1434            "unexpected error: {trailing}",
1435        );
1436    }
1437
1438    #[test]
1439    fn test_normalized_sglang_uses_page_size_alias_for_block_size() {
1440        let args = MockEngineArgs::builder()
1441            .engine_type(EngineType::Sglang)
1442            .sglang(Some(SglangArgs {
1443                page_size: Some(16),
1444                ..Default::default()
1445            }))
1446            .build()
1447            .unwrap()
1448            .normalized()
1449            .unwrap();
1450
1451        assert_eq!(args.block_size, 16);
1452    }
1453
1454    #[test]
1455    fn test_normalized_sglang_accepts_equal_block_size_and_page_size() {
1456        let args = MockEngineArgs::builder()
1457            .engine_type(EngineType::Sglang)
1458            .block_size(8)
1459            .sglang(Some(SglangArgs {
1460                page_size: Some(8),
1461                ..Default::default()
1462            }))
1463            .build()
1464            .unwrap()
1465            .normalized()
1466            .unwrap();
1467
1468        assert_eq!(args.block_size, 8);
1469    }
1470
1471    #[test]
1472    fn test_normalized_sglang_rejects_mismatched_block_size_and_page_size() {
1473        let error = MockEngineArgs::builder()
1474            .engine_type(EngineType::Sglang)
1475            .block_size(8)
1476            .sglang(Some(SglangArgs {
1477                page_size: Some(4),
1478                ..Default::default()
1479            }))
1480            .build()
1481            .unwrap()
1482            .normalized()
1483            .unwrap_err();
1484
1485        assert!(
1486            error
1487                .to_string()
1488                .contains("block_size and sglang.page_size to match"),
1489            "unexpected error: {error}",
1490        );
1491    }
1492
1493    #[test]
1494    fn test_normalized_rejects_out_of_range_aic_nextn() {
1495        // The mocker/replay JSON path must share AicPerfConfig's 1..=5 contract.
1496        for bad in [0_usize, 6, usize::MAX] {
1497            let err = MockEngineArgs::builder()
1498                .aic_nextn(Some(bad))
1499                .build()
1500                .unwrap()
1501                .normalized()
1502                .unwrap_err();
1503            assert!(
1504                err.to_string().contains("aic_nextn"),
1505                "unexpected error for nextn={bad}: {err}",
1506            );
1507        }
1508        MockEngineArgs::builder()
1509            .aic_nextn(Some(3))
1510            .build()
1511            .unwrap()
1512            .normalized()
1513            .expect("in-range aic_nextn should validate");
1514    }
1515
1516    #[test]
1517    fn test_normalized_rejects_zero_max_model_len() {
1518        let error = MockEngineArgs::builder()
1519            .max_model_len(Some(0))
1520            .build()
1521            .unwrap()
1522            .normalized()
1523            .unwrap_err();
1524
1525        assert!(
1526            error.to_string().contains("max_model_len"),
1527            "unexpected error: {error}",
1528        );
1529    }
1530
1531    #[test]
1532    fn test_mtp_defaults_and_json_round_trip() {
1533        let args = MockEngineArgs::builder()
1534            .aic_nextn(Some(3))
1535            .build()
1536            .unwrap()
1537            .normalized()
1538            .unwrap();
1539        assert_eq!(args.aic_nextn_accept_rates.as_deref(), Some("0.85,0.3,0"));
1540        assert_eq!(args.aic_mtp_seed, 42);
1541        assert_eq!(
1542            args.undiscounted_aic_accept_rates().as_deref(),
1543            Some("0,0,0")
1544        );
1545
1546        let json = serde_json::to_string(&args).unwrap();
1547        let round_trip = MockEngineArgs::from_json_str(&json).unwrap();
1548        assert_eq!(round_trip.aic_nextn, Some(3));
1549        assert_eq!(
1550            round_trip.aic_nextn_accept_rates.as_deref(),
1551            Some("0.85,0.3,0")
1552        );
1553        assert_eq!(round_trip.aic_mtp_seed, 42);
1554    }
1555
1556    #[test]
1557    fn test_mtp_rates_are_validated_before_normalization() {
1558        for rates in ["nan", "inf", "-0.1", "1.1", "bad"] {
1559            let err = MockEngineArgs::builder()
1560                .aic_nextn(Some(1))
1561                .aic_nextn_accept_rates(Some(rates.to_string()))
1562                .build()
1563                .unwrap()
1564                .normalized()
1565                .unwrap_err();
1566            assert!(
1567                err.to_string().contains("aic_nextn_accept_rates"),
1568                "unexpected error for rates={rates:?}: {err}"
1569            );
1570        }
1571    }
1572
1573    #[test]
1574    fn test_mtp_rates_are_padded_and_truncated_to_nextn() {
1575        let padded = MockEngineArgs::builder()
1576            .aic_nextn(Some(3))
1577            .aic_nextn_accept_rates(Some("1,0.5".to_string()))
1578            .build()
1579            .unwrap()
1580            .normalized()
1581            .unwrap();
1582        assert_eq!(padded.aic_nextn_accept_rates.as_deref(), Some("1,0.5,0"));
1583
1584        let truncated = MockEngineArgs::builder()
1585            .aic_nextn(Some(2))
1586            .aic_nextn_accept_rates(Some("1,0.5,0.25".to_string()))
1587            .build()
1588            .unwrap()
1589            .normalized()
1590            .unwrap();
1591        assert_eq!(truncated.aic_nextn_accept_rates.as_deref(), Some("1,0.5"));
1592    }
1593
1594    #[test]
1595    fn test_mtp_rejects_decode_speedup_ratio() {
1596        let err = MockEngineArgs::builder()
1597            .aic_nextn(Some(1))
1598            .decode_speedup_ratio(2.0)
1599            .build()
1600            .unwrap()
1601            .normalized()
1602            .unwrap_err();
1603        assert!(err.to_string().contains("decode_speedup_ratio=1.0"));
1604    }
1605
1606    #[test]
1607    fn test_normalized_sglang_defaults_block_size_to_one() {
1608        let args = MockEngineArgs::builder()
1609            .engine_type(EngineType::Sglang)
1610            .build()
1611            .unwrap()
1612            .normalized()
1613            .unwrap();
1614
1615        assert_eq!(args.block_size, 1);
1616    }
1617
1618    #[test]
1619    fn test_normalized_shared_scheduler_rejects_block_size_one_for_every_backend() {
1620        for engine_type in [EngineType::Vllm, EngineType::Trtllm] {
1621            let error = MockEngineArgs::builder()
1622                .engine_type(engine_type)
1623                .block_size(1)
1624                .build()
1625                .unwrap()
1626                .normalized()
1627                .unwrap_err();
1628            assert!(
1629                error
1630                    .to_string()
1631                    .contains("the vLLM/TRT-LLM scheduler requires block_size to be at least 2"),
1632                "engine_type={engine_type:?}, error={error:#}"
1633            );
1634        }
1635    }
1636
1637    #[test]
1638    fn test_normalized_sglang_accepts_block_size_one() {
1639        let args = MockEngineArgs::builder()
1640            .engine_type(EngineType::Sglang)
1641            .block_size(1)
1642            .build()
1643            .unwrap()
1644            .normalized()
1645            .unwrap();
1646
1647        assert_eq!(args.block_size, 1);
1648    }
1649
1650    #[test]
1651    fn test_from_json_file_normalizes_sglang_page_size() {
1652        let tempdir = tempfile::tempdir().unwrap();
1653        let path = tempdir.path().join("args.json");
1654        std::fs::write(
1655            &path,
1656            serde_json::to_string(&json!({
1657                "engine_type": "sglang",
1658                "sglang": {
1659                    "page_size": 32
1660                }
1661            }))
1662            .unwrap(),
1663        )
1664        .unwrap();
1665
1666        let args = MockEngineArgs::from_json_file(&path).unwrap();
1667        assert_eq!(args.block_size, 32);
1668    }
1669
1670    #[test]
1671    fn test_normalized_sglang_rejects_chunked_prefill_not_divisible_by_block_size() {
1672        let error = MockEngineArgs::builder()
1673            .engine_type(EngineType::Sglang)
1674            .block_size(4)
1675            .sglang(Some(SglangArgs {
1676                page_size: Some(4),
1677                chunked_prefill_size: Some(6),
1678                ..Default::default()
1679            }))
1680            .build()
1681            .unwrap()
1682            .normalized()
1683            .unwrap_err();
1684
1685        assert!(
1686            error
1687                .to_string()
1688                .contains("chunked_prefill_size to be divisible by block_size"),
1689            "unexpected error: {error}",
1690        );
1691    }
1692
1693    #[test]
1694    fn test_normalized_sglang_accepts_chunked_prefill_divisible_by_block_size() {
1695        let args = MockEngineArgs::builder()
1696            .engine_type(EngineType::Sglang)
1697            .block_size(4)
1698            .sglang(Some(SglangArgs {
1699                page_size: Some(4),
1700                chunked_prefill_size: Some(8),
1701                ..Default::default()
1702            }))
1703            .build()
1704            .unwrap()
1705            .normalized()
1706            .unwrap();
1707
1708        assert_eq!(args.block_size, 4);
1709    }
1710}