1use std::ops::Range;
13
14use serde::{Deserialize, Serialize};
15
16use crate::markers::{Missing, Provided};
17
18#[derive(Clone, Serialize, Deserialize, Debug)]
23pub struct VectorSearchRequest<F = Filter<serde_json::Value>> {
24 query: String,
26 samples: u64,
28 threshold: Option<f64>,
30 filter: Option<F>,
32}
33
34impl<Filter> VectorSearchRequest<Filter> {
35 pub fn builder() -> VectorSearchRequestBuilder<Filter> {
37 VectorSearchRequestBuilder::<Filter>::default()
38 }
39
40 pub fn query(&self) -> &str {
42 &self.query
43 }
44
45 pub fn samples(&self) -> u64 {
47 self.samples
48 }
49
50 pub fn threshold(&self) -> Option<f64> {
52 self.threshold
53 }
54
55 pub fn filter(&self) -> Option<&Filter> {
57 self.filter.as_ref()
58 }
59
60 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 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#[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 #[error("Filter serialization failed: {0}")]
111 Serialization(String),
112}
113
114pub 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#[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
145impl<P> Default for SqlCondition<P> {
148 fn default() -> Self {
149 Self::raw(String::new())
150 }
151}
152
153impl<P> SqlCondition<P> {
154 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 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 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 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 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 pub fn and(self, rhs: Self) -> Self {
213 self.combine("AND", rhs)
214 }
215
216 pub fn or(self, rhs: Self) -> Self {
218 self.combine("OR", rhs)
219 }
220
221 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 pub fn condition(&self) -> &str {
258 &self.condition
259 }
260
261 pub fn params(&self) -> &[P] {
263 &self.params
264 }
265
266 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 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 pub fn into_parts(self) -> (String, Vec<P>) {
298 (self.condition, self.params)
299 }
300}
301
302pub trait DynamicSearchFilter: SearchFilter + Sized {
310 fn from_dynamic_filter(filter: Filter<serde_json::Value>) -> Result<Self, FilterError>;
312
313 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#[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 fn eq(key: impl AsRef<str>, value: Self::Value) -> Self {
381 Self::Eq(key.as_ref().to_owned(), value)
382 }
383
384 fn gt(key: impl AsRef<str>, value: Self::Value) -> Self {
386 Self::Gt(key.as_ref().to_owned(), value)
387 }
388
389 fn lt(key: impl AsRef<str>, value: Self::Value) -> Self {
391 Self::Lt(key.as_ref().to_owned(), value)
392 }
393
394 fn and(self, rhs: Self) -> Self {
396 Self::And(self.into(), rhs.into())
397 }
398
399 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 pub fn interpret<F>(self) -> F
411 where
412 F: SearchFilter<Value = V>,
413 {
414 self.interpret_with(|v| v)
415 }
416
417 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 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 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 (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 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#[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 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 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 pub fn threshold(mut self, threshold: f64) -> Self {
549 self.threshold = Some(threshold);
550 self
551 }
552
553 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 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;