autoagents_core/vector_store/
mod.rs1pub 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}