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}