Skip to main content

autoagents_core/vector_store/
mod.rs

1pub use payload::{
2    NamedVectorPayloadDocument, PayloadDocument, PreparedNamedVectorPayloadDocument,
3    PreparedPayloadDocument, embed_documents_with_payload_fields, embed_named_payload_documents,
4    embed_payload_documents, mirrored_payload_fields, mirrored_payload_fields_for,
5};
6pub use request::VectorSearchRequest;
7
8use async_trait::async_trait;
9use serde::{Deserialize, Serialize};
10use std::collections::HashMap;
11use uuid::Uuid;
12
13use crate::document::Document;
14use crate::embeddings::{Embed, Embedding, EmbeddingError, SharedEmbeddingProvider};
15use crate::one_or_many::OneOrMany;
16use crate::vector_store::request::{FilterError, SearchFilter};
17
18pub mod in_memory_store;
19pub mod payload;
20pub mod request;
21
22pub const DEFAULT_VECTOR_NAME: &str = "default";
23
24#[derive(Debug, thiserror::Error)]
25pub enum VectorStoreError {
26    #[error("Embedding error: {0}")]
27    EmbeddingError(#[from] EmbeddingError),
28
29    #[error("Json error: {0}")]
30    JsonError(#[from] serde_json::Error),
31
32    #[error("Filter error: {0}")]
33    FilterError(#[from] FilterError),
34
35    #[error("Datastore error: {0}")]
36    DatastoreError(#[from] Box<dyn std::error::Error + Send + Sync + 'static>),
37
38    #[error("Error while building VectorSearchRequest: {0}")]
39    BuilderError(String),
40}
41
42#[async_trait]
43pub trait VectorStoreIndex: Send + Sync {
44    type Filter: SearchFilter + Send + Sync;
45
46    async fn insert_documents<T>(&self, documents: Vec<T>) -> Result<(), VectorStoreError>
47    where
48        T: Embed + Serialize + Send + Sync + Clone;
49
50    async fn insert_documents_with_ids<T>(
51        &self,
52        documents: Vec<(String, T)>,
53    ) -> Result<(), VectorStoreError>
54    where
55        T: Embed + Serialize + Send + Sync + Clone;
56
57    async fn top_n<T>(
58        &self,
59        req: VectorSearchRequest<Self::Filter>,
60    ) -> Result<Vec<(f64, String, T)>, VectorStoreError>
61    where
62        T: for<'de> Deserialize<'de> + Send + Sync;
63
64    async fn top_n_ids(
65        &self,
66        req: VectorSearchRequest<Self::Filter>,
67    ) -> Result<Vec<(f64, String)>, VectorStoreError>;
68
69    async fn insert_documents_with_named_vectors<T>(
70        &self,
71        documents: Vec<NamedVectorDocument<T>>,
72    ) -> Result<(), VectorStoreError>
73    where
74        T: Serialize + Send + Sync + Clone;
75}
76
77#[derive(Debug, Clone, Serialize, Deserialize)]
78pub struct VectorStoreOutput {
79    pub score: f64,
80    pub id: String,
81    pub document: Document,
82}
83
84#[derive(Debug, Clone)]
85pub struct PreparedDocument {
86    pub id: String,
87    pub raw: serde_json::Value,
88    pub embeddings: OneOrMany<Embedding>,
89}
90
91#[derive(Debug, Clone)]
92pub struct NamedVectorDocument<T> {
93    pub id: String,
94    pub raw: T,
95    pub vectors: HashMap<String, String>,
96}
97
98#[derive(Debug, Clone)]
99pub struct PreparedNamedVectorDocument {
100    pub id: String,
101    pub raw: serde_json::Value,
102    pub vectors: HashMap<String, Vec<f32>>,
103}
104
105pub async fn embed_documents<T>(
106    provider: &SharedEmbeddingProvider,
107    documents: Vec<(String, T)>,
108) -> Result<Vec<PreparedDocument>, VectorStoreError>
109where
110    T: Embed + Serialize + Send + Sync + Clone,
111{
112    let prepared =
113        embed_documents_with_payload_fields(provider, documents, std::iter::empty::<&str>())
114            .await?;
115    Ok(prepared
116        .into_iter()
117        .map(|doc| PreparedDocument {
118            id: doc.id,
119            raw: doc.raw,
120            embeddings: doc.embeddings,
121        })
122        .collect())
123}
124
125pub async fn embed_named_documents<T>(
126    provider: &SharedEmbeddingProvider,
127    documents: Vec<NamedVectorDocument<T>>,
128) -> Result<Vec<PreparedNamedVectorDocument>, VectorStoreError>
129where
130    T: Serialize + Send + Sync + Clone,
131{
132    let documents = documents
133        .into_iter()
134        .map(|doc| NamedVectorPayloadDocument {
135            id: doc.id,
136            raw: doc.raw,
137            vectors: doc.vectors,
138            payload_fields: HashMap::new(),
139        })
140        .collect();
141
142    let prepared = embed_named_payload_documents(provider, documents).await?;
143    Ok(prepared
144        .into_iter()
145        .map(|doc| PreparedNamedVectorDocument {
146            id: doc.id,
147            raw: doc.raw,
148            vectors: doc.vectors,
149        })
150        .collect())
151}
152
153pub fn normalize_id(id: Option<String>) -> String {
154    id.unwrap_or_else(|| Uuid::new_v4().to_string())
155}
156
157#[cfg(test)]
158mod tests {
159    use super::*;
160    use crate::document::Document;
161    use crate::embeddings::{Embed, EmbedError, TextEmbedder};
162    use autoagents_llm::embedding::EmbeddingProvider;
163    use autoagents_llm::error::LLMError;
164    use serde::Serialize;
165    use std::sync::Arc;
166
167    #[derive(Debug, Clone)]
168    struct DummyEmbeddingProvider {
169        vectors: Vec<Vec<f32>>,
170    }
171
172    #[async_trait::async_trait]
173    impl EmbeddingProvider for DummyEmbeddingProvider {
174        async fn embed(&self, _text: Vec<String>) -> Result<Vec<Vec<f32>>, LLMError> {
175            Ok(self.vectors.clone())
176        }
177    }
178
179    #[derive(Debug, Clone, Serialize)]
180    struct MultiPartDoc {
181        parts: Vec<String>,
182    }
183
184    impl Embed for MultiPartDoc {
185        fn embed(&self, embedder: &mut TextEmbedder) -> Result<(), EmbedError> {
186            for part in &self.parts {
187                embedder.embed(part.clone());
188            }
189            Ok(())
190        }
191    }
192
193    #[derive(Debug, Clone, Serialize)]
194    struct EmptyDoc;
195
196    impl Embed for EmptyDoc {
197        fn embed(&self, _embedder: &mut TextEmbedder) -> Result<(), EmbedError> {
198            Ok(())
199        }
200    }
201
202    #[test]
203    fn test_normalize_id_none_generates_uuid() {
204        let id = normalize_id(None);
205        assert!(!id.is_empty());
206        assert!(uuid::Uuid::parse_str(&id).is_ok());
207    }
208
209    #[test]
210    fn test_normalize_id_some_returns_value() {
211        let id = normalize_id(Some("custom-id".to_string()));
212        assert_eq!(id, "custom-id");
213    }
214
215    #[tokio::test]
216    async fn test_embed_documents_with_mock() {
217        use crate::tests::MockLLMProvider;
218        let provider: SharedEmbeddingProvider = Arc::new(MockLLMProvider {});
219        let docs = vec![("id1".to_string(), Document::new("hello"))];
220        let result = embed_documents(&provider, docs).await;
221        assert!(result.is_ok());
222        let prepared = result.unwrap();
223        assert_eq!(prepared.len(), 1);
224        assert_eq!(prepared[0].id, "id1");
225    }
226
227    #[tokio::test]
228    async fn test_embed_documents_empty_embedder() {
229        let provider: SharedEmbeddingProvider =
230            Arc::new(DummyEmbeddingProvider { vectors: vec![] });
231        let docs = vec![("id1".to_string(), EmptyDoc)];
232        let err = embed_documents(&provider, docs).await.unwrap_err();
233        assert!(err.to_string().contains("No content to embed"));
234    }
235
236    #[tokio::test]
237    async fn test_embed_documents_fewer_vectors_than_expected() {
238        let provider: SharedEmbeddingProvider = Arc::new(DummyEmbeddingProvider {
239            vectors: vec![vec![0.1_f32]],
240        });
241        let docs = vec![(
242            "id1".to_string(),
243            MultiPartDoc {
244                parts: vec!["a".to_string(), "b".to_string()],
245            },
246        )];
247        let err = embed_documents(&provider, docs).await.unwrap_err();
248        assert!(err.to_string().contains("fewer vectors"));
249    }
250
251    #[tokio::test]
252    async fn test_embed_named_documents_success() {
253        let provider: SharedEmbeddingProvider = Arc::new(DummyEmbeddingProvider {
254            vectors: vec![vec![0.1_f32], vec![0.2_f32]],
255        });
256        let docs = vec![NamedVectorDocument {
257            id: "doc-1".to_string(),
258            raw: "raw".to_string(),
259            vectors: HashMap::from([
260                ("title".to_string(), "hello".to_string()),
261                ("body".to_string(), "world".to_string()),
262            ]),
263        }];
264        let prepared = embed_named_documents(&provider, docs).await.unwrap();
265        assert_eq!(prepared.len(), 1);
266        assert_eq!(prepared[0].vectors.len(), 2);
267    }
268
269    #[tokio::test]
270    async fn test_embed_named_documents_empty_vectors() {
271        let provider: SharedEmbeddingProvider =
272            Arc::new(DummyEmbeddingProvider { vectors: vec![] });
273        let docs = vec![NamedVectorDocument {
274            id: "doc-1".to_string(),
275            raw: "raw".to_string(),
276            vectors: HashMap::new(),
277        }];
278        let err = embed_named_documents(&provider, docs).await.unwrap_err();
279        assert!(err.to_string().contains("No content to embed"));
280    }
281
282    #[tokio::test]
283    async fn test_embed_named_documents_fewer_vectors() {
284        let provider: SharedEmbeddingProvider = Arc::new(DummyEmbeddingProvider {
285            vectors: vec![vec![0.1_f32]],
286        });
287        let docs = vec![NamedVectorDocument {
288            id: "doc-1".to_string(),
289            raw: "raw".to_string(),
290            vectors: HashMap::from([
291                ("title".to_string(), "hello".to_string()),
292                ("body".to_string(), "world".to_string()),
293            ]),
294        }];
295        let err = embed_named_documents(&provider, docs).await.unwrap_err();
296        assert!(err.to_string().contains("fewer vectors"));
297    }
298}