Skip to main content

skippy_protocol/
lib.rs

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;
18/// Generation-scoped stage capability. A peer can advertise `stage-control`
19/// while still rejecting current-generation frames, so split planning gates on
20/// this exact token before sending current-generation control requests.
21pub 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;
30/// Maximum number of unresolved verify windows covered by native checkpoints.
31pub 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}