1use std::collections::BTreeSet;
4
5use sha2::{Digest, Sha256};
6
7use crate::{
8 EmbeddingContentDigest, EmbeddingNormalization, SearchArtifactError, VectorStoreLimits,
9 validate_vector, vector_schema,
10};
11
12#[derive(Clone, Debug, PartialEq)]
14pub struct EmbeddingBatchRow {
15 pub node_uuid: [u8; 16],
17 pub vector: Vec<f32>,
19}
20
21#[derive(Clone, Debug, PartialEq)]
23pub struct ValidatedEmbeddingBatch {
24 rows: Vec<EmbeddingBatchRow>,
25 dimension: usize,
26 content_digest: EmbeddingContentDigest,
27}
28
29impl ValidatedEmbeddingBatch {
30 #[must_use]
32 pub fn rows(&self) -> &[EmbeddingBatchRow] {
33 &self.rows
34 }
35
36 #[must_use]
38 pub const fn dimension(&self) -> usize {
39 self.dimension
40 }
41
42 #[must_use]
44 pub const fn content_digest(&self) -> EmbeddingContentDigest {
45 self.content_digest
46 }
47
48 #[must_use]
50 pub fn into_rows(self) -> Vec<EmbeddingBatchRow> {
51 self.rows
52 }
53}
54
55pub fn validate_embedding_batch<C>(
67 mut rows: Vec<EmbeddingBatchRow>,
68 eligible_nodes: &BTreeSet<[u8; 16]>,
69 dimension: usize,
70 normalization: EmbeddingNormalization,
71 limits: VectorStoreLimits,
72 mut checkpoint: C,
73) -> Result<ValidatedEmbeddingBatch, SearchArtifactError>
74where
75 C: FnMut() -> Result<(), SearchArtifactError>,
76{
77 checkpoint()?;
78 vector_schema(dimension, limits)?;
79 enforce_limit("embedding_rows", rows.len(), limits.stored_vectors)?;
80 enforce_limit(
81 "embedding_eligible_nodes",
82 eligible_nodes.len(),
83 limits.eligible_nodes,
84 )?;
85 let cells = rows
86 .len()
87 .checked_mul(dimension)
88 .ok_or_else(|| exhausted("embedding_vector_cells", limits.vector_cells))?;
89 enforce_limit("embedding_vector_cells", cells, limits.vector_cells)?;
90
91 rows.sort_unstable_by_key(|row| row.node_uuid);
92 for index in 0..rows.len() {
93 checkpoint()?;
94 if index != 0 && rows[index - 1].node_uuid == rows[index].node_uuid {
95 return Err(invalid("embedding batch", "contains duplicate node_uuid"));
96 }
97 if !eligible_nodes.contains(&rows[index].node_uuid) {
98 return Err(invalid(
99 "embedding batch",
100 "contains a node_uuid outside the eligible projection",
101 ));
102 }
103 if rows[index].vector.len() != dimension {
104 return Err(invalid(
105 "embedding batch",
106 format!(
107 "node_uuid has dimension {}, expected {dimension}",
108 rows[index].vector.len()
109 ),
110 ));
111 }
112 let squared_norm = validate_vector(&rows[index].vector, limits)?;
113 if normalization == EmbeddingNormalization::L2 {
114 normalize_l2(&mut rows[index].vector, squared_norm);
115 }
116 }
117
118 if rows.len() != eligible_nodes.len() {
119 return Err(invalid(
120 "embedding batch",
121 format!(
122 "missing {} eligible node_uuid rows",
123 eligible_nodes.len().saturating_sub(rows.len())
124 ),
125 ));
126 }
127
128 let mut hasher = Sha256::new();
129 for row in &rows {
130 checkpoint()?;
131 hasher.update(row.node_uuid);
132 for value in &row.vector {
133 hasher.update(value.to_le_bytes());
134 }
135 }
136 let content_digest = EmbeddingContentDigest::from_hex(&format!("{:x}", hasher.finalize()))?;
137 Ok(ValidatedEmbeddingBatch {
138 rows,
139 dimension,
140 content_digest,
141 })
142}
143
144#[allow(clippy::cast_possible_truncation)]
145fn normalize_l2(vector: &mut [f32], squared_norm: f64) {
146 let norm = squared_norm.sqrt();
149 for value in vector {
150 *value = (f64::from(*value) / norm) as f32;
151 }
152}
153
154fn enforce_limit(
155 resource: &'static str,
156 actual: usize,
157 limit: usize,
158) -> Result<(), SearchArtifactError> {
159 if actual > limit {
160 return Err(exhausted(resource, limit));
161 }
162 Ok(())
163}
164
165fn invalid(field: &'static str, reason: impl Into<String>) -> SearchArtifactError {
166 SearchArtifactError::InvalidSelector {
167 field,
168 reason: reason.into(),
169 }
170}
171
172fn exhausted(resource: &'static str, limit: usize) -> SearchArtifactError {
173 SearchArtifactError::ResourceExhausted {
174 resource,
175 limit: u64::try_from(limit).unwrap_or(u64::MAX),
176 }
177}
178
179#[cfg(test)]
180mod tests {
181 use std::sync::atomic::{AtomicUsize, Ordering};
182
183 use super::*;
184
185 const A: [u8; 16] = [1; 16];
186 const B: [u8; 16] = [2; 16];
187
188 fn row(node_uuid: [u8; 16], vector: &[f32]) -> EmbeddingBatchRow {
189 EmbeddingBatchRow {
190 node_uuid,
191 vector: vector.to_vec(),
192 }
193 }
194
195 fn eligible(values: &[[u8; 16]]) -> BTreeSet<[u8; 16]> {
196 values.iter().copied().collect()
197 }
198
199 #[test]
200 fn input_order_cannot_change_rows_or_little_endian_digest() {
201 let limits = VectorStoreLimits::default();
202 let expected_nodes = eligible(&[A, B]);
203 let left = validate_embedding_batch(
204 vec![row(B, &[3.0, 4.0]), row(A, &[1.0, 2.0])],
205 &expected_nodes,
206 2,
207 EmbeddingNormalization::None,
208 limits,
209 || Ok(()),
210 )
211 .unwrap();
212 let right = validate_embedding_batch(
213 vec![row(A, &[1.0, 2.0]), row(B, &[3.0, 4.0])],
214 &expected_nodes,
215 2,
216 EmbeddingNormalization::None,
217 limits,
218 || Ok(()),
219 )
220 .unwrap();
221
222 let mut expected_bytes = Vec::new();
223 expected_bytes.extend_from_slice(&A);
224 expected_bytes.extend_from_slice(&1.0_f32.to_le_bytes());
225 expected_bytes.extend_from_slice(&2.0_f32.to_le_bytes());
226 expected_bytes.extend_from_slice(&B);
227 expected_bytes.extend_from_slice(&3.0_f32.to_le_bytes());
228 expected_bytes.extend_from_slice(&4.0_f32.to_le_bytes());
229 assert_eq!(left, right);
230 assert_eq!(left.rows()[0].node_uuid, A);
231 assert_eq!(
232 left.content_digest(),
233 EmbeddingContentDigest::digest(&expected_bytes)
234 );
235 }
236
237 #[test]
238 fn normalization_contract_is_explicit() {
239 let complete = eligible(&[A]);
240 let preserved = validate_embedding_batch(
241 vec![row(A, &[3.0, 4.0])],
242 &complete,
243 2,
244 EmbeddingNormalization::None,
245 VectorStoreLimits::default(),
246 || Ok(()),
247 )
248 .unwrap();
249 let normalized = validate_embedding_batch(
250 vec![row(A, &[3.0, 4.0])],
251 &complete,
252 2,
253 EmbeddingNormalization::L2,
254 VectorStoreLimits::default(),
255 || Ok(()),
256 )
257 .unwrap();
258 assert_eq!(preserved.rows()[0].vector, vec![3.0, 4.0]);
259 assert_eq!(normalized.rows()[0].vector, vec![0.6, 0.8]);
260 assert_ne!(preserved.content_digest(), normalized.content_digest());
261 }
262
263 #[test]
264 fn coverage_and_vector_validation_fail_closed() {
265 let limits = VectorStoreLimits::default();
266 let complete = eligible(&[A, B]);
267 let invalid_batches = [
268 vec![row(A, &[1.0, 0.0]), row(A, &[0.0, 1.0])],
269 vec![row(A, &[1.0, 0.0]), row([3; 16], &[0.0, 1.0])],
270 vec![row(A, &[1.0, 0.0])],
271 vec![row(A, &[1.0]), row(B, &[0.0, 1.0])],
272 vec![row(A, &[f32::NAN, 1.0]), row(B, &[0.0, 1.0])],
273 vec![row(A, &[0.0, 0.0]), row(B, &[0.0, 1.0])],
274 ];
275 for rows in invalid_batches {
276 assert!(
277 validate_embedding_batch(
278 rows,
279 &complete,
280 2,
281 EmbeddingNormalization::None,
282 limits,
283 || Ok(())
284 )
285 .is_err()
286 );
287 }
288 }
289
290 #[test]
291 fn limits_and_cancellation_return_no_validated_batch() {
292 assert!(
293 validate_embedding_batch(
294 Vec::new(),
295 &BTreeSet::new(),
296 0,
297 EmbeddingNormalization::None,
298 VectorStoreLimits::default(),
299 || Ok(())
300 )
301 .is_err()
302 );
303
304 let mut limits = VectorStoreLimits::default();
305 limits.stored_vectors = 1;
306 assert!(matches!(
307 validate_embedding_batch(
308 vec![row(A, &[1.0]), row(B, &[1.0])],
309 &eligible(&[A, B]),
310 1,
311 EmbeddingNormalization::None,
312 limits,
313 || Ok(())
314 ),
315 Err(SearchArtifactError::ResourceExhausted {
316 resource: "embedding_rows",
317 ..
318 })
319 ));
320
321 let mut limits = VectorStoreLimits::default();
322 limits.eligible_nodes = 1;
323 assert!(matches!(
324 validate_embedding_batch(
325 vec![row(A, &[1.0]), row(B, &[1.0])],
326 &eligible(&[A, B]),
327 1,
328 EmbeddingNormalization::None,
329 limits,
330 || Ok(())
331 ),
332 Err(SearchArtifactError::ResourceExhausted {
333 resource: "embedding_eligible_nodes",
334 ..
335 })
336 ));
337
338 let mut limits = VectorStoreLimits::default();
339 limits.vector_cells = 1;
340 assert!(matches!(
341 validate_embedding_batch(
342 vec![row(A, &[1.0]), row(B, &[1.0])],
343 &eligible(&[A, B]),
344 1,
345 EmbeddingNormalization::None,
346 limits,
347 || Ok(())
348 ),
349 Err(SearchArtifactError::ResourceExhausted {
350 resource: "embedding_vector_cells",
351 ..
352 })
353 ));
354
355 let checkpoints = AtomicUsize::new(0);
356 assert!(matches!(
357 validate_embedding_batch(
358 vec![row(A, &[1.0]), row(B, &[1.0])],
359 &eligible(&[A, B]),
360 1,
361 EmbeddingNormalization::None,
362 VectorStoreLimits::default(),
363 || {
364 (checkpoints.fetch_add(1, Ordering::Relaxed) < 2)
365 .then_some(())
366 .ok_or(SearchArtifactError::Cancelled)
367 }
368 ),
369 Err(SearchArtifactError::Cancelled)
370 ));
371 }
372
373 #[test]
374 fn empty_complete_projection_retains_fixed_dimension() {
375 let batch = validate_embedding_batch(
376 Vec::new(),
377 &BTreeSet::new(),
378 3,
379 EmbeddingNormalization::L2,
380 VectorStoreLimits::default(),
381 || Ok(()),
382 )
383 .unwrap();
384 assert!(batch.rows().is_empty());
385 assert_eq!(batch.dimension(), 3);
386 assert_eq!(batch.content_digest(), EmbeddingContentDigest::digest(&[]));
387 }
388}