use crate::config::IndexConfig;
use crate::error::SearchError;
use crate::hit::SearchHit;
use crate::hnsw::VectorIndex;
use std::fmt;
pub trait EmbeddingProvider: Send + Sync {
type Error: fmt::Display + Send + Sync + 'static;
fn dimensions(&self) -> usize;
fn embed(&self, input: &str) -> Result<Vec<f32>, Self::Error>;
fn embed_batch(&self, inputs: &[&str]) -> Result<Vec<Vec<f32>>, Self::Error> {
inputs.iter().map(|input| self.embed(input)).collect()
}
}
#[derive(Debug)]
pub struct EmbeddingIndex {
index: VectorIndex,
}
impl EmbeddingIndex {
pub fn build<P>(
config: IndexConfig,
provider: &P,
texts: &[(u64, &str)],
) -> Result<Self, SearchError>
where
P: EmbeddingProvider,
{
if provider.dimensions() != config.dimensions {
return Err(SearchError::DimensionMismatch {
expected: config.dimensions,
actual: provider.dimensions(),
vector: None,
});
}
let inputs = texts.iter().map(|record| record.1).collect::<Vec<_>>();
let embeddings = provider
.embed_batch(&inputs)
.map_err(|error| SearchError::EmbeddingFailed(error.to_string()))?;
if embeddings.len() != texts.len() {
return Err(SearchError::EmbeddingFailed(format!(
"provider returned {} vectors for {} inputs",
embeddings.len(),
texts.len()
)));
}
let vectors = texts
.iter()
.zip(&embeddings)
.map(|((key, _), vector)| (*key, vector.as_slice()))
.collect::<Vec<_>>();
Ok(Self {
index: VectorIndex::build(config, &vectors)?,
})
}
#[must_use]
pub const fn as_vector_index(&self) -> &VectorIndex {
&self.index
}
#[must_use]
pub fn into_vector_index(self) -> VectorIndex {
self.index
}
pub fn search<P>(
&self,
provider: &P,
input: &str,
count: usize,
) -> Result<Vec<SearchHit>, SearchError>
where
P: EmbeddingProvider,
{
if provider.dimensions() != self.index.dimensions() {
return Err(SearchError::DimensionMismatch {
expected: self.index.dimensions(),
actual: provider.dimensions(),
vector: None,
});
}
let query = provider
.embed(input)
.map_err(|error| SearchError::EmbeddingFailed(error.to_string()))?;
self.index.search(&query, count)
}
pub fn search_batch<P>(
&self,
provider: &P,
inputs: &[&str],
count: usize,
) -> Result<Vec<Vec<SearchHit>>, SearchError>
where
P: EmbeddingProvider,
{
if provider.dimensions() != self.index.dimensions() {
return Err(SearchError::DimensionMismatch {
expected: self.index.dimensions(),
actual: provider.dimensions(),
vector: None,
});
}
let queries = provider
.embed_batch(inputs)
.map_err(|error| SearchError::EmbeddingFailed(error.to_string()))?;
if queries.len() != inputs.len() {
return Err(SearchError::EmbeddingFailed(format!(
"provider returned {} vectors for {} inputs",
queries.len(),
inputs.len()
)));
}
let borrowed = queries.iter().map(Vec::as_slice).collect::<Vec<_>>();
self.index.search_batch(&borrowed, count)
}
}