From 047464cf92e5a31d02a696f5158e45f7d34c67eb Mon Sep 17 00:00:00 2001 From: Paul Masurel Date: Mon, 28 Sep 2026 18:46:37 +0200 Subject: [PATCH] Accelerating jitexpr conditions using necessary conditions (#3129) * jitexpr * Added possible necessary conditions to jitexpr. A function makes it possible to infer a necessary query from an expression to match. We can then accelerate queries involving a calculated field by not even evaluating the expression on docs that trivially do not match. * CR comment * CR comment * Fixing unit tests --------- Co-authored-by: Paul Masurel --- jitexpr/Cargo.toml | 3 + jitexpr/src/ast/mod.rs | 4 + jitexpr/src/ast/presence.rs | 902 ++++++++++++++++++ jitexpr/src/types.rs | 1 - src/index/inverted_index_plugin.rs | 3 +- src/query/all_query.rs | 3 +- .../doc_predicate_query/function_predicate.rs | 7 +- .../doc_predicate_query/jitexpr_predicate.rs | 235 ++++- src/query/doc_predicate_query/mod.rs | 411 ++++++-- src/query/exist_query.rs | 2 +- 10 files changed, 1487 insertions(+), 84 deletions(-) create mode 100644 jitexpr/src/ast/presence.rs diff --git a/jitexpr/Cargo.toml b/jitexpr/Cargo.toml index e68d38323..85033708c 100644 --- a/jitexpr/Cargo.toml +++ b/jitexpr/Cargo.toml @@ -15,3 +15,6 @@ cranelift-native = "0.134.3" lru = "0.18.2" regex = "1" thiserror = "2.0.1" + +[dev-dependencies] +proptest = "1.7.0" diff --git a/jitexpr/src/ast/mod.rs b/jitexpr/src/ast/mod.rs index 9e5dd2f2b..de61743c1 100644 --- a/jitexpr/src/ast/mod.rs +++ b/jitexpr/src/ast/mod.rs @@ -1,11 +1,15 @@ mod infer_types; mod literal; +mod presence; mod serde; mod untyped_expr; pub use infer_types::{InferredTypeSet, TypeError, infer_types, infer_types_with_target}; pub(crate) use infer_types::{infer_type_with_variable_types, infer_types_aux}; pub use literal::{Literal, NonFiniteFloat}; +pub use presence::{ + ConditionSet, VariablePresenceCondition, required_presence, required_presence_for_true, +}; pub(crate) use serde::format_variable_name; pub use serde::{DeserializeError, deserialize, serialize}; pub use untyped_expr::UntypedExpr; diff --git a/jitexpr/src/ast/presence.rs b/jitexpr/src/ast/presence.rs new file mode 100644 index 000000000..edcf83fd0 --- /dev/null +++ b/jitexpr/src/ast/presence.rs @@ -0,0 +1,902 @@ +//! Necessary conditions on the presence of variables. +//! +//! Most functions return null as soon as one of their arguments is null. An expression can +//! therefore often only produce a value, or only evaluate to `true`, if some of its variables are +//! present. For instance, `(EQ (ADD a 1i64) b)` is null unless both `a` and `b` are present. +//! +//! A caller evaluating a predicate over many documents can use this to skip the documents missing +//! these variables without evaluating the expression. + +use std::sync::Arc; + +use crate::ast::{Function, Literal, UntypedExpr}; + +/// A boolean formula over the presence of variables. +/// +/// It is meant to be used as a necessary condition: it is implied by some property of an +/// expression (producing a value, or evaluating to `true`), but it does not imply it. +/// +/// As much as possible, we try to normalize these object, in order to simplify them. +#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] +pub enum VariablePresenceCondition { + /// Always satisfied. + Always, + /// Never satisfied. + Never, + /// Satisfied when the variable is present. + Present(Arc), + /// Satisfied when all of the conditions are satisfied. + All(ConditionSet), + /// Satisfied when at least one of the conditions is satisfied. + Any(ConditionSet), +} + +/// The children of a [`VariablePresenceCondition::All`] or [`VariablePresenceCondition::Any`] node. +/// +/// It can only be built through [`VariablePresenceCondition::all`] and +/// [`VariablePresenceCondition::any`], which +/// uphold the following hidden contract, on which the derived `Eq` and `Hash` rely: +/// - children are sorted and distinct, and there are at least two of them; +/// - no child is `Always`, `Never`, or a node of the same kind as the parent. +#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] +pub struct ConditionSet(Vec); + +impl ConditionSet { + /// Returns the children, in canonical order. + pub fn iter(&self) -> impl Iterator { + self.0.iter() + } +} + +impl VariablePresenceCondition { + pub fn all( + conditions: impl IntoIterator, + ) -> VariablePresenceCondition { + let mut children: Vec = Vec::new(); + for condition in conditions { + match condition { + VariablePresenceCondition::Always => {} + VariablePresenceCondition::Never => return VariablePresenceCondition::Never, + VariablePresenceCondition::All(grand_children) => children.extend(grand_children.0), + condition => children.push(condition), + } + } + children.sort(); + children.dedup(); + match children.len() { + 0 => VariablePresenceCondition::Always, + 1 => children.pop().unwrap(), + _ => VariablePresenceCondition::All(ConditionSet(children)), + } + } + + pub fn any( + conditions: impl IntoIterator, + ) -> VariablePresenceCondition { + let mut children: Vec = Vec::new(); + for condition in conditions { + match condition { + VariablePresenceCondition::Never => {} + VariablePresenceCondition::Always => return VariablePresenceCondition::Always, + VariablePresenceCondition::Any(grand_children) => children.extend(grand_children.0), + condition => children.push(condition), + } + } + children.sort(); + children.dedup(); + match children.len() { + 0 => VariablePresenceCondition::Never, + 1 => children.pop().unwrap(), + _ => VariablePresenceCondition::Any(ConditionSet(children)), + } + } + + /// Evaluates the condition, given the presence of each variable. + #[cfg(test)] + pub fn eval(&self, is_present: &mut impl FnMut(&str) -> bool) -> bool { + match self { + VariablePresenceCondition::Always => true, + VariablePresenceCondition::Never => false, + VariablePresenceCondition::Present(variable_name) => is_present(variable_name), + VariablePresenceCondition::All(conditions) => conditions + .iter() + .all(|condition| condition.eval(&mut *is_present)), + VariablePresenceCondition::Any(conditions) => conditions + .iter() + .any(|condition| condition.eval(&mut *is_present)), + } + } +} + +/// Returns a necessary presence condition for `expr` to evaluate to a non-null value. +pub fn required_presence(expr: &UntypedExpr) -> VariablePresenceCondition { + match expr { + UntypedExpr::Literal(_) => VariablePresenceCondition::Always, + UntypedExpr::Variable(variable_name) => { + VariablePresenceCondition::Present(variable_name.clone()) + } + UntypedExpr::FnCall { function, args } => required_presence_for_fn_call(*function, args), + } +} + +/// Returns a necessary presence condition for `expr` to evaluate to a present `true`. +pub fn required_presence_for_true(expr: &UntypedExpr) -> VariablePresenceCondition { + match expr { + UntypedExpr::Literal(Literal::Bool(true)) => VariablePresenceCondition::Always, + UntypedExpr::Literal(_) => VariablePresenceCondition::Never, + UntypedExpr::Variable(variable_name) => { + VariablePresenceCondition::Present(variable_name.clone()) + } + UntypedExpr::FnCall { function, args } => { + required_presence_for_true_for_fn_call(*function, args) + } + } +} + +fn required_presence_for_fn_call( + function: Function, + args: &[UntypedExpr], +) -> VariablePresenceCondition { + // null argument as "null in, null out" would make callers skip matching documents. + match function { + // A null argument makes the result null. + // + // AND belongs here: `(AND false none)` is null. + Function::Abs + | Function::Add + | Function::And + | Function::Ceil + | Function::Concat + | Function::Divide + | Function::Eq + | Function::Floor + | Function::Gt + | Function::GtEq + | Function::IntMod + | Function::Left + | Function::Lower + | Function::Lt + | Function::LtEq + | Function::Max + | Function::Min + | Function::Multiply + | Function::Pow + | Function::RegexpExtract + | Function::Right + | Function::Round + | Function::SplitAfter + | Function::SplitBefore + | Function::Sqrt + | Function::Substring + | Function::SubstringCount + | Function::Subtract + | Function::TextJoin + | Function::Trim + | Function::Upper => VariablePresenceCondition::all(args.iter().map(required_presence)), + // OR is null only if all of its arguments are null. + Function::Or => VariablePresenceCondition::any(args.iter().map(required_presence)), + // IF is null if its condition is null. Otherwise it takes the presence of the selected + // branch. + Function::If => { + let [condition, when_true, when_false] = args else { + return VariablePresenceCondition::Always; + }; + VariablePresenceCondition::all([ + required_presence(condition), + VariablePresenceCondition::any([ + required_presence(when_true), + required_presence(when_false), + ]), + ]) + } + // These functions always return a present value. + // + // REGEXP_LIKE returns `false` for a null input. + Function::IsNotNull + | Function::IsNull + | Function::Neq + | Function::Not + | Function::RegexpLike => VariablePresenceCondition::Always, + } +} + +fn required_presence_for_true_for_fn_call( + function: Function, + args: &[UntypedExpr], +) -> VariablePresenceCondition { + match function { + Function::And => { + VariablePresenceCondition::all(args.iter().map(required_presence_for_true)) + } + Function::Or => VariablePresenceCondition::any(args.iter().map(required_presence_for_true)), + Function::If => { + let [condition, when_true, when_false] = args else { + return VariablePresenceCondition::Always; + }; + VariablePresenceCondition::all([ + required_presence(condition), + VariablePresenceCondition::any([ + required_presence_for_true(when_true), + required_presence_for_true(when_false), + ]), + ]) + } + Function::IsNotNull => { + let [arg] = args else { + return VariablePresenceCondition::Always; + }; + required_presence(arg) + } + // REGEXP_LIKE returns `false` for a null input. + Function::RegexpLike => { + let Some(input) = args.first() else { + return VariablePresenceCondition::Always; + }; + required_presence(input) + } + // A `true` result is in particular a present result. This fallback is therefore correct + // for any function, including functions added later. + // + // It yields `Always` for NOT, NEQ, and IS_NULL, which are `true` when their + // argument is null. + _ => required_presence_for_fn_call(function, args), + } +} + +#[cfg(test)] +mod tests { + use std::collections::HashMap; + use std::collections::hash_map::DefaultHasher; + use std::hash::{Hash, Hasher}; + + use proptest::prelude::*; + use proptest::strategy::BoxedStrategy; + + use super::*; + use crate::ast::{InferredTypeSet, deserialize, infer_types_with_target}; + use crate::compile::{StringArena, compile}; + use crate::types::{VarType, VariableValue}; + + fn present(variable_name: &str) -> VariablePresenceCondition { + VariablePresenceCondition::Present(Arc::from(variable_name)) + } + + fn all(conditions: Vec) -> VariablePresenceCondition { + VariablePresenceCondition::all(conditions) + } + + fn any(conditions: Vec) -> VariablePresenceCondition { + VariablePresenceCondition::any(conditions) + } + + fn for_true(expr: &str) -> VariablePresenceCondition { + required_presence_for_true(&deserialize(expr).unwrap()) + } + + fn for_value(expr: &str) -> VariablePresenceCondition { + required_presence(&deserialize(expr).unwrap()) + } + + #[test] + fn test_all_simplification() { + assert_eq!( + VariablePresenceCondition::all([]), + VariablePresenceCondition::Always + ); + assert_eq!( + VariablePresenceCondition::all([VariablePresenceCondition::Always, present("a")]), + present("a") + ); + assert_eq!( + VariablePresenceCondition::all([present("a"), VariablePresenceCondition::Never]), + VariablePresenceCondition::Never + ); + assert_eq!( + VariablePresenceCondition::all([ + present("a"), + all(vec![present("b"), present("a")]), + present("c"), + ]), + all(vec![present("a"), present("b"), present("c")]) + ); + assert_eq!( + VariablePresenceCondition::all([any(vec![present("a"), present("b")]), present("c")]), + all(vec![any(vec![present("a"), present("b")]), present("c")]) + ); + } + + #[test] + fn test_any_simplification() { + assert_eq!( + VariablePresenceCondition::any([]), + VariablePresenceCondition::Never + ); + assert_eq!( + VariablePresenceCondition::any([VariablePresenceCondition::Never, present("a")]), + present("a") + ); + assert_eq!( + VariablePresenceCondition::any([present("a"), VariablePresenceCondition::Always]), + VariablePresenceCondition::Always + ); + assert_eq!( + VariablePresenceCondition::any([present("a"), any(vec![present("b"), present("a")])]), + any(vec![present("a"), present("b")]) + ); + } + + fn hash_of(condition: &VariablePresenceCondition) -> u64 { + let mut hasher = DefaultHasher::new(); + condition.hash(&mut hasher); + hasher.finish() + } + + fn assert_same(left: VariablePresenceCondition, right: VariablePresenceCondition) { + assert_eq!(left, right); + assert_eq!(hash_of(&left), hash_of(&right)); + } + + #[test] + fn test_canonical_order() { + let (a, b, c) = (present("a"), present("b"), present("c")); + assert_same( + VariablePresenceCondition::all([a.clone(), b.clone()]), + VariablePresenceCondition::all([b.clone(), a.clone()]), + ); + assert_same( + VariablePresenceCondition::any([c.clone(), a.clone(), b.clone()]), + VariablePresenceCondition::any([b.clone(), c.clone(), a.clone()]), + ); + assert_ne!( + VariablePresenceCondition::all([a.clone(), b.clone()]), + VariablePresenceCondition::any([a.clone(), b.clone()]) + ); + let VariablePresenceCondition::All(children) = + VariablePresenceCondition::all([c.clone(), a.clone()]) + else { + panic!("expected an All node"); + }; + assert_eq!(children.iter().collect::>(), vec![&a, &c]); + } + + #[test] + fn test_canonical_grouping_and_repetition() { + let (a, b, c) = (present("a"), present("b"), present("c")); + assert_same( + VariablePresenceCondition::all([ + a.clone(), + VariablePresenceCondition::all([b.clone(), c.clone()]), + ]), + VariablePresenceCondition::all([ + VariablePresenceCondition::all([c.clone(), a.clone()]), + b.clone(), + ]), + ); + assert_same( + VariablePresenceCondition::all([a.clone(), a.clone()]), + a.clone(), + ); + assert_same( + VariablePresenceCondition::any([ + VariablePresenceCondition::all([a.clone(), b.clone()]), + VariablePresenceCondition::all([b.clone(), a.clone()]), + ]), + VariablePresenceCondition::all([a.clone(), b.clone()]), + ); + } + + /// A condition tree built without any normalization. + #[derive(Clone, Debug)] + enum RawCondition { + Always, + Never, + Present(usize), + All(Vec), + Any(Vec), + } + + const RAW_VARIABLES: [&str; 4] = ["a", "b", "c", "d"]; + + impl RawCondition { + fn eval(&self, present_mask: u32) -> bool { + match self { + RawCondition::Always => true, + RawCondition::Never => false, + RawCondition::Present(variable_ord) => present_mask & (1 << variable_ord) != 0, + RawCondition::All(children) => { + children.iter().all(|child| child.eval(present_mask)) + } + RawCondition::Any(children) => { + children.iter().any(|child| child.eval(present_mask)) + } + } + } + + /// Builds the canonical condition, visiting children in reverse order if `reverse`. + fn build(&self, reverse: bool) -> VariablePresenceCondition { + let build_children = |children: &[RawCondition]| { + let mut built: Vec = + children.iter().map(|child| child.build(reverse)).collect(); + if reverse { + built.reverse(); + } + built + }; + match self { + RawCondition::Always => VariablePresenceCondition::Always, + RawCondition::Never => VariablePresenceCondition::Never, + RawCondition::Present(variable_ord) => present(RAW_VARIABLES[*variable_ord]), + RawCondition::All(children) => { + VariablePresenceCondition::all(build_children(children)) + } + RawCondition::Any(children) => { + VariablePresenceCondition::any(build_children(children)) + } + } + } + } + + fn raw_conditions() -> impl Strategy { + let leaf = prop_oneof![ + 1 => Just(RawCondition::Always), + 1 => Just(RawCondition::Never), + 6 => (0..RAW_VARIABLES.len()).prop_map(RawCondition::Present), + ]; + leaf.prop_recursive(4, 32, 4, |inner| { + prop_oneof![ + prop::collection::vec(inner.clone(), 0..4).prop_map(RawCondition::All), + prop::collection::vec(inner, 0..4).prop_map(RawCondition::Any), + ] + }) + } + + /// Checks the hidden contract of `ConditionSet`, recursively. + fn assert_canonical(condition: &VariablePresenceCondition) { + let (children, is_all) = match condition { + VariablePresenceCondition::All(children) => (children, true), + VariablePresenceCondition::Any(children) => (children, false), + _ => return, + }; + let children: Vec<&VariablePresenceCondition> = children.iter().collect(); + assert!(children.len() >= 2, "{condition:?}"); + assert!( + children.windows(2).all(|pair| pair[0] < pair[1]), + "{condition:?}" + ); + for child in &children { + assert_canonical(child); + match (child, is_all) { + (VariablePresenceCondition::Always | VariablePresenceCondition::Never, _) => { + panic!("neutral or absorbing child in {condition:?}") + } + (VariablePresenceCondition::All(_), true) + | (VariablePresenceCondition::Any(_), false) => { + panic!("same-kind child in {condition:?}") + } + _ => {} + } + } + } + + proptest! { + #[test] + fn proptest_canonical_form(raw in raw_conditions()) { + let condition = raw.build(false); + assert_canonical(&condition); + let reversed = raw.build(true); + prop_assert_eq!(&condition, &reversed); + prop_assert_eq!(hash_of(&condition), hash_of(&reversed)); + for present_mask in 0..(1u32 << RAW_VARIABLES.len()) { + let mut is_present = |variable_name: &str| { + let variable_ord = + RAW_VARIABLES.iter().position(|name| *name == variable_name).unwrap(); + present_mask & (1 << variable_ord) != 0 + }; + prop_assert_eq!(condition.eval(&mut is_present), raw.eval(present_mask)); + } + } + } + + fn presence_of<'a>(present_names: &'a [&'a str]) -> impl FnMut(&str) -> bool + 'a { + move |variable_name: &str| present_names.contains(&variable_name) + } + + #[test] + fn test_eval() { + let condition = all(vec![present("a"), any(vec![present("b"), present("c")])]); + assert!(condition.eval(&mut presence_of(&["a", "c"]))); + assert!(!condition.eval(&mut presence_of(&["a"]))); + assert!(!condition.eval(&mut presence_of(&["b", "c"]))); + assert!(VariablePresenceCondition::Always.eval(&mut presence_of(&[]))); + assert!(!VariablePresenceCondition::Never.eval(&mut presence_of(&["a"]))); + } + + #[test] + fn test_literals() { + assert_eq!(for_true("true"), VariablePresenceCondition::Always); + assert_eq!(for_true("false"), VariablePresenceCondition::Never); + assert_eq!(for_true("none"), VariablePresenceCondition::Never); + assert_eq!(for_value("none"), VariablePresenceCondition::Always); + assert_eq!(for_value("1u64"), VariablePresenceCondition::Always); + // `none` in a strict function is conservatively ignored. + assert_eq!(for_true("(EQ a none)"), present("a")); + } + + #[test] + fn test_variable() { + assert_eq!(for_true("flag"), present("flag")); + assert_eq!(for_value("a"), present("a")); + } + + #[test] + fn test_strict_functions() { + assert_eq!( + for_true("(EQ (ADD a 1i64) b)"), + all(vec![present("a"), present("b")]) + ); + assert_eq!(for_true("(GT (ABS a) 3i64)"), present("a")); + assert_eq!( + for_value(r#"(CONCAT "," "true" (UPPER a) (SUBSTRING b 0i64 2i64))"#), + all(vec![present("a"), present("b")]) + ); + assert_eq!(for_value(r#"(REGEXP_EXTRACT a "(x+)" 1u64)"#), present("a")); + assert_eq!(for_value("(ADD)"), VariablePresenceCondition::Always); + } + + #[test] + fn test_and_or() { + assert_eq!( + for_true("(AND (EQ a 1i64) (LT b 2i64) c)"), + all(vec![present("a"), present("b"), present("c")]) + ); + assert_eq!( + for_true("(OR (EQ a 1i64) (EQ b 2i64))"), + any(vec![present("a"), present("b")]) + ); + assert_eq!( + for_true("(AND (OR (EQ a 1i64) (EQ b 2i64)) (EQ c 3i64))"), + all(vec![any(vec![present("a"), present("b")]), present("c")]) + ); + // AND is null as soon as one of its arguments is null. + assert_eq!(for_value("(AND (NOT a) b)"), present("b")); + assert_eq!( + for_value("(OR (EQ a 1i64) (EQ b 2i64))"), + any(vec![present("a"), present("b")]) + ); + } + + #[test] + fn test_null_tolerant_functions() { + assert_eq!( + for_true("(NOT (EQ a 1i64))"), + VariablePresenceCondition::Always + ); + assert_eq!(for_true("(NEQ a 1i64)"), VariablePresenceCondition::Always); + assert_eq!(for_true("(IS_NULL a)"), VariablePresenceCondition::Always); + assert_eq!( + for_value("(IS_NOT_NULL a)"), + VariablePresenceCondition::Always + ); + assert_eq!( + for_true("(IS_NOT_NULL (ADD a b))"), + all(vec![present("a"), present("b")]) + ); + assert_eq!( + for_value(r#"(REGEXP_LIKE a "x")"#), + VariablePresenceCondition::Always + ); + assert_eq!(for_true(r#"(REGEXP_LIKE a "x")"#), present("a")); + assert_eq!( + for_true("(OR (EQ a 1i64) (IS_NULL b))"), + VariablePresenceCondition::Always + ); + assert_eq!(for_true("(AND (EQ a 1i64) (IS_NULL b))"), present("a")); + } + + #[test] + fn test_if() { + assert_eq!( + for_value("(IF c a b)"), + all(vec![present("c"), any(vec![present("a"), present("b")])]) + ); + assert_eq!( + for_true("(IF c (EQ a 1i64) (EQ b 1i64))"), + all(vec![present("c"), any(vec![present("a"), present("b")])]) + ); + assert_eq!(for_true("(IF c true false)"), present("c")); + assert_eq!( + for_true("(IF c false false)"), + VariablePresenceCondition::Never + ); + assert_eq!(for_value("(IF c 1i64 a)"), present("c")); + } + + // The property tests below check that the conditions are indeed necessary, by comparing them + // with the compiled expression over random inputs. + + const VARIABLES: [(&str, VarType); 7] = [ + ("b0", VarType::Bool), + ("b1", VarType::Bool), + ("n0", VarType::I64), + ("n1", VarType::I64), + ("f0", VarType::F64), + ("s0", VarType::Str), + ("s1", VarType::Str), + ]; + + /// Values are indexed by variable, in the order of `VARIABLES`. `None` means null. + type Assignment = Vec>; + + fn variable_value(var_type: VarType, value_ord: u8) -> VariableValue<'static> { + let value_ord = value_ord as usize; + match var_type { + VarType::Bool => VariableValue::some([true, false, true, false][value_ord]), + VarType::I64 => VariableValue::some([0i64, 1, -2, 3][value_ord]), + VarType::F64 => VariableValue::some([0.0f64, 1.5, -1.0, 2.0][value_ord]), + VarType::Str => VariableValue::some(["", "a", "ab,a", "ba"][value_ord]), + VarType::U64 | VarType::None => unreachable!(), + } + } + + fn assignments() -> impl Strategy> { + let value = prop_oneof![Just(None), (0u8..4).prop_map(Some)]; + prop::collection::vec(prop::collection::vec(value, VARIABLES.len()), 1..16) + } + + struct ExprStrategies { + boolean: BoxedStrategy, + number: BoxedStrategy, + string: BoxedStrategy, + } + + fn leaves() -> ExprStrategies { + let pick = |choices: &'static [&'static str]| { + prop::sample::select(choices) + .prop_map(str::to_string) + .boxed() + }; + ExprStrategies { + boolean: pick(&["b0", "b1", "true", "false", "none"]), + number: pick(&["n0", "n1", "f0", "0i64", "3i64", "-2i64", "1.5f64", "none"]), + string: pick(&["s0", "s1", r#""a""#, r#""""#, "none"]), + } + } + + fn unary(arg: &BoxedStrategy, template: &'static str) -> BoxedStrategy { + arg.clone() + .prop_map(move |arg| template.replace("$0", &arg)) + .boxed() + } + + fn binary( + left: &BoxedStrategy, + right: &BoxedStrategy, + template: &'static str, + ) -> BoxedStrategy { + (left.clone(), right.clone()) + .prop_map(move |(left, right)| template.replace("$0", &left).replace("$1", &right)) + .boxed() + } + + fn ternary( + first: &BoxedStrategy, + second: &BoxedStrategy, + third: &BoxedStrategy, + template: &'static str, + ) -> BoxedStrategy { + (first.clone(), second.clone(), third.clone()) + .prop_map(move |(first, second, third)| { + template + .replace("$0", &first) + .replace("$1", &second) + .replace("$2", &third) + }) + .boxed() + } + + /// Returns strategies generating well-typed expressions of the given depth. + fn exprs(depth: u32) -> ExprStrategies { + let leaves = leaves(); + if depth == 0 { + return leaves; + } + let ExprStrategies { + boolean: b, + number: n, + string: s, + } = exprs(depth - 1); + let any_kind = prop_oneof![b.clone(), n.clone(), s.clone()].boxed(); + let boolean = prop::strategy::Union::new(vec![ + leaves.boolean, + binary(&b, &b, "(AND $0 $1)"), + ternary(&b, &b, &b, "(AND $0 $1 $2)"), + binary(&b, &b, "(OR $0 $1)"), + ternary(&b, &b, &b, "(OR $0 $1 $2)"), + unary(&b, "(NOT $0)"), + unary(&any_kind, "(IS_NULL $0)"), + unary(&any_kind, "(IS_NOT_NULL $0)"), + binary(&n, &n, "(EQ $0 $1)"), + binary(&s, &s, "(EQ $0 $1)"), + binary(&b, &b, "(EQ $0 $1)"), + binary(&n, &n, "(NEQ $0 $1)"), + binary(&s, &s, "(NEQ $0 $1)"), + binary(&n, &n, "(LT $0 $1)"), + binary(&n, &n, "(LT_EQ $0 $1)"), + binary(&n, &n, "(GT $0 $1)"), + binary(&s, &s, "(GT_EQ $0 $1)"), + unary(&s, r#"(REGEXP_LIKE $0 "a")"#), + ternary(&b, &b, &b, "(IF $0 $1 $2)"), + ]) + .boxed(); + let number = prop::strategy::Union::new(vec![ + leaves.number, + unary(&n, "(ADD $0)"), + binary(&n, &n, "(ADD $0 $1)"), + binary(&n, &n, "(SUBTRACT $0 $1)"), + binary(&n, &n, "(MULTIPLY $0 $1)"), + binary(&n, &n, "(DIVIDE $0 $1)"), + binary(&n, &n, "(POW $0 $1)"), + binary(&n, &n, "(INT_MOD $0 $1)"), + binary(&n, &n, "(MIN $0 $1)"), + binary(&n, &n, "(MAX $0 $1)"), + unary(&n, "(ABS $0)"), + unary(&n, "(CEIL $0)"), + unary(&n, "(FLOOR $0)"), + unary(&n, "(SQRT $0)"), + unary(&n, "(ROUND $0)"), + unary(&n, "(ROUND $0 1i64)"), + // SUBSTRING_COUNT is not generated: its native implementation builds a slice from a + // null pointer when the haystack is null, which aborts debug builds. + ternary(&b, &n, &n, "(IF $0 $1 $2)"), + ]) + .boxed(); + let string = prop::strategy::Union::new(vec![ + leaves.string, + unary(&s, "(UPPER $0)"), + unary(&s, "(LOWER $0)"), + unary(&s, "(LEFT $0 1i64)"), + unary(&s, "(RIGHT $0 1i64)"), + unary(&s, "(SUBSTRING $0 0i64 1i64)"), + binary(&s, &s, r#"(CONCAT "," "false" $0 $1)"#), + binary(&s, &s, r#"(TEXT_JOIN "," "true" $0 $1)"#), + unary(&s, r#"(TRIM $0 "a" "both")"#), + unary(&s, r#"(SPLIT_AFTER $0 ",")"#), + unary(&s, r#"(SPLIT_BEFORE $0 "," 0i64)"#), + unary(&s, r#"(REGEXP_EXTRACT $0 "(a)b" 1u64)"#), + // IF is not generated for strings: with a null condition, it returns the selected + // branch instead of null, as the string pointer is not cleared. + ]) + .boxed(); + ExprStrategies { + boolean, + number, + string, + } + } + + /// Compiles `expr_str`, then checks that `required(expr)` holds for every assignment where + /// `holds(result)` is true. + /// + /// Following tantivy's fast field binding, a variable is bound only if its type is accepted by + /// type inference. Unbound variables are null, and therefore absent. + fn check_necessary_condition( + expr_str: &str, + target_type: InferredTypeSet, + assignments: &[Assignment], + required: fn(&UntypedExpr) -> VariablePresenceCondition, + holds: fn(VarType, VariableValue) -> bool, + ) -> Result<(), TestCaseError> { + let expr = deserialize(expr_str).unwrap(); + let Ok(inferred_types) = infer_types_with_target(&expr, target_type) else { + return Err(TestCaseError::reject("type inference failed")); + }; + let mut variable_types: HashMap<&str, VarType> = + HashMap::with_capacity(inferred_types.len()); + for (variable_name, accepted_types) in &inferred_types { + let (_, var_type) = VARIABLES + .iter() + .find(|(name, _)| name == variable_name) + .unwrap(); + if accepted_types.contains(*var_type) { + variable_types.insert(*variable_name, *var_type); + } + } + // Some expressions trip debug assertions of the compiler, unrelated to presence. For + // instance, `(SQRT (CEIL n0))` asks CEIL for a f64, while it always returns an i64. + let compile_result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| { + compile(&expr, &variable_types) + })); + let Ok(Ok(compiled_fn)) = compile_result else { + return Err(TestCaseError::reject("compilation failed")); + }; + let condition = required(&expr); + let mut string_arena = StringArena::default(); + for assignment in assignments { + let variable_ord = |variable_name: &str| { + VARIABLES + .iter() + .position(|(name, _)| *name == variable_name) + .unwrap() + }; + let args: Vec = compiled_fn + .inputs() + .iter() + .map( + |input| match assignment[variable_ord(&input.variable_name)] { + Some(value_ord) => variable_value(input.r#type, value_ord), + None => VariableValue::none(), + }, + ) + .collect(); + // SAFETY: Each slot follows the compiled input order, and uses the input type. + let result = unsafe { compiled_fn.call(&args, &mut string_arena) }; + if !holds(compiled_fn.result_type(), result) { + continue; + } + let mut is_present = |variable_name: &str| { + variable_types.contains_key(variable_name) + && assignment[variable_ord(variable_name)].is_some() + }; + prop_assert!( + condition.eval(&mut is_present), + "{expr_str} holds for {assignment:?}, but {condition:?} does not" + ); + } + Ok(()) + } + + fn is_true(result_type: VarType, result: VariableValue) -> bool { + // SAFETY: The union member is selected with the result type. + result_type == VarType::Bool && unsafe { result.as_bool() } == Some(true) + } + + fn is_present(result_type: VarType, result: VariableValue) -> bool { + // SAFETY: The union member is selected with the result type. + unsafe { + match result_type { + VarType::Bool => result.as_bool().is_some(), + VarType::F64 => result.as_f64().is_some(), + VarType::U64 => result.as_u64().is_some(), + VarType::I64 => result.as_i64().is_some(), + VarType::Str => result.as_str().is_some(), + VarType::None => false, + } + } + } + + proptest! { + // Compiler debug assertions reject a fraction of the generated expressions. + #![proptest_config(ProptestConfig { + max_global_rejects: 1 << 16, + ..ProptestConfig::with_cases(512) + })] + + #[test] + fn proptest_required_presence_for_true_is_necessary( + expr in exprs(3).boolean, + assignments in assignments(), + ) { + check_necessary_condition( + &expr, + InferredTypeSet::BOOLEAN, + &assignments, + required_presence_for_true, + is_true, + )?; + } + + #[test] + fn proptest_required_presence_is_necessary( + expr in prop_oneof![exprs(3).boolean, exprs(3).number, exprs(3).string], + assignments in assignments(), + ) { + check_necessary_condition( + &expr, + InferredTypeSet::ALL, + &assignments, + required_presence, + is_present, + )?; + } + } +} diff --git a/jitexpr/src/types.rs b/jitexpr/src/types.rs index 626e968f7..71e309739 100644 --- a/jitexpr/src/types.rs +++ b/jitexpr/src/types.rs @@ -370,7 +370,6 @@ impl<'a> From for VariableValue<'a> { #[cfg(test)] mod tests { use std::cmp::Ordering; - use std::collections::HashSet; use std::collections::hash_map::DefaultHasher; use std::hash::{Hash, Hasher}; diff --git a/src/index/inverted_index_plugin.rs b/src/index/inverted_index_plugin.rs index add130c75..bc9c55e58 100644 --- a/src/index/inverted_index_plugin.rs +++ b/src/index/inverted_index_plugin.rs @@ -712,12 +712,13 @@ fn write_postings_merge( Ok(()) } +#[cfg(not(feature = "compare_hash_only"))] #[cfg(test)] mod tests { + use super::compute_initial_table_size; #[test] - #[cfg(not(feature = "compare_hash_only"))] fn test_hashmap_size() { assert_eq!(compute_initial_table_size(100_000).unwrap(), 1 << 12); assert_eq!(compute_initial_table_size(1_000_000).unwrap(), 1 << 15); diff --git a/src/query/all_query.rs b/src/query/all_query.rs index 5431a3a1b..e4749ac24 100644 --- a/src/query/all_query.rs +++ b/src/query/all_query.rs @@ -47,7 +47,8 @@ pub struct AllScorer { impl AllScorer { /// Creates a new AllScorer with `max_doc` docs. pub fn new(max_doc: DocId) -> AllScorer { - AllScorer { doc: 0u32, max_doc } + let doc = if max_doc == 0u32 { TERMINATED } else { 0 }; + AllScorer { doc, max_doc } } } diff --git a/src/query/doc_predicate_query/function_predicate.rs b/src/query/doc_predicate_query/function_predicate.rs index 5d47e8282..acd5fd206 100644 --- a/src/query/doc_predicate_query/function_predicate.rs +++ b/src/query/doc_predicate_query/function_predicate.rs @@ -1,6 +1,7 @@ use super::{DocPredicate, SegmentDocPredicate}; use crate::index::SegmentReader; use crate::query::doc_predicate_query::ConstOrVariableSegmentPredicate; +use crate::query::AllScorer; use crate::DocId; /// Blanket [`SegmentDocPredicate`] implementation for any per-document @@ -53,7 +54,11 @@ where &self, segment_reader: &SegmentReader, ) -> crate::Result> { - (self.segment_predicate_factory)(segment_reader).map(ConstOrVariableSegmentPredicate::from) + let predicate = (self.segment_predicate_factory)(segment_reader)?; + Ok(ConstOrVariableSegmentPredicate::Variable { + predicate, + necessary_condition: Box::new(AllScorer::new(segment_reader.max_doc())), + }) } } diff --git a/src/query/doc_predicate_query/jitexpr_predicate.rs b/src/query/doc_predicate_query/jitexpr_predicate.rs index 94023ddcf..9698ac3fd 100644 --- a/src/query/doc_predicate_query/jitexpr_predicate.rs +++ b/src/query/doc_predicate_query/jitexpr_predicate.rs @@ -1,15 +1,21 @@ use std::collections::HashMap; use std::io; -use columnar::{ColumnType, DynamicColumn, StrColumn}; -use jitexpr::ast::{infer_types_with_target, InferredTypeSet, TypeError, UntypedExpr}; +use columnar::{ColumnIndex, ColumnType, DynamicColumn, StrColumn}; +use jitexpr::ast::{ + infer_types_with_target, required_presence_for_true, InferredTypeSet, TypeError, UntypedExpr, + VariablePresenceCondition, +}; use jitexpr::compile::{CompiledFnCtx, StringArena}; use jitexpr::types::{VarType, VariableValue}; use super::{DocPredicate, SegmentDocPredicate}; use crate::index::SegmentReader; use crate::query::doc_predicate_query::ConstOrVariableSegmentPredicate; -use crate::{DocId, TantivyError}; +use crate::query::exist_query::{ExistsColumnIndex, ExistsDocSet}; +use crate::query::union::SimpleUnion; +use crate::query::{AllScorer, EmptyScorer, Intersection}; +use crate::{DocId, DocSet, TantivyError, TERMINATED}; /// A [`DocPredicate`] that evaluates a boolean JIT expression against fast fields. /// @@ -25,17 +31,15 @@ use crate::{DocId, TantivyError}; /// /// Only a present `true` result matches. /// -/// ``` -/// use tantivy::jitexpr::ast::deserialize; -/// use tantivy::query::doc_predicate_query::{DocPredicateQuery, JitExprPredicate}; -/// -/// let expression = deserialize("(EQ (ADD price 1u64) 10u64)").unwrap(); -/// let query: DocPredicateQuery = JitExprPredicate::new(expression).unwrap().into(); -/// ``` +/// Documents missing the fields required for the expression to be `true` are skipped without +/// being evaluated. For instance, `(EQ (ADD price 1u64) 10u64)` is only evaluated on the +/// documents having a `price` value. #[derive(Clone, Debug)] pub struct JitExprPredicate { expression: UntypedExpr, inferred_inputs: Vec<(String, InferredTypeSet)>, + // A necessary condition, on the presence of the variables, for the expression to be `true`. + required_presence: VariablePresenceCondition, } impl JitExprPredicate { @@ -47,9 +51,11 @@ impl JitExprPredicate { .into_iter() .map(|(name, types)| (name.to_string(), types)) .collect(); + let required_presence = required_presence_for_true(&expression); Ok(Self { expression, inferred_inputs, + required_presence, }) } @@ -91,6 +97,17 @@ impl DocPredicate for JitExprPredicate { opened_columns.insert(name.as_str(), column); } + // The variables are bound to the columns opened above, and only to them: the presence + // of a variable is the presence of a value in its column. + let necessary_condition: Box = build_necessary_condition_docset( + &self.required_presence, + &opened_columns, + segment_reader.max_doc(), + ); + if necessary_condition.doc() == TERMINATED { + return Ok(ConstOrVariableSegmentPredicate::Const(false)); + } + let compiled_fn = segment_reader .index() .expr_compilation_cache() @@ -157,13 +174,68 @@ impl DocPredicate for JitExprPredicate { .filter(|column_opt| matches!(column_opt, Some(DynamicColumn::Str(_)))) .count(); let num_inputs = columns_opt.len(); - Ok(JitExprEvalState { + let predicate = JitExprEvalState { compiled: compiled_fn.context(), columns_opt, string_inputs: vec![String::new(); num_string_inputs], input_values: Vec::with_capacity(num_inputs), + }; + Ok(ConstOrVariableSegmentPredicate::Variable { + predicate, + necessary_condition, + }) + } +} + +/// Builds a docset off the variable presence condition. +/// +/// Variables are bound to the columns of `columns`. A variable missing from `columns` is null +/// for all documents. +/// +/// The returned `DocSet` is positioned on its first document. It is `TERMINATED` if and only if +/// no document of the segment satisfies the condition. +fn build_necessary_condition_docset( + condition: &VariablePresenceCondition, + columns: &HashMap<&str, DynamicColumn>, + max_doc: DocId, +) -> Box { + match condition { + VariablePresenceCondition::Always => Box::new(AllScorer::new(max_doc)), + VariablePresenceCondition::Never => Box::new(EmptyScorer), + VariablePresenceCondition::Present(variable_name) => { + let Some(column) = columns.get(variable_name.as_ref()) else { + return Box::new(EmptyScorer); + }; + let exists_column_index = match column.column_index() { + ColumnIndex::Empty { .. } => return Box::new(EmptyScorer), + ColumnIndex::Full => return Box::new(AllScorer::new(max_doc)), + ColumnIndex::Optional(optional_index) => { + ExistsColumnIndex::Optional(optional_index.clone()) + } + ColumnIndex::Multivalued(multivalued_index) => { + ExistsColumnIndex::Multivalued(multivalued_index.clone()) + } + }; + Box::new(ExistsDocSet::new(exists_column_index)) + } + VariablePresenceCondition::All(conditions) => { + let mut doc_sets: Vec> = conditions + .iter() + .map(|condition| build_necessary_condition_docset(condition, columns, max_doc)) + .collect(); + match doc_sets.len() { + 0 => Box::new(AllScorer::new(max_doc)), + 1 => doc_sets.pop().unwrap(), + _ => Box::new(Intersection::new(doc_sets, max_doc)), + } + } + VariablePresenceCondition::Any(conditions) => { + let doc_sets: Vec> = conditions + .iter() + .map(|condition| build_necessary_condition_docset(condition, columns, max_doc)) + .collect(); + Box::new(SimpleUnion::build(doc_sets)) } - .into()) } } @@ -239,19 +311,19 @@ impl<'a> Drop for ClearOnDrop<'a> { impl SegmentDocPredicate for JitExprEvalState { fn eval(&mut self, doc_id: DocId) -> bool { // Input_values is just a buffer we share to avoid allocations - let mut inputs_vec = ClearOnDrop::wrap(&mut self.input_values); + let inputs_vec = ClearOnDrop::wrap(&mut self.input_values); fill_input_values( &self.columns_opt, &mut self.string_inputs, - &mut inputs_vec.0, + inputs_vec.0, doc_id, ); // SAFETY: Columns follow compiled.inputs() and their types were checked // during setup. Each slot uses the matching union arm. String buffers // remain borrowed, and cannot be mutated, until this call finishes. - let eval_result: Option = unsafe { self.compiled.call(&inputs_vec.0).as_bool() }; + let eval_result: Option = unsafe { self.compiled.call(inputs_vec.0).as_bool() }; eval_result == Some(true) } @@ -317,10 +389,11 @@ fn load_str_input<'buffer>( #[cfg(test)] mod tests { use super::*; - use crate::collector::Count; + use crate::collector::{Count, DocSetCollector}; use crate::query::doc_predicate_query::DocPredicateQuery; - use crate::schema::{Schema, FAST, STORED, STRING}; - use crate::Index; + use crate::query::{EnableScoring, Query}; + use crate::schema::{Schema, FAST, INDEXED, STORED, STRING}; + use crate::{Index, TantivyDocument, Term}; fn create_index() -> Index { let mut schema_builder = Schema::builder(); @@ -589,6 +662,132 @@ mod tests { ); } + /// Two segments with sparse, multivalued, and segment-dependent columns, and deleted docs. + /// + /// `label` only has values in the second segment. + fn create_sparse_index() -> Index { + let mut schema_builder = Schema::builder(); + let id = schema_builder.add_u64_field("id", FAST | INDEXED); + let number = schema_builder.add_u64_field("number", FAST); + let score = schema_builder.add_i64_field("score", FAST); + let flag = schema_builder.add_bool_field("flag", FAST); + let label = schema_builder.add_text_field("label", STRING | FAST); + let tags = schema_builder.add_text_field("tags", STRING | FAST); + let index = Index::create_in_ram(schema_builder.build()); + let mut writer = index.writer_for_tests().unwrap(); + for segment_ord in 0..2u64 { + for i in 0..300u64 { + let mut doc = TantivyDocument::default(); + doc.add_u64(id, segment_ord * 1000 + i); + if i % 3 == 0 { + doc.add_u64(number, i); + } + if i % 5 == 0 { + doc.add_i64(score, (i % 7) as i64 - 3); + } + if i % 2 == 0 { + doc.add_bool(flag, i % 4 == 0); + } + if segment_ord == 1 && i % 4 == 0 { + doc.add_text(label, ["a", "b", "ab"][(i % 3) as usize]); + } + if i % 6 == 0 { + doc.add_text(tags, "x"); + doc.add_text(tags, "y"); + } else if i % 6 == 1 { + doc.add_text(tags, "y"); + } + writer.add_document(doc).unwrap(); + } + writer.commit().unwrap(); + } + writer.delete_term(Term::from_field_u64(id, 30)); + writer.delete_term(Term::from_field_u64(id, 1060)); + writer.commit().unwrap(); + index + } + + /// Returns a query matching the same documents as `query(expression)`, but requiring the + /// presence of no field, so that all documents are evaluated. + /// + /// `(NOT true)` is a present `false`: the disjunction is `true` if and only if `expression` is. + /// As `NOT` requires the presence of no field, neither does the disjunction. + fn query_without_required_presence(expression: &str) -> DocPredicateQuery { + query(&format!("(OR {expression} (NOT true))")) + } + + #[test] + fn test_required_presence_does_not_change_results() { + let index = create_sparse_index(); + let searcher = index.reader().unwrap().searcher(); + assert_eq!(searcher.segment_readers().len(), 2); + let expressions = [ + "true", + "false", + "flag", + "(EQ number 33u64)", + "(EQ number 30u64)", + "(IS_NOT_NULL number)", + "(IS_NULL score)", + "(NOT (EQ number 3u64))", + "(NEQ score 1i64)", + "(GT (ADD number score) 10i64)", + "(OR (EQ number 3u64) (EQ score 1i64))", + "(AND flag (IS_NOT_NULL label))", + r#"(EQ label "a")"#, + r#"(OR (EQ label "b") flag)"#, + r#"(REGEXP_LIKE label "b")"#, + r#"(EQ (UPPER tags) "X")"#, + "(IF flag (GT number 100u64) (LT score 0i64))", + "(IS_NOT_NULL (IF flag number score))", + "(AND (EQ missing 1i64) flag)", + "(OR (IS_NOT_NULL missing) (EQ score 2i64))", + "(OR (IS_NULL missing) (EQ score 2i64))", + ]; + for expression in expressions { + let expected = searcher + .search( + &query_without_required_presence(expression), + &DocSetCollector, + ) + .unwrap(); + let accelerated = searcher + .search(&query(expression), &DocSetCollector) + .unwrap(); + assert_eq!(accelerated, expected, "{expression}"); + } + // Sanity checks: the index does exercise the predicates. + let count = |expression: &str| searcher.search(&query(expression), &Count).unwrap(); + assert_eq!(count("(EQ number 33u64)"), 2); + // Doc 30 of the first segment is deleted. + assert_eq!(count("(EQ number 30u64)"), 1); + // `label` is "a" on 25 docs of the second segment, one of which (1060) is deleted. + assert_eq!(count(r#"(EQ label "a")"#), 24); + } + + #[test] + fn test_required_presence_restricts_evaluated_docs() { + let index = create_sparse_index(); + let searcher = index.reader().unwrap().searcher(); + let segment_reader = searcher.segment_reader(0); + let size_hint = |query: DocPredicateQuery| { + query + .weight(EnableScoring::disabled_from_searcher(&searcher)) + .unwrap() + .scorer(segment_reader, 1.0) + .unwrap() + .size_hint() + }; + // `number` has a value in one doc out of three. + assert_eq!(size_hint(query("(GT number 10u64)")), 100); + assert_eq!( + size_hint(query_without_required_presence("(GT number 10u64)")), + 300 + ); + // Nothing to require: all docs are evaluated. + assert_eq!(size_hint(query("(NOT (GT number 10u64))")), 300); + } + // THIS FAILS! due to our pick best possible column approach policy. // #[test] // fn test_multi_typed_field_picks_one() { diff --git a/src/query/doc_predicate_query/mod.rs b/src/query/doc_predicate_query/mod.rs index 11e79dfd9..e8c74e106 100644 --- a/src/query/doc_predicate_query/mod.rs +++ b/src/query/doc_predicate_query/mod.rs @@ -1,3 +1,4 @@ +use std::cmp::Ordering; use std::sync::Arc; mod function_predicate; @@ -65,67 +66,130 @@ impl Weight for DocPredicateQuery { } } -/// A [`DocSet`] that walks documents by repeatedly evaluating a -/// [`SegmentDocPredicate`], starting from doc `0`. +/// A [`DocSet`] that walks the documents of a necessary condition, and evaluates a +/// [`SegmentDocPredicate`] on each of them. +/// +/// Hidden contract: every document matching the predicate belongs to the necessary condition. +/// Documents outside of it are never evaluated, and are considered as not matching. +/// +/// Hidden contract: whenever the `DocPredicateDocSet` is in a valid state, the necessary condition +/// is in a valid state too, positioned on a matching document (or `TERMINATED`). The current doc +/// is therefore simply the necessary condition's current doc. pub struct DocPredicateDocSet { doc_predicate: TSegmentDocPredicate, - doc: DocId, - max_doc: DocId, + necessary_condition: Box, +} + +impl DocPredicateDocSet { + /// Creates a `DocPredicateDocSet` positioned on its first matching document. + fn new(doc_predicate: TSegmentDocPredicate, necessary_condition: Box) -> Self { + let first_candidate = necessary_condition.doc(); + let mut doc_set = DocPredicateDocSet { + doc_predicate, + necessary_condition, + }; + doc_set.find_match(first_candidate); + doc_set + } + + /// Creates a `DocPredicateDocSet`, and seeks it to `target`, following + /// [`Weight::scorer_danger`]'s contract. + /// + /// Documents before `target` are not evaluated. + fn new_seeked_to( + doc_predicate: TSegmentDocPredicate, + necessary_condition: Box, + target: DocId, + ) -> (SeekDangerResult, Self) { + let first_candidate = necessary_condition.doc(); + let mut doc_set = DocPredicateDocSet { + doc_predicate, + necessary_condition, + }; + if target >= TERMINATED { + if doc_set.necessary_condition.doc() < TERMINATED { + doc_set.necessary_condition.seek(TERMINATED); + } + return (SeekDangerResult::SeekLowerBound(TERMINATED), doc_set); + } + let seek_result = match first_candidate.cmp(&target) { + Ordering::Less => doc_set.seek_danger(target), + Ordering::Equal => doc_set.eval_candidate(target), + Ordering::Greater => SeekDangerResult::SeekLowerBound(first_candidate), + }; + (seek_result, doc_set) + } + + /// Evaluates the predicate on `candidate`. + /// + /// Hidden contract: the necessary condition is positioned on `candidate`. + fn eval_candidate(&mut self, candidate: DocId) -> SeekDangerResult { + if self.doc_predicate.eval(candidate) { + SeekDangerResult::Found + } else { + SeekDangerResult::SeekLowerBound(candidate + 1) + } + } + + /// Advances to the first matching document at or after `candidate`. + /// + /// Hidden contract: the necessary condition is positioned on `candidate`. + fn find_match(&mut self, mut candidate: DocId) -> DocId { + debug_assert_eq!(candidate, self.necessary_condition.doc()); + while candidate != TERMINATED && !self.doc_predicate.eval(candidate) { + candidate = self.necessary_condition.advance(); + } + candidate + } } impl DocSet for DocPredicateDocSet { fn advance(&mut self) -> DocId { - if self.doc == TERMINATED { + if self.doc() == TERMINATED { return TERMINATED; } - self.find_match(self.doc + 1) + let candidate = self.necessary_condition.advance(); + self.find_match(candidate) } fn seek(&mut self, target: DocId) -> DocId { - debug_assert!(target >= self.doc); - if self.doc == TERMINATED { - return TERMINATED; + let doc = self.doc(); + debug_assert!(target >= doc); + // In a valid state, the current doc is a match (or TERMINATED). + if doc >= target { + return doc; } - self.find_match(target) + let candidate = self.necessary_condition.seek(target); + self.find_match(candidate) } fn seek_danger(&mut self, target: DocId) -> SeekDangerResult { - if target >= self.max_doc { - self.doc = TERMINATED; - return SeekDangerResult::SeekLowerBound(TERMINATED); - } - if self.doc_predicate.eval(target) { - self.doc = target; - SeekDangerResult::Found - } else { - SeekDangerResult::SeekLowerBound(target + 1) + match self.necessary_condition.seek_danger(target) { + SeekDangerResult::Found => self.eval_candidate(target), + // Following `seek_danger`'s contract, we are now in an invalid state, and `doc()` may + // return anything until a subsequent `seek_danger` returns `Found`. + seek_lower_bound @ SeekDangerResult::SeekLowerBound(_) => seek_lower_bound, } } fn doc(&self) -> DocId { - self.doc + self.necessary_condition.doc() } fn size_hint(&self) -> u32 { - self.max_doc + self.necessary_condition.size_hint() } -} -impl DocPredicateDocSet { - fn find_match(&mut self, mut target: DocId) -> DocId { - loop { - match self.seek_danger(target) { - SeekDangerResult::Found => return target, - SeekDangerResult::SeekLowerBound(next_target) => { - if next_target >= TERMINATED { - return TERMINATED; - } - target = next_target; - } - } - } + fn cost(&self) -> u64 { + // `cost` is the method used to tell how costly it is to consume a DocSet entirely. + // + // This is used in intersection to have cheaper docset "lead" the intersection. + // + // Here, we naturally use a model where we use the cost of the necessary condition + // multiplied by some factor expressing how slow it is to evaluate an expression. + self.necessary_condition.cost() * self.doc_predicate.cost() } } @@ -156,14 +220,16 @@ impl DocPredicateBoxable for TDocPredicate { EmptyWeight.scorer(segment_reader, boost) } } - ConstOrVariableSegmentPredicate::Variable(doc_predicate) => { - let mut doc_set = DocPredicateDocSet { - doc_predicate, - doc: 0u32, - max_doc: segment_reader.max_doc(), - }; - doc_set.doc = doc_set.find_match(0); - Ok(Box::new(ConstScorer::new(doc_set, boost)) as Box) + ConstOrVariableSegmentPredicate::Variable { + predicate, + necessary_condition, + } => { + if necessary_condition.doc() >= segment_reader.max_doc() { + // The necessary condition is empty. + return EmptyWeight.scorer(segment_reader, boost); + } + let doc_set = DocPredicateDocSet::new(predicate, necessary_condition); + Ok(Box::new(ConstScorer::new(doc_set, boost))) } } } @@ -183,15 +249,13 @@ impl DocPredicateBoxable for TDocPredicate { EmptyWeight.scorer_danger(segment_reader, target, boost) } } - ConstOrVariableSegmentPredicate::Variable(doc_predicate) => { - let mut doc_set = DocPredicateDocSet { - doc_predicate, - doc: target, - max_doc: segment_reader.max_doc(), - }; - let seek_result = doc_set.seek_danger(target); - let scorer = Box::new(ConstScorer::new(doc_set, boost)) as Box; - Ok((seek_result, scorer)) + ConstOrVariableSegmentPredicate::Variable { + predicate, + necessary_condition, + } => { + let (seek_result, doc_set) = + DocPredicateDocSet::new_seeked_to(predicate, necessary_condition, target); + Ok((seek_result, Box::new(ConstScorer::new(doc_set, boost)))) } } } @@ -202,14 +266,20 @@ pub enum ConstOrVariableSegmentPredicate { /// Can be emitted to hint that a predicate will be always true or false on a segment. /// Returning Const instead of a variable is an optimization. Const(bool), - /// Just a regular SegmentDocPredicate. - Variable(P), -} - -impl From

for ConstOrVariableSegmentPredicate

{ - fn from(predicate: P) -> Self { - ConstOrVariableSegmentPredicate::Variable(predicate) - } + /// A regular SegmentDocPredicate, evaluated document by document. + Variable { + /// The predicate to evaluate. + predicate: P, + /// The [`DocSet`] of the documents on which `predicate` is evaluated. + /// + /// Hidden contract: it must contain every document for which `predicate.eval` returns + /// true. Documents outside of it are never evaluated, and are considered as not + /// matching. Use an [`AllScorer`](crate::query::AllScorer) to evaluate every document of + /// the segment. + /// + /// The `DocSet` must be positioned on its first document. + necessary_condition: Box, + }, } /// A per-query predicate that produces a [`SegmentDocPredicate`] for each @@ -236,12 +306,30 @@ pub trait DocPredicate: Send + Sync + 'static + std::fmt::Debug { pub trait SegmentDocPredicate: Send + 'static { /// Returns whether `doc_id` matches the predicate. fn eval(&mut self, doc_id: DocId) -> bool; + + /// Cost for the evaluation of a given predicate. + /// + /// This is used to infer the cost of consuming an associated `DocPredicateDocSet`. + /// This does not need to be accurate. It is only used by the intersection scorer + /// to choose which `DocSet` should "drive" the intersection. + /// + /// 1 is the time it takes to call `TermScorer::advance` (a few cycles). We defensively default + /// to 100. + fn cost(&self) -> u64 { + // We assume a default value of 100. + 100u64 + } } #[cfg(test)] pub(crate) mod tests { + use std::sync::atomic::{AtomicUsize, Ordering as AtomicOrdering}; + + use proptest::prelude::*; + use super::*; - use crate::collector::Count; + use crate::collector::{Count, DocSetCollector}; + use crate::query::VecDocSet; pub(crate) fn create_index_for_test(num_docs: u32) -> crate::Index { let schema_builder = crate::schema::Schema::builder(); @@ -314,6 +402,207 @@ pub(crate) mod tests { assert_eq!(scorer.doc(), 2); } + /// Matches even doc ids, and counts its evaluations. + struct EvenDocIds { + num_evals: Arc, + } + + impl SegmentDocPredicate for EvenDocIds { + fn eval(&mut self, doc_id: DocId) -> bool { + self.num_evals.fetch_add(1, AtomicOrdering::Relaxed); + doc_id.is_multiple_of(2) + } + } + + /// `EvenDocIds`, with a fixed necessary condition. + #[derive(Debug)] + struct EvenWithNecessaryCondition { + necessary_condition: Vec, + num_evals: Arc, + } + + impl EvenWithNecessaryCondition { + fn new(necessary_condition: Vec) -> Self { + EvenWithNecessaryCondition { + necessary_condition, + num_evals: Arc::default(), + } + } + } + + impl DocPredicate for EvenWithNecessaryCondition { + type SegmentDocPredicate = EvenDocIds; + + fn doc_predicate( + &self, + _segment_reader: &SegmentReader, + ) -> crate::Result> { + Ok(ConstOrVariableSegmentPredicate::Variable { + predicate: EvenDocIds { + num_evals: self.num_evals.clone(), + }, + necessary_condition: Box::new(VecDocSet::from(self.necessary_condition.clone())), + }) + } + } + + #[test] + fn test_necessary_condition_restricts_evaluations() { + let index = create_index_for_test(10); + let searcher = index.reader().unwrap().searcher(); + let predicate = EvenWithNecessaryCondition::new(vec![1, 2, 3, 4, 6, 9]); + let num_evals = predicate.num_evals.clone(); + let query: DocPredicateQuery = predicate.into(); + assert_eq!(searcher.search(&query, &DocSetCollector).unwrap().len(), 3); + assert_eq!(num_evals.load(AtomicOrdering::Relaxed), 6); + } + + #[test] + fn test_necessary_condition_size_hint_and_cost() { + let index = create_index_for_test(10); + let searcher = index.reader().unwrap().searcher(); + let query: DocPredicateQuery = + EvenWithNecessaryCondition::new(vec![1, 2, 3, 4, 6, 9]).into(); + let weight = query + .weight(EnableScoring::disabled_from_searcher(&searcher)) + .unwrap(); + let scorer = weight.scorer(searcher.segment_reader(0), 1.0).unwrap(); + assert_eq!(scorer.size_hint(), 6); + assert_eq!(scorer.cost(), 600); + // Without a necessary condition, all docs are candidates. + let scorer = even_doc_id_query() + .scorer(searcher.segment_reader(0), 1.0) + .unwrap(); + assert_eq!(scorer.size_hint(), 10); + assert_eq!(scorer.cost(), 1000); + } + + #[test] + fn test_necessary_condition_scorer_danger() { + let index = create_index_for_test(10); + let searcher = index.reader().unwrap().searcher(); + let segment_reader = searcher.segment_reader(0); + let scorer_danger = |necessary_condition: Vec, target: DocId| { + let predicate = EvenWithNecessaryCondition::new(necessary_condition); + let num_evals = predicate.num_evals.clone(); + let query: DocPredicateQuery = predicate.into(); + let (seek_result, scorer) = query.scorer_danger(segment_reader, target, 1.0).unwrap(); + (seek_result, scorer, num_evals.load(AtomicOrdering::Relaxed)) + }; + + // The necessary condition starts after the target. + let (seek_result, mut scorer, num_evals) = scorer_danger(vec![4, 6], 1); + assert_eq!(seek_result, SeekDangerResult::SeekLowerBound(4)); + assert_eq!(num_evals, 0); + assert_eq!(scorer.seek_danger(4), SeekDangerResult::Found); + assert_eq!(scorer.doc(), 4); + + // The target is the first candidate, and matches. + let (seek_result, scorer, _) = scorer_danger(vec![2, 6], 2); + assert_eq!(seek_result, SeekDangerResult::Found); + assert_eq!(scorer.doc(), 2); + + // The target is a candidate, but does not match. + let (seek_result, mut scorer, num_evals) = scorer_danger(vec![1, 3, 4], 3); + assert_eq!(seek_result, SeekDangerResult::SeekLowerBound(4)); + assert_eq!(num_evals, 1); + assert_eq!(scorer.seek_danger(4), SeekDangerResult::Found); + + // The target is not a candidate: it is not evaluated. + let (seek_result, _, num_evals) = scorer_danger(vec![1, 6], 2); + assert_eq!(seek_result, SeekDangerResult::SeekLowerBound(6)); + assert_eq!(num_evals, 0); + + // No match after the target. The lower bound can stop on a non-matching candidate. + let (seek_result, mut scorer, _) = scorer_danger(vec![1, 3], 2); + assert_eq!(seek_result, SeekDangerResult::SeekLowerBound(3)); + assert_eq!(scorer.seek_danger(3), SeekDangerResult::SeekLowerBound(4)); + assert_eq!( + scorer.seek_danger(4), + SeekDangerResult::SeekLowerBound(TERMINATED) + ); + } + + proptest! { + #[test] + fn proptest_necessary_condition_doc_set( + candidates in prop::collection::btree_set(0u32..200, 0..60), + modulo in 1u32..5, + targets in prop::collection::vec(0u32..220, 0..30), + advances in prop::collection::vec(any::(), 0..30), + ) { + let candidates: Vec = candidates.into_iter().collect(); + let expected: Vec = candidates + .iter() + .copied() + .filter(|doc| doc.is_multiple_of(modulo)) + .collect(); + let new_doc_set = || { + DocPredicateDocSet::new( + move |doc: DocId| doc.is_multiple_of(modulo), + Box::new(VecDocSet::from(candidates.clone())), + ) + }; + let first_match = |target: DocId| { + expected + .iter() + .copied() + .find(|doc| *doc >= target) + .unwrap_or(TERMINATED) + }; + + // advance + let mut doc_set = new_doc_set(); + let mut matches: Vec = Vec::new(); + while doc_set.doc() != TERMINATED { + matches.push(doc_set.doc()); + doc_set.advance(); + } + prop_assert_eq!(&matches, &expected); + + // interleaved seek and advance + let mut doc_set = new_doc_set(); + for (target, advance) in targets.iter().zip(advances.iter()) { + let target = (*target).max(doc_set.doc()); + if *advance && doc_set.doc() != TERMINATED { + let current = doc_set.doc(); + prop_assert_eq!(doc_set.advance(), first_match(current + 1)); + } else { + prop_assert_eq!(doc_set.seek(target), first_match(target)); + } + } + + // seek_danger, following its contract: strictly increasing targets, respecting + // the returned lower bounds. + let mut sorted_targets = targets.clone(); + sorted_targets.sort_unstable(); + sorted_targets.dedup(); + let mut doc_set = new_doc_set(); + let mut lower_bound = doc_set.doc(); + let mut previous_target = None; + for requested_target in sorted_targets { + let target = requested_target.max(lower_bound); + if previous_target.is_some_and(|previous| previous >= target) { + continue; + } + previous_target = Some(target); + let next_match = first_match(target); + match doc_set.seek_danger(target) { + SeekDangerResult::Found => { + prop_assert_eq!(next_match, target); + prop_assert_eq!(doc_set.doc(), target); + } + SeekDangerResult::SeekLowerBound(bound) => { + prop_assert!(next_match != target || target == TERMINATED); + prop_assert!(bound > target || target == TERMINATED); + prop_assert!(bound <= next_match); + lower_bound = bound; + } + } + } + } + } + #[test] fn test_doc_predicate_query_scorer_danger_target_past_max_doc() { let index = create_index_for_test(4); diff --git a/src/query/exist_query.rs b/src/query/exist_query.rs index fcda85fff..a0121fbe9 100644 --- a/src/query/exist_query.rs +++ b/src/query/exist_query.rs @@ -182,7 +182,7 @@ impl Weight for FastFieldExistsWeight { } } -enum ExistsColumnIndex { +pub(crate) enum ExistsColumnIndex { Optional(OptionalIndex), Multivalued(MultiValueIndex), }