Skip to main content

rig_helixdb/
lib.rs

1//! HelixDB vector store integration for Rig.
2//!
3//! This crate provides a small HTTP client for HelixDB query endpoints and a
4//! [`HelixDBVectorStore`] implementation of Rig's vector store traits.
5//!
6//! The root `rig` facade re-exports this crate as `rig::helixdb` when the
7//! `helixdb` feature is enabled.
8
9use 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/// A minimal HelixDB HTTP client for running generated Helix queries.
20#[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    /// Creates a HelixDB client using the default reqwest client.
30    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    /// Creates a HelixDB client using a caller-provided reqwest client.
35    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/// Errors returned by the HelixDB HTTP client.
51#[derive(Debug, thiserror::Error)]
52pub enum HelixError {
53    /// A request to HelixDB failed before a response body could be decoded.
54    #[error("error communicating with server: {0}")]
55    ReqwestError(#[from] reqwest::Error),
56
57    /// HelixDB returned a non-200 response.
58    #[error("got error from server: {details}")]
59    RemoteError {
60        /// Response body or status reason returned by HelixDB.
61        details: String,
62    },
63}
64
65/// Client interface used by [`HelixDBVectorStore`] to execute HelixDB queries.
66pub trait HelixDBClient {
67    /// Error type returned by this client.
68    type Err: std::error::Error;
69
70    /// Sends a query payload to a HelixDB endpoint and decodes the response body.
71    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
114/// A client for easily carrying out Rig-related vector store operations.
115///
116/// If you are unsure what type to use for the client, [`HelixDB`] is the typical default.
117///
118/// Usage:
119/// ```no_run
120/// use rig_core::client::{EmbeddingsClient, ProviderClient};
121/// use rig_helixdb::{HelixDB, HelixDBVectorStore};
122///
123/// # fn example() -> anyhow::Result<()> {
124/// let openai_model = rig_core::providers::openai::Client::from_env()?
125///     .embedding_model("text-embedding-ada-002");
126///
127/// let helixdb_client = HelixDB::new(None, Some(6969), None);
128/// let vector_store = HelixDBVectorStore::new(helixdb_client, openai_model.clone());
129/// # let _ = vector_store;
130/// # Ok(())
131/// # }
132/// ```
133pub struct HelixDBVectorStore<C, E> {
134    client: C,
135    model: E,
136}
137
138pub type HelixDBFilter = Filter<serde_json::Value>;
139
140/// The result of a query. Only used internally as this is a representative type required for the relevant HelixDB query (`VectorSearch`).
141#[derive(Deserialize, Serialize, Clone, Debug)]
142struct QueryResult {
143    id: String,
144    score: f64,
145    doc: String,
146    json_payload: String,
147}
148
149/// An input query. Only used internally as this is a representative type required for the relevant HelixDB query (`VectorSearch`).
150#[derive(Deserialize, Serialize, Clone, Debug)]
151struct QueryInput {
152    vector: Vec<f64>,
153    limit: u64,
154    threshold: f64,
155}
156
157impl QueryInput {
158    /// Makes a new instance of `QueryInput`.
159    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/// The shape of a `VectorSearch` query response.
169#[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    /// Creates a new HelixDB vector store.
180    pub fn new(client: C, model: E) -> Self {
181        Self { client, model }
182    }
183
184    /// Returns the underlying HelixDB client.
185    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    /// Embeds the query and runs the `VectorSearch` HelixDB query.
197    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                // HelixDB gives us the cosine distance, so we need to use `-(cosine_dist - 1)` to get the cosine similarity score.
295                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        // HelixDB gives us the cosine distance, so we need to use `-(cosine_dist - 1)` to get the cosine similarity score.
307        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}