use std::ops::Range;
use serde::{Deserialize, Serialize};
use crate::markers::{Missing, Provided};
#[derive(Clone, Serialize, Deserialize, Debug)]
pub struct VectorSearchRequest<F = Filter<serde_json::Value>> {
query: String,
samples: u64,
threshold: Option<f64>,
filter: Option<F>,
}
impl<Filter> VectorSearchRequest<Filter> {
pub fn builder() -> VectorSearchRequestBuilder<Filter> {
VectorSearchRequestBuilder::<Filter>::default()
}
pub fn query(&self) -> &str {
&self.query
}
pub fn samples(&self) -> u64 {
self.samples
}
pub fn threshold(&self) -> Option<f64> {
self.threshold
}
pub fn filter(&self) -> Option<&Filter> {
self.filter.as_ref()
}
pub fn map_filter<T, F>(self, f: F) -> VectorSearchRequest<T>
where
F: Fn(Filter) -> T,
{
VectorSearchRequest {
query: self.query,
samples: self.samples,
threshold: self.threshold,
filter: self.filter.map(f),
}
}
pub fn try_map_filter<T, F>(self, f: F) -> Result<VectorSearchRequest<T>, FilterError>
where
F: Fn(Filter) -> Result<T, FilterError>,
{
let filter = self.filter.map(f).transpose()?;
Ok(VectorSearchRequest {
query: self.query,
samples: self.samples,
threshold: self.threshold,
filter,
})
}
}
#[derive(Debug, Clone, thiserror::Error)]
pub enum FilterError {
#[error("Expected: {expected}, got: {got}")]
Expected { expected: String, got: String },
#[error("Cannot compile '{0}' to the backend's filter type")]
TypeError(String),
#[error("Missing field '{0}'")]
MissingField(String),
#[error("'{0}' must {1}")]
Must(String, String),
#[error("Filter serialization failed: {0}")]
Serialization(String),
}
pub trait SearchFilter {
type Value;
fn eq(key: impl AsRef<str>, value: Self::Value) -> Self;
fn gt(key: impl AsRef<str>, value: Self::Value) -> Self;
fn lt(key: impl AsRef<str>, value: Self::Value) -> Self;
fn and(self, rhs: Self) -> Self;
fn or(self, rhs: Self) -> Self;
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct SqlCondition<P> {
condition: String,
params: Vec<P>,
#[serde(default)]
placeholders: Vec<Range<usize>>,
}
impl<P> Default for SqlCondition<P> {
fn default() -> Self {
Self::raw(String::new())
}
}
impl<P> SqlCondition<P> {
pub fn binary(key: impl AsRef<str>, op: &str, placeholder: &str, value: P) -> Self {
let mut this = Self::raw(format!("{} {op} ", key.as_ref()));
this.push_placeholder(placeholder);
this.params.push(value);
this
}
pub fn list(key: impl AsRef<str>, op: &str, placeholder: &str, values: Vec<P>) -> Self {
let mut this = Self::raw(format!("{} {op} (", key.as_ref()));
for i in 0..values.len() {
if i > 0 {
this.condition.push_str(", ");
}
this.push_placeholder(placeholder);
}
this.condition.push(')');
this.params = values;
this
}
pub fn range(key: impl AsRef<str>, placeholder: &str, lo: P, hi: P) -> Self {
let key = key.as_ref();
let mut this = Self::raw(format!("{key} >= "));
this.push_placeholder(placeholder);
this.condition.push_str(&format!(" AND {key} <= "));
this.push_placeholder(placeholder);
this.params = vec![lo, hi];
this
}
pub fn between(key: impl AsRef<str>, placeholder: &str, lo: P, hi: P) -> Self {
let mut this = Self::raw(format!("{} between ", key.as_ref()));
this.push_placeholder(placeholder);
this.condition.push_str(" and ");
this.push_placeholder(placeholder);
this.params = vec![lo, hi];
this
}
pub fn raw(condition: impl Into<String>) -> Self {
Self {
condition: condition.into(),
params: Vec::new(),
placeholders: Vec::new(),
}
}
pub fn and(self, rhs: Self) -> Self {
self.combine("AND", rhs)
}
pub fn or(self, rhs: Self) -> Self {
self.combine("OR", rhs)
}
pub fn not(self) -> Self {
let mut this = Self::raw("NOT (");
this.append(self);
this.condition.push(')');
this
}
fn combine(self, joiner: &str, rhs: Self) -> Self {
let mut this = Self::raw("(");
this.append(self);
this.condition.push_str(&format!(") {joiner} ("));
this.append(rhs);
this.condition.push(')');
this
}
fn push_placeholder(&mut self, placeholder: &str) {
let start = self.condition.len();
self.condition.push_str(placeholder);
self.placeholders.push(start..self.condition.len());
}
fn append(&mut self, other: Self) {
let offset = self.condition.len();
self.condition.push_str(&other.condition);
self.params.extend(other.params);
self.placeholders.extend(
other
.placeholders
.into_iter()
.map(|range| range.start + offset..range.end + offset),
);
}
pub fn condition(&self) -> &str {
&self.condition
}
pub fn params(&self) -> &[P] {
&self.params
}
pub fn render_placeholders<D: std::fmt::Display>(
&self,
mut placeholder: impl FnMut(usize) -> D,
) -> String {
use std::fmt::Write as _;
let mut out = String::with_capacity(self.condition.len() + 2 * self.placeholders.len());
let mut copied = 0;
for (i, range) in self.placeholders.iter().enumerate() {
let (Some(before), Some(_)) = (
self.condition.get(copied..range.start),
self.condition.get(range.clone()),
) else {
continue;
};
out.push_str(before);
let _ = write!(out, "{}", placeholder(i));
copied = range.end;
}
out.push_str(self.condition.get(copied..).unwrap_or_default());
out
}
pub fn into_parts(self) -> (String, Vec<P>) {
(self.condition, self.params)
}
}
pub trait DynamicSearchFilter: SearchFilter + Sized {
fn from_dynamic_filter(filter: Filter<serde_json::Value>) -> Result<Self, FilterError>;
fn normalize_dynamic_document(document: serde_json::Value) -> serde_json::Value {
document
}
}
impl<F> DynamicSearchFilter for F
where
F: SearchFilter<Value = serde_json::Value>,
{
fn from_dynamic_filter(filter: Filter<serde_json::Value>) -> Result<Self, FilterError> {
Ok(filter.interpret())
}
fn normalize_dynamic_document(document: serde_json::Value) -> serde_json::Value {
prune_document(document).unwrap_or_default()
}
}
fn prune_document(document: serde_json::Value) -> Option<serde_json::Value> {
match document {
serde_json::Value::Object(mut map) => {
let new_map = map
.iter_mut()
.filter_map(|(key, value)| {
prune_document(value.take()).map(|value| (key.clone(), value))
})
.collect::<serde_json::Map<_, _>>();
Some(serde_json::Value::Object(new_map))
}
serde_json::Value::Array(vec) if vec.len() > 400 => None,
serde_json::Value::Array(vec) => Some(serde_json::Value::Array(
vec.into_iter().filter_map(prune_document).collect(),
)),
value => Some(value),
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum Filter<V>
where
V: std::fmt::Debug + Clone,
{
Eq(String, V),
Gt(String, V),
Lt(String, V),
And(Box<Self>, Box<Self>),
Or(Box<Self>, Box<Self>),
}
impl<V> SearchFilter for Filter<V>
where
V: std::fmt::Debug + Clone + Serialize + serde::de::DeserializeOwned,
{
type Value = V;
fn eq(key: impl AsRef<str>, value: Self::Value) -> Self {
Self::Eq(key.as_ref().to_owned(), value)
}
fn gt(key: impl AsRef<str>, value: Self::Value) -> Self {
Self::Gt(key.as_ref().to_owned(), value)
}
fn lt(key: impl AsRef<str>, value: Self::Value) -> Self {
Self::Lt(key.as_ref().to_owned(), value)
}
fn and(self, rhs: Self) -> Self {
Self::And(self.into(), rhs.into())
}
fn or(self, rhs: Self) -> Self {
Self::Or(self.into(), rhs.into())
}
}
impl<V> Filter<V>
where
V: std::fmt::Debug + Clone,
{
pub fn interpret<F>(self) -> F
where
F: SearchFilter<Value = V>,
{
self.interpret_with(|v| v)
}
pub fn interpret_with<F, W>(self, conv: impl Fn(V) -> W + Copy) -> F
where
F: SearchFilter<Value = W>,
{
match self.try_interpret(|v| Ok::<W, std::convert::Infallible>(conv(v))) {
Ok(filter) => filter,
Err(never) => match never {},
}
}
pub fn try_interpret<F, W, E>(self, conv: impl Fn(V) -> Result<W, E> + Copy) -> Result<F, E>
where
F: SearchFilter<Value = W>,
{
Ok(match self {
Self::Eq(key, val) => F::eq(key, conv(val)?),
Self::Gt(key, val) => F::gt(key, conv(val)?),
Self::Lt(key, val) => F::lt(key, conv(val)?),
Self::And(lhs, rhs) => F::and(lhs.try_interpret(conv)?, rhs.try_interpret(conv)?),
Self::Or(lhs, rhs) => F::or(lhs.try_interpret(conv)?, rhs.try_interpret(conv)?),
})
}
}
impl Filter<serde_json::Value> {
pub fn satisfies(&self, value: &serde_json::Value) -> bool {
use Filter::*;
use serde_json::{Value, Value::*};
use std::cmp::Ordering;
fn compare_pair(l: &Value, r: &Value) -> Option<Ordering> {
match (l, r) {
(Number(l), Number(r)) => {
if let (Some(l), Some(r)) = (l.as_i64(), r.as_i64()) {
Some(l.cmp(&r))
} else if let (Some(l), Some(r)) = (l.as_u64(), r.as_u64()) {
Some(l.cmp(&r))
} else {
l.as_f64()
.zip(r.as_f64())
.and_then(|(l, r)| l.partial_cmp(&r))
}
}
(String(l), String(r)) => Some(l.cmp(r)),
(Null, Null) => Some(Ordering::Equal),
(Bool(l), Bool(r)) => Some(l.cmp(r)),
_ => None,
}
}
match self {
Eq(k, v) => value
.get(k)
.is_some_and(|field| compare_pair(field, v) == Some(Ordering::Equal) || field == v),
Gt(k, v) => value
.get(k)
.and_then(|field| compare_pair(field, v))
.is_some_and(|ord| ord == Ordering::Greater),
Lt(k, v) => value
.get(k)
.and_then(|field| compare_pair(field, v))
.is_some_and(|ord| ord == Ordering::Less),
And(l, r) => l.satisfies(value) && r.satisfies(value),
Or(l, r) => l.satisfies(value) || r.satisfies(value),
}
}
}
#[derive(Clone, Serialize, Deserialize, Debug)]
pub struct VectorSearchRequestBuilder<F = Filter<serde_json::Value>, Q = Missing, S = Missing> {
query: Q,
samples: S,
threshold: Option<f64>,
filter: Option<F>,
}
impl<F> Default for VectorSearchRequestBuilder<F, Missing, Missing> {
fn default() -> Self {
Self {
query: Missing,
samples: Missing,
threshold: None,
filter: None,
}
}
}
impl<F, Q, S> VectorSearchRequestBuilder<F, Q, S>
where
F: SearchFilter,
{
pub fn query<T>(self, query: T) -> VectorSearchRequestBuilder<F, Provided<String>, S>
where
T: Into<String>,
{
VectorSearchRequestBuilder {
query: Provided(query.into()),
samples: self.samples,
threshold: self.threshold,
filter: self.filter,
}
}
pub fn samples(self, samples: u64) -> VectorSearchRequestBuilder<F, Q, Provided<u64>> {
VectorSearchRequestBuilder {
query: self.query,
samples: Provided(samples),
threshold: self.threshold,
filter: self.filter,
}
}
pub fn threshold(mut self, threshold: f64) -> Self {
self.threshold = Some(threshold);
self
}
pub fn filter(mut self, filter: F) -> Self {
self.filter = Some(filter);
self
}
}
impl<F> VectorSearchRequestBuilder<F, Provided<String>, Provided<u64>> {
pub fn build(self) -> VectorSearchRequest<F> {
VectorSearchRequest {
query: self.query.0,
samples: self.samples.0,
threshold: self.threshold,
filter: self.filter,
}
}
}
#[cfg(test)]
mod tests;