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 <paul.masurel@datadoghq.com>
This commit is contained in:
Paul Masurel
2026-09-28 18:46:37 +02:00
committed by GitHub
co-authored by Paul Masurel
parent b9125aad55
commit 047464cf92
10 changed files with 1487 additions and 84 deletions
+3
View File
@@ -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"
+4
View File
@@ -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;
+902
View File
@@ -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<str>),
/// 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<VariablePresenceCondition>);
impl ConditionSet {
/// Returns the children, in canonical order.
pub fn iter(&self) -> impl Iterator<Item = &VariablePresenceCondition> {
self.0.iter()
}
}
impl VariablePresenceCondition {
pub fn all(
conditions: impl IntoIterator<Item = VariablePresenceCondition>,
) -> VariablePresenceCondition {
let mut children: Vec<VariablePresenceCondition> = 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<Item = VariablePresenceCondition>,
) -> VariablePresenceCondition {
let mut children: Vec<VariablePresenceCondition> = 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 {
VariablePresenceCondition::all(conditions)
}
fn any(conditions: Vec<VariablePresenceCondition>) -> 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<_>>(), 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<RawCondition>),
Any(Vec<RawCondition>),
}
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<VariablePresenceCondition> =
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<Value = RawCondition> {
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<Option<u8>>;
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<Value = Vec<Assignment>> {
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<String>,
number: BoxedStrategy<String>,
string: BoxedStrategy<String>,
}
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<String>, template: &'static str) -> BoxedStrategy<String> {
arg.clone()
.prop_map(move |arg| template.replace("$0", &arg))
.boxed()
}
fn binary(
left: &BoxedStrategy<String>,
right: &BoxedStrategy<String>,
template: &'static str,
) -> BoxedStrategy<String> {
(left.clone(), right.clone())
.prop_map(move |(left, right)| template.replace("$0", &left).replace("$1", &right))
.boxed()
}
fn ternary(
first: &BoxedStrategy<String>,
second: &BoxedStrategy<String>,
third: &BoxedStrategy<String>,
template: &'static str,
) -> BoxedStrategy<String> {
(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<VariableValue> = 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,
)?;
}
}
}
-1
View File
@@ -370,7 +370,6 @@ impl<'a> From<VariablePrimitiveOpt> 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};
+2 -1
View File
@@ -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);
+2 -1
View File
@@ -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 }
}
}
@@ -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<ConstOrVariableSegmentPredicate<SegmentF>> {
(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())),
})
}
}
@@ -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<dyn DocSet> = 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<dyn DocSet> {
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<Box<dyn DocSet>> = 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<Box<dyn DocSet>> = 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<bool> = unsafe { self.compiled.call(&inputs_vec.0).as_bool() };
let eval_result: Option<bool> = 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() {
+350 -61
View File
@@ -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<TSegmentDocPredicate> {
doc_predicate: TSegmentDocPredicate,
doc: DocId,
max_doc: DocId,
necessary_condition: Box<dyn DocSet>,
}
impl<TSegmentDocPredicate: SegmentDocPredicate> DocPredicateDocSet<TSegmentDocPredicate> {
/// Creates a `DocPredicateDocSet` positioned on its first matching document.
fn new(doc_predicate: TSegmentDocPredicate, necessary_condition: Box<dyn DocSet>) -> 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<dyn DocSet>,
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<TSegmentDocPredicate: SegmentDocPredicate> DocSet
for DocPredicateDocSet<TSegmentDocPredicate>
{
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<TSegmentDocPredicate: SegmentDocPredicate> DocPredicateDocSet<TSegmentDocPredicate> {
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<TDocPredicate: DocPredicate> 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<dyn Scorer>)
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<TDocPredicate: DocPredicate> 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<dyn Scorer>;
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<P: SegmentDocPredicate> {
/// 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<P: SegmentDocPredicate> From<P> for ConstOrVariableSegmentPredicate<P> {
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<dyn DocSet>,
},
}
/// 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<AtomicUsize>,
}
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<DocId>,
num_evals: Arc<AtomicUsize>,
}
impl EvenWithNecessaryCondition {
fn new(necessary_condition: Vec<DocId>) -> Self {
EvenWithNecessaryCondition {
necessary_condition,
num_evals: Arc::default(),
}
}
}
impl DocPredicate for EvenWithNecessaryCondition {
type SegmentDocPredicate = EvenDocIds;
fn doc_predicate(
&self,
_segment_reader: &SegmentReader,
) -> crate::Result<ConstOrVariableSegmentPredicate<EvenDocIds>> {
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<DocId>, 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::<bool>(), 0..30),
) {
let candidates: Vec<DocId> = candidates.into_iter().collect();
let expected: Vec<DocId> = 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<DocId> = 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);
+1 -1
View File
@@ -182,7 +182,7 @@ impl Weight for FastFieldExistsWeight {
}
}
enum ExistsColumnIndex {
pub(crate) enum ExistsColumnIndex {
Optional(OptionalIndex),
Multivalued(MultiValueIndex),
}