Skip to main content

urna_format/reader/
validate.rs

1//! Validation hooks called during `from_bytes` after structural
2//! parsing succeeds. Each function focuses on one invariant; callers
3//! get a typed error pointing at the offending section / dtype.
4
5use super::UrnaView;
6use crate::encoding::{Int4EmbeddingsView, Int8EmbeddingsView, expected_embeddings_size};
7use crate::error::UrnaError;
8use crate::layout::{
9    REQUIRED_SECTIONS, SECTION_EMBEDDINGS, SECTION_ENCODING_FLOAT16, SECTION_ENCODING_FSST,
10    SECTION_ENCODING_INT4, SECTION_ENCODING_INT8, SECTION_ENCODING_INTPACK, SECTION_ENCODING_RAW,
11    SECTION_ENCODING_TXT_STREAMS, SECTION_ENCODING_ZSTD, SECTION_ENCODING_ZSTD_DICT,
12    SECTION_SEARCH_CONTRACT,
13};
14use crate::sections::decode_search_contract;
15
16impl UrnaView<'_> {
17    pub(super) fn check_required_sections(&self) -> crate::Result<()> {
18        for (id, name) in REQUIRED_SECTIONS {
19            if !self.section_table.iter().any(|e| e.section_id == *id) {
20                return Err(UrnaError::MissingRequiredSection(name));
21            }
22        }
23        Ok(())
24    }
25
26    pub(super) fn validate_embeddings_layout(&self) -> crate::Result<()> {
27        let entry = self.entry(SECTION_EMBEDDINGS)?;
28        let dim = self.header.embedding_dim as usize;
29        let n = self.header.n_embeddings as usize;
30        let dtype = self.manifest.dtype.as_str();
31
32        // Encoding/dtype consistency: float16 dtype implies float16 encoding,
33        // int8 dtype implies int8 encoding, int4 dtype implies int4 encoding,
34        // float32 dtype implies raw or zstd (zstd on embeddings is rejected
35        // separately by validate_encoding_for_section).
36        let valid_combo = matches!(
37            (dtype, entry.encoding),
38            ("float32", SECTION_ENCODING_RAW)
39                | ("float16", SECTION_ENCODING_FLOAT16)
40                | ("int8", SECTION_ENCODING_INT8)
41                | ("int4", SECTION_ENCODING_INT4)
42        );
43        if !valid_combo {
44            return Err(UrnaError::ManifestInvalid(format!(
45                "embeddings section encoding={} does not match dtype={}",
46                entry.encoding, dtype
47            )));
48        }
49
50        let want = expected_embeddings_size(dtype, n, dim).ok_or_else(|| {
51            UrnaError::UnsupportedDType(format!(
52                "unknown embeddings dtype {dtype}, or n={n} x dim={dim} overflows"
53            ))
54        })?;
55        let got = entry.size as usize;
56        if got != want {
57            return Err(UrnaError::EmbeddingSizeMismatch {
58                expected: want,
59                got,
60            });
61        }
62        Ok(())
63    }
64
65    /// When a space_table (0x15) is present, every listed space must have
66    /// its band section (0x20 + space_index) present with exactly the size
67    /// its (n_vectors, dim, dtype) imply, and n_vectors must match the
68    /// corpus chunk count (the bands are parallel per-chunk embeddings).
69    /// runs after the structural parse, so the band encoding is already
70    /// known legal (dtype encodings only, never zstd).
71    pub(super) fn validate_space_bands(&self) -> crate::Result<()> {
72        use crate::layout::{SECTION_SPACE_EMBEDDINGS_BASE, SECTION_SPACE_TABLE};
73        use crate::sections::decode_space_table;
74        if self.entry(SECTION_SPACE_TABLE).is_err() {
75            return Ok(());
76        }
77        let entries = decode_space_table(&self.decoded_section(SECTION_SPACE_TABLE)?)?;
78        let n_chunks = self.header.n_chunks;
79        for e in &entries {
80            if e.n_vectors != n_chunks {
81                return Err(UrnaError::ManifestInvalid(format!(
82                    "space {} lists {} vectors but the corpus has {} chunks",
83                    e.name, e.n_vectors, n_chunks
84                )));
85            }
86            let band_id = SECTION_SPACE_EMBEDDINGS_BASE + e.space_index as u32;
87            let band = self.entry(band_id).map_err(|_| {
88                UrnaError::ManifestInvalid(format!(
89                    "space {} lists band {:#x} but the section is missing",
90                    e.name, band_id
91                ))
92            })?;
93            let want =
94                expected_embeddings_size(e.dtype_str(), e.n_vectors as usize, e.dim as usize)
95                    .ok_or_else(|| {
96                        UrnaError::UnsupportedDType(format!(
97                            "space {} dtype code {}",
98                            e.name, e.dtype
99                        ))
100                    })?;
101            if band.size as usize != want {
102                return Err(UrnaError::EmbeddingSizeMismatch {
103                    expected: want,
104                    got: band.size as usize,
105                });
106            }
107        }
108        Ok(())
109    }
110
111    pub(super) fn validate_search_contract(&self) -> crate::Result<()> {
112        let bytes = self.decoded_section(SECTION_SEARCH_CONTRACT)?;
113        let contract = decode_search_contract(&bytes)?;
114        if contract.metric != self.manifest.metric {
115            return Err(UrnaError::UnsupportedMetric(format!(
116                "section says {} but manifest says {}",
117                contract.metric, self.manifest.metric
118            )));
119        }
120        if contract.score_type != self.manifest.score_type {
121            return Err(UrnaError::UnsupportedScoreType(format!(
122                "section says {} but manifest says {}",
123                contract.score_type, self.manifest.score_type
124            )));
125        }
126        if contract.normalize != self.manifest.normalize {
127            return Err(UrnaError::UnsupportedNormalize(format!(
128                "section says {} but manifest says {}",
129                contract.normalize, self.manifest.normalize
130            )));
131        }
132        if contract.index_type != self.manifest.index_type {
133            return Err(UrnaError::UnsupportedIndexType(format!(
134                "section says {} but manifest says {}",
135                contract.index_type, self.manifest.index_type
136            )));
137        }
138        if contract.rerank_policy != self.manifest.rerank_policy {
139            return Err(UrnaError::UnsupportedRerankPolicy(format!(
140                "section says {} but manifest says {}",
141                contract.rerank_policy, self.manifest.rerank_policy
142            )));
143        }
144        Ok(())
145    }
146
147    /// Walk the embeddings section and reject any NaN/Inf value. Works
148    /// for all supported dtypes.
149    pub fn validate_embeddings_values(&self) -> crate::Result<()> {
150        self.validate_embeddings_layout()?;
151        let entry = self.entry(SECTION_EMBEDDINGS)?;
152        let data = self.get_section_data(SECTION_EMBEDDINGS)?;
153        let n = self.header.n_embeddings as usize;
154        let dim = self.header.embedding_dim as usize;
155        validate_slab_values(entry.encoding, data, n, dim)
156    }
157}
158
159/// Reject any NaN/Inf in a fixed-stride vector slab of the given section
160/// encoding (raw f32, float16, int8, int4). The canonical embeddings, the
161/// `embeddings_fp` rerank slab and every multimodal band go through this
162/// before a kernel ever scores them: a NaN lane would turn a cosine into a
163/// NaN score, and a NaN score is what breaks a ranking (found by the
164/// mutation fuzz harness through the space bands).
165pub fn validate_slab_values(encoding: u32, data: &[u8], n: usize, dim: usize) -> crate::Result<()> {
166    match encoding {
167        SECTION_ENCODING_RAW => {
168            for chunk in data.chunks_exact(4) {
169                let v = f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]);
170                if !v.is_finite() {
171                    return Err(UrnaError::InvalidEmbeddingValue);
172                }
173            }
174        }
175        SECTION_ENCODING_FLOAT16 => {
176            for chunk in data.chunks_exact(2) {
177                let v = half::f16::from_le_bytes([chunk[0], chunk[1]]).to_f32();
178                if !v.is_finite() {
179                    return Err(UrnaError::InvalidEmbeddingValue);
180                }
181            }
182        }
183        SECTION_ENCODING_INT8 => {
184            // i8 cannot encode NaN/Inf; only the per-vector scales could.
185            let view = Int8EmbeddingsView::parse(data, n, dim)?;
186            if (0..view.n).any(|i| !view.scale(i).is_finite()) {
187                return Err(UrnaError::InvalidEmbeddingValue);
188            }
189        }
190        SECTION_ENCODING_INT4 => {
191            // 4-bit codes cannot encode NaN/Inf; only the per-group f16
192            // absmax scales could.
193            let view = Int4EmbeddingsView::parse(data, n, dim)?;
194            for i in 0..view.n {
195                if (0..view.blocks).any(|g| !view.group_scale(i, g).is_finite()) {
196                    return Err(UrnaError::InvalidEmbeddingValue);
197                }
198            }
199        }
200        other => {
201            return Err(UrnaError::UnsupportedSectionEncoding {
202                section_id: SECTION_EMBEDDINGS,
203                encoding: other,
204            });
205        }
206    }
207    Ok(())
208}
209
210/// Encoding rules: the embeddings section gets dtype-specific encodings
211/// (float16, int8, int4) and rejects zstd (we want SIMD-friendly mmap
212/// reads). the per-space vector bands (0x20-0x2F and the fp rerank band
213/// 0x30-0x3F) follow the SAME rule: fixed-stride slabs scored by the simd
214/// kernels, NEVER zstd. All other sections accept raw, zstd, intpack (the
215/// content_hash-preserving repack for chunk_ids / spans), txt_streams (the
216/// per-chunk-streams repack for chunks_canonical), zstd_dict (the
217/// trained-dictionary per-chunk-streams variant, decoded against section
218/// 0x0A), or fsst (the static-symbol-table per-chunk-streams variant). every
219/// non-embedding codec decodes byte-identically to the raw payload, so
220/// content_hash stays stable.
221pub(super) fn validate_encoding_for_section(section_id: u32, encoding: u32) -> crate::Result<()> {
222    use crate::layout::{
223        SECTION_SPACE_EMBEDDINGS_BASE, SECTION_SPACE_EMBEDDINGS_FP_BASE, SPACE_BAND_LEN,
224    };
225    let is_space_band = (SECTION_SPACE_EMBEDDINGS_BASE
226        ..SECTION_SPACE_EMBEDDINGS_BASE + SPACE_BAND_LEN)
227        .contains(&section_id)
228        || (SECTION_SPACE_EMBEDDINGS_FP_BASE..SECTION_SPACE_EMBEDDINGS_FP_BASE + SPACE_BAND_LEN)
229            .contains(&section_id);
230    let allowed = if section_id == SECTION_EMBEDDINGS || is_space_band {
231        matches!(
232            encoding,
233            SECTION_ENCODING_RAW
234                | SECTION_ENCODING_FLOAT16
235                | SECTION_ENCODING_INT8
236                | SECTION_ENCODING_INT4
237        )
238    } else {
239        matches!(
240            encoding,
241            SECTION_ENCODING_RAW
242                | SECTION_ENCODING_ZSTD
243                | SECTION_ENCODING_INTPACK
244                | SECTION_ENCODING_TXT_STREAMS
245                | SECTION_ENCODING_ZSTD_DICT
246                | SECTION_ENCODING_FSST
247        )
248    };
249    if !allowed {
250        return Err(UrnaError::UnsupportedSectionEncoding {
251            section_id,
252            encoding,
253        });
254    }
255    Ok(())
256}