1use std::future::Future;
10
11use reqwest::{Client, StatusCode};
12use rig_core::{
13 embeddings::EmbeddingModel,
14 vector_store::{InsertDocuments, VectorStoreError, VectorStoreIndex, request::Filter},
15 wasm_compat::{WasmCompatSend, WasmCompatSync},
16};
17use serde::{Deserialize, Serialize};
18
19#[derive(Debug, Clone)]
21pub struct HelixDB {
22 port: Option<u16>,
23 client: Client,
24 endpoint: String,
25 api_key: Option<String>,
26}
27
28impl HelixDB {
29 pub fn new(endpoint: Option<&str>, port: Option<u16>, api_key: Option<&str>) -> Self {
31 Self::with_client(endpoint, port, api_key, Client::new())
32 }
33
34 pub fn with_client(
36 endpoint: Option<&str>,
37 port: Option<u16>,
38 api_key: Option<&str>,
39 client: Client,
40 ) -> Self {
41 Self {
42 port,
43 client,
44 endpoint: endpoint.unwrap_or("http://localhost").to_string(),
45 api_key: api_key.map(ToString::to_string),
46 }
47 }
48}
49
50#[derive(Debug, thiserror::Error)]
52pub enum HelixError {
53 #[error("error communicating with server: {0}")]
55 ReqwestError(#[from] reqwest::Error),
56
57 #[error("got error from server: {details}")]
59 RemoteError {
60 details: String,
62 },
63}
64
65pub trait HelixDBClient {
67 type Err: std::error::Error;
69
70 fn query<T, R>(
72 &self,
73 endpoint: &str,
74 data: &T,
75 ) -> impl Future<Output = Result<R, Self::Err>> + WasmCompatSend
76 where
77 T: Serialize + WasmCompatSync,
78 R: for<'de> Deserialize<'de>;
79}
80
81impl HelixDBClient for HelixDB {
82 type Err = HelixError;
83
84 async fn query<T, R>(&self, endpoint: &str, data: &T) -> Result<R, HelixError>
85 where
86 T: Serialize + WasmCompatSync,
87 R: for<'de> Deserialize<'de>,
88 {
89 let port = self.port.map(|port| format!(":{port}")).unwrap_or_default();
90 let url = format!("{}{}/{}", self.endpoint, port, endpoint);
91
92 let mut request = self.client.post(&url).json(data);
93 if let Some(api_key) = &self.api_key {
94 request = request.header("x-api-key", api_key);
95 }
96
97 let response = request.send().await?;
98
99 match response.status() {
100 StatusCode::OK => response.json().await.map_err(Into::into),
101 code => match response.text().await {
102 Ok(details) => Err(HelixError::RemoteError { details }),
103 Err(_) => Err(HelixError::RemoteError {
104 details: code
105 .canonical_reason()
106 .map(ToString::to_string)
107 .unwrap_or_else(|| format!("unknown error with code: {code}")),
108 }),
109 },
110 }
111 }
112}
113
114pub struct HelixDBVectorStore<C, E> {
134 client: C,
135 model: E,
136}
137
138pub type HelixDBFilter = Filter<serde_json::Value>;
139
140#[derive(Deserialize, Serialize, Clone, Debug)]
142struct QueryResult {
143 id: String,
144 score: f64,
145 doc: String,
146 json_payload: String,
147}
148
149#[derive(Deserialize, Serialize, Clone, Debug)]
151struct QueryInput {
152 vector: Vec<f64>,
153 limit: u64,
154 threshold: f64,
155}
156
157impl QueryInput {
158 pub(crate) fn new(vector: Vec<f64>, limit: u64, threshold: f64) -> Self {
160 Self {
161 vector,
162 limit,
163 threshold,
164 }
165 }
166}
167
168#[derive(Serialize, Deserialize, Debug)]
170struct VecResult {
171 vec_docs: Vec<QueryResult>,
172}
173
174impl<C, E> HelixDBVectorStore<C, E>
175where
176 C: HelixDBClient + WasmCompatSend,
177 E: EmbeddingModel,
178{
179 pub fn new(client: C, model: E) -> Self {
181 Self { client, model }
182 }
183
184 pub fn client(&self) -> &C {
186 &self.client
187 }
188}
189
190impl<C, E> HelixDBVectorStore<C, E>
191where
192 C: HelixDBClient + WasmCompatSend + WasmCompatSync,
193 C::Err: std::error::Error + WasmCompatSend + WasmCompatSync + 'static,
194 E: EmbeddingModel + WasmCompatSend + WasmCompatSync,
195{
196 async fn vector_search(
198 &self,
199 req: &rig_core::vector_store::VectorSearchRequest<HelixDBFilter>,
200 ) -> Result<Vec<QueryResult>, VectorStoreError> {
201 let vector = self.model.embed_text(req.query()).await?.vec;
202
203 let query_input =
204 QueryInput::new(vector, req.samples(), req.threshold().unwrap_or_default());
205
206 let result: VecResult = self
207 .client
208 .query::<QueryInput, VecResult>("VectorSearch", &query_input)
209 .await
210 .map_err(VectorStoreError::datastore)?;
211
212 Ok(result.vec_docs)
213 }
214}
215
216impl<C, E> InsertDocuments for HelixDBVectorStore<C, E>
217where
218 C: HelixDBClient + WasmCompatSend + WasmCompatSync,
219 C::Err: std::error::Error + WasmCompatSend + WasmCompatSync + 'static,
220 E: EmbeddingModel + WasmCompatSend + WasmCompatSync,
221{
222 async fn insert_documents<Doc: Serialize + rig_core::Embed + WasmCompatSend>(
223 &self,
224 documents: Vec<(Doc, Vec<rig_core::embeddings::Embedding>)>,
225 ) -> Result<(), VectorStoreError> {
226 #[derive(Serialize, Deserialize, Clone, Debug, Default)]
227 struct QueryInput {
228 vector: Vec<f64>,
229 doc: String,
230 json_payload: String,
231 }
232
233 #[derive(Serialize, Deserialize, Clone, Debug, Default)]
234 struct QueryOutput {
235 doc: String,
236 }
237
238 let queries =
239 rig_core::vector_store::flatten_embedded(documents, |json_document, embedding| {
240 Ok(QueryInput {
241 vector: embedding.vec,
242 doc: embedding.document,
243 json_payload: serde_json::to_string(json_document)?,
244 })
245 })?;
246
247 for query in queries {
248 self.client
249 .query::<QueryInput, QueryOutput>("InsertVector", &query)
250 .await
251 .map_err(VectorStoreError::datastore)?;
252 }
253 Ok(())
254 }
255}
256
257impl<C, E> VectorStoreIndex for HelixDBVectorStore<C, E>
258where
259 C: HelixDBClient + WasmCompatSend + WasmCompatSync,
260 C::Err: std::error::Error + WasmCompatSend + WasmCompatSync + 'static,
261 E: EmbeddingModel + WasmCompatSend + WasmCompatSync,
262{
263 type Filter = HelixDBFilter;
264
265 async fn top_n<T: for<'a> serde::Deserialize<'a> + WasmCompatSend>(
266 &self,
267 req: rig_core::vector_store::VectorSearchRequest<HelixDBFilter>,
268 ) -> Result<Vec<(f64, String, T)>, rig_core::vector_store::VectorStoreError> {
269 let docs = self
270 .vector_search(&req)
271 .await?
272 .into_iter()
273 .filter(|x| {
274 let is_threshold = req
275 .threshold()
276 .map(|t| -(x.score - 1.) >= t)
277 .unwrap_or(true);
278
279 is_threshold
280 && req
281 .filter()
282 .clone()
283 .zip(serde_json::from_str(&x.json_payload).ok())
284 .map(
285 |(filter, payload): (Filter<serde_json::Value>, serde_json::Value)| {
286 filter.satisfies(&payload)
287 },
288 )
289 .unwrap_or(true)
290 })
291 .map(|x| {
292 let doc: T = serde_json::from_str(&x.json_payload)?;
293
294 Ok((-(x.score - 1.), x.id, doc))
296 })
297 .collect::<Result<Vec<_>, VectorStoreError>>()?;
298
299 Ok(docs)
300 }
301
302 async fn top_n_ids(
303 &self,
304 req: rig_core::vector_store::VectorSearchRequest<HelixDBFilter>,
305 ) -> Result<Vec<(f64, String)>, rig_core::vector_store::VectorStoreError> {
306 let docs = self
308 .vector_search(&req)
309 .await?
310 .into_iter()
311 .filter(|x| -(x.score - 1.) >= req.threshold().unwrap_or_default())
312 .map(|x| Ok((-(x.score - 1.), x.id)))
313 .collect::<Result<Vec<_>, VectorStoreError>>()?;
314
315 Ok(docs)
316 }
317}