Skip to main content

skippy_protocol/binary/
mod.rs

1mod activation;
2mod codec;
3mod types;
4
5pub use activation::{
6    activation_payload_multiplier_from_state_flags, activation_wire_bytes,
7    activation_wire_bytes_with_state_flags, encode_f32_activation_payload,
8    encode_f32_activation_payload_with_state_flags,
9};
10pub use codec::{
11    read_stage_message, recv_ready, recv_reply, send_ready, send_reply_ack,
12    send_reply_ack_with_stats, send_reply_message, send_reply_predicted,
13    send_reply_predicted_tokens_with_stats, send_reply_predicted_tokens_with_window_and_stats,
14    send_reply_predicted_with_stats, send_reply_predicted_with_tokens_and_stats,
15    send_reply_predicted_with_tokens_window_and_stats, write_stage_message,
16};
17pub use types::{
18    ACTIVATION_FLAG_GEMMA3N_ALTUP, ACTIVATION_FLAG_INKLING_MTP_EMBD, ACTIVATION_FLAG_RWKV7_V_FIRST,
19    LLAMA_TOKEN_NULL, MAX_STAGE_ACTIVATION_BYTES, MAX_STAGE_CHAT_SAMPLING_METADATA_BYTES,
20    MAX_STAGE_DECODED_ACTIVATION_BYTES, MAX_STAGE_LOGIT_BIAS, MAX_STAGE_PREDICTED_TOKENS,
21    MAX_STAGE_SIDEBAND_VALUES, MAX_STAGE_STATE_IMPORT_BYTES, READY_MAGIC,
22    STAGE_LOGIT_BIAS_WIRE_BYTES, STAGE_SAMPLING_CONFIG_BASE_BYTES, STAGE_STATE_HEADER_BYTES,
23    STAGE_STATE_VERSION, STAGE_WIRE_FIXED_HEADER_BYTES, StageLogitBias, StageNativeMtpDraft,
24    StageReply, StageReplyStats, StageReplyWindow, StageRequestEpoch, StageSamplingConfig,
25    StageStateHeader, StageWireMessage, WireActivationDType, WireMessageKind, WireReplyKind,
26    WireStagePhase, activation_frame_flags_from_state_flags,
27    activation_state_flags_from_frame_flags, state_flags,
28};
29
30pub(crate) fn invalid_data(message: &'static str) -> std::io::Error {
31    std::io::Error::new(std::io::ErrorKind::InvalidData, message)
32}
33
34pub(crate) fn invalid_input(message: &'static str) -> std::io::Error {
35    std::io::Error::new(std::io::ErrorKind::InvalidInput, message)
36}
37
38#[cfg(test)]
39mod tests {
40    use super::*;
41    use std::io::Cursor;
42
43    fn push_i32(bytes: &mut Vec<u8>, value: i32) {
44        bytes.extend_from_slice(&value.to_le_bytes());
45    }
46
47    fn push_u32(bytes: &mut Vec<u8>, value: u32) {
48        bytes.extend_from_slice(&value.to_le_bytes());
49    }
50
51    fn push_u64(bytes: &mut Vec<u8>, value: u64) {
52        bytes.extend_from_slice(&value.to_le_bytes());
53    }
54
55    fn push_state_header(bytes: &mut Vec<u8>, state: StageStateHeader) {
56        push_i32(bytes, state.version);
57        push_i32(bytes, state.seq_id);
58        push_i32(bytes, state.phase);
59        push_i32(bytes, state.flags);
60        push_i32(bytes, state.checkpoint_generation);
61        push_i32(bytes, state.prompt_token_count);
62        push_i32(bytes, state.decode_step);
63        push_i32(bytes, state.current_token);
64        push_i32(bytes, state.source_stage_index);
65        push_i32(bytes, state.reserved);
66    }
67
68    fn stage_frame_prefix(
69        kind: WireMessageKind,
70        token_count: i32,
71        token_sideband_count: i32,
72        position_sideband_count: i32,
73        state: StageStateHeader,
74    ) -> Vec<u8> {
75        let mut bytes = Vec::new();
76        push_i32(&mut bytes, kind as i32);
77        push_i32(&mut bytes, 0);
78        push_i32(&mut bytes, token_count);
79        push_i32(&mut bytes, token_sideband_count);
80        push_i32(&mut bytes, position_sideband_count);
81        push_state_header(&mut bytes, state);
82        push_u64(&mut bytes, 7);
83        push_u64(&mut bytes, 11);
84        bytes
85    }
86
87    fn assert_invalid_data<T: std::fmt::Debug>(result: std::io::Result<T>, expected: &str) {
88        let error = result.unwrap_err();
89        assert_eq!(error.kind(), std::io::ErrorKind::InvalidData);
90        assert_eq!(error.to_string(), expected);
91    }
92
93    #[test]
94    fn ready_round_trips() {
95        let mut bytes = Vec::new();
96        send_ready(&mut bytes).unwrap();
97        recv_ready(Cursor::new(bytes)).unwrap();
98    }
99
100    #[test]
101    fn reply_round_trips() {
102        let mut bytes = Vec::new();
103        send_reply_predicted(&mut bytes, 42).unwrap();
104        let reply = recv_reply(Cursor::new(bytes)).unwrap();
105        assert_eq!(reply.kind, WireReplyKind::PredictedToken);
106        assert_eq!(reply.predicted, 42);
107        assert_eq!(reply.predicted_tokens, vec![42]);
108        assert_eq!(reply.native_mtp_draft, None);
109    }
110
111    #[test]
112    fn reply_round_trips_typed_native_mtp_draft() {
113        let reply = StageReply {
114            kind: WireReplyKind::PredictedToken,
115            predicted: 42,
116            predicted_tokens: vec![42],
117            native_mtp_draft: Some(StageNativeMtpDraft {
118                token_ids: vec![43, 44],
119                proposal_compute_us: 12_345,
120            }),
121            window: StageReplyWindow::default(),
122            stats: StageReplyStats::default(),
123        };
124        let mut bytes = Vec::new();
125        send_reply_message(&mut bytes, &reply).unwrap();
126
127        assert_eq!(recv_reply(Cursor::new(bytes)).unwrap(), reply);
128    }
129
130    #[test]
131    fn predicted_token_reply_preserves_sideband_tokens() {
132        let mut bytes = Vec::new();
133        send_reply_predicted_with_tokens_and_stats(
134            &mut bytes,
135            42,
136            &[42, 43, 123],
137            StageReplyStats::default(),
138        )
139        .unwrap();
140        let reply = recv_reply(Cursor::new(bytes)).unwrap();
141        assert_eq!(reply.kind, WireReplyKind::PredictedToken);
142        assert_eq!(reply.predicted, 42);
143        assert_eq!(reply.predicted_tokens, vec![42, 43, 123]);
144    }
145
146    #[test]
147    fn reply_stats_preserve_prefill_edge_transport() {
148        let mut stats = StageReplyStats::default();
149        stats.observe_prefill_edge_transport(2, 12_000, 3_000, 1_048_576);
150        stats.observe_prefill_edge_transport(1, 4_000, 40_000, 524_288);
151
152        let mut bytes = Vec::new();
153        send_reply_predicted_with_stats(&mut bytes, 42, stats).unwrap();
154        let reply = recv_reply(Cursor::new(bytes)).unwrap();
155
156        assert_eq!(reply.stats.prefill_edge_write_us_max, 12_000);
157        assert_eq!(reply.stats.prefill_edge_wait_us_max, 40_000);
158        assert_eq!(reply.stats.prefill_edge_total_us_max, 44_000);
159        assert_eq!(reply.stats.prefill_edge_stage_index, 1);
160        assert_eq!(reply.stats.prefill_edge_activation_bytes_max, 524_288);
161        assert_eq!(reply.stats.prefill_edge_observation_count, 2);
162    }
163
164    #[test]
165    fn token_vector_reply_round_trips() {
166        let mut bytes = Vec::new();
167        send_reply_predicted_tokens_with_stats(&mut bytes, &[1, 2, 3], StageReplyStats::default())
168            .unwrap();
169        let reply = recv_reply(Cursor::new(bytes)).unwrap();
170        assert_eq!(reply.kind, WireReplyKind::PredictedTokens);
171        assert_eq!(reply.predicted, 1);
172        assert_eq!(reply.predicted_tokens, vec![1, 2, 3]);
173    }
174
175    #[test]
176    fn reply_window_metadata_round_trips() {
177        let mut bytes = Vec::new();
178        send_reply_predicted_tokens_with_window_and_stats(
179            &mut bytes,
180            &[1, 2, 3],
181            StageReplyWindow { window_id: 42 },
182            StageReplyStats::default(),
183        )
184        .unwrap();
185        let reply = recv_reply(Cursor::new(bytes)).unwrap();
186
187        assert_eq!(reply.kind, WireReplyKind::PredictedTokens);
188        assert_eq!(reply.predicted_tokens, vec![1, 2, 3]);
189        assert_eq!(reply.window.window_id, 42);
190    }
191
192    #[test]
193    fn reply_rejects_predicted_token_count_over_limit() {
194        let mut bytes = Vec::new();
195        push_i32(&mut bytes, WireReplyKind::PredictedTokens as i32);
196        push_i32(&mut bytes, 1);
197        push_i32(
198            &mut bytes,
199            i32::try_from(MAX_STAGE_PREDICTED_TOKENS + 1).unwrap(),
200        );
201
202        assert_invalid_data(
203            recv_reply(Cursor::new(bytes)),
204            "predicted token count exceeds maximum",
205        );
206    }
207
208    #[test]
209    fn stage_message_round_trips_f32() {
210        let mut state =
211            StageStateHeader::new(WireMessageKind::DecodeEmbd, WireActivationDType::F32);
212        state.checkpoint_generation = 3;
213        state.prompt_token_count = 1;
214        state.decode_step = 0;
215        state.current_token = 11;
216        state.source_stage_index = 0;
217        let activation = vec![1, 2, 3, 4, 5, 6, 7, 8];
218        let message = StageWireMessage {
219            kind: WireMessageKind::DecodeEmbd,
220            pos_start: 1,
221            token_count: 1,
222            state,
223            request_id: 7,
224            session_id: 11,
225            sampling: Some(StageSamplingConfig {
226                flags: 1,
227                seed: 42,
228                temperature: 0.8,
229                top_p: 0.9,
230                top_k: 40,
231                logit_bias: vec![StageLogitBias {
232                    token_id: 123,
233                    bias: -50.0,
234                }],
235                ..StageSamplingConfig::default()
236            }),
237            chat_sampling_metadata: None,
238            tokens: vec![11],
239            positions: Vec::new(),
240            activation: activation.clone(),
241            raw_bytes: Vec::new(),
242        };
243        let mut bytes = Vec::new();
244        write_stage_message(&mut bytes, &message, WireActivationDType::F32).unwrap();
245        let decoded = read_stage_message(Cursor::new(bytes), 2).unwrap();
246        assert_eq!(decoded.kind, WireMessageKind::DecodeEmbd);
247        assert_eq!(decoded.tokens, vec![11]);
248        assert_eq!(decoded.activation, activation);
249        assert_eq!(decoded.state.source_stage_index, 0);
250        assert_eq!(decoded.request_id, 7);
251        assert_eq!(decoded.session_id, 11);
252        assert_eq!(
253            decoded.request_epoch(),
254            StageRequestEpoch {
255                request_id: 7,
256                session_id: 11,
257                checkpoint_generation: 3,
258                prompt_token_count: 1,
259                decode_step: 0,
260            }
261        );
262        assert_ne!(decoded.state.flags & state_flags::SAMPLING, 0);
263        assert_eq!(decoded.state.flags & state_flags::CHAT_SAMPLING_METADATA, 0);
264        assert_eq!(decoded.chat_sampling_metadata, None);
265        let sampling = decoded.sampling.expect("sampling extension round-tripped");
266        assert_eq!(sampling.seed, 42);
267        assert_eq!(sampling.top_k, 40);
268        assert_eq!(sampling.logit_bias.len(), 1);
269        assert_eq!(sampling.logit_bias[0].token_id, 123);
270        assert_eq!(sampling.logit_bias[0].bias, -50.0);
271    }
272
273    #[test]
274    fn verify_window_message_round_trips_window_metadata() {
275        let mut state =
276            StageStateHeader::new(WireMessageKind::VerifyWindow, WireActivationDType::F32);
277        state.seq_id = 42;
278        state.prompt_token_count = 128;
279        state.decode_step = 7;
280        state.current_token = 1001;
281        let message = StageWireMessage {
282            kind: WireMessageKind::VerifyWindow,
283            pos_start: 135,
284            token_count: 4,
285            state,
286            request_id: 7,
287            session_id: 11,
288            sampling: None,
289            chat_sampling_metadata: None,
290            tokens: vec![1001, 1002, 1003, 1004],
291            positions: Vec::new(),
292            activation: Vec::new(),
293            raw_bytes: Vec::new(),
294        };
295
296        let mut bytes = Vec::new();
297        write_stage_message(&mut bytes, &message, WireActivationDType::F32).unwrap();
298        let decoded = read_stage_message(Cursor::new(bytes), 2).unwrap();
299
300        assert_eq!(decoded.kind, WireMessageKind::VerifyWindow);
301        assert_eq!(decoded.verify_window_id(), Some(42));
302        assert_eq!(decoded.verify_window_base_position(), Some(135));
303        assert_eq!(decoded.verify_window_token_count(), Some(4));
304        assert_eq!(decoded.authoritative_session_position(), Some(135));
305        assert_eq!(decoded.tokens, vec![1001, 1002, 1003, 1004]);
306        assert_eq!(decoded.state.decode_step, 7);
307    }
308
309    #[test]
310    fn only_decode_messages_carry_authoritative_session_positions() {
311        let mut decode = StageWireMessage {
312            kind: WireMessageKind::DecodeEmbd,
313            pos_start: 17,
314            token_count: 1,
315            state: StageStateHeader::new(WireMessageKind::DecodeEmbd, WireActivationDType::F32),
316            request_id: 1,
317            session_id: 2,
318            sampling: None,
319            chat_sampling_metadata: None,
320            tokens: vec![3],
321            positions: Vec::new(),
322            activation: Vec::new(),
323            raw_bytes: Vec::new(),
324        };
325        assert_eq!(decode.authoritative_session_position(), Some(17));
326
327        decode.pos_start = -1;
328        assert_eq!(decode.authoritative_session_position(), None);
329
330        decode.kind = WireMessageKind::PrefillEmbd;
331        assert_eq!(decode.authoritative_session_position(), None);
332    }
333
334    #[test]
335    fn stage_message_rejects_old_state_version() {
336        let mut state =
337            StageStateHeader::new(WireMessageKind::DecodeEmbd, WireActivationDType::F32);
338        state.version = STAGE_STATE_VERSION - 1;
339        let bytes = stage_frame_prefix(WireMessageKind::DecodeEmbd, 1, 0, 0, state);
340
341        assert_invalid_data(
342            read_stage_message(Cursor::new(bytes), 2),
343            "unsupported stage state version",
344        );
345    }
346
347    #[test]
348    fn stage_message_rejects_legacy_kind_10() {
349        let mut bytes = Vec::new();
350        push_i32(&mut bytes, 10);
351        push_i32(&mut bytes, 0);
352        push_i32(&mut bytes, 1);
353        push_i32(&mut bytes, 0);
354        push_i32(&mut bytes, 0);
355
356        assert_invalid_data(
357            read_stage_message(Cursor::new(bytes), 2),
358            "unknown stage message kind",
359        );
360    }
361
362    #[test]
363    fn stage_message_estimates_full_wire_transfer_bytes() {
364        let message = StageWireMessage {
365            kind: WireMessageKind::PrefillEmbd,
366            pos_start: 0,
367            token_count: 2,
368            state: StageStateHeader::new(WireMessageKind::PrefillEmbd, WireActivationDType::F32),
369            request_id: 7,
370            session_id: 11,
371            sampling: Some(StageSamplingConfig {
372                flags: 1,
373                logit_bias: vec![
374                    StageLogitBias {
375                        token_id: 1,
376                        bias: -1.0,
377                    },
378                    StageLogitBias {
379                        token_id: 2,
380                        bias: 1.0,
381                    },
382                ],
383                ..StageSamplingConfig::default()
384            }),
385            chat_sampling_metadata: Some("{}".to_string()),
386            tokens: vec![1, 2],
387            positions: vec![0],
388            activation: vec![0; 16],
389            raw_bytes: Vec::new(),
390        };
391
392        assert_eq!(
393            message.estimated_wire_bytes(),
394            STAGE_WIRE_FIXED_HEADER_BYTES
395                + STAGE_SAMPLING_CONFIG_BASE_BYTES
396                + 2 * STAGE_LOGIT_BIAS_WIRE_BYTES
397                + std::mem::size_of::<u32>()
398                + 2
399                + 3 * std::mem::size_of::<i32>()
400                + 16
401        );
402    }
403
404    #[test]
405    fn request_epoch_orders_only_matching_flows() {
406        let older = StageRequestEpoch {
407            request_id: 7,
408            session_id: 11,
409            checkpoint_generation: 1,
410            prompt_token_count: 8,
411            decode_step: 2,
412        };
413        let newer = StageRequestEpoch {
414            request_id: 7,
415            session_id: 11,
416            checkpoint_generation: 1,
417            prompt_token_count: 8,
418            decode_step: 3,
419        };
420        let different_session = StageRequestEpoch {
421            session_id: 12,
422            ..newer
423        };
424
425        assert!(older.same_flow(newer));
426        assert!(older.is_stale_for(newer));
427        assert!(!newer.is_stale_for(older));
428        assert!(!older.same_flow(different_session));
429        assert!(!older.is_stale_for(different_session));
430    }
431
432    #[test]
433    fn request_epoch_staleness_orders_generation_before_prompt_before_decode() {
434        let base = StageRequestEpoch {
435            request_id: 7,
436            session_id: 11,
437            checkpoint_generation: 1,
438            prompt_token_count: 8,
439            decode_step: 3,
440        };
441        let newer_checkpoint = StageRequestEpoch {
442            checkpoint_generation: 2,
443            prompt_token_count: 0,
444            decode_step: 0,
445            ..base
446        };
447        let newer_prompt = StageRequestEpoch {
448            prompt_token_count: 9,
449            decode_step: 0,
450            ..base
451        };
452        let newer_decode = StageRequestEpoch {
453            decode_step: 4,
454            ..base
455        };
456
457        assert!(base.same_flow(newer_checkpoint));
458        assert!(base.is_stale_for(newer_checkpoint));
459        assert!(!newer_checkpoint.is_stale_for(base));
460        assert!(base.is_stale_for(newer_prompt));
461        assert!(!newer_prompt.is_stale_for(base));
462        assert!(base.is_stale_for(newer_decode));
463        assert!(!newer_decode.is_stale_for(base));
464    }
465
466    #[test]
467    fn generation_config_round_trips_sampling_metadata() {
468        let message = StageWireMessage::configure_generation(
469            WireActivationDType::F32,
470            7,
471            11,
472            123,
473            Some(StageSamplingConfig {
474                flags: 1,
475                seed: 42,
476                temperature: 0.8,
477                top_p: 0.9,
478                top_k: 40,
479                ..StageSamplingConfig::default()
480            }),
481            Some("{\"grammar\":\"root ::= \\\"x\\\"\"}".to_string()),
482        );
483        let mut bytes = Vec::new();
484        write_stage_message(&mut bytes, &message, WireActivationDType::F32).unwrap();
485        let decoded = read_stage_message(Cursor::new(bytes), 2).unwrap();
486        assert_eq!(decoded.kind, WireMessageKind::ConfigureGeneration);
487        assert_eq!(decoded.token_count, 0);
488        assert_eq!(decoded.tokens, Vec::<i32>::new());
489        assert_eq!(decoded.activation, Vec::<u8>::new());
490        assert_eq!(decoded.request_id, 7);
491        assert_eq!(decoded.session_id, 11);
492        assert_eq!(decoded.state.prompt_token_count, 123);
493        assert_ne!(decoded.state.flags & state_flags::SAMPLING, 0);
494        assert_ne!(decoded.state.flags & state_flags::CHAT_SAMPLING_METADATA, 0);
495        assert_eq!(
496            decoded.chat_sampling_metadata.as_deref(),
497            Some("{\"grammar\":\"root ::= \\\"x\\\"\"}")
498        );
499        let sampling = decoded.sampling.expect("sampling extension round-tripped");
500        assert_eq!(sampling.seed, 42);
501        assert_eq!(sampling.top_k, 40);
502    }
503
504    #[test]
505    fn stage_message_rejects_sampling_metadata_length_over_limit() {
506        let mut state = StageStateHeader::new(
507            WireMessageKind::ConfigureGeneration,
508            WireActivationDType::F32,
509        );
510        state.flags |= state_flags::CHAT_SAMPLING_METADATA;
511        let mut bytes = stage_frame_prefix(WireMessageKind::ConfigureGeneration, 0, 0, 0, state);
512        push_u32(
513            &mut bytes,
514            u32::try_from(MAX_STAGE_CHAT_SAMPLING_METADATA_BYTES + 1).unwrap(),
515        );
516
517        assert_invalid_data(
518            read_stage_message(Cursor::new(bytes), 2048),
519            "chat sampling metadata length exceeds maximum",
520        );
521    }
522
523    #[test]
524    fn driver_origin_message_round_trips_without_activation() {
525        let mut state =
526            StageStateHeader::new(WireMessageKind::PrefillEmbd, WireActivationDType::F32);
527        state.prompt_token_count = 2;
528        state.current_token = 22;
529        state.source_stage_index = -1;
530        let message = StageWireMessage {
531            kind: WireMessageKind::PrefillEmbd,
532            pos_start: 0,
533            token_count: 2,
534            state,
535            request_id: 13,
536            session_id: 17,
537            sampling: None,
538            chat_sampling_metadata: None,
539            tokens: vec![11, 22],
540            positions: Vec::new(),
541            activation: Vec::new(),
542            raw_bytes: Vec::new(),
543        };
544        let mut bytes = Vec::new();
545        write_stage_message(&mut bytes, &message, WireActivationDType::F32).unwrap();
546        let decoded = read_stage_message(Cursor::new(bytes), 2048).unwrap();
547        assert_eq!(decoded.tokens, vec![11, 22]);
548        assert!(decoded.activation.is_empty());
549        assert_eq!(decoded.state.source_stage_index, -1);
550        assert_eq!(decoded.request_id, 13);
551        assert_eq!(decoded.session_id, 17);
552        assert_eq!(decoded.state.flags & state_flags::SAMPLING, 0);
553        assert!(decoded.sampling.is_none());
554    }
555
556    #[test]
557    fn stage_message_rejects_token_sideband_count_over_limit() {
558        let mut state =
559            StageStateHeader::new(WireMessageKind::PrefillEmbd, WireActivationDType::F32);
560        state.source_stage_index = -1;
561        let bytes = stage_frame_prefix(
562            WireMessageKind::PrefillEmbd,
563            0,
564            i32::try_from(MAX_STAGE_SIDEBAND_VALUES + 1).unwrap(),
565            0,
566            state,
567        );
568
569        assert_invalid_data(
570            read_stage_message(Cursor::new(bytes), 2048),
571            "token sideband count exceeds maximum",
572        );
573    }
574
575    #[test]
576    fn stage_message_rejects_position_sideband_count_over_limit() {
577        let mut state =
578            StageStateHeader::new(WireMessageKind::PrefillEmbd, WireActivationDType::F32);
579        state.source_stage_index = -1;
580        let bytes = stage_frame_prefix(
581            WireMessageKind::PrefillEmbd,
582            0,
583            0,
584            i32::try_from(MAX_STAGE_SIDEBAND_VALUES + 1).unwrap(),
585            state,
586        );
587
588        assert_invalid_data(
589            read_stage_message(Cursor::new(bytes), 2048),
590            "position sideband count exceeds maximum",
591        );
592    }
593
594    #[test]
595    fn prefill_wire_overhead_is_fixed_and_bounded() {
596        let mut state =
597            StageStateHeader::new(WireMessageKind::PrefillEmbd, WireActivationDType::F32);
598        state.prompt_token_count = 128;
599        state.current_token = 127;
600        state.source_stage_index = -1;
601        let tokens: Vec<i32> = (0..128).collect();
602        let message = StageWireMessage {
603            kind: WireMessageKind::PrefillEmbd,
604            pos_start: 0,
605            token_count: tokens.len() as i32,
606            state,
607            request_id: u64::MAX - 1,
608            session_id: u64::MAX,
609            sampling: None,
610            chat_sampling_metadata: None,
611            tokens,
612            positions: Vec::new(),
613            activation: Vec::new(),
614            raw_bytes: Vec::new(),
615        };
616        let mut bytes = Vec::new();
617        write_stage_message(&mut bytes, &message, WireActivationDType::F32).unwrap();
618
619        assert_eq!(STAGE_STATE_HEADER_BYTES, 40);
620        assert_eq!(STAGE_SAMPLING_CONFIG_BASE_BYTES, 40);
621        assert_eq!(STAGE_WIRE_FIXED_HEADER_BYTES, 76);
622        assert_eq!(
623            bytes.len(),
624            STAGE_WIRE_FIXED_HEADER_BYTES + message.tokens.len() * 4
625        );
626        const { assert!(STAGE_WIRE_FIXED_HEADER_BYTES <= 80) };
627    }
628
629    #[test]
630    fn verify_retirement_round_trips_exact_identity() {
631        let kind = WireMessageKind::RetireVerifyWindow;
632        let message = StageWireMessage {
633            kind,
634            pos_start: 128,
635            token_count: 8,
636            state: StageStateHeader::new(kind, WireActivationDType::F32),
637            request_id: 23,
638            session_id: 29,
639            sampling: None,
640            chat_sampling_metadata: None,
641            tokens: Vec::new(),
642            positions: Vec::new(),
643            activation: Vec::new(),
644            raw_bytes: Vec::new(),
645        };
646        let mut bytes = Vec::new();
647        write_stage_message(&mut bytes, &message, WireActivationDType::F32).unwrap();
648        let decoded = read_stage_message(Cursor::new(bytes), 2048).unwrap();
649
650        assert_eq!(decoded.kind, kind);
651        assert_eq!(decoded.pos_start, 128);
652        assert_eq!(decoded.token_count, 8);
653        assert!(decoded.state.matches_kind(kind));
654    }
655
656    #[test]
657    fn session_control_messages_are_fixed_header_only() {
658        let kind = WireMessageKind::TrimSession;
659        let message = StageWireMessage {
660            kind,
661            pos_start: 0,
662            token_count: 0,
663            state: StageStateHeader::new(kind, WireActivationDType::F32),
664            request_id: 23,
665            session_id: 29,
666            sampling: None,
667            chat_sampling_metadata: None,
668            tokens: Vec::new(),
669            positions: Vec::new(),
670            activation: Vec::new(),
671            raw_bytes: Vec::new(),
672        };
673        let mut bytes = Vec::new();
674        write_stage_message(&mut bytes, &message, WireActivationDType::F32).unwrap();
675        assert_eq!(bytes.len(), STAGE_WIRE_FIXED_HEADER_BYTES);
676        let decoded = read_stage_message(Cursor::new(bytes), 2048).unwrap();
677        assert_eq!(decoded.kind, kind);
678        assert_eq!(decoded.request_id, 23);
679        assert_eq!(decoded.session_id, 29);
680        assert!(decoded.tokens.is_empty());
681        assert!(decoded.activation.is_empty());
682    }
683
684    #[test]
685    fn state_import_message_round_trips_raw_bytes() {
686        let state = StageStateHeader::new(WireMessageKind::StateImport, WireActivationDType::F32);
687        let message = StageWireMessage {
688            kind: WireMessageKind::StateImport,
689            pos_start: 0,
690            token_count: 4,
691            state,
692            request_id: 31,
693            session_id: 37,
694            sampling: None,
695            chat_sampling_metadata: None,
696            tokens: Vec::new(),
697            positions: Vec::new(),
698            activation: Vec::new(),
699            raw_bytes: vec![1, 2, 3, 4],
700        };
701        let mut bytes = Vec::new();
702        write_stage_message(&mut bytes, &message, WireActivationDType::F32).unwrap();
703        let decoded = read_stage_message(Cursor::new(bytes), 2048).unwrap();
704        assert_eq!(decoded.kind, WireMessageKind::StateImport);
705        assert_eq!(decoded.raw_bytes, vec![1, 2, 3, 4]);
706        assert!(decoded.tokens.is_empty());
707        assert!(decoded.activation.is_empty());
708    }
709
710    #[test]
711    fn state_import_rejects_raw_byte_count_over_limit() {
712        let state = StageStateHeader::new(WireMessageKind::StateImport, WireActivationDType::F32);
713        let bytes = stage_frame_prefix(
714            WireMessageKind::StateImport,
715            i32::try_from(MAX_STAGE_STATE_IMPORT_BYTES + 1).unwrap(),
716            0,
717            0,
718            state,
719        );
720
721        assert_invalid_data(
722            read_stage_message(Cursor::new(bytes), 2048),
723            "state import byte count exceeds maximum",
724        );
725    }
726
727    #[test]
728    fn state_import_writer_rejects_raw_byte_count_mismatch() {
729        let state = StageStateHeader::new(WireMessageKind::StateImport, WireActivationDType::F32);
730        let message = StageWireMessage {
731            kind: WireMessageKind::StateImport,
732            pos_start: 0,
733            token_count: 8,
734            state,
735            request_id: 31,
736            session_id: 37,
737            sampling: None,
738            chat_sampling_metadata: None,
739            tokens: Vec::new(),
740            positions: Vec::new(),
741            activation: Vec::new(),
742            raw_bytes: vec![1, 2, 3, 4],
743        };
744        let mut bytes = Vec::new();
745        let error = write_stage_message(&mut bytes, &message, WireActivationDType::F32)
746            .expect_err("mismatched state import byte count should fail");
747        assert_eq!(error.kind(), std::io::ErrorKind::InvalidInput);
748        assert_eq!(error.to_string(), "state import raw byte count mismatch");
749    }
750
751    #[test]
752    fn state_export_message_round_trips_without_payload() {
753        let state = StageStateHeader::new(WireMessageKind::StateExport, WireActivationDType::F32);
754        let message = StageWireMessage {
755            kind: WireMessageKind::StateExport,
756            pos_start: 0,
757            token_count: 0,
758            state,
759            request_id: 41,
760            session_id: 43,
761            sampling: None,
762            chat_sampling_metadata: None,
763            tokens: Vec::new(),
764            positions: Vec::new(),
765            activation: Vec::new(),
766            raw_bytes: Vec::new(),
767        };
768        let mut bytes = Vec::new();
769        write_stage_message(&mut bytes, &message, WireActivationDType::F32).unwrap();
770        let decoded = read_stage_message(Cursor::new(bytes), 2048).unwrap();
771        assert_eq!(decoded.kind, WireMessageKind::StateExport);
772        assert!(decoded.raw_bytes.is_empty());
773        assert!(decoded.tokens.is_empty());
774        assert!(decoded.activation.is_empty());
775    }
776
777    #[test]
778    fn stage_message_rejects_activation_payload_over_limit() {
779        let mut state =
780            StageStateHeader::new(WireMessageKind::DecodeEmbd, WireActivationDType::F32);
781        state.source_stage_index = 0;
782        state.flags |= state_flags::GEMMA3N_ALTUP_SIDEBAND;
783        let token_count = i32::try_from(MAX_STAGE_ACTIVATION_BYTES / 4 / 4 / 1024 + 1).unwrap();
784        let bytes = stage_frame_prefix(WireMessageKind::DecodeEmbd, token_count, 0, 0, state);
785
786        assert_invalid_data(
787            read_stage_message(Cursor::new(bytes), 1024),
788            "activation payload byte count exceeds maximum",
789        );
790    }
791
792    #[test]
793    fn stage_message_rejects_f16_activation_when_decoded_payload_exceeds_limit() {
794        let mut state =
795            StageStateHeader::new(WireMessageKind::DecodeEmbd, WireActivationDType::F16);
796        state.source_stage_index = 0;
797        let n_embd = 65_536;
798        let token_count =
799            i32::try_from(MAX_STAGE_DECODED_ACTIVATION_BYTES / 4 / n_embd as usize + 1).unwrap();
800        let wire_bytes = activation_wire_bytes_with_state_flags(
801            WireActivationDType::F16,
802            token_count,
803            n_embd,
804            0,
805        )
806        .unwrap();
807        assert!(wire_bytes <= MAX_STAGE_ACTIVATION_BYTES);
808        let bytes = stage_frame_prefix(WireMessageKind::DecodeEmbd, token_count, 0, 0, state);
809
810        assert_invalid_data(
811            read_stage_message(Cursor::new(bytes), n_embd),
812            "decoded activation payload byte count exceeds maximum",
813        );
814    }
815
816    #[test]
817    fn stage_message_rejects_q8_activation_when_decoded_payload_exceeds_limit() {
818        let mut state = StageStateHeader::new(WireMessageKind::DecodeEmbd, WireActivationDType::Q8);
819        state.source_stage_index = 0;
820        let n_embd = 65_536;
821        let token_count =
822            i32::try_from(MAX_STAGE_DECODED_ACTIVATION_BYTES / 4 / n_embd as usize + 1).unwrap();
823        let wire_bytes =
824            activation_wire_bytes_with_state_flags(WireActivationDType::Q8, token_count, n_embd, 0)
825                .unwrap();
826        assert!(wire_bytes <= MAX_STAGE_ACTIVATION_BYTES);
827        let bytes = stage_frame_prefix(WireMessageKind::DecodeEmbd, token_count, 0, 0, state);
828
829        assert_invalid_data(
830            read_stage_message(Cursor::new(bytes), n_embd),
831            "decoded activation payload byte count exceeds maximum",
832        );
833    }
834
835    #[test]
836    fn stage_message_rejects_q8_sideband_activation_when_decoded_payload_exceeds_limit() {
837        let mut state = StageStateHeader::new(WireMessageKind::DecodeEmbd, WireActivationDType::Q8);
838        state.source_stage_index = 0;
839        state.flags |= state_flags::GEMMA3N_ALTUP_SIDEBAND;
840        let n_embd = 65_536;
841        let token_count =
842            i32::try_from(MAX_STAGE_DECODED_ACTIVATION_BYTES / 4 / 4 / n_embd as usize + 1)
843                .unwrap();
844        let wire_bytes = activation_wire_bytes_with_state_flags(
845            WireActivationDType::Q8,
846            token_count,
847            n_embd,
848            state.flags,
849        )
850        .unwrap();
851        assert!(wire_bytes <= MAX_STAGE_ACTIVATION_BYTES);
852        let bytes = stage_frame_prefix(WireMessageKind::DecodeEmbd, token_count, 0, 0, state);
853
854        assert_invalid_data(
855            read_stage_message(Cursor::new(bytes), n_embd),
856            "decoded activation payload byte count exceeds maximum",
857        );
858    }
859
860    #[test]
861    fn activation_encoding_rejects_decoded_payload_over_limit_before_compression() {
862        let n_embd = 65_536;
863        let token_count =
864            i32::try_from(MAX_STAGE_DECODED_ACTIVATION_BYTES / 4 / n_embd as usize + 1).unwrap();
865
866        assert_invalid_data(
867            encode_f32_activation_payload(WireActivationDType::F16, token_count, n_embd, &[]),
868            "decoded activation payload byte count exceeds maximum",
869        );
870    }
871
872    #[test]
873    fn q8_payload_decodes_to_f32_bytes() {
874        let mut payload = Vec::new();
875        payload.extend_from_slice(&0.5_f32.to_le_bytes());
876        payload.extend_from_slice(&[2_u8, 254_u8]);
877        let decoded = activation::decode_q8_to_f32_bytes(&payload, 1, 2).unwrap();
878        let first = f32::from_le_bytes(decoded[0..4].try_into().unwrap());
879        let second = f32::from_le_bytes(decoded[4..8].try_into().unwrap());
880        assert_eq!(first, 1.0);
881        assert_eq!(second, -1.0);
882    }
883
884    #[test]
885    fn f32_payload_encodes_to_q8_and_decodes() {
886        let mut input = Vec::new();
887        input.extend_from_slice(&1.0_f32.to_le_bytes());
888        input.extend_from_slice(&(-1.0_f32).to_le_bytes());
889        let encoded = encode_f32_activation_payload(WireActivationDType::Q8, 1, 2, &input).unwrap();
890        let decoded = activation::decode_q8_to_f32_bytes(&encoded, 1, 2).unwrap();
891        let first = f32::from_le_bytes(decoded[0..4].try_into().unwrap());
892        let second = f32::from_le_bytes(decoded[4..8].try_into().unwrap());
893        assert!((first - 1.0).abs() < 0.01);
894        assert!((second + 1.0).abs() < 0.01);
895    }
896
897    #[test]
898    fn rwkv7_sideband_activation_round_trips() {
899        let mut state =
900            StageStateHeader::new(WireMessageKind::DecodeEmbd, WireActivationDType::F32);
901        state.source_stage_index = 0;
902        state.flags |= state_flags::RWKV7_V_FIRST_SIDEBAND;
903        let mut activation = Vec::new();
904        for value in [1.0_f32, 2.0, 3.0, 4.0] {
905            activation.extend_from_slice(&value.to_le_bytes());
906        }
907        let message = StageWireMessage {
908            kind: WireMessageKind::DecodeEmbd,
909            pos_start: 0,
910            token_count: 1,
911            state,
912            request_id: 7,
913            session_id: 9,
914            sampling: None,
915            chat_sampling_metadata: None,
916            tokens: vec![42],
917            positions: Vec::new(),
918            activation,
919            raw_bytes: Vec::new(),
920        };
921        let mut bytes = Vec::new();
922        write_stage_message(&mut bytes, &message, WireActivationDType::F32).unwrap();
923        let decoded = read_stage_message(Cursor::new(bytes), 2).unwrap();
924        assert_eq!(decoded.activation.len(), 16);
925        assert_eq!(
926            activation_frame_flags_from_state_flags(decoded.state.flags),
927            ACTIVATION_FLAG_RWKV7_V_FIRST
928        );
929        assert_eq!(
930            decoded.activation_f32_payload(2).unwrap(),
931            message.activation
932        );
933    }
934
935    #[test]
936    fn inkling_mtp_embedding_sideband_activation_round_trips() {
937        let mut state =
938            StageStateHeader::new(WireMessageKind::PrefillEmbd, WireActivationDType::F32);
939        state.source_stage_index = 0;
940        state.flags |= state_flags::INKLING_MTP_EMBD_SIDEBAND;
941        let mut activation = Vec::new();
942        for value in [1.0_f32, 2.0, 3.0, 4.0] {
943            activation.extend_from_slice(&value.to_le_bytes());
944        }
945        let message = StageWireMessage {
946            kind: WireMessageKind::PrefillEmbd,
947            pos_start: 0,
948            token_count: 1,
949            state,
950            request_id: 7,
951            session_id: 9,
952            sampling: None,
953            chat_sampling_metadata: None,
954            tokens: Vec::new(),
955            positions: Vec::new(),
956            activation,
957            raw_bytes: Vec::new(),
958        };
959        let mut bytes = Vec::new();
960        write_stage_message(&mut bytes, &message, WireActivationDType::F32).unwrap();
961        let decoded = read_stage_message(Cursor::new(bytes), 2).unwrap();
962        assert_eq!(decoded.activation.len(), 16);
963        assert_eq!(
964            activation_frame_flags_from_state_flags(decoded.state.flags),
965            ACTIVATION_FLAG_INKLING_MTP_EMBD
966        );
967        assert_eq!(
968            activation_state_flags_from_frame_flags(ACTIVATION_FLAG_INKLING_MTP_EMBD),
969            state_flags::INKLING_MTP_EMBD_SIDEBAND
970        );
971        assert_eq!(
972            decoded.activation_f32_payload(2).unwrap(),
973            message.activation
974        );
975    }
976
977    #[test]
978    fn f32_activation_payload_can_be_moved_without_clone() {
979        let state = StageStateHeader::new(WireMessageKind::DecodeEmbd, WireActivationDType::F32);
980        let activation = vec![1_u8, 2, 3, 4, 5, 6, 7, 8];
981        let mut message = StageWireMessage {
982            kind: WireMessageKind::DecodeEmbd,
983            pos_start: 0,
984            token_count: 1,
985            state,
986            request_id: 7,
987            session_id: 9,
988            sampling: None,
989            chat_sampling_metadata: None,
990            tokens: vec![42],
991            positions: Vec::new(),
992            activation: activation.clone(),
993            raw_bytes: Vec::new(),
994        };
995
996        let payload = message.take_activation_f32_payload(2).unwrap();
997
998        assert_eq!(payload, activation);
999        assert!(message.activation.is_empty());
1000    }
1001
1002    #[test]
1003    fn f32_activation_payload_clone_helper_preserves_wire_payload() {
1004        let state = StageStateHeader::new(WireMessageKind::DecodeEmbd, WireActivationDType::F32);
1005        let activation = vec![1_u8, 2, 3, 4, 5, 6, 7, 8];
1006        let message = StageWireMessage {
1007            kind: WireMessageKind::DecodeEmbd,
1008            pos_start: 0,
1009            token_count: 1,
1010            state,
1011            request_id: 7,
1012            session_id: 9,
1013            sampling: None,
1014            chat_sampling_metadata: None,
1015            tokens: vec![42],
1016            positions: Vec::new(),
1017            activation: activation.clone(),
1018            raw_bytes: Vec::new(),
1019        };
1020
1021        let payload = message.activation_f32_payload(2).unwrap();
1022
1023        assert_eq!(payload, activation);
1024        assert_eq!(message.activation, activation);
1025    }
1026
1027    #[test]
1028    fn gemma3n_altup_sideband_activation_round_trips() {
1029        let mut state =
1030            StageStateHeader::new(WireMessageKind::DecodeEmbd, WireActivationDType::F32);
1031        state.source_stage_index = 0;
1032        state.flags |= state_flags::GEMMA3N_ALTUP_SIDEBAND;
1033        let mut activation = Vec::new();
1034        for value in 0..8 {
1035            activation.extend_from_slice(&(value as f32).to_le_bytes());
1036        }
1037        let message = StageWireMessage {
1038            kind: WireMessageKind::DecodeEmbd,
1039            pos_start: 0,
1040            token_count: 1,
1041            state,
1042            request_id: 7,
1043            session_id: 9,
1044            sampling: None,
1045            chat_sampling_metadata: None,
1046            tokens: vec![42],
1047            positions: Vec::new(),
1048            activation,
1049            raw_bytes: Vec::new(),
1050        };
1051        let mut bytes = Vec::new();
1052        write_stage_message(&mut bytes, &message, WireActivationDType::F32).unwrap();
1053        let decoded = read_stage_message(Cursor::new(bytes), 2).unwrap();
1054        assert_eq!(decoded.activation.len(), 32);
1055        assert_eq!(
1056            activation_frame_flags_from_state_flags(decoded.state.flags),
1057            ACTIVATION_FLAG_GEMMA3N_ALTUP
1058        );
1059        assert_eq!(
1060            decoded.activation_f32_payload(2).unwrap(),
1061            message.activation
1062        );
1063    }
1064}