weavatrix-search-vector 0.2.0

Persistent, mutable, bounded vector candidate search for Rust and Weavatrix
Documentation
use crate::{IndexConfig, SearchError, SearchHit, VectorIndex};
use std::fmt;

/// Pluggable text-to-vector provider.
///
/// The vector engine deliberately does not select a model, runtime, network
/// service, or tokenizer. Providers can wrap local models, hosted APIs, or
/// deterministic application-specific encoders.
pub trait EmbeddingProvider: Send + Sync {
    type Error: fmt::Display + Send + Sync + 'static;

    fn dimensions(&self) -> usize;

    /// Produces one dense embedding.
    ///
    /// # Errors
    ///
    /// Returns the provider-specific encoding or transport failure.
    fn embed(&self, input: &str) -> Result<Vec<f32>, Self::Error>;

    /// Produces embeddings in input order.
    ///
    /// # Errors
    ///
    /// Returns the first provider-specific failure.
    fn embed_batch(&self, inputs: &[&str]) -> Result<Vec<Vec<f32>>, Self::Error> {
        inputs.iter().map(|input| self.embed(input)).collect()
    }
}

/// Immutable vector index built from provider-generated embeddings.
#[derive(Debug)]
pub struct EmbeddingIndex {
    index: VectorIndex,
}

impl EmbeddingIndex {
    /// Embeds all texts as one provider batch and builds the vector index.
    ///
    /// # Errors
    ///
    /// Returns a provider failure, dimension mismatch, or vector-index build
    /// error.
    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
    }

    /// Embeds one query and returns nearest vector candidates.
    ///
    /// # Errors
    ///
    /// Returns a provider or vector-query error.
    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)
    }

    /// Embeds independent query strings as one provider batch and preserves
    /// input order.
    ///
    /// # Errors
    ///
    /// Returns a provider, query, allocation, or worker error.
    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)
    }
}