Skip to main content

rig_core/vector_store/
request.rs

1//! Vector-search requests and composable metadata filters.
2//!
3//! ```
4//! use rig_core::vector_store::request::{VectorSearchRequest, Filter, SearchFilter};
5//!
6//! let request: VectorSearchRequest = VectorSearchRequest::builder()
7//!     .query("guide").samples(5)
8//!     .filter(Filter::eq("category", serde_json::json!("manual"))).build();
9//! assert_eq!(request.samples(), 5);
10//! ```
11
12use serde::{Deserialize, Serialize};
13
14use super::VectorStoreError;
15use crate::markers::{Missing, Provided};
16
17/// A vector search request for querying a [`super::VectorStoreIndex`].
18///
19/// The type parameter `F` specifies the filter type (defaults to [`Filter<serde_json::Value>`]).
20/// Use [`VectorSearchRequest::builder()`] to construct instances.
21#[derive(Clone, Serialize, Deserialize, Debug)]
22pub struct VectorSearchRequest<F = Filter<serde_json::Value>> {
23    /// The query text to embed and search with.
24    query: String,
25    /// Maximum number of results to return.
26    samples: u64,
27    /// Minimum similarity score for results.
28    threshold: Option<f64>,
29    /// Backend-specific parameters as a JSON object.
30    additional_params: Option<serde_json::Value>,
31    /// Filter expression to narrow results by metadata.
32    filter: Option<F>,
33}
34
35impl<Filter> VectorSearchRequest<Filter> {
36    /// Creates a [`VectorSearchRequestBuilder`] which you can use to instantiate this struct.
37    pub fn builder() -> VectorSearchRequestBuilder<Filter> {
38        VectorSearchRequestBuilder::<Filter>::default()
39    }
40
41    /// The query to be embedded and used in similarity search.
42    pub fn query(&self) -> &str {
43        &self.query
44    }
45
46    /// Returns the maximum number of results to return.
47    pub fn samples(&self) -> u64 {
48        self.samples
49    }
50
51    /// Returns the optional similarity threshold.
52    pub fn threshold(&self) -> Option<f64> {
53        self.threshold
54    }
55
56    /// Returns a reference to the optional filter expression.
57    pub fn filter(&self) -> &Option<Filter> {
58        &self.filter
59    }
60
61    /// Transforms the filter type using the provided function.
62    ///
63    /// This is useful for converting between filter representations, such as
64    /// translating the canonical [`super::request::Filter`] to a backend-specific filter type.
65    pub fn map_filter<T, F>(self, f: F) -> VectorSearchRequest<T>
66    where
67        F: Fn(Filter) -> T,
68    {
69        VectorSearchRequest {
70            query: self.query,
71            samples: self.samples,
72            threshold: self.threshold,
73            additional_params: self.additional_params,
74            filter: self.filter.map(f),
75        }
76    }
77
78    /// Transforms the filter type using a provided function which can additionally return a result.
79    ///
80    /// Useful for converting between filter representations where the conversion can potentially fail (eg, unrepresentable or invalid values).
81    pub fn try_map_filter<T, F>(self, f: F) -> Result<VectorSearchRequest<T>, FilterError>
82    where
83        F: Fn(Filter) -> Result<T, FilterError>,
84    {
85        let filter = self.filter.map(f).transpose()?;
86
87        Ok(VectorSearchRequest {
88            query: self.query,
89            samples: self.samples,
90            threshold: self.threshold,
91            additional_params: self.additional_params,
92            filter,
93        })
94    }
95}
96
97/// Errors from constructing or converting filter expressions.
98#[derive(Debug, Clone, thiserror::Error)]
99pub enum FilterError {
100    #[error("Expected: {expected}, got: {got}")]
101    Expected { expected: String, got: String },
102
103    #[error("Cannot compile '{0}' to the backend's filter type")]
104    TypeError(String),
105
106    #[error("Missing field '{0}'")]
107    MissingField(String),
108
109    #[error("'{0}' must {1}")]
110    Must(String, String),
111
112    // NOTE: Uses String because `serde_json::Error` is not `Clone`.
113    #[error("Filter serialization failed: {0}")]
114    Serialization(String),
115}
116
117/// Trait for constructing filter expressions in vector search queries.
118///
119/// Uses [tagless final](https://nrinaudo.github.io/articles/tagless_final.html) encoding
120/// for backend-agnostic filters. Use `SearchFilter::eq(...)` etc. directly and let
121/// type inference resolve the concrete filter type.
122pub trait SearchFilter {
123    type Value;
124
125    fn eq(key: impl AsRef<str>, value: Self::Value) -> Self;
126    fn gt(key: impl AsRef<str>, value: Self::Value) -> Self;
127    fn lt(key: impl AsRef<str>, value: Self::Value) -> Self;
128    fn and(self, rhs: Self) -> Self;
129    fn or(self, rhs: Self) -> Self;
130}
131
132/// A rendered SQL-style condition together with its positional bind parameters.
133///
134/// Parameters retain left-to-right placeholder order. Keys, operators, and
135/// placeholders are interpolated verbatim; callers must validate or quote them
136/// for the target backend. Only values are separated as bind parameters.
137#[derive(Clone, Debug, Serialize, Deserialize)]
138pub struct SqlCondition<P> {
139    condition: String,
140    params: Vec<P>,
141}
142
143/// Hand-written so that `P` needs no [`Default`] of its own: a parameterless
144/// empty condition is meaningful for every parameter type.
145impl<P> Default for SqlCondition<P> {
146    fn default() -> Self {
147        Self {
148            condition: String::new(),
149            params: Vec::new(),
150        }
151    }
152}
153
154impl<P> SqlCondition<P> {
155    /// Renders `<key> <op> <placeholder>` bound to a single parameter, e.g.
156    /// `price >= $`.
157    pub fn binary(key: impl AsRef<str>, op: &str, placeholder: &str, value: P) -> Self {
158        Self {
159            condition: format!("{} {op} {placeholder}", key.as_ref()),
160            params: vec![value],
161        }
162    }
163
164    /// Renders `<key> <op> (<placeholder>, ...)` with one placeholder per value,
165    /// e.g. `id IN (?, ?)`.
166    pub fn list(key: impl AsRef<str>, op: &str, placeholder: &str, values: Vec<P>) -> Self {
167        let placeholders = vec![placeholder; values.len()].join(", ");
168
169        Self {
170            condition: format!("{} {op} ({placeholders})", key.as_ref()),
171            params: values,
172        }
173    }
174
175    /// Wraps an already-rendered, parameterless condition such as `id is null`.
176    pub fn raw(condition: impl Into<String>) -> Self {
177        Self {
178            condition: condition.into(),
179            params: Vec::new(),
180        }
181    }
182
183    /// Conjoins two conditions as `(lhs) AND (rhs)`, concatenating their parameters.
184    pub fn and(self, rhs: Self) -> Self {
185        self.combine("AND", rhs)
186    }
187
188    /// Disjoins two conditions as `(lhs) OR (rhs)`, concatenating their parameters.
189    pub fn or(self, rhs: Self) -> Self {
190        self.combine("OR", rhs)
191    }
192
193    /// Negates the condition as `NOT (condition)`, keeping its parameters.
194    pub fn not(self) -> Self {
195        Self {
196            condition: format!("NOT ({})", self.condition),
197            ..self
198        }
199    }
200
201    fn combine(self, joiner: &str, rhs: Self) -> Self {
202        Self {
203            condition: format!("({}) {joiner} ({})", self.condition, rhs.condition),
204            params: self.params.into_iter().chain(rhs.params).collect(),
205        }
206    }
207
208    /// The rendered condition, with placeholders as the caller supplied them.
209    pub fn condition(&self) -> &str {
210        &self.condition
211    }
212
213    /// The bind parameters, in placeholder order.
214    pub fn params(&self) -> &[P] {
215        &self.params
216    }
217
218    /// Consumes the condition, returning the rendered text and its parameters.
219    pub fn into_parts(self) -> (String, Vec<P>) {
220        (self.condition, self.params)
221    }
222}
223
224/// Converts the canonical JSON-valued [`Filter`] used by type-erased vector
225/// searches into a backend's native filter representation.
226///
227/// JSON-valued [`SearchFilter`] implementations receive this automatically.
228/// Backends with native value types implement the conversion once here rather
229/// than hand-writing both the bus's `RetrieveAdapter`
230/// methods.
231pub trait DynamicSearchFilter: SearchFilter + Sized {
232    /// Converts a canonical dynamic filter into this backend's filter type.
233    fn from_dynamic_filter(filter: Filter<serde_json::Value>) -> Result<Self, FilterError>;
234
235    /// Normalizes a document returned through the type-erased search surface.
236    ///
237    /// The default retains documents unchanged. The JSON-valued blanket
238    /// implementation recursively removes arrays longer than 400 elements;
239    /// a removed root array becomes null.
240    fn normalize_dynamic_document(document: serde_json::Value) -> serde_json::Value {
241        document
242    }
243}
244
245impl<F> DynamicSearchFilter for F
246where
247    F: SearchFilter<Value = serde_json::Value>,
248{
249    fn from_dynamic_filter(filter: Filter<serde_json::Value>) -> Result<Self, FilterError> {
250        Ok(filter.interpret())
251    }
252
253    fn normalize_dynamic_document(document: serde_json::Value) -> serde_json::Value {
254        prune_document(document).unwrap_or_default()
255    }
256}
257
258fn prune_document(document: serde_json::Value) -> Option<serde_json::Value> {
259    match document {
260        serde_json::Value::Object(mut map) => {
261            let new_map = map
262                .iter_mut()
263                .filter_map(|(key, value)| {
264                    prune_document(value.take()).map(|value| (key.clone(), value))
265                })
266                .collect::<serde_json::Map<_, _>>();
267
268            Some(serde_json::Value::Object(new_map))
269        }
270        serde_json::Value::Array(vec) if vec.len() > 400 => None,
271        serde_json::Value::Array(vec) => Some(serde_json::Value::Array(
272            vec.into_iter().filter_map(prune_document).collect(),
273        )),
274        value => Some(value),
275    }
276}
277
278/// Canonical, serializable filter representation.
279///
280/// Use for serialization, runtime inspection, or translating between backends via
281/// [`Filter::interpret`]. Prefer [`SearchFilter`] trait methods for writing queries.
282#[derive(Debug, Clone, Serialize, Deserialize)]
283#[serde(rename_all = "lowercase")]
284pub enum Filter<V>
285where
286    V: std::fmt::Debug + Clone,
287{
288    Eq(String, V),
289    Gt(String, V),
290    Lt(String, V),
291    And(Box<Self>, Box<Self>),
292    Or(Box<Self>, Box<Self>),
293}
294
295impl<V> SearchFilter for Filter<V>
296where
297    V: std::fmt::Debug + Clone + Serialize + serde::de::DeserializeOwned,
298{
299    type Value = V;
300
301    /// Select values where the entry at `key` is equal to `value`
302    fn eq(key: impl AsRef<str>, value: Self::Value) -> Self {
303        Self::Eq(key.as_ref().to_owned(), value)
304    }
305
306    /// Select values where the entry at `key` is greater than `value`
307    fn gt(key: impl AsRef<str>, value: Self::Value) -> Self {
308        Self::Gt(key.as_ref().to_owned(), value)
309    }
310
311    /// Select values where the entry at `key` is less than `value`
312    fn lt(key: impl AsRef<str>, value: Self::Value) -> Self {
313        Self::Lt(key.as_ref().to_owned(), value)
314    }
315
316    /// Select values where the entry satisfies `self` *and* `rhs`
317    fn and(self, rhs: Self) -> Self {
318        Self::And(self.into(), rhs.into())
319    }
320
321    /// Select values where the entry satisfies `self` *or* `rhs`
322    fn or(self, rhs: Self) -> Self {
323        Self::Or(self.into(), rhs.into())
324    }
325}
326
327impl<V> Filter<V>
328where
329    V: std::fmt::Debug + Clone,
330{
331    /// Converts this filter into a backend-specific filter type.
332    pub fn interpret<F>(self) -> F
333    where
334        F: SearchFilter<Value = V>,
335    {
336        self.interpret_with(|v| v)
337    }
338
339    /// Converts this filter into a backend-specific filter type, converting
340    /// each leaf value with `conv`.
341    pub fn interpret_with<F, W>(self, conv: impl Fn(V) -> W + Copy) -> F
342    where
343        F: SearchFilter<Value = W>,
344    {
345        match self.try_interpret(|v| Ok::<W, std::convert::Infallible>(conv(v))) {
346            Ok(filter) => filter,
347            Err(never) => match never {},
348        }
349    }
350
351    /// Converts this filter into a backend-specific filter type, converting
352    /// each leaf value with `conv`. Fails on the first value that fails to convert.
353    pub fn try_interpret<F, W, E>(self, conv: impl Fn(V) -> Result<W, E> + Copy) -> Result<F, E>
354    where
355        F: SearchFilter<Value = W>,
356    {
357        Ok(match self {
358            Self::Eq(key, val) => F::eq(key, conv(val)?),
359            Self::Gt(key, val) => F::gt(key, conv(val)?),
360            Self::Lt(key, val) => F::lt(key, conv(val)?),
361            Self::And(lhs, rhs) => F::and(lhs.try_interpret(conv)?, rhs.try_interpret(conv)?),
362            Self::Or(lhs, rhs) => F::or(lhs.try_interpret(conv)?, rhs.try_interpret(conv)?),
363        })
364    }
365}
366
367impl Filter<serde_json::Value> {
368    /// Tests whether a JSON document satisfies this filter.
369    ///
370    /// Looks up top-level object fields. Missing fields fail all leaves.
371    /// Equality accepts numeric comparison or structural JSON equality; ordering
372    /// compares only number, string, boolean, or null pairs of the same kind.
373    /// `And` and `Or` short-circuit.
374    pub fn satisfies(&self, value: &serde_json::Value) -> bool {
375        use Filter::*;
376        use serde_json::{Value, Value::*};
377        use std::cmp::Ordering;
378
379        fn compare_pair(l: &Value, r: &Value) -> Option<Ordering> {
380            match (l, r) {
381                // Prefer shared integer representations to avoid losing precision
382                // beyond 2^53; other numeric pairs fall back to f64.
383                (Number(l), Number(r)) => {
384                    if let (Some(l), Some(r)) = (l.as_i64(), r.as_i64()) {
385                        Some(l.cmp(&r))
386                    } else if let (Some(l), Some(r)) = (l.as_u64(), r.as_u64()) {
387                        Some(l.cmp(&r))
388                    } else {
389                        l.as_f64()
390                            .zip(r.as_f64())
391                            .and_then(|(l, r)| l.partial_cmp(&r))
392                    }
393                }
394                (String(l), String(r)) => Some(l.cmp(r)),
395                (Null, Null) => Some(Ordering::Equal),
396                (Bool(l), Bool(r)) => Some(l.cmp(r)),
397                _ => None,
398            }
399        }
400
401        match self {
402            // Numbers compare numerically so `5` matches `5.0`, consistent with
403            // `Gt`/`Lt`; other JSON types fall back to structural equality so
404            // strings/bools/arrays/objects still match exactly.
405            Eq(k, v) => value
406                .get(k)
407                .is_some_and(|field| compare_pair(field, v) == Some(Ordering::Equal) || field == v),
408            Gt(k, v) => value
409                .get(k)
410                .and_then(|field| compare_pair(field, v))
411                .is_some_and(|ord| ord == Ordering::Greater),
412            Lt(k, v) => value
413                .get(k)
414                .and_then(|field| compare_pair(field, v))
415                .is_some_and(|ord| ord == Ordering::Less),
416            And(l, r) => l.satisfies(value) && r.satisfies(value),
417            Or(l, r) => l.satisfies(value) || r.satisfies(value),
418        }
419    }
420}
421
422/// Builder for [`VectorSearchRequest`]. Requires `query` and `samples`.
423#[derive(Clone, Serialize, Deserialize, Debug)]
424pub struct VectorSearchRequestBuilder<F = Filter<serde_json::Value>, Q = Missing, S = Missing> {
425    query: Q,
426    samples: S,
427    threshold: Option<f64>,
428    additional_params: Option<serde_json::Value>,
429    filter: Option<F>,
430}
431
432impl<F> Default for VectorSearchRequestBuilder<F, Missing, Missing> {
433    fn default() -> Self {
434        Self {
435            query: Missing,
436            samples: Missing,
437            threshold: None,
438            additional_params: None,
439            filter: None,
440        }
441    }
442}
443
444impl<F, Q, S> VectorSearchRequestBuilder<F, Q, S>
445where
446    F: SearchFilter,
447{
448    /// Sets the query text. Required.
449    pub fn query<T>(self, query: T) -> VectorSearchRequestBuilder<F, Provided<String>, S>
450    where
451        T: Into<String>,
452    {
453        VectorSearchRequestBuilder {
454            query: Provided(query.into()),
455            samples: self.samples,
456            threshold: self.threshold,
457            additional_params: self.additional_params,
458            filter: self.filter,
459        }
460    }
461
462    /// Sets the maximum number of results. Required.
463    pub fn samples(self, samples: u64) -> VectorSearchRequestBuilder<F, Q, Provided<u64>> {
464        VectorSearchRequestBuilder {
465            query: self.query,
466            samples: Provided(samples),
467            threshold: self.threshold,
468            additional_params: self.additional_params,
469            filter: self.filter,
470        }
471    }
472
473    /// Sets the minimum similarity threshold.
474    pub fn threshold(mut self, threshold: f64) -> Self {
475        self.threshold = Some(threshold);
476        self
477    }
478
479    /// Replaces backend-specific parameters. Accepts any JSON value without validation.
480    pub fn additional_params(
481        mut self,
482        params: serde_json::Value,
483    ) -> Result<Self, VectorStoreError> {
484        self.additional_params = Some(params);
485        Ok(self)
486    }
487
488    /// Sets a filter expression.
489    pub fn filter(mut self, filter: F) -> Self {
490        self.filter = Some(filter);
491        self
492    }
493}
494
495impl<F> VectorSearchRequestBuilder<F, Provided<String>, Provided<u64>> {
496    /// Builds without validating query emptiness, sample count, or threshold.
497    pub fn build(self) -> VectorSearchRequest<F> {
498        VectorSearchRequest {
499            query: self.query.0,
500            samples: self.samples.0,
501            threshold: self.threshold,
502            additional_params: self.additional_params,
503            filter: self.filter,
504        }
505    }
506}
507
508#[cfg(test)]
509mod tests;