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