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}