use std::collections::HashSet;
use std::ops::Bound;
use super::literals::{try_expr_to_value, try_literal_to_value};
use crate::catalog::Distance;
use crate::exec::index::analysis::idiom_matches_containment;
use crate::expr::operator::{MatchesOperator, NearestNeighbor};
use crate::expr::visit::{MutVisitor, Visit, VisitMut, Visitor};
use crate::expr::{BinaryOperator, Cond, Expr, Idiom, Literal};
use crate::val::{Number, Value};
pub(crate) fn has_top_level_or(cond: Option<&Cond>) -> bool {
match cond {
Some(c) => matches!(
c.0,
Expr::Binary {
op: BinaryOperator::Or,
..
}
),
None => false,
}
}
pub(crate) fn has_knn_operator(expr: &Expr) -> bool {
scan_knn_operators(expr).found_any
}
pub(crate) fn has_knn_k_operator(expr: &Expr) -> bool {
scan_knn_operators(expr).found_k
}
pub(crate) fn has_knn_ktree_operator(expr: &Expr) -> bool {
scan_knn_operators(expr).found_ktree
}
fn scan_knn_operators(expr: &Expr) -> KnnOperatorChecker {
let mut checker = KnnOperatorChecker {
found_any: false,
found_k: false,
found_ktree: false,
};
let _ = checker.visit_expr(expr);
checker
}
struct KnnOperatorChecker {
found_any: bool,
found_k: bool,
found_ktree: bool,
}
impl Visitor for KnnOperatorChecker {
type Error = std::convert::Infallible;
fn visit_expr(&mut self, expr: &Expr) -> Result<(), Self::Error> {
if let Expr::Binary {
op: BinaryOperator::NearestNeighbor(nn),
..
} = expr
{
self.found_any = true;
match nn.as_ref() {
NearestNeighbor::K(..) => self.found_k = true,
NearestNeighbor::KTree(..) => self.found_ktree = true,
NearestNeighbor::Approximate(..) => {}
}
}
expr.visit(self)
}
fn visit_select(&mut self, _: &crate::expr::SelectStatement) -> Result<(), Self::Error> {
Ok(())
}
}
pub(crate) struct BruteForceKnnParams {
pub field: Idiom,
pub vector: BruteForceKnnVector,
pub k: u32,
pub distance: Distance,
}
pub(crate) enum BruteForceKnnVector {
Literal(Vec<Number>),
Deferred(Expr),
}
fn is_plan_time_computable(expr: &Expr) -> bool {
match expr {
Expr::Literal(Literal::Array(_)) => true,
Expr::Literal(
Literal::Integer(_)
| Literal::Bool(_)
| Literal::String(_)
| Literal::RecordId(_)
| Literal::Duration(_)
| Literal::Uuid(_)
| Literal::Datetime(_)
| Literal::None
| Literal::Null
| Literal::Decimal(_)
| Literal::Float(_),
) => true,
Expr::Param(_) | Expr::FunctionCall(_) => true,
Expr::Binary {
left,
right,
..
} => is_plan_time_computable(left) && is_plan_time_computable(right),
_ => false,
}
}
pub(crate) fn extract_bruteforce_knn(cond: &Cond) -> Option<BruteForceKnnParams> {
let mut expr = cond.0.clone();
let mut extractor = BruteForceKnnExtractor {
params: None,
};
let _ = extractor.visit_mut_expr(&mut expr);
extractor.params
}
pub(crate) fn strip_fts_condition(
cond: &Cond,
column: &Idiom,
operator: &MatchesOperator,
query: &str,
) -> Option<Cond> {
strip_and_simplify(cond.0.clone(), |e| {
let Expr::Binary {
left,
op: BinaryOperator::Matches(op),
right,
} = e
else {
return false;
};
op == operator
&& matches!(left.as_ref(), Expr::Idiom(idiom) if idiom == column)
&& matches!(
right.as_ref(),
Expr::Literal(lit)
if matches!(try_literal_to_value(lit), Some(Value::String(s)) if s.as_str() == query)
)
})
.map(Cond)
}
pub(crate) fn strip_knn_from_condition(cond: &Cond) -> Option<Cond> {
strip_and_simplify(cond.0.clone(), |e| {
matches!(
e,
Expr::Binary {
op: BinaryOperator::NearestNeighbor(nn),
..
} if matches!(nn.as_ref(), NearestNeighbor::K(..) | NearestNeighbor::Approximate(..))
)
})
.map(Cond)
}
pub(crate) fn strip_knn_and_matches_from_condition(cond: &Cond) -> Option<Cond> {
use crate::exec::index::analysis::IndexAnalyzer;
strip_and_simplify(cond.0.clone(), |e| {
matches!(
e,
Expr::Binary {
op: BinaryOperator::NearestNeighbor(nn),
..
} if matches!(nn.as_ref(), NearestNeighbor::K(..) | NearestNeighbor::Approximate(..))
) || IndexAnalyzer::expr_contains_matches(e)
})
.map(Cond)
}
fn strip_and_simplify<F>(mut expr: Expr, mut should_strip: F) -> Option<Expr>
where
F: FnMut(&Expr) -> bool,
{
fn walk<F: FnMut(&Expr) -> bool>(expr: &mut Expr, f: &mut F) {
if let Expr::Binary {
left,
op: BinaryOperator::And,
right,
} = expr
{
walk(left, f);
walk(right, f);
} else if f(expr) {
*expr = Expr::Literal(Literal::Bool(true));
}
}
walk(&mut expr, &mut should_strip);
let _ = BoolSimplifier.visit_mut_expr(&mut expr);
if matches!(expr, Expr::Literal(Literal::Bool(true))) {
None
} else {
Some(expr)
}
}
pub(crate) fn strip_index_conditions(
cond: &Cond,
access: &crate::exec::index::access_path::BTreeAccess,
cols: &[Idiom],
) -> Option<Cond> {
let matcher = IndexConditionMatcher {
cols,
access,
};
strip_and_simplify(cond.0.clone(), |e| match e {
Expr::Binary {
left,
op,
right,
} => matcher.matches_access(left, op, right),
_ => false,
})
.map(Cond)
}
pub(crate) fn strip_union_index_conditions(
cond: &Cond,
paths: &[crate::exec::index::access_path::AccessPath],
) -> Option<Cond> {
use crate::exec::index::access_path::{AccessPath, BTreeAccess};
let mut first_col: Option<&Idiom> = None;
let mut branch_values: HashSet<&Value> = HashSet::with_capacity(paths.len());
for path in paths {
let AccessPath::BTreeScan {
index_ref,
access,
..
} = path
else {
return Some(cond.clone());
};
let col = index_ref.definition().cols.first()?;
match first_col {
None => first_col = Some(col),
Some(prev) if prev == col => {}
Some(_) => return Some(cond.clone()),
}
match access {
BTreeAccess::Equality(v) => {
branch_values.insert(v);
}
BTreeAccess::Compound {
prefix,
range: None,
} if prefix.len() == 1 => {
branch_values.insert(&prefix[0]);
}
_ => return Some(cond.clone()),
}
}
let Some(col) = first_col else {
return Some(cond.clone());
};
strip_and_simplify(cond.0.clone(), |e| match e {
Expr::Binary {
left,
op,
right,
} => union_covers_leaf(col, &branch_values, left, op, right),
_ => false,
})
.map(Cond)
}
fn union_covers_leaf(
col: &Idiom,
branch_values: &HashSet<&crate::val::Value>,
left: &Expr,
op: &BinaryOperator,
right: &Expr,
) -> bool {
use crate::exec::index::analysis::idiom_matches_containment;
let (idiom, lit) = match op {
BinaryOperator::ContainAny => match (left, right) {
(Expr::Idiom(i), Expr::Literal(l)) => (i, l),
_ => return false,
},
BinaryOperator::AnyInside => match (left, right) {
(Expr::Literal(l), Expr::Idiom(i)) => (i, l),
_ => return false,
},
_ => return false,
};
if !idiom_matches_containment(idiom, col) {
return false;
}
let Some(Value::Array(arr)) = try_literal_to_value(lit) else {
return false;
};
if branch_values.is_empty() {
return false;
}
let leaf_values: HashSet<&Value> = arr.0.iter().collect();
branch_values.iter().all(|v| !v.is_nullish() && leaf_values.contains(*v))
}
struct IndexConditionMatcher<'a> {
cols: &'a [Idiom],
access: &'a crate::exec::index::access_path::BTreeAccess,
}
impl IndexConditionMatcher<'_> {
fn matches_access(&self, left: &Expr, op: &BinaryOperator, right: &Expr) -> bool {
use crate::exec::index::access_path::BTreeAccess;
let (idiom, value, effective_op, idiom_on_left) = match (left, right) {
(Expr::Idiom(i), rhs) => {
if let Some(v) = try_expr_to_value(rhs) {
(i, v, op.clone(), true)
} else {
return false;
}
}
(lhs, Expr::Idiom(i)) => {
if let Some(v) = try_expr_to_value(lhs) {
let flipped = match op {
BinaryOperator::LessThan => BinaryOperator::MoreThan,
BinaryOperator::LessThanEqual => BinaryOperator::MoreThanEqual,
BinaryOperator::MoreThan => BinaryOperator::LessThan,
BinaryOperator::MoreThanEqual => BinaryOperator::LessThanEqual,
other => other.clone(),
};
(i, v, flipped, false)
} else {
return false;
}
}
_ => return false,
};
let (effective_op, value) = if idiom_on_left
&& matches!(effective_op, BinaryOperator::Inside)
&& let Value::Array(arr) = &value
&& arr.len() == 1
{
(BinaryOperator::Equal, arr[0].clone())
} else {
(effective_op, value)
};
let is_equality =
matches!(effective_op, BinaryOperator::Equal | BinaryOperator::ExactEqual);
let is_containment = ((matches!(effective_op, BinaryOperator::Contain) && idiom_on_left)
|| (matches!(effective_op, BinaryOperator::Inside) && !idiom_on_left))
&& !value.is_nullish();
match self.access {
BTreeAccess::Compound {
prefix,
range,
} => {
if is_equality {
for (col, val) in self.cols.iter().zip(prefix.iter()) {
if idiom == col && value == *val {
return true;
}
}
}
if is_containment {
for (col, val) in self.cols.iter().zip(prefix.iter()) {
if idiom_matches_containment(idiom, col) && value == *val {
return true;
}
}
}
if let Some((range_op, range_val)) = range
&& let Some(col) = self.cols.get(prefix.len())
&& idiom == col
{
if effective_op == *range_op && value == *range_val {
return true;
}
if matches!(range_op, BinaryOperator::MoreThan)
&& matches!(range_val, Value::None)
&& effective_op == BinaryOperator::NotEqual
&& matches!(value, Value::None)
{
return true;
}
}
false
}
BTreeAccess::Equality(val) => {
let Some(col) = self.cols.first() else {
return false;
};
if is_equality && idiom == col && value == *val {
return true;
}
if is_containment && idiom_matches_containment(idiom, col) && value == *val {
return true;
}
false
}
BTreeAccess::Range {
range,
} => {
let Some(col) = self.cols.first() else {
return false;
};
if idiom != col {
return false;
}
if matches!(effective_op, BinaryOperator::Inside)
&& idiom_on_left
&& let Value::Range(r) = &value
&& **r == *range
{
return true;
}
match &range.start {
Bound::Included(x) => {
if effective_op == BinaryOperator::MoreThanEqual && value == *x {
return true;
}
}
Bound::Excluded(x) => {
if effective_op == BinaryOperator::MoreThan && value == *x {
return true;
}
if matches!(x, Value::None)
&& matches!(value, Value::None)
&& effective_op == BinaryOperator::NotEqual
{
return true;
}
}
Bound::Unbounded => {}
}
match &range.end {
Bound::Included(x) => {
if effective_op == BinaryOperator::LessThanEqual && value == *x {
return true;
}
}
Bound::Excluded(x) => {
if effective_op == BinaryOperator::LessThan && value == *x {
return true;
}
}
Bound::Unbounded => {}
}
false
}
_ => false,
}
}
}
struct BruteForceKnnExtractor {
params: Option<BruteForceKnnParams>,
}
impl MutVisitor for BruteForceKnnExtractor {
type Error = std::convert::Infallible;
fn visit_mut_expr(&mut self, expr: &mut Expr) -> Result<(), Self::Error> {
if self.params.is_some() {
return Ok(());
}
if let Expr::Binary {
left,
op: BinaryOperator::NearestNeighbor(nn),
right,
} = expr && let NearestNeighbor::K(k, dist) = nn.as_ref()
&& let Expr::Idiom(idiom) = left.as_ref()
{
let vector = if let Some(vector) = extract_literal_vector(right) {
Some(BruteForceKnnVector::Literal(vector))
} else if is_plan_time_computable(right) {
Some(BruteForceKnnVector::Deferred(right.as_ref().clone()))
} else {
None
};
if let Some(vector) = vector {
self.params = Some(BruteForceKnnParams {
field: idiom.clone(),
vector,
k: *k,
distance: dist.clone().into(),
});
*expr = Expr::Literal(Literal::Bool(true));
return Ok(());
}
}
expr.visit_mut(self)
}
fn visit_mut_select(
&mut self,
_: &mut crate::expr::SelectStatement,
) -> Result<(), Self::Error> {
Ok(())
}
}
struct BoolSimplifier;
impl MutVisitor for BoolSimplifier {
type Error = std::convert::Infallible;
fn visit_mut_expr(&mut self, expr: &mut Expr) -> Result<(), Self::Error> {
expr.visit_mut(self)?;
if let Expr::Binary {
left,
op: BinaryOperator::And,
right,
} = expr
{
let l_true = matches!(left.as_ref(), Expr::Literal(Literal::Bool(true)));
let r_true = matches!(right.as_ref(), Expr::Literal(Literal::Bool(true)));
match (l_true, r_true) {
(true, true) => *expr = Expr::Literal(Literal::Bool(true)),
(true, false) => {
let r = std::mem::replace(right.as_mut(), Expr::Literal(Literal::None));
*expr = r;
}
(false, true) => {
let l = std::mem::replace(left.as_mut(), Expr::Literal(Literal::None));
*expr = l;
}
_ => {}
}
}
Ok(())
}
fn visit_mut_select(
&mut self,
_: &mut crate::expr::SelectStatement,
) -> Result<(), Self::Error> {
Ok(())
}
}
fn extract_literal_vector(expr: &Expr) -> Option<Vec<Number>> {
match expr {
Expr::Literal(lit) => {
if let Literal::Array(arr) = lit {
let mut nums = Vec::with_capacity(arr.len());
for elem in arr.iter() {
match elem {
Expr::Literal(Literal::Integer(i)) => {
nums.push(Number::Int(*i));
}
Expr::Literal(Literal::Float(f)) => {
nums.push(Number::Float(*f));
}
Expr::Literal(Literal::Decimal(d)) => {
nums.push(Number::Decimal(*d));
}
_ => return None,
}
}
Some(nums)
} else {
None
}
}
_ => None,
}
}
pub(crate) fn extract_record_id_point_lookup(
cond: &Cond,
table_name: &surrealdb_strand::TableName,
) -> Option<Expr> {
find_id_equality_in_and_chain(&cond.0, table_name)
}
fn find_id_equality_in_and_chain(
expr: &Expr,
table_name: &surrealdb_strand::TableName,
) -> Option<Expr> {
match expr {
Expr::Binary {
left,
op: BinaryOperator::And,
right,
} => find_id_equality_in_and_chain(left, table_name)
.or_else(|| find_id_equality_in_and_chain(right, table_name)),
Expr::Binary {
left,
op: BinaryOperator::Equal | BinaryOperator::ExactEqual,
right,
} => check_id_recordid_pair(left, right, table_name)
.or_else(|| check_id_recordid_pair(right, left, table_name)),
_ => None,
}
}
fn check_id_recordid_pair(
idiom_side: &Expr,
lit_side: &Expr,
table_name: &surrealdb_strand::TableName,
) -> Option<Expr> {
if let Expr::Idiom(idiom) = idiom_side
&& idiom.is_id()
&& let Expr::Literal(Literal::RecordId(rid)) = lit_side
&& &rid.table == table_name
&& !matches!(rid.key, crate::expr::RecordIdKeyLit::Range(_))
{
Some(lit_side.clone())
} else {
None
}
}
pub(crate) fn is_value_source_expr(expr: &Expr) -> bool {
match expr {
Expr::Literal(Literal::Array(_)) => true,
Expr::Literal(Literal::String(_))
| Expr::Literal(Literal::Integer(_))
| Expr::Literal(Literal::Float(_))
| Expr::Literal(Literal::Decimal(_))
| Expr::Literal(Literal::Bool(_))
| Expr::Literal(Literal::None)
| Expr::Literal(Literal::Null) => true,
Expr::Table(_) => false,
Expr::Literal(Literal::RecordId(_)) => false,
Expr::Param(_) => false,
Expr::Select(_) => false,
_ => false,
}
}
pub(crate) fn all_value_sources(sources: &[Expr]) -> bool {
!sources.is_empty() && sources.iter().all(is_value_source_expr)
}
pub(crate) fn extract_matches_context(
cond: &Cond,
ctx: Option<&crate::ctx::FrozenContext>,
) -> crate::exec::function::MatchesContext {
let mut collector = MatchesCollector(crate::exec::function::MatchesContext::new(), ctx);
let _ = collector.visit_expr(&cond.0);
collector.0
}
struct MatchesCollector<'a>(
crate::exec::function::MatchesContext,
Option<&'a crate::ctx::FrozenContext>,
);
impl Visitor for MatchesCollector<'_> {
type Error = std::convert::Infallible;
fn visit_expr(&mut self, expr: &Expr) -> Result<(), Self::Error> {
if let Expr::Binary {
left,
op: BinaryOperator::Matches(matches_op),
right,
} = expr && let Expr::Idiom(idiom) = left.as_ref()
{
let query_str = match right.as_ref() {
Expr::Literal(Literal::String(s)) => Some(s.as_str().to_owned()),
Expr::Param(param) => {
self.1.and_then(|ctx| {
ctx.value(param.as_str()).and_then(|v| {
if let crate::val::Value::String(s) = v {
Some(s.as_str().to_owned())
} else {
None
}
})
})
}
_ => None,
};
if let Some(query) = query_str {
let match_ref = matches_op.rf.unwrap_or(0);
self.0.insert(
match_ref,
crate::exec::function::MatchInfo {
idiom: idiom.clone(),
query,
},
);
}
}
expr.visit(self)
}
fn visit_select(&mut self, _: &crate::expr::SelectStatement) -> Result<(), Self::Error> {
Ok(())
}
}
pub(crate) fn extract_table_from_matches(
matches_context: &crate::exec::function::MatchesContext,
) -> surrealdb_strand::TableName {
if let Some(table) = matches_context.table() {
return table.clone();
}
surrealdb_strand::TableName::from("unknown".to_string())
}
#[cfg(test)]
mod tests {
use std::str::FromStr;
use std::sync::Arc;
use surrealdb_strand::Strand;
use surrealdb_types::ToSql;
use super::*;
use crate::catalog::{Index, IndexDefinition, IndexId};
use crate::exec::index::access_path::{AccessPath, BTreeAccess, IndexRef};
use crate::kvs::Direction;
#[test]
fn fts_stripper_consumes_only_the_scans_own_probe() {
let title = Idiom::from_str("title").expect("valid idiom");
let op = |sql: &str| match parse_cond(sql).0 {
Expr::Binary {
op: BinaryOperator::Matches(op),
..
} => op,
other => panic!("expected a MATCHES, got {}", other.to_sql()),
};
let (any, first) = (op("title @@ 'x'"), op("title @1@ 'x'"));
let strip = |sql: &str, op: &MatchesOperator| {
strip_fts_condition(&parse_cond(sql), &title, op, "x").map(|c| c.0.to_sql())
};
assert_eq!(strip("title @@ 'x' AND a = 1", &any).as_deref(), Some("a = 1"));
assert_eq!(strip("title @@ 'x'", &any), None);
assert_eq!(
strip("title @1@ 'x' AND body @2@ 'y'", &first).as_deref(),
Some("body @2@ 'y'")
);
assert_eq!(
strip("title @1@ 'x' AND title @2@ 'x'", &first).as_deref(),
Some("title @2@ 'x'")
);
assert_eq!(
strip("title @1@ 'x' AND title @1@ $q", &first).as_deref(),
Some("title @1@ $q")
);
}
fn parse_cond(snippet: &str) -> Cond {
let src = format!("SELECT * FROM t WHERE {snippet}");
let mut exprs = crate::syn::parse(&src).expect("parse").expressions;
assert_eq!(exprs.len(), 1, "expected one statement from {src:?}");
match exprs.remove(0).into() {
crate::expr::TopLevelExpr::Expr(Expr::Select(s)) => s.cond.expect("WHERE"),
other => panic!("expected SELECT, got {other:?}"),
}
}
fn union_paths(col: &str, values: &[&str]) -> Vec<AccessPath> {
let def = IndexDefinition {
index_id: IndexId(1),
name: Strand::from("ix"),
table_name: "t".into(),
cols: vec![Idiom::from_str(col).expect("valid idiom")],
index: Index::Idx,
count_cond: None,
comment: None,
prepare_remove: false,
format_version: 1,
};
let indexes: Arc<[IndexDefinition]> = Arc::from(vec![def].into_boxed_slice());
values
.iter()
.map(|v| AccessPath::BTreeScan {
index_ref: IndexRef::new(Arc::clone(&indexes), 0),
access: BTreeAccess::Equality(Value::from(*v)),
direction: Direction::Forward,
})
.collect()
}
fn residual(cond: &str, col: &str, branches: &[&str]) -> Option<String> {
strip_union_index_conditions(&parse_cond(cond), &union_paths(col, branches))
.map(|c| c.0.to_sql())
}
#[test]
fn strips_leaf_matching_the_branch_values() {
assert_eq!(residual("tags CONTAINSANY ['a', 'b']", "tags.*", &["a", "b"]), None);
assert_eq!(residual("['a', 'b'] ANYINSIDE tags", "tags.*", &["a", "b"]), None);
}
#[test]
fn strips_leaf_wider_than_the_branch_values() {
assert_eq!(residual("tags CONTAINSANY ['a', 'b', 'c']", "tags.*", &["a", "b"]), None);
}
#[test]
fn keeps_leaf_narrower_than_the_branch_values() {
assert_eq!(
residual(
"tags CONTAINSANY ['a'] AND (tags CONTAINS 'a' OR tags CONTAINS 'b')",
"tags.*",
&["a", "b"],
)
.as_deref(),
Some("tags CONTAINSANY ['a'] AND (tags CONTAINS 'a' OR tags CONTAINS 'b')"),
);
assert_eq!(
residual(
"tags CONTAINSANY ['a', 'b'] AND tags CONTAINSANY ['a']",
"tags.*",
&["a", "b"]
)
.as_deref(),
Some("tags CONTAINSANY ['a']"),
);
}
#[test]
fn keeps_leaf_disjoint_from_the_branch_values() {
assert_eq!(
residual("tags CONTAINSANY ['z']", "tags.*", &["a", "b"]).as_deref(),
Some("tags CONTAINSANY ['z']"),
);
}
#[test]
fn keeps_empty_containsany_leaf() {
assert_eq!(
residual("tags CONTAINSANY [] AND tags CONTAINSANY ['a']", "tags.*", &["a"]).as_deref(),
Some("tags CONTAINSANY []"),
);
}
#[test]
fn keeps_containsall_leaf() {
assert_eq!(
residual("tags CONTAINSALL ['a', 'b']", "tags.*", &["a", "b"]).as_deref(),
Some("tags CONTAINSALL ['a', 'b']"),
);
}
#[test]
fn keeps_leaf_on_a_different_field() {
assert_eq!(
residual("other CONTAINSANY ['a', 'b']", "tags.*", &["a", "b"]).as_deref(),
Some("other CONTAINSANY ['a', 'b']"),
);
}
#[test]
fn keeps_leaf_when_index_column_is_not_an_array_flatten() {
assert_eq!(
residual("tags CONTAINSANY ['a', 'b']", "tags", &["a", "b"]).as_deref(),
Some("tags CONTAINSANY ['a', 'b']"),
);
}
}