use http::StatusCode;
pub use request::VectorSearchRequest;
use serde::{Deserialize, Serialize, de::DeserializeOwned};
use serde_json::{Value, json};
use crate::error::ProviderError;
use crate::{
Embed,
embeddings::Embedding,
tool::PortableTool,
vector_store::request::{FilterError, SearchFilter},
wasm_compat::{WasmCompatSend, WasmCompatSync},
};
pub mod builder;
pub mod in_memory_store;
pub mod lsh;
pub mod request;
#[derive(Debug, thiserror::Error)]
pub enum VectorStoreError {
#[error("Embedding error: {0}")]
EmbeddingError(#[from] ProviderError),
#[error("Json error: {0}")]
JsonError(#[from] serde_json::Error),
#[cfg(not(target_family = "wasm"))]
#[error("Datastore error: {0}")]
DatastoreError(#[from] Box<dyn std::error::Error + Send + Sync + 'static>),
#[error("Filter error: {0}")]
FilterError(#[from] FilterError),
#[cfg(target_family = "wasm")]
#[error("Datastore error: {0}")]
DatastoreError(#[from] Box<dyn std::error::Error + 'static>),
#[error("Missing Id: {0}")]
MissingIdError(String),
#[error("HTTP request error: {0}")]
Http(#[from] crate::http_client::Error),
#[error("External call to API returned an error. Error code: {0} Message: {1}")]
ExternalAPIError(StatusCode, String),
#[error("Requested {requested} samples, but this vector store returns at most {max}")]
SamplesOutOfRange {
requested: u64,
max: u64,
},
}
impl VectorStoreError {
#[cfg(not(target_family = "wasm"))]
pub fn datastore(e: impl std::error::Error + Send + Sync + 'static) -> Self {
Self::DatastoreError(Box::new(e))
}
#[cfg(target_family = "wasm")]
pub fn datastore(e: impl std::error::Error + 'static) -> Self {
Self::DatastoreError(Box::new(e))
}
}
pub fn flatten_embedded<Doc: Serialize, R>(
documents: Vec<(Doc, Vec<Embedding>)>,
mut f: impl FnMut(&Value, Embedding) -> Result<R, VectorStoreError>,
) -> Result<Vec<R>, VectorStoreError> {
let mut records = Vec::new();
for (document, embeddings) in documents {
let json_document = serde_json::to_value(&document)?;
for embedding in embeddings {
records.push(f(&json_document, embedding)?);
}
}
Ok(records)
}
pub trait InsertDocuments: WasmCompatSend + WasmCompatSync {
fn insert_documents<Doc: Serialize + Embed + WasmCompatSend>(
&self,
documents: Vec<(Doc, Vec<Embedding>)>,
) -> impl std::future::Future<Output = Result<(), VectorStoreError>> + WasmCompatSend;
}
pub trait VectorStoreIndex: WasmCompatSend + WasmCompatSync {
type Filter: SearchFilter + WasmCompatSend + WasmCompatSync;
fn top_n<T: DeserializeOwned + WasmCompatSend>(
&self,
req: VectorSearchRequest<Self::Filter>,
) -> impl std::future::Future<Output = Result<Vec<VectorSearchResult<T>>, VectorStoreError>>
+ WasmCompatSend;
fn top_n_ids(
&self,
req: VectorSearchRequest<Self::Filter>,
) -> impl std::future::Future<Output = Result<Vec<VectorSearchIdResult>, VectorStoreError>>
+ WasmCompatSend;
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct VectorSearchResult<T> {
pub score: f64,
pub id: String,
pub document: T,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct VectorSearchIdResult {
pub score: f64,
pub id: String,
}
impl<T, F> PortableTool for T
where
F: SearchFilter<Value = serde_json::Value>
+ WasmCompatSend
+ WasmCompatSync
+ serde::de::DeserializeOwned,
T: VectorStoreIndex<Filter = F>,
{
const NAME: &'static str = "search_vector_store";
type Error = VectorStoreError;
type Args = VectorSearchRequest<F>;
type Output = Vec<VectorSearchResult<Value>>;
fn description(&self) -> String {
"Retrieves the most relevant documents from a vector store based on a query.".to_string()
}
fn parameters(&self) -> serde_json::Value {
json!({
"type": "object",
"properties": {
"query": {
"type": "string",
"description": "The query string to search for relevant documents in the vector store."
},
"samples": {
"type": "integer",
"description": "The maximum number of samples / documents to retrieve.",
"default": 5,
"minimum": 1
},
"threshold": {
"type": "number",
"description": "Similarity search threshold. If present, any result with a distance less than this may be omitted from the final result."
}
},
"required": ["query", "samples"]
})
}
async fn call(&self, args: Self::Args) -> Result<Self::Output, Self::Error> {
self.top_n(args).await
}
}
#[derive(Clone, Debug, Default)]
pub enum IndexStrategy {
#[default]
BruteForce,
LSH {
num_tables: usize,
num_hyperplanes: usize,
},
}
#[cfg(test)]
mod tests;