Skip to main content

autoagents_core/vector_store/
payload.rs

1use std::collections::{HashMap, HashSet};
2
3use serde::Serialize;
4
5use crate::embeddings::{Embed, Embedding, EmbeddingError, SharedEmbeddingProvider, TextEmbedder};
6use crate::one_or_many::OneOrMany;
7
8use super::{NamedVectorDocument, VectorStoreError};
9
10#[derive(Debug, Clone)]
11pub struct PayloadDocument<T> {
12    pub id: String,
13    pub raw: T,
14    pub payload_fields: HashMap<String, serde_json::Value>,
15}
16
17#[derive(Debug, Clone)]
18pub struct PreparedPayloadDocument {
19    pub id: String,
20    pub raw: serde_json::Value,
21    pub payload_fields: HashMap<String, serde_json::Value>,
22    pub embeddings: OneOrMany<Embedding>,
23}
24
25#[derive(Debug, Clone)]
26pub struct NamedVectorPayloadDocument<T> {
27    pub id: String,
28    pub raw: T,
29    pub vectors: HashMap<String, String>,
30    pub payload_fields: HashMap<String, serde_json::Value>,
31}
32
33#[derive(Debug, Clone)]
34pub struct PreparedNamedVectorPayloadDocument {
35    pub id: String,
36    pub raw: serde_json::Value,
37    pub payload_fields: HashMap<String, serde_json::Value>,
38    pub vectors: HashMap<String, Vec<f32>>,
39}
40
41impl<T> PayloadDocument<T> {
42    pub fn new(id: impl Into<String>, raw: T) -> Self {
43        Self {
44            id: id.into(),
45            raw,
46            payload_fields: HashMap::new(),
47        }
48    }
49
50    pub fn with_payload_fields(
51        mut self,
52        payload_fields: HashMap<String, serde_json::Value>,
53    ) -> Self {
54        self.payload_fields = payload_fields;
55        self
56    }
57}
58
59impl<T> PayloadDocument<T>
60where
61    T: Serialize,
62{
63    pub fn with_mirrored_payload_fields(
64        mut self,
65        fields: impl IntoIterator<Item = impl AsRef<str>>,
66    ) -> Result<Self, serde_json::Error> {
67        self.payload_fields = mirrored_payload_fields_for(&self.raw, fields)?;
68        Ok(self)
69    }
70}
71
72impl<T> NamedVectorDocument<T> {
73    pub fn with_payload_fields(
74        self,
75        payload_fields: HashMap<String, serde_json::Value>,
76    ) -> NamedVectorPayloadDocument<T> {
77        NamedVectorPayloadDocument {
78            id: self.id,
79            raw: self.raw,
80            vectors: self.vectors,
81            payload_fields,
82        }
83    }
84}
85
86impl<T> NamedVectorPayloadDocument<T> {
87    pub fn new(id: impl Into<String>, raw: T, vectors: HashMap<String, String>) -> Self {
88        Self {
89            id: id.into(),
90            raw,
91            vectors,
92            payload_fields: HashMap::new(),
93        }
94    }
95
96    pub fn with_payload_fields(
97        mut self,
98        payload_fields: HashMap<String, serde_json::Value>,
99    ) -> Self {
100        self.payload_fields = payload_fields;
101        self
102    }
103}
104
105impl<T> NamedVectorPayloadDocument<T>
106where
107    T: Serialize,
108{
109    pub fn with_mirrored_payload_fields(
110        mut self,
111        fields: impl IntoIterator<Item = impl AsRef<str>>,
112    ) -> Result<Self, serde_json::Error> {
113        self.payload_fields = mirrored_payload_fields_for(&self.raw, fields)?;
114        Ok(self)
115    }
116}
117
118pub fn mirrored_payload_fields(
119    raw: &serde_json::Value,
120    fields: impl IntoIterator<Item = impl AsRef<str>>,
121) -> HashMap<String, serde_json::Value> {
122    let Some(raw_object) = raw.as_object() else {
123        return HashMap::new();
124    };
125
126    let mut mirrored = HashMap::new();
127    let mut seen = HashSet::new();
128    for field in fields {
129        let field = field.as_ref();
130        if !seen.insert(field.to_string()) {
131            continue;
132        }
133
134        if field == "raw" || field == "source_id" {
135            continue;
136        }
137
138        if let Some(value) = raw_object.get(field) {
139            mirrored.insert(field.to_string(), value.clone());
140        }
141    }
142
143    mirrored
144}
145
146pub fn mirrored_payload_fields_for<T>(
147    raw: &T,
148    fields: impl IntoIterator<Item = impl AsRef<str>>,
149) -> Result<HashMap<String, serde_json::Value>, serde_json::Error>
150where
151    T: Serialize,
152{
153    let raw = serde_json::to_value(raw)?;
154    Ok(mirrored_payload_fields(&raw, fields))
155}
156
157pub async fn embed_documents_with_payload_fields<T, I, S>(
158    provider: &SharedEmbeddingProvider,
159    documents: Vec<(String, T)>,
160    payload_fields: I,
161) -> Result<Vec<PreparedPayloadDocument>, VectorStoreError>
162where
163    T: Embed + Serialize + Send + Sync + Clone,
164    I: IntoIterator<Item = S>,
165    S: AsRef<str>,
166{
167    let mut all_texts = Vec::new();
168    let mut ranges = Vec::new();
169    let mut raws = Vec::new();
170    let mut ids = Vec::new();
171    let payload_field_names = payload_fields
172        .into_iter()
173        .map(|field| field.as_ref().to_string())
174        .collect::<Vec<_>>();
175    let mut mirrored_payloads = Vec::new();
176
177    for (id, doc) in documents.iter() {
178        let mut embedder = TextEmbedder::default();
179        doc.embed(&mut embedder).map_err(|err| {
180            VectorStoreError::EmbeddingError(EmbeddingError::EmbedFailure(err.to_string()))
181        })?;
182
183        if embedder.is_empty() {
184            return Err(VectorStoreError::EmbeddingError(EmbeddingError::Empty));
185        }
186
187        let start = all_texts.len();
188        let count = embedder.len();
189        all_texts.extend(embedder.into_parts());
190        ranges.push((start, count));
191        let raw = serde_json::to_value(doc)?;
192        mirrored_payloads.push(mirrored_payload_fields(&raw, &payload_field_names));
193        raws.push(raw);
194        ids.push(id.clone());
195    }
196
197    let vectors = provider
198        .embed(all_texts.clone())
199        .await
200        .map_err(EmbeddingError::Provider)?;
201
202    let mut prepared = Vec::with_capacity(ids.len());
203    let mut vectors_iter = vectors.into_iter();
204    let mut expected_start = 0usize;
205    for (((id, raw), payload_fields), (start, count)) in
206        ids.into_iter().zip(raws).zip(mirrored_payloads).zip(ranges)
207    {
208        if start != expected_start {
209            return Err(VectorStoreError::EmbeddingError(
210                EmbeddingError::EmbedFailure("embedding ranges are inconsistent".into()),
211            ));
212        }
213
214        let mut embeddings = Vec::with_capacity(count);
215        for offset in 0..count {
216            let Some(vector) = vectors_iter.next() else {
217                return Err(VectorStoreError::EmbeddingError(
218                    EmbeddingError::EmbedFailure(
219                        "embedding provider returned fewer vectors than expected".into(),
220                    ),
221                ));
222            };
223
224            embeddings.push(Embedding {
225                document: all_texts[start + offset].clone(),
226                vec: vector.into(),
227            });
228        }
229        expected_start += count;
230
231        prepared.push(PreparedPayloadDocument {
232            id,
233            raw,
234            payload_fields,
235            embeddings: OneOrMany::from(embeddings),
236        });
237    }
238
239    Ok(prepared)
240}
241
242pub async fn embed_payload_documents<T>(
243    provider: &SharedEmbeddingProvider,
244    documents: Vec<PayloadDocument<T>>,
245) -> Result<Vec<PreparedPayloadDocument>, VectorStoreError>
246where
247    T: Embed + Serialize + Send + Sync + Clone,
248{
249    let mut all_texts = Vec::new();
250    let mut ranges = Vec::new();
251    let mut raws = Vec::new();
252    let mut ids = Vec::new();
253    let mut mirrored_payloads = Vec::new();
254
255    for doc in documents.iter() {
256        let mut embedder = TextEmbedder::default();
257        doc.raw.embed(&mut embedder).map_err(|err| {
258            VectorStoreError::EmbeddingError(EmbeddingError::EmbedFailure(err.to_string()))
259        })?;
260
261        if embedder.is_empty() {
262            return Err(VectorStoreError::EmbeddingError(EmbeddingError::Empty));
263        }
264
265        let start = all_texts.len();
266        let count = embedder.len();
267        all_texts.extend(embedder.into_parts());
268        ranges.push((start, count));
269        raws.push(serde_json::to_value(&doc.raw)?);
270        mirrored_payloads.push(doc.payload_fields.clone());
271        ids.push(doc.id.clone());
272    }
273
274    let vectors = provider
275        .embed(all_texts.clone())
276        .await
277        .map_err(EmbeddingError::Provider)?;
278
279    let mut prepared = Vec::with_capacity(ids.len());
280    let mut vectors_iter = vectors.into_iter();
281    let mut expected_start = 0usize;
282    for (((id, raw), payload_fields), (start, count)) in
283        ids.into_iter().zip(raws).zip(mirrored_payloads).zip(ranges)
284    {
285        if start != expected_start {
286            return Err(VectorStoreError::EmbeddingError(
287                EmbeddingError::EmbedFailure("embedding ranges are inconsistent".into()),
288            ));
289        }
290
291        let mut embeddings = Vec::with_capacity(count);
292        for offset in 0..count {
293            let Some(vector) = vectors_iter.next() else {
294                return Err(VectorStoreError::EmbeddingError(
295                    EmbeddingError::EmbedFailure(
296                        "embedding provider returned fewer vectors than expected".into(),
297                    ),
298                ));
299            };
300
301            embeddings.push(Embedding {
302                document: all_texts[start + offset].clone(),
303                vec: vector.into(),
304            });
305        }
306        expected_start += count;
307
308        prepared.push(PreparedPayloadDocument {
309            id,
310            raw,
311            payload_fields,
312            embeddings: OneOrMany::from(embeddings),
313        });
314    }
315
316    Ok(prepared)
317}
318
319pub async fn embed_named_payload_documents<T>(
320    provider: &SharedEmbeddingProvider,
321    documents: Vec<NamedVectorPayloadDocument<T>>,
322) -> Result<Vec<PreparedNamedVectorPayloadDocument>, VectorStoreError>
323where
324    T: Serialize + Send + Sync + Clone,
325{
326    let mut all_texts = Vec::new();
327    let mut ranges = Vec::new();
328    let mut raws = Vec::new();
329    let mut ids = Vec::new();
330    let mut names_by_doc = Vec::new();
331    let mut mirrored_payloads = Vec::new();
332
333    for doc in documents {
334        if doc.vectors.is_empty() {
335            return Err(VectorStoreError::EmbeddingError(EmbeddingError::Empty));
336        }
337
338        let mut names = Vec::with_capacity(doc.vectors.len());
339        let start = all_texts.len();
340
341        for (name, text) in doc.vectors {
342            names.push(name);
343            all_texts.push(text);
344        }
345
346        ranges.push((start, names.len()));
347        names_by_doc.push(names);
348        let raw = serde_json::to_value(doc.raw)?;
349        mirrored_payloads.push(doc.payload_fields);
350        raws.push(raw);
351        ids.push(doc.id);
352    }
353
354    let vectors = provider
355        .embed(all_texts.clone())
356        .await
357        .map_err(EmbeddingError::Provider)?;
358
359    let mut prepared = Vec::with_capacity(ids.len());
360    let mut vectors_iter = vectors.into_iter();
361    let mut expected_start = 0usize;
362    for ((((id, raw), payload_fields), (start, count)), names) in ids
363        .into_iter()
364        .zip(raws)
365        .zip(mirrored_payloads)
366        .zip(ranges)
367        .zip(names_by_doc)
368    {
369        if start != expected_start {
370            return Err(VectorStoreError::EmbeddingError(
371                EmbeddingError::EmbedFailure("embedding ranges are inconsistent".into()),
372            ));
373        }
374
375        let mut mapped = HashMap::with_capacity(count);
376        for name in names.into_iter() {
377            let Some(vector) = vectors_iter.next() else {
378                return Err(VectorStoreError::EmbeddingError(
379                    EmbeddingError::EmbedFailure(
380                        "embedding provider returned fewer vectors than expected".into(),
381                    ),
382                ));
383            };
384            mapped.insert(name, vector);
385        }
386        expected_start += count;
387
388        prepared.push(PreparedNamedVectorPayloadDocument {
389            id,
390            raw,
391            payload_fields,
392            vectors: mapped,
393        });
394    }
395
396    Ok(prepared)
397}
398
399#[cfg(test)]
400mod tests {
401    use super::*;
402    use crate::embeddings::{EmbedError, TextEmbedder};
403    use autoagents_llm::embedding::EmbeddingProvider;
404    use autoagents_llm::error::LLMError;
405    use std::sync::Arc;
406
407    #[derive(Debug, Clone, Serialize)]
408    struct IndexedDoc {
409        workspace_id: &'static str,
410        title: &'static str,
411        body: &'static str,
412    }
413
414    impl crate::embeddings::Embed for IndexedDoc {
415        fn embed(&self, embedder: &mut TextEmbedder) -> Result<(), EmbedError> {
416            embedder.embed(self.title);
417            embedder.embed(self.body);
418            Ok(())
419        }
420    }
421
422    #[derive(Debug, Clone)]
423    struct DummyEmbeddingProvider {
424        vectors: Vec<Vec<f32>>,
425    }
426
427    #[async_trait::async_trait]
428    impl EmbeddingProvider for DummyEmbeddingProvider {
429        async fn embed(&self, _text: Vec<String>) -> Result<Vec<Vec<f32>>, LLMError> {
430            Ok(self.vectors.clone())
431        }
432    }
433
434    #[test]
435    fn test_mirrored_payload_fields_extracts_selected_root_keys() {
436        let raw = serde_json::json!({
437            "workspace_id": "ws-1",
438            "project_id": "proj-1",
439            "file_path": "src/lib.rs",
440            "body": "very large text"
441        });
442
443        let mirrored = mirrored_payload_fields(&raw, ["workspace_id", "file_path", "missing"]);
444        assert_eq!(mirrored.len(), 2);
445        assert_eq!(
446            mirrored.get("workspace_id"),
447            Some(&serde_json::json!("ws-1"))
448        );
449        assert_eq!(
450            mirrored.get("file_path"),
451            Some(&serde_json::json!("src/lib.rs"))
452        );
453    }
454
455    #[test]
456    fn test_mirrored_payload_fields_for_serializable_value() {
457        #[derive(Serialize)]
458        struct IndexedDoc {
459            workspace_id: &'static str,
460            project_id: &'static str,
461            body: &'static str,
462        }
463
464        let mirrored = mirrored_payload_fields_for(
465            &IndexedDoc {
466                workspace_id: "ws-1",
467                project_id: "proj-1",
468                body: "large text",
469            },
470            ["workspace_id", "project_id"],
471        )
472        .unwrap();
473
474        assert_eq!(
475            mirrored.get("workspace_id"),
476            Some(&serde_json::json!("ws-1"))
477        );
478        assert_eq!(
479            mirrored.get("project_id"),
480            Some(&serde_json::json!("proj-1"))
481        );
482        assert!(!mirrored.contains_key("body"));
483    }
484
485    #[test]
486    fn test_payload_document_builders_cover_manual_and_mirrored_fields() {
487        let doc = PayloadDocument::new(
488            "doc-1",
489            IndexedDoc {
490                workspace_id: "ws-1",
491                title: "Title",
492                body: "Body",
493            },
494        )
495        .with_payload_fields(HashMap::from([(
496            "workspace_id".to_string(),
497            serde_json::json!("manual"),
498        )]));
499        assert_eq!(doc.id, "doc-1");
500        assert_eq!(
501            doc.payload_fields["workspace_id"],
502            serde_json::json!("manual")
503        );
504
505        let mirrored = PayloadDocument::new(
506            "doc-2",
507            IndexedDoc {
508                workspace_id: "ws-2",
509                title: "Second",
510                body: "Document",
511            },
512        )
513        .with_mirrored_payload_fields(["workspace_id", "raw", "source_id", "workspace_id"])
514        .expect("mirrored payload should build");
515        assert_eq!(mirrored.payload_fields.len(), 1);
516        assert_eq!(
517            mirrored.payload_fields["workspace_id"],
518            serde_json::json!("ws-2")
519        );
520    }
521
522    #[test]
523    fn test_named_vector_payload_document_builders_cover_manual_and_mirrored_fields() {
524        let base = NamedVectorDocument {
525            id: "doc-1".to_string(),
526            raw: IndexedDoc {
527                workspace_id: "ws-1",
528                title: "Title",
529                body: "Body",
530            },
531            vectors: HashMap::from([
532                ("title".to_string(), "Title".to_string()),
533                ("body".to_string(), "Body".to_string()),
534            ]),
535        };
536        let payload_doc = base.clone().with_payload_fields(HashMap::from([(
537            "workspace_id".to_string(),
538            serde_json::json!("ws-1"),
539        )]));
540        assert_eq!(payload_doc.vectors.len(), 2);
541        assert_eq!(
542            payload_doc.payload_fields["workspace_id"],
543            serde_json::json!("ws-1")
544        );
545
546        let mirrored = NamedVectorPayloadDocument::new("doc-2", base.raw, base.vectors)
547            .with_mirrored_payload_fields(["workspace_id"])
548            .expect("mirrored payload should build");
549        assert_eq!(
550            mirrored.payload_fields["workspace_id"],
551            serde_json::json!("ws-1")
552        );
553    }
554
555    #[tokio::test]
556    async fn test_embed_documents_with_payload_fields_success_and_short_vector_error() {
557        let provider: SharedEmbeddingProvider = Arc::new(DummyEmbeddingProvider {
558            vectors: vec![vec![0.1_f32], vec![0.2_f32]],
559        });
560        let docs = vec![(
561            "doc-1".to_string(),
562            IndexedDoc {
563                workspace_id: "ws-1",
564                title: "Title",
565                body: "Body",
566            },
567        )];
568
569        let prepared = embed_documents_with_payload_fields(&provider, docs, ["workspace_id"])
570            .await
571            .expect("documents should embed");
572        assert_eq!(prepared.len(), 1);
573        assert_eq!(prepared[0].id, "doc-1");
574        assert_eq!(
575            prepared[0].payload_fields["workspace_id"],
576            serde_json::json!("ws-1")
577        );
578        assert_eq!(prepared[0].embeddings.len(), 2);
579
580        let short_provider: SharedEmbeddingProvider = Arc::new(DummyEmbeddingProvider {
581            vectors: vec![vec![0.1_f32]],
582        });
583        let err = embed_documents_with_payload_fields(
584            &short_provider,
585            vec![(
586                "doc-2".to_string(),
587                IndexedDoc {
588                    workspace_id: "ws-2",
589                    title: "Another",
590                    body: "Entry",
591                },
592            )],
593            ["workspace_id"],
594        )
595        .await
596        .expect_err("short embedding response should fail");
597        assert!(err.to_string().contains("fewer vectors"));
598    }
599
600    #[tokio::test]
601    async fn test_embed_payload_documents_and_named_payload_documents_cover_success_and_empty() {
602        let provider: SharedEmbeddingProvider = Arc::new(DummyEmbeddingProvider {
603            vectors: vec![vec![0.1_f32], vec![0.2_f32], vec![0.3_f32], vec![0.4_f32]],
604        });
605
606        let prepared = embed_payload_documents(
607            &provider,
608            vec![
609                PayloadDocument::new(
610                    "doc-1",
611                    IndexedDoc {
612                        workspace_id: "ws-1",
613                        title: "Title",
614                        body: "Body",
615                    },
616                )
617                .with_payload_fields(HashMap::from([(
618                    "workspace_id".to_string(),
619                    serde_json::json!("ws-1"),
620                )])),
621            ],
622        )
623        .await
624        .expect("payload documents should embed");
625        assert_eq!(prepared.len(), 1);
626        assert_eq!(
627            prepared[0].payload_fields["workspace_id"],
628            serde_json::json!("ws-1")
629        );
630
631        let named = embed_named_payload_documents(
632            &provider,
633            vec![
634                NamedVectorPayloadDocument::new(
635                    "doc-2",
636                    IndexedDoc {
637                        workspace_id: "ws-2",
638                        title: "Named",
639                        body: "Vector",
640                    },
641                    HashMap::from([
642                        ("title".to_string(), "Named".to_string()),
643                        ("body".to_string(), "Vector".to_string()),
644                    ]),
645                )
646                .with_payload_fields(HashMap::from([(
647                    "workspace_id".to_string(),
648                    serde_json::json!("ws-2"),
649                )])),
650            ],
651        )
652        .await
653        .expect("named payload documents should embed");
654        assert_eq!(named.len(), 1);
655        assert_eq!(named[0].vectors.len(), 2);
656        assert_eq!(
657            named[0].payload_fields["workspace_id"],
658            serde_json::json!("ws-2")
659        );
660
661        let err = embed_named_payload_documents(
662            &provider,
663            vec![NamedVectorPayloadDocument::new(
664                "doc-3",
665                IndexedDoc {
666                    workspace_id: "ws-3",
667                    title: "Empty",
668                    body: "Vectors",
669                },
670                HashMap::new(),
671            )],
672        )
673        .await
674        .expect_err("empty named vectors should fail");
675        assert!(err.to_string().contains("No content to embed"));
676    }
677}