weavatrix_search_vector/
embedding.rs1use crate::config::IndexConfig;
2use crate::error::SearchError;
3use crate::hit::SearchHit;
4use crate::hnsw::VectorIndex;
5use std::fmt;
6
7pub trait EmbeddingProvider: Send + Sync {
13 type Error: fmt::Display + Send + Sync + 'static;
14
15 fn dimensions(&self) -> usize;
16
17 fn embed(&self, input: &str) -> Result<Vec<f32>, Self::Error>;
23
24 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#[derive(Debug)]
36pub struct EmbeddingIndex {
37 index: VectorIndex,
38}
39
40impl EmbeddingIndex {
41 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 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 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}