Skip to main content

weavatrix_search_vector/
embedding.rs

1use crate::config::IndexConfig;
2use crate::error::SearchError;
3use crate::hit::SearchHit;
4use crate::hnsw::VectorIndex;
5use std::fmt;
6
7/// Pluggable text-to-vector provider.
8///
9/// The vector engine deliberately does not select a model, runtime, network
10/// service, or tokenizer. Providers can wrap local models, hosted APIs, or
11/// deterministic application-specific encoders.
12pub trait EmbeddingProvider: Send + Sync {
13    type Error: fmt::Display + Send + Sync + 'static;
14
15    fn dimensions(&self) -> usize;
16
17    /// Produces one dense embedding.
18    ///
19    /// # Errors
20    ///
21    /// Returns the provider-specific encoding or transport failure.
22    fn embed(&self, input: &str) -> Result<Vec<f32>, Self::Error>;
23
24    /// Produces embeddings in input order.
25    ///
26    /// # Errors
27    ///
28    /// Returns the first provider-specific failure.
29    fn embed_batch(&self, inputs: &[&str]) -> Result<Vec<Vec<f32>>, Self::Error> {
30        inputs.iter().map(|input| self.embed(input)).collect()
31    }
32}
33
34/// Immutable vector index built from provider-generated embeddings.
35#[derive(Debug)]
36pub struct EmbeddingIndex {
37    index: VectorIndex,
38}
39
40impl EmbeddingIndex {
41    /// Embeds all texts as one provider batch and builds the vector index.
42    ///
43    /// # Errors
44    ///
45    /// Returns a provider failure, dimension mismatch, or vector-index build
46    /// error.
47    pub fn build<P>(
48        config: IndexConfig,
49        provider: &P,
50        texts: &[(u64, &str)],
51    ) -> Result<Self, SearchError>
52    where
53        P: EmbeddingProvider,
54    {
55        if provider.dimensions() != config.dimensions {
56            return Err(SearchError::DimensionMismatch {
57                expected: config.dimensions,
58                actual: provider.dimensions(),
59                vector: None,
60            });
61        }
62        let inputs = texts.iter().map(|record| record.1).collect::<Vec<_>>();
63        let embeddings = provider
64            .embed_batch(&inputs)
65            .map_err(|error| SearchError::EmbeddingFailed(error.to_string()))?;
66        if embeddings.len() != texts.len() {
67            return Err(SearchError::EmbeddingFailed(format!(
68                "provider returned {} vectors for {} inputs",
69                embeddings.len(),
70                texts.len()
71            )));
72        }
73        let vectors = texts
74            .iter()
75            .zip(&embeddings)
76            .map(|((key, _), vector)| (*key, vector.as_slice()))
77            .collect::<Vec<_>>();
78        Ok(Self {
79            index: VectorIndex::build(config, &vectors)?,
80        })
81    }
82
83    #[must_use]
84    pub const fn as_vector_index(&self) -> &VectorIndex {
85        &self.index
86    }
87
88    #[must_use]
89    pub fn into_vector_index(self) -> VectorIndex {
90        self.index
91    }
92
93    /// Embeds one query and returns nearest vector candidates.
94    ///
95    /// # Errors
96    ///
97    /// Returns a provider or vector-query error.
98    pub fn search<P>(
99        &self,
100        provider: &P,
101        input: &str,
102        count: usize,
103    ) -> Result<Vec<SearchHit>, SearchError>
104    where
105        P: EmbeddingProvider,
106    {
107        if provider.dimensions() != self.index.dimensions() {
108            return Err(SearchError::DimensionMismatch {
109                expected: self.index.dimensions(),
110                actual: provider.dimensions(),
111                vector: None,
112            });
113        }
114        let query = provider
115            .embed(input)
116            .map_err(|error| SearchError::EmbeddingFailed(error.to_string()))?;
117        self.index.search(&query, count)
118    }
119
120    /// Embeds independent query strings as one provider batch and preserves
121    /// input order.
122    ///
123    /// # Errors
124    ///
125    /// Returns a provider, query, allocation, or worker error.
126    pub fn search_batch<P>(
127        &self,
128        provider: &P,
129        inputs: &[&str],
130        count: usize,
131    ) -> Result<Vec<Vec<SearchHit>>, SearchError>
132    where
133        P: EmbeddingProvider,
134    {
135        if provider.dimensions() != self.index.dimensions() {
136            return Err(SearchError::DimensionMismatch {
137                expected: self.index.dimensions(),
138                actual: provider.dimensions(),
139                vector: None,
140            });
141        }
142        let queries = provider
143            .embed_batch(inputs)
144            .map_err(|error| SearchError::EmbeddingFailed(error.to_string()))?;
145        if queries.len() != inputs.len() {
146            return Err(SearchError::EmbeddingFailed(format!(
147                "provider returned {} vectors for {} inputs",
148                queries.len(),
149                inputs.len()
150            )));
151        }
152        let borrowed = queries.iter().map(Vec::as_slice).collect::<Vec<_>>();
153        self.index.search_batch(&borrowed, count)
154    }
155}