1use crate::{
4 CollectionHandle, Error, JsonObject, Point, PointId, Result, ScoredPoint, Store, WriteResult,
5};
6use serde::Serialize;
7use serde_json::Value;
8#[cfg(feature = "fastembed")]
9use std::sync::Mutex;
10
11#[cfg(feature = "fastembed")]
12pub use fastembed::{EmbeddingModel as FastEmbedModel, TextInitOptions as FastEmbedInitOptions};
13
14pub trait Embedder {
16 fn model_id(&self) -> &str;
18
19 fn embed(&self, input: &[String]) -> Result<Vec<Vec<f32>>>;
21}
22
23#[cfg(feature = "fastembed")]
28pub struct FastEmbedder {
29 model: Mutex<fastembed::TextEmbedding>,
30 model_id: String,
31}
32
33#[cfg(feature = "fastembed")]
34impl FastEmbedder {
35 pub fn try_new() -> Result<Self> {
37 Self::try_with_model(FastEmbedModel::default())
38 }
39
40 pub fn try_with_model(model: FastEmbedModel) -> Result<Self> {
42 Self::try_from_options(FastEmbedInitOptions::new(model))
43 }
44
45 pub fn try_from_options(options: FastEmbedInitOptions) -> Result<Self> {
47 let model_id = format!("fastembed/{}@5.17.3", options.model_name);
48 let model = fastembed::TextEmbedding::try_new(options)
49 .map_err(|error| Error::Embedding(error.to_string()))?;
50 Ok(Self {
51 model: Mutex::new(model),
52 model_id,
53 })
54 }
55}
56
57#[cfg(feature = "fastembed")]
58impl Embedder for FastEmbedder {
59 fn model_id(&self) -> &str {
60 &self.model_id
61 }
62
63 fn embed(&self, input: &[String]) -> Result<Vec<Vec<f32>>> {
64 self.model
65 .lock()
66 .map_err(|_| Error::Embedding("FastEmbed model lock was poisoned".into()))?
67 .embed(input, None)
68 .map_err(|error| Error::Embedding(error.to_string()))
69 }
70}
71
72#[derive(Clone, Debug, PartialEq)]
74pub struct Document {
75 pub id: PointId,
77 pub text: String,
79 pub metadata: JsonObject,
81}
82
83impl Document {
84 pub fn new(id: impl Into<PointId>, text: impl Into<String>) -> Self {
86 Self {
87 id: id.into(),
88 text: text.into(),
89 metadata: JsonObject::new(),
90 }
91 }
92
93 pub fn with_metadata(mut self, metadata: impl Serialize) -> Result<Self> {
95 match serde_json::to_value(metadata)? {
96 Value::Object(metadata) => {
97 self.metadata = metadata;
98 Ok(self)
99 }
100 _ => Err(Error::Invalid(
101 "document metadata must serialize to a JSON object".into(),
102 )),
103 }
104 }
105}
106
107#[derive(Clone, Debug)]
109pub struct TextCollection<E> {
110 collection: CollectionHandle,
111 embedder: E,
112 model_id: String,
113}
114
115impl Store {
116 pub fn text_collection<E: Embedder>(
118 &self,
119 name: impl Into<String>,
120 embedder: E,
121 ) -> Result<TextCollection<E>> {
122 TextCollection::new(self.collection(name), embedder)
123 }
124}
125
126impl<E: Embedder> TextCollection<E> {
127 fn new(collection: CollectionHandle, embedder: E) -> Result<Self> {
128 let model_id = embedder.model_id().trim().to_owned();
129 if model_id.is_empty() {
130 return Err(Error::Invalid(
131 "embedding model identity must not be empty".into(),
132 ));
133 }
134 if let Ok(existing) = collection.advanced() {
135 let actual = existing.info()?.config.vector_space;
136 if actual.as_deref() != Some(model_id.as_str()) {
137 return Err(Error::Invalid(format!(
138 "collection uses vector space {actual:?}, expected {model_id:?}"
139 )));
140 }
141 }
142 Ok(Self {
143 collection,
144 embedder,
145 model_id,
146 })
147 }
148
149 pub fn upsert_documents(
151 &self,
152 documents: impl IntoIterator<Item = Document>,
153 ) -> Result<WriteResult> {
154 let documents: Vec<Document> = documents.into_iter().collect();
155 if documents.is_empty() {
156 return Err(Error::Invalid("document batch must not be empty".into()));
157 }
158 let input: Vec<String> = documents
159 .iter()
160 .map(|document| document.text.clone())
161 .collect();
162 let vectors = self.embedder.embed(&input)?;
163 if vectors.len() != documents.len() {
164 return Err(Error::Invalid(format!(
165 "embedder returned {} vectors for {} documents",
166 vectors.len(),
167 documents.len()
168 )));
169 }
170 let points = documents
171 .into_iter()
172 .zip(vectors)
173 .map(|(document, vector)| {
174 let mut payload = document.metadata;
175 if payload
176 .insert("document".into(), Value::String(document.text))
177 .is_some()
178 {
179 return Err(Error::Invalid(
180 "document metadata reserves the key \"document\"".into(),
181 ));
182 }
183 Ok(Point {
184 id: document.id,
185 vector,
186 payload,
187 })
188 })
189 .collect::<Result<Vec<_>>>()?;
190 self.collection
191 .upsert_with_vector_space(points, Some(&self.model_id))
192 }
193
194 pub fn search_text(&self, text: impl Into<String>, limit: usize) -> Result<Vec<ScoredPoint>> {
196 let vectors = self.embedder.embed(&[text.into()])?;
197 let mut vectors = vectors.into_iter();
198 let vector = vectors
199 .next()
200 .ok_or_else(|| Error::Invalid("embedder returned no query vector".into()))?;
201 if vectors.next().is_some() {
202 return Err(Error::Invalid(
203 "embedder returned multiple vectors for one query".into(),
204 ));
205 }
206 self.collection.search(vector, limit)
207 }
208
209 pub fn vectors(&self) -> &CollectionHandle {
211 &self.collection
212 }
213}