Skip to main content

graphforge_storage/
embedding_batch.rs

1//! Complete, canonical UUID/vector batches for embedding generations.
2
3use std::collections::BTreeSet;
4
5use sha2::{Digest, Sha256};
6
7use crate::{
8    EmbeddingContentDigest, EmbeddingNormalization, SearchArtifactError, VectorStoreLimits,
9    validate_vector, vector_schema,
10};
11
12/// One caller- or producer-supplied UUID/vector row before validation.
13#[derive(Clone, Debug, PartialEq)]
14pub struct EmbeddingBatchRow {
15    /// Stable graph identity; numeric execution surrogates are never accepted.
16    pub node_uuid: [u8; 16],
17    /// Fixed-dimension Float32 coordinates.
18    pub vector: Vec<f32>,
19}
20
21/// A complete validated batch in canonical raw-UUID order.
22#[derive(Clone, Debug, PartialEq)]
23pub struct ValidatedEmbeddingBatch {
24    rows: Vec<EmbeddingBatchRow>,
25    dimension: usize,
26    content_digest: EmbeddingContentDigest,
27}
28
29impl ValidatedEmbeddingBatch {
30    /// Canonical rows sorted by raw UUID bytes.
31    #[must_use]
32    pub fn rows(&self) -> &[EmbeddingBatchRow] {
33        &self.rows
34    }
35
36    /// Fixed vector width retained even when the complete projection is empty.
37    #[must_use]
38    pub const fn dimension(&self) -> usize {
39        self.dimension
40    }
41
42    /// SHA-256 of canonical UUID and little-endian Float32 bytes.
43    #[must_use]
44    pub const fn content_digest(&self) -> EmbeddingContentDigest {
45        self.content_digest
46    }
47
48    /// Consume the validated wrapper and return canonical rows.
49    #[must_use]
50    pub fn into_rows(self) -> Vec<EmbeddingBatchRow> {
51        self.rows
52    }
53}
54
55/// Validate and canonicalize one complete embedding-space batch.
56///
57/// `eligible_nodes` is the resolved source projection for this generation. The
58/// returned batch covers it exactly: every row must be eligible and every
59/// eligible UUID must have one row. `checkpoint` is called before validation,
60/// once per row, and once per digest row so cancellation never returns a
61/// partial validated value.
62///
63/// # Errors
64/// Rejects invalid dimensions/vectors, duplicate or ineligible UUIDs,
65/// incomplete coverage, configured resource exhaustion, and cancellation.
66pub 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    // The compatibility contract deliberately persists Float32 coordinates;
147    // accumulate the norm safely in Float64, then round each result to Float32.
148    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}