1use serde::{Deserialize, Serialize};
13
14use super::VectorStoreError;
15use crate::markers::{Missing, Provided};
16
17#[derive(Clone, Serialize, Deserialize, Debug)]
22pub struct VectorSearchRequest<F = Filter<serde_json::Value>> {
23 query: String,
25 samples: u64,
27 threshold: Option<f64>,
29 additional_params: Option<serde_json::Value>,
31 filter: Option<F>,
33}
34
35impl<Filter> VectorSearchRequest<Filter> {
36 pub fn builder() -> VectorSearchRequestBuilder<Filter> {
38 VectorSearchRequestBuilder::<Filter>::default()
39 }
40
41 pub fn query(&self) -> &str {
43 &self.query
44 }
45
46 pub fn samples(&self) -> u64 {
48 self.samples
49 }
50
51 pub fn threshold(&self) -> Option<f64> {
53 self.threshold
54 }
55
56 pub fn filter(&self) -> &Option<Filter> {
58 &self.filter
59 }
60
61 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 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#[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 #[error("Filter serialization failed: {0}")]
114 Serialization(String),
115}
116
117pub 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#[derive(Clone, Debug, Serialize, Deserialize)]
138pub struct SqlCondition<P> {
139 condition: String,
140 params: Vec<P>,
141}
142
143impl<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 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 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 pub fn raw(condition: impl Into<String>) -> Self {
177 Self {
178 condition: condition.into(),
179 params: Vec::new(),
180 }
181 }
182
183 pub fn and(self, rhs: Self) -> Self {
185 self.combine("AND", rhs)
186 }
187
188 pub fn or(self, rhs: Self) -> Self {
190 self.combine("OR", rhs)
191 }
192
193 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 pub fn condition(&self) -> &str {
210 &self.condition
211 }
212
213 pub fn params(&self) -> &[P] {
215 &self.params
216 }
217
218 pub fn into_parts(self) -> (String, Vec<P>) {
220 (self.condition, self.params)
221 }
222}
223
224pub trait DynamicSearchFilter: SearchFilter + Sized {
232 fn from_dynamic_filter(filter: Filter<serde_json::Value>) -> Result<Self, FilterError>;
234
235 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#[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 fn eq(key: impl AsRef<str>, value: Self::Value) -> Self {
303 Self::Eq(key.as_ref().to_owned(), value)
304 }
305
306 fn gt(key: impl AsRef<str>, value: Self::Value) -> Self {
308 Self::Gt(key.as_ref().to_owned(), value)
309 }
310
311 fn lt(key: impl AsRef<str>, value: Self::Value) -> Self {
313 Self::Lt(key.as_ref().to_owned(), value)
314 }
315
316 fn and(self, rhs: Self) -> Self {
318 Self::And(self.into(), rhs.into())
319 }
320
321 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 pub fn interpret<F>(self) -> F
333 where
334 F: SearchFilter<Value = V>,
335 {
336 self.interpret_with(|v| v)
337 }
338
339 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 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 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 (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 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#[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 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 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 pub fn threshold(mut self, threshold: f64) -> Self {
475 self.threshold = Some(threshold);
476 self
477 }
478
479 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 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 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;