rig_core/vector_store/
mod.rs1use http::StatusCode;
13pub use request::VectorSearchRequest;
14use serde::{Deserialize, Serialize, de::DeserializeOwned};
15use serde_json::{Value, json};
16
17use crate::error::ProviderError;
18use crate::{
19 Embed,
20 embeddings::Embedding,
21 tool::PortableTool,
22 vector_store::request::{FilterError, SearchFilter},
23 wasm_compat::{WasmCompatSend, WasmCompatSync},
24};
25
26pub mod builder;
27pub mod in_memory_store;
28pub mod lsh;
29pub mod request;
30
31#[derive(Debug, thiserror::Error)]
33pub enum VectorStoreError {
34 #[error("Embedding error: {0}")]
36 EmbeddingError(#[from] ProviderError),
37
38 #[error("Json error: {0}")]
40 JsonError(#[from] serde_json::Error),
41
42 #[cfg(not(target_family = "wasm"))]
43 #[error("Datastore error: {0}")]
45 DatastoreError(#[from] Box<dyn std::error::Error + Send + Sync + 'static>),
46
47 #[error("Filter error: {0}")]
49 FilterError(#[from] FilterError),
50
51 #[cfg(target_family = "wasm")]
52 #[error("Datastore error: {0}")]
54 DatastoreError(#[from] Box<dyn std::error::Error + 'static>),
55
56 #[error("Missing Id: {0}")]
58 MissingIdError(String),
59
60 #[error("HTTP request error: {0}")]
67 Http(#[from] crate::http_client::Error),
68
69 #[error("External call to API returned an error. Error code: {0} Message: {1}")]
71 ExternalAPIError(StatusCode, String),
72
73 #[error("Requested {requested} samples, but this vector store returns at most {max}")]
75 SamplesOutOfRange {
76 requested: u64,
78 max: u64,
80 },
81}
82
83impl VectorStoreError {
84 #[cfg(not(target_family = "wasm"))]
89 pub fn datastore(e: impl std::error::Error + Send + Sync + 'static) -> Self {
90 Self::DatastoreError(Box::new(e))
91 }
92
93 #[cfg(target_family = "wasm")]
95 pub fn datastore(e: impl std::error::Error + 'static) -> Self {
96 Self::DatastoreError(Box::new(e))
97 }
98}
99
100pub fn flatten_embedded<Doc: Serialize, R>(
107 documents: Vec<(Doc, Vec<Embedding>)>,
108 mut f: impl FnMut(&Value, Embedding) -> Result<R, VectorStoreError>,
109) -> Result<Vec<R>, VectorStoreError> {
110 let mut records = Vec::new();
111 for (document, embeddings) in documents {
112 let json_document = serde_json::to_value(&document)?;
113 for embedding in embeddings {
114 records.push(f(&json_document, embedding)?);
115 }
116 }
117 Ok(records)
118}
119
120pub trait InsertDocuments: WasmCompatSend + WasmCompatSync {
122 fn insert_documents<Doc: Serialize + Embed + WasmCompatSend>(
127 &self,
128 documents: Vec<(Doc, Vec<Embedding>)>,
129 ) -> impl std::future::Future<Output = Result<(), VectorStoreError>> + WasmCompatSend;
130}
131
132pub trait VectorStoreIndex: WasmCompatSend + WasmCompatSync {
134 type Filter: SearchFilter + WasmCompatSend + WasmCompatSync;
136
137 fn top_n<T: DeserializeOwned + WasmCompatSend>(
139 &self,
140 req: VectorSearchRequest<Self::Filter>,
141 ) -> impl std::future::Future<Output = Result<Vec<VectorSearchResult<T>>, VectorStoreError>>
142 + WasmCompatSend;
143
144 fn top_n_ids(
146 &self,
147 req: VectorSearchRequest<Self::Filter>,
148 ) -> impl std::future::Future<Output = Result<Vec<VectorSearchIdResult>, VectorStoreError>>
149 + WasmCompatSend;
150}
151
152#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
154pub struct VectorSearchResult<T> {
155 pub score: f64,
157 pub id: String,
159 pub document: T,
161}
162
163#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
165pub struct VectorSearchIdResult {
166 pub score: f64,
168 pub id: String,
170}
171
172impl<T, F> PortableTool for T
173where
174 F: SearchFilter<Value = serde_json::Value>
175 + WasmCompatSend
176 + WasmCompatSync
177 + serde::de::DeserializeOwned,
178 T: VectorStoreIndex<Filter = F>,
179{
180 const NAME: &'static str = "search_vector_store";
181 type Error = VectorStoreError;
182 type Args = VectorSearchRequest<F>;
183 type Output = Vec<VectorSearchResult<Value>>;
184
185 fn description(&self) -> String {
186 "Retrieves the most relevant documents from a vector store based on a query.".to_string()
187 }
188
189 fn parameters(&self) -> serde_json::Value {
190 json!({
191 "type": "object",
192 "properties": {
193 "query": {
194 "type": "string",
195 "description": "The query string to search for relevant documents in the vector store."
196 },
197 "samples": {
198 "type": "integer",
199 "description": "The maximum number of samples / documents to retrieve.",
200 "default": 5,
201 "minimum": 1
202 },
203 "threshold": {
204 "type": "number",
205 "description": "Similarity search threshold. If present, any result with a distance less than this may be omitted from the final result."
206 }
207 },
208 "required": ["query", "samples"]
209 })
210 }
211
212 async fn call(&self, args: Self::Args) -> Result<Self::Output, Self::Error> {
213 self.top_n(args).await
214 }
215}
216
217#[derive(Clone, Debug, Default)]
219pub enum IndexStrategy {
220 #[default]
222 BruteForce,
223
224 LSH {
226 num_tables: usize,
228 num_hyperplanes: usize,
230 },
231}
232
233#[cfg(test)]
234mod tests;