1use 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
17pub 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 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#[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
57pub trait RawKvEventSink: Send + Sync {
59 fn publish(&self, event: RawKvEvent) -> anyhow::Result<()>;
60
61 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#[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 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 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
154pub use aisimulate_core::replay::ForwardPassSnapshot;
156
157pub trait FpmSink: Send + Sync {
160 fn publish(&self, snapshot: ForwardPassSnapshot) -> anyhow::Result<()>;
161}
162
163#[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
183pub use aisimulate_core::replay::DirectRequest;
187
188#[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 pub completed: bool,
197 #[serde(default)]
201 pub rejected: bool,
202 #[serde(default, skip_serializing_if = "Option::is_none")]
203 pub handoff_delay_ms: Option<f64>,
204 #[serde(default, skip_serializing_if = "Option::is_none")]
207 pub cached_tokens: Option<usize>,
208}
209
210#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
212#[serde(rename_all = "lowercase")]
213pub enum PreemptionMode {
214 #[default]
216 Lifo,
217 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#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
237#[serde(rename_all = "lowercase")]
238pub enum EngineType {
239 #[default]
241 Vllm,
242 Sglang,
244 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#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
266#[serde(rename_all = "lowercase")]
267pub enum WorkerType {
268 #[default]
270 Aggregated,
271 Prefill,
273 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#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
294#[serde(rename_all = "snake_case")]
295pub enum KvTransferTimingMode {
296 #[default]
298 FullPrompt,
299 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#[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 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 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#[derive(Debug, Clone, Serialize, Deserialize, Validate, Default)]
354pub struct SglangArgs {
355 pub schedule_policy: Option<String>,
357 #[validate(range(min = 1))]
359 pub page_size: Option<usize>,
360 #[validate(range(min = 1))]
362 pub max_prefill_tokens: Option<usize>,
363 #[validate(range(min = 1))]
365 pub chunked_prefill_size: Option<usize>,
366 #[validate(range(min = 1))]
368 pub clip_max_new_tokens: Option<usize>,
369 #[validate(range(min = 0.0, max = 1.0))]
371 pub schedule_conservativeness: Option<f64>,
372}
373
374#[derive(Debug, Clone, Serialize, Deserialize, Validate, Default)]
379pub struct TrtllmArgs {
380 pub capacity_scheduler_policy: Option<String>,
383}
384
385#[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#[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 #[builder(default = "EngineType::Vllm")]
515 pub engine_type: EngineType,
516
517 #[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 #[builder(default = "None")]
531 #[validate(range(min = 1))]
532 pub max_model_len: Option<usize>,
533
534 #[builder(default = Some(256))]
536 #[validate(range(min = 1))]
537 pub max_num_seqs: Option<usize>,
538
539 #[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 #[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 #[builder(default = "None")]
568 #[validate(range(min = 0.0))]
569 pub startup_time: Option<f64>,
570
571 #[builder(default = "WorkerType::Aggregated")]
573 pub worker_type: WorkerType,
574
575 #[builder(default = "None")]
577 pub planner_profile_data: Option<PathBuf>,
578
579 #[serde(skip)]
581 #[builder(default = "Arc::new(PerfModel::default())")]
582 pub perf_model: Arc<PerfModel>,
583
584 #[serde(skip)]
588 #[builder(default = "None")]
589 pub aic_backend: Option<String>,
590
591 #[serde(skip)]
593 #[builder(default = "None")]
594 pub aic_system: Option<String>,
595
596 #[serde(skip)]
600 #[builder(default = "None")]
601 pub aic_backend_version: Option<String>,
602
603 #[serde(skip)]
606 #[builder(default = "None")]
607 pub aic_tp_size: Option<usize>,
608
609 #[serde(skip)]
611 #[builder(default = "None")]
612 pub aic_model_path: Option<String>,
613
614 #[serde(skip)]
617 #[builder(default = "None")]
618 pub aic_moe_tp_size: Option<usize>,
619
620 #[serde(skip)]
623 #[builder(default = "None")]
624 pub aic_moe_ep_size: Option<usize>,
625
626 #[serde(skip)]
630 #[builder(default = "None")]
631 pub aic_attention_dp_size: Option<usize>,
632
633 #[serde(skip)]
635 #[builder(default = "None")]
636 pub aic_gemm_dtype: Option<String>,
637
638 #[serde(skip)]
640 #[builder(default = "None")]
641 pub aic_moe_dtype: Option<String>,
642
643 #[serde(skip)]
645 #[builder(default = "None")]
646 pub aic_fmha_dtype: Option<String>,
647
648 #[serde(skip)]
650 #[builder(default = "None")]
651 pub aic_kv_cache_dtype: Option<String>,
652
653 #[serde(skip)]
655 #[builder(default = "None")]
656 pub aic_comm_dtype: Option<String>,
657
658 #[builder(default = "None")]
662 #[validate(range(min = 1, max = 5))]
663 pub aic_nextn: Option<usize>,
664
665 #[builder(default = "None")]
668 pub aic_nextn_accept_rates: Option<String>,
669
670 #[builder(default = "42")]
673 pub aic_mtp_seed: u64,
674
675 #[builder(default = "None")]
677 #[validate(range(min = 0.0, max = 1.0))]
678 pub gpu_memory_utilization: Option<f64>,
679
680 #[builder(default = "None")]
682 #[validate(range(min = 0.0, max = 1.0))]
683 pub mem_fraction_static: Option<f64>,
684
685 #[builder(default = "None")]
691 #[validate(range(min = 0.0, max = 1.0))]
692 pub free_gpu_memory_fraction: Option<f64>,
693
694 #[builder(default = "false")]
696 pub enable_local_indexer: bool,
697
698 #[builder(default = "None")]
702 pub bootstrap_port: Option<u16>,
703
704 #[builder(default = "300_000")]
706 #[validate(range(min = 1))]
707 pub handoff_session_timeout_ms: u64,
708
709 #[builder(default = "None")]
712 pub kv_bytes_per_token: Option<usize>,
713
714 #[builder(default = "None")]
716 #[serde(skip_serializing_if = "Option::is_none")]
717 pub kv_cache_bytes_per_token: Option<usize>,
718
719 #[builder(default = "None")]
723 #[validate(range(min = 0.0))]
724 pub kv_transfer_bandwidth: Option<f64>,
725
726 #[builder(default = "KvTransferTimingMode::FullPrompt")]
729 pub kv_transfer_timing_mode: KvTransferTimingMode,
730
731 #[builder(default = "None")]
734 pub reasoning: Option<ReasoningConfig>,
735
736 #[builder(default = "None")]
740 pub response_replay_trace_path: Option<PathBuf>,
741
742 #[builder(default = "None")]
746 pub zmq_kv_events_port: Option<u16>,
747
748 #[builder(default = "None")]
753 pub zmq_replay_port: Option<u16>,
754
755 #[builder(default)]
758 pub preemption_mode: PreemptionMode,
759
760 #[builder(default = "None")]
762 pub router_queue_policy: Option<RouterQueuePolicy>,
763
764 #[builder(default = "None")]
766 pub sglang: Option<SglangArgs>,
767
768 #[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 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 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 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 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}