Skip to main content

fib_quant/kv/
stream.rs

1//! Streaming KV-cache encoder for incremental token-by-token construction.
2//!
3//! [`KvStreamEncoder`] allows building a [`KvEncodedTensorV1`] one token at a
4//! time instead of requiring the full tensor upfront. Pages are flushed
5//! automatically when they reach `tokens_per_page` from the page geometry.
6//!
7//! The streaming encoder produces output identical to [`encode_kv_tensor`] for
8//! the same input data, shape, layout, and profile. The source digest is
9//! accumulated incrementally via [`blake3::Hasher`] and finalized in
10//! [`KvStreamEncoder::finish`].
11//!
12//! # Limitations
13//!
14//! The current implementation supports `batch = 1, layers = 1, kv_heads = 1`
15//! shapes (the common single-stream inference case). Each call to
16//! [`KvStreamEncoder::append_token`] produces exactly one encoded block for
17//! the matching role.
18
19use crate::{FibQuantError, FibQuantizer, Result};
20
21use super::{
22    block::{KvBlockEncodingV1, KvEncodedBlockV1},
23    codec::KvEncodedTensorV1,
24    layout::KvCacheLayoutV1,
25    page::KvEncodedPageV1,
26    profile::{KvAxisPolicyV1, KvCompressionProfileV1, KvFallbackModeV1},
27    receipt::{now_unix_seconds, KvCompressionReceiptV1, KvOperationKindV1, KV_RECEIPT_SCHEMA},
28    shape::{KvRole, KvTensorShapeV1},
29};
30
31/// Receipt returned after appending a single token to the stream encoder.
32#[derive(Debug, Clone, PartialEq)]
33pub struct AppendReceipt {
34    /// Global token index that was appended.
35    pub token_index: u32,
36    /// Encoding used for the key block: `"fib_quant"` or `"raw"`.
37    pub key_block_encoding: String,
38    /// Encoding used for the value block: `"fib_quant"` or `"raw"`.
39    pub value_block_encoding: String,
40    /// Compressed byte size of the key block.
41    pub key_compressed_bytes: usize,
42    /// Compressed byte size of the value block.
43    pub value_compressed_bytes: usize,
44    /// Fallback reasons (if any) for this token.
45    pub fallback_reasons: Vec<String>,
46}
47
48/// Incremental/streaming KV-cache encoder.
49///
50/// Construct with [`KvStreamEncoder::new`], feed tokens one at a time with
51/// [`KvStreamEncoder::append_token`], and finalize with
52/// [`KvStreamEncoder::finish`] to obtain a [`KvEncodedTensorV1`].
53///
54/// The encoder accumulates a BLAKE3 source digest over all appended vectors
55/// using an incremental [`blake3::Hasher`]. Pages are buffered as pending
56/// (blocks + metadata) and materialized into [`KvEncodedPageV1`] objects in
57/// `finish()` once the full source digest is available, ensuring that page
58/// digests match the batch [`super::codec::encode_kv_tensor`] path.
59pub struct KvStreamEncoder {
60    /// Logical tensor shape.
61    shape: KvTensorShapeV1,
62    /// Physical layout.
63    layout: KvCacheLayoutV1,
64    /// Compression profile.
65    profile: KvCompressionProfileV1,
66    /// FibQuant quantizer built from the profile.
67    quantizer: FibQuantizer,
68    /// Blocks being accumulated for the current page.
69    current_page_blocks: Vec<KvEncodedBlockV1>,
70    /// Token index of the first token in the current page.
71    current_page_token_start: u32,
72    /// Materialized pages (filled in `finish()`).
73    completed_pages: Vec<KvEncodedPageV1>,
74    /// Incremental source digest accumulator.
75    source_digest_state: blake3::Hasher,
76
77    // ── private tracking state ───────────────────────────────────────────
78    /// Next token index to assign (also equals number of tokens appended so far).
79    token_index: u32,
80    /// Next page id to assign.
81    page_id: u32,
82    /// Running count of compressed (non-fallback) blocks.
83    compressed_blocks: u32,
84    /// Running count of raw fallback blocks.
85    raw_fallback_blocks: u32,
86    /// Deduplicated fallback reasons across all tokens.
87    fallback_reasons: Vec<String>,
88    /// Pending pages: (page_id, token_start, token_count, blocks).
89    /// Materialized into `completed_pages` in `finish()`.
90    pending_pages: Vec<(u32, u32, u32, Vec<KvEncodedBlockV1>)>,
91}
92
93impl KvStreamEncoder {
94    /// Create a new streaming encoder.
95    ///
96    /// Validates the shape, layout, and profile, then builds a quantizer from
97    /// the profile. The source digest accumulator is initialized with the
98    /// domain tag and expected total element count so that incremental
99    /// hashing produces the same result as [`super::receipt::kv_tensor_digest`].
100    ///
101    /// Currently requires `batch = 1, layers = 1, kv_heads = 1`.
102    pub fn new(
103        shape: KvTensorShapeV1,
104        layout: KvCacheLayoutV1,
105        profile: KvCompressionProfileV1,
106    ) -> Result<Self> {
107        shape.validate()?;
108        layout.validate_for_shape(&shape)?;
109        profile.validate_for_shape(&shape)?;
110
111        // The streaming encoder produces one block per token (one key vector
112        // or one value vector). Multi-batch/multi-layer/multi-head shapes
113        // would require multiple vectors per token and are not supported.
114        if shape.batch != 1 || shape.layers != 1 || shape.kv_heads != 1 {
115            return Err(FibQuantError::DependencyUnsupported(
116                "streaming encoder currently supports batch=1, layers=1, kv_heads=1 only".into(),
117            ));
118        }
119
120        let quantizer = build_quantizer(&profile)?;
121
122        // Initialize the incremental BLAKE3 hasher with the same prefix that
123        // `kv_tensor_digest` uses, so the finalized digest matches.
124        let mut source_digest_state = blake3::Hasher::new();
125        source_digest_state.update(b"fib_quant_kv_tensor_f32_v1");
126        source_digest_state.update(&[0]);
127        let total_elements = shape.element_count()? as u64;
128        source_digest_state.update(&total_elements.to_le_bytes());
129
130        Ok(Self {
131            shape,
132            layout,
133            profile,
134            quantizer,
135            current_page_blocks: Vec::new(),
136            current_page_token_start: 0,
137            completed_pages: Vec::new(),
138            source_digest_state,
139            token_index: 0,
140            page_id: 0,
141            compressed_blocks: 0,
142            raw_fallback_blocks: 0,
143            fallback_reasons: Vec::new(),
144            pending_pages: Vec::new(),
145        })
146    }
147
148    /// Append one token's key and value vectors.
149    ///
150    /// Both vectors must have length `head_dim` and contain only finite values.
151    /// Only the vector matching `shape.role` is encoded into a block; the other
152    /// is reported as `"raw"` in the receipt with zero compressed bytes (it is
153    /// not stored in this encoder's output).
154    ///
155    /// When the current page reaches `tokens_per_page` blocks, the blocks are
156    /// flushed to a pending page buffer (materialized in `finish()`).
157    pub fn append_token(
158        &mut self,
159        key_vector: &[f32],
160        value_vector: &[f32],
161    ) -> Result<AppendReceipt> {
162        let head_dim = self.shape.head_dim as usize;
163
164        // Validate vector lengths.
165        if key_vector.len() != head_dim {
166            return Err(FibQuantError::CorruptPayload(format!(
167                "key_vector has {} elements, expected head_dim={}",
168                key_vector.len(),
169                head_dim
170            )));
171        }
172        if value_vector.len() != head_dim {
173            return Err(FibQuantError::CorruptPayload(format!(
174                "value_vector has {} elements, expected head_dim={}",
175                value_vector.len(),
176                head_dim
177            )));
178        }
179
180        // Check for non-finite values.
181        if key_vector.iter().any(|v| !v.is_finite()) {
182            return Err(FibQuantError::CorruptPayload(
183                "key_vector contains non-finite value".into(),
184            ));
185        }
186        if value_vector.iter().any(|v| !v.is_finite()) {
187            return Err(FibQuantError::CorruptPayload(
188                "value_vector contains non-finite value".into(),
189            ));
190        }
191
192        // Guard against appending more tokens than the shape declares.
193        if self.token_index >= self.shape.tokens {
194            return Err(FibQuantError::CorruptPayload(format!(
195                "stream encoder received token {} but shape only has {} tokens",
196                self.token_index, self.shape.tokens
197            )));
198        }
199
200        let token = self.token_index;
201
202        // Select the vector matching the encoder's role and feed it to the
203        // incremental source digest.
204        let (encode_vector, other_vector) = match self.shape.role {
205            KvRole::Key => (key_vector, value_vector),
206            KvRole::Value => (value_vector, key_vector),
207        };
208        for v in encode_vector {
209            self.source_digest_state.update(&v.to_le_bytes());
210        }
211        // The non-matching vector is NOT included in the source digest — the
212        // batch encoder only hashes the single-role tensor values.
213        let _ = other_vector; // suppress unused warning
214
215        // Encode the matching vector into a block.
216        let block_id = self.current_page_blocks.len() as u32;
217        let protected = self
218            .profile
219            .protected_policy
220            .is_protected(&self.shape, 0, 0, token);
221
222        let block = if protected {
223            KvEncodedBlockV1::raw(
224                block_id,
225                0,
226                0,
227                0,
228                token,
229                encode_vector.to_vec(),
230                self.profile.page_geometry.encoded_block_bytes,
231                "protected_region",
232            )
233        } else {
234            self.encode_vector_block(block_id, token, encode_vector)?
235        };
236
237        // Track stats.
238        let mut token_fallback_reasons = Vec::new();
239        if block.raw_fallback {
240            self.raw_fallback_blocks += 1;
241            if !self.fallback_reasons.contains(&block.reason) {
242                self.fallback_reasons.push(block.reason.clone());
243            }
244            token_fallback_reasons.push(block.reason.clone());
245        } else {
246            self.compressed_blocks += 1;
247        }
248
249        // Determine encoding type and compressed bytes for the receipt.
250        let (encoding_type, compressed_bytes) = match &block.encoding {
251            KvBlockEncodingV1::FibQuant { code } => ("fib_quant", code.compact_size()),
252            KvBlockEncodingV1::RawF32 { values } => {
253                ("raw", values.len() * std::mem::size_of::<f32>())
254            }
255        };
256
257        self.current_page_blocks.push(block);
258        self.token_index += 1;
259
260        // Flush page when full.
261        let tokens_per_page = self.profile.page_geometry.tokens_per_page;
262        let tokens_in_current_page = self.token_index - self.current_page_token_start;
263        if tokens_in_current_page >= tokens_per_page {
264            self.flush_page();
265        }
266
267        // Build the receipt, reporting both key and value status.
268        let (key_encoding, key_bytes, value_encoding, value_bytes) = match self.shape.role {
269            KvRole::Key => (encoding_type, compressed_bytes, "raw", 0),
270            KvRole::Value => ("raw", 0, encoding_type, compressed_bytes),
271        };
272
273        Ok(AppendReceipt {
274            token_index: token,
275            key_block_encoding: key_encoding.to_string(),
276            value_block_encoding: value_encoding.to_string(),
277            key_compressed_bytes: key_bytes,
278            value_compressed_bytes: value_bytes,
279            fallback_reasons: token_fallback_reasons,
280        })
281    }
282
283    /// Finalize the stream and produce the encoded tensor.
284    ///
285    /// Flushes any remaining blocks as the final page, computes the source
286    /// digest, materializes all pending pages, and builds the compression
287    /// receipt.
288    ///
289    /// Returns an error if the number of appended tokens does not match
290    /// `shape.tokens`.
291    pub fn finish(mut self) -> Result<KvEncodedTensorV1> {
292        // Validate token count.
293        if self.token_index == 0 {
294            return Err(FibQuantError::CorruptPayload(
295                "stream encoder finished without any appended tokens".into(),
296            ));
297        }
298        if self.token_index != self.shape.tokens {
299            return Err(FibQuantError::CorruptPayload(format!(
300                "stream encoder finished with {} tokens, expected {}",
301                self.token_index, self.shape.tokens
302            )));
303        }
304
305        // Flush remaining blocks as the last page.
306        if !self.current_page_blocks.is_empty() {
307            self.flush_page();
308        }
309
310        // Finalize the incremental source digest.
311        let source_digest = format!("blake3:{}", self.source_digest_state.finalize().to_hex());
312        let profile_digest = self.profile.digest(&self.shape)?;
313
314        // Materialize all pending pages.
315        for (page_id, token_start, token_count, blocks) in self.pending_pages.drain(..) {
316            let page = KvEncodedPageV1::new(
317                page_id,
318                token_start,
319                token_count,
320                source_digest.clone(),
321                profile_digest.clone(),
322                &self.shape,
323                self.profile.page_geometry.clone(),
324                blocks,
325            )?;
326            self.completed_pages.push(page);
327        }
328
329        // Build the compression receipt.
330        let page_digests = self
331            .completed_pages
332            .iter()
333            .map(|p| p.page_digest.clone())
334            .collect();
335
336        let receipt = KvCompressionReceiptV1 {
337            schema_version: KV_RECEIPT_SCHEMA.into(),
338            operation_kind: KvOperationKindV1::Compress,
339            source_digest,
340            profile_digest,
341            shape_digest: self.shape.digest()?,
342            page_digests,
343            codebook_digest: self.profile.codebook_digest.clone(),
344            rotation_digest: self.profile.rotation_digest.clone(),
345            encoded_pages: self.completed_pages.len() as u32,
346            compressed_blocks: self.compressed_blocks,
347            raw_fallback_blocks: self.raw_fallback_blocks,
348            fallback_reasons: std::mem::take(&mut self.fallback_reasons),
349            recorded_unix_seconds: now_unix_seconds(),
350        };
351
352        Ok(KvEncodedTensorV1 {
353            shape: self.shape,
354            layout: self.layout,
355            profile: self.profile,
356            pages: std::mem::take(&mut self.completed_pages),
357            receipt,
358        })
359    }
360
361    /// Flush the current page blocks to the pending buffer.
362    fn flush_page(&mut self) {
363        if self.current_page_blocks.is_empty() {
364            return;
365        }
366        let token_start = self.current_page_token_start;
367        let token_count = self.token_index - token_start;
368        let blocks = std::mem::take(&mut self.current_page_blocks);
369        self.pending_pages
370            .push((self.page_id, token_start, token_count, blocks));
371        self.page_id += 1;
372        self.current_page_token_start = self.token_index;
373    }
374
375    /// Encode a single vector block, replicating the per-token logic from
376    /// `codec::encode_vector_block`.
377    fn encode_vector_block(
378        &self,
379        block_id: u32,
380        token: u32,
381        vector: &[f32],
382    ) -> Result<KvEncodedBlockV1> {
383        match self.profile.axis_policy {
384            KvAxisPolicyV1::Raw => Ok(KvEncodedBlockV1::raw(
385                block_id,
386                0,
387                0,
388                0,
389                token,
390                vector.to_vec(),
391                self.profile.page_geometry.encoded_block_bytes,
392                "raw_axis_policy",
393            )),
394            KvAxisPolicyV1::PerToken => match self.quantizer.encode(vector) {
395                Ok(code) => Ok(KvEncodedBlockV1::fib_quant(
396                    block_id,
397                    0,
398                    0,
399                    0,
400                    token,
401                    code,
402                    self.profile.page_geometry.encoded_block_bytes,
403                    "fib_quant_per_token",
404                )),
405                Err(err) if self.profile.fallback_policy.mode == KvFallbackModeV1::KeepRaw => {
406                    Ok(KvEncodedBlockV1::raw(
407                        block_id,
408                        0,
409                        0,
410                        0,
411                        token,
412                        vector.to_vec(),
413                        self.profile.page_geometry.encoded_block_bytes,
414                        format!("encode_fallback:{err}"),
415                    ))
416                }
417                Err(err) => Err(err),
418            },
419            KvAxisPolicyV1::PerChannel | KvAxisPolicyV1::RoleAwareKiviStyle => {
420                if self.profile.fallback_policy.mode == KvFallbackModeV1::KeepRaw {
421                    Ok(KvEncodedBlockV1::raw(
422                        block_id,
423                        0,
424                        0,
425                        0,
426                        token,
427                        vector.to_vec(),
428                        self.profile.page_geometry.encoded_block_bytes,
429                        "unsupported_axis_raw_fallback",
430                    ))
431                } else {
432                    Err(FibQuantError::DependencyUnsupported(
433                        "CPU reference codec supports per-token FibQuant compression only".into(),
434                    ))
435                }
436            }
437        }
438    }
439}
440
441/// Build a quantizer from a compression profile, verifying the codebook digest.
442///
443/// Replicates the private `build_quantizer` in `codec.rs` since that function
444/// is not exported.
445fn build_quantizer(profile: &KvCompressionProfileV1) -> Result<FibQuantizer> {
446    let quantizer = FibQuantizer::new(profile.fib_profile.clone())?;
447    if quantizer.codebook().codebook_digest != profile.codebook_digest {
448        return Err(FibQuantError::CodebookDigestMismatch {
449            expected: quantizer.codebook().codebook_digest.clone(),
450            actual: profile.codebook_digest.clone(),
451        });
452    }
453    Ok(quantizer)
454}
455
456// ────────────────────────────────────────────────────────────────────────────
457//  Tests
458// ────────────────────────────────────────────────────────────────────────────
459
460#[cfg(test)]
461mod tests {
462    use super::super::codec::encode_kv_tensor;
463    use super::super::layout::KvPageGeometryV1;
464    use super::super::profile::KvAxisPolicyV1;
465    use super::super::shape::{KvAttentionKind, KvDType, KvRopeState};
466    use super::*;
467    use crate::profile::FibQuantProfileV1;
468
469    /// Build a shape/layout/profile suitable for the streaming encoder.
470    ///
471    /// Shape: batch=1, layers=1, kv_heads=1, tokens=3, head_dim=8
472    /// Page geometry: tokens_per_page=2 (so 2 pages — boundary at token 2)
473    /// FibQuant profile: k=4, N=32, ambient_dim=8
474    fn build_test_parts() -> (
475        KvTensorShapeV1,
476        KvCacheLayoutV1,
477        KvCompressionProfileV1,
478        Vec<f32>,
479    ) {
480        let shape = KvTensorShapeV1::new(
481            KvRole::Key,
482            KvAttentionKind::Mha,
483            1, // batch
484            1, // layers
485            1, // kv_heads
486            1, // query_heads (== kv_heads for MHA)
487            3, // tokens
488            8, // head_dim
489            KvDType::F32,
490            KvRopeState::PreRope,
491        );
492        let layout = KvCacheLayoutV1::canonical(&shape).expect("canonical layout");
493        let fib_profile =
494            FibQuantProfileV1::paper_default(8, 4, 32, 42).expect("build fib profile");
495        let quantizer = FibQuantizer::new(fib_profile.clone()).expect("build quantizer");
496        let page_geometry = KvPageGeometryV1::new(2, 8, 64); // tokens_per_page=2
497        let profile = KvCompressionProfileV1::from_parts(
498            "test-stream-profile",
499            &shape,
500            fib_profile,
501            quantizer.codebook().codebook_digest.clone(),
502            KvAxisPolicyV1::PerToken,
503            page_geometry,
504        )
505        .expect("build kv profile");
506
507        // Deterministic input values: 3 tokens * 8 head_dim = 24 values.
508        let total = shape.element_count().expect("element count");
509        let values: Vec<f32> = (0..total).map(|i| (i as f32) * 0.1).collect();
510
511        (shape, layout, profile, values)
512    }
513
514    #[test]
515    fn stream_matches_batch_encode() {
516        let (shape, layout, profile, values) = build_test_parts();
517
518        // Batch encode.
519        let batch_result =
520            encode_kv_tensor(shape.clone(), layout.clone(), profile.clone(), &values)
521                .expect("batch encode");
522
523        // Stream encode: one token at a time.
524        let mut encoder = KvStreamEncoder::new(shape.clone(), layout.clone(), profile.clone())
525            .expect("build stream encoder");
526        let head_dim = shape.head_dim as usize;
527        for token in 0..shape.tokens {
528            let start = token as usize * head_dim;
529            let key_slice = &values[start..start + head_dim];
530            // Pass the same slice as value_vector; only key is encoded for
531            // a Key-role shape.
532            encoder
533                .append_token(key_slice, key_slice)
534                .expect("append token");
535        }
536        let stream_result = encoder.finish().expect("stream finish");
537
538        // Compare pages and blocks (page digests depend on source digest
539        // which should match since we hash the same values in the same order).
540        assert_eq!(stream_result.pages.len(), batch_result.pages.len());
541        for (stream_page, batch_page) in stream_result.pages.iter().zip(batch_result.pages.iter()) {
542            assert_eq!(stream_page.page_id, batch_page.page_id);
543            assert_eq!(stream_page.token_start, batch_page.token_start);
544            assert_eq!(stream_page.token_count, batch_page.token_count);
545            assert_eq!(
546                stream_page.source_tensor_digest,
547                batch_page.source_tensor_digest
548            );
549            assert_eq!(
550                stream_page.page_digest, batch_page.page_digest,
551                "page digest mismatch for page {}",
552                stream_page.page_id
553            );
554            assert_eq!(
555                stream_page.encoded_blocks.len(),
556                batch_page.encoded_blocks.len()
557            );
558            for (sb, bb) in stream_page
559                .encoded_blocks
560                .iter()
561                .zip(batch_page.encoded_blocks.iter())
562            {
563                assert_eq!(sb.block_id, bb.block_id);
564                assert_eq!(sb.token, bb.token);
565                assert_eq!(sb.raw_fallback, bb.raw_fallback);
566                assert_eq!(
567                    sb.encoding, bb.encoding,
568                    "block encoding mismatch for block {} (token {})",
569                    sb.block_id, sb.token
570                );
571            }
572        }
573
574        // Compare receipt fields (except recorded_unix_seconds which depends
575        // on wall-clock time).
576        assert_eq!(
577            stream_result.receipt.source_digest,
578            batch_result.receipt.source_digest
579        );
580        assert_eq!(
581            stream_result.receipt.profile_digest,
582            batch_result.receipt.profile_digest
583        );
584        assert_eq!(
585            stream_result.receipt.shape_digest,
586            batch_result.receipt.shape_digest
587        );
588        assert_eq!(
589            stream_result.receipt.page_digests,
590            batch_result.receipt.page_digests
591        );
592        assert_eq!(
593            stream_result.receipt.codebook_digest,
594            batch_result.receipt.codebook_digest
595        );
596        assert_eq!(
597            stream_result.receipt.rotation_digest,
598            batch_result.receipt.rotation_digest
599        );
600        assert_eq!(
601            stream_result.receipt.encoded_pages,
602            batch_result.receipt.encoded_pages
603        );
604        assert_eq!(
605            stream_result.receipt.compressed_blocks,
606            batch_result.receipt.compressed_blocks
607        );
608        assert_eq!(
609            stream_result.receipt.raw_fallback_blocks,
610            batch_result.receipt.raw_fallback_blocks
611        );
612        assert_eq!(
613            stream_result.receipt.fallback_reasons,
614            batch_result.receipt.fallback_reasons
615        );
616    }
617
618    #[test]
619    fn append_receipt_fields_correct() {
620        let (shape, layout, profile, values) = build_test_parts();
621        let mut encoder =
622            KvStreamEncoder::new(shape, layout, profile).expect("build stream encoder");
623        let head_dim = 8;
624
625        // Token 0 — should be in page 0.
626        let r0 = encoder
627            .append_token(&values[0..head_dim], &values[0..head_dim])
628            .expect("append token 0");
629        assert_eq!(r0.token_index, 0);
630        // Key role → key_block_encoding should be "fib_quant" (if compression
631        // succeeds) or "raw" (if fallback). Value should be "raw" with 0 bytes.
632        assert!(
633            r0.key_block_encoding == "fib_quant" || r0.key_block_encoding == "raw",
634            "unexpected key_block_encoding: {}",
635            r0.key_block_encoding
636        );
637        assert_eq!(r0.value_block_encoding, "raw");
638        assert_eq!(r0.value_compressed_bytes, 0);
639        if r0.key_block_encoding == "fib_quant" {
640            assert!(r0.key_compressed_bytes > 0);
641        } else {
642            assert_eq!(
643                r0.key_compressed_bytes,
644                head_dim * std::mem::size_of::<f32>()
645            );
646        }
647
648        // Token 1 — fills page 0 (tokens_per_page=2).
649        let r1 = encoder
650            .append_token(
651                &values[head_dim..2 * head_dim],
652                &values[head_dim..2 * head_dim],
653            )
654            .expect("append token 1");
655        assert_eq!(r1.token_index, 1);
656
657        // Token 2 — starts page 1.
658        let r2 = encoder
659            .append_token(
660                &values[2 * head_dim..3 * head_dim],
661                &values[2 * head_dim..3 * head_dim],
662            )
663            .expect("append token 2");
664        assert_eq!(r2.token_index, 2);
665
666        // Fallback reasons should be empty if all tokens compressed.
667        if r0.key_block_encoding == "fib_quant"
668            && r1.key_block_encoding == "fib_quant"
669            && r2.key_block_encoding == "fib_quant"
670        {
671            assert!(r0.fallback_reasons.is_empty());
672            assert!(r1.fallback_reasons.is_empty());
673            assert!(r2.fallback_reasons.is_empty());
674        }
675    }
676
677    #[test]
678    fn empty_stream_finish_returns_error() {
679        let (shape, layout, profile, _) = build_test_parts();
680        let encoder = KvStreamEncoder::new(shape, layout, profile).expect("build encoder");
681        let err = encoder.finish().unwrap_err();
682        assert!(
683            matches!(err, FibQuantError::CorruptPayload(ref msg)
684                if msg.contains("without any appended tokens")),
685            "expected empty-stream error, got: {err:?}"
686        );
687    }
688
689    #[test]
690    fn stream_decode_roundtrip() {
691        let (shape, layout, profile, values) = build_test_parts();
692
693        // Batch encode for reference.
694        let batch_encoded =
695            encode_kv_tensor(shape.clone(), layout.clone(), profile.clone(), &values)
696                .expect("batch encode");
697        let batch_decoded =
698            super::super::codec::decode_kv_pages(&batch_encoded).expect("batch decode");
699
700        // Stream encode.
701        let mut encoder = KvStreamEncoder::new(shape.clone(), layout.clone(), profile.clone())
702            .expect("build stream encoder");
703        let head_dim = shape.head_dim as usize;
704        for token in 0..shape.tokens {
705            let start = token as usize * head_dim;
706            encoder
707                .append_token(
708                    &values[start..start + head_dim],
709                    &values[start..start + head_dim],
710                )
711                .expect("append token");
712        }
713        let encoded = encoder.finish().expect("stream finish");
714
715        // Decode the stream output.
716        let decoded = super::super::codec::decode_kv_pages(&encoded).expect("decode");
717        assert_eq!(decoded.values.len(), values.len());
718
719        // The stream-decoded values should match the batch-decoded values
720        // exactly (same encoder, same quantizer → same codes).
721        assert_eq!(
722            decoded.values, batch_decoded.values,
723            "stream decode must match batch decode"
724        );
725    }
726
727    #[test]
728    fn too_many_tokens_returns_error() {
729        let (shape, layout, profile, values) = build_test_parts();
730        let mut encoder = KvStreamEncoder::new(shape, layout, profile).expect("build encoder");
731        let head_dim = 8;
732        // Append all 3 tokens.
733        for token in 0..3 {
734            let start = token as usize * head_dim;
735            encoder
736                .append_token(
737                    &values[start..start + head_dim],
738                    &values[start..start + head_dim],
739                )
740                .expect("append token");
741        }
742        // Try appending a 4th token.
743        let extra = vec![0.0f32; head_dim];
744        let err = encoder.append_token(&extra, &extra).unwrap_err();
745        assert!(
746            matches!(err, FibQuantError::CorruptPayload(ref msg)
747                if msg.contains("but shape only has")),
748            "expected too-many-tokens error, got: {err:?}"
749        );
750    }
751
752    #[test]
753    fn partial_stream_returns_error() {
754        let (shape, layout, profile, values) = build_test_parts();
755        let mut encoder = KvStreamEncoder::new(shape, layout, profile).expect("build encoder");
756        let head_dim = 8;
757        // Append only 2 of 3 tokens.
758        for token in 0..2 {
759            let start = token as usize * head_dim;
760            encoder
761                .append_token(
762                    &values[start..start + head_dim],
763                    &values[start..start + head_dim],
764                )
765                .expect("append token");
766        }
767        let err = encoder.finish().unwrap_err();
768        assert!(
769            matches!(err, FibQuantError::CorruptPayload(ref msg)
770                if msg.contains("finished with 2 tokens, expected 3")),
771            "expected partial-stream error, got: {err:?}"
772        );
773    }
774
775    #[test]
776    fn wrong_vector_length_returns_error() {
777        let (shape, layout, profile, _) = build_test_parts();
778        let mut encoder = KvStreamEncoder::new(shape, layout, profile).expect("build encoder");
779        let short = vec![0.0f32; 4]; // wrong length (should be 8)
780        let err = encoder.append_token(&short, &short).unwrap_err();
781        assert!(
782            matches!(err, FibQuantError::CorruptPayload(ref msg)
783                if msg.contains("key_vector has 4 elements")),
784            "expected wrong-length error, got: {err:?}"
785        );
786    }
787
788    #[test]
789    fn non_finite_value_returns_error() {
790        let (shape, layout, profile, _) = build_test_parts();
791        let mut encoder = KvStreamEncoder::new(shape, layout, profile).expect("build encoder");
792        let nan_vec = vec![f32::NAN; 8];
793        let err = encoder.append_token(&nan_vec, &nan_vec).unwrap_err();
794        assert!(
795            matches!(err, FibQuantError::CorruptPayload(ref msg)
796                if msg.contains("non-finite")),
797            "expected non-finite error, got: {err:?}"
798        );
799    }
800}