1use crate::{
4 CollectionHandle, DeleteSelector, Error, Filter, JsonObject, MutationResult, Point, PointId,
5 Query, QueryParams, Result, ScoredPoint, SnapshotMutation, Store, WriteResult,
6};
7use serde::{Deserialize, Serialize};
8use serde_json::Value;
9#[cfg(feature = "fastembed")]
10use std::sync::Mutex;
11#[cfg(feature = "fastembed")]
12use std::{env, path::PathBuf};
13
14#[cfg(feature = "fastembed")]
15pub use fastembed::{EmbeddingModel as FastEmbedModel, TextInitOptions as FastEmbedInitOptions};
16
17pub trait Embedder {
19 fn model_id(&self) -> &str;
21
22 fn embed(&self, input: &[String]) -> Result<Vec<Vec<f32>>>;
24}
25
26#[cfg(feature = "fastembed")]
31pub struct FastEmbedder {
32 model: Mutex<fastembed::TextEmbedding>,
33 model_id: String,
34}
35
36#[cfg(feature = "fastembed")]
37impl FastEmbedder {
38 pub fn try_new() -> Result<Self> {
40 Self::try_with_model(FastEmbedModel::default())
41 }
42
43 pub fn try_with_model(model: FastEmbedModel) -> Result<Self> {
45 let mut options = FastEmbedInitOptions::new(model);
46 if env::var_os("FASTEMBED_CACHE_DIR").is_none() {
47 if let Some(cache) = default_fastembed_cache_dir() {
48 options = options.with_cache_dir(cache);
49 }
50 }
51 Self::try_from_options(options)
52 }
53
54 pub fn try_from_options(options: FastEmbedInitOptions) -> Result<Self> {
56 let model_id = format!("fastembed/{}@5.17.3", options.model_name);
57 let model = fastembed::TextEmbedding::try_new(options)
58 .map_err(|error| Error::Embedding(error.to_string()))?;
59 Ok(Self {
60 model: Mutex::new(model),
61 model_id,
62 })
63 }
64}
65
66#[cfg(feature = "fastembed")]
67fn default_fastembed_cache_dir() -> Option<PathBuf> {
68 if let Some(cache) = env::var_os("XDG_CACHE_HOME") {
69 return Some(PathBuf::from(cache).join("git-vdb/fastembed"));
70 }
71 if let Some(cache) = env::var_os("LOCALAPPDATA") {
72 return Some(PathBuf::from(cache).join("git-vdb/fastembed"));
73 }
74 env::var_os("HOME").map(|home| {
75 let home = PathBuf::from(home);
76 if cfg!(target_os = "macos") {
77 home.join("Library/Caches/git-vdb/fastembed")
78 } else {
79 home.join(".cache/git-vdb/fastembed")
80 }
81 })
82}
83
84#[cfg(feature = "fastembed")]
85impl Embedder for FastEmbedder {
86 fn model_id(&self) -> &str {
87 &self.model_id
88 }
89
90 fn embed(&self, input: &[String]) -> Result<Vec<Vec<f32>>> {
91 self.model
92 .lock()
93 .map_err(|_| Error::Embedding("FastEmbed model lock was poisoned".into()))?
94 .embed(input, None)
95 .map_err(|error| Error::Embedding(error.to_string()))
96 }
97}
98
99#[derive(Clone, Debug, PartialEq)]
101pub struct Document {
102 pub id: PointId,
104 pub text: String,
106 pub metadata: JsonObject,
108}
109
110#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
112pub struct DocumentHit {
113 pub id: PointId,
115 pub document: String,
117 pub metadata: JsonObject,
119 pub score: f32,
121}
122
123#[derive(Clone, Debug, Serialize, Deserialize)]
125pub struct TextQuery {
126 pub text: String,
128 pub limit: usize,
130 pub filter: Option<Filter>,
132 pub params: QueryParams,
134}
135
136impl TextQuery {
137 pub fn new(text: impl Into<String>) -> Self {
139 Self {
140 text: text.into(),
141 limit: 10,
142 filter: None,
143 params: QueryParams::default(),
144 }
145 }
146
147 #[must_use]
149 pub fn limit(mut self, limit: usize) -> Self {
150 self.limit = limit;
151 self
152 }
153
154 #[must_use]
156 pub fn with_filter(mut self, filter: Filter) -> Self {
157 self.filter = Some(filter);
158 self
159 }
160
161 #[must_use]
163 pub fn with_params(mut self, params: QueryParams) -> Self {
164 self.params = params;
165 self
166 }
167}
168
169impl Document {
170 pub fn new(id: impl Into<PointId>, text: impl Into<String>) -> Self {
172 Self {
173 id: id.into(),
174 text: text.into(),
175 metadata: JsonObject::new(),
176 }
177 }
178
179 pub fn with_metadata(mut self, metadata: impl Serialize) -> Result<Self> {
181 match serde_json::to_value(metadata)? {
182 Value::Object(metadata) => {
183 self.metadata = metadata;
184 Ok(self)
185 }
186 _ => Err(Error::Invalid(
187 "document metadata must serialize to a JSON object".into(),
188 )),
189 }
190 }
191}
192
193#[derive(Clone, Debug)]
195pub struct TextCollection<E> {
196 collection: CollectionHandle,
197 embedder: E,
198 model_id: String,
199}
200
201impl Store {
202 pub fn text_collection<E: Embedder>(
204 &self,
205 name: impl Into<String>,
206 embedder: E,
207 ) -> Result<TextCollection<E>> {
208 TextCollection::new(self.collection(name), embedder)
209 }
210}
211
212impl<E: Embedder> TextCollection<E> {
213 fn new(collection: CollectionHandle, embedder: E) -> Result<Self> {
214 let model_id = embedder.model_id().trim().to_owned();
215 if model_id.is_empty() {
216 return Err(Error::Invalid(
217 "embedding model identity must not be empty".into(),
218 ));
219 }
220 if let Ok(existing) = collection.advanced() {
221 let actual = existing.info()?.config.vector_space;
222 if actual.as_deref() != Some(model_id.as_str()) {
223 return Err(Error::Invalid(format!(
224 "collection uses vector space {actual:?}, expected {model_id:?}"
225 )));
226 }
227 }
228 Ok(Self {
229 collection,
230 embedder,
231 model_id,
232 })
233 }
234
235 pub fn upsert_documents(
237 &self,
238 documents: impl IntoIterator<Item = Document>,
239 ) -> Result<WriteResult> {
240 let points = self.embed_documents(documents)?;
241 self.collection
242 .upsert_with_vector_space(points, Some(&self.model_id))
243 }
244
245 pub fn replace_documents(
250 &self,
251 filter: Filter,
252 documents: impl IntoIterator<Item = Document>,
253 ) -> Result<MutationResult> {
254 let points = self.embed_documents(documents)?;
255 let mut mutations = Vec::with_capacity(points.len() + 1);
256 mutations.push(SnapshotMutation::delete_filter(filter));
257 mutations.extend(points.into_iter().map(SnapshotMutation::upsert));
258 self.collection.apply(mutations)
259 }
260
261 pub fn delete(&self, selector: DeleteSelector) -> Result<WriteResult> {
263 self.collection.delete(selector)
264 }
265
266 pub fn search_text(&self, text: impl Into<String>, limit: usize) -> Result<Vec<ScoredPoint>> {
268 let vectors = self.embedder.embed(&[text.into()])?;
269 let mut vectors = vectors.into_iter();
270 let vector = vectors
271 .next()
272 .ok_or_else(|| Error::Invalid("embedder returned no query vector".into()))?;
273 if vectors.next().is_some() {
274 return Err(Error::Invalid(
275 "embedder returned multiple vectors for one query".into(),
276 ));
277 }
278 self.collection.search(vector, limit)
279 }
280
281 pub fn query(&self, query: TextQuery) -> Result<Vec<DocumentHit>> {
283 self.query_batch([query])?
284 .pop()
285 .ok_or_else(|| Error::Invalid("text query batch returned no result".into()))
286 }
287
288 pub fn query_batch(
290 &self,
291 queries: impl IntoIterator<Item = TextQuery>,
292 ) -> Result<Vec<Vec<DocumentHit>>> {
293 let queries: Vec<TextQuery> = queries.into_iter().collect();
294 if queries.is_empty() {
295 return Ok(Vec::new());
296 }
297 let input: Vec<String> = queries.iter().map(|query| query.text.clone()).collect();
298 let vectors = self.embedder.embed(&input)?;
299 if vectors.len() != queries.len() {
300 return Err(Error::Invalid(format!(
301 "embedder returned {} vectors for {} text queries",
302 vectors.len(),
303 queries.len()
304 )));
305 }
306 let vector_queries = queries
307 .into_iter()
308 .zip(vectors)
309 .map(|(text, vector)| Query {
310 vector,
311 limit: text.limit,
312 filter: text.filter,
313 with_payload: true,
314 expected_vector_space: Some(self.model_id.clone()),
315 params: text.params,
316 ..Query::default()
317 });
318 self.collection
319 .query_batch(vector_queries)?
320 .into_iter()
321 .map(|result| result.points.into_iter().map(document_hit).collect())
322 .collect()
323 }
324
325 pub fn vectors(&self) -> &CollectionHandle {
327 &self.collection
328 }
329
330 fn embed_documents(&self, documents: impl IntoIterator<Item = Document>) -> Result<Vec<Point>> {
331 let documents: Vec<Document> = documents.into_iter().collect();
332 if documents.is_empty() {
333 return Err(Error::Invalid("document batch must not be empty".into()));
334 }
335 let input: Vec<String> = documents
336 .iter()
337 .map(|document| document.text.clone())
338 .collect();
339 let vectors = self.embedder.embed(&input)?;
340 if vectors.len() != documents.len() {
341 return Err(Error::Invalid(format!(
342 "embedder returned {} vectors for {} documents",
343 vectors.len(),
344 documents.len()
345 )));
346 }
347 documents
348 .into_iter()
349 .zip(vectors)
350 .map(|(document, vector)| {
351 let mut payload = document.metadata;
352 if payload
353 .insert("document".into(), Value::String(document.text))
354 .is_some()
355 {
356 return Err(Error::Invalid(
357 "document metadata reserves the key \"document\"".into(),
358 ));
359 }
360 Ok(Point {
361 id: document.id,
362 vector,
363 payload,
364 })
365 })
366 .collect()
367 }
368}
369
370fn document_hit(point: ScoredPoint) -> Result<DocumentHit> {
371 let mut metadata = point
372 .payload
373 .ok_or_else(|| Error::Corrupt("text query result omitted its payload".into()))?;
374 let document = metadata
375 .remove("document")
376 .and_then(|value| value.as_str().map(str::to_owned))
377 .ok_or_else(|| Error::Corrupt("text point has no string document field".into()))?;
378 Ok(DocumentHit {
379 id: point.id,
380 document,
381 metadata,
382 score: point.score,
383 })
384}