1use serde::{Deserialize, Serialize};
2
3pub mod tokenizer;
4
5pub mod binary;
6pub mod proto {
7 pub mod stage {
8 include!(concat!(env!("OUT_DIR"), "/skippy.stage.v1.rs"));
9 }
10}
11
12pub const SCHEMA_VERSION: u32 = 1;
13pub const STAGE_ALPN_V2: &[u8] = b"skippy-stage/2";
14pub const STAGE_SUBPROTOCOL_NAME: &str = "skippy-stage";
15pub const STAGE_SUBPROTOCOL_MAJOR: u32 = 2;
16pub const STAGE_SUBPROTOCOL_FEATURE_STAGE_CONTROL: &str = "stage-control";
17pub const STAGE_PROTOCOL_GENERATION: u32 = 4;
18pub const STAGE_SUBPROTOCOL_FEATURE_STAGE_PROTOCOL_GENERATION_V4: &str = "stage-generation-4";
22pub const STAGE_SUBPROTOCOL_FEATURE_STAGE_GENERATION: &str =
23 STAGE_SUBPROTOCOL_FEATURE_STAGE_PROTOCOL_GENERATION_V4;
24pub const STAGE_SUBPROTOCOL_FEATURE_ARTIFACT_TRANSFER: &str = "artifact-transfer";
25pub const STAGE_SUBPROTOCOL_FEATURE_STATUS_LIST: &str = "status-list";
26pub const STAGE_STREAM_CONTROL: u8 = 0x01;
27pub const STAGE_STREAM_TRANSPORT: u8 = 0x02;
28pub const STAGE_STREAM_ARTIFACT_TRANSFER: u8 = 0x03;
29pub const MAX_STAGE_FRAME_BYTES: usize = 8 * 1024 * 1024;
30pub const MAX_VERIFY_WINDOW_PIPELINE_DEPTH: usize = 64;
32
33#[derive(Debug, Clone, PartialEq, Eq)]
34pub enum StageFrameError {
35 BadGeneration { got: u32 },
36 InvalidEndpointId { got: usize },
37 InvalidArtifactDigestLength { got: usize },
38 InvalidArtifactPath,
39 InvalidArtifactOffset,
40 MissingStageControlCommand,
41 MissingStageControlResponse,
42 MissingStageTransportTarget,
43 MissingStageArtifactTarget,
44}
45
46impl std::fmt::Display for StageFrameError {
47 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
48 match self {
49 StageFrameError::BadGeneration { got } => write!(
50 f,
51 "bad skippy stage generation: expected {}, got {}",
52 STAGE_PROTOCOL_GENERATION, got
53 ),
54 StageFrameError::InvalidEndpointId { got } => {
55 write!(f, "invalid endpoint_id length: expected 32, got {got}")
56 }
57 StageFrameError::InvalidArtifactDigestLength { got } => write!(
58 f,
59 "invalid artifact sha256 length: expected 64 hex chars, got {got}"
60 ),
61 StageFrameError::InvalidArtifactPath => {
62 write!(f, "artifact relative_path must be a safe relative path")
63 }
64 StageFrameError::InvalidArtifactOffset => {
65 write!(f, "artifact offset exceeds expected artifact size")
66 }
67 StageFrameError::MissingStageControlCommand => {
68 write!(f, "stage control command is required but missing")
69 }
70 StageFrameError::MissingStageControlResponse => {
71 write!(f, "stage control response is required but missing")
72 }
73 StageFrameError::MissingStageTransportTarget => {
74 write!(f, "stage transport target is required but missing")
75 }
76 StageFrameError::MissingStageArtifactTarget => {
77 write!(f, "stage artifact transfer target is required but missing")
78 }
79 }
80 }
81}
82
83impl std::error::Error for StageFrameError {}
84
85pub fn validate_stage_control_request(
86 frame: &proto::stage::StageControlRequest,
87) -> Result<(), StageFrameError> {
88 validate_generation(frame.r#gen)?;
89 validate_endpoint_id(frame.requester_id.len())?;
90 if frame.command.is_none() {
91 return Err(StageFrameError::MissingStageControlCommand);
92 }
93 Ok(())
94}
95
96pub fn validate_stage_control_response(
97 frame: &proto::stage::StageControlResponse,
98) -> Result<(), StageFrameError> {
99 validate_generation(frame.r#gen)?;
100 if frame.response.is_none() {
101 return Err(StageFrameError::MissingStageControlResponse);
102 }
103 Ok(())
104}
105
106pub fn validate_stage_transport_open(
107 frame: &proto::stage::StageTransportOpen,
108) -> Result<(), StageFrameError> {
109 validate_generation(frame.r#gen)?;
110 validate_endpoint_id(frame.requester_id.len())?;
111 if frame.topology_id.is_empty() || frame.run_id.is_empty() || frame.stage_id.is_empty() {
112 return Err(StageFrameError::MissingStageTransportTarget);
113 }
114 Ok(())
115}
116
117pub fn validate_stage_artifact_transfer_request(
118 frame: &proto::stage::StageArtifactTransferRequest,
119) -> Result<(), StageFrameError> {
120 validate_generation(frame.r#gen)?;
121 validate_endpoint_id(frame.requester_id.len())?;
122 if frame.topology_id.is_empty()
123 || frame.run_id.is_empty()
124 || frame.stage_id.is_empty()
125 || !frame.package_ref.starts_with("hf://")
126 {
127 return Err(StageFrameError::MissingStageArtifactTarget);
128 }
129 validate_artifact_digest(&frame.manifest_sha256)?;
130 if let Some(expected_sha) = frame.expected_sha256.as_deref() {
131 validate_artifact_digest(expected_sha)?;
132 }
133 if frame.expected_size.is_some_and(|size| frame.offset > size) {
134 return Err(StageFrameError::InvalidArtifactOffset);
135 }
136 validate_safe_relative_artifact_path(&frame.relative_path)?;
137 Ok(())
138}
139
140pub fn validate_stage_artifact_transfer_response(
141 frame: &proto::stage::StageArtifactTransferResponse,
142) -> Result<(), StageFrameError> {
143 validate_generation(frame.r#gen)?;
144 if let Some(sha256) = frame.sha256.as_deref() {
145 validate_artifact_digest(sha256)?;
146 }
147 Ok(())
148}
149
150fn validate_generation(r#gen: u32) -> Result<(), StageFrameError> {
151 if r#gen != STAGE_PROTOCOL_GENERATION {
152 return Err(StageFrameError::BadGeneration { got: r#gen });
153 }
154 Ok(())
155}
156
157fn validate_artifact_digest(value: &str) -> Result<(), StageFrameError> {
158 if value.len() != 64 || !value.chars().all(|ch| ch.is_ascii_hexdigit()) {
159 return Err(StageFrameError::InvalidArtifactDigestLength { got: value.len() });
160 }
161 Ok(())
162}
163
164fn validate_safe_relative_artifact_path(path: &str) -> Result<(), StageFrameError> {
165 use std::path::{Component, Path};
166
167 if path.trim().is_empty() {
168 return Err(StageFrameError::InvalidArtifactPath);
169 }
170 let path = Path::new(path);
171 let mut components = path.components();
172 let Some(first) = components.next() else {
173 return Err(StageFrameError::InvalidArtifactPath);
174 };
175 if !matches!(first, Component::Normal(_))
176 || !components.all(|component| matches!(component, Component::Normal(_)))
177 {
178 return Err(StageFrameError::InvalidArtifactPath);
179 }
180 Ok(())
181}
182
183fn validate_endpoint_id(len: usize) -> Result<(), StageFrameError> {
184 if len != 32 {
185 return Err(StageFrameError::InvalidEndpointId { got: len });
186 }
187 Ok(())
188}
189
190#[derive(Debug, Clone, Copy, PartialEq, Eq)]
191pub enum MessageKind {
192 Ready,
193 PrefillChunk,
194 FinalPrefillChunk,
195 DecodeToken,
196 StateImport,
197 StateExport,
198 Ack,
199 TokenReply,
200 Stop,
201 Error,
202}
203
204#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)]
205pub struct StageIdentity {
206 pub run_id: String,
207 pub request_id: String,
208 pub session_id: String,
209 pub topology_id: String,
210 pub stage_id: String,
211 pub stage_index: u32,
212}
213
214#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)]
215#[serde(rename_all = "kebab-case")]
216pub enum LoadMode {
217 RuntimeSlice,
218 LayerPackage,
219 ArtifactSlice,
220}
221
222#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Deserialize, Serialize)]
223#[serde(rename_all = "snake_case")]
224pub enum FlashAttentionType {
225 #[default]
226 Auto,
227 Disabled,
228 Enabled,
229}
230
231#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)]
232pub struct StageConfig {
233 pub run_id: String,
234 pub topology_id: String,
235 pub model_id: String,
236 #[serde(default)]
237 pub package_ref: Option<String>,
238 #[serde(default)]
239 pub manifest_sha256: Option<String>,
240 #[serde(default)]
241 pub source_model_path: Option<String>,
242 #[serde(default)]
243 pub source_model_sha256: Option<String>,
244 #[serde(default)]
245 pub source_model_bytes: Option<u64>,
246 #[serde(default)]
247 pub materialized_path: Option<String>,
248 #[serde(default)]
249 pub materialized_pinned: bool,
250 #[serde(default)]
251 pub model_path: Option<String>,
252 #[serde(default)]
253 pub projector_path: Option<String>,
254 pub stage_id: String,
255 pub stage_index: u32,
256 pub layer_start: u32,
257 pub layer_end: u32,
258 #[serde(default = "default_ctx_size")]
259 pub ctx_size: u32,
260 #[serde(default = "default_lane_count")]
261 pub lane_count: u32,
262 #[serde(default)]
263 pub n_batch: Option<u32>,
264 #[serde(default)]
265 pub n_ubatch: Option<u32>,
266 #[serde(default)]
267 pub n_gpu_layers: i32,
268 #[serde(default)]
269 pub mmap: Option<bool>,
270 #[serde(default)]
271 pub mlock: bool,
272 #[serde(default = "default_cache_type")]
273 pub cache_type_k: String,
274 #[serde(default = "default_cache_type")]
275 pub cache_type_v: String,
276 #[serde(default)]
277 pub flash_attn_type: FlashAttentionType,
278 #[serde(default)]
279 pub filter_tensors_on_load: bool,
280 #[serde(default)]
281 pub selected_device: Option<StageDevice>,
282 #[serde(default)]
283 pub kv_cache: Option<StageKvCacheConfig>,
284 #[serde(default = "default_native_mtp_enabled")]
285 pub native_mtp_enabled: bool,
286 pub load_mode: LoadMode,
287 pub bind_addr: String,
288 #[serde(default)]
289 pub upstream: Option<PeerConfig>,
290 #[serde(default)]
291 pub downstream: Option<PeerConfig>,
292}
293
294#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)]
295pub struct StageDevice {
296 pub backend_device: String,
297 #[serde(default)]
298 pub stable_id: Option<String>,
299 #[serde(default)]
300 pub index: Option<usize>,
301 #[serde(default)]
302 pub vram_bytes: Option<u64>,
303}
304
305#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)]
306#[serde(rename_all = "kebab-case")]
307pub enum StageKvCacheMode {
308 Disabled,
309 Auto,
310 Record,
311 LookupRecord,
312}
313
314#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize, Serialize)]
315#[serde(rename_all = "kebab-case")]
316pub enum StageKvCachePayload {
317 Auto,
318 ResidentKv,
319 KvRecurrent,
320 FullState,
321}
322
323#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)]
324pub struct StageKvCacheConfig {
325 #[serde(default = "default_kv_cache_mode")]
326 pub mode: StageKvCacheMode,
327 #[serde(default = "default_kv_cache_payload")]
328 pub payload: StageKvCachePayload,
329 #[serde(default = "default_kv_cache_max_entries")]
330 pub max_entries: usize,
331 #[serde(default)]
332 pub max_bytes: u64,
333 #[serde(default = "default_kv_cache_min_tokens")]
334 pub min_tokens: u64,
335 #[serde(default = "default_kv_cache_shared_stride_tokens")]
336 pub shared_prefix_stride_tokens: u64,
337 #[serde(default = "default_kv_cache_shared_record_limit")]
338 pub shared_prefix_record_limit: u64,
339}
340
341fn default_kv_cache_mode() -> StageKvCacheMode {
342 StageKvCacheMode::Auto
343}
344
345fn default_kv_cache_payload() -> StageKvCachePayload {
346 StageKvCachePayload::Auto
347}
348
349fn default_kv_cache_max_entries() -> usize {
350 64
351}
352
353fn default_kv_cache_min_tokens() -> u64 {
354 64
355}
356
357fn default_kv_cache_shared_stride_tokens() -> u64 {
358 128
359}
360
361fn default_kv_cache_shared_record_limit() -> u64 {
362 2
363}
364
365fn default_ctx_size() -> u32 {
366 512
367}
368
369fn default_lane_count() -> u32 {
370 4
371}
372
373fn default_cache_type() -> String {
374 "f16".to_string()
375}
376
377fn default_native_mtp_enabled() -> bool {
378 true
379}
380
381#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)]
382pub struct PeerConfig {
383 pub stage_id: String,
384 pub stage_index: u32,
385 pub endpoint: String,
386}
387
388#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)]
389pub struct StageTopology {
390 pub topology_id: String,
391 pub model_id: String,
392 pub stages: Vec<StageTopologyEntry>,
393}
394
395#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)]
396pub struct StageTopologyEntry {
397 pub stage_id: String,
398 pub stage_index: u32,
399 pub host: Option<String>,
400 pub endpoint: String,
401 pub layer_start: u32,
402 pub layer_end: u32,
403 pub load_mode: LoadMode,
404}
405
406#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize, Serialize)]
407#[serde(rename_all = "snake_case")]
408pub enum ActivationDType {
409 Unknown,
410 F32,
411 F16,
412 Bf16,
413}
414
415#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize, Serialize)]
416#[serde(rename_all = "snake_case")]
417pub enum ActivationLayout {
418 Opaque,
419 TokenMajor,
420}
421
422#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)]
423pub struct ActivationDescriptor {
424 pub version: u32,
425 pub dtype: ActivationDType,
426 pub layout: ActivationLayout,
427 pub producer_stage_index: i32,
428 pub layer_start: i32,
429 pub layer_end: i32,
430 pub token_count: u32,
431 pub sequence_count: u32,
432 pub payload_bytes: u64,
433 #[serde(default)]
434 pub flags: u64,
435 #[serde(default)]
436 pub payload_sha256: Option<String>,
437}
438
439#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)]
440#[serde(tag = "message_type", rename_all = "snake_case")]
441pub enum StageMessage {
442 Ready(ReadyMessage),
443 PrefillChunk(PrefillChunkMessage),
444 FinalPrefillChunk(FinalPrefillChunkMessage),
445 DecodeToken(DecodeTokenMessage),
446 StateImport(StateImportMessage),
447 StateExport(StateExportMessage),
448 Ack(AckMessage),
449 TokenReply(TokenReplyMessage),
450 Stop(StopMessage),
451 Error(ErrorMessage),
452}
453
454#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)]
455pub struct MessageBase {
456 pub schema_version: u32,
457 pub run_id: String,
458 pub request_id: String,
459 pub session_id: String,
460 pub stage_id: String,
461 pub stage_index: u32,
462 pub topology_id: String,
463 #[serde(default)]
464 pub model_id: Option<String>,
465 #[serde(default)]
466 pub tokenizer_id: Option<String>,
467 #[serde(default)]
468 pub chat_template_id: Option<String>,
469 #[serde(default)]
470 pub seq: Option<u64>,
471}
472
473#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)]
474pub struct ReadyMessage {
475 #[serde(flatten)]
476 pub base: MessageBase,
477 pub layer_start: u32,
478 pub layer_end: u32,
479}
480
481#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)]
482pub struct PrefillChunkMessage {
483 #[serde(flatten)]
484 pub base: MessageBase,
485 pub token_ids: Vec<i32>,
486 pub prompt_token_start: u32,
487 #[serde(default)]
488 pub activation_dtype: Option<String>,
489 #[serde(default)]
490 pub activation_bytes: Option<u64>,
491 #[serde(default)]
492 pub activation: Option<ActivationDescriptor>,
493 #[serde(default)]
494 pub activation_ref: Option<String>,
495 #[serde(default)]
496 pub is_final: Option<bool>,
497}
498
499#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)]
500pub struct FinalPrefillChunkMessage {
501 #[serde(flatten)]
502 pub base: MessageBase,
503 pub token_ids: Vec<i32>,
504 pub prompt_token_start: u32,
505 pub is_final: bool,
506 #[serde(default)]
507 pub activation_dtype: Option<String>,
508 #[serde(default)]
509 pub activation_bytes: Option<u64>,
510 #[serde(default)]
511 pub activation: Option<ActivationDescriptor>,
512 #[serde(default)]
513 pub activation_ref: Option<String>,
514}
515
516#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)]
517pub struct DecodeTokenMessage {
518 #[serde(flatten)]
519 pub base: MessageBase,
520 pub token_id: i32,
521 pub decode_index: u32,
522 #[serde(default)]
523 pub activation_dtype: Option<String>,
524 #[serde(default)]
525 pub activation_bytes: Option<u64>,
526 #[serde(default)]
527 pub activation: Option<ActivationDescriptor>,
528 #[serde(default)]
529 pub activation_ref: Option<String>,
530}
531
532#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)]
533pub struct StateImportMessage {
534 #[serde(flatten)]
535 pub base: MessageBase,
536 pub layer_start: u32,
537 pub layer_end: u32,
538 pub state_bytes: u64,
539 #[serde(default)]
540 pub state_sha256: Option<String>,
541}
542
543#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)]
544pub struct StateExportMessage {
545 #[serde(flatten)]
546 pub base: MessageBase,
547 pub layer_start: u32,
548 pub layer_end: u32,
549}
550
551#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)]
552pub struct AckMessage {
553 #[serde(flatten)]
554 pub base: MessageBase,
555 pub acked_seq: u64,
556}
557
558#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)]
559pub struct TokenReplyMessage {
560 #[serde(flatten)]
561 pub base: MessageBase,
562 pub token_id: i32,
563 #[serde(default)]
564 pub decode_index: Option<u32>,
565}
566
567#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)]
568pub struct StopMessage {
569 #[serde(flatten)]
570 pub base: MessageBase,
571}
572
573#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)]
574pub struct ErrorMessage {
575 #[serde(flatten)]
576 pub base: MessageBase,
577 pub error_code: String,
578 pub error_message: String,
579}
580
581impl StageConfig {
582 pub fn ready_message(&self) -> StageMessage {
583 StageMessage::Ready(ReadyMessage {
584 base: MessageBase {
585 schema_version: SCHEMA_VERSION,
586 run_id: self.run_id.clone(),
587 request_id: "stage-ready".to_string(),
588 session_id: "stage-lifecycle".to_string(),
589 stage_id: self.stage_id.clone(),
590 stage_index: self.stage_index,
591 topology_id: self.topology_id.clone(),
592 model_id: Some(self.model_id.clone()),
593 tokenizer_id: None,
594 chat_template_id: None,
595 seq: Some(0),
596 },
597 layer_start: self.layer_start,
598 layer_end: self.layer_end,
599 })
600 }
601}
602
603impl StageMessage {
604 pub fn base(&self) -> &MessageBase {
605 match self {
606 Self::Ready(message) => &message.base,
607 Self::PrefillChunk(message) => &message.base,
608 Self::FinalPrefillChunk(message) => &message.base,
609 Self::DecodeToken(message) => &message.base,
610 Self::StateImport(message) => &message.base,
611 Self::StateExport(message) => &message.base,
612 Self::Ack(message) => &message.base,
613 Self::TokenReply(message) => &message.base,
614 Self::Stop(message) => &message.base,
615 Self::Error(message) => &message.base,
616 }
617 }
618
619 pub fn kind(&self) -> MessageKind {
620 match self {
621 Self::Ready(_) => MessageKind::Ready,
622 Self::PrefillChunk(_) => MessageKind::PrefillChunk,
623 Self::FinalPrefillChunk(_) => MessageKind::FinalPrefillChunk,
624 Self::DecodeToken(_) => MessageKind::DecodeToken,
625 Self::StateImport(_) => MessageKind::StateImport,
626 Self::StateExport(_) => MessageKind::StateExport,
627 Self::Ack(_) => MessageKind::Ack,
628 Self::TokenReply(_) => MessageKind::TokenReply,
629 Self::Stop(_) => MessageKind::Stop,
630 Self::Error(_) => MessageKind::Error,
631 }
632 }
633
634 pub fn ack_for(&self, stage: &StageConfig) -> StageMessage {
635 let base = self.base();
636 StageMessage::Ack(AckMessage {
637 base: MessageBase {
638 schema_version: SCHEMA_VERSION,
639 run_id: base.run_id.clone(),
640 request_id: base.request_id.clone(),
641 session_id: base.session_id.clone(),
642 stage_id: stage.stage_id.clone(),
643 stage_index: stage.stage_index,
644 topology_id: stage.topology_id.clone(),
645 model_id: Some(stage.model_id.clone()),
646 tokenizer_id: base.tokenizer_id.clone(),
647 chat_template_id: base.chat_template_id.clone(),
648 seq: base.seq,
649 },
650 acked_seq: base.seq.unwrap_or(0),
651 })
652 }
653}
654
655#[cfg(test)]
656mod tests {
657 use prost::Message as _;
658
659 use super::proto::stage::{
660 CancelPrepareStage, GetLayerInventory, GetStageStatus, LayerInventory, LayerRange,
661 LoadStage, PrepareStage, PrepareStageAccepted, SourceModelKind,
662 StageArtifactTransferRequest, StageArtifactTransferResponse, StageControlRequest,
663 StageControlResponse, StagePreparationState, StagePreparationStatus, StageReady,
664 StageRuntimeState, StageStatus, StageStatusAck, StageStatusList, StageStatusUpdate,
665 StageTransportOpen, StageWireDType, StopStage, stage_control_request,
666 stage_control_response,
667 };
668 use super::{
669 STAGE_PROTOCOL_GENERATION, STAGE_SUBPROTOCOL_FEATURE_STAGE_PROTOCOL_GENERATION_V4,
670 StageFrameError, validate_stage_artifact_transfer_request,
671 validate_stage_artifact_transfer_response, validate_stage_control_request,
672 validate_stage_control_response, validate_stage_transport_open,
673 };
674
675 #[test]
676 fn stage_protocol_generation_feature_names_current_generation() {
677 assert_eq!(
678 STAGE_SUBPROTOCOL_FEATURE_STAGE_PROTOCOL_GENERATION_V4,
679 format!("stage-generation-{STAGE_PROTOCOL_GENERATION}")
680 );
681 }
682
683 #[test]
684 fn stage_control_request_validates_generation_sender_and_command() {
685 let frame = StageControlRequest {
686 r#gen: STAGE_PROTOCOL_GENERATION,
687 requester_id: vec![9u8; 32],
688 command: Some(stage_control_request::Command::GetStageStatus(
689 GetStageStatus {
690 topology_id: Some("topology-a".to_string()),
691 run_id: Some("run-a".to_string()),
692 stage_id: Some("stage-0".to_string()),
693 },
694 )),
695 };
696 validate_stage_control_request(&frame).unwrap();
697
698 let load = StageControlRequest {
699 command: Some(stage_control_request::Command::LoadStage(LoadStage {
700 topology_id: "topology-a".to_string(),
701 run_id: "run-a".to_string(),
702 model_id: "qwen".to_string(),
703 backend: "skippy".to_string(),
704 package_ref: "hf://repo/model".to_string(),
705 manifest_sha256: "a5".repeat(32),
706 stage_id: "stage-0".to_string(),
707 layer_end: 16,
708 activation_width: 4096,
709 projector_path: Some("/models/mmproj.gguf".to_string()),
710 ..Default::default()
711 })),
712 ..frame.clone()
713 };
714 let decoded = StageControlRequest::decode(load.encode_to_vec().as_slice()).unwrap();
715 match decoded.command {
716 Some(stage_control_request::Command::LoadStage(load)) => {
717 assert_eq!(load.projector_path.as_deref(), Some("/models/mmproj.gguf"));
718 }
719 other => panic!("expected LoadStage, got {other:?}"),
720 }
721
722 let stop = StageControlRequest {
723 command: Some(stage_control_request::Command::StopStage(StopStage {
724 topology_id: "topology-a".to_string(),
725 run_id: "run-a".to_string(),
726 stage_id: "stage-0".to_string(),
727 shutdown_generation: 7,
728 coordinator_term: 7,
729 })),
730 ..frame.clone()
731 };
732 validate_stage_control_request(&stop).unwrap();
733
734 let inventory = StageControlRequest {
735 command: Some(stage_control_request::Command::GetLayerInventory(
736 GetLayerInventory {
737 model_id: "qwen".to_string(),
738 package_ref: "hf://repo/model".to_string(),
739 manifest_sha256: "a5".repeat(32),
740 },
741 )),
742 ..frame.clone()
743 };
744 validate_stage_control_request(&inventory).unwrap();
745
746 let prepare = StageControlRequest {
747 command: Some(stage_control_request::Command::PrepareStage(PrepareStage {
748 load_stage: Some(LoadStage {
749 topology_id: "topology-a".to_string(),
750 run_id: "run-a".to_string(),
751 model_id: "qwen".to_string(),
752 backend: "skippy".to_string(),
753 package_ref: "gguf:///model.gguf".to_string(),
754 manifest_sha256: "direct-gguf:1:model.gguf".to_string(),
755 stage_id: "stage-1".to_string(),
756 layer_start: 8,
757 layer_end: 16,
758 ..Default::default()
759 }),
760 coordinator_id: Some(vec![8u8; 32]),
761 })),
762 ..frame.clone()
763 };
764 validate_stage_control_request(&prepare).unwrap();
765
766 let status_update = StageControlRequest {
767 command: Some(stage_control_request::Command::StageStatusUpdate(
768 StageStatusUpdate {
769 status: Some(StagePreparationStatus {
770 topology_id: "topology-a".to_string(),
771 run_id: "run-a".to_string(),
772 model_id: "qwen".to_string(),
773 backend: "skippy".to_string(),
774 package_ref: "gguf:///model.gguf".to_string(),
775 manifest_sha256: "direct-gguf:1:model.gguf".to_string(),
776 stage_id: "stage-1".to_string(),
777 stage_index: 1,
778 layer_start: 8,
779 layer_end: 16,
780 state: StagePreparationState::Loading as i32,
781 bytes_done: Some(10),
782 bytes_total: Some(20),
783 shutdown_generation: 7,
784 ..Default::default()
785 }),
786 },
787 )),
788 ..frame.clone()
789 };
790 validate_stage_control_request(&status_update).unwrap();
791
792 let cancel = StageControlRequest {
793 command: Some(stage_control_request::Command::CancelPrepareStage(
794 CancelPrepareStage {
795 topology_id: "topology-a".to_string(),
796 run_id: "run-a".to_string(),
797 stage_id: "stage-1".to_string(),
798 shutdown_generation: 8,
799 },
800 )),
801 ..frame.clone()
802 };
803 validate_stage_control_request(&cancel).unwrap();
804
805 let missing_command = StageControlRequest {
806 command: None,
807 ..frame.clone()
808 };
809 assert!(matches!(
810 validate_stage_control_request(&missing_command),
811 Err(StageFrameError::MissingStageControlCommand)
812 ));
813
814 let wrong_gen = StageControlRequest { r#gen: 1, ..frame };
815 assert!(matches!(
816 validate_stage_control_request(&wrong_gen),
817 Err(StageFrameError::BadGeneration { got: 1 })
818 ));
819 }
820
821 #[test]
822 fn stage_control_response_validates_generation_and_response() {
823 let frame = StageControlResponse {
824 r#gen: STAGE_PROTOCOL_GENERATION,
825 response: Some(stage_control_response::Response::StageReady(StageReady {
826 accepted: true,
827 status: Some(StageStatus {
828 topology_id: "topology-a".to_string(),
829 run_id: "run-a".to_string(),
830 model_id: "qwen".to_string(),
831 backend: "skippy".to_string(),
832 stage_id: "stage-0".to_string(),
833 stage_index: 0,
834 layer_start: 0,
835 layer_end: 16,
836 state: StageRuntimeState::Ready as i32,
837 bind_addr: "127.0.0.1:0".to_string(),
838 activation_width: 4096,
839 wire_dtype: StageWireDType::StageWireDtypeF16 as i32,
840 shutdown_generation: 7,
841 ctx_size: 8192,
842 lane_count: 2,
843 projector_path: Some("/models/mmproj.gguf".to_string()),
844 ..Default::default()
845 }),
846 error: None,
847 })),
848 };
849 let decoded = StageControlResponse::decode(frame.encode_to_vec().as_slice()).unwrap();
850 validate_stage_control_response(&decoded).unwrap();
851 match decoded.response {
852 Some(stage_control_response::Response::StageReady(ready)) => {
853 let status = ready.status.expect("stage-ready status");
854 assert_eq!(
855 status.projector_path.as_deref(),
856 Some("/models/mmproj.gguf")
857 );
858 assert_eq!(status.lane_count, 2);
859 }
860 other => panic!("expected StageReady, got {other:?}"),
861 }
862
863 let inventory_response = StageControlResponse {
864 response: Some(stage_control_response::Response::LayerInventory(
865 LayerInventory {
866 model_id: "qwen".to_string(),
867 package_ref: "hf://repo/model".to_string(),
868 manifest_sha256: "a5".repeat(32),
869 layer_count: 16,
870 source_model_path: Some("/model.gguf".to_string()),
871 source_model_bytes: Some(1024),
872 source_model_kind: SourceModelKind::PlainGguf as i32,
873 ready_ranges: vec![LayerRange {
874 layer_start: 0,
875 layer_end: 8,
876 }],
877 ..Default::default()
878 },
879 )),
880 ..frame.clone()
881 };
882 validate_stage_control_response(&inventory_response).unwrap();
883
884 let prepare_response = StageControlResponse {
885 response: Some(stage_control_response::Response::PrepareStageAccepted(
886 PrepareStageAccepted {
887 accepted: true,
888 status: Some(StagePreparationStatus {
889 topology_id: "topology-a".to_string(),
890 run_id: "run-a".to_string(),
891 model_id: "qwen".to_string(),
892 backend: "skippy".to_string(),
893 package_ref: "hf://repo/model".to_string(),
894 manifest_sha256: "a5".repeat(32),
895 stage_id: "stage-1".to_string(),
896 stage_index: 1,
897 layer_start: 8,
898 layer_end: 16,
899 state: StagePreparationState::Assigned as i32,
900 shutdown_generation: 7,
901 ..Default::default()
902 }),
903 error: None,
904 },
905 )),
906 ..frame.clone()
907 };
908 validate_stage_control_response(&prepare_response).unwrap();
909
910 let ack_response = StageControlResponse {
911 response: Some(stage_control_response::Response::StageStatusAck(
912 StageStatusAck {
913 accepted: true,
914 error: None,
915 },
916 )),
917 ..frame.clone()
918 };
919 validate_stage_control_response(&ack_response).unwrap();
920
921 let status_list_response = StageControlResponse {
922 response: Some(stage_control_response::Response::StageStatuses(
923 StageStatusList {
924 statuses: vec![StageStatus {
925 topology_id: "topology-a".to_string(),
926 run_id: "run-a".to_string(),
927 model_id: "qwen".to_string(),
928 backend: "skippy".to_string(),
929 stage_id: "stage-0".to_string(),
930 stage_index: 0,
931 layer_start: 0,
932 layer_end: 16,
933 state: StageRuntimeState::Ready as i32,
934 bind_addr: "127.0.0.1:51234".to_string(),
935 activation_width: 4096,
936 wire_dtype: StageWireDType::StageWireDtypeF16 as i32,
937 shutdown_generation: 7,
938 ctx_size: 8192,
939 lane_count: 2,
940 ..Default::default()
941 }],
942 },
943 )),
944 ..frame.clone()
945 };
946 validate_stage_control_response(&status_list_response).unwrap();
947
948 let missing_response = StageControlResponse {
949 response: None,
950 ..frame.clone()
951 };
952 assert!(matches!(
953 validate_stage_control_response(&missing_response),
954 Err(StageFrameError::MissingStageControlResponse)
955 ));
956
957 let wrong_gen = StageControlResponse { r#gen: 1, ..frame };
958 assert!(matches!(
959 validate_stage_control_response(&wrong_gen),
960 Err(StageFrameError::BadGeneration { got: 1 })
961 ));
962 }
963
964 #[test]
965 fn stage_transport_open_validates_generation_sender_and_target() {
966 let frame = StageTransportOpen {
967 r#gen: STAGE_PROTOCOL_GENERATION,
968 requester_id: vec![7u8; 32],
969 topology_id: "topology-a".to_string(),
970 run_id: "run-a".to_string(),
971 stage_id: "stage-1".to_string(),
972 };
973 validate_stage_transport_open(&frame).unwrap();
974
975 let missing_target = StageTransportOpen {
976 stage_id: String::new(),
977 ..frame.clone()
978 };
979 assert!(matches!(
980 validate_stage_transport_open(&missing_target),
981 Err(StageFrameError::MissingStageTransportTarget)
982 ));
983
984 let wrong_gen = StageTransportOpen { r#gen: 1, ..frame };
985 assert!(matches!(
986 validate_stage_transport_open(&wrong_gen),
987 Err(StageFrameError::BadGeneration { got: 1 })
988 ));
989 }
990
991 #[test]
992 fn stage_artifact_transfer_frames_validate_skippy_owned_contract() {
993 let request = StageArtifactTransferRequest {
994 r#gen: STAGE_PROTOCOL_GENERATION,
995 requester_id: vec![7u8; 32],
996 topology_id: "topology-a".to_string(),
997 run_id: "run-a".to_string(),
998 stage_id: "stage-0".to_string(),
999 package_ref: "hf://meshllm/demo-layers@abc123".to_string(),
1000 manifest_sha256: "a".repeat(64),
1001 relative_path: "layers/layer-000.gguf".to_string(),
1002 offset: 0,
1003 expected_size: Some(8),
1004 expected_sha256: Some("b".repeat(64)),
1005 };
1006 let decoded =
1007 StageArtifactTransferRequest::decode(request.encode_to_vec().as_slice()).unwrap();
1008 validate_stage_artifact_transfer_request(&decoded).unwrap();
1009 assert_eq!(decoded.stage_id, "stage-0");
1010
1011 let mut unsafe_path = request.clone();
1012 unsafe_path.relative_path = "../layer.gguf".to_string();
1013 assert!(matches!(
1014 validate_stage_artifact_transfer_request(&unsafe_path),
1015 Err(StageFrameError::InvalidArtifactPath)
1016 ));
1017
1018 let mut bad_offset = request.clone();
1019 bad_offset.offset = 9;
1020 assert!(matches!(
1021 validate_stage_artifact_transfer_request(&bad_offset),
1022 Err(StageFrameError::InvalidArtifactOffset)
1023 ));
1024
1025 let mut missing_target = request.clone();
1026 missing_target.topology_id.clear();
1027 assert!(matches!(
1028 validate_stage_artifact_transfer_request(&missing_target),
1029 Err(StageFrameError::MissingStageArtifactTarget)
1030 ));
1031
1032 let response = StageArtifactTransferResponse {
1033 r#gen: STAGE_PROTOCOL_GENERATION,
1034 accepted: true,
1035 total_size: 8,
1036 sha256: Some("b".repeat(64)),
1037 error: None,
1038 };
1039 let decoded =
1040 StageArtifactTransferResponse::decode(response.encode_to_vec().as_slice()).unwrap();
1041 validate_stage_artifact_transfer_response(&decoded).unwrap();
1042
1043 let bad_response_sha = StageArtifactTransferResponse {
1044 sha256: Some("not-a-sha".to_string()),
1045 ..response
1046 };
1047 assert!(matches!(
1048 validate_stage_artifact_transfer_response(&bad_response_sha),
1049 Err(StageFrameError::InvalidArtifactDigestLength { .. })
1050 ));
1051 }
1052}