From 97a5aa4ce45f6297cd1563b4249be2d90a41f7ab Mon Sep 17 00:00:00 2001 From: "dependabot[bot]" <49699333+dependabot[bot]@users.noreply.github.com> Date: Mon, 7 Sep 2026 20:02:30 +0000 Subject: [PATCH 01/49] Update datasketches requirement from 0.3.0 to 0.5.0 Updates the requirements on [datasketches](https://github.com/apache/datasketches-rust) to permit the latest version. - [Changelog](https://github.com/apache/datasketches-rust/blob/main/CHANGELOG.md) - [Commits](https://github.com/apache/datasketches-rust/compare/0.3.0...0.5.0) --- updated-dependencies: - dependency-name: datasketches dependency-version: 0.5.0 dependency-type: direct:production ... Signed-off-by: dependabot[bot] --- Cargo.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/Cargo.toml b/Cargo.toml index 8d74f858d..85a061c62 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -84,7 +84,7 @@ tantivy-bitpacker = { version = "0.10", path = "./bitpacker" } common = { version = "0.11", path = "./common/", package = "tantivy-common" } tokenizer-api = { version = "0.7", path = "./tokenizer-api", package = "tantivy-tokenizer-api" } sketches-ddsketch = { version = "0.4", features = ["use_serde"] } -datasketches = { version = "0.3.0", features = ["hll"] } +datasketches = { version = "0.5.0", features = ["hll"] } futures-util = { version = "0.3.28", optional = true } futures-channel = { version = "0.3.28", optional = true } fnv = "1.0.7" From 94a116f103d4a2f464de0e2035a9393b5864e31f Mon Sep 17 00:00:00 2001 From: Paul Masurel Date: Tue, 15 Sep 2026 16:28:53 +0200 Subject: [PATCH 02/49] (Calculated fields) Add feature-gated JIT expression document predicate (#3081) --- Cargo.toml | 4 + jitexpr/src/lib.rs | 2 + src/lib.rs | 2 + .../doc_predicate_query/function_predicate.rs | 8 +- .../doc_predicate_query/jitexpr_predicate.rs | 634 ++++++++++++++++++ src/query/doc_predicate_query/mod.rs | 81 ++- 6 files changed, 710 insertions(+), 21 deletions(-) create mode 100644 src/query/doc_predicate_query/jitexpr_predicate.rs diff --git a/Cargo.toml b/Cargo.toml index a4ff435ae..ea6abb0fa 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -76,6 +76,9 @@ measure_time = "0.9.0" arc-swap = "1.5.0" bon = "3.3.1" +# EXPERIMENTAL. The API is likely to change in the near future. +jitexpr = { version = "0.1", path = "./jitexpr", optional = true } + columnar = { version = "0.7", path = "./columnar", package = "tantivy-columnar" } sstable = { version = "0.7", path = "./sstable", package = "tantivy-sstable", optional = true } stacker = { version = "0.7", path = "./stacker", package = "tantivy-stacker" } @@ -149,6 +152,7 @@ failpoints = ["fail", "fail/failpoints"] unstable = [] # useful for benches. quickwit = ["sstable", "futures-util", "futures-channel"] +jitexpr = ["dep:jitexpr"] # Compares only the hash of a string when indexing data. # Increases indexing speed, but may lead to extremely rare missing terms, when there's a hash collision. diff --git a/jitexpr/src/lib.rs b/jitexpr/src/lib.rs index b6aa1517c..9d9d88dbf 100644 --- a/jitexpr/src/lib.rs +++ b/jitexpr/src/lib.rs @@ -1,3 +1,5 @@ +//! EXPERIMENTAL. The API is likely to change in the near future. + pub mod ast; pub mod compile; pub mod types; diff --git a/src/lib.rs b/src/lib.rs index 53a4fb10f..7537e3cc5 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -220,6 +220,8 @@ use std::fmt; pub use census::{Inventory, TrackedObject}; pub use common::{f64_to_u64, i64_to_u64, u64_to_f64, u64_to_i64, HasLen}; +#[cfg(feature = "jitexpr")] +pub use jitexpr; use once_cell::sync::Lazy; use serde::{Deserialize, Serialize}; diff --git a/src/query/doc_predicate_query/function_predicate.rs b/src/query/doc_predicate_query/function_predicate.rs index ac7192580..5d47e8282 100644 --- a/src/query/doc_predicate_query/function_predicate.rs +++ b/src/query/doc_predicate_query/function_predicate.rs @@ -1,5 +1,6 @@ use super::{DocPredicate, SegmentDocPredicate}; use crate::index::SegmentReader; +use crate::query::doc_predicate_query::ConstOrVariableSegmentPredicate; use crate::DocId; /// Blanket [`SegmentDocPredicate`] implementation for any per-document @@ -48,8 +49,11 @@ where { type SegmentDocPredicate = SegmentF; - fn doc_predicate(&self, segment_reader: &SegmentReader) -> crate::Result { - (self.segment_predicate_factory)(segment_reader) + fn doc_predicate( + &self, + segment_reader: &SegmentReader, + ) -> crate::Result> { + (self.segment_predicate_factory)(segment_reader).map(ConstOrVariableSegmentPredicate::from) } } diff --git a/src/query/doc_predicate_query/jitexpr_predicate.rs b/src/query/doc_predicate_query/jitexpr_predicate.rs new file mode 100644 index 000000000..09b1c2cfa --- /dev/null +++ b/src/query/doc_predicate_query/jitexpr_predicate.rs @@ -0,0 +1,634 @@ +use std::collections::HashMap; +use std::io; + +use columnar::{ColumnType, DynamicColumn, StrColumn}; +use jitexpr::ast::{infer_types_with_target, InferredTypeSet, TypeError, UntypedExpr}; +use jitexpr::compile::{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}; + +/// A [`DocPredicate`] that evaluates a boolean JIT expression against fast fields. +/// +/// Requires the `jitexpr` feature. Variable names are resolved as fast-field names +/// for each segment, supporting boolean, numeric, and string columns. Missing or +/// incompatible columns are left unbound, so the compiler treats them as `None`. +/// +/// For bound columns, multivalued documents contribute their first value, and a +/// document missing any input will still be evaluated with null in place of the input. +/// +/// We "fast path" cases where we detect the expression will always evaluate to true or false +/// on a segment. (e.g. if all variable columns are missing). +/// +/// 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(); +/// ``` +#[derive(Clone, Debug)] +pub struct JitExprPredicate { + expression: UntypedExpr, + inferred_inputs: Vec<(String, InferredTypeSet)>, +} + +impl JitExprPredicate { + /// Creates a predicate after inferring its inputs and requiring a boolean result. + pub fn new(expression: UntypedExpr) -> Result { + let inferred_types: HashMap<&str, InferredTypeSet> = + infer_types_with_target(&expression, InferredTypeSet::BOOLEAN)?; + let inferred_inputs: Vec<(String, InferredTypeSet)> = inferred_types + .into_iter() + .map(|(name, types)| (name.to_string(), types)) + .collect(); + Ok(Self { + expression, + inferred_inputs, + }) + } + + /// Returns the expression evaluated by this predicate. + pub fn expression(&self) -> &UntypedExpr { + &self.expression + } +} + +impl DocPredicate for JitExprPredicate { + type SegmentDocPredicate = JitExprEvalState; + + fn doc_predicate( + &self, + segment_reader: &SegmentReader, + ) -> crate::Result> { + let mut variable_types = HashMap::with_capacity(self.inferred_inputs.len()); + let mut opened_columns: HashMap<&str, DynamicColumn> = + HashMap::with_capacity(self.inferred_inputs.len()); + + // We pick a single column for each variable name. NOTE this CAN yield to unexpected results + // for some expression (e.g. (IS_NULL "mycol")). + // For instance, a document could be matching in one segment, and not matching if it + // was in another segment, just because the presence of column with the same name + // and different type could interfere. + for (name, accepted_types) in &self.inferred_inputs { + let Some(column) = open_input_column(segment_reader, name, *accepted_types)? else { + // If we do not have a valid column for that expression, we do not + // fill the HashMap at all. + // + // The compiler will replace the expression and make it behave like the null + // literal. + continue; + }; + let Some(var_type) = var_type_for_column_type(column.column_type()) else { + continue; + }; + variable_types.insert(name.as_str(), var_type); + opened_columns.insert(name.as_str(), column); + } + + let compiled_fn = + compile(&self.expression, &variable_types).map_err(|compilation_err| { + TantivyError::InvalidArgument(format!( + "the expression compilation failed {:?}. error: {compilation_err}", + self.expression + )) + })?; + + // If the function is by nature const, or if all of its variable are known to be null + // (because we don't have such columns), then we eval the value only once optimize + // + // TODO optimize further when the columns have a single value (full + min/max value) + if variable_types.is_empty() || compiled_fn.inputs().is_empty() { + // We have no variables! + // This means eventual inputs are not. Let's return a const predicate. + let inputs: Vec = + std::iter::repeat_n(VariableValue::none(), compiled_fn.inputs().len()).collect(); + let mut string_arena = StringArena::default(); + let result = unsafe { compiled_fn.call(&inputs[..], &mut string_arena) }; + let const_bool = unsafe { result.as_bool() }.unwrap_or(false); + return Ok(ConstOrVariableSegmentPredicate::Const(const_bool)); + } + + // We ended up with an expression that could not resolve to anything apparently. + if compiled_fn.result_type() == VarType::None { + return Ok(ConstOrVariableSegmentPredicate::Const(false)); + } + + if compiled_fn.result_type() != VarType::Bool { + // This should never happen: we passed a target inferred type of Bool, + // so we should have either Bool or None. + return Err(TantivyError::InvalidArgument(format!( + "the expression is not a predicate {}", + self.expression + ))); + } + + // The compiler owns the definitive ABI order. Do not rely on inference + // or HashMap iteration order when building the argument slots. + let mut columns_opt = Vec::with_capacity(compiled_fn.inputs().len()); + for input in compiled_fn.inputs() { + let column_opt: Option = + opened_columns.remove(input.variable_name.as_ref()); + if let Some(column) = column_opt { + if var_type_for_column_type(column.column_type()) != Some(input.r#type) { + return Err(TantivyError::InternalError(format!( + "compiled input `{}` expects {:?}, but its column has type {}", + input.variable_name, + input.r#type, + column.column_type() + ))); + } + columns_opt.push(Some(column)); + } else { + columns_opt.push(None); + } + } + // There is one reusable buffer per string column, in ABI order. + let num_string_inputs = columns_opt + .iter() + .filter(|column_opt| matches!(column_opt, Some(DynamicColumn::Str(_)))) + .count(); + let num_inputs = columns_opt.len(); + Ok(JitExprEvalState { + compiled: compiled_fn.context(), + columns_opt, + string_inputs: vec![String::new(); num_string_inputs], + input_values: Vec::with_capacity(num_inputs), + } + .into()) + } +} + +fn open_input_column( + reader: &SegmentReader, + name: &str, + accepted_types: InferredTypeSet, +) -> io::Result> { + let Ok(column_handles) = reader.fast_fields().dynamic_column_handles(name) else { + // If the call to dynamic_column_handles fails (for instance because the column is not a + // fast field) we choose to act as if the column was absent. + return Ok(None); + }; + for handle in column_handles { + // We return the first column that could be accepted + let Some(var_type) = var_type_for_column_type(handle.column_type()) else { + continue; + }; + if accepted_types.contains(var_type) { + return Ok(Some(handle.open()?)); + } + } + Ok(None) +} + +fn var_type_for_column_type(column_type: ColumnType) -> Option { + match column_type { + ColumnType::Bool => Some(VarType::Bool), + ColumnType::I64 => Some(VarType::I64), + ColumnType::U64 => Some(VarType::U64), + ColumnType::F64 => Some(VarType::F64), + ColumnType::Str => Some(VarType::Str), + ColumnType::Bytes | ColumnType::IpAddr | ColumnType::DateTime => None, + } +} + +/// The [`SegmentDocPredicate`] produced by [`JitExprPredicate`] for one segment. +pub struct JitExprEvalState { + compiled: CompiledFnCtx, + columns_opt: Vec>, + // One reusable buffer per string column. + string_inputs: Vec, + // Reusable argument slots. + // + // Hidden contract: this vector is always empty between evaluations, so the + // `'static` lifetime is a placeholder for an unused element type rather + // than a claim about any stored string. Only its capacity carries over, + // which is what makes restoring the `'static` type after an evaluation + // sound. `eval` is responsible for upholding this on every return path. + input_values: Vec>, +} + +/// A wrapper to make sure the variable value buffer is cleared even if the evaluation +/// panicked. +struct ClearOnDrop<'a>(&'a mut Vec>); + +impl<'a> ClearOnDrop<'a> { + fn wrap(input_values: &'a mut Vec>) -> Self { + debug_assert!(input_values.is_empty()); + // Input_values is just a buffer we share to avoid allocations + let lower_lifetime_input_values: &mut Vec> = + unsafe { std::mem::transmute(input_values) }; + ClearOnDrop(lower_lifetime_input_values) + } +} + +impl<'a> Drop for ClearOnDrop<'a> { + fn drop(&mut self) { + self.0.clear(); + } +} + +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); + + fill_input_values( + &self.columns_opt, + &mut self.string_inputs, + &mut 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() }; + + eval_result == Some(true) + } +} + +fn fill_input_values<'buffer>( + columns: &[Option], + string_inputs: &'buffer mut [String], + input_values: &mut Vec>, + doc_id: DocId, +) { + debug_assert!(input_values.is_empty()); + let mut string_inputs = string_inputs.iter_mut(); + for column_opt in columns { + let Some(column) = column_opt else { + // The full column is absent. We treat it as None. + input_values.push(VariableValue::none()); + continue; + }; + let input: Option = match column { + DynamicColumn::Bool(column) => column.first(doc_id).map(VariableValue::from), + DynamicColumn::I64(column) => column.first(doc_id).map(VariableValue::from), + DynamicColumn::U64(column) => column.first(doc_id).map(VariableValue::from), + DynamicColumn::F64(column) => column.first(doc_id).map(VariableValue::from), + DynamicColumn::Str(column) => { + let string_input = string_inputs + .next() + .expect("every string column has a string input buffer"); + load_str_input(column, doc_id, string_input).map(VariableValue::from) + } + DynamicColumn::Bytes(_) | DynamicColumn::IpAddr(_) | DynamicColumn::DateTime(_) => { + unreachable!("unsupported columns are filtered before compilation") + } + }; + // If the value is someone absent, we set the input to none/null. + input_values.push(input.unwrap_or(VariableValue::none())); + } +} + +/// Loads the first value of a string column for `doc_id` into `buffer`. +/// +/// Missing values return None. +/// +/// This function may panic if the dictionary is corrupted or if the column +/// contains term ords that do not exist in the dictionary. +fn load_str_input<'buffer>( + column: &StrColumn, + doc_id: DocId, + buffer: &'buffer mut String, +) -> Option<&'buffer str> { + buffer.clear(); + let term_ord = column.ords().first(doc_id)?; + // SegmentDocPredicate::eval cannot return I/O errors; an unreadable + // dictionary therefore panics. + // TODO this is terribly inefficient: we need at least some caching. + let found = column + .ord_to_str(term_ord, buffer) + .expect("fast-field string dictionary is corrupted"); + assert!(found); + Some(buffer.as_str()) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::collector::Count; + use crate::query::doc_predicate_query::DocPredicateQuery; + use crate::schema::{Schema, FAST, STORED, STRING}; + use crate::Index; + + fn create_index() -> Index { + let mut schema_builder = Schema::builder(); + let number = schema_builder.add_u64_field("number", FAST); + let flag = schema_builder.add_bool_field("flag", FAST); + let _notfast = schema_builder.add_bool_field("notfast", STORED); + let label = schema_builder.add_text_field("label", STRING | FAST); + let index = Index::create_in_ram(schema_builder.build()); + let mut writer = index.writer_for_tests().unwrap(); + writer + .add_document(doc!(number => 1u64, flag => true, label => "one")) + .unwrap(); + writer + .add_document(doc!(number => 2u64, flag => false, label => "two")) + .unwrap(); + writer + .add_document(doc!(number => 3u64, flag => true, label => "three")) + .unwrap(); + writer.add_document(doc!(number => 4u64)).unwrap(); + writer.commit().unwrap(); + index + } + + fn query(expression: &str) -> DocPredicateQuery { + JitExprPredicate::new(jitexpr::ast::deserialize(expression).unwrap()) + .unwrap() + .into() + } + + #[test] + fn test_constructor_requires_boolean_expression() { + let expression = jitexpr::ast::deserialize("(ADD number 1u64)").unwrap(); + assert!(JitExprPredicate::new(expression).is_err()); + } + + #[test] + fn test_simple() { + let index = create_index(); + let searcher = index.reader().unwrap().searcher(); + assert_eq!( + searcher.search(&query("(EQ number 2i64)"), &Count).unwrap(), + 1 + ); + assert_eq!( + searcher + .search(&query("(EQ (ADD number 1u64) 3u64)"), &Count) + .unwrap(), + 1 + ); + } + + #[test] + fn test_simple_string_ref_predicate() { + let index = create_index(); + let searcher = index.reader().unwrap().searcher(); + assert_eq!( + searcher + .search( + &query(r#"(EQ (REGEXP_EXTRACT label "(.).*" 1u64) "o")"#), + &Count + ) + .unwrap(), + 1 + ); + } + + #[test] + fn test_simple_built_string_predicate() { + let index = create_index(); + let searcher = index.reader().unwrap().searcher(); + assert_eq!( + searcher + .search(&query(r#"(EQ (UPPER label) "TWO")"#), &Count) + .unwrap(), + 1 + ); + } + + #[test] + fn test_simple_missing_field() { + let index = create_index(); + let searcher = index.reader().unwrap().searcher(); + assert_eq!( + searcher + .search(&query(r#"(EQ missing_field true)"#), &Count) + .unwrap(), + 0 + ); + } + + #[test] + fn test_simple_missing_field_is_not_null() { + let index = create_index(); + let searcher = index.reader().unwrap().searcher(); + assert_eq!( + searcher + .search(&query(r#"(IS_NOT_NULL missing_field)"#), &Count) + .unwrap(), + 0 + ); + } + + #[test] + fn test_simple_missing_field_is_null() { + let index = create_index(); + let searcher = index.reader().unwrap().searcher(); + assert_eq!( + searcher + .search(&query(r#"(IS_NULL missing_field)"#), &Count) + .unwrap(), + 4 + ); + } + + #[test] + fn test_simple_notfast() { + let index = create_index(); + let searcher = index.reader().unwrap().searcher(); + assert_eq!( + searcher + .search(&query(r#"(EQ notfast true)"#), &Count) + .unwrap(), + 0 + ); + } + + #[test] + fn test_boolean_and_string_inputs_follow_compiled_order() { + let index = create_index(); + let searcher = index.reader().unwrap().searcher(); + assert_eq!(searcher.search(&query("flag"), &Count).unwrap(), 2); + assert_eq!( + searcher + .search(&query(r#"(EQ label "three")"#), &Count) + .unwrap(), + 1 + ); + // Inference sorts names, but the ABI follows expression order: label, flag. + assert_eq!( + searcher + .search(&query(r#"(EQ (EQ label "three") flag)"#), &Count) + .unwrap(), + 2 + ); + } + + #[test] + fn test_constant_predicates() { + let index = create_index(); + let searcher = index.reader().unwrap().searcher(); + assert_eq!(searcher.search(&query("true"), &Count).unwrap(), 4); + assert_eq!(searcher.search(&query("false"), &Count).unwrap(), 0); + } + + #[test] + fn test_missing_and_incompatible_columns() { + let index = create_index(); + let searcher = index.reader().unwrap().searcher(); + assert_eq!(searcher.search(&query("missing"), &Count).unwrap(), 0); + + // A column missing from the segment is compiled as None. + assert_eq!( + searcher + .search(&query("(IS_NULL missing)"), &Count) + .unwrap(), + 4 + ); + // A column missing from the segment is compiled as None. + assert_eq!( + searcher + .search(&query("(IS_NOT_NULL missing)"), &Count) + .unwrap(), + 0 + ); + assert_eq!( + searcher.search(&query("(IS_NULL flag)"), &Count).unwrap(), + 1 + ); + assert_eq!( + searcher + .search(&query("(IS_NOT_NULL flag)"), &Count) + .unwrap(), + 3 + ); + assert_eq!( + searcher + .search(&query("(EQ (ADD label 1i64) 2i64)"), &Count) + .unwrap(), + 0 + ); + assert_eq!( + searcher + .search(&query("(IS_NULL (ADD label 1i64))"), &Count) + .unwrap(), + 4 + ); + } + + #[test] + fn test_signed_and_float_columns() { + let mut schema_builder = Schema::builder(); + let signed = schema_builder.add_i64_field("signed", FAST); + let float = schema_builder.add_f64_field("float", FAST); + let index = Index::create_in_ram(schema_builder.build()); + let mut writer = index.writer_for_tests().unwrap(); + writer + .add_document(doc!(signed => -2i64, float => 1.5f64)) + .unwrap(); + writer + .add_document(doc!(signed => 3i64, float => 2.5f64)) + .unwrap(); + writer.add_document(doc!()).unwrap(); + writer.commit().unwrap(); + let searcher = index.reader().unwrap().searcher(); + assert_eq!( + searcher + .search(&query("(EQ signed -2i64)"), &Count) + .unwrap(), + 1 + ); + assert_eq!( + searcher + .search(&query("(EQ signed -2f64)"), &Count) + .unwrap(), + 1 + ); + assert_eq!( + searcher + .search(&query("(EQ float 1.5f64)"), &Count) + .unwrap(), + 1 + ); + assert_eq!( + searcher.search(&query("(EQ float 1i64)"), &Count).unwrap(), + 0 + ); + } + + #[test] + fn test_multivalued_columns_use_first_value() { + let mut schema_builder = Schema::builder(); + let number = schema_builder.add_u64_field("number", FAST); + let label = schema_builder.add_text_field("label", STRING | FAST); + let index = Index::create_in_ram(schema_builder.build()); + let mut writer = index.writer_for_tests().unwrap(); + writer + .add_document(doc!(number => 1u64, number => 2u64, + label => "first", label => "second")) + .unwrap(); + writer + .add_document(doc!(number => 2u64, number => 1u64, + label => "second", label => "first")) + .unwrap(); + writer.add_document(doc!()).unwrap(); + writer.commit().unwrap(); + let searcher = index.reader().unwrap().searcher(); + assert_eq!( + searcher.search(&query("(EQ number 1u64)"), &Count).unwrap(), + 1 + ); + assert_eq!( + searcher + .search(&query(r#"(EQ label "first")"#), &Count) + .unwrap(), + 1 + ); + } + + // THIS FAILS! due to our pick best possible column approach policy. + // #[test] + // fn test_multi_typed_field_picks_one() { + // let mut schema_builder = Schema::builder(); + // let json = schema_builder.add_json_field("json", FAST); + // let index = Index::create_in_ram(schema_builder.build()); + // let mut writer = index.writer_for_tests().unwrap(); + // writer + // .add_document(doc!(json => serde_json::json!({"myfield": 2u64}))) + // .unwrap(); + // writer + // .add_document(doc!(json => serde_json::json!({"myfield": "b"}))) + // .unwrap(); + // writer.commit().unwrap(); + // let searcher = index.reader().unwrap().searcher(); + // assert_eq!( + // searcher + // .search(&query(r#"(IS_NULL json.myfield)"#), &Count) + // .unwrap(), + // 2 // assertion fails, expected 2 got 1 + // ); + // } + + // THIS FAILS DUE TO EQ infer_types being too lenient. + // #[test] + // fn test_multi_typed_field_eq_too_lenient_failing() { + // let mut schema_builder = Schema::builder(); + // let json = schema_builder.add_json_field("json", FAST); + // let index = Index::create_in_ram(schema_builder.build()); + // let mut writer = index.writer_for_tests().unwrap(); + // writer + // .add_document(doc!(json => serde_json::json!({"myfield": 2u64}))) + // .unwrap(); + // writer + // .add_document(doc!(json => serde_json::json!({"myfield": "b"}))) + // .unwrap(); + // writer.commit().unwrap(); + // let searcher = index.reader().unwrap().searcher(); + // assert_eq!( + // searcher + // .search(&query(r#"(EQ json.myfield "b")"#), &Count) + // .unwrap(), + // 1 // assertion fails, expected 2 got 1 + // ); + // } +} diff --git a/src/query/doc_predicate_query/mod.rs b/src/query/doc_predicate_query/mod.rs index 5346253de..11e79dfd9 100644 --- a/src/query/doc_predicate_query/mod.rs +++ b/src/query/doc_predicate_query/mod.rs @@ -1,13 +1,19 @@ use std::sync::Arc; mod function_predicate; +#[cfg(feature = "jitexpr")] +mod jitexpr_predicate; pub use function_predicate::FunctionPredicate; +#[cfg(feature = "jitexpr")] +pub use jitexpr_predicate::{JitExprEvalState, JitExprPredicate}; use crate::docset::{SeekDangerResult, TERMINATED}; use crate::index::SegmentReader; use crate::query::explanation::does_not_match; -use crate::query::{ConstScorer, EnableScoring, Explanation, Query, Scorer, Weight}; +use crate::query::{ + AllWeight, ConstScorer, EmptyWeight, EnableScoring, Explanation, Query, Scorer, Weight, +}; use crate::{DocId, DocSet, Score}; /// A query that evaluates, for each DocId, whether it matches or not. @@ -141,14 +147,25 @@ pub trait DocPredicateBoxable: std::fmt::Debug + 'static + Send + Sync { impl DocPredicateBoxable for TDocPredicate { fn scorer(&self, segment_reader: &SegmentReader, boost: f32) -> crate::Result> { - let doc_predicate = self.doc_predicate(segment_reader)?; - 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) + let const_or_variable_segment_predicate = self.doc_predicate(segment_reader)?; + match const_or_variable_segment_predicate { + ConstOrVariableSegmentPredicate::Const(always_match) => { + if always_match { + AllWeight.scorer(segment_reader, boost) + } else { + 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) + } + } } fn scorer_danger( @@ -157,15 +174,41 @@ impl DocPredicateBoxable for TDocPredicate { target: DocId, boost: f32, ) -> crate::Result<(SeekDangerResult, Box)> { - let doc_predicate = self.doc_predicate(segment_reader)?; - 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)) + let const_or_variable_segment_predicate = self.doc_predicate(segment_reader)?; + match const_or_variable_segment_predicate { + ConstOrVariableSegmentPredicate::Const(always_match) => { + if always_match { + AllWeight.scorer_danger(segment_reader, target, boost) + } else { + 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)) + } + } + } +} + +/// Represents a segment predicate. +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) } } @@ -186,7 +229,7 @@ pub trait DocPredicate: Send + Sync + 'static + std::fmt::Debug { fn doc_predicate( &self, segment_reader: &SegmentReader, - ) -> crate::Result; + ) -> crate::Result>; } /// The per-segment predicate produced by a [`DocPredicate`]. From 34edbf782197c5be52d844695d25cf3621bbdcb8 Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Sun, 13 Sep 2026 18:23:36 +0200 Subject: [PATCH 03/49] Fix query parser panics on malformed input Treat a fieldless exists leaf as match-all and make the lenient set parser stop when parsing no longer consumes input. Add public QueryParser proptests covering both regressions. Findings and original fixes by Oleksii Syniakov in osyniakov/tantivy#4. Co-authored-by: Oleksii Syniakov <1282756+osyniakov@users.noreply.github.com> --- query-grammar/src/query_grammar.rs | 14 +++++++++++--- query-grammar/src/user_input_ast.rs | 5 +++-- src/query/query_parser/query_parser.rs | 25 +++++++++++++++++++++++++ 3 files changed, 39 insertions(+), 5 deletions(-) diff --git a/query-grammar/src/query_grammar.rs b/query-grammar/src/query_grammar.rs index aaa7800e8..ca3905f75 100644 --- a/query-grammar/src/query_grammar.rs +++ b/query-grammar/src/query_grammar.rs @@ -687,12 +687,20 @@ fn set_infallible(mut inp: &str) -> JResult<&str, UserInputLeaf> { return Ok((inp, (res, errs))); } errs.append(&mut space_error); - // TODO - // here we do the assumption term_or_phrase_infallible always consume something if the - // first byte is not `)` or ' '. If it did not, we would end up looping. let (rest, (delim_term, mut err)) = simple_term_infallible("]")(inp)?; errs.append(&mut err); + if rest.len() == inp.len() { + errs.push(LenientErrorInternal { + pos: inp.len(), + message: "missing ]".to_string(), + }); + let res = UserInputLeaf::Set { + field: None, + elements, + }; + return Ok((inp, (res, errs))); + } if let Some((_, term)) = delim_term { elements.push(term); } diff --git a/query-grammar/src/user_input_ast.rs b/query-grammar/src/user_input_ast.rs index b607e56cd..0451b244d 100644 --- a/query-grammar/src/user_input_ast.rs +++ b/query-grammar/src/user_input_ast.rs @@ -47,8 +47,9 @@ impl UserInputLeaf { upper, }, UserInputLeaf::Set { field: _, elements } => UserInputLeaf::Set { field, elements }, - UserInputLeaf::Exists { field: _ } => UserInputLeaf::Exists { - field: field.expect("Exist query without a field isn't allowed"), + UserInputLeaf::Exists { field: _ } => match field { + Some(field) => UserInputLeaf::Exists { field }, + None => UserInputLeaf::All, }, UserInputLeaf::Regex { field: _, pattern } => UserInputLeaf::Regex { field, pattern }, } diff --git a/src/query/query_parser/query_parser.rs b/src/query/query_parser/query_parser.rs index bc97f1456..f4a592f5a 100644 --- a/src/query/query_parser/query_parser.rs +++ b/src/query/query_parser/query_parser.rs @@ -1106,6 +1106,7 @@ fn convert_to_query(fuzzy: &FxHashMap, logical_ast: LogicalAst) -> #[cfg(test)] mod test { use matches::assert_matches; + use proptest::prelude::*; use super::super::logical_ast::*; use super::{QueryParser, QueryParserError}; @@ -1171,6 +1172,30 @@ mod test { make_query_parser_with_default_fields(&["title", "text"]) } + proptest! { + #[test] + fn test_query_parser_does_not_panic_after_match_all( + suffix in prop::sample::select(vec!['\u{b}', '\u{c}', '\u{85}']) + ) { + let query_parser = make_query_parser(); + let query = format!("*{suffix}"); + let _ = query_parser.parse_query(&query); + let _ = query_parser.parse_query_lenient(&query); + } + + #[test] + fn test_lenient_query_parser_makes_progress_in_invalid_sets( + field in proptest::option::of("[a-z]{1,4}"), + invalid_char in prop::sample::select(vec!['\0', '\u{b}', '\u{c}', '\u{7f}']), + ) { + let query_parser = make_query_parser(); + let field = field.map(|field| format!("{field}:")).unwrap_or_default(); + let query = format!("{field}IN [{invalid_char}"); + let (_, errors) = query_parser.parse_query_lenient(&query); + prop_assert!(!errors.is_empty()); + } + } + fn parse_query_to_logical_ast_with_default_fields( query: &str, default_conjunction: bool, From 54b82e582fdc500f2e1e8b0ef7750e36c45d6e77 Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Sun, 13 Sep 2026 19:12:01 +0200 Subject: [PATCH 04/49] Report unsupported whitespace as a syntax error --- query-grammar/src/query_grammar.rs | 17 +++++++++-------- src/query/query_parser/query_parser.rs | 5 +++-- 2 files changed, 12 insertions(+), 10 deletions(-) diff --git a/query-grammar/src/query_grammar.rs b/query-grammar/src/query_grammar.rs index ca3905f75..3d0d58b64 100644 --- a/query-grammar/src/query_grammar.rs +++ b/query-grammar/src/query_grammar.rs @@ -325,11 +325,10 @@ fn exists(inp: &str) -> IResult<&str, UserInputLeaf> { multispace0, char('*'), peek(alt(( + value("", multispace1), value( "", - satisfy(|c: char| { - c.is_whitespace() || (ESCAPE_IN_WORD.contains(&c) && c != '\\') - }), + satisfy(|c: char| ESCAPE_IN_WORD.contains(&c) && c != '\\'), ), eof, ))), @@ -345,11 +344,10 @@ fn exists_precond(inp: &str) -> IResult<&str, (), ()> { multispace0, char('*'), peek(alt(( + value("", multispace1), value( "", - satisfy(|c: char| { - c.is_whitespace() || (ESCAPE_IN_WORD.contains(&c) && c != '\\') - }), + satisfy(|c: char| ESCAPE_IN_WORD.contains(&c) && c != '\\'), ), eof, ))), // we need to check this isn't a wildcard query @@ -1133,11 +1131,14 @@ pub fn parse_to_ast(inp: &str) -> IResult<&str, UserInputAst> { } pub fn parse_to_ast_lenient(query_str: &str) -> (UserInputAst, Vec) { - if query_str.trim().is_empty() { + if query_str + .chars() + .all(|c| matches!(c, ' ' | '\t' | '\r' | '\n')) + { return (UserInputAst::Clause(Vec::new()), Vec::new()); } let (left, (res, mut errors)) = ast_infallible(query_str).unwrap(); - if !left.trim().is_empty() { + if !left.is_empty() { errors.push(LenientErrorInternal { pos: left.len(), message: "unparsed end of query".to_string(), diff --git a/src/query/query_parser/query_parser.rs b/src/query/query_parser/query_parser.rs index f4a592f5a..05de1e776 100644 --- a/src/query/query_parser/query_parser.rs +++ b/src/query/query_parser/query_parser.rs @@ -1179,8 +1179,9 @@ mod test { ) { let query_parser = make_query_parser(); let query = format!("*{suffix}"); - let _ = query_parser.parse_query(&query); - let _ = query_parser.parse_query_lenient(&query); + prop_assert!(query_parser.parse_query(&query).is_err()); + let (_, errors) = query_parser.parse_query_lenient(&query); + prop_assert!(!errors.is_empty()); } #[test] From b06d8e9a04bc47b2ee5a08fb47faf8fbaf8a821d Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Fri, 11 Sep 2026 11:27:22 +0200 Subject: [PATCH 05/49] Track active RamDirectory writer memory Include VecWriter allocations in total_mem_usage so callers can account for memory before writers are flushed or terminated. --- src/directory/ram_directory.rs | 113 +++++++++++++++++++++++---------- 1 file changed, 81 insertions(+), 32 deletions(-) diff --git a/src/directory/ram_directory.rs b/src/directory/ram_directory.rs index 7ea6db9e0..a9c5ae437 100644 --- a/src/directory/ram_directory.rs +++ b/src/directory/ram_directory.rs @@ -1,6 +1,7 @@ use std::collections::HashMap; use std::io::{self, BufWriter, Cursor, Write}; use std::path::{Path, PathBuf}; +use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::{Arc, RwLock}; use std::{fmt, result}; @@ -14,6 +15,43 @@ use crate::directory::{ WatchHandle, WritePtr, }; +const MEMORY_USAGE_UPDATE_THRESHOLD: usize = 10_000; + +struct MemoryUsageTracker { + shared_usage: Arc, + reported_bytes: usize, + unreported_bytes: usize, +} + +impl MemoryUsageTracker { + fn new(shared_usage: Arc) -> Self { + Self { + shared_usage, + reported_bytes: 0, + unreported_bytes: 0, + } + } + + fn add(&mut self, num_bytes: usize) { + self.unreported_bytes += num_bytes; + if self.unreported_bytes >= MEMORY_USAGE_UPDATE_THRESHOLD { + self.shared_usage + .fetch_add(self.unreported_bytes, Ordering::Relaxed); + self.reported_bytes += self.unreported_bytes; + self.unreported_bytes = 0; + } + } +} + +impl Drop for MemoryUsageTracker { + fn drop(&mut self) { + if self.reported_bytes > 0 { + self.shared_usage + .fetch_sub(self.reported_bytes, Ordering::Relaxed); + } + } +} + /// Writer associated with the [`RamDirectory`]. /// /// The Writer just writes a buffer. @@ -21,15 +59,19 @@ struct VecWriter { path: PathBuf, shared_directory: RamDirectory, data: Cursor>, + memory_usage: MemoryUsageTracker, is_flushed: bool, } impl VecWriter { fn new(path_buf: PathBuf, shared_directory: RamDirectory) -> VecWriter { + let memory_usage = + MemoryUsageTracker::new(Arc::clone(&shared_directory.active_writer_mem_usage)); VecWriter { path: path_buf, data: Cursor::new(Vec::new()), shared_directory, + memory_usage, is_flushed: true, } } @@ -51,7 +93,10 @@ impl Drop for VecWriter { impl Write for VecWriter { fn write(&mut self, buf: &[u8]) -> io::Result { self.is_flushed = false; + let previous_capacity = self.data.get_ref().capacity(); self.data.write_all(buf)?; + let capacity = self.data.get_ref().capacity(); + self.memory_usage.add(capacity - previous_capacity); Ok(buf.len()) } @@ -121,6 +166,7 @@ impl fmt::Debug for RamDirectory { #[derive(Clone, Default)] pub struct RamDirectory { fs: Arc>, + active_writer_mem_usage: Arc, } impl RamDirectory { @@ -129,24 +175,12 @@ impl RamDirectory { Self::default() } - /// Deep clones the directory. + /// Returns the size of the files and an estimate of active writer allocations. /// - /// Ulterior writes on one of the copy - /// will not affect the other copy. - pub fn deep_clone(&self) -> RamDirectory { - let inner_clone = InnerDirectory { - fs: self.fs.read().unwrap().fs.clone(), - watch_router: Default::default(), - }; - RamDirectory { - fs: Arc::new(RwLock::new(inner_clone)), - } - } - - /// Returns the sum of the size of the different files - /// in the [`RamDirectory`]. + /// Active writer allocations are reported in 10 kB increments. pub fn total_mem_usage(&self) -> usize { self.fs.read().unwrap().total_mem_usage() + + self.active_writer_mem_usage.load(Ordering::Relaxed) } /// Write a copy of all of the files saved in the [`RamDirectory`] in the target [`Directory`]. @@ -206,7 +240,9 @@ impl Directory for RamDirectory { if exists { Err(OpenWriteError::FileAlreadyExists(path_buf)) } else { - Ok(BufWriter::new(Box::new(vec_writer))) + // The writer's allocation is tracked by `RamDirectory`; an additional buffer would not + // be included in `total_mem_usage()`. + Ok(BufWriter::with_capacity(0, Box::new(vec_writer))) } } @@ -244,7 +280,8 @@ mod tests { use std::io::Write; use std::path::Path; - use super::RamDirectory; + use super::{RamDirectory, MEMORY_USAGE_UPDATE_THRESHOLD}; + use crate::directory::TerminatingWrite; use crate::Directory; #[test] @@ -265,21 +302,33 @@ mod tests { } #[test] - fn test_ram_directory_deep_clone() { + fn test_active_writer_memory_usage() { let dir = RamDirectory::default(); - let test = Path::new("test"); - let test2 = Path::new("test2"); - dir.atomic_write(test, b"firstwrite").unwrap(); - let dir_clone = dir.deep_clone(); - assert_eq!( - dir_clone.atomic_read(test).unwrap(), - dir.atomic_read(test).unwrap() - ); - dir.atomic_write(test, b"original").unwrap(); - dir_clone.atomic_write(test, b"clone").unwrap(); - dir_clone.atomic_write(test2, b"clone2").unwrap(); - assert_eq!(dir.atomic_read(test).unwrap(), b"original"); - assert_eq!(&dir_clone.atomic_read(test).unwrap(), b"clone"); - assert_eq!(&dir_clone.atomic_read(test2).unwrap(), b"clone2"); + let path = Path::new("file"); + let mut writer = dir.open_write(path).unwrap(); + assert_eq!(dir.total_mem_usage(), 0); + + writer.write_all(&[0u8]).unwrap(); + assert_eq!(dir.total_mem_usage(), 0); + writer + .write_all(&vec![0u8; MEMORY_USAGE_UPDATE_THRESHOLD]) + .unwrap(); + let first_capacity = dir.total_mem_usage(); + assert!(first_capacity >= MEMORY_USAGE_UPDATE_THRESHOLD); + + writer.write_all(&vec![0u8; first_capacity + 1]).unwrap(); + let grown_capacity = dir.total_mem_usage(); + assert!(grown_capacity > first_capacity); + + writer.flush().unwrap(); + let file_len = 1 + MEMORY_USAGE_UPDATE_THRESHOLD + first_capacity + 1; + assert_eq!(dir.total_mem_usage(), grown_capacity + file_len); + assert_eq!(dir.clone().total_mem_usage(), dir.total_mem_usage()); + + writer.terminate().unwrap(); + assert_eq!(dir.total_mem_usage(), file_len); + + dir.delete(path).unwrap(); + assert_eq!(dir.total_mem_usage(), 0); } } From 98d91f043632d036800877a85d81bbcc4f1bdfe8 Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Fri, 11 Sep 2026 20:25:15 +0200 Subject: [PATCH 06/49] Remove intermediate `flush` from directory contract Intermediate writes are an implementation detail. From Tantivy side we only decide when a file is finished. This change is done to remove flush overhead for VecWriter. On finalization we just move the Vec to the directory. Rename TerminatingWrite to FinishableWrite to reflect these semantics. --- benches/merge_segments.rs | 10 +-- .../src/column_index/multivalued_index.rs | 2 +- .../src/columnar/merge/merge_dict_column.rs | 2 +- columnar/src/columnar/writer/mod.rs | 2 +- common/src/lib.rs | 2 +- common/src/writer.rs | 42 +++++----- src/directory/composite_file.rs | 14 +--- src/directory/directory.rs | 35 +++----- src/directory/footer.rs | 18 ++--- src/directory/managed_directory.rs | 6 +- src/directory/mmap_directory/mod.rs | 14 ++-- src/directory/mod.rs | 4 +- src/directory/ram_directory.rs | 81 ++++++++++--------- src/directory/tests.rs | 19 ++--- src/fastfield/alive_bitset.rs | 2 +- src/fastfield/mod.rs | 26 +++--- src/fastfield/plugin.rs | 6 +- src/fieldnorm/serializer.rs | 1 - src/indexer/index_writer.rs | 4 +- src/plugin.rs | 4 +- src/positions/serializer.rs | 6 +- src/postings/serializer.rs | 1 - src/store/store_compressor.rs | 4 +- src/termdict/tests.rs | 8 +- sstable/src/index/mod.rs | 2 +- sstable/src/lib.rs | 2 +- tests/failpoints/mod.rs | 4 +- 27 files changed, 153 insertions(+), 168 deletions(-) diff --git a/benches/merge_segments.rs b/benches/merge_segments.rs index 7e911f7ec..d57fddf74 100644 --- a/benches/merge_segments.rs +++ b/benches/merge_segments.rs @@ -16,7 +16,7 @@ use rand::rngs::StdRng; use rand::SeedableRng; use tantivy::directory::error::{DeleteError, OpenReadError, OpenWriteError}; use tantivy::directory::{ - AntiCallToken, Directory, FileHandle, OwnedBytes, TerminatingWrite, WatchCallback, WatchHandle, + AntiCallToken, Directory, FileHandle, FinishableWrite, OwnedBytes, WatchCallback, WatchHandle, WritePtr, }; use tantivy::indexer::{merge_filtered_segments, NoMergePolicy}; @@ -40,8 +40,8 @@ impl Write for NullWriter { } } -impl TerminatingWrite for NullWriter { - fn terminate_ref(&mut self, _token: AntiCallToken) -> io::Result<()> { +impl FinishableWrite for NullWriter { + fn finish_ref(&mut self, _token: AntiCallToken) -> io::Result<()> { Ok(()) } } @@ -63,8 +63,8 @@ impl Write for InMemoryWriter { } } -impl TerminatingWrite for InMemoryWriter { - fn terminate_ref(&mut self, _token: AntiCallToken) -> io::Result<()> { +impl FinishableWrite for InMemoryWriter { + fn finish_ref(&mut self, _token: AntiCallToken) -> io::Result<()> { let bytes = OwnedBytes::new(std::mem::take(&mut self.buffer)); self.blobs.write().unwrap().insert(self.path.clone(), bytes); Ok(()) diff --git a/columnar/src/column_index/multivalued_index.rs b/columnar/src/column_index/multivalued_index.rs index ad7efd363..ff6cea5aa 100644 --- a/columnar/src/column_index/multivalued_index.rs +++ b/columnar/src/column_index/multivalued_index.rs @@ -33,7 +33,7 @@ pub fn serialize_multivalued_index( } = doc_ids_with_values; serialize_optional_index(&**non_null_row_ids, *num_rows, &mut count_writer)?; let optional_len = count_writer.written_bytes() as u32; - let output = count_writer.finish(); + let output = count_writer.into_inner(); serialize_u64_based_column_values( &**start_offsets, &[CodecType::Bitpacked, CodecType::Linear], diff --git a/columnar/src/columnar/merge/merge_dict_column.rs b/columnar/src/columnar/merge/merge_dict_column.rs index ae2724b02..768b19d36 100644 --- a/columnar/src/columnar/merge/merge_dict_column.rs +++ b/columnar/src/columnar/merge/merge_dict_column.rs @@ -22,7 +22,7 @@ pub fn merge_bytes_or_str_column( // TODO !!! Remove useless terms. let term_ord_mapping = serialize_merged_dict(bytes_columns, merge_row_order, &mut output)?; let dictionary_num_bytes: u32 = output.written_bytes() as u32; - let output = output.finish(); + let output = output.into_inner(); let remapped_term_ordinals_values = RemappedTermOrdinalsValues { bytes_columns, term_ord_mapping: &term_ord_mapping, diff --git a/columnar/src/columnar/writer/mod.rs b/columnar/src/columnar/writer/mod.rs index 999ccd058..8f803f1b0 100644 --- a/columnar/src/columnar/writer/mod.rs +++ b/columnar/src/columnar/writer/mod.rs @@ -554,7 +554,7 @@ fn serialize_bytes_or_str_column( let term_id_mapping: TermIdMapping = dictionary_builder.serialize(arena, &mut counting_writer)?; let dictionary_num_bytes: u32 = counting_writer.written_bytes() as u32; - let mut wrt = counting_writer.finish(); + let mut wrt = counting_writer.into_inner(); let operation_iterator = operation_it.map(|symbol: ColumnOperation| { // We map unordered ids to ordered ids. match symbol { diff --git a/common/src/lib.rs b/common/src/lib.rs index 4e64af11c..f97c92d7f 100644 --- a/common/src/lib.rs +++ b/common/src/lib.rs @@ -24,7 +24,7 @@ pub use serialize::{BinarySerializable, DeserializeFrom, FixedSize}; pub use vint::{ VInt, VIntU128, read_u32_vint, read_u32_vint_no_advance, serialize_vint_u32, write_u32_vint, }; -pub use writer::{AntiCallToken, CountingWriter, TerminatingWrite}; +pub use writer::{AntiCallToken, CountingWriter, FinishableWrite}; /// Has length trait pub trait HasLen { diff --git a/common/src/writer.rs b/common/src/writer.rs index 8056b11d8..625bddc6f 100644 --- a/common/src/writer.rs +++ b/common/src/writer.rs @@ -21,7 +21,7 @@ impl CountingWriter { /// Returns the underlying write object. /// Note that this method does not trigger any flushing. #[inline] - pub fn finish(self) -> W { + pub fn into_inner(self) -> W { self.underlying } } @@ -47,15 +47,15 @@ impl Write for CountingWriter { } } -impl TerminatingWrite for CountingWriter { +impl FinishableWrite for CountingWriter { #[inline] - fn terminate_ref(&mut self, token: AntiCallToken) -> io::Result<()> { - self.underlying.terminate_ref(token) + fn finish_ref(&mut self, token: AntiCallToken) -> io::Result<()> { + self.underlying.finish_ref(token) } } /// Struct used to prevent from calling -/// [`terminate_ref`](TerminatingWrite::terminate_ref) directly +/// [`finish_ref`](FinishableWrite::finish_ref) directly /// /// The point is that while the type is public, it cannot be built by anyone /// outside of this module. @@ -64,33 +64,35 @@ pub struct AntiCallToken(()); /// Trait used to indicate when no more write need to be done on a writer /// /// Thread-safety is enforced at the call sites that require it. -pub trait TerminatingWrite: Write { - /// Indicate that the writer will no longer be used. Internally call terminate_ref. - fn terminate(mut self) -> io::Result<()> +pub trait FinishableWrite: Write { + /// Finishes the writer, consuming it. Internally calls [`FinishableWrite::finish_ref`]. + fn finish(mut self) -> io::Result<()> where Self: Sized { - self.terminate_ref(AntiCallToken(())) + self.finish_ref(AntiCallToken(())) } /// You should implement this function to define custom behavior. /// This function should flush any buffer it may hold. - fn terminate_ref(&mut self, _: AntiCallToken) -> io::Result<()>; + fn finish_ref(&mut self, _: AntiCallToken) -> io::Result<()>; } -impl TerminatingWrite for Box { - fn terminate_ref(&mut self, token: AntiCallToken) -> io::Result<()> { - self.as_mut().terminate_ref(token) +impl FinishableWrite for Box { + fn finish_ref(&mut self, token: AntiCallToken) -> io::Result<()> { + self.as_mut().finish_ref(token) } } -impl TerminatingWrite for BufWriter { - fn terminate_ref(&mut self, a: AntiCallToken) -> io::Result<()> { - self.flush()?; - self.get_mut().terminate_ref(a) +impl FinishableWrite for BufWriter { + fn finish_ref(&mut self, a: AntiCallToken) -> io::Result<()> { + if !self.buffer().is_empty() { + self.flush()?; + } + self.get_mut().finish_ref(a) } } -impl TerminatingWrite for &mut Vec { - fn terminate_ref(&mut self, _a: AntiCallToken) -> io::Result<()> { +impl FinishableWrite for &mut Vec { + fn finish_ref(&mut self, _a: AntiCallToken) -> io::Result<()> { self.flush() } } @@ -109,7 +111,7 @@ mod test { let bytes = (0u8..10u8).collect::>(); counting_writer.write_all(&bytes).unwrap(); let len = counting_writer.written_bytes(); - let buffer_restituted: Vec = counting_writer.finish(); + let buffer_restituted: Vec = counting_writer.into_inner(); assert_eq!(len, 10u64); assert_eq!(buffer_restituted.len(), 10); } diff --git a/src/directory/composite_file.rs b/src/directory/composite_file.rs index 93e063880..fe6effd16 100644 --- a/src/directory/composite_file.rs +++ b/src/directory/composite_file.rs @@ -4,7 +4,7 @@ use std::ops::Range; use common::{BinarySerializable, CountingWriter, HasLen, VInt}; -use crate::directory::{FileSlice, TerminatingWrite, WritePtr}; +use crate::directory::{FileSlice, FinishableWrite, WritePtr}; use crate::schema::{Field, Schema}; use crate::space_usage::{FieldUsage, PerFieldSpaceUsage}; @@ -40,7 +40,7 @@ pub struct CompositeWrite { offsets: Vec<(FileAddr, u64)>, } -impl CompositeWrite { +impl CompositeWrite { /// Crate a new API writer that writes a composite file /// in a given write. pub fn wrap(w: W) -> CompositeWrite { @@ -81,7 +81,7 @@ impl CompositeWrite { let footer_len = (self.write.written_bytes() - footer_offset) as u32; footer_len.serialize(&mut self.write)?; - self.write.terminate() + self.write.finish() } } @@ -183,7 +183,6 @@ impl CompositeFile { #[cfg(test)] mod test { - use std::io::Write; use std::path::Path; use common::{BinarySerializable, VInt}; @@ -201,10 +200,8 @@ mod test { let mut composite_write = CompositeWrite::wrap(w); let mut write_0 = composite_write.for_field(Field::from_field_id(0u32)); VInt(32431123u64).serialize(&mut write_0)?; - write_0.flush()?; let mut write_4 = composite_write.for_field(Field::from_field_id(4u32)); VInt(2).serialize(&mut write_4)?; - write_4.flush()?; composite_write.close()?; } { @@ -243,13 +240,10 @@ mod test { let mut composite_write = CompositeWrite::wrap(w); let mut write = composite_write.for_field_with_idx(Field::from_field_id(1u32), 0); VInt(32431123u64).serialize(&mut write)?; - write.flush()?; - let write = composite_write.for_field_with_idx(Field::from_field_id(1u32), 1); - write.flush()?; + composite_write.for_field_with_idx(Field::from_field_id(1u32), 1); let mut write = composite_write.for_field_with_idx(Field::from_field_id(0u32), 0); VInt(1_000_000).serialize(&mut write)?; - write.flush()?; composite_write.close()?; } diff --git a/src/directory/directory.rs b/src/directory/directory.rs index 0bc4b7f95..75e57aedf 100644 --- a/src/directory/directory.rs +++ b/src/directory/directory.rs @@ -1,4 +1,3 @@ -use std::io::Write; use std::path::{Path, PathBuf}; use std::sync::Arc; use std::time::Duration; @@ -6,7 +5,9 @@ use std::{fmt, io, thread}; use crate::directory::directory_lock::Lock; use crate::directory::error::{DeleteError, LockError, OpenReadError, OpenWriteError}; -use crate::directory::{FileHandle, FileSlice, WatchCallback, WatchHandle, WritePtr}; +use crate::directory::{ + FileHandle, FileSlice, FinishableWrite, WatchCallback, WatchHandle, WritePtr, +}; /// Retry the logic of acquiring locks is pretty simple. /// We just retry `n` times after a given `duratio`, both @@ -75,11 +76,11 @@ fn try_acquire_lock( filepath: &Path, directory: &dyn Directory, ) -> Result { - let mut write = directory.open_write(filepath).map_err(|e| match e { + let write = directory.open_write(filepath).map_err(|e| match e { OpenWriteError::FileAlreadyExists(_) => TryAcquireLockError::FileExists, OpenWriteError::IoError { io_error, .. } => TryAcquireLockError::IoError(io_error), })?; - write.flush().map_err(TryAcquireLockError::from)?; + write.finish().map_err(TryAcquireLockError::from)?; Ok(DirectoryLock::from(Box::new(DirectoryLockGuard { directory: directory.box_clone(), path: filepath.to_owned(), @@ -138,27 +139,15 @@ pub trait Directory: DirectoryClone + fmt::Debug + Send + Sync + 'static { /// Opens a writer for the *virtual file* associated with /// a [`Path`]. /// - /// Right after this call, for the span of the execution of the program - /// the file should be created and any subsequent call to - /// [`Directory::open_read()`] for the same path should return - /// a [`FileSlice`]. + /// Depending on the directory implementation, [`Directory::sync_directory()`] may be required + /// after finishing the writer to ensure that the file is durably created. /// - /// However, depending on the directory implementation, - /// it might be required to call [`Directory::sync_directory()`] to ensure - /// that the file is durably created. - /// (The semantics here are the same when dealing with - /// a POSIX filesystem.) + /// Write operations may be aggressively buffered. The client must call + /// [`FinishableWrite::finish()`] to finalize the file and make all writes available to + /// subsequent reads. The directory implementation owns its buffering strategy; clients should + /// not rely on `flush()` making an incomplete file available. /// - /// Write operations may be aggressively buffered. - /// The client of this trait is responsible for calling flush - /// to ensure that subsequent `read` operations - /// will take into account preceding `write` operations. - /// - /// Flush operation should also be persistent. - /// - /// The user shall not rely on [`Drop`] triggering `flush`. - /// Note that [`RamDirectory`][crate::directory::RamDirectory] will - /// panic! if `flush` was not called. + /// The user shall not rely on [`Drop`] finalizing the file. /// /// The file may not previously exist. fn open_write(&self, path: &Path) -> Result; diff --git a/src/directory/footer.rs b/src/directory/footer.rs index bffa2f2cf..a71caf72c 100644 --- a/src/directory/footer.rs +++ b/src/directory/footer.rs @@ -12,7 +12,7 @@ use crc32fast::Hasher; use serde::{Deserialize, Serialize}; use crate::directory::error::Incompatibility; -use crate::directory::{AntiCallToken, FileSlice, TerminatingWrite}; +use crate::directory::{AntiCallToken, FileSlice, FinishableWrite}; use crate::{Version, INDEX_FORMAT_OLDEST_SUPPORTED_VERSION, INDEX_FORMAT_VERSION}; const FOOTER_MAX_LEN: u32 = 50_000; @@ -125,14 +125,14 @@ impl Footer { } } -pub(crate) struct FooterProxy { - /// always Some except after terminate call +pub(crate) struct FooterProxy { + /// Always `Some` except after `finish()` is called. hasher: Option, - /// always Some except after terminate call + /// Always `Some` except after `finish()` is called. writer: Option, } -impl FooterProxy { +impl FooterProxy { pub fn new(writer: W) -> Self { FooterProxy { hasher: Some(Hasher::new()), @@ -141,7 +141,7 @@ impl FooterProxy { } } -impl Write for FooterProxy { +impl Write for FooterProxy { fn write(&mut self, buf: &[u8]) -> io::Result { let count = self.writer.as_mut().unwrap().write(buf)?; self.hasher.as_mut().unwrap().update(&buf[..count]); @@ -153,13 +153,13 @@ impl Write for FooterProxy { } } -impl TerminatingWrite for FooterProxy { - fn terminate_ref(&mut self, _: AntiCallToken) -> io::Result<()> { +impl FinishableWrite for FooterProxy { + fn finish_ref(&mut self, _: AntiCallToken) -> io::Result<()> { let crc32 = self.hasher.take().unwrap().finalize(); let footer = Footer::new(crc32); let mut writer = self.writer.take().unwrap(); footer.append_footer(&mut writer)?; - writer.terminate() + writer.finish() } } diff --git a/src/directory/managed_directory.rs b/src/directory/managed_directory.rs index a6b9718b6..33eb01e5d 100644 --- a/src/directory/managed_directory.rs +++ b/src/directory/managed_directory.rs @@ -344,7 +344,7 @@ mod tests_mmap_specific { use tempfile::TempDir; - use crate::directory::{Directory, ManagedDirectory, MmapDirectory, TerminatingWrite}; + use crate::directory::{Directory, FinishableWrite, ManagedDirectory, MmapDirectory}; #[test] fn test_managed_directory() { @@ -357,7 +357,7 @@ mod tests_mmap_specific { let mmap_directory = MmapDirectory::open(&tempdir_path).unwrap(); let mut managed_directory = ManagedDirectory::wrap(Box::new(mmap_directory)).unwrap(); let write_file = managed_directory.open_write(test_path1).unwrap(); - write_file.terminate().unwrap(); + write_file.finish().unwrap(); managed_directory .atomic_write(test_path2, &[0u8, 1u8]) .unwrap(); @@ -392,7 +392,7 @@ mod tests_mmap_specific { let mut managed_directory = ManagedDirectory::wrap(Box::new(mmap_directory)).unwrap(); let mut write = managed_directory.open_write(test_path1).unwrap(); write.write_all(&[0u8, 1u8]).unwrap(); - write.terminate().unwrap(); + write.finish().unwrap(); assert!(managed_directory.exists(test_path1).unwrap()); let _mmap_read = managed_directory.open_read(test_path1).unwrap(); diff --git a/src/directory/mmap_directory/mod.rs b/src/directory/mmap_directory/mod.rs index 0370a8b9e..d1e8b20d2 100644 --- a/src/directory/mmap_directory/mod.rs +++ b/src/directory/mmap_directory/mod.rs @@ -22,7 +22,7 @@ use crate::directory::error::{ DeleteError, LockError, OpenDirectoryError, OpenReadError, OpenWriteError, }; use crate::directory::{ - AntiCallToken, Directory, DirectoryLock, FileHandle, Lock, OwnedBytes, TerminatingWrite, + AntiCallToken, Directory, DirectoryLock, FileHandle, FinishableWrite, Lock, OwnedBytes, WatchCallback, WatchHandle, WritePtr, }; @@ -338,8 +338,8 @@ impl Write for SafeFileWriter { } } -impl TerminatingWrite for SafeFileWriter { - fn terminate_ref(&mut self, _: AntiCallToken) -> io::Result<()> { +impl FinishableWrite for SafeFileWriter { + fn finish_ref(&mut self, _: AntiCallToken) -> io::Result<()> { self.0.flush()?; self.0.sync_data()?; Ok(()) @@ -432,7 +432,7 @@ impl Directory for MmapDirectory { .create_new(true) .open(full_path); - let mut file = open_res.map_err(|io_err| { + let file = open_res.map_err(|io_err| { if io_err.kind() == io::ErrorKind::AlreadyExists { OpenWriteError::FileAlreadyExists(path.to_path_buf()) } else { @@ -440,16 +440,12 @@ impl Directory for MmapDirectory { } })?; - // making sure the file is created. - file.flush() - .map_err(|io_error| OpenWriteError::wrap_io_error(io_error, path.to_path_buf()))?; - // Note we actually do not sync the parent directory here. // // A newly created file, may, in some case, be created and even flushed to disk. // and then lost... // - // The file will only be durably written after we terminate AND + // The file will only be durably written after we finish AND // sync_directory() is called. let writer = SafeFileWriter::new(file); diff --git a/src/directory/mod.rs b/src/directory/mod.rs index ab1164cf8..1f7ba6e89 100644 --- a/src/directory/mod.rs +++ b/src/directory/mod.rs @@ -19,7 +19,7 @@ use std::io::BufWriter; use std::path::PathBuf; pub use common::file_slice::{FileHandle, FileSlice}; -pub use common::{AntiCallToken, OwnedBytes, TerminatingWrite}; +pub use common::{AntiCallToken, FinishableWrite, OwnedBytes}; pub use self::composite_file::{CompositeFile, CompositeWrite}; pub use self::directory::{Directory, DirectoryClone, DirectoryLock}; @@ -52,7 +52,7 @@ pub use self::mmap_directory::MmapDirectory; /// /// `WritePtr` are required to implement both Write /// and Seek. -pub type WritePtr = BufWriter>; +pub type WritePtr = BufWriter>; #[cfg(test)] mod tests; diff --git a/src/directory/ram_directory.rs b/src/directory/ram_directory.rs index a9c5ae437..ee19fcd0f 100644 --- a/src/directory/ram_directory.rs +++ b/src/directory/ram_directory.rs @@ -11,7 +11,7 @@ use super::FileHandle; use crate::core::META_FILEPATH; use crate::directory::error::{DeleteError, OpenReadError, OpenWriteError}; use crate::directory::{ - AntiCallToken, Directory, FileSlice, TerminatingWrite, WatchCallback, WatchCallbackList, + AntiCallToken, Directory, FileSlice, FinishableWrite, WatchCallback, WatchCallbackList, WatchHandle, WritePtr, }; @@ -41,26 +41,32 @@ impl MemoryUsageTracker { self.unreported_bytes = 0; } } + + fn finish(&mut self) { + if self.reported_bytes > 0 { + self.shared_usage + .fetch_sub(self.reported_bytes, Ordering::Relaxed); + self.reported_bytes = 0; + } + self.unreported_bytes = 0; + } } impl Drop for MemoryUsageTracker { fn drop(&mut self) { - if self.reported_bytes > 0 { - self.shared_usage - .fetch_sub(self.reported_bytes, Ordering::Relaxed); - } + self.finish(); } } /// Writer associated with the [`RamDirectory`]. /// -/// The Writer just writes a buffer. +/// The writer stores its buffer in the directory when finished. struct VecWriter { path: PathBuf, shared_directory: RamDirectory, data: Cursor>, memory_usage: MemoryUsageTracker, - is_flushed: bool, + is_finished: bool, } impl VecWriter { @@ -72,16 +78,16 @@ impl VecWriter { data: Cursor::new(Vec::new()), shared_directory, memory_usage, - is_flushed: true, + is_finished: true, } } } impl Drop for VecWriter { fn drop(&mut self) { - if !self.is_flushed { + if !self.is_finished { warn!( - "You forgot to flush {:?} before its writer got Drop. Do not rely on drop. This \ + "You forgot to finish {:?} before its writer got Drop. Do not rely on drop. This \ also occurs when the indexer crashed, so you may want to check the logs for the \ root cause.", self.path @@ -92,7 +98,7 @@ impl Drop for VecWriter { impl Write for VecWriter { fn write(&mut self, buf: &[u8]) -> io::Result { - self.is_flushed = false; + self.is_finished = false; let previous_capacity = self.data.get_ref().capacity(); self.data.write_all(buf)?; let capacity = self.data.get_ref().capacity(); @@ -100,17 +106,21 @@ impl Write for VecWriter { Ok(buf.len()) } + /// Nothing to flush since the data is stored in memory. The memory usage is updated on each + /// write. fn flush(&mut self) -> io::Result<()> { - self.is_flushed = true; - let mut fs = self.shared_directory.fs.write().unwrap(); - fs.write(self.path.clone(), self.data.get_ref()); Ok(()) } } -impl TerminatingWrite for VecWriter { - fn terminate_ref(&mut self, _: AntiCallToken) -> io::Result<()> { - self.flush() +impl FinishableWrite for VecWriter { + fn finish_ref(&mut self, _: AntiCallToken) -> io::Result<()> { + self.is_finished = true; + let data = std::mem::take(self.data.get_mut()); + let mut fs = self.shared_directory.fs.write().unwrap(); + fs.write_owned(self.path.clone(), data); + self.memory_usage.finish(); + Ok(()) } } @@ -122,8 +132,11 @@ struct InnerDirectory { impl InnerDirectory { fn write(&mut self, path: PathBuf, data: &[u8]) -> bool { - let data = FileSlice::from(data.to_vec()); - self.fs.insert(path, data).is_some() + self.write_owned(path, data.to_vec()) + } + + fn write_owned(&mut self, path: PathBuf, data: Vec) -> bool { + self.fs.insert(path, FileSlice::from(data)).is_some() } fn open_read(&self, path: &Path) -> Result { @@ -162,7 +175,7 @@ impl fmt::Debug for RamDirectory { /// A Directory storing everything in anonymous memory. /// /// It is mainly meant for unit testing. -/// Writes are only made visible upon flushing. +/// Writes are only made visible upon finishing the writer. #[derive(Clone, Default)] pub struct RamDirectory { fs: Arc>, @@ -194,7 +207,7 @@ impl RamDirectory { for (path, file) in wlock.fs.iter() { let mut dest_wrt = dest.open_write(path)?; dest_wrt.write_all(file.read_bytes()?.as_slice())?; - dest_wrt.terminate()?; + dest_wrt.finish()?; } Ok(()) } @@ -232,18 +245,14 @@ impl Directory for RamDirectory { } fn open_write(&self, path: &Path) -> Result { - let mut fs = self.fs.write().unwrap(); let path_buf = PathBuf::from(path); - let vec_writer = VecWriter::new(path_buf.clone(), self.clone()); - let exists = fs.write(path_buf.clone(), &[]); - // force the creation of the file to mimic the MMap directory. - if exists { - Err(OpenWriteError::FileAlreadyExists(path_buf)) - } else { - // The writer's allocation is tracked by `RamDirectory`; an additional buffer would not - // be included in `total_mem_usage()`. - Ok(BufWriter::with_capacity(0, Box::new(vec_writer))) + if self.fs.read().unwrap().exists(path) { + return Err(OpenWriteError::FileAlreadyExists(path_buf)); } + let vec_writer = VecWriter::new(path_buf, self.clone()); + // The writer's allocation is tracked by `RamDirectory`; an additional buffer would not + // be included in `total_mem_usage()`. + Ok(BufWriter::with_capacity(0, Box::new(vec_writer))) } fn atomic_read(&self, path: &Path) -> Result, OpenReadError> { @@ -281,7 +290,7 @@ mod tests { use std::path::Path; use super::{RamDirectory, MEMORY_USAGE_UPDATE_THRESHOLD}; - use crate::directory::TerminatingWrite; + use crate::directory::FinishableWrite; use crate::Directory; #[test] @@ -294,7 +303,7 @@ mod tests { assert!(directory.atomic_write(path_atomic, msg_atomic).is_ok()); let mut wrt = directory.open_write(path_seq).unwrap(); assert!(wrt.write_all(msg_seq).is_ok()); - assert!(wrt.flush().is_ok()); + assert!(wrt.finish().is_ok()); let directory_copy = RamDirectory::create(); assert!(directory.persist(&directory_copy).is_ok()); assert_eq!(directory_copy.atomic_read(path_atomic).unwrap(), msg_atomic); @@ -321,11 +330,11 @@ mod tests { assert!(grown_capacity > first_capacity); writer.flush().unwrap(); - let file_len = 1 + MEMORY_USAGE_UPDATE_THRESHOLD + first_capacity + 1; - assert_eq!(dir.total_mem_usage(), grown_capacity + file_len); + assert_eq!(dir.total_mem_usage(), grown_capacity); assert_eq!(dir.clone().total_mem_usage(), dir.total_mem_usage()); - writer.terminate().unwrap(); + let file_len = 1 + MEMORY_USAGE_UPDATE_THRESHOLD + first_capacity + 1; + writer.finish().unwrap(); assert_eq!(dir.total_mem_usage(), file_len); dir.delete(path).unwrap(); diff --git a/src/directory/tests.rs b/src/directory/tests.rs index a2c8473ce..cfc524860 100644 --- a/src/directory/tests.rs +++ b/src/directory/tests.rs @@ -120,11 +120,10 @@ mod ram_directory_tests { fn test_simple(directory: &dyn Directory) -> crate::Result<()> { let test_path: &'static Path = Path::new("some_path_for_test"); let mut write_file = directory.open_write(test_path)?; - assert!(directory.exists(test_path).unwrap()); write_file.write_all(&[4])?; write_file.write_all(&[3])?; write_file.write_all(&[7, 3, 5])?; - write_file.flush()?; + write_file.finish()?; let read_file = directory.open_read(test_path)?.read_bytes()?; assert_eq!(read_file.as_slice(), &[4u8, 3u8, 7u8, 3u8, 5u8]); mem::drop(read_file); @@ -135,7 +134,7 @@ fn test_simple(directory: &dyn Directory) -> crate::Result<()> { fn test_rewrite_forbidden(directory: &dyn Directory) -> crate::Result<()> { let test_path: &'static Path = Path::new("some_path_for_test"); - directory.open_write(test_path)?; + directory.open_write(test_path)?.finish()?; assert!(directory.exists(test_path).unwrap()); assert!(directory.open_write(test_path).is_err()); assert!(directory.delete(test_path).is_ok()); @@ -144,13 +143,11 @@ fn test_rewrite_forbidden(directory: &dyn Directory) -> crate::Result<()> { fn test_write_create_the_file(directory: &dyn Directory) { let test_path: &'static Path = Path::new("some_path_for_test"); - { - assert!(directory.open_read(test_path).is_err()); - let _w = directory.open_write(test_path).unwrap(); - assert!(directory.exists(test_path).unwrap()); - assert!(directory.open_read(test_path).is_ok()); - assert!(directory.delete(test_path).is_ok()); - } + assert!(directory.open_read(test_path).is_err()); + directory.open_write(test_path).unwrap().finish().unwrap(); + assert!(directory.exists(test_path).unwrap()); + assert!(directory.open_read(test_path).is_ok()); + assert!(directory.delete(test_path).is_ok()); } fn test_directory_delete(directory: &dyn Directory) -> crate::Result<()> { @@ -158,7 +155,7 @@ fn test_directory_delete(directory: &dyn Directory) -> crate::Result<()> { assert!(directory.open_read(test_path).is_err()); let mut write_file = directory.open_write(test_path)?; write_file.write_all(&[1, 2, 3, 4])?; - write_file.flush()?; + write_file.finish()?; { let read_handle = directory.open_read(test_path)?.read_bytes()?; assert_eq!(read_handle.as_slice(), &[1u8, 2u8, 3u8, 4u8]); diff --git a/src/fastfield/alive_bitset.rs b/src/fastfield/alive_bitset.rs index bbdc82a45..3ca05ed13 100644 --- a/src/fastfield/alive_bitset.rs +++ b/src/fastfield/alive_bitset.rs @@ -8,7 +8,7 @@ use crate::DocId; /// Write an alive `BitSet` /// /// where `alive_bitset` is the set of alive `DocId`. -/// Warning: this function does not call terminate. The caller is in charge of +/// Warning: this function does not call `finish()`. The caller is in charge of /// closing the writer properly. pub fn write_alive_bitset(alive_bitset: &BitSet, writer: &mut T) -> io::Result<()> { alive_bitset.serialize(writer)?; diff --git a/src/fastfield/mod.rs b/src/fastfield/mod.rs index d56dc27a8..724dbcfaf 100644 --- a/src/fastfield/mod.rs +++ b/src/fastfield/mod.rs @@ -82,7 +82,7 @@ mod tests { use std::path::Path; use columnar::StrColumn; - use common::{ByteCount, DateTimePrecision, HasLen, TerminatingWrite}; + use common::{ByteCount, DateTimePrecision, FinishableWrite, HasLen}; use once_cell::sync::Lazy; use rand::prelude::SliceRandom; use rand::rngs::StdRng; @@ -132,7 +132,7 @@ mod tests { .add_document(&doc!(*FIELD=>2u64)) .unwrap(); fast_field_writers.serialize(&mut write, None).unwrap(); - write.terminate().unwrap(); + write.finish().unwrap(); } let file = directory.open_read(path).unwrap(); @@ -183,7 +183,7 @@ mod tests { .add_document(&doc!(*FIELD=>215u64)) .unwrap(); fast_field_writers.serialize(&mut write, None).unwrap(); - write.terminate().unwrap(); + write.finish().unwrap(); } let file = directory.open_read(path).unwrap(); assert_eq!(file.len(), 108); @@ -216,7 +216,7 @@ mod tests { .unwrap(); } fast_field_writers.serialize(&mut write, None).unwrap(); - write.terminate().unwrap(); + write.finish().unwrap(); } let file = directory.open_read(path).unwrap(); assert_eq!(file.len(), 81); @@ -248,7 +248,7 @@ mod tests { .unwrap(); } fast_field_writers.serialize(&mut write, None).unwrap(); - write.terminate().unwrap(); + write.finish().unwrap(); } let file = directory.open_read(path).unwrap(); assert_eq!(file.len(), 4476); @@ -281,7 +281,7 @@ mod tests { fast_field_writers.add_document(&doc).unwrap(); } fast_field_writers.serialize(&mut write, None).unwrap(); - write.terminate().unwrap(); + write.finish().unwrap(); } let file = directory.open_read(path).unwrap(); assert_eq!(file.len(), 252); @@ -320,7 +320,7 @@ mod tests { let doc = TantivyDocument::default(); fast_field_writers.add_document(&doc).unwrap(); fast_field_writers.serialize(&mut write, None).unwrap(); - write.terminate().unwrap(); + write.finish().unwrap(); } let file = directory.open_read(path).unwrap(); @@ -353,7 +353,7 @@ mod tests { let doc = TantivyDocument::default(); fast_field_writers.add_document(&doc).unwrap(); fast_field_writers.serialize(&mut write, None).unwrap(); - write.terminate().unwrap(); + write.finish().unwrap(); } let file = directory.open_read(path).unwrap(); @@ -390,7 +390,7 @@ mod tests { fast_field_writers.add_document(&doc!(*FIELD=>x)).unwrap(); } fast_field_writers.serialize(&mut write, None).unwrap(); - write.terminate().unwrap(); + write.finish().unwrap(); } let file = directory.open_read(path).unwrap(); let fast_field_readers = FastFieldReaders::open(file, SCHEMA.clone()).unwrap(); @@ -775,7 +775,7 @@ mod tests { .add_document(&doc!(field=>false)) .unwrap(); fast_field_writers.serialize(&mut write, None).unwrap(); - write.terminate().unwrap(); + write.finish().unwrap(); } let file = directory.open_read(path).unwrap(); assert_eq!(file.len(), 84); @@ -807,7 +807,7 @@ mod tests { .unwrap(); } fast_field_writers.serialize(&mut write, None).unwrap(); - write.terminate().unwrap(); + write.finish().unwrap(); } let file = directory.open_read(path).unwrap(); assert_eq!(file.len(), 96); @@ -832,7 +832,7 @@ mod tests { let doc = TantivyDocument::default(); fast_field_writers.add_document(&doc).unwrap(); fast_field_writers.serialize(&mut write, None).unwrap(); - write.terminate().unwrap(); + write.finish().unwrap(); } let file = directory.open_read(path).unwrap(); assert_eq!(file.len(), 86); @@ -860,7 +860,7 @@ mod tests { fast_field_writers.add_document(doc).unwrap(); } fast_field_writers.serialize(&mut write, None).unwrap(); - write.terminate().unwrap(); + write.finish().unwrap(); } Ok(directory) } diff --git a/src/fastfield/plugin.rs b/src/fastfield/plugin.rs index eb33b6816..5131f51d5 100644 --- a/src/fastfield/plugin.rs +++ b/src/fastfield/plugin.rs @@ -9,7 +9,7 @@ use std::collections::BTreeMap; use columnar::{ ColumnType, ColumnarReader, MergeRowOrder, RowAddr, ShuffleMergeOrder, StackMergeOrder, }; -use common::TerminatingWrite; +use common::FinishableWrite; use measure_time::debug_time; use crate::directory::{Directory, WritePtr}; @@ -60,7 +60,7 @@ impl SegmentPlugin for FastFieldsPlugin { &mut fast_field_wrt, )?; - fast_field_wrt.terminate()?; + fast_field_wrt.finish()?; Ok(()) } @@ -131,7 +131,7 @@ impl PluginWriter for FastFieldsPluginWriter { self.writer .serialize(&mut self.fast_field_write, doc_id_map) .map_err(|e| crate::TantivyError::InternalError(e.to_string()))?; - self.fast_field_write.terminate()?; + self.fast_field_write.finish()?; Ok(()) } diff --git a/src/fieldnorm/serializer.rs b/src/fieldnorm/serializer.rs index 316b4cfad..5326ea494 100644 --- a/src/fieldnorm/serializer.rs +++ b/src/fieldnorm/serializer.rs @@ -22,7 +22,6 @@ impl FieldNormsSerializer { pub fn serialize_field(&mut self, field: Field, fieldnorms_data: &[u8]) -> io::Result<()> { let write = self.composite_write.for_field(field); write.write_all(fieldnorms_data)?; - write.flush()?; Ok(()) } diff --git a/src/indexer/index_writer.rs b/src/indexer/index_writer.rs index 22a251ccf..cb0862907 100644 --- a/src/indexer/index_writer.rs +++ b/src/indexer/index_writer.rs @@ -9,7 +9,7 @@ use smallvec::smallvec; use super::operation::{AddOperation, UserOperation}; use super::segment_updater::SegmentUpdater; use super::{AddBatch, AddBatchReceiver, AddBatchSender, PreparedCommit}; -use crate::directory::{DirectoryLock, GarbageCollectionResult, TerminatingWrite}; +use crate::directory::{DirectoryLock, FinishableWrite, GarbageCollectionResult}; use crate::error::TantivyError; use crate::fastfield::write_alive_bitset; use crate::index::{Index, Segment, SegmentComponent, SegmentId, SegmentMeta, SegmentReader}; @@ -172,7 +172,7 @@ pub fn advance_deletes( segment = segment.with_delete_meta(num_deleted_docs, target_opstamp); let mut alive_doc_file = segment.open_write(SegmentComponent::Delete)?; write_alive_bitset(&alive_bitset, &mut alive_doc_file)?; - alive_doc_file.terminate()?; + alive_doc_file.finish()?; } segment_entry.set_meta(segment.meta().clone()); diff --git a/src/plugin.rs b/src/plugin.rs index b66c56124..a0f040656 100644 --- a/src/plugin.rs +++ b/src/plugin.rs @@ -192,7 +192,7 @@ mod tests { let mut write = ctx.target_segment.open_write(component)?; use std::io::Write; write.write_all(&MARKER.to_le_bytes())?; - common::TerminatingWrite::terminate(write)?; + common::FinishableWrite::finish(write)?; Ok(()) } } @@ -258,7 +258,7 @@ mod tests { write.write_all(&(payload.len() as u32).to_le_bytes())?; write.write_all(payload)?; } - common::TerminatingWrite::terminate(write)?; + common::FinishableWrite::finish(write)?; Ok(()) } diff --git a/src/positions/serializer.rs b/src/positions/serializer.rs index f41923e8b..617033909 100644 --- a/src/positions/serializer.rs +++ b/src/positions/serializer.rs @@ -86,8 +86,8 @@ impl PositionSerializer { Ok(()) } - /// Close the positions for this term and flushes the data. - pub fn close(mut self) -> io::Result<()> { - self.positions_wrt.flush() + /// Close the positions for this field. + pub fn close(self) -> io::Result<()> { + Ok(()) } } diff --git a/src/postings/serializer.rs b/src/postings/serializer.rs index f44d14d93..4fbd42620 100644 --- a/src/postings/serializer.rs +++ b/src/postings/serializer.rs @@ -249,7 +249,6 @@ impl<'a, W: Write> FieldSerializer<'a, W> { if let Some(positions_serializer) = self.positions_serializer_opt { positions_serializer.close()?; } - self.postings_write.flush()?; self.term_dictionary_builder.finish()?; Ok(()) } diff --git a/src/store/store_compressor.rs b/src/store/store_compressor.rs index 20211b25a..683b6d30b 100644 --- a/src/store/store_compressor.rs +++ b/src/store/store_compressor.rs @@ -3,7 +3,7 @@ use std::sync::mpsc::{sync_channel, Receiver, SyncSender}; use std::thread::JoinHandle; use std::{io, thread}; -use common::{BinarySerializable, CountingWriter, TerminatingWrite}; +use common::{BinarySerializable, CountingWriter, FinishableWrite}; use super::DOC_STORE_VERSION; use crate::directory::WritePtr; @@ -151,7 +151,7 @@ impl BlockCompressorImpl { ); self.offset_index_writer.serialize_into(&mut self.writer)?; docstore_footer.serialize(&mut self.writer)?; - self.writer.terminate() + self.writer.finish() } } diff --git a/src/termdict/tests.rs b/src/termdict/tests.rs index 71b3f1c3e..050dd4bc4 100644 --- a/src/termdict/tests.rs +++ b/src/termdict/tests.rs @@ -2,7 +2,7 @@ use std::path::PathBuf; use std::{io, str}; use super::{TermDictionary, TermDictionaryBuilder, TermStreamer}; -use crate::directory::{Directory, FileSlice, RamDirectory, TerminatingWrite}; +use crate::directory::{Directory, FileSlice, FinishableWrite, RamDirectory}; use crate::postings::TermInfo; const BLOCK_SIZE: usize = 1_500; @@ -41,7 +41,7 @@ fn test_term_ordinals() -> crate::Result<()> { for term in COUNTRIES.iter() { term_dictionary_builder.insert(term.as_bytes(), &make_term_info(0u64))?; } - term_dictionary_builder.finish()?.terminate()?; + term_dictionary_builder.finish()?.finish()?; } let term_file = directory.open_read(&path)?; let term_dict: TermDictionary = TermDictionary::open(term_file)?; @@ -63,7 +63,7 @@ fn test_term_dictionary_simple() -> crate::Result<()> { let mut term_dictionary_builder = TermDictionaryBuilder::create(write)?; term_dictionary_builder.insert("abc".as_bytes(), &make_term_info(34u64))?; term_dictionary_builder.insert("abcd".as_bytes(), &make_term_info(346u64))?; - term_dictionary_builder.finish()?.terminate()?; + term_dictionary_builder.finish()?.finish()?; } let file = directory.open_read(&path)?; let term_dict: TermDictionary = TermDictionary::open(file)?; @@ -412,7 +412,7 @@ fn test_automaton_search() -> crate::Result<()> { for term in COUNTRIES.iter() { term_dictionary_builder.insert(term.as_bytes(), &make_term_info(0u64))?; } - term_dictionary_builder.finish()?.terminate()?; + term_dictionary_builder.finish()?.finish()?; } let file = directory.open_read(&path)?; let term_dict: TermDictionary = TermDictionary::open(file)?; diff --git a/sstable/src/index/mod.rs b/sstable/src/index/mod.rs index f927379be..b70bdf369 100644 --- a/sstable/src/index/mod.rs +++ b/sstable/src/index/mod.rs @@ -265,7 +265,7 @@ impl SSTableIndexBuilder { } let counting_writer = map_builder.into_inner().map_err(fst_error_to_io_error)?; let written_bytes = counting_writer.written_bytes(); - let mut wrt = counting_writer.finish(); + let mut wrt = counting_writer.into_inner(); let mut block_store_writer = v3::BlockAddrStoreWriter::new(); for block in &self.blocks { diff --git a/sstable/src/lib.rs b/sstable/src/lib.rs index 1f6bd14c7..082fc2f25 100644 --- a/sstable/src/lib.rs +++ b/sstable/src/lib.rs @@ -357,7 +357,7 @@ where SSTABLE_VERSION.serialize(&mut wrt)?; - let wrt = wrt.finish(); + let wrt = wrt.into_inner(); Ok(wrt.into_inner()?) } } diff --git a/tests/failpoints/mod.rs b/tests/failpoints/mod.rs index 213c86628..be6ad5b7f 100644 --- a/tests/failpoints/mod.rs +++ b/tests/failpoints/mod.rs @@ -1,6 +1,6 @@ use std::path::Path; -use tantivy::directory::{Directory, ManagedDirectory, RamDirectory, TerminatingWrite}; +use tantivy::directory::{Directory, FinishableWrite, ManagedDirectory, RamDirectory}; use tantivy::schema::{Schema, TEXT}; use tantivy::{doc, Index, IndexWriter, Term}; @@ -15,7 +15,7 @@ fn test_failpoints_managed_directory_gc_if_delete_fails() { managed_directory .open_write(test_path) .unwrap() - .terminate() + .finish() .unwrap(); assert!(managed_directory.exists(test_path).unwrap()); // triggering gc and setting the delete operation to fail. From f2a106df1c135f42afaecdae0365e60f41ab81a9 Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Mon, 14 Sep 2026 10:39:10 +0200 Subject: [PATCH 07/49] Prevent concurrent RamDirectory writers Reserve paths when writers are opened and release reservations when writers finish or drop. --- src/directory/ram_directory.rs | 27 +++++++++++++++++++++++---- src/directory/tests.rs | 4 +++- 2 files changed, 26 insertions(+), 5 deletions(-) diff --git a/src/directory/ram_directory.rs b/src/directory/ram_directory.rs index ee19fcd0f..9c7a74e4c 100644 --- a/src/directory/ram_directory.rs +++ b/src/directory/ram_directory.rs @@ -1,4 +1,4 @@ -use std::collections::HashMap; +use std::collections::{HashMap, HashSet}; use std::io::{self, BufWriter, Cursor, Write}; use std::path::{Path, PathBuf}; use std::sync::atomic::{AtomicUsize, Ordering}; @@ -91,7 +91,10 @@ impl Drop for VecWriter { also occurs when the indexer crashed, so you may want to check the logs for the \ root cause.", self.path - ) + ); + } + if let Ok(mut fs) = self.shared_directory.fs.write() { + fs.active_writers.remove(&self.path); } } } @@ -115,10 +118,11 @@ impl Write for VecWriter { impl FinishableWrite for VecWriter { fn finish_ref(&mut self, _: AntiCallToken) -> io::Result<()> { - self.is_finished = true; let data = std::mem::take(self.data.get_mut()); let mut fs = self.shared_directory.fs.write().unwrap(); + fs.active_writers.remove(&self.path); fs.write_owned(self.path.clone(), data); + self.is_finished = true; self.memory_usage.finish(); Ok(()) } @@ -127,6 +131,7 @@ impl FinishableWrite for VecWriter { #[derive(Default)] struct InnerDirectory { fs: HashMap, + active_writers: HashSet, watch_router: WatchCallbackList, } @@ -246,9 +251,11 @@ impl Directory for RamDirectory { fn open_write(&self, path: &Path) -> Result { let path_buf = PathBuf::from(path); - if self.fs.read().unwrap().exists(path) { + let mut fs = self.fs.write().unwrap(); + if fs.exists(path) || !fs.active_writers.insert(path_buf.clone()) { return Err(OpenWriteError::FileAlreadyExists(path_buf)); } + drop(fs); let vec_writer = VecWriter::new(path_buf, self.clone()); // The writer's allocation is tracked by `RamDirectory`; an additional buffer would not // be included in `total_mem_usage()`. @@ -310,6 +317,18 @@ mod tests { assert_eq!(directory_copy.atomic_read(path_seq).unwrap(), msg_seq); } + #[test] + fn test_dropped_writer_releases_path() { + let dir = RamDirectory::default(); + let path = Path::new("file"); + let writer = dir.open_write(path).unwrap(); + assert!(dir.open_write(path).is_err()); + + drop(writer); + dir.open_write(path).unwrap().finish().unwrap(); + assert!(dir.exists(path).unwrap()); + } + #[test] fn test_active_writer_memory_usage() { let dir = RamDirectory::default(); diff --git a/src/directory/tests.rs b/src/directory/tests.rs index cfc524860..1d3b01c8a 100644 --- a/src/directory/tests.rs +++ b/src/directory/tests.rs @@ -134,7 +134,9 @@ fn test_simple(directory: &dyn Directory) -> crate::Result<()> { fn test_rewrite_forbidden(directory: &dyn Directory) -> crate::Result<()> { let test_path: &'static Path = Path::new("some_path_for_test"); - directory.open_write(test_path)?.finish()?; + let writer = directory.open_write(test_path)?; + assert!(directory.open_write(test_path).is_err()); + writer.finish()?; assert!(directory.exists(test_path).unwrap()); assert!(directory.open_write(test_path).is_err()); assert!(directory.delete(test_path).is_ok()); From e46d44c38833e7d56ebae02c7695a78305decf34 Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Mon, 14 Sep 2026 10:39:11 +0200 Subject: [PATCH 08/49] Finalize mmap writers in directory tests Use finish instead of flush before reading completed files and clarify when file data is synchronized. --- src/directory/mmap_directory/mod.rs | 18 ++++++------------ 1 file changed, 6 insertions(+), 12 deletions(-) diff --git a/src/directory/mmap_directory/mod.rs b/src/directory/mmap_directory/mod.rs index d1e8b20d2..f2751bb9b 100644 --- a/src/directory/mmap_directory/mod.rs +++ b/src/directory/mmap_directory/mod.rs @@ -318,8 +318,7 @@ impl Drop for ReleaseLockFile { } } -/// This Write wraps a File, but has the specificity of -/// call `sync_all` on flush. +/// Wraps a file and syncs its data when the writer is finished. struct SafeFileWriter(File); impl SafeFileWriter { @@ -555,10 +554,7 @@ mod tests { // In that case the directory returns a SharedVecSlice. let mmap_directory = MmapDirectory::create_from_tempdir().unwrap(); let path = PathBuf::from("test"); - { - let mut w = mmap_directory.open_write(&path).unwrap(); - w.flush().unwrap(); - } + mmap_directory.open_write(&path).unwrap().finish().unwrap(); let readonlymap = mmap_directory.open_read(&path).unwrap(); assert_eq!(readonlymap.len(), 0); } @@ -574,12 +570,10 @@ mod tests { let paths: Vec = (0..num_paths) .map(|i| PathBuf::from(&*format!("file_{i}"))) .collect(); - { - for path in &paths { - let mut w = mmap_directory.open_write(path).unwrap(); - w.write_all(content).unwrap(); - w.flush().unwrap(); - } + for path in &paths { + let mut w = mmap_directory.open_write(path).unwrap(); + w.write_all(content).unwrap(); + w.finish().unwrap(); } let mut keep = vec![]; From 2ab4fb80f02153b0518acf2bc6de621c1b4346cc Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Mon, 14 Sep 2026 11:56:41 +0200 Subject: [PATCH 09/49] Buffer RamDirectory writes Use the standard BufWriter capacity to avoid forwarding every small write directly to the underlying VecWriter. --- src/directory/ram_directory.rs | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/src/directory/ram_directory.rs b/src/directory/ram_directory.rs index 9c7a74e4c..388c1d307 100644 --- a/src/directory/ram_directory.rs +++ b/src/directory/ram_directory.rs @@ -257,9 +257,7 @@ impl Directory for RamDirectory { } drop(fs); let vec_writer = VecWriter::new(path_buf, self.clone()); - // The writer's allocation is tracked by `RamDirectory`; an additional buffer would not - // be included in `total_mem_usage()`. - Ok(BufWriter::with_capacity(0, Box::new(vec_writer))) + Ok(BufWriter::new(Box::new(vec_writer))) } fn atomic_read(&self, path: &Path) -> Result, OpenReadError> { From 1c8a4801b3b58b2d251d29e0c26480760d8529db Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Mon, 14 Sep 2026 12:17:33 +0200 Subject: [PATCH 10/49] Shrink RamDirectory files before storing --- src/directory/ram_directory.rs | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/src/directory/ram_directory.rs b/src/directory/ram_directory.rs index 388c1d307..59cbff625 100644 --- a/src/directory/ram_directory.rs +++ b/src/directory/ram_directory.rs @@ -118,7 +118,8 @@ impl Write for VecWriter { impl FinishableWrite for VecWriter { fn finish_ref(&mut self, _: AntiCallToken) -> io::Result<()> { - let data = std::mem::take(self.data.get_mut()); + let mut data = std::mem::take(self.data.get_mut()); + data.shrink_to_fit(); let mut fs = self.shared_directory.fs.write().unwrap(); fs.active_writers.remove(&self.path); fs.write_owned(self.path.clone(), data); From bbc1b9e1d1f5d906d54a70e07d76b1e407b4fe15 Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Mon, 14 Sep 2026 13:00:16 +0200 Subject: [PATCH 11/49] Restore TerminatingWrite API --- benches/merge_segments.rs | 10 ++--- .../src/column_index/multivalued_index.rs | 2 +- .../src/columnar/merge/merge_dict_column.rs | 2 +- columnar/src/columnar/writer/mod.rs | 2 +- common/src/lib.rs | 2 +- common/src/writer.rs | 38 +++++++++---------- src/directory/composite_file.rs | 6 +-- src/directory/directory.rs | 8 ++-- src/directory/footer.rs | 18 ++++----- src/directory/managed_directory.rs | 6 +-- src/directory/mmap_directory/mod.rs | 14 +++---- src/directory/mod.rs | 4 +- src/directory/ram_directory.rs | 22 +++++------ src/directory/tests.rs | 8 ++-- src/fastfield/alive_bitset.rs | 2 +- src/fastfield/mod.rs | 26 ++++++------- src/fastfield/plugin.rs | 6 +-- src/indexer/index_writer.rs | 4 +- src/plugin.rs | 4 +- src/store/store_compressor.rs | 4 +- src/termdict/tests.rs | 8 ++-- sstable/src/index/mod.rs | 2 +- sstable/src/lib.rs | 2 +- tests/failpoints/mod.rs | 4 +- 24 files changed, 102 insertions(+), 102 deletions(-) diff --git a/benches/merge_segments.rs b/benches/merge_segments.rs index d57fddf74..c2fcbb6f1 100644 --- a/benches/merge_segments.rs +++ b/benches/merge_segments.rs @@ -16,7 +16,7 @@ use rand::rngs::StdRng; use rand::SeedableRng; use tantivy::directory::error::{DeleteError, OpenReadError, OpenWriteError}; use tantivy::directory::{ - AntiCallToken, Directory, FileHandle, FinishableWrite, OwnedBytes, WatchCallback, WatchHandle, + AntiCallToken, Directory, FileHandle, TerminatingWrite, OwnedBytes, WatchCallback, WatchHandle, WritePtr, }; use tantivy::indexer::{merge_filtered_segments, NoMergePolicy}; @@ -40,8 +40,8 @@ impl Write for NullWriter { } } -impl FinishableWrite for NullWriter { - fn finish_ref(&mut self, _token: AntiCallToken) -> io::Result<()> { +impl TerminatingWrite for NullWriter { + fn terminate_ref(&mut self, _token: AntiCallToken) -> io::Result<()> { Ok(()) } } @@ -63,8 +63,8 @@ impl Write for InMemoryWriter { } } -impl FinishableWrite for InMemoryWriter { - fn finish_ref(&mut self, _token: AntiCallToken) -> io::Result<()> { +impl TerminatingWrite for InMemoryWriter { + fn terminate_ref(&mut self, _token: AntiCallToken) -> io::Result<()> { let bytes = OwnedBytes::new(std::mem::take(&mut self.buffer)); self.blobs.write().unwrap().insert(self.path.clone(), bytes); Ok(()) diff --git a/columnar/src/column_index/multivalued_index.rs b/columnar/src/column_index/multivalued_index.rs index ff6cea5aa..ad7efd363 100644 --- a/columnar/src/column_index/multivalued_index.rs +++ b/columnar/src/column_index/multivalued_index.rs @@ -33,7 +33,7 @@ pub fn serialize_multivalued_index( } = doc_ids_with_values; serialize_optional_index(&**non_null_row_ids, *num_rows, &mut count_writer)?; let optional_len = count_writer.written_bytes() as u32; - let output = count_writer.into_inner(); + let output = count_writer.finish(); serialize_u64_based_column_values( &**start_offsets, &[CodecType::Bitpacked, CodecType::Linear], diff --git a/columnar/src/columnar/merge/merge_dict_column.rs b/columnar/src/columnar/merge/merge_dict_column.rs index 768b19d36..ae2724b02 100644 --- a/columnar/src/columnar/merge/merge_dict_column.rs +++ b/columnar/src/columnar/merge/merge_dict_column.rs @@ -22,7 +22,7 @@ pub fn merge_bytes_or_str_column( // TODO !!! Remove useless terms. let term_ord_mapping = serialize_merged_dict(bytes_columns, merge_row_order, &mut output)?; let dictionary_num_bytes: u32 = output.written_bytes() as u32; - let output = output.into_inner(); + let output = output.finish(); let remapped_term_ordinals_values = RemappedTermOrdinalsValues { bytes_columns, term_ord_mapping: &term_ord_mapping, diff --git a/columnar/src/columnar/writer/mod.rs b/columnar/src/columnar/writer/mod.rs index 8f803f1b0..999ccd058 100644 --- a/columnar/src/columnar/writer/mod.rs +++ b/columnar/src/columnar/writer/mod.rs @@ -554,7 +554,7 @@ fn serialize_bytes_or_str_column( let term_id_mapping: TermIdMapping = dictionary_builder.serialize(arena, &mut counting_writer)?; let dictionary_num_bytes: u32 = counting_writer.written_bytes() as u32; - let mut wrt = counting_writer.into_inner(); + let mut wrt = counting_writer.finish(); let operation_iterator = operation_it.map(|symbol: ColumnOperation| { // We map unordered ids to ordered ids. match symbol { diff --git a/common/src/lib.rs b/common/src/lib.rs index f97c92d7f..4e64af11c 100644 --- a/common/src/lib.rs +++ b/common/src/lib.rs @@ -24,7 +24,7 @@ pub use serialize::{BinarySerializable, DeserializeFrom, FixedSize}; pub use vint::{ VInt, VIntU128, read_u32_vint, read_u32_vint_no_advance, serialize_vint_u32, write_u32_vint, }; -pub use writer::{AntiCallToken, CountingWriter, FinishableWrite}; +pub use writer::{AntiCallToken, CountingWriter, TerminatingWrite}; /// Has length trait pub trait HasLen { diff --git a/common/src/writer.rs b/common/src/writer.rs index 625bddc6f..f871a9b9e 100644 --- a/common/src/writer.rs +++ b/common/src/writer.rs @@ -21,7 +21,7 @@ impl CountingWriter { /// Returns the underlying write object. /// Note that this method does not trigger any flushing. #[inline] - pub fn into_inner(self) -> W { + pub fn finish(self) -> W { self.underlying } } @@ -47,15 +47,15 @@ impl Write for CountingWriter { } } -impl FinishableWrite for CountingWriter { +impl TerminatingWrite for CountingWriter { #[inline] - fn finish_ref(&mut self, token: AntiCallToken) -> io::Result<()> { - self.underlying.finish_ref(token) + fn terminate_ref(&mut self, token: AntiCallToken) -> io::Result<()> { + self.underlying.terminate_ref(token) } } /// Struct used to prevent from calling -/// [`finish_ref`](FinishableWrite::finish_ref) directly +/// [`terminate_ref`](TerminatingWrite::terminate_ref) directly /// /// The point is that while the type is public, it cannot be built by anyone /// outside of this module. @@ -64,35 +64,35 @@ pub struct AntiCallToken(()); /// Trait used to indicate when no more write need to be done on a writer /// /// Thread-safety is enforced at the call sites that require it. -pub trait FinishableWrite: Write { - /// Finishes the writer, consuming it. Internally calls [`FinishableWrite::finish_ref`]. - fn finish(mut self) -> io::Result<()> +pub trait TerminatingWrite: Write { + /// Indicates that the writer will no longer be used. Internally calls `terminate_ref`. + fn terminate(mut self) -> io::Result<()> where Self: Sized { - self.finish_ref(AntiCallToken(())) + self.terminate_ref(AntiCallToken(())) } /// You should implement this function to define custom behavior. /// This function should flush any buffer it may hold. - fn finish_ref(&mut self, _: AntiCallToken) -> io::Result<()>; + fn terminate_ref(&mut self, _: AntiCallToken) -> io::Result<()>; } -impl FinishableWrite for Box { - fn finish_ref(&mut self, token: AntiCallToken) -> io::Result<()> { - self.as_mut().finish_ref(token) +impl TerminatingWrite for Box { + fn terminate_ref(&mut self, token: AntiCallToken) -> io::Result<()> { + self.as_mut().terminate_ref(token) } } -impl FinishableWrite for BufWriter { - fn finish_ref(&mut self, a: AntiCallToken) -> io::Result<()> { +impl TerminatingWrite for BufWriter { + fn terminate_ref(&mut self, a: AntiCallToken) -> io::Result<()> { if !self.buffer().is_empty() { self.flush()?; } - self.get_mut().finish_ref(a) + self.get_mut().terminate_ref(a) } } -impl FinishableWrite for &mut Vec { - fn finish_ref(&mut self, _a: AntiCallToken) -> io::Result<()> { +impl TerminatingWrite for &mut Vec { + fn terminate_ref(&mut self, _a: AntiCallToken) -> io::Result<()> { self.flush() } } @@ -111,7 +111,7 @@ mod test { let bytes = (0u8..10u8).collect::>(); counting_writer.write_all(&bytes).unwrap(); let len = counting_writer.written_bytes(); - let buffer_restituted: Vec = counting_writer.into_inner(); + let buffer_restituted: Vec = counting_writer.finish(); assert_eq!(len, 10u64); assert_eq!(buffer_restituted.len(), 10); } diff --git a/src/directory/composite_file.rs b/src/directory/composite_file.rs index fe6effd16..f3d3d4cd6 100644 --- a/src/directory/composite_file.rs +++ b/src/directory/composite_file.rs @@ -4,7 +4,7 @@ use std::ops::Range; use common::{BinarySerializable, CountingWriter, HasLen, VInt}; -use crate::directory::{FileSlice, FinishableWrite, WritePtr}; +use crate::directory::{FileSlice, TerminatingWrite, WritePtr}; use crate::schema::{Field, Schema}; use crate::space_usage::{FieldUsage, PerFieldSpaceUsage}; @@ -40,7 +40,7 @@ pub struct CompositeWrite { offsets: Vec<(FileAddr, u64)>, } -impl CompositeWrite { +impl CompositeWrite { /// Crate a new API writer that writes a composite file /// in a given write. pub fn wrap(w: W) -> CompositeWrite { @@ -81,7 +81,7 @@ impl CompositeWrite { let footer_len = (self.write.written_bytes() - footer_offset) as u32; footer_len.serialize(&mut self.write)?; - self.write.finish() + self.write.terminate() } } diff --git a/src/directory/directory.rs b/src/directory/directory.rs index 75e57aedf..46b3c669f 100644 --- a/src/directory/directory.rs +++ b/src/directory/directory.rs @@ -6,7 +6,7 @@ use std::{fmt, io, thread}; use crate::directory::directory_lock::Lock; use crate::directory::error::{DeleteError, LockError, OpenReadError, OpenWriteError}; use crate::directory::{ - FileHandle, FileSlice, FinishableWrite, WatchCallback, WatchHandle, WritePtr, + FileHandle, FileSlice, TerminatingWrite, WatchCallback, WatchHandle, WritePtr, }; /// Retry the logic of acquiring locks is pretty simple. @@ -80,7 +80,7 @@ fn try_acquire_lock( OpenWriteError::FileAlreadyExists(_) => TryAcquireLockError::FileExists, OpenWriteError::IoError { io_error, .. } => TryAcquireLockError::IoError(io_error), })?; - write.finish().map_err(TryAcquireLockError::from)?; + write.terminate().map_err(TryAcquireLockError::from)?; Ok(DirectoryLock::from(Box::new(DirectoryLockGuard { directory: directory.box_clone(), path: filepath.to_owned(), @@ -140,10 +140,10 @@ pub trait Directory: DirectoryClone + fmt::Debug + Send + Sync + 'static { /// a [`Path`]. /// /// Depending on the directory implementation, [`Directory::sync_directory()`] may be required - /// after finishing the writer to ensure that the file is durably created. + /// after terminating the writer to ensure that the file is durably created. /// /// Write operations may be aggressively buffered. The client must call - /// [`FinishableWrite::finish()`] to finalize the file and make all writes available to + /// [`TerminatingWrite::terminate()`] to finalize the file and make all writes available to /// subsequent reads. The directory implementation owns its buffering strategy; clients should /// not rely on `flush()` making an incomplete file available. /// diff --git a/src/directory/footer.rs b/src/directory/footer.rs index a71caf72c..d2ff52f9a 100644 --- a/src/directory/footer.rs +++ b/src/directory/footer.rs @@ -12,7 +12,7 @@ use crc32fast::Hasher; use serde::{Deserialize, Serialize}; use crate::directory::error::Incompatibility; -use crate::directory::{AntiCallToken, FileSlice, FinishableWrite}; +use crate::directory::{AntiCallToken, FileSlice, TerminatingWrite}; use crate::{Version, INDEX_FORMAT_OLDEST_SUPPORTED_VERSION, INDEX_FORMAT_VERSION}; const FOOTER_MAX_LEN: u32 = 50_000; @@ -125,14 +125,14 @@ impl Footer { } } -pub(crate) struct FooterProxy { - /// Always `Some` except after `finish()` is called. +pub(crate) struct FooterProxy { + /// Always `Some` except after `terminate()` is called. hasher: Option, - /// Always `Some` except after `finish()` is called. + /// Always `Some` except after `terminate()` is called. writer: Option, } -impl FooterProxy { +impl FooterProxy { pub fn new(writer: W) -> Self { FooterProxy { hasher: Some(Hasher::new()), @@ -141,7 +141,7 @@ impl FooterProxy { } } -impl Write for FooterProxy { +impl Write for FooterProxy { fn write(&mut self, buf: &[u8]) -> io::Result { let count = self.writer.as_mut().unwrap().write(buf)?; self.hasher.as_mut().unwrap().update(&buf[..count]); @@ -153,13 +153,13 @@ impl Write for FooterProxy { } } -impl FinishableWrite for FooterProxy { - fn finish_ref(&mut self, _: AntiCallToken) -> io::Result<()> { +impl TerminatingWrite for FooterProxy { + fn terminate_ref(&mut self, _: AntiCallToken) -> io::Result<()> { let crc32 = self.hasher.take().unwrap().finalize(); let footer = Footer::new(crc32); let mut writer = self.writer.take().unwrap(); footer.append_footer(&mut writer)?; - writer.finish() + writer.terminate() } } diff --git a/src/directory/managed_directory.rs b/src/directory/managed_directory.rs index 33eb01e5d..a6b9718b6 100644 --- a/src/directory/managed_directory.rs +++ b/src/directory/managed_directory.rs @@ -344,7 +344,7 @@ mod tests_mmap_specific { use tempfile::TempDir; - use crate::directory::{Directory, FinishableWrite, ManagedDirectory, MmapDirectory}; + use crate::directory::{Directory, ManagedDirectory, MmapDirectory, TerminatingWrite}; #[test] fn test_managed_directory() { @@ -357,7 +357,7 @@ mod tests_mmap_specific { let mmap_directory = MmapDirectory::open(&tempdir_path).unwrap(); let mut managed_directory = ManagedDirectory::wrap(Box::new(mmap_directory)).unwrap(); let write_file = managed_directory.open_write(test_path1).unwrap(); - write_file.finish().unwrap(); + write_file.terminate().unwrap(); managed_directory .atomic_write(test_path2, &[0u8, 1u8]) .unwrap(); @@ -392,7 +392,7 @@ mod tests_mmap_specific { let mut managed_directory = ManagedDirectory::wrap(Box::new(mmap_directory)).unwrap(); let mut write = managed_directory.open_write(test_path1).unwrap(); write.write_all(&[0u8, 1u8]).unwrap(); - write.finish().unwrap(); + write.terminate().unwrap(); assert!(managed_directory.exists(test_path1).unwrap()); let _mmap_read = managed_directory.open_read(test_path1).unwrap(); diff --git a/src/directory/mmap_directory/mod.rs b/src/directory/mmap_directory/mod.rs index f2751bb9b..1707260f2 100644 --- a/src/directory/mmap_directory/mod.rs +++ b/src/directory/mmap_directory/mod.rs @@ -22,7 +22,7 @@ use crate::directory::error::{ DeleteError, LockError, OpenDirectoryError, OpenReadError, OpenWriteError, }; use crate::directory::{ - AntiCallToken, Directory, DirectoryLock, FileHandle, FinishableWrite, Lock, OwnedBytes, + AntiCallToken, Directory, DirectoryLock, FileHandle, Lock, OwnedBytes, TerminatingWrite, WatchCallback, WatchHandle, WritePtr, }; @@ -318,7 +318,7 @@ impl Drop for ReleaseLockFile { } } -/// Wraps a file and syncs its data when the writer is finished. +/// Wraps a file and syncs its data when the writer is terminated. struct SafeFileWriter(File); impl SafeFileWriter { @@ -337,8 +337,8 @@ impl Write for SafeFileWriter { } } -impl FinishableWrite for SafeFileWriter { - fn finish_ref(&mut self, _: AntiCallToken) -> io::Result<()> { +impl TerminatingWrite for SafeFileWriter { + fn terminate_ref(&mut self, _: AntiCallToken) -> io::Result<()> { self.0.flush()?; self.0.sync_data()?; Ok(()) @@ -444,7 +444,7 @@ impl Directory for MmapDirectory { // A newly created file, may, in some case, be created and even flushed to disk. // and then lost... // - // The file will only be durably written after we finish AND + // The file will only be durably written after we terminate AND // sync_directory() is called. let writer = SafeFileWriter::new(file); @@ -554,7 +554,7 @@ mod tests { // In that case the directory returns a SharedVecSlice. let mmap_directory = MmapDirectory::create_from_tempdir().unwrap(); let path = PathBuf::from("test"); - mmap_directory.open_write(&path).unwrap().finish().unwrap(); + mmap_directory.open_write(&path).unwrap().terminate().unwrap(); let readonlymap = mmap_directory.open_read(&path).unwrap(); assert_eq!(readonlymap.len(), 0); } @@ -573,7 +573,7 @@ mod tests { for path in &paths { let mut w = mmap_directory.open_write(path).unwrap(); w.write_all(content).unwrap(); - w.finish().unwrap(); + w.terminate().unwrap(); } let mut keep = vec![]; diff --git a/src/directory/mod.rs b/src/directory/mod.rs index 1f7ba6e89..ab1164cf8 100644 --- a/src/directory/mod.rs +++ b/src/directory/mod.rs @@ -19,7 +19,7 @@ use std::io::BufWriter; use std::path::PathBuf; pub use common::file_slice::{FileHandle, FileSlice}; -pub use common::{AntiCallToken, FinishableWrite, OwnedBytes}; +pub use common::{AntiCallToken, OwnedBytes, TerminatingWrite}; pub use self::composite_file::{CompositeFile, CompositeWrite}; pub use self::directory::{Directory, DirectoryClone, DirectoryLock}; @@ -52,7 +52,7 @@ pub use self::mmap_directory::MmapDirectory; /// /// `WritePtr` are required to implement both Write /// and Seek. -pub type WritePtr = BufWriter>; +pub type WritePtr = BufWriter>; #[cfg(test)] mod tests; diff --git a/src/directory/ram_directory.rs b/src/directory/ram_directory.rs index 59cbff625..2e6609353 100644 --- a/src/directory/ram_directory.rs +++ b/src/directory/ram_directory.rs @@ -11,7 +11,7 @@ use super::FileHandle; use crate::core::META_FILEPATH; use crate::directory::error::{DeleteError, OpenReadError, OpenWriteError}; use crate::directory::{ - AntiCallToken, Directory, FileSlice, FinishableWrite, WatchCallback, WatchCallbackList, + AntiCallToken, Directory, FileSlice, TerminatingWrite, WatchCallback, WatchCallbackList, WatchHandle, WritePtr, }; @@ -60,7 +60,7 @@ impl Drop for MemoryUsageTracker { /// Writer associated with the [`RamDirectory`]. /// -/// The writer stores its buffer in the directory when finished. +/// The writer stores its buffer in the directory when terminated. struct VecWriter { path: PathBuf, shared_directory: RamDirectory, @@ -87,7 +87,7 @@ impl Drop for VecWriter { fn drop(&mut self) { if !self.is_finished { warn!( - "You forgot to finish {:?} before its writer got Drop. Do not rely on drop. This \ + "You forgot to terminate {:?} before its writer got Drop. Do not rely on drop. This \ also occurs when the indexer crashed, so you may want to check the logs for the \ root cause.", self.path @@ -116,8 +116,8 @@ impl Write for VecWriter { } } -impl FinishableWrite for VecWriter { - fn finish_ref(&mut self, _: AntiCallToken) -> io::Result<()> { +impl TerminatingWrite for VecWriter { + fn terminate_ref(&mut self, _: AntiCallToken) -> io::Result<()> { let mut data = std::mem::take(self.data.get_mut()); data.shrink_to_fit(); let mut fs = self.shared_directory.fs.write().unwrap(); @@ -181,7 +181,7 @@ impl fmt::Debug for RamDirectory { /// A Directory storing everything in anonymous memory. /// /// It is mainly meant for unit testing. -/// Writes are only made visible upon finishing the writer. +/// Writes are only made visible upon terminating the writer. #[derive(Clone, Default)] pub struct RamDirectory { fs: Arc>, @@ -213,7 +213,7 @@ impl RamDirectory { for (path, file) in wlock.fs.iter() { let mut dest_wrt = dest.open_write(path)?; dest_wrt.write_all(file.read_bytes()?.as_slice())?; - dest_wrt.finish()?; + dest_wrt.terminate()?; } Ok(()) } @@ -296,7 +296,7 @@ mod tests { use std::path::Path; use super::{RamDirectory, MEMORY_USAGE_UPDATE_THRESHOLD}; - use crate::directory::FinishableWrite; + use crate::directory::TerminatingWrite; use crate::Directory; #[test] @@ -309,7 +309,7 @@ mod tests { assert!(directory.atomic_write(path_atomic, msg_atomic).is_ok()); let mut wrt = directory.open_write(path_seq).unwrap(); assert!(wrt.write_all(msg_seq).is_ok()); - assert!(wrt.finish().is_ok()); + assert!(wrt.terminate().is_ok()); let directory_copy = RamDirectory::create(); assert!(directory.persist(&directory_copy).is_ok()); assert_eq!(directory_copy.atomic_read(path_atomic).unwrap(), msg_atomic); @@ -324,7 +324,7 @@ mod tests { assert!(dir.open_write(path).is_err()); drop(writer); - dir.open_write(path).unwrap().finish().unwrap(); + dir.open_write(path).unwrap().terminate().unwrap(); assert!(dir.exists(path).unwrap()); } @@ -352,7 +352,7 @@ mod tests { assert_eq!(dir.clone().total_mem_usage(), dir.total_mem_usage()); let file_len = 1 + MEMORY_USAGE_UPDATE_THRESHOLD + first_capacity + 1; - writer.finish().unwrap(); + writer.terminate().unwrap(); assert_eq!(dir.total_mem_usage(), file_len); dir.delete(path).unwrap(); diff --git a/src/directory/tests.rs b/src/directory/tests.rs index 1d3b01c8a..a9b599910 100644 --- a/src/directory/tests.rs +++ b/src/directory/tests.rs @@ -123,7 +123,7 @@ fn test_simple(directory: &dyn Directory) -> crate::Result<()> { write_file.write_all(&[4])?; write_file.write_all(&[3])?; write_file.write_all(&[7, 3, 5])?; - write_file.finish()?; + write_file.terminate()?; let read_file = directory.open_read(test_path)?.read_bytes()?; assert_eq!(read_file.as_slice(), &[4u8, 3u8, 7u8, 3u8, 5u8]); mem::drop(read_file); @@ -136,7 +136,7 @@ fn test_rewrite_forbidden(directory: &dyn Directory) -> crate::Result<()> { let test_path: &'static Path = Path::new("some_path_for_test"); let writer = directory.open_write(test_path)?; assert!(directory.open_write(test_path).is_err()); - writer.finish()?; + writer.terminate()?; assert!(directory.exists(test_path).unwrap()); assert!(directory.open_write(test_path).is_err()); assert!(directory.delete(test_path).is_ok()); @@ -146,7 +146,7 @@ fn test_rewrite_forbidden(directory: &dyn Directory) -> crate::Result<()> { fn test_write_create_the_file(directory: &dyn Directory) { let test_path: &'static Path = Path::new("some_path_for_test"); assert!(directory.open_read(test_path).is_err()); - directory.open_write(test_path).unwrap().finish().unwrap(); + directory.open_write(test_path).unwrap().terminate().unwrap(); assert!(directory.exists(test_path).unwrap()); assert!(directory.open_read(test_path).is_ok()); assert!(directory.delete(test_path).is_ok()); @@ -157,7 +157,7 @@ fn test_directory_delete(directory: &dyn Directory) -> crate::Result<()> { assert!(directory.open_read(test_path).is_err()); let mut write_file = directory.open_write(test_path)?; write_file.write_all(&[1, 2, 3, 4])?; - write_file.finish()?; + write_file.terminate()?; { let read_handle = directory.open_read(test_path)?.read_bytes()?; assert_eq!(read_handle.as_slice(), &[1u8, 2u8, 3u8, 4u8]); diff --git a/src/fastfield/alive_bitset.rs b/src/fastfield/alive_bitset.rs index 3ca05ed13..ed2f8dfe8 100644 --- a/src/fastfield/alive_bitset.rs +++ b/src/fastfield/alive_bitset.rs @@ -8,7 +8,7 @@ use crate::DocId; /// Write an alive `BitSet` /// /// where `alive_bitset` is the set of alive `DocId`. -/// Warning: this function does not call `finish()`. The caller is in charge of +/// Warning: this function does not call `terminate()`. The caller is in charge of /// closing the writer properly. pub fn write_alive_bitset(alive_bitset: &BitSet, writer: &mut T) -> io::Result<()> { alive_bitset.serialize(writer)?; diff --git a/src/fastfield/mod.rs b/src/fastfield/mod.rs index 724dbcfaf..d56dc27a8 100644 --- a/src/fastfield/mod.rs +++ b/src/fastfield/mod.rs @@ -82,7 +82,7 @@ mod tests { use std::path::Path; use columnar::StrColumn; - use common::{ByteCount, DateTimePrecision, FinishableWrite, HasLen}; + use common::{ByteCount, DateTimePrecision, HasLen, TerminatingWrite}; use once_cell::sync::Lazy; use rand::prelude::SliceRandom; use rand::rngs::StdRng; @@ -132,7 +132,7 @@ mod tests { .add_document(&doc!(*FIELD=>2u64)) .unwrap(); fast_field_writers.serialize(&mut write, None).unwrap(); - write.finish().unwrap(); + write.terminate().unwrap(); } let file = directory.open_read(path).unwrap(); @@ -183,7 +183,7 @@ mod tests { .add_document(&doc!(*FIELD=>215u64)) .unwrap(); fast_field_writers.serialize(&mut write, None).unwrap(); - write.finish().unwrap(); + write.terminate().unwrap(); } let file = directory.open_read(path).unwrap(); assert_eq!(file.len(), 108); @@ -216,7 +216,7 @@ mod tests { .unwrap(); } fast_field_writers.serialize(&mut write, None).unwrap(); - write.finish().unwrap(); + write.terminate().unwrap(); } let file = directory.open_read(path).unwrap(); assert_eq!(file.len(), 81); @@ -248,7 +248,7 @@ mod tests { .unwrap(); } fast_field_writers.serialize(&mut write, None).unwrap(); - write.finish().unwrap(); + write.terminate().unwrap(); } let file = directory.open_read(path).unwrap(); assert_eq!(file.len(), 4476); @@ -281,7 +281,7 @@ mod tests { fast_field_writers.add_document(&doc).unwrap(); } fast_field_writers.serialize(&mut write, None).unwrap(); - write.finish().unwrap(); + write.terminate().unwrap(); } let file = directory.open_read(path).unwrap(); assert_eq!(file.len(), 252); @@ -320,7 +320,7 @@ mod tests { let doc = TantivyDocument::default(); fast_field_writers.add_document(&doc).unwrap(); fast_field_writers.serialize(&mut write, None).unwrap(); - write.finish().unwrap(); + write.terminate().unwrap(); } let file = directory.open_read(path).unwrap(); @@ -353,7 +353,7 @@ mod tests { let doc = TantivyDocument::default(); fast_field_writers.add_document(&doc).unwrap(); fast_field_writers.serialize(&mut write, None).unwrap(); - write.finish().unwrap(); + write.terminate().unwrap(); } let file = directory.open_read(path).unwrap(); @@ -390,7 +390,7 @@ mod tests { fast_field_writers.add_document(&doc!(*FIELD=>x)).unwrap(); } fast_field_writers.serialize(&mut write, None).unwrap(); - write.finish().unwrap(); + write.terminate().unwrap(); } let file = directory.open_read(path).unwrap(); let fast_field_readers = FastFieldReaders::open(file, SCHEMA.clone()).unwrap(); @@ -775,7 +775,7 @@ mod tests { .add_document(&doc!(field=>false)) .unwrap(); fast_field_writers.serialize(&mut write, None).unwrap(); - write.finish().unwrap(); + write.terminate().unwrap(); } let file = directory.open_read(path).unwrap(); assert_eq!(file.len(), 84); @@ -807,7 +807,7 @@ mod tests { .unwrap(); } fast_field_writers.serialize(&mut write, None).unwrap(); - write.finish().unwrap(); + write.terminate().unwrap(); } let file = directory.open_read(path).unwrap(); assert_eq!(file.len(), 96); @@ -832,7 +832,7 @@ mod tests { let doc = TantivyDocument::default(); fast_field_writers.add_document(&doc).unwrap(); fast_field_writers.serialize(&mut write, None).unwrap(); - write.finish().unwrap(); + write.terminate().unwrap(); } let file = directory.open_read(path).unwrap(); assert_eq!(file.len(), 86); @@ -860,7 +860,7 @@ mod tests { fast_field_writers.add_document(doc).unwrap(); } fast_field_writers.serialize(&mut write, None).unwrap(); - write.finish().unwrap(); + write.terminate().unwrap(); } Ok(directory) } diff --git a/src/fastfield/plugin.rs b/src/fastfield/plugin.rs index 5131f51d5..eb33b6816 100644 --- a/src/fastfield/plugin.rs +++ b/src/fastfield/plugin.rs @@ -9,7 +9,7 @@ use std::collections::BTreeMap; use columnar::{ ColumnType, ColumnarReader, MergeRowOrder, RowAddr, ShuffleMergeOrder, StackMergeOrder, }; -use common::FinishableWrite; +use common::TerminatingWrite; use measure_time::debug_time; use crate::directory::{Directory, WritePtr}; @@ -60,7 +60,7 @@ impl SegmentPlugin for FastFieldsPlugin { &mut fast_field_wrt, )?; - fast_field_wrt.finish()?; + fast_field_wrt.terminate()?; Ok(()) } @@ -131,7 +131,7 @@ impl PluginWriter for FastFieldsPluginWriter { self.writer .serialize(&mut self.fast_field_write, doc_id_map) .map_err(|e| crate::TantivyError::InternalError(e.to_string()))?; - self.fast_field_write.finish()?; + self.fast_field_write.terminate()?; Ok(()) } diff --git a/src/indexer/index_writer.rs b/src/indexer/index_writer.rs index cb0862907..22a251ccf 100644 --- a/src/indexer/index_writer.rs +++ b/src/indexer/index_writer.rs @@ -9,7 +9,7 @@ use smallvec::smallvec; use super::operation::{AddOperation, UserOperation}; use super::segment_updater::SegmentUpdater; use super::{AddBatch, AddBatchReceiver, AddBatchSender, PreparedCommit}; -use crate::directory::{DirectoryLock, FinishableWrite, GarbageCollectionResult}; +use crate::directory::{DirectoryLock, GarbageCollectionResult, TerminatingWrite}; use crate::error::TantivyError; use crate::fastfield::write_alive_bitset; use crate::index::{Index, Segment, SegmentComponent, SegmentId, SegmentMeta, SegmentReader}; @@ -172,7 +172,7 @@ pub fn advance_deletes( segment = segment.with_delete_meta(num_deleted_docs, target_opstamp); let mut alive_doc_file = segment.open_write(SegmentComponent::Delete)?; write_alive_bitset(&alive_bitset, &mut alive_doc_file)?; - alive_doc_file.finish()?; + alive_doc_file.terminate()?; } segment_entry.set_meta(segment.meta().clone()); diff --git a/src/plugin.rs b/src/plugin.rs index a0f040656..b66c56124 100644 --- a/src/plugin.rs +++ b/src/plugin.rs @@ -192,7 +192,7 @@ mod tests { let mut write = ctx.target_segment.open_write(component)?; use std::io::Write; write.write_all(&MARKER.to_le_bytes())?; - common::FinishableWrite::finish(write)?; + common::TerminatingWrite::terminate(write)?; Ok(()) } } @@ -258,7 +258,7 @@ mod tests { write.write_all(&(payload.len() as u32).to_le_bytes())?; write.write_all(payload)?; } - common::FinishableWrite::finish(write)?; + common::TerminatingWrite::terminate(write)?; Ok(()) } diff --git a/src/store/store_compressor.rs b/src/store/store_compressor.rs index 683b6d30b..20211b25a 100644 --- a/src/store/store_compressor.rs +++ b/src/store/store_compressor.rs @@ -3,7 +3,7 @@ use std::sync::mpsc::{sync_channel, Receiver, SyncSender}; use std::thread::JoinHandle; use std::{io, thread}; -use common::{BinarySerializable, CountingWriter, FinishableWrite}; +use common::{BinarySerializable, CountingWriter, TerminatingWrite}; use super::DOC_STORE_VERSION; use crate::directory::WritePtr; @@ -151,7 +151,7 @@ impl BlockCompressorImpl { ); self.offset_index_writer.serialize_into(&mut self.writer)?; docstore_footer.serialize(&mut self.writer)?; - self.writer.finish() + self.writer.terminate() } } diff --git a/src/termdict/tests.rs b/src/termdict/tests.rs index 050dd4bc4..71b3f1c3e 100644 --- a/src/termdict/tests.rs +++ b/src/termdict/tests.rs @@ -2,7 +2,7 @@ use std::path::PathBuf; use std::{io, str}; use super::{TermDictionary, TermDictionaryBuilder, TermStreamer}; -use crate::directory::{Directory, FileSlice, FinishableWrite, RamDirectory}; +use crate::directory::{Directory, FileSlice, RamDirectory, TerminatingWrite}; use crate::postings::TermInfo; const BLOCK_SIZE: usize = 1_500; @@ -41,7 +41,7 @@ fn test_term_ordinals() -> crate::Result<()> { for term in COUNTRIES.iter() { term_dictionary_builder.insert(term.as_bytes(), &make_term_info(0u64))?; } - term_dictionary_builder.finish()?.finish()?; + term_dictionary_builder.finish()?.terminate()?; } let term_file = directory.open_read(&path)?; let term_dict: TermDictionary = TermDictionary::open(term_file)?; @@ -63,7 +63,7 @@ fn test_term_dictionary_simple() -> crate::Result<()> { let mut term_dictionary_builder = TermDictionaryBuilder::create(write)?; term_dictionary_builder.insert("abc".as_bytes(), &make_term_info(34u64))?; term_dictionary_builder.insert("abcd".as_bytes(), &make_term_info(346u64))?; - term_dictionary_builder.finish()?.finish()?; + term_dictionary_builder.finish()?.terminate()?; } let file = directory.open_read(&path)?; let term_dict: TermDictionary = TermDictionary::open(file)?; @@ -412,7 +412,7 @@ fn test_automaton_search() -> crate::Result<()> { for term in COUNTRIES.iter() { term_dictionary_builder.insert(term.as_bytes(), &make_term_info(0u64))?; } - term_dictionary_builder.finish()?.finish()?; + term_dictionary_builder.finish()?.terminate()?; } let file = directory.open_read(&path)?; let term_dict: TermDictionary = TermDictionary::open(file)?; diff --git a/sstable/src/index/mod.rs b/sstable/src/index/mod.rs index b70bdf369..f927379be 100644 --- a/sstable/src/index/mod.rs +++ b/sstable/src/index/mod.rs @@ -265,7 +265,7 @@ impl SSTableIndexBuilder { } let counting_writer = map_builder.into_inner().map_err(fst_error_to_io_error)?; let written_bytes = counting_writer.written_bytes(); - let mut wrt = counting_writer.into_inner(); + let mut wrt = counting_writer.finish(); let mut block_store_writer = v3::BlockAddrStoreWriter::new(); for block in &self.blocks { diff --git a/sstable/src/lib.rs b/sstable/src/lib.rs index 082fc2f25..1f6bd14c7 100644 --- a/sstable/src/lib.rs +++ b/sstable/src/lib.rs @@ -357,7 +357,7 @@ where SSTABLE_VERSION.serialize(&mut wrt)?; - let wrt = wrt.into_inner(); + let wrt = wrt.finish(); Ok(wrt.into_inner()?) } } diff --git a/tests/failpoints/mod.rs b/tests/failpoints/mod.rs index be6ad5b7f..213c86628 100644 --- a/tests/failpoints/mod.rs +++ b/tests/failpoints/mod.rs @@ -1,6 +1,6 @@ use std::path::Path; -use tantivy::directory::{Directory, FinishableWrite, ManagedDirectory, RamDirectory}; +use tantivy::directory::{Directory, ManagedDirectory, RamDirectory, TerminatingWrite}; use tantivy::schema::{Schema, TEXT}; use tantivy::{doc, Index, IndexWriter, Term}; @@ -15,7 +15,7 @@ fn test_failpoints_managed_directory_gc_if_delete_fails() { managed_directory .open_write(test_path) .unwrap() - .finish() + .terminate() .unwrap(); assert!(managed_directory.exists(test_path).unwrap()); // triggering gc and setting the delete operation to fail. From a9d3f1a02694b9cc463a4532be9ab4efa5de18e3 Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Mon, 14 Sep 2026 13:09:38 +0200 Subject: [PATCH 12/49] Run nightly fmt --- benches/merge_segments.rs | 2 +- src/directory/mmap_directory/mod.rs | 6 +++++- src/directory/ram_directory.rs | 6 +++--- src/directory/tests.rs | 6 +++++- 4 files changed, 14 insertions(+), 6 deletions(-) diff --git a/benches/merge_segments.rs b/benches/merge_segments.rs index c2fcbb6f1..7e911f7ec 100644 --- a/benches/merge_segments.rs +++ b/benches/merge_segments.rs @@ -16,7 +16,7 @@ use rand::rngs::StdRng; use rand::SeedableRng; use tantivy::directory::error::{DeleteError, OpenReadError, OpenWriteError}; use tantivy::directory::{ - AntiCallToken, Directory, FileHandle, TerminatingWrite, OwnedBytes, WatchCallback, WatchHandle, + AntiCallToken, Directory, FileHandle, OwnedBytes, TerminatingWrite, WatchCallback, WatchHandle, WritePtr, }; use tantivy::indexer::{merge_filtered_segments, NoMergePolicy}; diff --git a/src/directory/mmap_directory/mod.rs b/src/directory/mmap_directory/mod.rs index 1707260f2..061a4ca5f 100644 --- a/src/directory/mmap_directory/mod.rs +++ b/src/directory/mmap_directory/mod.rs @@ -554,7 +554,11 @@ mod tests { // In that case the directory returns a SharedVecSlice. let mmap_directory = MmapDirectory::create_from_tempdir().unwrap(); let path = PathBuf::from("test"); - mmap_directory.open_write(&path).unwrap().terminate().unwrap(); + mmap_directory + .open_write(&path) + .unwrap() + .terminate() + .unwrap(); let readonlymap = mmap_directory.open_read(&path).unwrap(); assert_eq!(readonlymap.len(), 0); } diff --git a/src/directory/ram_directory.rs b/src/directory/ram_directory.rs index 2e6609353..87b5b46a9 100644 --- a/src/directory/ram_directory.rs +++ b/src/directory/ram_directory.rs @@ -87,9 +87,9 @@ impl Drop for VecWriter { fn drop(&mut self) { if !self.is_finished { warn!( - "You forgot to terminate {:?} before its writer got Drop. Do not rely on drop. This \ - also occurs when the indexer crashed, so you may want to check the logs for the \ - root cause.", + "You forgot to terminate {:?} before its writer got Drop. Do not rely on drop. \ + This also occurs when the indexer crashed, so you may want to check the logs for \ + the root cause.", self.path ); } diff --git a/src/directory/tests.rs b/src/directory/tests.rs index a9b599910..c920b375f 100644 --- a/src/directory/tests.rs +++ b/src/directory/tests.rs @@ -146,7 +146,11 @@ fn test_rewrite_forbidden(directory: &dyn Directory) -> crate::Result<()> { fn test_write_create_the_file(directory: &dyn Directory) { let test_path: &'static Path = Path::new("some_path_for_test"); assert!(directory.open_read(test_path).is_err()); - directory.open_write(test_path).unwrap().terminate().unwrap(); + directory + .open_write(test_path) + .unwrap() + .terminate() + .unwrap(); assert!(directory.exists(test_path).unwrap()); assert!(directory.open_read(test_path).is_ok()); assert!(directory.delete(test_path).is_ok()); From dfe904667e49c2da9f8dbcea40ff2aefa15e2cb3 Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Mon, 14 Sep 2026 13:13:17 +0200 Subject: [PATCH 13/49] Minimize unrelated changes --- common/src/writer.rs | 6 ++---- src/directory/directory.rs | 20 +++++++++++++------- src/directory/footer.rs | 4 ++-- src/directory/mmap_directory/mod.rs | 3 ++- src/directory/ram_directory.rs | 4 +--- src/fastfield/alive_bitset.rs | 2 +- src/positions/serializer.rs | 6 +++--- 7 files changed, 24 insertions(+), 21 deletions(-) diff --git a/common/src/writer.rs b/common/src/writer.rs index f871a9b9e..8056b11d8 100644 --- a/common/src/writer.rs +++ b/common/src/writer.rs @@ -65,7 +65,7 @@ pub struct AntiCallToken(()); /// /// Thread-safety is enforced at the call sites that require it. pub trait TerminatingWrite: Write { - /// Indicates that the writer will no longer be used. Internally calls `terminate_ref`. + /// Indicate that the writer will no longer be used. Internally call terminate_ref. fn terminate(mut self) -> io::Result<()> where Self: Sized { self.terminate_ref(AntiCallToken(())) @@ -84,9 +84,7 @@ impl TerminatingWrite for Box { impl TerminatingWrite for BufWriter { fn terminate_ref(&mut self, a: AntiCallToken) -> io::Result<()> { - if !self.buffer().is_empty() { - self.flush()?; - } + self.flush()?; self.get_mut().terminate_ref(a) } } diff --git a/src/directory/directory.rs b/src/directory/directory.rs index 46b3c669f..9e3a7b14d 100644 --- a/src/directory/directory.rs +++ b/src/directory/directory.rs @@ -139,15 +139,21 @@ pub trait Directory: DirectoryClone + fmt::Debug + Send + Sync + 'static { /// Opens a writer for the *virtual file* associated with /// a [`Path`]. /// - /// Depending on the directory implementation, [`Directory::sync_directory()`] may be required - /// after terminating the writer to ensure that the file is durably created. + /// After the writer is terminated, the file should be created and any subsequent call to + /// [`Directory::open_read()`] for the same path should return a [`FileSlice`]. /// - /// Write operations may be aggressively buffered. The client must call - /// [`TerminatingWrite::terminate()`] to finalize the file and make all writes available to - /// subsequent reads. The directory implementation owns its buffering strategy; clients should - /// not rely on `flush()` making an incomplete file available. + /// However, depending on the directory implementation, + /// it might be required to call [`Directory::sync_directory()`] to ensure + /// that the file is durably created. + /// (The semantics here are the same when dealing with + /// a POSIX filesystem.) /// - /// The user shall not rely on [`Drop`] finalizing the file. + /// Write operations may be aggressively buffered. + /// The client of this trait is responsible for calling terminate + /// to ensure that subsequent `read` operations + /// will take into account preceding `write` operations. + /// + /// The user shall not rely on [`Drop`] triggering terminate. /// /// The file may not previously exist. fn open_write(&self, path: &Path) -> Result; diff --git a/src/directory/footer.rs b/src/directory/footer.rs index d2ff52f9a..bffa2f2cf 100644 --- a/src/directory/footer.rs +++ b/src/directory/footer.rs @@ -126,9 +126,9 @@ impl Footer { } pub(crate) struct FooterProxy { - /// Always `Some` except after `terminate()` is called. + /// always Some except after terminate call hasher: Option, - /// Always `Some` except after `terminate()` is called. + /// always Some except after terminate call writer: Option, } diff --git a/src/directory/mmap_directory/mod.rs b/src/directory/mmap_directory/mod.rs index 061a4ca5f..1c33f0358 100644 --- a/src/directory/mmap_directory/mod.rs +++ b/src/directory/mmap_directory/mod.rs @@ -318,7 +318,8 @@ impl Drop for ReleaseLockFile { } } -/// Wraps a file and syncs its data when the writer is terminated. +/// This Write wraps a File, but has the specificity of +/// calling `sync_all` on terminate. struct SafeFileWriter(File); impl SafeFileWriter { diff --git a/src/directory/ram_directory.rs b/src/directory/ram_directory.rs index 87b5b46a9..bfa3d0ecb 100644 --- a/src/directory/ram_directory.rs +++ b/src/directory/ram_directory.rs @@ -60,7 +60,7 @@ impl Drop for MemoryUsageTracker { /// Writer associated with the [`RamDirectory`]. /// -/// The writer stores its buffer in the directory when terminated. +/// The Writer just writes a buffer. struct VecWriter { path: PathBuf, shared_directory: RamDirectory, @@ -109,8 +109,6 @@ impl Write for VecWriter { Ok(buf.len()) } - /// Nothing to flush since the data is stored in memory. The memory usage is updated on each - /// write. fn flush(&mut self) -> io::Result<()> { Ok(()) } diff --git a/src/fastfield/alive_bitset.rs b/src/fastfield/alive_bitset.rs index ed2f8dfe8..bbdc82a45 100644 --- a/src/fastfield/alive_bitset.rs +++ b/src/fastfield/alive_bitset.rs @@ -8,7 +8,7 @@ use crate::DocId; /// Write an alive `BitSet` /// /// where `alive_bitset` is the set of alive `DocId`. -/// Warning: this function does not call `terminate()`. The caller is in charge of +/// Warning: this function does not call terminate. The caller is in charge of /// closing the writer properly. pub fn write_alive_bitset(alive_bitset: &BitSet, writer: &mut T) -> io::Result<()> { alive_bitset.serialize(writer)?; diff --git a/src/positions/serializer.rs b/src/positions/serializer.rs index 617033909..f41923e8b 100644 --- a/src/positions/serializer.rs +++ b/src/positions/serializer.rs @@ -86,8 +86,8 @@ impl PositionSerializer { Ok(()) } - /// Close the positions for this field. - pub fn close(self) -> io::Result<()> { - Ok(()) + /// Close the positions for this term and flushes the data. + pub fn close(mut self) -> io::Result<()> { + self.positions_wrt.flush() } } From e56587f3dfe14135d770c810dde3fc7424a23696 Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Wed, 16 Sep 2026 16:49:20 +0200 Subject: [PATCH 14/49] Rename memory usage tracker finish to release --- src/directory/ram_directory.rs | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/src/directory/ram_directory.rs b/src/directory/ram_directory.rs index bfa3d0ecb..db1fe924f 100644 --- a/src/directory/ram_directory.rs +++ b/src/directory/ram_directory.rs @@ -42,7 +42,7 @@ impl MemoryUsageTracker { } } - fn finish(&mut self) { + fn release(&mut self) { if self.reported_bytes > 0 { self.shared_usage .fetch_sub(self.reported_bytes, Ordering::Relaxed); @@ -54,7 +54,7 @@ impl MemoryUsageTracker { impl Drop for MemoryUsageTracker { fn drop(&mut self) { - self.finish(); + self.release(); } } @@ -122,7 +122,7 @@ impl TerminatingWrite for VecWriter { fs.active_writers.remove(&self.path); fs.write_owned(self.path.clone(), data); self.is_finished = true; - self.memory_usage.finish(); + self.memory_usage.release(); Ok(()) } } From 1680885e683d8492d2c47cb8034dac3ff5dc8418 Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Wed, 16 Sep 2026 17:07:14 +0200 Subject: [PATCH 15/49] Report active RamDirectory writers as existing --- src/directory/ram_directory.rs | 2 +- src/directory/tests.rs | 7 ++----- 2 files changed, 3 insertions(+), 6 deletions(-) diff --git a/src/directory/ram_directory.rs b/src/directory/ram_directory.rs index db1fe924f..2ce9d39d1 100644 --- a/src/directory/ram_directory.rs +++ b/src/directory/ram_directory.rs @@ -158,7 +158,7 @@ impl InnerDirectory { } fn exists(&self, path: &Path) -> bool { - self.fs.contains_key(path) + self.fs.contains_key(path) || self.active_writers.contains(path) } fn watch(&mut self, watch_handle: WatchCallback) -> WatchHandle { diff --git a/src/directory/tests.rs b/src/directory/tests.rs index c920b375f..dd6d5de6a 100644 --- a/src/directory/tests.rs +++ b/src/directory/tests.rs @@ -146,12 +146,9 @@ fn test_rewrite_forbidden(directory: &dyn Directory) -> crate::Result<()> { fn test_write_create_the_file(directory: &dyn Directory) { let test_path: &'static Path = Path::new("some_path_for_test"); assert!(directory.open_read(test_path).is_err()); - directory - .open_write(test_path) - .unwrap() - .terminate() - .unwrap(); + let writer = directory.open_write(test_path).unwrap(); assert!(directory.exists(test_path).unwrap()); + writer.terminate().unwrap(); assert!(directory.open_read(test_path).is_ok()); assert!(directory.delete(test_path).is_ok()); } From e69acca220f872a9fd09990c71360508724d2ef4 Mon Sep 17 00:00:00 2001 From: David Yaffe Date: Wed, 16 Sep 2026 14:55:28 -0400 Subject: [PATCH 16/49] Fix cardinality calls for datasketches 0.5 --- src/aggregation/metric/cardinality.rs | 13 +++++++------ 1 file changed, 7 insertions(+), 6 deletions(-) diff --git a/src/aggregation/metric/cardinality.rs b/src/aggregation/metric/cardinality.rs index ca0a00fb2..0dc326418 100644 --- a/src/aggregation/metric/cardinality.rs +++ b/src/aggregation/metric/cardinality.rs @@ -168,7 +168,7 @@ impl CouponCache { if should_use_dense { // We don't really care about the value here. We will populate all the values we will // read anyway. - let uninitialized_coupon = Coupon::from_hash(0); + let uninitialized_coupon = Coupon::from_value(0); let mut coupon_map: Vec = vec![uninitialized_coupon; highest_term_ord as usize + 1]; @@ -557,7 +557,7 @@ fn build_coupon_cache( let mut coupons: Vec = Vec::with_capacity(term_ords.len()); let all_term_ords_found: bool = dictionary.sorted_ords_to_term_cb(&term_ords, |term_bytes| { - let coupon: Coupon = Coupon::from_hash(term_bytes); + let coupon: Coupon = Coupon::from_value(term_bytes); coupons.push(coupon); })?; assert!(all_term_ords_found); @@ -566,14 +566,14 @@ fn build_coupon_cache( // we populate the cache with the missing key too (if any). let missing_coupon_opt: Option = missing_value_opt.map(|missing_key| { if let Key::Str(missing_value_str) = missing_key { - Coupon::from_hash(missing_value_str.as_bytes()) + Coupon::from_value(missing_value_str.as_bytes()) } else { // See https://github.com/quickwit-oss/tantivy/issues/2891 // A missing key with a type different from Str will not work as intended // for the moment. // // Right now this is just a partial workaround. - Coupon::from_hash("__tantivy_missing_non_str__".as_bytes()) + Coupon::from_value("__tantivy_missing_non_str__".as_bytes()) } }); Ok(CouponCache::new(term_ords, coupons, missing_coupon_opt)) @@ -826,7 +826,8 @@ impl<'de> Deserialize<'de> for CardinalityCollector { impl CardinalityCollector { fn new(salt: u8) -> Self { Self { - sketch: HllSketch::new(LG_K, HllType::Hll8), + sketch: HllSketch::new(LG_K, HllType::Hll8) + .expect("LG_K is within the supported range"), salt, } } @@ -854,7 +855,7 @@ impl CardinalityCollector { } pub(crate) fn merge_fruits(&mut self, right: CardinalityCollector) -> crate::Result<()> { - let mut union = HllUnion::new(LG_K); + let mut union = HllUnion::new(LG_K).expect("LG_K is within the supported range"); union.update(&self.sketch); union.update(&right.sketch); self.sketch = union.to_sketch(HllType::Hll8); From 1496f58963849213e01e3c48cb06229f57febbbb Mon Sep 17 00:00:00 2001 From: Paul Masurel Date: Thu, 17 Sep 2026 10:34:15 +0200 Subject: [PATCH 17/49] Added jitexpr compilation cache (#3107) * Added jitexpr compilation cache * Added jitexpr to the CI test matrix --- .github/workflows/test.yml | 4 +- jitexpr/Cargo.toml | 1 + jitexpr/examples/basic.rs | 5 +- jitexpr/src/ast/infer_types.rs | 2 +- jitexpr/src/ast/literal.rs | 45 +-- jitexpr/src/ast/mod.rs | 2 +- jitexpr/src/ast/serde.rs | 40 ++- jitexpr/src/compile/cache.rs | 306 ++++++++++++++++++ jitexpr/src/compile/compile_fn_builder.rs | 14 +- jitexpr/src/compile/error.rs | 8 +- jitexpr/src/compile/mod.rs | 4 +- jitexpr/src/compile/typed_expr.rs | 26 +- jitexpr/src/compile/typed_expr_serialize.rs | 2 +- jitexpr/src/functions/add.rs | 22 +- jitexpr/src/functions/is_null.rs | 27 +- jitexpr/src/functions/left.rs | 2 +- jitexpr/src/functions/mod.rs | 2 +- jitexpr/src/functions/right.rs | 2 +- jitexpr/src/functions/round.rs | 9 +- jitexpr/src/functions/split_after.rs | 2 +- jitexpr/src/functions/split_before.rs | 2 +- jitexpr/src/functions/substring.rs | 2 +- jitexpr/src/types.rs | 193 ++++++++++- src/index/index.rs | 23 +- src/index/segment_reader.rs | 5 + .../doc_predicate_query/jitexpr_predicate.rs | 9 +- 26 files changed, 667 insertions(+), 92 deletions(-) create mode 100644 jitexpr/src/compile/cache.rs diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index bbb967922..d9e468d7c 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -70,8 +70,8 @@ jobs: strategy: matrix: features: - - { label: "all", flags: "mmap,stopwords,lz4-compression,zstd-compression,failpoints,stemmer" } - - { label: "quickwit", flags: "mmap,quickwit,failpoints" } + - { label: "all", flags: "mmap,stopwords,lz4-compression,zstd-compression,failpoints,stemmer,jitexpr" } + - { label: "quickwit", flags: "mmap,quickwit,failpoints,jitexpr" } - { label: "none", flags: "" } name: test-${{ matrix.features.label}} diff --git a/jitexpr/Cargo.toml b/jitexpr/Cargo.toml index e770cb802..e68d38323 100644 --- a/jitexpr/Cargo.toml +++ b/jitexpr/Cargo.toml @@ -12,5 +12,6 @@ cranelift = "0.134.3" cranelift-jit = "0.134.3" cranelift-module = "0.134.3" cranelift-native = "0.134.3" +lru = "0.18.2" regex = "1" thiserror = "2.0.1" diff --git a/jitexpr/examples/basic.rs b/jitexpr/examples/basic.rs index bc63bd7e1..d41dcc079 100644 --- a/jitexpr/examples/basic.rs +++ b/jitexpr/examples/basic.rs @@ -2,7 +2,7 @@ use std::collections::HashMap; use std::error::Error; use std::sync::Arc; -use jitexpr::ast::{Function, InferredTypeSet, UntypedExpr, infer_types}; +use jitexpr::ast::{Function, InferredTypeSet, Literal, UntypedExpr, infer_types}; use jitexpr::compile::{CompiledFn, CompiledFnCtx, compile}; use jitexpr::types::{VarType, VariableValue}; @@ -13,7 +13,8 @@ fn main() -> Result<(), Box> { Function::Add, vec![ UntypedExpr::variable("my_col"), - UntypedExpr::literal(1.0f64), + // A float literal must be finite, so the conversion is fallible. + UntypedExpr::literal(Literal::try_from(1.0f64)?), ], )?; diff --git a/jitexpr/src/ast/infer_types.rs b/jitexpr/src/ast/infer_types.rs index 74b8191c3..9abe9bcb9 100644 --- a/jitexpr/src/ast/infer_types.rs +++ b/jitexpr/src/ast/infer_types.rs @@ -136,7 +136,7 @@ impl std::fmt::Display for InferredTypeSet { } } -#[derive(Debug, thiserror::Error)] +#[derive(Debug, thiserror::Error, Clone)] pub enum TypeError { #[error(transparent)] InvalidFnCall(#[from] InvalidFnCall), diff --git a/jitexpr/src/ast/literal.rs b/jitexpr/src/ast/literal.rs index 535290da8..4b3f8c6d6 100644 --- a/jitexpr/src/ast/literal.rs +++ b/jitexpr/src/ast/literal.rs @@ -1,6 +1,7 @@ use std::sync::Arc; use crate::ast::InferredTypeSet; +use crate::types::SafeF64; #[cfg(test)] use crate::types::VarType; @@ -11,7 +12,7 @@ pub enum Literal { Bool(bool), U64(u64), I64(i64), - F64(f64), + F64(SafeF64), String(Arc), } @@ -43,10 +44,11 @@ impl Literal { ..InferredTypeSet::NONE }, Literal::F64(value) => { + let value: f64 = value.get(); let is_integral: bool = value.fract() == 0.0; InferredTypeSet { - i64: is_integral && *value >= i64::MIN as f64 && *value < -(i64::MIN as f64), - u64: is_integral && *value >= 0.0 && *value < u64::MAX as f64, + i64: is_integral && value >= i64::MIN as f64 && value < -(i64::MIN as f64), + u64: is_integral && value >= 0.0 && value < u64::MAX as f64, f64: true, ..InferredTypeSet::NONE } @@ -86,9 +88,18 @@ impl From for Literal { } } -impl From for Literal { - fn from(value: f64) -> Self { - Literal::F64(value) +/// A float could not become a literal because it was NaN or infinite. +#[derive(Clone, Copy, Debug, Eq, PartialEq, thiserror::Error)] +#[error("an f64 literal must be finite, and neither NaN nor infinite")] +pub struct NonFiniteFloat; + +impl TryFrom for Literal { + type Error = NonFiniteFloat; + + /// Fails for NaN and for either infinity: [`SafeF64`] holds only finite + /// values, so those have no literal representation. + fn try_from(value: f64) -> Result { + SafeF64::new(value).map(Literal::F64).ok_or(NonFiniteFloat) } } @@ -108,6 +119,10 @@ impl From<&str> for Literal { mod tests { use super::*; + fn f64_literal(val: f64) -> Literal { + Literal::try_from(val).unwrap() + } + #[test] fn test_literal_types_depend_on_representable_value() { let i64_f64 = InferredTypeSet { @@ -125,8 +140,8 @@ mod tests { assert_eq!(Literal::I64(1).types(), InferredTypeSet::NUMERICAL); assert_eq!(Literal::I64(-1).types(), i64_f64); assert_eq!(Literal::U64(1 << 63).types(), u64_f64); - assert_eq!(Literal::F64(1.2).types(), InferredTypeSet::F64); - assert_eq!(Literal::F64(1.0).types(), InferredTypeSet::NUMERICAL); + assert_eq!(f64_literal(1.2).types(), InferredTypeSet::F64); + assert_eq!(f64_literal(1.0).types(), InferredTypeSet::NUMERICAL); } #[test] @@ -162,22 +177,16 @@ mod tests { } #[test] - fn test_f64_literal_types_handle_integer_boundaries_and_special_values() { + fn test_f64_literal_types_handle_integer_boundaries() { assert_eq!( - Literal::F64(2f64.powi(63)).types(), + f64_literal(2f64.powi(63)).types(), InferredTypeSet { u64: true, f64: true, ..InferredTypeSet::NONE } ); - assert_eq!(Literal::F64(2f64.powi(64)).types(), InferredTypeSet::F64); - assert_eq!(Literal::F64(-0.0).types(), InferredTypeSet::NUMERICAL); - assert_eq!(Literal::F64(f64::NAN).types(), InferredTypeSet::F64); - assert_eq!(Literal::F64(f64::INFINITY).types(), InferredTypeSet::F64); - assert_eq!( - Literal::F64(f64::NEG_INFINITY).types(), - InferredTypeSet::F64 - ); + assert_eq!(f64_literal(2f64.powi(64)).types(), InferredTypeSet::F64); + assert_eq!(f64_literal(-0.0).types(), InferredTypeSet::NUMERICAL); } } diff --git a/jitexpr/src/ast/mod.rs b/jitexpr/src/ast/mod.rs index fe5028a54..9e5dd2f2b 100644 --- a/jitexpr/src/ast/mod.rs +++ b/jitexpr/src/ast/mod.rs @@ -5,7 +5,7 @@ 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; +pub use literal::{Literal, NonFiniteFloat}; pub(crate) use serde::format_variable_name; pub use serde::{DeserializeError, deserialize, serialize}; pub use untyped_expr::UntypedExpr; diff --git a/jitexpr/src/ast/serde.rs b/jitexpr/src/ast/serde.rs index 705ad8d6c..39ae800dc 100644 --- a/jitexpr/src/ast/serde.rs +++ b/jitexpr/src/ast/serde.rs @@ -27,6 +27,7 @@ use std::fmt; use std::sync::Arc; use crate::ast::{Function, Literal, UntypedExpr}; +use crate::types::SafeF64; /// Serializes an untyped expression into its canonical Lisp-like form. pub fn serialize(expr: &UntypedExpr) -> String { @@ -118,7 +119,7 @@ fn format_literal(literal: &Literal, formatter: &mut fmt::Formatter) -> fmt::Res Literal::Bool(value) => write!(formatter, "{value}"), Literal::U64(value) => write!(formatter, "{value}u64"), Literal::I64(value) => write!(formatter, "{value}i64"), - Literal::F64(value) => write!(formatter, "{value}f64"), + Literal::F64(value) => write!(formatter, "{}f64", value.get()), Literal::String(value) => format_quoted(value, '"', formatter), } } @@ -273,7 +274,9 @@ fn parse_literal_atom(atom: &str) -> Option { } if let Some(value_str) = atom.strip_suffix("f64") { let val = value_str.parse::().ok()?; - return Some(Literal::F64(val)); + // Yields `None` for a non-finite float, which `parse_atom` reports as + // an error rather than letting it fall through to a variable name. + return SafeF64::new(val).map(Literal::F64); } None } @@ -372,16 +375,19 @@ impl<'a> Parser<'a> { let atom_offset = self.offset; let atom = self.take_atom(); if let Some(literal) = parse_literal_atom(atom) { - if let Literal::F64(value) = &literal - && !value.is_finite() - { - return Err(DeserializeError::new( - atom_offset, - format!("f64 literal `{atom}` must be finite"), - )); - } return Ok(UntypedExpr::Literal(literal)); } + // An `f64`-suffixed atom that parses as a float but produced no literal + // is non-finite: `SafeF64` cannot hold it, and it must not be mistaken + // for a variable name. + if let Some(value_str) = atom.strip_suffix("f64") + && value_str.parse::().is_ok() + { + return Err(DeserializeError::new( + atom_offset, + format!("f64 literal `{atom}` must be finite"), + )); + } Ok(UntypedExpr::Variable(Arc::from(atom))) } @@ -529,8 +535,14 @@ mod tests { (UntypedExpr::literal(false), "false"), (UntypedExpr::literal(u64::MAX), "18446744073709551615u64"), (UntypedExpr::literal(i64::MIN), "-9223372036854775808i64"), - (UntypedExpr::literal(1.5f64), "1.5f64"), - (UntypedExpr::literal(1.0f64), "1f64"), + ( + UntypedExpr::literal(Literal::try_from(1.5f64).unwrap()), + "1.5f64", + ), + ( + UntypedExpr::literal(Literal::try_from(1.0f64).unwrap()), + "1f64", + ), ]; for (expr, expected) in cases { @@ -586,12 +598,12 @@ mod tests { f64::from_bits(1), -0.0, ] { - let serialized = serialize(&UntypedExpr::literal(value)); + let serialized = serialize(&UntypedExpr::literal(Literal::try_from(value).unwrap())); let UntypedExpr::Literal(Literal::F64(parsed)) = deserialize(&serialized).unwrap() else { panic!("expected an f64 literal"); }; - assert_eq!(parsed.to_bits(), value.to_bits()); + assert_eq!(parsed.get().to_bits(), value.to_bits()); } } diff --git a/jitexpr/src/compile/cache.rs b/jitexpr/src/compile/cache.rs new file mode 100644 index 000000000..54326b027 --- /dev/null +++ b/jitexpr/src/compile/cache.rs @@ -0,0 +1,306 @@ +use std::collections::HashMap; +use std::num::NonZeroUsize; +use std::sync::{Arc, Mutex, OnceLock}; + +use lru::LruCache; + +use super::{CompileError, CompiledFn}; +use crate::ast::UntypedExpr; +use crate::types::VarType; + +/// The outcome of compiling one key, shared by every caller of that key. +type CompilationResult = Result, CompileError>; + +/// A per-key cell, written once by whichever caller compiles the expression. +/// +/// Initializing it is what serializes concurrent compilations of the same key. +type CompilationSlot = Arc>; + +/// A bounded, thread-safe cache of JIT-compiled expressions. +/// The cache is cheap to clone: every clone shares one set of entries. +/// The cache does not allocate on creation. Allocation happens on the first usage. +#[derive(Clone)] +pub struct ExprCompilationCache { + inner: Arc>, +} + +struct ExprCompilationCacheInner { + capacity: usize, + // We use Option here to lazily allocate on the first insertion. + entries: Option>, +} + +impl ExprCompilationCacheInner { + fn entries(&mut self) -> Option<&mut LruCache> { + if self.entries.is_none() { + let non_zero_capacity = NonZeroUsize::new(self.capacity)?; + self.entries = Some(LruCache::new(non_zero_capacity)); + } + self.entries.as_mut() + } +} + +/// Identifies a compilation: an expression plus the types it was compiled for. +#[derive(PartialEq, Eq, Hash)] +struct ExprCacheKey { + expr: String, + /// The variable types, sorted by variable name so the key does not depend + /// on the caller's `HashMap` iteration order. + var_types: Box<[(String, VarType)]>, +} + +impl ExprCacheKey { + fn new(untyped_expr: &UntypedExpr, var_types: &HashMap<&str, VarType>) -> ExprCacheKey { + let mut sorted_var_types: Vec<(String, VarType)> = Vec::with_capacity(var_types.len()); + for (variable_name, var_type) in var_types { + sorted_var_types.push((variable_name.to_string(), *var_type)); + } + sorted_var_types.sort_unstable(); + ExprCacheKey { + expr: untyped_expr.to_string(), + var_types: sorted_var_types.into_boxed_slice(), + } + } +} + +impl ExprCompilationCache { + /// A capacity of 0 means disabled. + pub fn with_capacity(capacity: usize) -> ExprCompilationCache { + ExprCompilationCache { + inner: Arc::new(Mutex::new(ExprCompilationCacheInner { + capacity, + entries: None, + })), + } + } + + /// Creates a cache that memoizes nothing and allocates nothing. + pub fn disabled() -> ExprCompilationCache { + ExprCompilationCache::with_capacity(0) + } + + /// Returns false for a cache that memoizes nothing. + pub fn is_enabled(&self) -> bool { + self.inner.lock().unwrap().capacity > 0 + } + + /// Returns the number of compiled expressions currently retained. + pub fn len(&self) -> usize { + let mut inner_guard = self.inner.lock().unwrap(); + let Some(entries) = inner_guard.entries() else { + return 0; + }; + entries.len() + } + + /// Returns true if the cache retains no compiled expression. + pub fn is_empty(&self) -> bool { + self.len() == 0 + } + + /// Returns the expression compiled for `var_types`, compiling it on a miss. + /// + /// Concurrent callers asking for the same expression and types block until + /// the first of them is done, so an expression is normally compiled once. + /// An expression evicted while a compilation is in flight may be compiled + /// again by a later caller. + /// + /// A compilation failure is cached like a success and handed to later + /// callers. + pub fn compile( + &self, + untyped_expr: &UntypedExpr, + var_types: &HashMap<&str, VarType>, + ) -> Result, CompileError> { + if let Some(slot) = self.slot_opt(untyped_expr, var_types) { + // Initializing the cell ensures we cannot have two threads compiling + // the same function at the same time. + slot.get_or_init(|| super::compile(untyped_expr, var_types)) + .clone() + } else { + // no caching + super::compile(untyped_expr, var_types) + } + } + + fn slot_opt( + &self, + untyped_expr: &UntypedExpr, + var_types: &HashMap<&str, VarType>, + ) -> Option { + let key = ExprCacheKey::new(untyped_expr, var_types); + // That function does take the lock but only does trivial things that + // cannot panick before releasing it. + let mut inner_guard = self.inner.lock().unwrap(); + let entries = inner_guard.entries()?; + let slot: CompilationSlot = entries.get_or_insert(key, CompilationSlot::default).clone(); + Some(slot) + } +} + +#[cfg(test)] +mod tests { + use std::sync::Barrier; + + use super::*; + use crate::ast::Function; + + #[test] + fn test_cache_hit_returns_the_same_compiled_fn() { + let cache = ExprCompilationCache::with_capacity(64); + let untyped_expr = UntypedExpr::variable("flag"); + let variable_types = HashMap::from([("flag", VarType::Bool)]); + let first = cache.compile(&untyped_expr, &variable_types).unwrap(); + let second = cache.compile(&untyped_expr, &variable_types).unwrap(); + assert!(Arc::ptr_eq(&first, &second)); + assert_eq!(cache.len(), 1); + } + + #[test] + fn test_var_types_are_part_of_the_key() { + let cache = ExprCompilationCache::with_capacity(64); + let untyped_expr = UntypedExpr::variable("value"); + + let as_u64 = cache + .compile(&untyped_expr, &HashMap::from([("value", VarType::U64)])) + .unwrap(); + let as_i64 = cache + .compile(&untyped_expr, &HashMap::from([("value", VarType::I64)])) + .unwrap(); + + assert!(!Arc::ptr_eq(&as_u64, &as_i64)); + assert_eq!(as_u64.result_type(), VarType::U64); + assert_eq!(as_i64.result_type(), VarType::I64); + assert_eq!(cache.len(), 2); + } + + #[test] + fn test_var_types_order_is_not_part_of_the_key() { + let untyped_expr = Function::Add + .call(vec![ + UntypedExpr::variable("arg1"), + UntypedExpr::variable("arg2"), + UntypedExpr::variable("arg3"), + UntypedExpr::variable("arg4"), + ]) + .unwrap(); + let var_args = HashMap::from([ + ("arg2", VarType::I64), + ("arg4", VarType::Str), + ("arg3", VarType::F64), + ("arg1", VarType::U64), + ]); + let key = ExprCacheKey::new(&untyped_expr, &var_args); + assert_eq!(key.var_types.len(), 4); + assert_eq!(key.var_types[0].0, "arg1"); + assert_eq!(key.var_types[1].0, "arg2"); + assert_eq!(key.var_types[2].0, "arg3"); + assert_eq!(key.var_types[3].0, "arg4"); + } + + #[test] + fn test_equal_expressions_built_separately_share_an_entry() { + let cache = ExprCompilationCache::with_capacity(64); + let variable_types = HashMap::from([("value", VarType::U64)]); + let make_expr = || { + Function::Add + .call(vec![ + UntypedExpr::variable("value"), + UntypedExpr::literal(1u64), + ]) + .unwrap() + }; + + let first = cache.compile(&make_expr(), &variable_types).unwrap(); + let second = cache.compile(&make_expr(), &variable_types).unwrap(); + + assert!(Arc::ptr_eq(&first, &second)); + assert_eq!(cache.len(), 1); + } + + #[test] + fn test_least_recently_used_entry_is_evicted() { + let cache = ExprCompilationCache::with_capacity(1); + let variable_types = HashMap::from([("flag", VarType::Bool)]); + let first_expr = UntypedExpr::variable("flag"); + let second_expr = Function::Not + .call(vec![UntypedExpr::variable("flag")]) + .unwrap(); + + let first = cache.compile(&first_expr, &variable_types).unwrap(); + assert_eq!(cache.len(), 1); + cache.compile(&second_expr, &variable_types).unwrap(); + assert_eq!(cache.len(), 1); + let first_again = cache.compile(&first_expr, &variable_types).unwrap(); + + assert!(!Arc::ptr_eq(&first, &first_again)); + assert_eq!(cache.len(), 1); + } + + #[test] + fn test_disabled_cache_memoizes_nothing() { + let cache = ExprCompilationCache::disabled(); + let untyped_expr = UntypedExpr::variable("flag"); + let variable_types = HashMap::from([("flag", VarType::Bool)]); + + let first = cache.compile(&untyped_expr, &variable_types).unwrap(); + let second = cache.compile(&untyped_expr, &variable_types).unwrap(); + + assert!(!cache.is_enabled()); + assert!(!Arc::ptr_eq(&first, &second)); + assert_eq!(cache.len(), 0); + assert!(cache.is_empty()); + } + + #[test] + fn test_capacity_0_means_disabled() { + assert!(ExprCompilationCache::with_capacity(1).is_enabled()); + assert!(!ExprCompilationCache::with_capacity(0).is_enabled()); + assert!(!ExprCompilationCache::disabled().is_enabled()); + } + + #[test] + fn test_compilation_failures_are_memoized() { + let cache = ExprCompilationCache::with_capacity(16); + // `(` is not a valid regular expression, which is rejected at compile time. + let untyped_expr = crate::ast::deserialize(r#"(REGEXP_EXTRACT "a" "*(" 1u64)"#).unwrap(); + let variable_types = HashMap::new(); + assert!(cache.is_empty()); + assert!(cache.compile(&untyped_expr, &variable_types).is_err()); + assert_eq!(cache.len(), 1); + } + + #[test] + fn test_concurrent_callers_compile_once() { + const NUM_THREADS: usize = 8; + + let cache = ExprCompilationCache::with_capacity(NUM_THREADS); + let barrier = Barrier::new(NUM_THREADS); + let untyped_expr = UntypedExpr::variable("flag"); + + let compiled_fns: Vec> = std::thread::scope(|scope| { + let handles: Vec<_> = (0..NUM_THREADS) + .map(|_| { + let cache = cache.clone(); + let untyped_expr = &untyped_expr; + let barrier = &barrier; + scope.spawn(move || { + let variable_types = HashMap::from([("flag", VarType::Bool)]); + barrier.wait(); + cache.compile(untyped_expr, &variable_types).unwrap() + }) + }) + .collect(); + handles + .into_iter() + .map(|handle| handle.join().unwrap()) + .collect() + }); + + // Identical pointers can only come from a single compilation. + for compiled_fn in &compiled_fns { + assert!(Arc::ptr_eq(&compiled_fns[0], compiled_fn)); + } + assert_eq!(cache.len(), 1); + } +} diff --git a/jitexpr/src/compile/compile_fn_builder.rs b/jitexpr/src/compile/compile_fn_builder.rs index 1df012ee5..cf6a870fa 100644 --- a/jitexpr/src/compile/compile_fn_builder.rs +++ b/jitexpr/src/compile/compile_fn_builder.rs @@ -15,7 +15,7 @@ use super::{ }; use crate::ast::{InferredTypeSet, Literal, UntypedExpr}; use crate::functions::{declare_native_functions, register_jit_symbols}; -use crate::types::VarType; +use crate::types::{SafeF64, VarType}; pub(crate) struct CompileFnBuilder<'types, 'names> { variable_types: &'types HashMap<&'names str, VarType>, @@ -109,8 +109,8 @@ impl<'types, 'names> CompileFnBuilder<'types, 'names> { match literal { Literal::U64(value) => TypedLiteral::I64(*value as i64), Literal::I64(value) => TypedLiteral::I64(*value), - Literal::F64(value) if f64_to_i64_lossless(*value).is_some() => { - TypedLiteral::I64(f64_to_i64_lossless(*value).unwrap()) + Literal::F64(value) if f64_to_i64_lossless(value.get()).is_some() => { + TypedLiteral::I64(f64_to_i64_lossless(value.get()).unwrap()) } _ => panic!("cannot coerce literal {literal:?} to i64"), } @@ -118,15 +118,15 @@ impl<'types, 'names> CompileFnBuilder<'types, 'names> { match literal { Literal::U64(value) => TypedLiteral::U64(*value), Literal::I64(value) => TypedLiteral::U64(*value as u64), - Literal::F64(value) if f64_to_u64_lossless(*value).is_some() => { - TypedLiteral::U64(f64_to_u64_lossless(*value).unwrap()) + Literal::F64(value) if f64_to_u64_lossless(value.get()).is_some() => { + TypedLiteral::U64(f64_to_u64_lossless(value.get()).unwrap()) } _ => panic!("cannot coerce literal {literal:?} to u64"), } } else if intersection.contains(VarType::F64) { match literal { - Literal::U64(value) => TypedLiteral::F64(*value as f64), - Literal::I64(value) => TypedLiteral::F64(*value as f64), + Literal::U64(value) => TypedLiteral::F64(SafeF64::from_integer(*value)), + Literal::I64(value) => TypedLiteral::F64(SafeF64::from_integer(*value)), Literal::F64(value) => TypedLiteral::F64(*value), _ => panic!("cannot coerce literal {literal:?} to f64"), } diff --git a/jitexpr/src/compile/error.rs b/jitexpr/src/compile/error.rs index 2e6a48110..76f9339b9 100644 --- a/jitexpr/src/compile/error.rs +++ b/jitexpr/src/compile/error.rs @@ -1,12 +1,14 @@ +use std::sync::Arc; + use crate::ast::{Function, InvalidFnCall, TypeError}; use crate::types::VarType; -#[derive(Debug, thiserror::Error)] +#[derive(Debug, Clone, thiserror::Error)] pub enum CompileError { #[error("type inference failed: {0}")] TypeInference(#[from] TypeError), #[error("JIT compilation failed: {0}")] - Module(#[source] Box), + Module(#[source] Arc), #[error("cannot coerce an expression from {from_type:?} to {target:?}")] UnsupportedCoercion { from_type: VarType, target: VarType }, #[error("cannot compile {function:?} with result type {return_type:?}")] @@ -26,6 +28,6 @@ pub enum CompileError { impl From for CompileError { fn from(error: cranelift_module::ModuleError) -> Self { - CompileError::Module(Box::new(error)) + CompileError::Module(Arc::new(error)) } } diff --git a/jitexpr/src/compile/mod.rs b/jitexpr/src/compile/mod.rs index 7a2051e4f..65b83ec70 100644 --- a/jitexpr/src/compile/mod.rs +++ b/jitexpr/src/compile/mod.rs @@ -1,3 +1,4 @@ +mod cache; mod compile_fn_builder; mod compiled_fn; mod error; @@ -8,6 +9,7 @@ mod typed_expr_serialize; use std::collections::HashMap; use std::sync::Arc; +pub use cache::ExprCompilationCache; pub(crate) use compile_fn_builder::CompileFnBuilder; pub use compiled_fn::{CompiledFn, CompiledFnCtx}; use cranelift::codegen::ir::{ @@ -170,7 +172,7 @@ fn lower_literal( builder .ins() .f64const(cranelift::codegen::ir::immediates::Ieee64::with_bits( - value.to_bits(), + value.get().to_bits(), )) } TypedLiteral::String(value) => { diff --git a/jitexpr/src/compile/typed_expr.rs b/jitexpr/src/compile/typed_expr.rs index e0507040b..3b81d0ffb 100644 --- a/jitexpr/src/compile/typed_expr.rs +++ b/jitexpr/src/compile/typed_expr.rs @@ -3,7 +3,7 @@ use std::sync::Arc; #[cfg(test)] use crate::ast::Literal; use crate::functions::FnCallEnum; -use crate::types::VarType; +use crate::types::{SafeF64, VarType}; #[derive(Clone, PartialEq)] pub struct TypedVariable { @@ -34,29 +34,27 @@ impl TypedExpr { TypedExprAst::Literal(TypedLiteral::I64(value as i64)) } (TypedExprAst::Literal(TypedLiteral::U64(value)), VarType::F64) => { - TypedExprAst::Literal(TypedLiteral::F64(value as f64)) + TypedExprAst::Literal(TypedLiteral::F64(SafeF64::from_integer(value))) } (TypedExprAst::Literal(TypedLiteral::I64(value)), VarType::U64) if value >= 0 => { TypedExprAst::Literal(TypedLiteral::U64(value as u64)) } (TypedExprAst::Literal(TypedLiteral::I64(value)), VarType::F64) => { - TypedExprAst::Literal(TypedLiteral::F64(value as f64)) + TypedExprAst::Literal(TypedLiteral::F64(SafeF64::from_integer(value))) } (TypedExprAst::Literal(TypedLiteral::F64(value)), VarType::U64) - if value.is_finite() - && value.fract() == 0.0 - && value >= 0.0 - && value < u64::MAX as f64 => + if value.get().fract() == 0.0 + && value.get() >= 0.0 + && value.get() < u64::MAX as f64 => { - TypedExprAst::Literal(TypedLiteral::U64(value as u64)) + TypedExprAst::Literal(TypedLiteral::U64(value.get() as u64)) } (TypedExprAst::Literal(TypedLiteral::F64(value)), VarType::I64) - if value.is_finite() - && value.fract() == 0.0 - && value >= i64::MIN as f64 - && value < -(i64::MIN as f64) => + if value.get().fract() == 0.0 + && value.get() >= i64::MIN as f64 + && value.get() < -(i64::MIN as f64) => { - TypedExprAst::Literal(TypedLiteral::I64(value as i64)) + TypedExprAst::Literal(TypedLiteral::I64(value.get() as i64)) } (ast, target_type) => TypedExprAst::Coerce { target_type, @@ -99,7 +97,7 @@ pub(crate) enum TypedLiteral { Bool(bool), U64(u64), I64(i64), - F64(f64), + F64(SafeF64), String(Arc), } diff --git a/jitexpr/src/compile/typed_expr_serialize.rs b/jitexpr/src/compile/typed_expr_serialize.rs index dbc6b7571..65a431c77 100644 --- a/jitexpr/src/compile/typed_expr_serialize.rs +++ b/jitexpr/src/compile/typed_expr_serialize.rs @@ -46,7 +46,7 @@ fn format_literal(literal: &TypedLiteral, formatter: &mut fmt::Formatter) -> fmt TypedLiteral::Bool(value) => write!(formatter, "{value}"), TypedLiteral::U64(value) => write!(formatter, "{value}u64"), TypedLiteral::I64(value) => write!(formatter, "{value}i64"), - TypedLiteral::F64(value) => write!(formatter, "{value}f64"), + TypedLiteral::F64(value) => write!(formatter, "{}f64", value.get()), TypedLiteral::String(value) => format_string_literal(value, formatter), } } diff --git a/jitexpr/src/functions/add.rs b/jitexpr/src/functions/add.rs index 6182ae397..740af0265 100644 --- a/jitexpr/src/functions/add.rs +++ b/jitexpr/src/functions/add.rs @@ -181,7 +181,10 @@ mod tests { fn test_infer_types_rejects_string_argument() { let expr = UntypedExpr::new_fn_call( Function::Add, - vec![UntypedExpr::literal(1.0), UntypedExpr::literal("hello")], + vec![ + UntypedExpr::literal(Literal::try_from(1.0).unwrap()), + UntypedExpr::literal("hello"), + ], ) .unwrap(); let error = infer_types(&expr).unwrap_err(); @@ -387,7 +390,7 @@ mod tests { vec![ UntypedExpr::variable("myfield"), UntypedExpr::literal(-2i64), - UntypedExpr::literal(0.5f64), + UntypedExpr::literal(Literal::try_from(0.5f64).unwrap()), ], ) .unwrap(); @@ -428,7 +431,10 @@ mod tests { fn test_compile_u64_to_float_coercion_is_unsigned() { let expression = UntypedExpr::new_fn_call( Function::Add, - vec![UntypedExpr::variable("x"), UntypedExpr::literal(0.5f64)], + vec![ + UntypedExpr::variable("x"), + UntypedExpr::literal(Literal::try_from(0.5f64).unwrap()), + ], ) .unwrap(); let variable_types = HashMap::from([("x", VarType::U64)]); @@ -467,7 +473,10 @@ mod tests { fn test_compile_can_coerce_variable_when_necessary() { let expression = UntypedExpr::new_fn_call( Function::Add, - vec![UntypedExpr::variable("x"), UntypedExpr::literal(1.2f64)], + vec![ + UntypedExpr::variable("x"), + UntypedExpr::literal(Literal::try_from(1.2f64).unwrap()), + ], ) .unwrap(); let variable_types = HashMap::from([("x", VarType::U64)]); @@ -497,7 +506,10 @@ mod tests { #[test] fn test_no_variable_works() { - let args = vec![UntypedExpr::literal(1.2f64), UntypedExpr::literal(1u64)]; + let args = vec![ + UntypedExpr::literal(Literal::try_from(1.2f64).unwrap()), + UntypedExpr::literal(1u64), + ]; let variable_types = HashMap::new(); let typed_expr = crate::typed_expr_from_str("(ADD 1.2f64 1u64)", &variable_types); assert_eq!(typed_expr.return_type, VarType::F64); diff --git a/jitexpr/src/functions/is_null.rs b/jitexpr/src/functions/is_null.rs index b7bf84d68..3ef9645fa 100644 --- a/jitexpr/src/functions/is_null.rs +++ b/jitexpr/src/functions/is_null.rs @@ -98,7 +98,7 @@ impl From for FnCallEnum { #[cfg(test)] mod tests { use super::*; - use crate::ast::{InvalidFnCall, deserialize, infer_types}; + use crate::ast::{InvalidFnCall, Literal, deserialize, infer_types}; use crate::compile::compile; use crate::functions::ArgumentCount; use crate::types::VariableValue; @@ -164,12 +164,25 @@ mod tests { #[test] fn test_non_finite_values_are_present() { - // The textual parser rejects non-finite literals, but programmatically - // constructed expressions can still contain them. They are not null. - for value in [f64::NAN, f64::INFINITY, f64::NEG_INFINITY] { - let expression = - UntypedExpr::new_fn_call(Function::IsNull, vec![UntypedExpr::literal(value)]) - .unwrap(); + // `SafeF64` bars non-finite *literals*, but arithmetic still reaches + // non-finite *values* at runtime. Those are present, hence not null. + let float_literal = |value: f64| UntypedExpr::literal(Literal::try_from(value).unwrap()); + let multiply = |left: UntypedExpr, right: UntypedExpr| { + UntypedExpr::new_fn_call(Function::Multiply, vec![left, right]).unwrap() + }; + + // `f64::MAX * f64::MAX` overflows to an infinity, and subtracting two + // like-signed infinities yields NaN. + let positive_infinity = multiply(float_literal(f64::MAX), float_literal(f64::MAX)); + let negative_infinity = multiply(float_literal(f64::MIN), float_literal(f64::MAX)); + let not_a_number = UntypedExpr::new_fn_call( + Function::Subtract, + vec![positive_infinity.clone(), positive_infinity.clone()], + ) + .unwrap(); + + for non_finite in [positive_infinity, negative_infinity, not_a_number] { + let expression = UntypedExpr::new_fn_call(Function::IsNull, vec![non_finite]).unwrap(); let mut compiled = compile(&expression, &HashMap::new()).unwrap().context(); // SAFETY: The expression has no runtime inputs and returns a boolean. assert_eq!(unsafe { compiled.call(&[]).as_bool() }, Some(false)); diff --git a/jitexpr/src/functions/left.rs b/jitexpr/src/functions/left.rs index f194314b6..d0ab71078 100644 --- a/jitexpr/src/functions/left.rs +++ b/jitexpr/src/functions/left.rs @@ -38,7 +38,7 @@ fn constant_length(expression: &UntypedExpr) -> Result, super::Inv Ok(match literal { Literal::I64(value) => usize::try_from(*value).ok(), Literal::U64(value) => usize::try_from(*value).ok(), - Literal::F64(value) => usize::try_from(*value as i64).ok(), + Literal::F64(value) => usize::try_from(value.get() as i64).ok(), Literal::None => None, Literal::Bool(_) | Literal::String(_) => None, }) diff --git a/jitexpr/src/functions/mod.rs b/jitexpr/src/functions/mod.rs index 734d46a99..1c49825b0 100644 --- a/jitexpr/src/functions/mod.rs +++ b/jitexpr/src/functions/mod.rs @@ -599,7 +599,7 @@ impl FnCallEnum { } /// Error representing an invalid function call. -#[derive(Debug, Eq, PartialEq, thiserror::Error)] +#[derive(Debug, Eq, PartialEq, thiserror::Error, Clone)] pub enum InvalidFnCall { #[error("invalid number of arguments: expected {expected}, got {provided}")] InvalidNumberOfArguments { diff --git a/jitexpr/src/functions/right.rs b/jitexpr/src/functions/right.rs index 1c4efb2fd..e2c638f9d 100644 --- a/jitexpr/src/functions/right.rs +++ b/jitexpr/src/functions/right.rs @@ -39,7 +39,7 @@ fn constant_length(expression: &UntypedExpr) -> Result, super::Inv Ok(match literal { Literal::I64(value) => usize::try_from(*value).ok(), Literal::U64(value) => usize::try_from(*value).ok(), - Literal::F64(value) => usize::try_from(*value as i64).ok(), + Literal::F64(value) => usize::try_from(value.get() as i64).ok(), Literal::None => None, Literal::Bool(_) | Literal::String(_) => None, }) diff --git a/jitexpr/src/functions/round.rs b/jitexpr/src/functions/round.rs index d526a2dee..28e0c1071 100644 --- a/jitexpr/src/functions/round.rs +++ b/jitexpr/src/functions/round.rs @@ -57,12 +57,11 @@ fn constant_precision( Literal::I64(value) => Some(*value), Literal::U64(value) => i64::try_from(*value).ok(), Literal::F64(value) - if value.is_finite() - && value.fract() == 0.0 - && *value >= i64::MIN as f64 - && *value < -(i64::MIN as f64) => + if value.get().fract() == 0.0 + && value.get() >= i64::MIN as f64 + && value.get() < -(i64::MIN as f64) => { - Some(*value as i64) + Some(value.get() as i64) } Literal::None => None, Literal::F64(_) | Literal::Bool(_) | Literal::String(_) => None, diff --git a/jitexpr/src/functions/split_after.rs b/jitexpr/src/functions/split_after.rs index 31705ecc9..5f0b3c0a9 100644 --- a/jitexpr/src/functions/split_after.rs +++ b/jitexpr/src/functions/split_after.rs @@ -52,7 +52,7 @@ fn constant_occurrence( Ok(match literal { Literal::I64(value) => usize::try_from(*value).ok(), Literal::U64(value) => usize::try_from(*value).ok(), - Literal::F64(value) => usize::try_from(*value as i64).ok(), + Literal::F64(value) => usize::try_from(value.get() as i64).ok(), Literal::None => None, Literal::Bool(_) | Literal::String(_) => None, }) diff --git a/jitexpr/src/functions/split_before.rs b/jitexpr/src/functions/split_before.rs index 67415f6fd..4e75aadbc 100644 --- a/jitexpr/src/functions/split_before.rs +++ b/jitexpr/src/functions/split_before.rs @@ -51,7 +51,7 @@ fn constant_occurrence( Ok(match literal { Literal::I64(value) => usize::try_from(*value).ok(), Literal::U64(value) => usize::try_from(*value).ok(), - Literal::F64(value) => usize::try_from(*value as i64).ok(), + Literal::F64(value) => usize::try_from(value.get() as i64).ok(), Literal::None => None, Literal::Bool(_) | Literal::String(_) => None, }) diff --git a/jitexpr/src/functions/substring.rs b/jitexpr/src/functions/substring.rs index 5d2e2e92d..e15a6d18e 100644 --- a/jitexpr/src/functions/substring.rs +++ b/jitexpr/src/functions/substring.rs @@ -53,7 +53,7 @@ fn constant_usize( .ok() .and_then(|value| usize::try_from(value).ok()), Literal::F64(value) if literal.types().contains(VarType::I64) => { - usize::try_from(*value as i64).ok() + usize::try_from(value.get() as i64).ok() } Literal::None => None, Literal::Bool(_) | Literal::F64(_) | Literal::String(_) => None, diff --git a/jitexpr/src/types.rs b/jitexpr/src/types.rs index 18a8cb957..2b534bd97 100644 --- a/jitexpr/src/types.rs +++ b/jitexpr/src/types.rs @@ -1,5 +1,9 @@ //! Source types and nullable runtime value representations. +use std::cmp::Ordering; +use std::fmt; +use std::hash::{Hash, Hasher}; + /// A value type supported by compiled expressions. #[derive(Clone, Copy, Debug, Eq, Hash, PartialEq, Ord, PartialOrd)] pub enum VarType { @@ -11,6 +15,83 @@ pub enum VarType { None, } +/// Wraps a f64 that is not inf nor Nan. +/// +/// Excluding NaN is what makes the `Eq` and `Ord` impls below sound: every +/// remaining value compares equal to itself, and no comparison is undefined. +/// +/// Equality and ordering are *bitwise*, via [`f64::total_cmp`], so `-0.0` and +/// `0.0` are distinct and `-0.0` sorts first, even though IEEE equality calls +/// them equal. Callers that key generated code on a literal need that +/// distinction: the two lower to different machine constants, and the sign of +/// a zero is observable in a result. +#[derive(Copy, Clone)] +pub struct SafeF64(f64); + +impl SafeF64 { + /// Returns `None` for NaN and for either infinity. + pub fn new(val: f64) -> Option { + if val.is_nan() || val.is_infinite() { + None + } else { + Some(SafeF64(val)) + } + } + + /// Converts an integer to the nearest `f64`. + /// + /// Total, hence infallible: every `i64` and `u64` magnitude is far inside + /// `f64`'s finite range, so no non-finite value can come out. The + /// conversion may still round, exactly as `as f64` would. + pub fn from_integer(value: impl Into) -> SafeF64 { + SafeF64(value.into() as f64) + } + + /// Returns the wrapped value, which is guaranteed finite. + pub fn get(self) -> f64 { + self.0 + } +} + +/// Forwards to the wrapped `f64`, so a `SafeF64` is indistinguishable from the +/// number it holds. +impl fmt::Debug for SafeF64 { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + fmt::Debug::fmt(&self.0, formatter) + } +} + +impl PartialEq for SafeF64 { + fn eq(&self, other: &SafeF64) -> bool { + self.0.to_bits() == other.0.to_bits() + } +} + +// Sound because NaN is excluded, so equality is reflexive. +impl Eq for SafeF64 {} + +impl Ord for SafeF64 { + fn cmp(&self, other: &SafeF64) -> Ordering { + // `total_cmp` returns `Equal` exactly when the bit patterns match, so + // this order is consistent with `PartialEq` above. + self.0.total_cmp(&other.0) + } +} + +impl PartialOrd for SafeF64 { + fn partial_cmp(&self, other: &SafeF64) -> Option { + Some(self.cmp(other)) + } +} + +/// Hashes the bit pattern, which is what [`PartialEq`] compares. Hashing the +/// value any other way would break the `Hash`/`Eq` agreement for signed zero. +impl Hash for SafeF64 { + fn hash(&self, hasher: &mut H) { + self.0.to_bits().hash(hasher); + } +} + /// The payload of a primitive runtime value. /// /// This union is deliberately untagged. The corresponding @@ -291,7 +372,117 @@ impl<'a> From for VariableValue<'a> { #[cfg(test)] mod tests { - use crate::types::{VariablePrimitive, VariablePrimitiveOpt, VariableValue}; + use std::cmp::Ordering; + use std::collections::hash_map::DefaultHasher; + use std::hash::{Hash, Hasher}; + + use crate::types::{SafeF64, VariablePrimitive, VariablePrimitiveOpt, VariableValue}; + + fn safe(val: f64) -> SafeF64 { + SafeF64::new(val).unwrap() + } + + fn hash_of(val: SafeF64) -> u64 { + let mut hasher = DefaultHasher::new(); + val.hash(&mut hasher); + hasher.finish() + } + + /// Values spanning both zero signs and both extremes. + fn sample_values() -> Vec { + [ + -0.0f64, + 0.0, + 1.0, + -1.0, + 1.5, + -1.5, + f64::MIN, + f64::MAX, + f64::MIN_POSITIVE, + ] + .into_iter() + .map(safe) + .collect() + } + + #[test] + fn test_safe_f64_rejects_non_finite() { + assert!(SafeF64::new(f64::NAN).is_none()); + assert!(SafeF64::new(-f64::NAN).is_none()); + assert!(SafeF64::new(f64::INFINITY).is_none()); + assert!(SafeF64::new(f64::NEG_INFINITY).is_none()); + assert_eq!(SafeF64::new(1.5).unwrap().get(), 1.5); + assert_eq!(SafeF64::new(f64::MAX).unwrap().get(), f64::MAX); + } + + #[test] + fn test_safe_f64_eq_and_ord_laws() { + let values = sample_values(); + for left in &values { + // Reflexivity is what NaN would have broken, making `Eq` unsound. + assert_eq!(left, left); + assert_eq!(left.cmp(left), Ordering::Equal); + for right in &values { + assert_eq!(left == right, right == left, "symmetry"); + assert_eq!(left.cmp(right), right.cmp(left).reverse(), "antisymmetry"); + // `Ord` must agree with `Eq`, and `Hash` with both. + assert_eq!( + left == right, + left.cmp(right) == Ordering::Equal, + "cmp agrees with eq" + ); + if left == right { + assert_eq!(hash_of(*left), hash_of(*right), "equal values hash equally"); + } + for third in &values { + if left == right && right == third { + assert_eq!(left, third, "transitivity of eq"); + } + if left <= right && right <= third { + assert!(left <= third, "transitivity of ord"); + } + } + } + } + } + + #[test] + fn test_safe_f64_distinguishes_signed_zero() { + // IEEE equality calls these equal, but they lower to different machine + // constants, so the key-facing comparison must keep them apart. + assert_eq!(-0.0f64, 0.0f64); + assert_ne!(safe(-0.0), safe(0.0)); + assert_eq!(safe(-0.0).cmp(&safe(0.0)), Ordering::Less); + assert_ne!(hash_of(safe(-0.0)), hash_of(safe(0.0))); + } + + #[test] + fn test_safe_f64_sorts_in_numeric_order() { + let mut values = sample_values(); + values.sort(); + let sorted: Vec = values.iter().map(|value| value.get()).collect(); + assert_eq!( + sorted, + vec![ + f64::MIN, + -1.5, + -1.0, + -0.0, + 0.0, + f64::MIN_POSITIVE, + 1.0, + 1.5, + f64::MAX + ] + ); + } + + #[test] + fn test_safe_f64_debug_is_transparent() { + assert_eq!(format!("{:?}", safe(1.5)), format!("{:?}", 1.5f64)); + assert_eq!(format!("{:?}", safe(-0.0)), format!("{:?}", -0.0f64)); + } #[test] fn test_runtime_value_layouts() { diff --git a/src/index/index.rs b/src/index/index.rs index 81c631cbd..0fc6993fb 100644 --- a/src/index/index.rs +++ b/src/index/index.rs @@ -6,6 +6,9 @@ use std::path::PathBuf; use std::sync::Arc; use std::thread::available_parallelism; +#[cfg(feature = "jitexpr")] +use jitexpr::compile::ExprCompilationCache; + use super::segment::Segment; use super::segment_reader::merge_field_meta_data; use super::{FieldMetadata, IndexSettings}; @@ -382,6 +385,8 @@ pub struct Index { fast_field_tokenizers: TokenizerManager, inventory: SegmentMetaInventory, custom_plugins: Vec>, + #[cfg(feature = "jitexpr")] + expr_compilation_cache: ExprCompilationCache, } impl Index { @@ -502,6 +507,9 @@ impl Index { executor: Executor::single_thread(), inventory, custom_plugins: Vec::new(), + // We default to a capacity of 64, but it only allocates if used. + #[cfg(feature = "jitexpr")] + expr_compilation_cache: ExprCompilationCache::with_capacity(64), } } @@ -910,8 +918,21 @@ impl Index { } } +#[cfg(feature = "jitexpr")] +impl Index { + /// Setter for the expression compilation cache. + pub fn set_expr_compilation_cache(&mut self, cache: ExprCompilationCache) { + self.expr_compilation_cache = cache; + } + + /// Accessor for the expression compilation cache. + pub fn expr_compilation_cache(&self) -> &ExprCompilationCache { + &self.expr_compilation_cache + } +} + impl fmt::Debug for Index { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result { write!(f, "Index({:?})", self.directory) } } diff --git a/src/index/segment_reader.rs b/src/index/segment_reader.rs index 5b6ea0dfc..0aa0c4a96 100644 --- a/src/index/segment_reader.rs +++ b/src/index/segment_reader.rs @@ -70,6 +70,11 @@ impl SegmentReader { &self.schema } + /// Returns the index this segment belongs to. + pub fn index(&self) -> &Index { + &self.index + } + /// Return the number of documents that have been /// deleted in the segment. pub fn num_deleted_docs(&self) -> DocId { diff --git a/src/query/doc_predicate_query/jitexpr_predicate.rs b/src/query/doc_predicate_query/jitexpr_predicate.rs index 09b1c2cfa..94023ddcf 100644 --- a/src/query/doc_predicate_query/jitexpr_predicate.rs +++ b/src/query/doc_predicate_query/jitexpr_predicate.rs @@ -3,7 +3,7 @@ use std::io; use columnar::{ColumnType, DynamicColumn, StrColumn}; use jitexpr::ast::{infer_types_with_target, InferredTypeSet, TypeError, UntypedExpr}; -use jitexpr::compile::{compile, CompiledFnCtx, StringArena}; +use jitexpr::compile::{CompiledFnCtx, StringArena}; use jitexpr::types::{VarType, VariableValue}; use super::{DocPredicate, SegmentDocPredicate}; @@ -91,8 +91,11 @@ impl DocPredicate for JitExprPredicate { opened_columns.insert(name.as_str(), column); } - let compiled_fn = - compile(&self.expression, &variable_types).map_err(|compilation_err| { + let compiled_fn = segment_reader + .index() + .expr_compilation_cache() + .compile(&self.expression, &variable_types) + .map_err(|compilation_err| { TantivyError::InvalidArgument(format!( "the expression compilation failed {:?}. error: {compilation_err}", self.expression From 20d7f72f2c959f2b768e9ee88b8c0cfcdab842a4 Mon Sep 17 00:00:00 2001 From: Paul Masurel Date: Thu, 17 Sep 2026 12:29:31 +0200 Subject: [PATCH 18/49] Addressing -0.0 problem in SafeF64 (#3109) Co-authored-by: Paul Masurel --- jitexpr/src/ast/serde.rs | 10 ++----- jitexpr/src/functions/sqrt.rs | 7 +++-- jitexpr/src/types.rs | 56 ++++++++++++++++++----------------- 3 files changed, 37 insertions(+), 36 deletions(-) diff --git a/jitexpr/src/ast/serde.rs b/jitexpr/src/ast/serde.rs index 39ae800dc..959edff15 100644 --- a/jitexpr/src/ast/serde.rs +++ b/jitexpr/src/ast/serde.rs @@ -591,13 +591,9 @@ mod tests { #[test] fn test_finite_float_edge_values_round_trip() { - for value in [ - f64::MIN, - f64::MAX, - f64::MIN_POSITIVE, - f64::from_bits(1), - -0.0, - ] { + // `-0.0` is deliberately absent: `SafeF64` folds it to `0.0`, so it + // serializes as `0f64` and cannot round-trip. + for value in [f64::MIN, f64::MAX, f64::MIN_POSITIVE, f64::from_bits(1)] { let serialized = serialize(&UntypedExpr::literal(Literal::try_from(value).unwrap())); let UntypedExpr::Literal(Literal::F64(parsed)) = deserialize(&serialized).unwrap() else { diff --git a/jitexpr/src/functions/sqrt.rs b/jitexpr/src/functions/sqrt.rs index b914f48ac..cadaf986b 100644 --- a/jitexpr/src/functions/sqrt.rs +++ b/jitexpr/src/functions/sqrt.rs @@ -145,8 +145,11 @@ mod tests { assert_eq!(eval("(SQRT none)"), None); assert_eq!(eval("(SQRT -1i64)"), None); - let negative_zero = eval("(SQRT -0f64)").unwrap(); - assert_eq!(negative_zero.to_bits(), (-0.0f64).to_bits()); + // `SafeF64` folds `-0.0` on construction, so the literal carries a + // positive zero and IEEE's `sqrt(-0.0) == -0.0` is unreachable here. + // The sign still exists for runtime values, as `abs.rs` exercises. + let zero = eval("(SQRT -0f64)").unwrap(); + assert_eq!(zero.to_bits(), 0.0f64.to_bits()); } #[test] diff --git a/jitexpr/src/types.rs b/jitexpr/src/types.rs index 2b534bd97..626e968f7 100644 --- a/jitexpr/src/types.rs +++ b/jitexpr/src/types.rs @@ -15,24 +15,21 @@ pub enum VarType { None, } -/// Wraps a f64 that is not inf nor Nan. -/// -/// Excluding NaN is what makes the `Eq` and `Ord` impls below sound: every -/// remaining value compares equal to itself, and no comparison is undefined. -/// -/// Equality and ordering are *bitwise*, via [`f64::total_cmp`], so `-0.0` and -/// `0.0` are distinct and `-0.0` sorts first, even though IEEE equality calls -/// them equal. Callers that key generated code on a literal need that -/// distinction: the two lower to different machine constants, and the sign of -/// a zero is observable in a result. +/// Wraps a f64 that is not inf, nor Nan, nor neg 0. #[derive(Copy, Clone)] pub struct SafeF64(f64); impl SafeF64 { /// Returns `None` for NaN and for either infinity. + /// + /// A negative zero is accepted, and folded to `0.0`. pub fn new(val: f64) -> Option { if val.is_nan() || val.is_infinite() { - None + return None; + } + if val == 0.0 { + // True of both zeros; `0.0` is the representative. + Some(SafeF64(0.0)) } else { Some(SafeF64(val)) } @@ -63,7 +60,7 @@ impl fmt::Debug for SafeF64 { impl PartialEq for SafeF64 { fn eq(&self, other: &SafeF64) -> bool { - self.0.to_bits() == other.0.to_bits() + self.0 == other.0 } } @@ -71,21 +68,21 @@ impl PartialEq for SafeF64 { impl Eq for SafeF64 {} impl Ord for SafeF64 { + #[inline(always)] fn cmp(&self, other: &SafeF64) -> Ordering { - // `total_cmp` returns `Equal` exactly when the bit patterns match, so - // this order is consistent with `PartialEq` above. - self.0.total_cmp(&other.0) + self.partial_cmp(&other).unwrap() } } impl PartialOrd for SafeF64 { + #[inline(always)] fn partial_cmp(&self, other: &SafeF64) -> Option { - Some(self.cmp(other)) + self.0.partial_cmp(&other.0) } } -/// Hashes the bit pattern, which is what [`PartialEq`] compares. Hashing the -/// value any other way would break the `Hash`/`Eq` agreement for signed zero. +/// Hashes the bit pattern. Because we remove NaN Inf and -0.0, +/// this is consistent with equality. (x == y => hash(x) == hash(y)). impl Hash for SafeF64 { fn hash(&self, hasher: &mut H) { self.0.to_bits().hash(hasher); @@ -373,6 +370,7 @@ 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}; @@ -388,7 +386,7 @@ mod tests { hasher.finish() } - /// Values spanning both zero signs and both extremes. + /// Values spanning both zeros -- which fold together -- and both extremes. fn sample_values() -> Vec { [ -0.0f64, @@ -448,13 +446,18 @@ mod tests { } #[test] - fn test_safe_f64_distinguishes_signed_zero() { - // IEEE equality calls these equal, but they lower to different machine - // constants, so the key-facing comparison must keep them apart. - assert_eq!(-0.0f64, 0.0f64); - assert_ne!(safe(-0.0), safe(0.0)); - assert_eq!(safe(-0.0).cmp(&safe(0.0)), Ordering::Less); - assert_ne!(hash_of(safe(-0.0)), hash_of(safe(0.0))); + fn test_safe_f64_folds_negative_zero() { + // The fold happens on the way in, so there is a single zero to compare. + assert_eq!(safe(-0.0).get().to_bits(), 0.0f64.to_bits()); + assert!(!safe(-0.0).get().is_sign_negative()); + assert_eq!( + SafeF64::from_integer(0i64).get().to_bits(), + 0.0f64.to_bits() + ); + + assert_eq!(safe(-0.0), safe(0.0)); + assert_eq!(safe(-0.0).cmp(&safe(0.0)), Ordering::Equal); + assert_eq!(hash_of(safe(-0.0)), hash_of(safe(0.0))); } #[test] @@ -481,7 +484,6 @@ mod tests { #[test] fn test_safe_f64_debug_is_transparent() { assert_eq!(format!("{:?}", safe(1.5)), format!("{:?}", 1.5f64)); - assert_eq!(format!("{:?}", safe(-0.0)), format!("{:?}", -0.0f64)); } #[test] From b91a405ee973f86a3a8b9bf139f53db21b94cf9c Mon Sep 17 00:00:00 2001 From: Paul Masurel Date: Fri, 18 Sep 2026 07:27:01 +0200 Subject: [PATCH 19/49] Make aggregation request-data types crate-internal (#3110) Hiding the per segment *AggReqData members FilterAggReqData::reqandMultiTermsAggReqData::sub_aggregations` were dead code so they were removed. --- src/aggregation/agg_data.rs | 30 +++++++++---------- src/aggregation/bucket/composite/accessors.rs | 8 ++--- src/aggregation/bucket/composite/mod.rs | 3 +- src/aggregation/bucket/filter.rs | 12 ++++---- src/aggregation/bucket/histogram/histogram.rs | 16 +++++----- src/aggregation/bucket/multi_terms/mod.rs | 14 ++++----- src/aggregation/bucket/range.rs | 12 ++++---- src/aggregation/bucket/term_agg/mod.rs | 22 +++++++------- src/aggregation/bucket/term_missing_agg.rs | 8 ++--- src/aggregation/metric/cardinality.rs | 14 ++++----- src/aggregation/metric/mod.rs | 16 +++++----- src/aggregation/metric/percentiles.rs | 2 +- src/aggregation/metric/top_hits.rs | 12 ++++---- 13 files changed, 82 insertions(+), 87 deletions(-) diff --git a/src/aggregation/agg_data.rs b/src/aggregation/agg_data.rs index 475cef2ea..bcb0acedd 100644 --- a/src/aggregation/agg_data.rs +++ b/src/aggregation/agg_data.rs @@ -37,8 +37,8 @@ use crate::{SegmentOrdinal, SegmentReader}; /// It is passed to the collectors during collection. pub struct AggregationsSegmentCtx { /// Request data for each aggregation type. - pub per_request: PerRequestAggSegCtx, - pub context: AggContextParams, + pub(crate) per_request: PerRequestAggSegCtx, + pub(crate) context: AggContextParams, pub(crate) column_block_accessor: ColumnBlockAccessor, } @@ -115,28 +115,28 @@ impl AggregationsSegmentCtx { #[derive(Default)] pub struct PerRequestAggSegCtx { /// TermsAggReqData contains the request data for a terms aggregation. - pub term_req_data: Vec, + pub(crate) term_req_data: Vec, /// HistogramAggReqData contains the request data for a histogram aggregation. - pub histogram_req_data: Vec, + pub(crate) histogram_req_data: Vec, /// RangeAggReqData contains the request data for a range aggregation. - pub range_req_data: Vec, + pub(crate) range_req_data: Vec, /// FilterAggReqData contains the request data for a filter aggregation. - pub filter_req_data: Vec, + pub(crate) filter_req_data: Vec, /// Shared by avg, min, max, sum, stats, extended_stats, count - pub stats_metric_req_data: Vec, + pub(crate) stats_metric_req_data: Vec, /// CardinalityAggReqData contains the request data for a cardinality aggregation. - pub cardinality_req_data: Vec, + pub(crate) cardinality_req_data: Vec, /// TopHitsAggReqData contains the request data for a top_hits aggregation. - pub top_hits_req_data: Vec, + pub(crate) top_hits_req_data: Vec, /// MissingTermAggReqData contains the request data for a missing term aggregation. - pub missing_term_req_data: Vec, + pub(crate) missing_term_req_data: Vec, /// CompositeAggReqData contains the request data for a composite aggregation. - pub composite_req_data: Vec, + pub(crate) composite_req_data: Vec, /// MultiTermsAggReqData contains the request data for a multi_terms aggregation. - pub multi_terms_req_data: Vec, + pub(crate) multi_terms_req_data: Vec, /// Request tree used to build collectors. - pub agg_tree: Vec, + pub(crate) agg_tree: Vec, } impl PerRequestAggSegCtx { @@ -690,7 +690,6 @@ fn build_nodes( let idx_in_req_data = data.push_filter_req_data(FilterAggReqData { name: agg_name.to_string(), - req: filter_req.clone(), segment_reader: reader.clone(), evaluator, is_top_level, @@ -827,7 +826,6 @@ fn build_multi_terms_nodes( req: req.clone(), fields, missing_accessors, - sub_aggregations: sub_aggs.clone(), is_top_level, }); let children = build_children(sub_aggs, reader, segment_ordinal, data)?; @@ -1125,7 +1123,7 @@ fn build_terms_or_cardinality_nodes( missing_value_for_accessor, name: agg_name.to_string(), req: TermsAggregationInternal::from_req(req), - sug_aggregations: sub_aggs.clone(), + sub_aggregations: sub_aggs.clone(), allowed_term_ids, is_top_level, }); diff --git a/src/aggregation/bucket/composite/accessors.rs b/src/aggregation/bucket/composite/accessors.rs index 005700bed..c21e3a065 100644 --- a/src/aggregation/bucket/composite/accessors.rs +++ b/src/aggregation/bucket/composite/accessors.rs @@ -17,13 +17,13 @@ use crate::{SegmentReader, TantivyError}; /// Contains all information required by the SegmentCompositeCollector to perform the /// composite aggregation on a segment. #[derive(Debug, Clone)] -pub struct CompositeAggReqData { +pub(crate) struct CompositeAggReqData { /// The name of the aggregation. - pub name: String, + pub(crate) name: String, /// The normalized term aggregation request. - pub req: CompositeAggregation, + pub(crate) req: CompositeAggregation, /// Accessors for each source, each source can have multiple accessors (columns). - pub composite_accessors: Vec, + pub(crate) composite_accessors: Vec, } impl CompositeAggReqData { diff --git a/src/aggregation/bucket/composite/mod.rs b/src/aggregation/bucket/composite/mod.rs index 1ea39f8a3..b7e5dd70a 100644 --- a/src/aggregation/bucket/composite/mod.rs +++ b/src/aggregation/bucket/composite/mod.rs @@ -15,8 +15,9 @@ use rustc_hash::FxHashMap; use serde::{Deserialize, Serialize}; use crate::aggregation::agg_result::CompositeKey; +pub(crate) use crate::aggregation::bucket::composite::accessors::CompositeAggReqData; pub use crate::aggregation::bucket::composite::accessors::{ - CompositeAccessor, CompositeAggReqData, CompositeSourceAccessors, PrecomputedDateInterval, + CompositeAccessor, CompositeSourceAccessors, PrecomputedDateInterval, }; pub use crate::aggregation::bucket::composite::collector::SegmentCompositeCollector; use crate::aggregation::bucket::composite::numeric_types::num_cmp::{ diff --git a/src/aggregation/bucket/filter.rs b/src/aggregation/bucket/filter.rs index 9e39edd41..0c08d08c7 100644 --- a/src/aggregation/bucket/filter.rs +++ b/src/aggregation/bucket/filter.rs @@ -398,19 +398,17 @@ impl PartialEq for FilterAggregation { /// Request data for filter aggregation /// This struct holds the per-segment data needed to execute a filter aggregation #[derive(Clone)] -pub struct FilterAggReqData { +pub(crate) struct FilterAggReqData { /// The name of the filter aggregation - pub name: String, - /// The filter aggregation - pub req: FilterAggregation, + pub(crate) name: String, /// The segment reader - pub segment_reader: SegmentReader, + pub(crate) segment_reader: SegmentReader, /// Document evaluator for the filter query (precomputed BitSet). /// Wrapped in `Rc` so cloning the request data does not duplicate the (potentially large) /// underlying BitSet. - pub evaluator: Rc, + pub(crate) evaluator: Rc, /// True if this filter aggregation is at the top level of the aggregation tree (not nested). - pub is_top_level: bool, + pub(crate) is_top_level: bool, } impl FilterAggReqData { diff --git a/src/aggregation/bucket/histogram/histogram.rs b/src/aggregation/bucket/histogram/histogram.rs index e97a9ca5c..196713730 100644 --- a/src/aggregation/bucket/histogram/histogram.rs +++ b/src/aggregation/bucket/histogram/histogram.rs @@ -22,21 +22,21 @@ use crate::TantivyError; /// Contains all information required by the SegmentHistogramCollector to perform the /// histogram or date_histogram aggregation on a segment. #[derive(Debug, Clone)] -pub struct HistogramAggReqData { +pub(crate) struct HistogramAggReqData { /// The column accessor to access the fast field values. - pub accessor: Column, + pub(crate) accessor: Column, /// The field type of the fast field. - pub field_type: ColumnType, + pub(crate) field_type: ColumnType, /// The name of the aggregation. - pub name: String, + pub(crate) name: String, /// The histogram aggregation request. - pub req: HistogramAggregation, + pub(crate) req: HistogramAggregation, /// True if this is a date_histogram aggregation. - pub is_date_histogram: bool, + pub(crate) is_date_histogram: bool, /// The bounds to limit the buckets to. - pub bounds: HistogramBounds, + pub(crate) bounds: HistogramBounds, /// The offset used to calculate the bucket position. - pub offset: f64, + pub(crate) offset: f64, } impl HistogramAggReqData { /// Estimate the memory consumption of this struct in bytes. diff --git a/src/aggregation/bucket/multi_terms/mod.rs b/src/aggregation/bucket/multi_terms/mod.rs index 5b3af5f08..2cdfbb1f1 100644 --- a/src/aggregation/bucket/multi_terms/mod.rs +++ b/src/aggregation/bucket/multi_terms/mod.rs @@ -158,23 +158,21 @@ impl MultiTermsFieldAccessor { /// Per-request data bundle passed to the segment collector. #[derive(Debug, Clone)] -pub struct MultiTermsAggReqData { +pub(crate) struct MultiTermsAggReqData { /// Aggregation name used to look up this entry in the result tree. - pub name: String, + pub(crate) name: String, /// Original request (needed for final-result conversion). - pub req: MultiTermsAggregation, + pub(crate) req: MultiTermsAggregation, /// One typed accessor per field listed in `req.terms`. - pub fields: Vec, + pub(crate) fields: Vec, /// Missing-value handling corresponding to `fields`. Only the designated physical accessor /// choice for each requested field carries `Some`, preventing duplicate missing buckets when /// type-specific collectors are merged. - pub missing_accessors: Vec>, - /// Sub-aggregation descriptor (empty when no sub-aggs). - pub sub_aggregations: Aggregations, + pub(crate) missing_accessors: Vec>, /// True if this multi_terms aggregation is at the top level of the aggregation tree /// (not nested). Used to gate the Vec/Paged packed-key storage tiers, which assume a /// bounded number of parent buckets (mirrors [`TermsAggReqData::is_top_level`]). - pub is_top_level: bool, + pub(crate) is_top_level: bool, } impl MultiTermsAggReqData { diff --git a/src/aggregation/bucket/range.rs b/src/aggregation/bucket/range.rs index a66ea8225..5857d5a42 100644 --- a/src/aggregation/bucket/range.rs +++ b/src/aggregation/bucket/range.rs @@ -24,17 +24,17 @@ use crate::TantivyError; /// Contains all information required by the SegmentRangeCollector to perform the /// range aggregation on a segment. #[derive(Debug, Clone)] -pub struct RangeAggReqData { +pub(crate) struct RangeAggReqData { /// The column accessor to access the fast field values. - pub accessor: Column, + pub(crate) accessor: Column, /// The type of the fast field. - pub field_type: ColumnType, + pub(crate) field_type: ColumnType, /// The range aggregation request. - pub req: RangeAggregation, + pub(crate) req: RangeAggregation, /// The name of the aggregation. - pub name: String, + pub(crate) name: String, /// Whether this is a top-level aggregation. - pub is_top_level: bool, + pub(crate) is_top_level: bool, } impl RangeAggReqData { diff --git a/src/aggregation/bucket/term_agg/mod.rs b/src/aggregation/bucket/term_agg/mod.rs index 08ec2f093..694811866 100644 --- a/src/aggregation/bucket/term_agg/mod.rs +++ b/src/aggregation/bucket/term_agg/mod.rs @@ -35,25 +35,25 @@ mod flattened_term_histogram; /// Contains all information required by the SegmentTermCollector to perform the /// terms aggregation on a segment. #[derive(Debug, Clone)] -pub struct TermsAggReqData { +pub(crate) struct TermsAggReqData { /// The column accessor to access the fast field values. - pub accessor: Column, + pub(crate) accessor: Column, /// The type of the column. - pub column_type: ColumnType, + pub(crate) column_type: ColumnType, /// The string dictionary column if the field is of type text. - pub str_dict_column: Option, + pub(crate) str_dict_column: Option, /// The missing value as u64 value. - pub missing_value_for_accessor: Option, + pub(crate) missing_value_for_accessor: Option, /// Used to build the correct nested result when we have an empty result. - pub sug_aggregations: Aggregations, + pub(crate) sub_aggregations: Aggregations, /// The name of the aggregation. - pub name: String, + pub(crate) name: String, /// The normalized term aggregation request. - pub req: TermsAggregationInternal, + pub(crate) req: TermsAggregationInternal, /// Preloaded allowed term ords (string columns only). If set, only ords present are collected. - pub allowed_term_ids: Option, + pub(crate) allowed_term_ids: Option, /// True if this terms aggregation is at the top level of the aggregation tree (not nested). - pub is_top_level: bool, + pub(crate) is_top_level: bool, } impl TermsAggReqData { @@ -1390,7 +1390,7 @@ where // TODO: Handle rev streaming for descending sorting by keys let mut stream = term_dict.stream()?; let empty_sub_aggregation = - IntermediateAggregationResults::empty_from_req(&term_req.sug_aggregations); + IntermediateAggregationResults::empty_from_req(&term_req.sub_aggregations); while stream.advance() { if dict.len() >= term_req.req.segment_size as usize { break; diff --git a/src/aggregation/bucket/term_missing_agg.rs b/src/aggregation/bucket/term_missing_agg.rs index 72dacfd07..1da26d856 100644 --- a/src/aggregation/bucket/term_missing_agg.rs +++ b/src/aggregation/bucket/term_missing_agg.rs @@ -20,13 +20,13 @@ use crate::aggregation::BucketId; /// - The field is not text and missing is provided as string (we cannot use the numeric missing /// value optimization) #[derive(Default)] -pub struct MissingTermAggReqData { +pub(crate) struct MissingTermAggReqData { /// The accessors to check for existence of a value. - pub accessors: Vec<(Column, ColumnType)>, + pub(crate) accessors: Vec<(Column, ColumnType)>, /// The name of the aggregation. - pub name: String, + pub(crate) name: String, /// The original terms aggregation request. - pub req: TermsAggregation, + pub(crate) req: TermsAggregation, } impl MissingTermAggReqData { diff --git a/src/aggregation/metric/cardinality.rs b/src/aggregation/metric/cardinality.rs index ca0a00fb2..d605125e2 100644 --- a/src/aggregation/metric/cardinality.rs +++ b/src/aggregation/metric/cardinality.rs @@ -92,19 +92,19 @@ pub struct CardinalityAggregationReq { /// Contains all information required by the SegmentCardinalityCollector to perform the /// cardinality aggregation on a segment. -pub struct CardinalityAggReqData { +pub(crate) struct CardinalityAggReqData { /// The column accessor to access the fast field values. - pub accessor: Column, + pub(crate) accessor: Column, /// The column_type of the field. - pub column_type: ColumnType, + pub(crate) column_type: ColumnType, /// The string dictionary column if the field is of type string. - pub str_dict_column: Option, + pub(crate) str_dict_column: Option, /// The missing value normalized to the internal u64 representation of the field type. - pub missing_value_for_accessor: Option, + pub(crate) missing_value_for_accessor: Option, /// The name of the aggregation. - pub name: String, + pub(crate) name: String, /// The aggregation request. - pub req: CardinalityAggregationReq, + pub(crate) req: CardinalityAggregationReq, } impl CardinalityAggReqData { diff --git a/src/aggregation/metric/mod.rs b/src/aggregation/metric/mod.rs index 05ff861f2..56ca63f01 100644 --- a/src/aggregation/metric/mod.rs +++ b/src/aggregation/metric/mod.rs @@ -48,21 +48,21 @@ use crate::schema::OwnedValue; /// Contains all information required by metric aggregations like avg, min, max, sum, stats, /// extended_stats, count, percentiles. #[repr(C)] -pub struct MetricAggReqData { +pub(crate) struct MetricAggReqData { /// True if the field is of number or date type. - pub is_number_or_date_type: bool, + pub(crate) is_number_or_date_type: bool, /// The type of the field. - pub field_type: ColumnType, + pub(crate) field_type: ColumnType, /// The missing value normalized to the internal u64 representation of the field type. - pub missing_u64: Option, + pub(crate) missing_u64: Option, /// The column accessor to access the fast field values. - pub accessor: Column, + pub(crate) accessor: Column, /// Used when converting to intermediate result - pub collecting_for: StatsType, + pub(crate) collecting_for: StatsType, /// The missing value - pub missing: Option, + pub(crate) missing: Option, /// The name of the aggregation. - pub name: String, + pub(crate) name: String, } impl MetricAggReqData { diff --git a/src/aggregation/metric/percentiles.rs b/src/aggregation/metric/percentiles.rs index 9b9c878a6..4d213c927 100644 --- a/src/aggregation/metric/percentiles.rs +++ b/src/aggregation/metric/percentiles.rs @@ -139,7 +139,7 @@ pub(crate) struct SegmentPercentilesCollector { /// The missing value normalized to the internal u64 representation of the field type. pub missing_u64: Option, /// The column accessor to access the fast field values. - pub accessor: Column, + pub(crate) accessor: Column, } #[derive(Clone, Serialize, Deserialize)] diff --git a/src/aggregation/metric/top_hits.rs b/src/aggregation/metric/top_hits.rs index 77e2856a4..4cf1cf4eb 100644 --- a/src/aggregation/metric/top_hits.rs +++ b/src/aggregation/metric/top_hits.rs @@ -24,17 +24,17 @@ use crate::{DocAddress, DocId, SegmentOrdinal}; /// Contains all information required by the TopHitsSegmentCollector to perform the /// top_hits aggregation on a segment. #[derive(Default)] -pub struct TopHitsAggReqData { +pub(crate) struct TopHitsAggReqData { /// The accessors to access the fast field values. - pub accessors: Vec<(Column, ColumnType)>, + pub(crate) accessors: Vec<(Column, ColumnType)>, /// The accessors to access the fast field values for retrieving document fields. - pub value_accessors: HashMap>, + pub(crate) value_accessors: HashMap>, /// The ordinal of the segment this request data is for. - pub segment_ordinal: SegmentOrdinal, + pub(crate) segment_ordinal: SegmentOrdinal, /// The name of the aggregation. - pub name: String, + pub(crate) name: String, /// The top_hits aggregation request. - pub req: TopHitsAggregationReq, + pub(crate) req: TopHitsAggregationReq, } impl TopHitsAggReqData { From 507e1760aada7aa9cee52d5920ea1588a2acdec4 Mon Sep 17 00:00:00 2001 From: Paul Masurel Date: Tue, 22 Sep 2026 10:47:49 +0200 Subject: [PATCH 20/49] Minor refactor of the cardinality aggregation. (#3120) The Str path and the non string path are very different. This PR isolates them as different segment aggregation collector. Co-authored-by: Paul Masurel --- src/aggregation/agg_data.rs | 39 +- src/aggregation/metric/cardinality.rs | 1407 ----------------- src/aggregation/metric/cardinality/mod.rs | 609 +++++++ .../metric/cardinality/numeric_collector.rs | 255 +++ .../metric/cardinality/str_collector.rs | 397 +++++ .../cardinality/term_ord_accumulator.rs | 343 ++++ 6 files changed, 1607 insertions(+), 1443 deletions(-) delete mode 100644 src/aggregation/metric/cardinality.rs create mode 100644 src/aggregation/metric/cardinality/mod.rs create mode 100644 src/aggregation/metric/cardinality/numeric_collector.rs create mode 100644 src/aggregation/metric/cardinality/str_collector.rs create mode 100644 src/aggregation/metric/cardinality/term_ord_accumulator.rs diff --git a/src/aggregation/agg_data.rs b/src/aggregation/agg_data.rs index bcb0acedd..16a08c0e4 100644 --- a/src/aggregation/agg_data.rs +++ b/src/aggregation/agg_data.rs @@ -22,9 +22,8 @@ use crate::aggregation::bucket::{ use crate::aggregation::metric::{ build_segment_stats_collector, AverageAggregation, CardinalityAggReqData, CardinalityAggregationReq, CountAggregation, ExtendedStatsAggregation, MaxAggregation, - MetricAggReqData, MinAggregation, SegmentCardinalityCollector, SegmentExtendedStatsCollector, - SegmentPercentilesCollector, StatsAggregation, StatsType, SumAggregation, TermOrdSet, - TopHitsAggReqData, TopHitsSegmentCollector, BITSET_MAX_TERM_ORD, + MetricAggReqData, MinAggregation, SegmentExtendedStatsCollector, SegmentPercentilesCollector, + StatsAggregation, StatsType, SumAggregation, TopHitsAggReqData, TopHitsSegmentCollector, }; use crate::aggregation::segment_agg_result::{ GenericSegmentAggregationResultsCollector, SegmentAggregationCollector, @@ -286,39 +285,7 @@ pub(crate) fn build_segment_agg_collector( Ok(Box::new(TermMissingAgg::new(req, node)?)) } AggKind::Cardinality => { - let req_data = req.get_cardinality_req_data(node.idx_in_req_data); - // For str columns, choose the per-bucket entries representation - // based on the segment's column.max_value(): - // * small (< BITSET_MAX_TERM_ORD): `BitSet`, pre-allocated, no promotion machinery. - // * large: `TermOrdSet` (sparse FxHashSet that promotes to a paged bitset). - // For non-str columns the `entries` field is unused (values go - // straight into the HLL sketch); we still pick `TermOrdSet` - // because its empty Sparse(FxHashSet) costs nothing. - let is_str = req_data.column_type == ColumnType::Str; - let max_term_ord_inclusive = if is_str { - req_data.accessor.max_value() - } else { - 0 - }; - let collector: Box = - if is_str && max_term_ord_inclusive < BITSET_MAX_TERM_ORD { - Box::new(SegmentCardinalityCollector::::from_req( - req_data.column_type, - node.idx_in_req_data, - req_data.accessor.clone(), - req_data.missing_value_for_accessor, - max_term_ord_inclusive, - )) - } else { - Box::new(SegmentCardinalityCollector::::from_req( - req_data.column_type, - node.idx_in_req_data, - req_data.accessor.clone(), - req_data.missing_value_for_accessor, - max_term_ord_inclusive, - )) - }; - Ok(collector) + crate::aggregation::metric::build_segment_cardinality_collector(req, node) } AggKind::StatsKind(stats_type) => { let req_data = &mut req.per_request.stats_metric_req_data[node.idx_in_req_data]; diff --git a/src/aggregation/metric/cardinality.rs b/src/aggregation/metric/cardinality.rs deleted file mode 100644 index d605125e2..000000000 --- a/src/aggregation/metric/cardinality.rs +++ /dev/null @@ -1,1407 +0,0 @@ -use std::fmt::Debug; -use std::hash::Hash; -use std::io; - -use columnar::column_values::CompactSpaceU64Accessor; -use columnar::{Column, ColumnType, Dictionary, StrColumn}; -use common::{BitSet, TinySet}; -use datasketches::hll::{Coupon, HllSketch, HllType, HllUnion}; -use rustc_hash::{FxBuildHasher, FxHashMap, FxHashSet}; -use serde::{Deserialize, Deserializer, Serialize, Serializer}; - -use crate::aggregation::agg_data::AggregationsSegmentCtx; -use crate::aggregation::intermediate_agg_result::{ - IntermediateAggregationResult, IntermediateAggregationResults, IntermediateMetricResult, -}; -use crate::aggregation::segment_agg_result::SegmentAggregationCollector; -use crate::aggregation::*; -use crate::TantivyError; - -/// Log2 of the number of registers for the HLL sketch. -/// 2^11 = 2048 registers, giving ~2.3% relative error and ~1KB per sketch (Hll4). -const LG_K: u8 = 11; - -/// Promote FxHashSet -> PagedBitset at ~3% density (`len * 32 > -/// dict_num_terms`). Past this point the bitset (~`dict_num_terms / 7.5` -/// bytes) is smaller than the hashset (~10 B/entry minimum) and avoids -/// the per-insert hash. -const PROMOTION_RATIO: u64 = 32; - -/// # Cardinality -/// -/// The cardinality aggregation allows for computing an estimate -/// of the number of different values in a data set based on the -/// Apache DataSketches HyperLogLog algorithm. This is particularly useful for -/// understanding the uniqueness of values in a large dataset where counting -/// each unique value individually would be computationally expensive. -/// -/// For example, you might use a cardinality aggregation to estimate the number -/// of unique visitors to a website by aggregating on a field that contains -/// user IDs or session IDs. -/// -/// To use the cardinality aggregation, you'll need to provide a field to -/// aggregate on. The following example demonstrates a request for the cardinality -/// of the "user_id" field: -/// -/// ```JSON -/// { -/// "cardinality": { -/// "field": "user_id" -/// } -/// } -/// ``` -/// -/// This request will return an estimate of the number of unique values in the -/// "user_id" field. -/// -/// ## Missing Values -/// -/// The `missing` parameter defines how documents that are missing a value should be treated. -/// By default, documents without a value for the specified field are ignored. However, you can -/// specify a default value for these documents using the `missing` parameter. This can be useful -/// when you want to include documents with missing values in the aggregation. -/// -/// For example, the following request treats documents with missing values in the "user_id" -/// field as if they had a value of "unknown": -/// -/// ```JSON -/// { -/// "cardinality": { -/// "field": "user_id", -/// "missing": "unknown" -/// } -/// } -/// ``` -/// -/// # Estimation Accuracy -/// -/// The cardinality aggregation provides an approximate count, which is usually -/// accurate within a small error range. This trade-off allows for efficient -/// computation even on very large datasets. -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -pub struct CardinalityAggregationReq { - /// The field name to compute the percentiles on. - pub field: String, - /// The missing parameter defines how documents that are missing a value should be treated. - /// By default they will be ignored but it is also possible to treat them as if they had a - /// value. Examples in JSON format: - /// { "field": "my_numbers", "missing": "10.0" } - #[serde(skip_serializing_if = "Option::is_none", default)] - pub missing: Option, -} - -/// Contains all information required by the SegmentCardinalityCollector to perform the -/// cardinality aggregation on a segment. -pub(crate) struct CardinalityAggReqData { - /// The column accessor to access the fast field values. - pub(crate) accessor: Column, - /// The column_type of the field. - pub(crate) column_type: ColumnType, - /// The string dictionary column if the field is of type string. - pub(crate) str_dict_column: Option, - /// The missing value normalized to the internal u64 representation of the field type. - pub(crate) missing_value_for_accessor: Option, - /// The name of the aggregation. - pub(crate) name: String, - /// The aggregation request. - pub(crate) req: CardinalityAggregationReq, -} - -impl CardinalityAggReqData { - /// Estimate the memory consumption of this struct in bytes. - pub fn get_memory_consumption(&self) -> usize { - std::mem::size_of::() - } -} - -impl CardinalityAggregationReq { - /// Creates a new [`CardinalityAggregationReq`] instance from a field name. - pub fn from_field_name(field_name: String) -> Self { - Self { - field: field_name, - missing: None, - } - } - /// Returns the field name the aggregation is computed on. - pub fn field_name(&self) -> &str { - &self.field - } -} - -/// A CouponCache is here to cache the mapping term ordinal -> coupon (see above). -/// The idea is that we do not want to fetch terms associated to several term ordinals, -/// several times due to the fact that we have several buckets. -enum CouponCache { - Dense { - coupon_map: Vec, - missing_coupon_opt: Option, - }, - Sparse { - coupon_map: FxHashMap, - missing_coupon_opt: Option, - }, -} - -impl CouponCache { - fn new( - term_ords: Vec, - coupons: Vec, - missing_coupon_opt: Option, - ) -> CouponCache { - let num_terms = term_ords.len(); - assert_eq!(num_terms, coupons.len()); - if term_ords.is_empty() { - return CouponCache::Dense { - coupon_map: Vec::new(), - missing_coupon_opt, - }; - } - let highest_term_ord = term_ords.last().copied().unwrap_or(0u64); - // We prefer the dense implementation, if it is not too wasteful. - // There are two cases for which we can use it. - // 1- if the data is small. - // 2- if the data is not necessarily small, but due to a high occupancy ratio, the RAM usage - // is not that much bigger than if we had used a HashSet. (occupancy ratio + extra - // metadata ~ x2.25) - let should_use_dense = - highest_term_ord < 1_000_000u64 || highest_term_ord < num_terms as u64 * 3u64; - if should_use_dense { - // We don't really care about the value here. We will populate all the values we will - // read anyway. - let uninitialized_coupon = Coupon::from_hash(0); - let mut coupon_map: Vec = - vec![uninitialized_coupon; highest_term_ord as usize + 1]; - - for (term_ord, coupon) in term_ords.into_iter().zip(coupons) { - coupon_map[term_ord as usize] = coupon; - } - CouponCache::Dense { - coupon_map, - missing_coupon_opt, - } - } else { - let coupon_map: FxHashMap = term_ords.into_iter().zip(coupons).collect(); - CouponCache::Sparse { - coupon_map, - missing_coupon_opt, - } - } - } -} - -// ================================================================= -// PagedBitset: a sparse bitset indexed by term_ord. -// -// Used as the dense alternative to FxHashSet once a string -// cardinality bucket has accumulated enough unique term ordinals. -// Memory is bounded to (touched pages) * (page bytes), not -// (max_term_ord / 8). -// -// Page geometry mirrors `PagedTermMap` in `term_agg.rs`: 1024 ords -// per page, lazy `Vec>>` directory. -// ================================================================= -const BITSET_PAGE_SHIFT: u32 = 10; -const BITSET_PAGE_BITS: u64 = 1u64 << BITSET_PAGE_SHIFT; // 1024 -const BITSET_PAGE_MASK: u64 = BITSET_PAGE_BITS - 1; -const BITSET_WORDS_PER_PAGE: usize = (BITSET_PAGE_BITS / 64) as usize; // 16 - -#[derive(Clone)] -struct PagedBitsetPage { - words: [TinySet; BITSET_WORDS_PER_PAGE], -} - -impl PagedBitsetPage { - fn new() -> Self { - Self { - words: [TinySet::empty(); BITSET_WORDS_PER_PAGE], - } - } -} - -pub(crate) struct PagedBitset { - pages: Vec>>, - /// Cached number of set bits, maintained on insert. - count: u64, -} - -impl PagedBitset { - /// Allocates a directory big enough to hold ords up to and including - /// `max_term_ord`. Pages are allocated lazily on first set. - fn with_max_term_ord(max_term_ord: u64) -> Self { - let max_page_idx = (max_term_ord >> BITSET_PAGE_SHIFT) as usize; - let num_pages = max_page_idx + 1; - Self { - pages: vec![None; num_pages], - count: 0, - } - } - - #[inline] - fn insert(&mut self, term_ord: u64) { - let page_idx = (term_ord >> BITSET_PAGE_SHIFT) as usize; - let intra = term_ord & BITSET_PAGE_MASK; - let word_idx = (intra >> 6) as usize; - let bit_idx = (intra & 63) as u32; - - let page = match &mut self.pages[page_idx] { - Some(p) => p, - None => { - self.pages[page_idx] = Some(Box::new(PagedBitsetPage::new())); - self.pages[page_idx].as_mut().unwrap() - } - }; - if page.words[word_idx].insert_mut(bit_idx) { - self.count += 1; - } - } - - /// Number of set bits. O(1). - #[inline] - fn len(&self) -> u64 { - self.count - } - - /// Iterate set ords in ascending order. - fn iter_sorted(&self) -> impl Iterator + '_ { - self.pages - .iter() - .enumerate() - .filter_map(|(page_idx, page_opt)| page_opt.as_ref().map(|p| (page_idx, p))) - .flat_map(|(page_idx, page)| { - let page_base_ord = (page_idx as u64) << BITSET_PAGE_SHIFT; - page.words - .iter() - .enumerate() - .flat_map(move |(word_idx, &word)| { - let word_base_ord = page_base_ord + (word_idx as u64) * 64; - word.into_iter() - .map(move |bit| word_base_ord + u64::from(bit)) - }) - }) - } -} - -/// Threshold below which we use `BitSet` instead of `TermOrdSet`. -/// -/// Both `BitSet` and `FxHashSet` have the same 32-byte struct, so the comparison is heap only: -/// * `BitSet` at T=256: 5 `TinySet` words covering 258 bits (with the missing-value sentinel) = -/// 40 bytes. -/// * `FxHashSet` after one insert: 4-bucket hashbrown table ≈ 56 bytes -pub(crate) const BITSET_MAX_TERM_ORD: u64 = 256; - -// ================================================================= -// TermOrdAccumulator: per-bucket abstraction over the entries set. -// -// Implementations: -// - `BitSet` (from `common`): used when `column.max_value()` is small (< BITSET_MAX_TERM_ORD). -// Pre-allocated, no promotion. -// - `TermOrdSet`: adaptive, starts as FxHashSet and promotes to a paged bitset when occupancy -// crosses the density threshold (only if promotion is enabled — typically gated on top-level -// aggregation). -// -// The trait lets `SegmentCardinalityCollector` be generic over the choice -// so the hot collect() loop monomorphizes to a direct call (no enum -// dispatch per insert). -// ================================================================= -pub(crate) trait TermOrdAccumulator: Sized { - /// Construct an empty accumulator. - /// `max_term_ord_inclusive` is the largest term_ord that may be - /// inserted (used to size pre-allocated bitsets and the dense bitset - /// on promotion). - fn new(max_term_ord_inclusive: u64) -> Self; - fn insert(&mut self, term_ord: u64); - /// Bulk insert. Implementations may override to hoist any inner - /// dispatch outside the loop. Default loops `insert`. - #[inline] - fn extend_from_iter>(&mut self, ords: I) { - for ord in ords { - self.insert(ord); - } - } - /// Hook called once per ingested block. Adaptive impls use this to - /// decide on sparse->dense promotion. - fn maybe_compact(&mut self) {} - fn len(&self) -> usize; - fn iter_ords(&self) -> impl Iterator + '_; -} - -impl TermOrdAccumulator for BitSet { - #[inline] - fn new(max_term_ord_inclusive: u64) -> Self { - // `BitSet::with_max_value(M)` accepts ords in [0, M). - // We need ords up to and including `max_term_ord_inclusive`, plus - // the missing-value sentinel `column.max_value() + 1`. - BitSet::with_max_value((max_term_ord_inclusive + 2) as u32) - } - #[inline] - fn insert(&mut self, term_ord: u64) { - BitSet::insert(self, term_ord as u32); - } - #[inline] - fn len(&self) -> usize { - BitSet::len(self) - } - fn iter_ords(&self) -> impl Iterator + '_ { - // `BitSet` itself doesn't expose iteration, but - // `BitSet::tinyset(bucket)` does. Walk per-bucket and yield each - // set bit. The capacity is `max_value()`; iterating to - // `div_ceil(64)` covers every possible ord exactly once. - let num_buckets = self.max_value().div_ceil(64); - (0..num_buckets).flat_map(move |bucket| { - let chunk_base = u64::from(bucket) * 64; - self.tinyset(bucket) - .into_iter() - .map(move |bit| chunk_base + u64::from(bit)) - }) - } -} - -// ================================================================= -// TermOrdSet: adaptive sparse->dense accumulator. -// -// Starts as an FxHashSet (cheap when few ords are seen). When occupancy -// crosses `len * PROMOTION_RATIO > max_term_ord_inclusive`, drains into -// a `PagedBitset` and continues dense. Promotion is one-way. -// ================================================================= -pub(crate) struct TermOrdSet { - inner: TermOrdSetInner, - /// Largest term_ord that may be inserted. Used for both sizing the - /// dense bitset on promotion and as the promotion-threshold reference. - max_term_ord_inclusive: u64, -} - -enum TermOrdSetInner { - Sparse(FxHashSet), - Dense(PagedBitset), -} - -impl TermOrdAccumulator for TermOrdSet { - fn new(max_term_ord_inclusive: u64) -> Self { - Self { - inner: TermOrdSetInner::Sparse(FxHashSet::default()), - max_term_ord_inclusive, - } - } - - #[inline] - fn insert(&mut self, term_ord: u64) { - match &mut self.inner { - TermOrdSetInner::Sparse(set) => { - set.insert(term_ord); - } - TermOrdSetInner::Dense(bitset) => bitset.insert(term_ord), - } - } - - /// Hoist the Sparse/Dense match outside the per-ord loop so that a - /// block of inserts dispatches once. - fn extend_from_iter>(&mut self, ords: I) { - match &mut self.inner { - TermOrdSetInner::Sparse(set) => { - for ord in ords { - set.insert(ord); - } - } - TermOrdSetInner::Dense(bitset) => { - for ord in ords { - bitset.insert(ord); - } - } - } - } - - fn maybe_compact(&mut self) { - let TermOrdSetInner::Sparse(set) = &mut self.inner else { - return; - }; - if set.len() as u64 * PROMOTION_RATIO <= self.max_term_ord_inclusive { - return; - } - // Size for ord <= max_term_ord_inclusive plus the missing sentinel - // (column.max_value() + 1, which may equal max_term_ord_inclusive - // when the column references every dictionary term). - let mut bitset = PagedBitset::with_max_term_ord(self.max_term_ord_inclusive + 1); - let set = std::mem::take(set); - for ord in set { - bitset.insert(ord); - } - self.inner = TermOrdSetInner::Dense(bitset); - } - - fn len(&self) -> usize { - match &self.inner { - TermOrdSetInner::Sparse(set) => set.len(), - TermOrdSetInner::Dense(bitset) => bitset.len() as usize, - } - } - - fn iter_ords(&self) -> impl Iterator + '_ { - match &self.inner { - TermOrdSetInner::Sparse(set) => itertools::Either::Left(set.iter().copied()), - TermOrdSetInner::Dense(bitset) => itertools::Either::Right(bitset.iter_sorted()), - } - } -} - -pub(crate) struct SegmentCardinalityCollector { - /// Buckets are Some(_) until they get consumed by into_intermediate_results(). - buckets: Vec>>, - accessor_idx: usize, - /// The column accessor to access the fast field values. - accessor: Column, - /// The column_type of the field. - column_type: ColumnType, - /// The missing value normalized to the internal u64 representation of the field type. - missing_value_for_accessor: Option, - coupon_cache: Option, - /// Largest term_ord that may be inserted into a bucket. For str columns - /// this is `accessor.max_value()`; for non-str columns this is unused - /// (no inserts go into `entries`) and set to 0. - max_term_ord_inclusive: u64, -} - -impl Debug for SegmentCardinalityCollector { - fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result { - f.debug_struct("SegmentCardinalityCollector") - .field("column_type", &self.column_type) - .field( - "missing_value_for_accessor", - &self.missing_value_for_accessor, - ) - .finish() - } -} - -/// Per-bucket state. Shape depends on column kind: str columns dedup -/// term ords and only build the HLL sketch at finalization (saves the -/// ~96 B `CardinalityCollector` per bucket during collect); numeric/IpAddr -/// columns feed the sketch directly during collect. -pub(crate) enum SegmentCardinalityCollectorBucket { - Str(S), - Numeric(CardinalityCollector), -} -impl SegmentCardinalityCollectorBucket { - #[inline(always)] - pub fn new(column_type: ColumnType, max_term_ord_inclusive: u64) -> Self { - if column_type == ColumnType::Str { - Self::Str(S::new(max_term_ord_inclusive)) - } else { - Self::Numeric(CardinalityCollector::new(column_type as u8)) - } - } - - // Returns a intermediate metric result. - // - // If the column is not str, the values have been added to the - // sketch during collection. - // - // If the column is str, then the values are dictionary encoded - // and have not been added to the sketch yet. - // We need to resolves the term ords accumulated in the str entries - // with the coupon cache, and append the results to a fresh sketch. - fn into_intermediate_metric_result( - self, - coupon_cache_opt: Option<&CouponCache>, - ) -> crate::Result { - let cardinality = match self { - Self::Str(entries) => { - let mut cardinality = CardinalityCollector::new(ColumnType::Str as u8); - if let Some(coupon_cache) = coupon_cache_opt { - // Sketch must be empty for str columns: coupons are appended here - // from the term_ord set (and not directly during collection). - assert!(cardinality.sketch.is_empty()); - append_to_sketch(&entries, coupon_cache, &mut cardinality); - } - cardinality - } - Self::Numeric(cardinality) => cardinality, - }; - Ok(IntermediateMetricResult::Cardinality(cardinality)) - } -} - -/// Builds a coupon cache from the given buckets, dictionary, and optional missing value. -/// Returns a mapping from term_ord to the hash (coupon) of the associated term. -fn build_coupon_cache( - buckets: &[Option>], - dictionary: &Dictionary, - missing_value_opt: Option<&Key>, -) -> io::Result { - // Caller restricts this to str cardinality collectors, so every - // present bucket must be the `Str` variant. Pass 1 validates and - // computes the capacity hint; pass 2 inserts. - let mut max_bucket_len = 0usize; - for bucket in buckets.iter().flatten() { - match bucket { - SegmentCardinalityCollectorBucket::Str(entries) => { - max_bucket_len = max_bucket_len.max(entries.len()); - } - SegmentCardinalityCollectorBucket::Numeric(_) => { - return Err(io::Error::other( - "build_coupon_cache invoked with a non-str bucket", - )); - } - } - } - let mut term_ords_set = FxHashSet::with_capacity_and_hasher(max_bucket_len * 2, FxBuildHasher); - for bucket in buckets.iter().flatten() { - if let SegmentCardinalityCollectorBucket::Str(entries) = bucket { - term_ords_set.extend(entries.iter_ords()); - } - } - let mut term_ords: Vec = term_ords_set.into_iter().collect(); - term_ords.sort_unstable(); - - term_ords.pop_if(|highest_term_ord| *highest_term_ord >= dictionary.num_terms() as u64); - - let mut coupons: Vec = Vec::with_capacity(term_ords.len()); - let all_term_ords_found: bool = - dictionary.sorted_ords_to_term_cb(&term_ords, |term_bytes| { - let coupon: Coupon = Coupon::from_hash(term_bytes); - coupons.push(coupon); - })?; - assert!(all_term_ords_found); - - // Regardless of whether or not there is effectively a missing value in one of the buckets, - // we populate the cache with the missing key too (if any). - let missing_coupon_opt: Option = missing_value_opt.map(|missing_key| { - if let Key::Str(missing_value_str) = missing_key { - Coupon::from_hash(missing_value_str.as_bytes()) - } else { - // See https://github.com/quickwit-oss/tantivy/issues/2891 - // A missing key with a type different from Str will not work as intended - // for the moment. - // - // Right now this is just a partial workaround. - Coupon::from_hash("__tantivy_missing_non_str__".as_bytes()) - } - }); - Ok(CouponCache::new(term_ords, coupons, missing_coupon_opt)) -} - -fn append_to_sketch( - term_ords: &S, - coupon_cache: &CouponCache, - sketch: &mut CardinalityCollector, -) { - match coupon_cache { - CouponCache::Dense { - coupon_map, - missing_coupon_opt, - } => { - for term_ord in term_ords.iter_ords() { - if let Some(coupon) = coupon_map - .get(term_ord as usize) - .copied() - .or(*missing_coupon_opt) - { - sketch.insert_coupon(coupon); - } - } - } - CouponCache::Sparse { - coupon_map, - missing_coupon_opt, - } => { - for term_ord in term_ords.iter_ords() { - if let Some(coupon) = coupon_map.get(&term_ord).copied().or(*missing_coupon_opt) { - sketch.insert_coupon(coupon); - } - } - } - } -} - -impl SegmentCardinalityCollector { - pub fn from_req( - column_type: ColumnType, - accessor_idx: usize, - accessor: Column, - missing_value_for_accessor: Option, - max_term_ord_inclusive: u64, - ) -> Self { - Self { - buckets: Vec::new(), - column_type, - accessor_idx, - accessor, - missing_value_for_accessor, - coupon_cache: None, - max_term_ord_inclusive, - } - } - - fn fetch_block_with_field( - &mut self, - docs: &[crate::DocId], - agg_data: &mut AggregationsSegmentCtx, - ) { - agg_data.column_block_accessor.fetch_block_with_missing( - docs, - &self.accessor, - self.missing_value_for_accessor, - ); - } -} - -impl SegmentAggregationCollector - for SegmentCardinalityCollector -{ - fn add_intermediate_aggregation_result( - &mut self, - agg_data: &AggregationsSegmentCtx, - results: &mut IntermediateAggregationResults, - bucket_id: BucketId, - ) -> crate::Result<()> { - self.prepare_max_bucket(bucket_id, agg_data)?; - let req_data = &agg_data.get_cardinality_req_data(self.accessor_idx); - // Strings are dictionary encoded. Fetching the terms associated to strings - // is expensive. For this reason, we do that once for all buckets and cache the results - // here. - if let Some(str_dict_column) = &req_data.str_dict_column { - // Ensure the coupon cache is populated. - // A mapping from term_ord to the hash of the associated term. - // The missing value sentinel will be associated to the hash of the missing value if - // any. - if self.coupon_cache.is_none() { - self.coupon_cache = Some(build_coupon_cache( - &self.buckets, - str_dict_column.dictionary(), - req_data.req.missing.as_ref(), - )?); - } - } - let name = req_data.name.to_string(); - // take the bucket in buckets and replace it with a new empty one - let Some(bucket) = self.buckets[bucket_id as usize].take() else { - return Err(crate::TantivyError::InternalError( - "the same bucket should not be finalized twice.".to_string(), - )); - }; - let intermediate_result = - bucket.into_intermediate_metric_result(self.coupon_cache.as_ref())?; - results.push( - name, - IntermediateAggregationResult::Metric(intermediate_result), - )?; - - Ok(()) - } - - fn collect( - &mut self, - parent_bucket_id: BucketId, - docs: &[crate::DocId], - agg_data: &mut AggregationsSegmentCtx, - ) -> crate::Result<()> { - self.fetch_block_with_field(docs, agg_data); - let Some(bucket) = &mut self.buckets[parent_bucket_id as usize].as_mut() else { - return Err(crate::TantivyError::InternalError( - "collection should not happen after finalization".to_string(), - )); - }; - let col_block_accessor = &agg_data.column_block_accessor; - match bucket { - SegmentCardinalityCollectorBucket::Str(entries) => { - // Promotion check runs on the pre-block state: the first call - // sees an empty set (no-op), and the last block of inserts - // doesn't trigger a promotion of a set we won't grow further. - // The trait dispatches once per block (via `extend_from_iter`) - // for adaptive variants and inlines to a tight loop for the - // BitSet path. - entries.maybe_compact(); - entries.extend_from_iter(col_block_accessor.iter_vals()); - } - SegmentCardinalityCollectorBucket::Numeric(cardinality) => { - if self.column_type == ColumnType::IpAddr { - let compact_space_accessor = self - .accessor - .values - .clone() - .downcast_arc::() - .map_err(|_| { - TantivyError::AggregationError( - crate::aggregation::AggregationError::InternalError( - "Type mismatch: Could not downcast to CompactSpaceU64Accessor" - .to_string(), - ), - ) - })?; - for val in col_block_accessor.iter_vals() { - let val: u128 = compact_space_accessor.compact_to_u128(val as u32); - cardinality.insert(val); - } - } else { - for val in col_block_accessor.iter_vals() { - cardinality.insert(val); - } - } - } - } - - Ok(()) - } - - fn prepare_max_bucket( - &mut self, - max_bucket: BucketId, - _agg_data: &AggregationsSegmentCtx, - ) -> crate::Result<()> { - if max_bucket as usize >= self.buckets.len() { - let column_type = self.column_type; - let max_term_ord_inclusive = self.max_term_ord_inclusive; - self.buckets.resize_with(max_bucket as usize + 1, || { - Some(SegmentCardinalityCollectorBucket::::new( - column_type, - max_term_ord_inclusive, - )) - }); - } - Ok(()) - } - - fn compute_metric_value( - &self, - bucket_id: BucketId, - sub_agg_name: &str, - sub_agg_property: &str, - agg_data: &AggregationsSegmentCtx, - ) -> Option { - let req_data = &agg_data.get_cardinality_req_data(self.accessor_idx); - if req_data.name != sub_agg_name || !sub_agg_property.is_empty() { - return None; - } - let bucket = self.buckets.get(bucket_id as usize)?.as_ref()?; - // For string columns the sketch isn't built until finalization; the - // term_ord set's len is the exact distinct count. For numeric columns - // the sketch is populated during collect. - match bucket { - SegmentCardinalityCollectorBucket::Str(entries) => Some(entries.len() as f64), - SegmentCardinalityCollectorBucket::Numeric(cardinality) => { - Some(cardinality.sketch.estimate().trunc()) - } - } - } -} - -#[derive(Clone, Debug)] -/// The cardinality collector used during segment collection and for merging results. -/// Uses Apache DataSketches HLL (lg_k=11, Hll4) for compact binary serialization -/// and cross-language compatibility (e.g. Java `datasketches` library). -pub struct CardinalityCollector { - sketch: HllSketch, - /// Salt derived from `ColumnType`, used to differentiate values of different column types - /// that map to the same u64 (e.g. bool `false` = 0 vs i64 `0`). - /// Not serialized — only needed during insertion, not after sketch registers are populated. - salt: u8, -} - -impl Default for CardinalityCollector { - fn default() -> Self { - Self::new(0) - } -} - -impl PartialEq for CardinalityCollector { - fn eq(&self, _other: &Self) -> bool { - false - } -} - -impl Serialize for CardinalityCollector { - fn serialize(&self, serializer: S) -> Result { - let bytes = self.sketch.serialize(); - serializer.serialize_bytes(&bytes) - } -} - -impl<'de> Deserialize<'de> for CardinalityCollector { - fn deserialize>(deserializer: D) -> Result { - let bytes: Vec = Deserialize::deserialize(deserializer)?; - let sketch = HllSketch::deserialize(&bytes).map_err(serde::de::Error::custom)?; - Ok(Self { sketch, salt: 0 }) - } -} - -impl CardinalityCollector { - fn new(salt: u8) -> Self { - Self { - sketch: HllSketch::new(LG_K, HllType::Hll8), - salt, - } - } - - /// Insert a value into the HLL sketch, salted by the column type. - /// The salt ensures that identical u64 values from different column types - /// (e.g. bool `false` vs i64 `0`) are counted as distinct. - fn insert(&mut self, value: T) { - self.sketch.update((self.salt, value)); - } - - fn insert_coupon(&mut self, coupon: Coupon) { - self.sketch.update_with_coupon(coupon); - } - - /// Compute the final cardinality estimate. - pub fn finalize(self) -> Option { - Some(self.sketch.estimate().trunc()) - } - - /// Serialize the HLL sketch to its compact binary representation. - /// The format is cross-language compatible with Apache DataSketches (Java, C++, Python). - pub fn to_sketch_bytes(&self) -> Vec { - self.sketch.serialize() - } - - pub(crate) fn merge_fruits(&mut self, right: CardinalityCollector) -> crate::Result<()> { - let mut union = HllUnion::new(LG_K); - union.update(&self.sketch); - union.update(&right.sketch); - self.sketch = union.to_sketch(HllType::Hll8); - Ok(()) - } -} - -#[cfg(test)] -mod tests { - - use std::net::IpAddr; - use std::str::FromStr; - - use columnar::MonotonicallyMappableToU64; - - use crate::aggregation::agg_req::Aggregations; - use crate::aggregation::tests::{exec_request, get_test_index_from_terms}; - use crate::schema::{IntoIpv6Addr, Schema, FAST, STRING}; - use crate::Index; - - #[test] - fn cardinality_aggregation_test_empty_index() -> crate::Result<()> { - let values = vec![]; - let index = get_test_index_from_terms(false, &values)?; - let agg_req: Aggregations = serde_json::from_value(json!({ - "cardinality": { - "cardinality": { - "field": "string_id", - } - }, - })) - .unwrap(); - - let res = exec_request(agg_req, &index)?; - assert_eq!(res["cardinality"]["value"], 0.0); - - Ok(()) - } - - #[test] - fn cardinality_aggregation_test_single_segment() -> crate::Result<()> { - cardinality_aggregation_test_merge_segment(true) - } - #[test] - fn cardinality_aggregation_test() -> crate::Result<()> { - cardinality_aggregation_test_merge_segment(false) - } - fn cardinality_aggregation_test_merge_segment(merge_segments: bool) -> crate::Result<()> { - let segment_and_terms = vec![ - vec!["terma"], - vec!["termb"], - vec!["termc"], - vec!["terma"], - vec!["terma"], - vec!["terma"], - vec!["termb"], - vec!["terma"], - ]; - let index = get_test_index_from_terms(merge_segments, &segment_and_terms)?; - let agg_req: Aggregations = serde_json::from_value(json!({ - "cardinality": { - "cardinality": { - "field": "string_id", - } - }, - })) - .unwrap(); - - let res = exec_request(agg_req, &index)?; - assert_eq!(res["cardinality"]["value"], 3.0); - - Ok(()) - } - - /// Build a single-segment string-cardinality index with 32 unique terms. - /// `column.max_value() = 31` is well below `BITSET_MAX_TERM_ORD`, - /// so the bucket exercises the `BitSet` path end to end. - #[test] - fn cardinality_aggregation_test_str_bitset() -> crate::Result<()> { - let terms: Vec = (0..32).map(|i| format!("term_{i}")).collect(); - let term_refs: Vec> = terms.iter().map(|t| vec![t.as_str()]).collect::>(); - // single segment so we have a single dictionary of 32 terms. - let index = get_test_index_from_terms(true, &term_refs)?; - - let agg_req: Aggregations = serde_json::from_value(json!({ - "cardinality": { - "cardinality": { "field": "string_id" } - }, - })) - .unwrap(); - - let res = exec_request(agg_req, &index)?; - assert_eq!(res["cardinality"]["value"], 32.0); - Ok(()) - } - - /// `BitSet` path with a `missing` parameter: the column-level missing - /// sentinel (`column.max_value() + 1`) flows into the bitset, the - /// dict lookup filter at finalization drops it, and the missing - /// coupon is applied separately. - #[test] - fn cardinality_aggregation_test_str_bitset_with_missing() { - let mut schema_builder = Schema::builder(); - let name_field = schema_builder.add_text_field("name", STRING | FAST); - let index = Index::create_in_ram(schema_builder.build()); - let mut writer = index.writer_for_tests().unwrap(); - for i in 0..16 { - let term = format!("t{i:02}"); - writer.add_document(doc!(name_field => term)).unwrap(); - } - // One empty doc, exercising the missing sentinel. - writer.add_document(doc!()).unwrap(); - writer.commit().unwrap(); - - let agg_req: Aggregations = serde_json::from_value(json!({ - "cardinality": { - "cardinality": { - "field": "name", - "missing": "MISSING_SENTINEL_KEY", - } - }, - })) - .unwrap(); - - let res = exec_request(agg_req, &index).unwrap(); - // 16 distinct real terms + 1 distinct "missing" value = 17. - assert_eq!(res["cardinality"]["value"], 17.0); - } - - /// Unit-test the PagedBitset itself: cross-page inserts produce sorted - /// iteration, len() matches the inserted set, and duplicates are - /// idempotent. - #[test] - fn paged_bitset_basic() { - use super::PagedBitset; - // Span several pages: BITSET_PAGE_BITS = 1024, so ords > 1024 land - // on the second page, > 2048 on the third, etc. - let ords = [0u64, 1, 63, 64, 1023, 1024, 1025, 4096, 4097, 9999, 10_000]; - let max_ord = *ords.iter().max().unwrap(); - let mut bitset = PagedBitset::with_max_term_ord(max_ord); - for &ord in &ords { - bitset.insert(ord); - // Idempotent: inserting again must not increase count. - bitset.insert(ord); - } - assert_eq!(bitset.len(), ords.len() as u64); - let collected: Vec = bitset.iter_sorted().collect(); - let mut expected: Vec = ords.to_vec(); - expected.sort_unstable(); - assert_eq!(collected, expected); - } - - /// Unit-test `TermOrdSet`: starts Sparse, promotes to Dense on - /// `maybe_compact` once the density threshold is crossed, and - /// `iter_ords()` yields the same set in either state. Ords spanning - /// multiple paged-bitset pages exercise the Dense iter ordering. - #[test] - fn term_ord_set_promotes_on_maybe_compact() { - use super::{TermOrdAccumulator, TermOrdSet, PROMOTION_RATIO}; - // Pick max so promotion needs few inserts: len * RATIO > max with - // RATIO=32 and max=64 trips at len=3 (3*32=96 > 64). - let max_term_ord = 64u64; - let mut set = ::new(max_term_ord); - // Two inserts: should stay Sparse after maybe_compact (2 * RATIO = 64, not > 64). - set.insert(0); - set.insert(7); - set.maybe_compact(); - assert_eq!(set.len(), 2); - - // Third insert promotes on next maybe_compact. - set.insert(20); - assert_eq!(set.len(), 3); - // Sanity check: at len=3, 3 * PROMOTION_RATIO = 96 > 64. - assert!(3u64 * PROMOTION_RATIO > max_term_ord); - set.maybe_compact(); - - // Post-promotion: extending continues to work. - set.insert(15); - set.insert(15); // dup - assert_eq!(set.len(), 4); - - let mut collected: Vec = set.iter_ords().collect(); - collected.sort_unstable(); - assert_eq!(collected, vec![0, 7, 15, 20]); - } - - /// Unit-test the `BitSet` impl of `TermOrdAccumulator`: insert, - /// dedup, and iter_ords order. - #[test] - fn bitset_accumulator_basic() { - use common::BitSet; - - use super::TermOrdAccumulator; - let mut set = ::new(255); - for ord in [0u64, 1, 63, 64, 65, 128, 200, 200, 0] { - ::insert(&mut set, ord); - } - assert_eq!(::len(&set), 7); - let collected: Vec = set.iter_ords().collect(); - assert_eq!(collected, vec![0, 1, 63, 64, 65, 128, 200]); - } - - #[test] - fn cardinality_aggregation_u64() -> crate::Result<()> { - let mut schema_builder = Schema::builder(); - let id_field = schema_builder.add_u64_field("id", FAST); - let index = Index::create_in_ram(schema_builder.build()); - { - let mut writer = index.writer_for_tests()?; - writer.add_document(doc!(id_field => 1u64))?; - writer.add_document(doc!(id_field => 2u64))?; - writer.add_document(doc!(id_field => 3u64))?; - writer.add_document(doc!())?; - writer.commit()?; - } - - let agg_req: Aggregations = serde_json::from_value(json!({ - "cardinality": { - "cardinality": { - "field": "id", - "missing": 0u64 - }, - } - })) - .unwrap(); - - let res = exec_request(agg_req, &index)?; - assert_eq!(res["cardinality"]["value"], 4.0); - - Ok(()) - } - - #[test] - fn cardinality_aggregation_ip_addr() -> crate::Result<()> { - let mut schema_builder = Schema::builder(); - let field = schema_builder.add_ip_addr_field("ip_field", FAST); - let index = Index::create_in_ram(schema_builder.build()); - { - let mut writer = index.writer_for_tests()?; - // IpV6 loopback - writer.add_document(doc!(field=>IpAddr::from_str("::1").unwrap().into_ipv6_addr()))?; - writer.add_document(doc!(field=>IpAddr::from_str("::1").unwrap().into_ipv6_addr()))?; - // IpV4 - writer.add_document( - doc!(field=>IpAddr::from_str("127.0.0.1").unwrap().into_ipv6_addr()), - )?; - writer.commit()?; - } - - let agg_req: Aggregations = serde_json::from_value(json!({ - "cardinality": { - "cardinality": { - "field": "ip_field" - }, - } - })) - .unwrap(); - - let res = exec_request(agg_req, &index)?; - assert_eq!(res["cardinality"]["value"], 2.0); - - Ok(()) - } - - #[test] - fn cardinality_aggregation_bytes_excluded_from_accessors() -> crate::Result<()> { - // `Bytes` columns are opened as raw per-segment dictionary ordinals (like `Str`), but - // unlike `Str`, cardinality has no dictionary-resolution path for them: it would hash - // the raw ordinal directly, which are segment dependant. Ignore bytes values instead of - // counting them wrong. - let mut schema_builder = Schema::builder(); - let field = schema_builder.add_bytes_field("raw", FAST); - let index = Index::create_in_ram(schema_builder.build()); - { - let mut writer = index.writer_for_tests()?; - writer.add_document(doc!(field => vec![1u8]))?; - writer.add_document(doc!(field => vec![2u8]))?; - writer.commit()?; - writer.add_document(doc!(field => vec![3u8]))?; - writer.add_document(doc!(field => vec![4u8]))?; - writer.commit()?; - } - - let agg_req: Aggregations = serde_json::from_value(json!({ - "cardinality": { - "cardinality": { - "field": "raw" - }, - } - })) - .unwrap(); - - let res = exec_request(agg_req, &index)?; - assert_eq!(res["cardinality"]["value"], 0.0); - - Ok(()) - } - - #[test] - fn cardinality_aggregation_json() -> crate::Result<()> { - let mut schema_builder = Schema::builder(); - let field = schema_builder.add_json_field("json", FAST); - let index = Index::create_in_ram(schema_builder.build()); - { - let mut writer = index.writer_for_tests()?; - writer.add_document(doc!(field => json!({"value": false})))?; - writer.add_document(doc!(field => json!({"value": true})))?; - writer.add_document(doc!(field => json!({"value": i64::from_u64(0u64)})))?; - writer.add_document(doc!(field => json!({"value": i64::from_u64(1u64)})))?; - writer.commit()?; - } - - let agg_req: Aggregations = serde_json::from_value(json!({ - "cardinality": { - "cardinality": { - "field": "json.value" - }, - } - })) - .unwrap(); - - let res = exec_request(agg_req, &index)?; - assert_eq!(res["cardinality"]["value"], 4.0); - - Ok(()) - } - - /// A JSON path that resolves to both a Str column and a numeric column - /// produces two collector instances per segment — one with `Str` buckets - /// and one with `Numeric` buckets. Their `IntermediateMetricResult`s must - /// merge into the union cardinality. - #[test] - fn cardinality_aggregation_json_str_and_numeric() -> crate::Result<()> { - let mut schema_builder = Schema::builder(); - let field = schema_builder.add_json_field("json", FAST); - let index = Index::create_in_ram(schema_builder.build()); - { - let mut writer = index.writer_for_tests()?; - writer.add_document(doc!(field => json!({"value": "hello"})))?; - writer.add_document(doc!(field => json!({"value": "world"})))?; - writer.add_document(doc!(field => json!({"value": "hello"})))?; // dup str - writer.add_document(doc!(field => json!({"value": i64::from_u64(7u64)})))?; - writer.add_document(doc!(field => json!({"value": i64::from_u64(42u64)})))?; - writer.add_document(doc!(field => json!({"value": i64::from_u64(7u64)})))?; // dup num - writer.commit()?; - } - - let agg_req: Aggregations = serde_json::from_value(json!({ - "cardinality": { - "cardinality": { - "field": "json.value" - }, - } - })) - .unwrap(); - - let res = exec_request(agg_req, &index)?; - // 4 distinct values: "hello", "world", 7, 42. - assert_eq!(res["cardinality"]["value"], 4.0); - - Ok(()) - } - - #[test] - fn cardinality_collector_serde_roundtrip() { - use super::CardinalityCollector; - - let mut collector = CardinalityCollector::default(); - collector.insert("hello"); - collector.insert("world"); - collector.insert("hello"); // duplicate - - let serialized = serde_json::to_vec(&collector).unwrap(); - let deserialized: CardinalityCollector = serde_json::from_slice(&serialized).unwrap(); - - let original_estimate = collector.finalize().unwrap(); - let roundtrip_estimate = deserialized.finalize().unwrap(); - assert_eq!(original_estimate, roundtrip_estimate); - assert_eq!(original_estimate, 2.0); - } - - #[test] - fn cardinality_collector_merge() { - use super::CardinalityCollector; - - let mut left = CardinalityCollector::default(); - left.insert("a"); - left.insert("b"); - - let mut right = CardinalityCollector::default(); - right.insert("b"); - right.insert("c"); - - left.merge_fruits(right).unwrap(); - let estimate = left.finalize().unwrap(); - assert_eq!(estimate, 3.0); - } - - /// Verifies that merging two small sketches (both in List/Set coupon mode) - /// produces an exact result — i.e. the HllUnion does not unnecessarily - /// promote to the full HLL array when the combined cardinality is small. - #[test] - fn cardinality_collector_merge_stays_exact_for_small_sets() { - use super::CardinalityCollector; - - let mut left = CardinalityCollector::default(); - for i in 0u64..50 { - left.insert(i); - } - - let mut right = CardinalityCollector::default(); - for i in 30u64..100 { - right.insert(i); - } - - left.merge_fruits(right).unwrap(); - let estimate = left.finalize().unwrap(); - // 100 distinct values (0..100). Both sketches are in Set mode (< 192 coupons), - // so the union should stay in coupon mode and give an exact count. - assert_eq!(estimate, 100.0); - } - - #[test] - fn cardinality_collector_serialize_deserialize_binary() { - use datasketches::hll::HllSketch; - - use super::CardinalityCollector; - - let mut collector = CardinalityCollector::default(); - collector.insert("apple"); - collector.insert("banana"); - collector.insert("cherry"); - - let bytes = collector.to_sketch_bytes(); - let deserialized = HllSketch::deserialize(&bytes).unwrap(); - assert!((deserialized.estimate() - 3.0).abs() < 0.01); - } - - /// Tests that the `missing` parameter correctly counts a single empty document - /// for both u64 and str columns. - #[test] - fn cardinality_aggregation_missing_value_single_empty_doc() { - let mut schema_builder = Schema::builder(); - let id_field = schema_builder.add_u64_field("id", FAST); - let name_field = schema_builder.add_text_field("name", STRING | FAST); - let index = Index::create_in_ram(schema_builder.build()); - let mut writer = index.writer_for_tests().unwrap(); - writer - .add_document(doc!(id_field=>1u64,name_field=>"some_name")) - .unwrap(); - writer.add_document(doc!()).unwrap(); - writer.commit().unwrap(); - - { - // int colum with missing value non redundant - let agg_req: Aggregations = serde_json::from_value(json!({ - "cardinality": { - "cardinality": { - "field": "id", - "missing": 42u64 - }, - } - })) - .unwrap(); - let res = exec_request(agg_req, &index).unwrap(); - assert_eq!(res["cardinality"]["value"], 2.0); - } - - { - // int colum with missing value redundant - let agg_req: Aggregations = serde_json::from_value(json!({ - "cardinality": { - "cardinality": { - "field": "id", - "missing": 1u64 - }, - } - })) - .unwrap(); - let res = exec_request(agg_req, &index).unwrap(); - assert_eq!(res["cardinality"]["value"], 1.0); - } - - { - // str colum with missing value non redundant - // With more than one segment, this is not well handled. - let agg_req: Aggregations = serde_json::from_value(json!({ - "cardinality": { - "cardinality": { - "field": "name", - "missing": "other_name" - }, - } - })) - .unwrap(); - let res = exec_request(agg_req, &index).unwrap(); - assert_eq!(res["cardinality"]["value"], 2.0); - } - - { - // str colum with missing value redundant - let agg_req: Aggregations = serde_json::from_value(json!({ - "cardinality": { - "cardinality": { - "field": "name", - "missing": "some_name" - }, - } - })) - .unwrap(); - let res = exec_request(agg_req, &index).unwrap(); - assert_eq!(res["cardinality"]["value"], 1.0); - } - - { - // str column with missing value with a number type. - let agg_req: Aggregations = serde_json::from_value(json!({ - "cardinality": { - "cardinality": { - "field": "name", - "missing": 3, - }, - } - })) - .unwrap(); - let res = exec_request(agg_req, &index).unwrap(); - assert_eq!(res["cardinality"]["value"], 2.0); - } - } - - #[test] - fn cardinality_collector_salt_differentiates_types() { - use super::CardinalityCollector; - - // Without salt, same u64 value from different column types would collide - let mut collector_bool = CardinalityCollector::new(5); // e.g. ColumnType::Bool - collector_bool.insert(0u64); // false - collector_bool.insert(1u64); // true - - let mut collector_i64 = CardinalityCollector::new(2); // e.g. ColumnType::I64 - collector_i64.insert(0u64); - collector_i64.insert(1u64); - - // Merge them - collector_bool.merge_fruits(collector_i64).unwrap(); - let estimate = collector_bool.finalize().unwrap(); - // Should be 4 because salt makes (5, 0) != (2, 0) and (5, 1) != (2, 1) - assert_eq!(estimate, 4.0); - } -} diff --git a/src/aggregation/metric/cardinality/mod.rs b/src/aggregation/metric/cardinality/mod.rs new file mode 100644 index 000000000..72d43de18 --- /dev/null +++ b/src/aggregation/metric/cardinality/mod.rs @@ -0,0 +1,609 @@ +//! Cardinality aggregation. +//! +//! * [`str_collector`] holds the `ColumnType::Str` segment collector, which accumulates term +//! ordinals and resolves them to HLL coupons at finalization. +//! * [`numeric_collector`] holds the segment collector for every other column type, which feeds +//! the HLL sketch directly during collection. +//! +//! Both segment collectors converge on the same +//! [`IntermediateMetricResult::Cardinality`] payload, so results coming from a +//! str column and from a numeric column (e.g. a JSON path that resolves to +//! both) merge through the single [`CardinalityCollector::merge_fruits`] +//! implementation here. + +mod numeric_collector; +mod str_collector; +mod term_ord_accumulator; + +use std::hash::Hash; + +use columnar::{Column, ColumnType, StrColumn}; +use common::BitSet; +use datasketches::hll::{Coupon, HllSketch, HllType, HllUnion}; +pub(crate) use numeric_collector::SegmentNumericCardinalityCollector; +use serde::{Deserialize, Deserializer, Serialize, Serializer}; +pub(crate) use str_collector::SegmentStrCardinalityCollector; +pub(crate) use term_ord_accumulator::{TermOrdSet, BITSET_MAX_TERM_ORD}; + +use crate::aggregation::agg_data::{AggRefNode, AggregationsSegmentCtx}; +use crate::aggregation::segment_agg_result::SegmentAggregationCollector; +use crate::aggregation::*; + +/// Log2 of the number of registers for the HLL sketch. +/// 2^11 = 2048 registers, giving ~2.3% relative error and ~1KB per sketch (Hll4). +const LG_K: u8 = 11; + +/// # Cardinality +/// +/// The cardinality aggregation allows for computing an estimate +/// of the number of different values in a data set based on the +/// Apache DataSketches HyperLogLog algorithm. This is particularly useful for +/// understanding the uniqueness of values in a large dataset where counting +/// each unique value individually would be computationally expensive. +/// +/// For example, you might use a cardinality aggregation to estimate the number +/// of unique visitors to a website by aggregating on a field that contains +/// user IDs or session IDs. +/// +/// To use the cardinality aggregation, you'll need to provide a field to +/// aggregate on. The following example demonstrates a request for the cardinality +/// of the "user_id" field: +/// +/// ```JSON +/// { +/// "cardinality": { +/// "field": "user_id" +/// } +/// } +/// ``` +/// +/// This request will return an estimate of the number of unique values in the +/// "user_id" field. +/// +/// ## Missing Values +/// +/// The `missing` parameter defines how documents that are missing a value should be treated. +/// By default, documents without a value for the specified field are ignored. However, you can +/// specify a default value for these documents using the `missing` parameter. This can be useful +/// when you want to include documents with missing values in the aggregation. +/// +/// For example, the following request treats documents with missing values in the "user_id" +/// field as if they had a value of "unknown": +/// +/// ```JSON +/// { +/// "cardinality": { +/// "field": "user_id", +/// "missing": "unknown" +/// } +/// } +/// ``` +/// +/// # Estimation Accuracy +/// +/// The cardinality aggregation provides an approximate count, which is usually +/// accurate within a small error range. This trade-off allows for efficient +/// computation even on very large datasets. +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct CardinalityAggregationReq { + /// The field name to compute the percentiles on. + pub field: String, + /// The missing parameter defines how documents that are missing a value should be treated. + /// By default they will be ignored but it is also possible to treat them as if they had a + /// value. Examples in JSON format: + /// { "field": "my_numbers", "missing": "10.0" } + #[serde(skip_serializing_if = "Option::is_none", default)] + pub missing: Option, +} + +/// Contains all information required by the segment cardinality collectors to perform the +/// cardinality aggregation on a segment. +pub(crate) struct CardinalityAggReqData { + /// The column accessor to access the fast field values. + pub(crate) accessor: Column, + /// The column_type of the field. + pub(crate) column_type: ColumnType, + /// The string dictionary column if the field is of type string. + pub(crate) str_dict_column: Option, + /// The missing value normalized to the internal u64 representation of the field type. + pub(crate) missing_value_for_accessor: Option, + /// The name of the aggregation. + pub(crate) name: String, + /// The aggregation request. + pub(crate) req: CardinalityAggregationReq, +} + +impl CardinalityAggReqData { + /// Estimate the memory consumption of this struct in bytes. + pub fn get_memory_consumption(&self) -> usize { + std::mem::size_of::() + } +} + +impl CardinalityAggregationReq { + /// Creates a new [`CardinalityAggregationReq`] instance from a field name. + pub fn from_field_name(field_name: String) -> Self { + Self { + field: field_name, + missing: None, + } + } + /// Returns the field name the aggregation is computed on. + pub fn field_name(&self) -> &str { + &self.field + } +} + +#[derive(Clone, Debug)] +/// The cardinality collector used during segment collection and for merging results. +/// Uses Apache DataSketches HLL (lg_k=11, Hll4) for compact binary serialization +/// and cross-language compatibility (e.g. Java `datasketches` library). +pub struct CardinalityCollector { + sketch: HllSketch, + /// Salt derived from `ColumnType`, used to differentiate values of different column types + /// that map to the same u64 (e.g. bool `false` = 0 vs i64 `0`). + /// Not serialized — only needed during insertion, not after sketch registers are populated. + salt: u8, +} + +impl Default for CardinalityCollector { + fn default() -> Self { + Self::new(0) + } +} + +impl PartialEq for CardinalityCollector { + fn eq(&self, _other: &Self) -> bool { + false + } +} + +impl Serialize for CardinalityCollector { + fn serialize(&self, serializer: S) -> Result { + let bytes = self.sketch.serialize(); + serializer.serialize_bytes(&bytes) + } +} + +impl<'de> Deserialize<'de> for CardinalityCollector { + fn deserialize>(deserializer: D) -> Result { + let bytes: Vec = Deserialize::deserialize(deserializer)?; + let sketch = HllSketch::deserialize(&bytes).map_err(serde::de::Error::custom)?; + Ok(Self { sketch, salt: 0 }) + } +} + +impl CardinalityCollector { + fn new(salt: u8) -> Self { + Self { + sketch: HllSketch::new(LG_K, HllType::Hll8), + salt, + } + } + + /// Insert a value into the HLL sketch, salted by the column type. + /// The salt ensures that identical u64 values from different column types + /// (e.g. bool `false` vs i64 `0`) are counted as distinct. + fn insert(&mut self, value: impl Hash) { + self.sketch.update((self.salt, value)); + } + + fn insert_coupon(&mut self, coupon: Coupon) { + self.sketch.update_with_coupon(coupon); + } + + /// Compute the final cardinality estimate. + pub fn finalize(self) -> Option { + Some(self.sketch.estimate().trunc()) + } + + /// Serialize the HLL sketch to its compact binary representation. + /// The format is cross-language compatible with Apache DataSketches (Java, C++, Python). + pub fn to_sketch_bytes(&self) -> Vec { + self.sketch.serialize() + } + + pub(crate) fn merge_fruits(&mut self, right: CardinalityCollector) -> crate::Result<()> { + let mut union = HllUnion::new(LG_K); + union.update(&self.sketch); + union.update(&right.sketch); + self.sketch = union.to_sketch(HllType::Hll8); + Ok(()) + } +} + +/// Builds the segment collector for a cardinality aggregation. +/// +/// str and non-str columns use two entirely different collectors: str +/// accumulates term ordinals and resolves them into HLL coupons at +/// finalization, non-str feeds the HLL sketch directly. Both produce the same +/// [`IntermediateMetricResult::Cardinality`], so they merge uniformly. +pub(crate) fn build_segment_cardinality_collector( + req: &mut AggregationsSegmentCtx, + node: &AggRefNode, +) -> crate::Result> { + let req_data = req.get_cardinality_req_data(node.idx_in_req_data); + if req_data.column_type != ColumnType::Str { + return Ok(Box::new(SegmentNumericCardinalityCollector::from_req( + req_data.column_type, + node.idx_in_req_data, + req_data.accessor.clone(), + req_data.missing_value_for_accessor, + )?)); + } + // For str columns, we need to collect the set of term ordinals encounterred. + // We choose a different representation depending on the number of maximum + // number of terms. + // * small (< BITSET_MAX_TERM_ORD): `BitSet`, pre-allocated. + // * large: `TermOrdSet` (sparse HashSet that promotes to a paged bitset). + let max_term_ord_inclusive = req_data.accessor.max_value(); + if max_term_ord_inclusive < BITSET_MAX_TERM_ORD { + Ok(Box::new( + SegmentStrCardinalityCollector::::from_req( + node.idx_in_req_data, + req_data.accessor.clone(), + req_data.missing_value_for_accessor, + max_term_ord_inclusive, + ), + )) + } else { + Ok(Box::new( + SegmentStrCardinalityCollector::::from_req( + node.idx_in_req_data, + req_data.accessor.clone(), + req_data.missing_value_for_accessor, + max_term_ord_inclusive, + ), + )) + } +} + +#[cfg(test)] +mod tests { + use columnar::MonotonicallyMappableToU64; + + use crate::aggregation::agg_req::Aggregations; + use crate::aggregation::tests::{exec_request, get_test_index_from_terms}; + use crate::schema::{Schema, FAST, STRING}; + use crate::Index; + + #[test] + fn cardinality_aggregation_test_empty_index() -> crate::Result<()> { + let values = vec![]; + let index = get_test_index_from_terms(false, &values)?; + let agg_req: Aggregations = serde_json::from_value(json!({ + "cardinality": { + "cardinality": { + "field": "string_id", + } + }, + })) + .unwrap(); + + let res = exec_request(agg_req, &index)?; + assert_eq!(res["cardinality"]["value"], 0.0); + + Ok(()) + } + + #[test] + fn cardinality_aggregation_test_single_segment() -> crate::Result<()> { + cardinality_aggregation_test_merge_segment(true) + } + #[test] + fn cardinality_aggregation_test() -> crate::Result<()> { + cardinality_aggregation_test_merge_segment(false) + } + fn cardinality_aggregation_test_merge_segment(merge_segments: bool) -> crate::Result<()> { + let segment_and_terms = vec![ + vec!["terma"], + vec!["termb"], + vec!["termc"], + vec!["terma"], + vec!["terma"], + vec!["terma"], + vec!["termb"], + vec!["terma"], + ]; + let index = get_test_index_from_terms(merge_segments, &segment_and_terms)?; + let agg_req: Aggregations = serde_json::from_value(json!({ + "cardinality": { + "cardinality": { + "field": "string_id", + } + }, + })) + .unwrap(); + + let res = exec_request(agg_req, &index)?; + assert_eq!(res["cardinality"]["value"], 3.0); + + Ok(()) + } + + #[test] + fn cardinality_aggregation_bytes_excluded_from_accessors() -> crate::Result<()> { + // `Bytes` columns are opened as raw per-segment dictionary ordinals (like `Str`), but + // unlike `Str`, cardinality has no dictionary-resolution path for them: it would hash + // the raw ordinal directly, which are segment dependant. Ignore bytes values instead of + // counting them wrong. + let mut schema_builder = Schema::builder(); + let field = schema_builder.add_bytes_field("raw", FAST); + let index = Index::create_in_ram(schema_builder.build()); + { + let mut writer = index.writer_for_tests()?; + writer.add_document(doc!(field => vec![1u8]))?; + writer.add_document(doc!(field => vec![2u8]))?; + writer.commit()?; + writer.add_document(doc!(field => vec![3u8]))?; + writer.add_document(doc!(field => vec![4u8]))?; + writer.commit()?; + } + + let agg_req: Aggregations = serde_json::from_value(json!({ + "cardinality": { + "cardinality": { + "field": "raw" + }, + } + })) + .unwrap(); + + let res = exec_request(agg_req, &index)?; + assert_eq!(res["cardinality"]["value"], 0.0); + + Ok(()) + } + + #[test] + fn cardinality_aggregation_json() -> crate::Result<()> { + let mut schema_builder = Schema::builder(); + let field = schema_builder.add_json_field("json", FAST); + let index = Index::create_in_ram(schema_builder.build()); + { + let mut writer = index.writer_for_tests()?; + writer.add_document(doc!(field => json!({"value": false})))?; + writer.add_document(doc!(field => json!({"value": true})))?; + writer.add_document(doc!(field => json!({"value": i64::from_u64(0u64)})))?; + writer.add_document(doc!(field => json!({"value": i64::from_u64(1u64)})))?; + writer.commit()?; + } + + let agg_req: Aggregations = serde_json::from_value(json!({ + "cardinality": { + "cardinality": { + "field": "json.value" + }, + } + })) + .unwrap(); + + let res = exec_request(agg_req, &index)?; + assert_eq!(res["cardinality"]["value"], 4.0); + + Ok(()) + } + + /// A JSON path that resolves to both a Str column and a numeric column + /// produces two collector instances per segment — one with `Str` buckets + /// and one with `Numeric` buckets. Their `IntermediateMetricResult`s must + /// merge into the union cardinality. + #[test] + fn cardinality_aggregation_json_str_and_numeric() -> crate::Result<()> { + let mut schema_builder = Schema::builder(); + let field = schema_builder.add_json_field("json", FAST); + let index = Index::create_in_ram(schema_builder.build()); + { + let mut writer = index.writer_for_tests()?; + writer.add_document(doc!(field => json!({"value": "hello"})))?; + writer.add_document(doc!(field => json!({"value": "world"})))?; + writer.add_document(doc!(field => json!({"value": "hello"})))?; // dup str + writer.add_document(doc!(field => json!({"value": i64::from_u64(7u64)})))?; + writer.add_document(doc!(field => json!({"value": i64::from_u64(42u64)})))?; + writer.add_document(doc!(field => json!({"value": i64::from_u64(7u64)})))?; // dup num + writer.commit()?; + } + + let agg_req: Aggregations = serde_json::from_value(json!({ + "cardinality": { + "cardinality": { + "field": "json.value" + }, + } + })) + .unwrap(); + + let res = exec_request(agg_req, &index)?; + // 4 distinct values: "hello", "world", 7, 42. + assert_eq!(res["cardinality"]["value"], 4.0); + + Ok(()) + } + + #[test] + fn cardinality_collector_serde_roundtrip() { + use super::CardinalityCollector; + + let mut collector = CardinalityCollector::default(); + collector.insert("hello"); + collector.insert("world"); + collector.insert("hello"); // duplicate + + let serialized = serde_json::to_vec(&collector).unwrap(); + let deserialized: CardinalityCollector = serde_json::from_slice(&serialized).unwrap(); + + let original_estimate = collector.finalize().unwrap(); + let roundtrip_estimate = deserialized.finalize().unwrap(); + assert_eq!(original_estimate, roundtrip_estimate); + assert_eq!(original_estimate, 2.0); + } + + #[test] + fn cardinality_collector_merge() { + use super::CardinalityCollector; + + let mut left = CardinalityCollector::default(); + left.insert("a"); + left.insert("b"); + + let mut right = CardinalityCollector::default(); + right.insert("b"); + right.insert("c"); + + left.merge_fruits(right).unwrap(); + let estimate = left.finalize().unwrap(); + assert_eq!(estimate, 3.0); + } + + /// Verifies that merging two small sketches (both in List/Set coupon mode) + /// produces an exact result — i.e. the HllUnion does not unnecessarily + /// promote to the full HLL array when the combined cardinality is small. + #[test] + fn cardinality_collector_merge_stays_exact_for_small_sets() { + use super::CardinalityCollector; + + let mut left = CardinalityCollector::default(); + for i in 0u64..50 { + left.insert(i); + } + + let mut right = CardinalityCollector::default(); + for i in 30u64..100 { + right.insert(i); + } + + left.merge_fruits(right).unwrap(); + let estimate = left.finalize().unwrap(); + // 100 distinct values (0..100). Both sketches are in Set mode (< 192 coupons), + // so the union should stay in coupon mode and give an exact count. + assert_eq!(estimate, 100.0); + } + + #[test] + fn cardinality_collector_serialize_deserialize_binary() { + use datasketches::hll::HllSketch; + + use super::CardinalityCollector; + + let mut collector = CardinalityCollector::default(); + collector.insert("apple"); + collector.insert("banana"); + collector.insert("cherry"); + + let bytes = collector.to_sketch_bytes(); + let deserialized = HllSketch::deserialize(&bytes).unwrap(); + assert!((deserialized.estimate() - 3.0).abs() < 0.01); + } + + /// Tests that the `missing` parameter correctly counts a single empty document + /// for both u64 and str columns. + #[test] + fn cardinality_aggregation_missing_value_single_empty_doc() { + let mut schema_builder = Schema::builder(); + let id_field = schema_builder.add_u64_field("id", FAST); + let name_field = schema_builder.add_text_field("name", STRING | FAST); + let index = Index::create_in_ram(schema_builder.build()); + let mut writer = index.writer_for_tests().unwrap(); + writer + .add_document(doc!(id_field=>1u64,name_field=>"some_name")) + .unwrap(); + writer.add_document(doc!()).unwrap(); + writer.commit().unwrap(); + + { + // int colum with missing value non redundant + let agg_req: Aggregations = serde_json::from_value(json!({ + "cardinality": { + "cardinality": { + "field": "id", + "missing": 42u64 + }, + } + })) + .unwrap(); + let res = exec_request(agg_req, &index).unwrap(); + assert_eq!(res["cardinality"]["value"], 2.0); + } + + { + // int colum with missing value redundant + let agg_req: Aggregations = serde_json::from_value(json!({ + "cardinality": { + "cardinality": { + "field": "id", + "missing": 1u64 + }, + } + })) + .unwrap(); + let res = exec_request(agg_req, &index).unwrap(); + assert_eq!(res["cardinality"]["value"], 1.0); + } + + { + // str colum with missing value non redundant + // With more than one segment, this is not well handled. + let agg_req: Aggregations = serde_json::from_value(json!({ + "cardinality": { + "cardinality": { + "field": "name", + "missing": "other_name" + }, + } + })) + .unwrap(); + let res = exec_request(agg_req, &index).unwrap(); + assert_eq!(res["cardinality"]["value"], 2.0); + } + + { + // str colum with missing value redundant + let agg_req: Aggregations = serde_json::from_value(json!({ + "cardinality": { + "cardinality": { + "field": "name", + "missing": "some_name" + }, + } + })) + .unwrap(); + let res = exec_request(agg_req, &index).unwrap(); + assert_eq!(res["cardinality"]["value"], 1.0); + } + + { + // str column with missing value with a number type. + let agg_req: Aggregations = serde_json::from_value(json!({ + "cardinality": { + "cardinality": { + "field": "name", + "missing": 3, + }, + } + })) + .unwrap(); + let res = exec_request(agg_req, &index).unwrap(); + assert_eq!(res["cardinality"]["value"], 2.0); + } + } + + #[test] + fn cardinality_collector_salt_differentiates_types() { + use super::CardinalityCollector; + + // Without salt, same u64 value from different column types would collide + let mut collector_bool = CardinalityCollector::new(5); // e.g. ColumnType::Bool + collector_bool.insert(0u64); // false + collector_bool.insert(1u64); // true + + let mut collector_i64 = CardinalityCollector::new(2); // e.g. ColumnType::I64 + collector_i64.insert(0u64); + collector_i64.insert(1u64); + + // Merge them + collector_bool.merge_fruits(collector_i64).unwrap(); + let estimate = collector_bool.finalize().unwrap(); + // Should be 4 because salt makes (5, 0) != (2, 0) and (5, 1) != (2, 1) + assert_eq!(estimate, 4.0); + } +} diff --git a/src/aggregation/metric/cardinality/numeric_collector.rs b/src/aggregation/metric/cardinality/numeric_collector.rs new file mode 100644 index 000000000..072f922a6 --- /dev/null +++ b/src/aggregation/metric/cardinality/numeric_collector.rs @@ -0,0 +1,255 @@ +//! Segment collector for `cardinality` over any non-str column +//! (numeric, bool, date, IpAddr). +//! +//! Unlike the str case there is no dictionary to resolve, so values go +//! straight into a per-bucket HLL sketch during collection. The produced +//! [`CardinalityCollector`] is the same type the str collector produces, so +//! intermediate merging stays in the parent module. + +use std::fmt::Debug; +use std::sync::Arc; + +use columnar::column_values::CompactSpaceU64Accessor; +use columnar::{Column, ColumnType}; + +use super::CardinalityCollector; +use crate::aggregation::agg_data::AggregationsSegmentCtx; +use crate::aggregation::intermediate_agg_result::{ + IntermediateAggregationResult, IntermediateAggregationResults, IntermediateMetricResult, +}; +use crate::aggregation::segment_agg_result::SegmentAggregationCollector; +use crate::aggregation::*; +use crate::TantivyError; + +/// Segment collector for `cardinality` over any non-str column +/// (numeric, bool, date, IpAddr). +/// +/// Hidden contract: `column_type` must not be `ColumnType::Str`. Values are +/// inserted into the HLL sketch during collection, so the sketch of a bucket +/// is already complete when the bucket is finalized. +pub(crate) struct SegmentNumericCardinalityCollector { + /// Buckets are Some(_) until they get consumed by + /// `add_intermediate_aggregation_result`. + buckets: Vec>, + accessor_idx: usize, + /// The column accessor to access the fast field values. + accessor: Column, + /// The column_type of the field. + column_type: ColumnType, + /// Set iff `column_type == ColumnType::IpAddr`. Resolved once at + /// construction: the raw column values are compact-space codes that must + /// be expanded to their u128 ip representation before hashing. + compact_space_accessor: Option>, + /// The missing value normalized to the internal u64 representation of the field type. + missing_value_for_accessor: Option, +} + +impl Debug for SegmentNumericCardinalityCollector { + fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result { + f.debug_struct("SegmentNumericCardinalityCollector") + .field("column_type", &self.column_type) + .field( + "missing_value_for_accessor", + &self.missing_value_for_accessor, + ) + .finish() + } +} + +impl SegmentNumericCardinalityCollector { + pub fn from_req( + column_type: ColumnType, + accessor_idx: usize, + accessor: Column, + missing_value_for_accessor: Option, + ) -> crate::Result { + assert_ne!(column_type, ColumnType::Str); + let compact_space_accessor = if column_type == ColumnType::IpAddr { + let compact_space_accessor = accessor + .values + .clone() + .downcast_arc::() + .map_err(|_| { + TantivyError::AggregationError( + crate::aggregation::AggregationError::InternalError( + "Type mismatch: Could not downcast to CompactSpaceU64Accessor" + .to_string(), + ), + ) + })?; + Some(compact_space_accessor) + } else { + None + }; + Ok(Self { + buckets: Vec::new(), + accessor_idx, + accessor, + column_type, + compact_space_accessor, + missing_value_for_accessor, + }) + } +} + +impl SegmentAggregationCollector for SegmentNumericCardinalityCollector { + fn add_intermediate_aggregation_result( + &mut self, + agg_data: &AggregationsSegmentCtx, + results: &mut IntermediateAggregationResults, + bucket_id: BucketId, + ) -> crate::Result<()> { + self.prepare_max_bucket(bucket_id, agg_data)?; + let name = agg_data + .get_cardinality_req_data(self.accessor_idx) + .name + .to_string(); + // take the bucket in buckets and replace it with a new empty one + let Some(cardinality) = self.buckets[bucket_id as usize].take() else { + return Err(crate::TantivyError::InternalError( + "the same bucket should not be finalized twice.".to_string(), + )); + }; + results.push( + name, + IntermediateAggregationResult::Metric(IntermediateMetricResult::Cardinality( + cardinality, + )), + )?; + Ok(()) + } + + fn collect( + &mut self, + parent_bucket_id: BucketId, + docs: &[crate::DocId], + agg_data: &mut AggregationsSegmentCtx, + ) -> crate::Result<()> { + agg_data.column_block_accessor.fetch_block_with_missing( + docs, + &self.accessor, + self.missing_value_for_accessor, + ); + let cardinality = self.buckets[parent_bucket_id as usize] + .as_mut() + .ok_or_else(|| { + crate::TantivyError::InternalError( + "collection should not happen after finalization".to_string(), + ) + })?; + let col_block_accessor = &agg_data.column_block_accessor; + if let Some(compact_space_accessor) = self.compact_space_accessor.as_ref() { + for val in col_block_accessor.iter_vals() { + let val: u128 = compact_space_accessor.compact_to_u128(val as u32); + cardinality.insert(val); + } + } else { + for val in col_block_accessor.iter_vals() { + cardinality.insert(val); + } + } + Ok(()) + } + + fn prepare_max_bucket( + &mut self, + max_bucket: BucketId, + _agg_data: &AggregationsSegmentCtx, + ) -> crate::Result<()> { + if max_bucket as usize >= self.buckets.len() { + let column_type = self.column_type; + self.buckets.resize_with(max_bucket as usize + 1, || { + Some(CardinalityCollector::new(column_type as u8)) + }); + } + Ok(()) + } + + fn compute_metric_value( + &self, + bucket_id: BucketId, + sub_agg_name: &str, + sub_agg_property: &str, + agg_data: &AggregationsSegmentCtx, + ) -> Option { + let req_data = &agg_data.get_cardinality_req_data(self.accessor_idx); + if req_data.name != sub_agg_name || !sub_agg_property.is_empty() { + return None; + } + let cardinality = self.buckets.get(bucket_id as usize)?.as_ref()?; + Some(cardinality.sketch.estimate().trunc()) + } +} + +#[cfg(test)] +mod tests { + use std::net::IpAddr; + use std::str::FromStr; + + use crate::aggregation::agg_req::Aggregations; + use crate::aggregation::tests::exec_request; + use crate::schema::{IntoIpv6Addr, Schema, FAST}; + use crate::Index; + + #[test] + fn cardinality_aggregation_u64() -> crate::Result<()> { + let mut schema_builder = Schema::builder(); + let id_field = schema_builder.add_u64_field("id", FAST); + let index = Index::create_in_ram(schema_builder.build()); + { + let mut writer = index.writer_for_tests()?; + writer.add_document(doc!(id_field => 1u64))?; + writer.add_document(doc!(id_field => 2u64))?; + writer.add_document(doc!(id_field => 3u64))?; + writer.add_document(doc!())?; + writer.commit()?; + } + + let agg_req: Aggregations = serde_json::from_value(json!({ + "cardinality": { + "cardinality": { + "field": "id", + "missing": 0u64 + }, + } + })) + .unwrap(); + + let res = exec_request(agg_req, &index)?; + assert_eq!(res["cardinality"]["value"], 4.0); + + Ok(()) + } + + #[test] + fn cardinality_aggregation_ip_addr() -> crate::Result<()> { + let mut schema_builder = Schema::builder(); + let field = schema_builder.add_ip_addr_field("ip_field", FAST); + let index = Index::create_in_ram(schema_builder.build()); + { + let mut writer = index.writer_for_tests()?; + // IpV6 loopback + writer.add_document(doc!(field=>IpAddr::from_str("::1").unwrap().into_ipv6_addr()))?; + writer.add_document(doc!(field=>IpAddr::from_str("::1").unwrap().into_ipv6_addr()))?; + // IpV4 + writer.add_document( + doc!(field=>IpAddr::from_str("127.0.0.1").unwrap().into_ipv6_addr()), + )?; + writer.commit()?; + } + + let agg_req: Aggregations = serde_json::from_value(json!({ + "cardinality": { + "cardinality": { + "field": "ip_field" + }, + } + })) + .unwrap(); + + let res = exec_request(agg_req, &index)?; + assert_eq!(res["cardinality"]["value"], 2.0); + + Ok(()) + } +} diff --git a/src/aggregation/metric/cardinality/str_collector.rs b/src/aggregation/metric/cardinality/str_collector.rs new file mode 100644 index 000000000..0ff5af74e --- /dev/null +++ b/src/aggregation/metric/cardinality/str_collector.rs @@ -0,0 +1,397 @@ +//! Segment collector for `cardinality` over a `ColumnType::Str` column. +//! +//! Strings are dictionary encoded, and resolving a term ordinal to its bytes +//! is the expensive part. So instead of hashing values during collection, the +//! collector accumulates *term ordinals* per bucket and, at finalization, +//! builds one shared term_ord -> coupon cache for every bucket at once. The +//! coupons are then appended into a [`CardinalityCollector`], which is the +//! same type the numeric collector produces — merging across segments and +//! across column kinds stays in the parent module. + +use std::fmt::Debug; +use std::io; + +use columnar::{Column, ColumnType, Dictionary}; +use datasketches::hll::Coupon; +use rustc_hash::{FxBuildHasher, FxHashMap, FxHashSet}; + +use super::term_ord_accumulator::TermOrdAccumulator; +use super::CardinalityCollector; +use crate::aggregation::agg_data::AggregationsSegmentCtx; +use crate::aggregation::intermediate_agg_result::{ + IntermediateAggregationResult, IntermediateAggregationResults, IntermediateMetricResult, +}; +use crate::aggregation::segment_agg_result::SegmentAggregationCollector; +use crate::aggregation::*; + +/// A CouponCache is here to cache the mapping term ordinal -> coupon (see above). +/// The idea is that we do not want to fetch terms associated to several term ordinals, +/// several times due to the fact that we have several buckets. +enum CouponCache { + Dense { + coupon_map: Vec, + missing_coupon_opt: Option, + }, + Sparse { + coupon_map: FxHashMap, + missing_coupon_opt: Option, + }, +} + +impl CouponCache { + fn new( + term_ords: Vec, + coupons: Vec, + missing_coupon_opt: Option, + ) -> CouponCache { + let num_terms = term_ords.len(); + assert_eq!(num_terms, coupons.len()); + if term_ords.is_empty() { + return CouponCache::Dense { + coupon_map: Vec::new(), + missing_coupon_opt, + }; + } + let highest_term_ord = term_ords.last().copied().unwrap_or(0u64); + // We prefer the dense implementation, if it is not too wasteful. + // There are two cases for which we can use it. + // 1- if the data is small. + // 2- if the data is not necessarily small, but due to a high occupancy ratio, the RAM usage + // is not that much bigger than if we had used a HashSet. (occupancy ratio + extra + // metadata ~ x2.25) + let should_use_dense = + highest_term_ord < 1_000_000u64 || highest_term_ord < num_terms as u64 * 3u64; + if should_use_dense { + // We don't really care about the value here. We will populate all the values we will + // read anyway. + let uninitialized_coupon = Coupon::from_hash(0); + let mut coupon_map: Vec = + vec![uninitialized_coupon; highest_term_ord as usize + 1]; + + for (term_ord, coupon) in term_ords.into_iter().zip(coupons) { + coupon_map[term_ord as usize] = coupon; + } + CouponCache::Dense { + coupon_map, + missing_coupon_opt, + } + } else { + let coupon_map: FxHashMap = term_ords.into_iter().zip(coupons).collect(); + CouponCache::Sparse { + coupon_map, + missing_coupon_opt, + } + } + } +} + +/// Segment collector for `cardinality` over a `ColumnType::Str` column. +/// +/// Hidden contract: the column passed at construction must be a str column +/// whose values are term ordinals of the associated dictionary, and +/// `max_term_ord_inclusive` must be `accessor.max_value()`. The missing +/// sentinel `accessor.max_value() + 1` may additionally be inserted when +/// `missing_value_for_accessor` is set, hence accumulators size for +/// `max_term_ord_inclusive + 1`. +pub(crate) struct SegmentStrCardinalityCollector { + /// Buckets are Some(_) until they get consumed by + /// `add_intermediate_aggregation_result`. + buckets: Vec>, + accessor_idx: usize, + /// The column accessor to access the fast field values (term ordinals). + accessor: Column, + /// The missing value normalized to the internal u64 representation of the field type. + missing_value_for_accessor: Option, + /// Lazily built at finalization time, shared by every bucket. + coupon_cache: Option, + /// Largest term_ord that may be inserted into a bucket, i.e. + /// `accessor.max_value()`. + max_term_ord_inclusive: u64, +} + +impl Debug for SegmentStrCardinalityCollector { + fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result { + f.debug_struct("SegmentStrCardinalityCollector") + .field("num_buckets", &self.buckets.len()) + .field( + "missing_value_for_accessor", + &self.missing_value_for_accessor, + ) + .finish() + } +} + +/// Builds a coupon cache from the given buckets, dictionary, and optional missing value. +/// Returns a mapping from term_ord to the hash (coupon) of the associated term. +fn build_coupon_cache( + buckets: &[Option], + dictionary: &Dictionary, + missing_value_opt: Option<&Key>, +) -> io::Result { + // Pass 1 computes the capacity hint, pass 2 inserts. + let mut max_bucket_len = 0usize; + for bucket in buckets.iter().flatten() { + max_bucket_len = max_bucket_len.max(bucket.len()); + } + let mut term_ords_set = FxHashSet::with_capacity_and_hasher(max_bucket_len * 2, FxBuildHasher); + for bucket in buckets.iter().flatten() { + term_ords_set.extend(bucket.iter_ords()); + } + let mut term_ords: Vec = term_ords_set.into_iter().collect(); + term_ords.sort_unstable(); + + term_ords.pop_if(|highest_term_ord| *highest_term_ord >= dictionary.num_terms() as u64); + + let mut coupons: Vec = Vec::with_capacity(term_ords.len()); + let all_term_ords_found: bool = + dictionary.sorted_ords_to_term_cb(&term_ords, |term_bytes| { + let coupon: Coupon = Coupon::from_hash(term_bytes); + coupons.push(coupon); + })?; + assert!(all_term_ords_found); + + // Regardless of whether or not there is effectively a missing value in one of the buckets, + // we populate the cache with the missing key too (if any). + let missing_coupon_opt: Option = missing_value_opt.map(|missing_key| { + if let Key::Str(missing_value_str) = missing_key { + Coupon::from_hash(missing_value_str.as_bytes()) + } else { + // See https://github.com/quickwit-oss/tantivy/issues/2891 + // A missing key with a type different from Str will not work as intended + // for the moment. + // + // Right now this is just a partial workaround. + Coupon::from_hash("__tantivy_missing_non_str__".as_bytes()) + } + }); + Ok(CouponCache::new(term_ords, coupons, missing_coupon_opt)) +} + +fn append_to_sketch( + term_ords: &impl TermOrdAccumulator, + coupon_cache: &CouponCache, + sketch: &mut CardinalityCollector, +) { + match coupon_cache { + CouponCache::Dense { + coupon_map, + missing_coupon_opt, + } => { + if let Some(missing_coupon) = missing_coupon_opt { + for term_ord in term_ords.iter_ords() { + let coupon: Coupon = coupon_map + .get(term_ord as usize) + .copied() + .unwrap_or(*missing_coupon); + sketch.insert_coupon(coupon); + } + } else { + for term_ord in term_ords.iter_ords() { + if let Some(coupon) = coupon_map.get(term_ord as usize).copied() { + sketch.insert_coupon(coupon); + } + } + } + } + CouponCache::Sparse { + coupon_map, + missing_coupon_opt, + } => { + for term_ord in term_ords.iter_ords() { + if let Some(coupon) = coupon_map.get(&term_ord).copied().or(*missing_coupon_opt) { + sketch.insert_coupon(coupon); + } + } + } + } +} + +impl SegmentStrCardinalityCollector { + pub fn from_req( + accessor_idx: usize, + accessor: Column, + missing_value_for_accessor: Option, + max_term_ord_inclusive: u64, + ) -> Self { + Self { + buckets: Vec::new(), + accessor_idx, + accessor, + missing_value_for_accessor, + coupon_cache: None, + max_term_ord_inclusive, + } + } +} + +impl SegmentAggregationCollector + for SegmentStrCardinalityCollector +{ + fn add_intermediate_aggregation_result( + &mut self, + agg_data: &AggregationsSegmentCtx, + results: &mut IntermediateAggregationResults, + bucket_id: BucketId, + ) -> crate::Result<()> { + self.prepare_max_bucket(bucket_id, agg_data)?; + let req_data = &agg_data.get_cardinality_req_data(self.accessor_idx); + let Some(str_dict_column) = &req_data.str_dict_column else { + return Err(crate::TantivyError::InternalError( + "a str cardinality collector requires a str dictionary column".to_string(), + )); + }; + // Strings are dictionary encoded. Fetching the terms associated to strings + // is expensive. For this reason, we do that once for all buckets and cache the results + // here. + // + // The cache maps a term_ord to the hash of the associated term. The missing value + // sentinel will be associated to the hash of the missing value if any. + if self.coupon_cache.is_none() { + self.coupon_cache = Some(build_coupon_cache( + &self.buckets, + str_dict_column.dictionary(), + req_data.req.missing.as_ref(), + )?); + } + let name = req_data.name.to_string(); + // take the bucket in buckets and replace it with a new empty one + let Some(term_ords) = self.buckets[bucket_id as usize].take() else { + return Err(crate::TantivyError::InternalError( + "the same bucket should not be finalized twice.".to_string(), + )); + }; + let mut cardinality = CardinalityCollector::new(ColumnType::Str as u8); + if let Some(coupon_cache) = self.coupon_cache.as_ref() { + append_to_sketch(&term_ords, coupon_cache, &mut cardinality); + } + results.push( + name, + IntermediateAggregationResult::Metric(IntermediateMetricResult::Cardinality( + cardinality, + )), + )?; + + Ok(()) + } + + fn collect( + &mut self, + parent_bucket_id: BucketId, + docs: &[crate::DocId], + agg_data: &mut AggregationsSegmentCtx, + ) -> crate::Result<()> { + agg_data.column_block_accessor.fetch_block_with_missing( + docs, + &self.accessor, + self.missing_value_for_accessor, + ); + let Some(term_ords) = self.buckets[parent_bucket_id as usize].as_mut() else { + return Err(crate::TantivyError::InternalError( + "collection should not happen after finalization".to_string(), + )); + }; + // Promotion check runs on the pre-block state: the first call + // sees an empty set (no-op), and the last block of inserts + // doesn't trigger a promotion of a set we won't grow further. + // The trait dispatches once per block (via `extend_from_iter`) + // for adaptive variants and inlines to a tight loop for the + // BitSet path. + term_ords.maybe_compact(); + term_ords.extend_from_iter(agg_data.column_block_accessor.iter_vals()); + Ok(()) + } + + fn prepare_max_bucket( + &mut self, + max_bucket: BucketId, + _agg_data: &AggregationsSegmentCtx, + ) -> crate::Result<()> { + if max_bucket as usize >= self.buckets.len() { + let max_term_ord_inclusive = self.max_term_ord_inclusive; + self.buckets.resize_with(max_bucket as usize + 1, || { + Some(S::new(max_term_ord_inclusive)) + }); + } + Ok(()) + } + + fn compute_metric_value( + &self, + bucket_id: BucketId, + sub_agg_name: &str, + sub_agg_property: &str, + agg_data: &AggregationsSegmentCtx, + ) -> Option { + let req_data = &agg_data.get_cardinality_req_data(self.accessor_idx); + if req_data.name != sub_agg_name || !sub_agg_property.is_empty() { + return None; + } + // The sketch isn't built until finalization; the term_ord set's len is + // the exact distinct count. + let term_ords = self.buckets.get(bucket_id as usize)?.as_ref()?; + Some(term_ords.len() as f64) + } +} + +#[cfg(test)] +mod tests { + use crate::aggregation::agg_req::Aggregations; + use crate::aggregation::tests::{exec_request, get_test_index_from_terms}; + use crate::schema::{Schema, FAST, STRING}; + use crate::Index; + + /// Build a single-segment string-cardinality index with 32 unique terms. + /// `column.max_value() = 31` is well below `BITSET_MAX_TERM_ORD`, + /// so the bucket exercises the `BitSet` path end to end. + #[test] + fn cardinality_aggregation_test_str_bitset() -> crate::Result<()> { + let terms: Vec = (0..32).map(|i| format!("term_{i}")).collect(); + let term_refs: Vec> = terms.iter().map(|t| vec![t.as_str()]).collect::>(); + // single segment so we have a single dictionary of 32 terms. + let index = get_test_index_from_terms(true, &term_refs)?; + + let agg_req: Aggregations = serde_json::from_value(json!({ + "cardinality": { + "cardinality": { "field": "string_id" } + }, + })) + .unwrap(); + + let res = exec_request(agg_req, &index)?; + assert_eq!(res["cardinality"]["value"], 32.0); + Ok(()) + } + + /// `BitSet` path with a `missing` parameter: the column-level missing + /// sentinel (`column.max_value() + 1`) flows into the bitset, the + /// dict lookup filter at finalization drops it, and the missing + /// coupon is applied separately. + #[test] + fn cardinality_aggregation_test_str_bitset_with_missing() { + let mut schema_builder = Schema::builder(); + let name_field = schema_builder.add_text_field("name", STRING | FAST); + let index = Index::create_in_ram(schema_builder.build()); + let mut writer = index.writer_for_tests().unwrap(); + for i in 0..16 { + let term = format!("t{i:02}"); + writer.add_document(doc!(name_field => term)).unwrap(); + } + // One empty doc, exercising the missing sentinel. + writer.add_document(doc!()).unwrap(); + writer.commit().unwrap(); + + let agg_req: Aggregations = serde_json::from_value(json!({ + "cardinality": { + "cardinality": { + "field": "name", + "missing": "MISSING_SENTINEL_KEY", + } + }, + })) + .unwrap(); + + let res = exec_request(agg_req, &index).unwrap(); + // 16 distinct real terms + 1 distinct "missing" value = 17. + assert_eq!(res["cardinality"]["value"], 17.0); + } +} diff --git a/src/aggregation/metric/cardinality/term_ord_accumulator.rs b/src/aggregation/metric/cardinality/term_ord_accumulator.rs new file mode 100644 index 000000000..f35d6148c --- /dev/null +++ b/src/aggregation/metric/cardinality/term_ord_accumulator.rs @@ -0,0 +1,343 @@ +//! Per-bucket term-ordinal accumulators used by the str cardinality +//! collector. +//! +//! A str cardinality bucket needs a *set of term ordinals*, and the right +//! representation depends on how large the segment's dictionary slice is: +//! +//! * [`BitSet`] (from `common`): used when `column.max_value()` is small (< +//! [`BITSET_MAX_TERM_ORD`]). Pre-allocated, no promotion machinery. +//! * [`TermOrdSet`]: adaptive. Starts as an `FxHashSet` (cheap when few ords are seen) and +//! promotes to a [`PagedBitset`] once occupancy crosses the density threshold. +//! +//! Both are exposed through the [`TermOrdAccumulator`] trait so that +//! `SegmentStrCardinalityCollector` can be generic over the choice and the hot +//! `collect()` loop monomorphizes to a direct call (no enum dispatch per +//! insert). + +use common::{BitSet, TinySet}; +use rustc_hash::FxHashSet; + +/// Promote FxHashSet -> PagedBitset at ~3% density (`len * 32 > +/// dict_num_terms`). Past this point the bitset (~`dict_num_terms / 7.5` +/// bytes) is smaller than the hashset (~10 B/entry minimum) and avoids +/// the per-insert hash. +const PROMOTION_RATIO: u64 = 32; + +// ================================================================= +// PagedBitset: a sparse bitset indexed by term_ord. +// +// Used as the dense alternative to FxHashSet once a string +// cardinality bucket has accumulated enough unique term ordinals. +// Memory is bounded to (touched pages) * (page bytes), not +// (max_term_ord / 8). +// +// Page geometry mirrors `PagedTermMap` in `term_agg.rs`: 1024 ords +// per page, lazy `Vec>>` directory. +// ================================================================= +const BITSET_PAGE_SHIFT: u32 = 10; +const BITSET_PAGE_BITS: u64 = 1u64 << BITSET_PAGE_SHIFT; // 1024 +const BITSET_PAGE_MASK: u64 = BITSET_PAGE_BITS - 1; +const BITSET_WORDS_PER_PAGE: usize = (BITSET_PAGE_BITS / 64) as usize; // 16 + +#[derive(Clone)] +struct PagedBitsetPage { + words: [TinySet; BITSET_WORDS_PER_PAGE], +} + +impl PagedBitsetPage { + fn new() -> Self { + Self { + words: [TinySet::empty(); BITSET_WORDS_PER_PAGE], + } + } +} + +pub(crate) struct PagedBitset { + pages: Vec>>, + /// Cached number of set bits, maintained on insert. + count: u64, +} + +impl PagedBitset { + /// Allocates a directory big enough to hold ords up to and including + /// `max_term_ord`. Pages are allocated lazily on first set. + fn with_max_term_ord(max_term_ord: u64) -> Self { + let max_page_idx = (max_term_ord >> BITSET_PAGE_SHIFT) as usize; + let num_pages = max_page_idx + 1; + Self { + pages: vec![None; num_pages], + count: 0, + } + } + + #[inline] + fn insert(&mut self, term_ord: u64) { + let page_idx = (term_ord >> BITSET_PAGE_SHIFT) as usize; + let intra = term_ord & BITSET_PAGE_MASK; + let word_idx = (intra >> 6) as usize; + let bit_idx = (intra & 63) as u32; + + let page = match &mut self.pages[page_idx] { + Some(p) => p, + None => { + self.pages[page_idx] = Some(Box::new(PagedBitsetPage::new())); + self.pages[page_idx].as_mut().unwrap() + } + }; + if page.words[word_idx].insert_mut(bit_idx) { + self.count += 1; + } + } + + /// Number of set bits. O(1). + #[inline] + fn len(&self) -> u64 { + self.count + } + + /// Iterate set ords in ascending order. + fn iter_sorted(&self) -> impl Iterator + '_ { + self.pages + .iter() + .enumerate() + .filter_map(|(page_idx, page_opt)| page_opt.as_ref().map(|p| (page_idx, p))) + .flat_map(|(page_idx, page)| { + let page_base_ord = (page_idx as u64) << BITSET_PAGE_SHIFT; + page.words + .iter() + .enumerate() + .flat_map(move |(word_idx, &word)| { + let word_base_ord = page_base_ord + (word_idx as u64) * 64; + word.into_iter() + .map(move |bit| word_base_ord + u64::from(bit)) + }) + }) + } +} + +/// Threshold below which we use `BitSet` instead of `TermOrdSet`. +/// +/// Both `BitSet` and `FxHashSet` have the same 32-byte struct, so the comparison is heap only: +/// * `BitSet` at T=256: 5 `TinySet` words covering 258 bits (with the missing-value sentinel) = +/// 40 bytes. +/// * `FxHashSet` after one insert: 4-bucket hashbrown table ≈ 56 bytes +pub(crate) const BITSET_MAX_TERM_ORD: u64 = 256; + +// ================================================================= +// TermOrdAccumulator: per-bucket abstraction over the entries set. +// +// Implementations: +// - `BitSet` (from `common`): used when `column.max_value()` is small (< BITSET_MAX_TERM_ORD). +// Pre-allocated, no promotion. +// - `TermOrdSet`: adaptive, starts as FxHashSet and promotes to a paged bitset when occupancy +// crosses the density threshold (only if promotion is enabled — typically gated on top-level +// aggregation). +// +// The trait lets `SegmentStrCardinalityCollector` be generic over the choice +// so the hot collect() loop monomorphizes to a direct call (no enum +// dispatch per insert). +// ================================================================= +pub(crate) trait TermOrdAccumulator: Sized { + /// Construct an empty accumulator. + /// `max_term_ord_inclusive` is the largest term_ord that may be + /// inserted (used to size pre-allocated bitsets and the dense bitset + /// on promotion). + fn new(max_term_ord_inclusive: u64) -> Self; + fn insert(&mut self, term_ord: u64); + fn extend_from_iter(&mut self, ords: impl IntoIterator); + /// Hook called once per ingested block. Adaptive impls use this to + /// decide on sparse->dense promotion. + fn maybe_compact(&mut self) {} + fn len(&self) -> usize; + fn iter_ords(&self) -> impl Iterator + '_; +} + +impl TermOrdAccumulator for BitSet { + #[inline] + fn new(max_term_ord_inclusive: u64) -> Self { + // `BitSet::with_max_value(M)` accepts ords in [0, M). + // We need ords up to and including `max_term_ord_inclusive`, plus + // the missing-value sentinel `column.max_value() + 1`. + BitSet::with_max_value((max_term_ord_inclusive + 2) as u32) + } + #[inline] + fn insert(&mut self, term_ord: u64) { + BitSet::insert(self, term_ord as u32); + } + #[inline] + fn len(&self) -> usize { + BitSet::len(self) + } + fn iter_ords(&self) -> impl Iterator + '_ { + // `BitSet` itself doesn't expose iteration, but + // `BitSet::tinyset(bucket)` does. Walk per-bucket and yield each + // set bit. The capacity is `max_value()`; iterating to + // `div_ceil(64)` covers every possible ord exactly once. + let num_buckets = self.max_value().div_ceil(64); + (0..num_buckets).flat_map(move |bucket| { + let chunk_base = u64::from(bucket) * 64; + self.tinyset(bucket) + .into_iter() + .map(move |bit| chunk_base + u64::from(bit)) + }) + } + #[inline(never)] //< required to not have a perf regression + fn extend_from_iter(&mut self, ords: impl IntoIterator) { + for ord in ords { + ::insert(self, ord); + } + } +} + +// TermOrdSet: adaptive sparse->dense accumulator. +// +// Starts as an HashSet (cheap when few ords are seen). When occupancy +// crosses `len * PROMOTION_RATIO > max_term_ord_inclusive`, drains into +// a `PagedBitset` and continues dense. +pub(crate) struct TermOrdSet { + inner: TermOrdSetInner, + /// Largest term_ord that may be inserted. Used for both sizing the + /// dense bitset on promotion and as the promotion-threshold reference. + max_term_ord_inclusive: u64, +} + +enum TermOrdSetInner { + Sparse(FxHashSet), + Dense(PagedBitset), +} + +impl TermOrdAccumulator for TermOrdSet { + fn new(max_term_ord_inclusive: u64) -> Self { + Self { + inner: TermOrdSetInner::Sparse(FxHashSet::default()), + max_term_ord_inclusive, + } + } + + #[inline] + fn insert(&mut self, term_ord: u64) { + match &mut self.inner { + TermOrdSetInner::Sparse(set) => { + set.insert(term_ord); + } + TermOrdSetInner::Dense(bitset) => bitset.insert(term_ord), + } + } + + fn extend_from_iter(&mut self, ords: impl IntoIterator) { + match &mut self.inner { + TermOrdSetInner::Sparse(set) => { + set.extend(ords); + } + TermOrdSetInner::Dense(bitset) => { + for ord in ords { + bitset.insert(ord); + } + } + } + } + + fn maybe_compact(&mut self) { + let TermOrdSetInner::Sparse(set) = &mut self.inner else { + return; + }; + if set.len() as u64 * PROMOTION_RATIO <= self.max_term_ord_inclusive { + return; + } + let mut bitset = PagedBitset::with_max_term_ord(self.max_term_ord_inclusive + 1); + let set = std::mem::take(set); + for ord in set { + bitset.insert(ord); + } + self.inner = TermOrdSetInner::Dense(bitset); + } + + fn len(&self) -> usize { + match &self.inner { + TermOrdSetInner::Sparse(set) => set.len(), + TermOrdSetInner::Dense(bitset) => bitset.len() as usize, + } + } + + fn iter_ords(&self) -> impl Iterator + '_ { + match &self.inner { + TermOrdSetInner::Sparse(set) => itertools::Either::Left(set.iter().copied()), + TermOrdSetInner::Dense(bitset) => itertools::Either::Right(bitset.iter_sorted()), + } + } +} + +#[cfg(test)] +mod tests { + use common::BitSet; + + use super::{PagedBitset, TermOrdAccumulator, TermOrdSet, PROMOTION_RATIO}; + + /// Unit-test the PagedBitset itself: cross-page inserts produce sorted + /// iteration, len() matches the inserted set, and duplicates are + /// idempotent. + #[test] + fn paged_bitset_basic() { + // Span several pages: BITSET_PAGE_BITS = 1024, so ords > 1024 land + // on the second page, > 2048 on the third, etc. + let ords = [0u64, 1, 63, 64, 1023, 1024, 1025, 4096, 4097, 9999, 10_000]; + let max_ord = *ords.iter().max().unwrap(); + let mut bitset = PagedBitset::with_max_term_ord(max_ord); + for &ord in &ords { + bitset.insert(ord); + // Idempotent: inserting again must not increase count. + bitset.insert(ord); + } + assert_eq!(bitset.len(), ords.len() as u64); + let collected: Vec = bitset.iter_sorted().collect(); + let mut expected: Vec = ords.to_vec(); + expected.sort_unstable(); + assert_eq!(collected, expected); + } + + /// Unit-test `TermOrdSet`: starts Sparse, promotes to Dense on + /// `maybe_compact` once the density threshold is crossed, and + /// `iter_ords()` yields the same set in either state. Ords spanning + /// multiple paged-bitset pages exercise the Dense iter ordering. + #[test] + fn term_ord_set_promotes_on_maybe_compact() { + // Pick max so promotion needs few inserts: len * RATIO > max with + // RATIO=32 and max=64 trips at len=3 (3*32=96 > 64). + let max_term_ord = 64u64; + let mut set = ::new(max_term_ord); + // Two inserts: should stay Sparse after maybe_compact (2 * RATIO = 64, not > 64). + set.insert(0); + set.insert(7); + set.maybe_compact(); + assert_eq!(set.len(), 2); + + // Third insert promotes on next maybe_compact. + set.insert(20); + assert_eq!(set.len(), 3); + // Sanity check: at len=3, 3 * PROMOTION_RATIO = 96 > 64. + assert!(3u64 * PROMOTION_RATIO > max_term_ord); + set.maybe_compact(); + + // Post-promotion: extending continues to work. + set.insert(15); + set.insert(15); // dup + assert_eq!(set.len(), 4); + + let mut collected: Vec = set.iter_ords().collect(); + collected.sort_unstable(); + assert_eq!(collected, vec![0, 7, 15, 20]); + } + + /// Unit-test the `BitSet` impl of `TermOrdAccumulator`: insert, + /// dedup, and iter_ords order. + #[test] + fn bitset_accumulator_basic() { + let mut set = ::new(255); + for ord in [0u64, 1, 63, 64, 65, 128, 200, 200, 0] { + ::insert(&mut set, ord); + } + assert_eq!(::len(&set), 7); + let collected: Vec = set.iter_ords().collect(); + assert_eq!(collected, vec![0, 1, 63, 64, 65, 128, 200]); + } +} From 0f512fc3e200eba3958b279b4eac1c0615235b62 Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Sat, 19 Sep 2026 16:20:25 +0200 Subject: [PATCH 21/49] Benchmark multi-terms aggregations across many segments Run the existing many-segment aggregation benchmarks at 100 and 1,000 segments with one million total documents. Add multi-terms and nested terms cases for status/Zipf and high-cardinality/Zipf combinations, including top-500 requests, to expose merge costs across many segments. The nested top-500 case limits outer buckets, while multi-terms limits tuples globally. Run both groups with: cargo bench --bench agg_bench -- _segments --- benches/agg_bench.rs | 35 +++++++++++++++++++++++++++++------ 1 file changed, 29 insertions(+), 6 deletions(-) diff --git a/benches/agg_bench.rs b/benches/agg_bench.rs index 0ffd4322c..4ebd7cb45 100644 --- a/benches/agg_bench.rs +++ b/benches/agg_bench.rs @@ -54,23 +54,46 @@ fn main() { bench_agg(&mut runner, &index, execute_filtered); } - bench_many_segments(); + for num_segments in [100, 1_000] { + bench_many_segments(num_segments); + } } -fn bench_many_segments() { +fn bench_many_segments(num_segments: usize) { let mut runner = BenchRunner::new(); runner.add_plugin(PeakMemAllocPlugin::new(GLOBAL)); - let index = get_test_index_bench_with_num_segments(Cardinality::Full, 100).unwrap(); + runner.config().set_num_iter_for_group(1); + let index = get_test_index_bench_with_num_segments(Cardinality::Full, num_segments).unwrap(); let mut group = runner.new_group(); - group.set_name("100_segments"); + group.set_name(format!("{num_segments}_segments")); + let mut multi_terms_top500 = multi_terms_many_and_zipf_1000(); + multi_terms_top500["mt"]["multi_terms"]["size"] = json!(500); + let mut nested_terms_top500 = nested_terms_many_and_zipf_1000(); + // This limits outer buckets, unlike the global tuple limit for multi_terms. + nested_terms_top500["my_texts"]["terms"]["size"] = json!(500); + let reader = index.reader().unwrap(); + let searcher = reader.searcher(); for (benchmark_name, agg_req) in [ ("terms_7", terms_on_field("text_few_terms_status")), ("terms_zipfs_1000", terms_on_field("text_1000_terms_zipf")), ("terms_150_000", terms_on_field("text_many_terms")), ("terms_all_unique", terms_on_field("text_all_unique_terms")), + benchmark_config!(nested_terms_status_and_zipf_1000), + benchmark_config!(multi_terms_status_and_zipf_1000), + benchmark_config!(nested_terms_many_and_zipf_1000), + benchmark_config!(multi_terms_many_and_zipf_1000), + ( + "nested_terms_many_and_zipf_1000_top500", + nested_terms_top500, + ), + ("multi_terms_many_and_zipf_1000_top500", multi_terms_top500), ] { - group.register_with_input(benchmark_name, &index, move |index| { - execute_agg(index, agg_req.clone()) + let searcher = searcher.clone(); + group.register_with_input(benchmark_name, &(), move |_| { + let agg_req: Aggregations = serde_json::from_value(agg_req.clone()).unwrap(); + let collector = get_collector(agg_req); + + black_box(searcher.search(&AllQuery, &collector).unwrap()); }); } group.run(); From 085a8327d32df2a7475ae5d349e4cd517c3536fe Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Sat, 19 Sep 2026 16:23:58 +0200 Subject: [PATCH 22/49] Merge incoming aggregation buckets into the accumulator Probe the accumulated map only for incoming keys rather than rehashing every accumulated key on each merge. This removes quadratic work when folding many disjoint multi-terms results. The same helper serves range and composite buckets; existing bucket values still merge left-to-right and the wire representation and pruning rules are unchanged. Cover overlapping/disjoint tuple keys, empty inputs, recursive range subaggregations, postcard round trips, error/count bookkeeping, and merging after pruning. Same benchmark and configuration as 0062d2f0d (Apple M4 Max, rustc 1.98.0): Median milliseconds for 10 / 100 / 1000 inputs: shared 8: 0.006365 / 0.0656 / 0.6449 disjoint 16: 0.0126 / 0.1486 / 1.6882 disjoint 160: 0.1336 / 1.5558 / 20.4019 At 1000 inputs this is 52.4x and 61.6x faster for the disjoint cases, with unchanged measured peak allocation. Shared-key control remains within approximately 2% of baseline. Validation: 308 aggregation tests pass with default features and 308 with quickwit; changed-file rustfmt and git diff checks pass. Clippy --lib --bench agg_bench passes with the pre-existing clippy::drop_non_drop warning allowed (unchanged drop(add_document) in the benchmark). --- src/aggregation/intermediate_agg_result.rs | 119 +++++++++++++++++++-- 1 file changed, 110 insertions(+), 9 deletions(-) diff --git a/src/aggregation/intermediate_agg_result.rs b/src/aggregation/intermediate_agg_result.rs index ceb149b1c..fd9877e33 100644 --- a/src/aggregation/intermediate_agg_result.rs +++ b/src/aggregation/intermediate_agg_result.rs @@ -1213,19 +1213,20 @@ trait MergeFruits { fn merge_fruits(&mut self, other: Self) -> crate::Result<()>; } -fn merge_maps( +fn merge_maps( entries_left: &mut FxHashMap, - mut entries_right: FxHashMap, + entries_right: FxHashMap, ) -> crate::Result<()> { - for (name, entry_left) in entries_left.iter_mut() { - if let Some(entry_right) = entries_right.remove(name) { - entry_left.merge_fruits(entry_right)?; + // Visit incoming entries, not the growing accumulator, so folding many results + // does not repeatedly hash all previously merged keys. + for (key, entry_right) in entries_right { + match entries_left.entry(key) { + Entry::Occupied(mut entry) => entry.get_mut().merge_fruits(entry_right)?, + Entry::Vacant(entry) => { + entry.insert(entry_right); + } } } - - for (key, res) in entries_right.into_iter() { - entries_left.entry(key).or_insert(res); - } Ok(()) } @@ -1698,6 +1699,106 @@ mod tests { assert_range_trees_eq(&tree_left, &tree_expected); } + #[test] + fn test_multi_terms_repeated_merge_and_prune() { + let key = |id: u64| { + vec![ + IntermediateKey::Str(format!("host-{}", id % 2)), + IntermediateKey::U64(id / 2), + ] + }; + let make_result = + |data: &[(u64, u64)], source: &str| IntermediateBucketResult::MultiTerms { + buckets: IntermediateMultiTermsBucketResult { + entries: data + .iter() + .map(|&(id, count)| { + ( + key(id), + IntermediateTermBucketEntry { + doc_count: count, + sub_aggregation: get_sub_test_tree(&[ + ("shared".to_string(), count), + (source.to_string(), count), + ]), + }, + ) + }) + .collect(), + sum_other_doc_count: 2, + doc_count_error_upper_bound: 1, + }, + }; + let mut merged = IntermediateBucketResult::MultiTerms { + buckets: Default::default(), + }; + merged + .merge_fruits(make_result(&[(0, 3), (1, 5)], "first")) + .unwrap(); + merged + .merge_fruits(make_result(&[(1, 7), (2, 11)], "second")) + .unwrap(); + // A distributed fold may resume after serializing an intermediate result. + merged = postcard::from_bytes(&postcard::to_allocvec(&merged).unwrap()).unwrap(); + merged.merge_fruits(make_result(&[], "empty")).unwrap(); + merged + .merge_fruits(make_result(&[(0, 13), (3, 17)], "third")) + .unwrap(); + + let IntermediateBucketResult::MultiTerms { buckets } = &mut merged else { + panic!("expected multi_terms"); + }; + assert_eq!(buckets.entries.len(), 4); + assert_eq!(buckets.sum_other_doc_count, 8); + assert_eq!(buckets.doc_count_error_upper_bound, 4); + for (id, count, sources) in [ + (0, 16, vec![("first", 3), ("third", 13)]), + (1, 12, vec![("first", 5), ("second", 7)]), + (2, 11, vec![("second", 11)]), + (3, 17, vec![("third", 17)]), + ] { + let entry = &buckets.entries[&key(id)]; + assert_eq!(entry.doc_count, count); + let mut expected = vec![("shared".to_string(), count)]; + expected.extend( + sources + .into_iter() + .map(|(name, count)| (name.to_string(), count)), + ); + assert_range_trees_eq(&entry.sub_aggregation, &get_sub_test_tree(&expected)); + } + + let req = serde_json::from_value(serde_json::json!({ + "terms": [{"field": "host"}, {"field": "path"}], + "size": 1, + "segment_size": 2 + })) + .unwrap(); + buckets + .prune_intermediate_results(&req, &Default::default(), PruneMode::Intermediate) + .unwrap(); + assert_eq!(buckets.entries.len(), 2); + assert_eq!(buckets.sum_other_doc_count, 31); // 8 + 12 + 11 + assert_eq!(buckets.doc_count_error_upper_bound, 16); // 4 + cutoff 12 + + // A previously pruned key can be inserted again in a subsequent fold. + merged + .merge_fruits(make_result(&[(2, 19)], "fourth")) + .unwrap(); + let IntermediateBucketResult::MultiTerms { buckets } = &mut merged else { + unreachable!(); + }; + assert_eq!(buckets.entries.len(), 3); + assert_eq!(buckets.entries[&key(2)].doc_count, 19); + buckets + .prune_intermediate_results(&req, &Default::default(), PruneMode::Final) + .unwrap(); + assert_eq!(buckets.entries.len(), 1); + assert_eq!(buckets.entries[&key(2)].doc_count, 19); + assert_eq!(buckets.sum_other_doc_count, 66); // 31 + 2 + 16 + 17 + assert_eq!(buckets.doc_count_error_upper_bound, 17); // 16 + 1, no final cutoff + } + #[test] fn test_prune_intermediate_results_finalizer_size() { use crate::aggregation::bucket::TermsAggregation; From 3655e6f8cb399d3092f2b899e0f389d9d1971509 Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Sun, 20 Sep 2026 19:51:17 +0200 Subject: [PATCH 23/49] Batch multi-terms dictionary lookups after pruning Resolve retained string ordinals in sorted batches per field instead of decoding dictionary blocks separately for each tuple component. Reuse decoding for repeated ordinals while preserving tuple order and existing numeric and missing-value handling. --- src/aggregation/bucket/multi_terms/mod.rs | 96 +++++++++++++++-------- 1 file changed, 63 insertions(+), 33 deletions(-) diff --git a/src/aggregation/bucket/multi_terms/mod.rs b/src/aggregation/bucket/multi_terms/mod.rs index 2cdfbb1f1..2db369d0c 100644 --- a/src/aggregation/bucket/multi_terms/mod.rs +++ b/src/aggregation/bucket/multi_terms/mod.rs @@ -250,11 +250,8 @@ trait MultiTermsPacking: Clone + Debug + 'static { fn push_full_values(&self, keys: &mut [Self::PackingType], field_idx: usize, values: I) where I: IntoIterator; - fn unpack( - &self, - key: &Self::PackingType, - req_data: &MultiTermsAggReqData, - ) -> crate::Result>; + /// Extract one field's raw value, allowing dictionary lookups to be batched by field. + fn unpack_value(&self, key: &Self::PackingType, field_idx: usize) -> u64; } #[derive(Clone, Debug)] @@ -286,18 +283,8 @@ impl MultiTermsPacking for U64ArrayKeyPacking { } } - fn unpack( - &self, - key: &Self::PackingType, - req_data: &MultiTermsAggReqData, - ) -> crate::Result> { - key.iter() - .zip(req_data.fields.iter()) - .zip(req_data.missing_accessors.iter()) - .map(|((value, field_acc), missing)| { - resolve_key_value(*value, field_acc, missing.as_ref()) - }) - .collect() + fn unpack_value(&self, key: &Self::PackingType, field_idx: usize) -> u64 { + key[field_idx] } } @@ -330,20 +317,10 @@ impl MultiTermsPacking for PackedU64KeyPacking { } } - fn unpack( - &self, - key: &Self::PackingType, - req_data: &MultiTermsAggReqData, - ) -> crate::Result> { - self.packs - .iter() - .zip(req_data.fields.iter()) - .zip(req_data.missing_accessors.iter()) - .map(|((pack, field), missing)| { - let offset = key.checked_shr(pack.shift).unwrap_or(0) & pack.mask; - resolve_key_value(offset + pack.min_value, field, missing.as_ref()) - }) - .collect() + fn unpack_value(&self, key: &Self::PackingType, field_idx: usize) -> u64 { + let pack = self.packs[field_idx]; + let offset = key.checked_shr(pack.shift).unwrap_or(0) & pack.mask; + offset + pack.min_value } } @@ -733,8 +710,8 @@ where let mut result_entries: FxHashMap, IntermediateTermBucketEntry> = FxHashMap::with_capacity_and_hasher(entries.len(), Default::default()); - for entry in entries { - let intermediate_key = packing.unpack(&entry.key, req_data)?; + let keys = resolve_bucket_keys(packing, &entries, req_data)?; + for (entry, intermediate_key) in entries.into_iter().zip(keys) { let mut sub_aggregation_res = IntermediateAggregationResults::default(); if let Some(sub_agg_collector) = sub_agg_collector.as_deref_mut() { sub_agg_collector.add_intermediate_aggregation_result( @@ -1095,6 +1072,59 @@ where })) } +/// Resolve only the retained candidates, batching string ordinals by field. Sorted lookups +/// decode each dictionary block once and reuse decoded terms for repeated ordinals. +fn resolve_bucket_keys( + packing: &P, + entries: &[MultiTermsBucketEntry], + req_data: &MultiTermsAggReqData, +) -> crate::Result>> { + let mut keys: Vec<_> = (0..entries.len()) + .map(|_| Vec::with_capacity(req_data.fields.len())) + .collect(); + for (field_idx, (field, missing)) in req_data + .fields + .iter() + .zip(&req_data.missing_accessors) + .enumerate() + { + let mut ords_and_positions = Vec::new(); + for (position, entry) in entries.iter().enumerate() { + let value = packing.unpack_value(&entry.key, field_idx); + if field.column_type == ColumnType::Str + && !missing.as_ref().is_some_and(|m| m.missing_value == value) + { + ords_and_positions.push((value, position)); + } else { + keys[position].push(resolve_key_value(value, field, missing.as_ref())?); + } + } + if ords_and_positions.is_empty() { + continue; + } + ords_and_positions.sort_unstable(); + let (ords, positions): (Vec<_>, Vec<_>) = ords_and_positions.into_iter().unzip(); + let mut positions = positions.into_iter(); + let fallback_dict = Dictionary::empty(); + let dictionary = field + .str_dict_column + .as_ref() + .map(|column| column.dictionary()) + .unwrap_or(&fallback_dict); + let all_found = dictionary.sorted_ords_to_term_cb(&ords, |term| { + keys[positions.next().unwrap()].push(IntermediateKey::Str( + String::from_utf8(term.to_vec()).expect("term dict returned non-UTF-8"), + )); + })?; + if !all_found { + return Err(TantivyError::InternalError( + "multi_terms string ordinal not found in dictionary".to_string(), + )); + } + } + Ok(keys) +} + /// Resolve one raw fast-field value, recognizing the configured missing encoding first. /// /// When the encoding reuses a string term ordinal, both paths resolve identically by construction. From 62f3d4b4dda900ec520890462449c1ac7114a10d Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Mon, 21 Sep 2026 08:58:27 +0200 Subject: [PATCH 24/49] Apply multi-terms cutoff before final bucket conversion Avoid formatting keys and finalizing sub-aggregations for buckets discarded by the final size limit. Preallocate the retained bucket vector and reuse the cutoff helper with intermediate tuple keys. Preserve final-key ordering and finalize the ordering metric when sorting by a sub-aggregation. Cover mixed key types, empty sums, and discarded document counts in tests. --- src/aggregation/bucket/multi_terms/mod.rs | 233 +++++++++++++++------- src/aggregation/bucket/term_agg/mod.rs | 2 +- 2 files changed, 158 insertions(+), 77 deletions(-) diff --git a/src/aggregation/bucket/multi_terms/mod.rs b/src/aggregation/bucket/multi_terms/mod.rs index 2db369d0c..7146740e7 100644 --- a/src/aggregation/bucket/multi_terms/mod.rs +++ b/src/aggregation/bucket/multi_terms/mod.rs @@ -1,3 +1,4 @@ +use std::collections::hash_map::Entry; use std::fmt::Debug; use std::net::Ipv6Addr; use std::sync::Arc; @@ -724,12 +725,12 @@ where // Distinct encoded keys can resolve to the same public key (for example a synthetic // missing string and a real date). Merge rather than overwriting either contribution. match result_entries.entry(intermediate_key) { - std::collections::hash_map::Entry::Occupied(mut occupied) => { + Entry::Occupied(mut occupied) => { let existing = occupied.get_mut(); existing.doc_count += doc_count; existing.sub_aggregation.merge_fruits(sub_aggregation_res)?; } - std::collections::hash_map::Entry::Vacant(vacant) => { + Entry::Vacant(vacant) => { vacant.insert(IntermediateTermBucketEntry { doc_count, sub_aggregation: sub_aggregation_res, @@ -1360,86 +1361,99 @@ impl IntermediateMultiTermsBucketResult { let req = MultiTermsAggregationInternal::from_req(req); - let mut buckets: Vec = self - .entries - .into_iter() - .filter(|(_, e)| e.doc_count >= req.min_doc_count) - .map(|(key_vec, entry)| { - let key_as_string = key_vec - .iter() - .map(|k| match k { - // Bool keys need special-casing: `Key` form is numeric (1/0), but - // `key_as_string` must still carry the "true"/"false" string form. - IntermediateKey::Bool(b) => b.to_string(), - other => Key::from(other.clone()).to_string(), - }) - .collect::>() - .join("|"); - let keys: Vec = key_vec.into_iter().map(Key::from).collect(); - Ok(MultiTermsBucketEntry { - key_as_string, - key: keys, - doc_count: entry.doc_count, - sub_aggregation: entry - .sub_aggregation - .into_final_result_internal(sub_aggregation_req, limits)?, - }) - }) - .collect::>()?; + let mut entries = Vec::with_capacity(self.entries.len()); + entries.extend( + self.entries + .into_iter() + .filter(|(_, e)| e.doc_count >= req.min_doc_count), + ); - // Sort by order. + // Select the final buckets before formatting keys or finalizing their sub-aggregations. match &req.order.target { OrderTarget::Count => { if req.order.order == Order::Desc { - buckets.sort_unstable_by_key(|b| std::cmp::Reverse(b.doc_count)); + entries.sort_unstable_by_key(|(_, entry)| std::cmp::Reverse(entry.doc_count)); } else { - buckets.sort_unstable_by_key(|b| b.doc_count); + entries.sort_unstable_by_key(|(_, entry)| entry.doc_count); } } OrderTarget::Key => { - buckets.sort_by(|left, right| { - let cmp = left - .key - .iter() - .zip(right.key.iter()) - .find_map(|(l, r)| { - let c = l.partial_cmp(r)?; - if c != std::cmp::Ordering::Equal { - Some(c) - } else { - None - } - }) - .unwrap_or(std::cmp::Ordering::Equal); - if req.order.order == Order::Asc { - cmp - } else { - cmp.reverse() - } - }); + // Final keys have different ordering from intermediate keys (e.g. IPs are + // strings and bools are u64s). Cache the conversion rather than cloning per + // comparison. + let key = |(key_vec, _): &(Vec, IntermediateTermBucketEntry)| { + key_vec.iter().cloned().map(Key::from).collect::>() + }; + if req.order.order == Order::Asc { + entries.sort_by_cached_key(key); + } else { + entries.sort_by_cached_key(|entry| std::cmp::Reverse(key(entry))); + } } OrderTarget::SubAggregation(name) => { let (agg_name, agg_property) = get_agg_name_and_property(name); - let mut buckets_with_val = buckets + let mut entries_with_val = entries .into_iter() - .map(|bucket| { - let val = bucket + .map(|entry| { + let sub_req = sub_aggregation_req.get(agg_name).ok_or_else(|| { + TantivyError::InternalError(format!( + "Can't find aggregation {agg_name:?} in sub-aggregations" + )) + })?; + // Only finalize the ordering metric. Its final value can differ from + // the intermediate value, e.g. an empty sum defaults to zero. + let metric = entry + .1 .sub_aggregation + .aggs_res + .get(agg_name) + .cloned() + .unwrap_or_else(|| { + crate::aggregation::intermediate_agg_result::empty_from_req(sub_req) + }); + let val = metric + .into_final_result(sub_req, limits)? .get_value_from_aggregation(agg_name, agg_property)? .unwrap_or(f64::MIN); - Ok((bucket, val)) + Ok((entry, val)) }) .collect::>>()?; - buckets_with_val.sort_by(|(_, v1), (_, v2)| match req.order.order { + entries_with_val.sort_by(|(_, v1), (_, v2)| match req.order.order { Order::Desc => v2.total_cmp(v1), Order::Asc => v1.total_cmp(v2), }); - buckets = buckets_with_val.into_iter().map(|(b, _)| b).collect(); + entries = entries_with_val + .into_iter() + .map(|(entry, _)| entry) + .collect(); } } let (_before_cutoff, sum_other_from_final) = - cut_off_buckets(&mut buckets, req.size as usize, None); + cut_off_buckets(&mut entries, req.size as usize, None); + + let mut buckets = Vec::with_capacity(entries.len()); + for (key_vec, entry) in entries { + let key_as_string = key_vec + .iter() + .map(|k| match k { + // Bool keys need special-casing: `Key` form is numeric (1/0), but + // `key_as_string` must still carry the "true"/"false" string form. + IntermediateKey::Bool(b) => b.to_string(), + other => Key::from(other.clone()).to_string(), + }) + .collect::>() + .join("|"); + let keys: Vec = key_vec.into_iter().map(Key::from).collect(); + buckets.push(MultiTermsBucketEntry { + key_as_string, + key: keys, + doc_count: entry.doc_count, + sub_aggregation: entry + .sub_aggregation + .into_final_result_internal(sub_aggregation_req, limits)?, + }); + } let doc_count_error_upper_bound = if req.show_term_doc_count_error { Some(self.doc_count_error_upper_bound) @@ -1581,6 +1595,63 @@ mod tests { Ok(()) } + #[test] + fn test_multi_terms_final_key_order() -> crate::Result<()> { + let keys = [ + IntermediateKey::IpAddr("::ffff:10.0.0.2".parse().unwrap()), + IntermediateKey::IpAddr("::ffff:10.0.0.10".parse().unwrap()), + IntermediateKey::Str("0".to_string()), + IntermediateKey::Bool(false), + IntermediateKey::I64(-1), + IntermediateKey::U64(2), + IntermediateKey::F64(1.5), + ]; + let expected = ["0", "10.0.0.10", "10.0.0.2", "-1", "false", "2", "1.5"]; + for order in ["asc", "desc"] { + let req = serde_json::from_value(json!({ + "terms": [{"field": "value"}], + "size": 6, + "order": {"_key": order} + }))?; + let intermediate = IntermediateMultiTermsBucketResult { + entries: keys + .iter() + .cloned() + .map(|key| { + ( + vec![key], + IntermediateTermBucketEntry { + doc_count: 1, + sub_aggregation: Default::default(), + }, + ) + }) + .collect(), + ..Default::default() + }; + let result = intermediate.into_final_result( + &req, + &Default::default(), + &mut Default::default(), + )?; + let result = serde_json::to_value(result)?; + let actual: Vec<_> = result["buckets"] + .as_array() + .unwrap() + .iter() + .map(|bucket| bucket["key_as_string"].as_str().unwrap()) + .collect(); + let mut expected = expected.to_vec(); + if order == "desc" { + expected.reverse(); + } + expected.truncate(6); + assert_eq!(actual, expected); + assert_eq!(result["sum_other_doc_count"], 1); + } + Ok(()) + } + #[test] fn test_multi_terms_min_doc_count() -> crate::Result<()> { let index = build_two_field_index( @@ -1631,7 +1702,7 @@ mod tests { let buckets = res["mt"]["buckets"].as_array().unwrap(); assert_eq!(buckets.len(), 2); assert_eq!(buckets[0]["key_as_string"], "rock|A"); - assert!(res["mt"]["sum_other_doc_count"].as_u64().unwrap() > 0); + assert_eq!(res["mt"]["sum_other_doc_count"], 1); Ok(()) } @@ -2734,7 +2805,7 @@ mod tests { } #[test] - fn test_multi_terms_missing_subagg_value_sorts_last_at_segment_cutoff() -> crate::Result<()> { + fn test_multi_terms_missing_subagg_value_at_cutoff() -> crate::Result<()> { let mut schema_builder = Schema::builder(); let genre_field = schema_builder.add_text_field("genre", STRING | FAST); let product_field = schema_builder.add_text_field("product", STRING | FAST); @@ -2751,23 +2822,33 @@ mod tests { writer.commit()?; } - let agg_req: Aggregations = serde_json::from_value(json!({ - "mt": { - "multi_terms": { - "terms": [{"field": "genre"}, {"field": "product"}], - "size": 1, - "segment_size": 1, - "order": {"avg_delta": "desc"} - }, - "aggs": { - "avg_delta": {"avg": {"field": "delta"}} + for (metric, segment_size, expected) in [ + (json!({"avg": {"field": "delta"}}), 1, "rock|A"), + (json!({"avg": {"field": "delta"}}), 2, "rock|A"), + // At final cutoff, an empty sum is zero unless none_if_no_match is set. + (json!({"sum": {"field": "delta"}}), 2, "rock|B"), + ( + json!({"sum": {"field": "delta", "none_if_no_match": true}}), + 2, + "rock|A", + ), + ] { + let agg_req: Aggregations = serde_json::from_value(json!({ + "mt": { + "multi_terms": { + "terms": [{"field": "genre"}, {"field": "product"}], + "size": 1, + "segment_size": segment_size, + "order": {"metric": "desc"} + }, + "aggs": {"metric": metric} } - } - }))?; - let res = exec_request(agg_req, &index)?; - let buckets = res["mt"]["buckets"].as_array().unwrap(); - assert_eq!(buckets.len(), 1); - assert_eq!(buckets[0]["key_as_string"], "rock|A"); + }))?; + let res = exec_request(agg_req, &index)?; + let buckets = res["mt"]["buckets"].as_array().unwrap(); + assert_eq!(buckets.len(), 1); + assert_eq!(buckets[0]["key_as_string"], expected); + } Ok(()) } diff --git a/src/aggregation/bucket/term_agg/mod.rs b/src/aggregation/bucket/term_agg/mod.rs index 694811866..15eda32a6 100644 --- a/src/aggregation/bucket/term_agg/mod.rs +++ b/src/aggregation/bucket/term_agg/mod.rs @@ -1554,7 +1554,7 @@ pub(crate) trait GetDocCount { fn doc_count(&self) -> u64; } -impl GetDocCount for (String, IntermediateTermBucketEntry) { +impl GetDocCount for (K, IntermediateTermBucketEntry) { fn doc_count(&self) -> u64 { self.1.doc_count } From 7d3a3d6bfece8dd82d4557d53264886d1f50e31c Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Tue, 22 Sep 2026 13:01:49 +0200 Subject: [PATCH 25/49] Use unwrap_or(false) for missing bucket value check --- src/aggregation/bucket/multi_terms/mod.rs | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/src/aggregation/bucket/multi_terms/mod.rs b/src/aggregation/bucket/multi_terms/mod.rs index 7146740e7..ff6ec2819 100644 --- a/src/aggregation/bucket/multi_terms/mod.rs +++ b/src/aggregation/bucket/multi_terms/mod.rs @@ -1093,7 +1093,10 @@ fn resolve_bucket_keys( for (position, entry) in entries.iter().enumerate() { let value = packing.unpack_value(&entry.key, field_idx); if field.column_type == ColumnType::Str - && !missing.as_ref().is_some_and(|m| m.missing_value == value) + && !missing + .as_ref() + .map(|m| m.missing_value == value) + .unwrap_or(false) { ords_and_positions.push((value, position)); } else { From e229de6db748c6b21747a99a18220446e76104e1 Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Tue, 22 Sep 2026 13:24:00 +0200 Subject: [PATCH 26/49] Collect decoded multi-term keys before assigning bucket positions --- src/aggregation/bucket/multi_terms/mod.rs | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/src/aggregation/bucket/multi_terms/mod.rs b/src/aggregation/bucket/multi_terms/mod.rs index ff6ec2819..d8467961b 100644 --- a/src/aggregation/bucket/multi_terms/mod.rs +++ b/src/aggregation/bucket/multi_terms/mod.rs @@ -1108,15 +1108,15 @@ fn resolve_bucket_keys( } ords_and_positions.sort_unstable(); let (ords, positions): (Vec<_>, Vec<_>) = ords_and_positions.into_iter().unzip(); - let mut positions = positions.into_iter(); let fallback_dict = Dictionary::empty(); let dictionary = field .str_dict_column .as_ref() .map(|column| column.dictionary()) .unwrap_or(&fallback_dict); + let mut decoded_keys = Vec::with_capacity(ords.len()); let all_found = dictionary.sorted_ords_to_term_cb(&ords, |term| { - keys[positions.next().unwrap()].push(IntermediateKey::Str( + decoded_keys.push(IntermediateKey::Str( String::from_utf8(term.to_vec()).expect("term dict returned non-UTF-8"), )); })?; @@ -1125,6 +1125,9 @@ fn resolve_bucket_keys( "multi_terms string ordinal not found in dictionary".to_string(), )); } + for (position, key) in positions.into_iter().zip(decoded_keys) { + keys[position].push(key); + } } Ok(keys) } From be8a1d09867fab0ac942cffa1dbb91382ded2a96 Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Fri, 11 Sep 2026 21:23:30 +0200 Subject: [PATCH 27/49] Add 26-bit date histogram aggregation benchmark --- benches/agg_bench.rs | 78 ++++++++++++++++++++++++++++++-------------- 1 file changed, 53 insertions(+), 25 deletions(-) diff --git a/benches/agg_bench.rs b/benches/agg_bench.rs index 4ebd7cb45..aba0adf5b 100644 --- a/benches/agg_bench.rs +++ b/benches/agg_bench.rs @@ -11,13 +11,13 @@ use tantivy::aggregation::AggregationCollector; use tantivy::indexer::NoMergePolicy; use tantivy::query::{AllQuery, Query, TermQuery}; use tantivy::schema::{IndexRecordOption, Schema, TextFieldIndexing, FAST, STRING}; -use tantivy::{doc, DateTime, Index, Term}; +use tantivy::{doc, DateTime, Index, Searcher, Term}; #[global_allocator] pub static GLOBAL: &PeakMemAlloc = &INSTRUMENTED_SYSTEM; type AggregationRequest = serde_json::Value; -type AggregationExecutor = fn(&Index, AggregationRequest); +type AggregationExecutor = fn(&Searcher, AggregationRequest); type BenchmarkConfig = (&'static str, AggregationRequest); type BenchmarkGroup = (&'static str, AggregationExecutor, Vec); @@ -40,6 +40,8 @@ fn main() { for (input_name, cardinality) in inputs { let index = get_test_index_bench(cardinality).unwrap(); + let reader = index.reader().unwrap(); + let searcher = reader.searcher(); // On sparse this will not effectively filter anything. This should simulate co-located // data which are sparse. // So for sparse aggregation, although the value is sparse all values in the aggregation @@ -51,7 +53,7 @@ fn main() { execute_agg_filtered }; runner.set_name(input_name); - bench_agg(&mut runner, &index, execute_filtered); + bench_agg(&mut runner, &searcher, execute_filtered); } for num_segments in [100, 1_000] { @@ -64,6 +66,8 @@ fn bench_many_segments(num_segments: usize) { runner.add_plugin(PeakMemAllocPlugin::new(GLOBAL)); runner.config().set_num_iter_for_group(1); let index = get_test_index_bench_with_num_segments(Cardinality::Full, num_segments).unwrap(); + let reader = index.reader().unwrap(); + let searcher = reader.searcher(); let mut group = runner.new_group(); group.set_name(format!("{num_segments}_segments")); let mut multi_terms_top500 = multi_terms_many_and_zipf_1000(); @@ -71,8 +75,6 @@ fn bench_many_segments(num_segments: usize) { let mut nested_terms_top500 = nested_terms_many_and_zipf_1000(); // This limits outer buckets, unlike the global tuple limit for multi_terms. nested_terms_top500["my_texts"]["terms"]["size"] = json!(500); - let reader = index.reader().unwrap(); - let searcher = reader.searcher(); for (benchmark_name, agg_req) in [ ("terms_7", terms_on_field("text_few_terms_status")), ("terms_zipfs_1000", terms_on_field("text_1000_terms_zipf")), @@ -88,12 +90,8 @@ fn bench_many_segments(num_segments: usize) { ), ("multi_terms_many_and_zipf_1000_top500", multi_terms_top500), ] { - let searcher = searcher.clone(); - group.register_with_input(benchmark_name, &(), move |_| { - let agg_req: Aggregations = serde_json::from_value(agg_req.clone()).unwrap(); - let collector = get_collector(agg_req); - - black_box(searcher.search(&AllQuery, &collector).unwrap()); + group.register_with_input(benchmark_name, &searcher, move |searcher| { + execute_agg(searcher, agg_req.clone()) }); } group.run(); @@ -105,7 +103,11 @@ fn terms_on_field(field: &str) -> AggregationRequest { }) } -fn bench_agg(runner: &mut BenchRunner, index: &Index, execute_filtered: AggregationExecutor) { +fn bench_agg( + runner: &mut BenchRunner, + searcher: &Searcher, + execute_filtered: AggregationExecutor, +) { let multi_terms_vs_nested = vec![ benchmark_config!(nested_terms_status_and_zipf_1000), benchmark_config!(multi_terms_status_and_zipf_1000), @@ -163,6 +165,7 @@ fn bench_agg(runner: &mut BenchRunner, index: &Index, execute_filtered: Aggregat benchmark_config!(terms_status_with_histogram), benchmark_config!(terms_zipf_1000_with_histogram), benchmark_config!(terms_status_with_date_histogram), + benchmark_config!(terms_status_with_date_histogram_26_bits), benchmark_config!(terms_status_with_date_histogram_single_bucket), benchmark_config!(terms_status_with_date_histogram_4_buckets), benchmark_config!(terms_status_with_date_histogram_8_buckets), @@ -215,8 +218,8 @@ fn bench_agg(runner: &mut BenchRunner, index: &Index, execute_filtered: Aggregat let mut group = runner.new_group(); group.set_name(group_name); for (benchmark_name, agg_req) in configs { - group.register_with_input(benchmark_name, index, move |index| { - execute(index, agg_req.clone()) + group.register_with_input(benchmark_name, searcher, move |searcher| { + execute(searcher, agg_req.clone()) }); } group.run(); @@ -551,6 +554,17 @@ fn terms_status_with_date_histogram() -> AggregationRequest { }) } +fn terms_status_with_date_histogram_26_bits() -> AggregationRequest { + json!({ + "my_texts": { + "terms": { "field": "text_few_terms_status" }, + "aggs": { + "over_time": { "date_histogram": { "field": "timestamp_26_bits", "fixed_interval": "134h" } } + } + } + }) +} + /// Same flattened terms × date_histogram, but with `hard_bounds`. The timestamps span 0..120h; the /// bounds drop only the first and last hour (ms: 1h=3_600_000, 119h=428_400_000), so almost every /// doc is in-bounds. This exercises the collector's hard-bounds path: `bounds.contains` runs per @@ -858,34 +872,36 @@ fn multi_terms_status_and_zipf_1000_avg_sub_agg() -> AggregationRequest { }) } -fn execute_agg(index: &Index, agg_req: AggregationRequest) { - execute_agg_with_query(index, agg_req, &AllQuery); +fn execute_agg(searcher: &Searcher, agg_req: AggregationRequest) { + execute_agg_with_query(searcher, agg_req, &AllQuery); } -fn execute_agg_filtered(index: &Index, agg_req: AggregationRequest) { - let filter_field = index.schema().get_field("filter_field").unwrap(); +fn execute_agg_filtered(searcher: &Searcher, agg_req: AggregationRequest) { + let filter_field = searcher.schema().get_field("filter_field").unwrap(); let filter_query = TermQuery::new( Term::from_field_text(filter_field, "a"), IndexRecordOption::Basic, ); - execute_agg_with_query(index, agg_req, &filter_query); + execute_agg_with_query(searcher, agg_req, &filter_query); } -fn execute_agg_filtered_on_single_term(index: &Index, agg_req: AggregationRequest) { - let filter_field = index.schema().get_field("single_term").unwrap(); +fn execute_agg_filtered_on_single_term(searcher: &Searcher, agg_req: AggregationRequest) { + let filter_field = searcher.schema().get_field("single_term").unwrap(); let filter_query = TermQuery::new( Term::from_field_text(filter_field, "single_term"), IndexRecordOption::Basic, ); - execute_agg_with_query(index, agg_req, &filter_query); + execute_agg_with_query(searcher, agg_req, &filter_query); } -fn execute_agg_with_query(index: &Index, agg_req: AggregationRequest, query: &dyn Query) { +fn execute_agg_with_query( + searcher: &Searcher, + agg_req: AggregationRequest, + query: &dyn Query, +) { let agg_req: Aggregations = serde_json::from_value(agg_req).unwrap(); let collector = get_collector(agg_req); - let reader = index.reader().unwrap(); - let searcher = reader.searcher(); black_box(searcher.search(query, &collector).unwrap()); } @@ -1081,6 +1097,7 @@ fn get_test_index_bench_with_num_segments( let score_field_f64 = schema_builder.add_f64_field("score_f64", score_fieldtype.clone()); let score_field_i64 = schema_builder.add_i64_field("score_i64", score_fieldtype); let date_field = schema_builder.add_date_field("timestamp", FAST); + let date_26_bits_field = schema_builder.add_date_field("timestamp_26_bits", FAST); let schema = schema_builder.build(); if reuse_index && std::path::Path::new(&index_dir).try_exists()? { @@ -1133,6 +1150,7 @@ fn get_test_index_bench_with_num_segments( { let mut rng = StdRng::from_seed([1u8; 32]); let mut filter_rng = StdRng::from_seed([2u8; 32]); + let mut timestamp_26_bits_rng = StdRng::from_seed([3u8; 32]); let mut index_writer = index.writer_with_num_threads(1, 400_000_000)?; if num_segments > 1 { index_writer.set_merge_policy(Box::new(NoMergePolicy)); @@ -1207,6 +1225,7 @@ fn get_test_index_bench_with_num_segments( let _val_max = 1_000_000.0; const SPAN_MS: i64 = 120 * 3600 * 1000; // 120 hours in ms const NOISE_MS: i64 = 2 * 3600 * 1000; // ±2h noise + const MAX_26_BIT_TIMESTAMP_SECS: i64 = (1 << 26) - 1; for i in 0..doc_with_value { let val: f64 = rng.random_range(0.0..1_000_000.0); let json = if rng.random_bool(0.1) { @@ -1218,6 +1237,14 @@ fn get_test_index_bench_with_num_segments( let base_ms = (i as i64 * SPAN_MS) / doc_with_value as i64; let noise_ms = rng.random_range(-NOISE_MS..NOISE_MS); let ts_ms = (base_ms + noise_ms).clamp(0, SPAN_MS); + // Force the endpoints and randomize the interior so the column uses a 26-bit packed + // representation rather than the blockwise-linear codec. + let ts_26_bits_secs = match i { + 0 => 0, + 1 => 1, + i if i + 1 == doc_with_value => MAX_26_BIT_TIMESTAMP_SECS, + _ => timestamp_26_bits_rng.random_range(0..=MAX_26_BIT_TIMESTAMP_SECS), + }; add_document(doc!( single_term => "single_term", text_field => "cool", @@ -1232,6 +1259,7 @@ fn get_test_index_bench_with_num_segments( score_field_f64 => lg_norm.sample(&mut rng), score_field_i64 => val as i64, date_field => DateTime::from_timestamp_millis(ts_ms), + date_26_bits_field => DateTime::from_timestamp_secs(ts_26_bits_secs), ))?; if cardinality == Cardinality::OptionalSparse { for _ in 0..20 { From 4a660f59d2dfe9d99f912ecf40691045ef0b8f07 Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Mon, 14 Sep 2026 12:20:38 +0200 Subject: [PATCH 28/49] Faster columnar data fetching --- bitpacker/src/bitpacker.rs | 109 ++++++++++++++++- columnar/src/column/mod.rs | 4 +- .../src/column_values/monotonic_column.rs | 45 ++++++- .../src/column_values/monotonic_mapping.rs | 15 +++ .../column_values/monotonic_mapping_u128.rs | 5 + columnar/src/column_values/u128_based/mod.rs | 2 +- .../src/column_values/u64_based/bitpacked.rs | 110 +++++++++++++++++- columnar/src/column_values/u64_based/mod.rs | 14 ++- columnar/src/dynamic_column.rs | 9 +- 9 files changed, 288 insertions(+), 25 deletions(-) diff --git a/bitpacker/src/bitpacker.rs b/bitpacker/src/bitpacker.rs index 1efd2755e..ce18c3415 100644 --- a/bitpacker/src/bitpacker.rs +++ b/bitpacker/src/bitpacker.rs @@ -75,6 +75,7 @@ impl BitUnpacker { /// The bitunpacker works by doing an unaligned read of 8 bytes. /// For this reason, values of `num_bits` between /// [57..63] are forbidden. + #[inline] pub fn new(num_bits: u8) -> BitUnpacker { assert!(num_bits <= 7 * 8 || num_bits == 64); let mask: u64 = if num_bits == 64 { @@ -101,7 +102,7 @@ impl BitUnpacker { return 0; } let bit_shift = addr_in_bits & 7; - return self.get_slow_path(addr, bit_shift as u32, data); + return Self::get_slow_path(self.mask, addr, bit_shift as u32, data); } let bit_shift = addr_in_bits & 7; let bytes: [u8; 8] = (&data[addr..addr + 8]).try_into().unwrap(); @@ -110,8 +111,83 @@ impl BitUnpacker { val_shifted & self.mask } + /// Decodes consecutive values into `output`. + /// + /// Panics if the requested bits are outside `data`. + #[inline(always)] + pub fn get_range(&self, start_idx: u32, data: &[u8], output: &mut [u64]) { + if output.is_empty() { + return; + } + if self.num_bits == 0 { + output.fill(0); + return; + } + let start_idx = start_idx as usize; + let end_bit = (start_idx + output.len()) * self.num_bits; + debug_assert!( + end_bit.div_ceil(8) <= data.len(), + "Requested range is out of bounds" + ); + + // Only the last few values may need a partial eight-byte load. + let fast_len = if data.len() >= 8 { + let last_full_load_bit = (data.len() - 8) * 8 + 7; + let full_values = last_full_load_bit / self.num_bits + 1; + (full_values.max(start_idx) - start_idx).min(output.len()) + } else { + 0 + }; + let (fast, tail) = output.split_at_mut(fast_len); + let load = |bit_addr: usize| { + // SAFETY: only called for values in fast, whose eight-byte loads fit in data. + let packed = unsafe { + data.as_ptr() + .add(bit_addr >> 3) + .cast::() + .read_unaligned() + }; + u64::from_le(packed) >> (bit_addr & 7) + }; + let mut bit_addr = start_idx * self.num_bits; + let mut chunks = fast.chunks_exact_mut(4); + for chunk in &mut chunks { + // Four values plus at most seven leading bits fit in one load. + // At 16 bits, values are byte-aligned, so there are no leading bits. + if self.num_bits <= 14 || self.num_bits == 16 { + let packed = load(bit_addr); + for (i, out) in chunk.iter_mut().enumerate() { + *out = (packed >> (i * self.num_bits)) & self.mask; + } + } else if self.num_bits <= 28 || self.num_bits == 32 { + // Two values plus at most seven leading bits fit in one load. + // At 32 bits, values are byte-aligned, so there are no leading bits. + for (pair_idx, pair) in chunk.chunks_exact_mut(2).enumerate() { + let packed = load(bit_addr + pair_idx * 2 * self.num_bits); + pair[0] = packed & self.mask; + pair[1] = (packed >> self.num_bits) & self.mask; + } + } else { + for (i, out) in chunk.iter_mut().enumerate() { + *out = load(bit_addr + i * self.num_bits) & self.mask; + } + } + bit_addr += 4 * self.num_bits; + } + for out in chunks.into_remainder() { + *out = load(bit_addr) & self.mask; + bit_addr += self.num_bits; + } + for out in tail { + *out = Self::get_slow_path(self.mask, bit_addr >> 3, (bit_addr & 7) as u32, data); + bit_addr += self.num_bits; + } + } + + // Pass the mask by value so specialized callers don't need to materialize a + // temporary BitUnpacker on the stack just to pass &self to this non-inlined helper. #[inline(never)] - fn get_slow_path(&self, addr: usize, bit_shift: u32, data: &[u8]) -> u64 { + fn get_slow_path(mask: u64, addr: usize, bit_shift: u32, data: &[u8]) -> u64 { let mut bytes: [u8; 8] = [0u8; 8]; let available_bytes = data.len() - addr; // This function is meant to only be called if we did not have 8 bytes to load. @@ -119,7 +195,7 @@ impl BitUnpacker { bytes[..available_bytes].copy_from_slice(&data[addr..]); let val_unshifted_unmasked: u64 = u64::from_le_bytes(bytes); let val_shifted = val_unshifted_unmasked >> bit_shift; - val_shifted & self.mask + val_shifted & mask } // Decodes the range of bitpacked `u32` values with idx @@ -315,6 +391,33 @@ mod test { assert!(val <= max_val); assert_eq!(bitunpacker.get(i as u32, &buffer), val); } + for start in 0..=vals.len() { + let remaining = vals.len() - start; + for len in [0, remaining.min(1), remaining / 2, remaining] { + let mut output = vec![u64::MAX; len]; + bitunpacker.get_range(start as u32, &buffer, &mut output); + assert_eq!(output, vals[start..start + len]); + } + } + } + + #[test] + fn test_get_range_all_bit_widths() { + for num_bits in (0..=56).chain(std::iter::once(64)) { + let mask = u64::MAX.checked_shr(64 - num_bits as u32).unwrap_or(0); + for len in [0, 1, 2, 7, 8, 9, 31, 32, 33, 63, 64, 65, 255, 256, 257] { + let vals: Vec = (0..len) + .map(|i| (i as u64).wrapping_mul(0x9e3779b97f4a7c15) & mask) + .collect(); + test_bitpacker_aux(num_bits, &vals); + } + } + } + + #[test] + #[should_panic(expected = "Requested range is out of bounds")] + fn test_get_range_out_of_bounds() { + BitUnpacker::new(3).get_range(2, &[0], &mut [0]); } proptest::proptest! { diff --git a/columnar/src/column/mod.rs b/columnar/src/column/mod.rs index f6a50b45f..92a08ad33 100644 --- a/columnar/src/column/mod.rs +++ b/columnar/src/column/mod.rs @@ -46,10 +46,10 @@ impl Column { impl Column { pub fn to_u64_monotonic(self) -> Column { - let values = Arc::new(monotonic_map_column( + let values = monotonic_map_column( self.values, StrictlyMonotonicMappingToInternal::::new(), - )); + ); Column { index: self.index, values, diff --git a/columnar/src/column_values/monotonic_column.rs b/columnar/src/column_values/monotonic_column.rs index 35de3787a..d5927982c 100644 --- a/columnar/src/column_values/monotonic_column.rs +++ b/columnar/src/column_values/monotonic_column.rs @@ -1,6 +1,8 @@ +use std::any::{Any, TypeId}; use std::fmt::Debug; use std::marker::PhantomData; use std::ops::{Range, RangeInclusive}; +use std::sync::Arc; use crate::ColumnValues; use crate::column_values::monotonic_mapping::StrictlyMonotonicFn; @@ -29,18 +31,26 @@ struct MonotonicMappingColumn { pub fn monotonic_map_column( from_column: C, monotonic_mapping: T, -) -> impl ColumnValues +) -> Arc> where C: ColumnValues + 'static, T: StrictlyMonotonicFn + Send + Sync + 'static, Input: PartialOrd + Debug + Send + Sync + Clone + 'static, Output: PartialOrd + Debug + Send + Sync + Clone + 'static, { - MonotonicMappingColumn { + // Preserve specialized codec methods (notably get_range) for identity mappings. + if T::IS_IDENTITY && TypeId::of::() == TypeId::of::() { + let column: Arc> = Arc::new(from_column); + return (&column as &dyn Any) + .downcast_ref::>>() + .unwrap() + .clone(); + } + Arc::new(MonotonicMappingColumn { from_column, monotonic_mapping, _phantom: PhantomData, - } + }) } impl ColumnValues for MonotonicMappingColumn @@ -104,6 +114,35 @@ mod tests { StrictlyMonotonicMappingInverter, StrictlyMonotonicMappingToInternal, }; + #[test] + fn test_u128_identity() { + let column: Arc> = monotonic_map_column( + VecColumn::from(vec![u128::MAX]), + StrictlyMonotonicMappingInverter::from( + StrictlyMonotonicMappingToInternal::::new(), + ), + ); + assert_eq!(column.get_val(0), u128::MAX); + } + + #[test] + fn test_same_type_non_identity() { + struct Shift; + impl StrictlyMonotonicFn for Shift { + fn mapping(&self, value: u64) -> u64 { + value + 1 + } + + fn inverse(&self, value: u64) -> u64 { + value - 1 + } + } + let column = monotonic_map_column(VecColumn::from(vec![1u64, 2, 3]), Shift); + let mut output = [0; 3]; + column.get_range(0, &mut output); + assert_eq!(output, [2, 3, 4]); + } + #[test] fn test_monotonic_mapping_iter() { let vals: Vec = (0..100u64).map(|el| el * 10).collect(); diff --git a/columnar/src/column_values/monotonic_mapping.rs b/columnar/src/column_values/monotonic_mapping.rs index 4626053ed..f83f6d118 100644 --- a/columnar/src/column_values/monotonic_mapping.rs +++ b/columnar/src/column_values/monotonic_mapping.rs @@ -9,6 +9,9 @@ use crate::RowId; /// Monotonic maps a value to u64 value space. /// Monotonic mapping enables `PartialOrd` on u64 space without conversion to original space. pub trait MonotonicallyMappableToU64: 'static + PartialOrd + Debug + Copy + Send + Sync { + /// Whether conversion to and from u64 leaves values unchanged. + const IS_IDENTITY: bool = false; + /// Converts a value to u64. /// /// Internally all fast field values are encoded as u64. @@ -32,6 +35,10 @@ pub trait MonotonicallyMappableToU64: 'static + PartialOrd + Debug + Copy + Send /// so a value can be converted back to its original domain (e.g. ip address or f64) from its /// internal representation. pub trait StrictlyMonotonicFn { + /// Whether both mapping directions leave values unchanged. + /// Only used to bypass mapping when the input and output types also match. + const IS_IDENTITY: bool = false; + /// Strictly monotonically maps the value from External to Internal. fn mapping(&self, inp: External) -> Internal; /// Inverse of `mapping`. Maps the value from Internal to External. @@ -58,6 +65,8 @@ impl From for StrictlyMonotonicMappingInverter { impl StrictlyMonotonicFn for StrictlyMonotonicMappingInverter where T: StrictlyMonotonicFn { + const IS_IDENTITY: bool = T::IS_IDENTITY; + #[inline(always)] fn mapping(&self, val: To) -> From { self.orig_mapping.inverse(val) @@ -86,6 +95,8 @@ impl StrictlyMonotonicFn for StrictlyMonotonicMappingToInternal where T: MonotonicallyMappableToU128 { + const IS_IDENTITY: bool = External::IS_IDENTITY; + #[inline(always)] fn mapping(&self, inp: External) -> u128 { External::to_u128(inp) @@ -101,6 +112,8 @@ impl StrictlyMonotonicFn for StrictlyMonotonicMappingToInternal where T: MonotonicallyMappableToU64 { + const IS_IDENTITY: bool = External::IS_IDENTITY; + #[inline(always)] fn mapping(&self, inp: External) -> u64 { External::to_u64(inp) @@ -113,6 +126,8 @@ where T: MonotonicallyMappableToU64 } impl MonotonicallyMappableToU64 for u64 { + const IS_IDENTITY: bool = true; + #[inline(always)] fn to_u64(self) -> u64 { self diff --git a/columnar/src/column_values/monotonic_mapping_u128.rs b/columnar/src/column_values/monotonic_mapping_u128.rs index 9e16dc58c..5b6f5192e 100644 --- a/columnar/src/column_values/monotonic_mapping_u128.rs +++ b/columnar/src/column_values/monotonic_mapping_u128.rs @@ -4,6 +4,9 @@ use std::net::Ipv6Addr; /// Monotonic maps a value to u128 value space /// Monotonic mapping enables `PartialOrd` on u128 space without conversion to original space. pub trait MonotonicallyMappableToU128: 'static + PartialOrd + Copy + Debug + Send + Sync { + /// Whether conversion to and from u128 leaves values unchanged. + const IS_IDENTITY: bool = false; + /// Converts a value to u128. /// /// Internally all fast field values are encoded as u64. @@ -17,6 +20,8 @@ pub trait MonotonicallyMappableToU128: 'static + PartialOrd + Copy + Debug + Sen } impl MonotonicallyMappableToU128 for u128 { + const IS_IDENTITY: bool = true; + fn to_u128(self) -> u128 { self } diff --git a/columnar/src/column_values/u128_based/mod.rs b/columnar/src/column_values/u128_based/mod.rs index d26f5ce35..b08026980 100644 --- a/columnar/src/column_values/u128_based/mod.rs +++ b/columnar/src/column_values/u128_based/mod.rs @@ -108,7 +108,7 @@ pub fn open_u128_mapped( let reader = CompactSpaceDecompressor::open(bytes)?; let inverted: StrictlyMonotonicMappingInverter> = StrictlyMonotonicMappingToInternal::::new().into(); - Ok(Arc::new(monotonic_map_column(reader, inverted))) + Ok(monotonic_map_column(reader, inverted)) } /// Returns the u64 representation of the u128 data. diff --git a/columnar/src/column_values/u64_based/bitpacked.rs b/columnar/src/column_values/u64_based/bitpacked.rs index 71319cbec..25bec8063 100644 --- a/columnar/src/column_values/u64_based/bitpacked.rs +++ b/columnar/src/column_values/u64_based/bitpacked.rs @@ -1,18 +1,18 @@ use std::io::{self, Write}; use std::num::NonZeroU64; use std::ops::{Range, RangeInclusive}; +use std::sync::Arc; use common::{BinarySerializable, OwnedBytes}; use fastdivide::DividerU64; use tantivy_bitpacker::{BitPacker, BitUnpacker, compute_num_bits}; use crate::column_values::u64_based::{ColumnCodec, ColumnCodecEstimator, ColumnStats}; -use crate::{ColumnValues, RowId}; +use crate::{ColumnValues, MonotonicallyMappableToU64, RowId}; -/// Depending on the field type, a different -/// fast field is required. +/// A bitpacked column reader. `u8::MAX` uses the bit width stored in the column. #[derive(Clone)] -pub struct BitpackedReader { +pub struct BitpackedReader { data: OwnedBytes, bit_unpacker: BitUnpacker, stats: ColumnStats, @@ -48,10 +48,37 @@ fn transform_range_before_linear_transformation( Some(start_before_gcd_multiplication..=end_before_gcd_multiplication) } -impl ColumnValues for BitpackedReader { +impl BitpackedReader { + #[inline(always)] + fn unpacker(&self) -> BitUnpacker { + if NUM_BITS == u8::MAX { + self.bit_unpacker + } else { + BitUnpacker::new(NUM_BITS) + } + } +} + +impl ColumnValues for BitpackedReader { #[inline(always)] fn get_val(&self, doc: u32) -> u64 { - self.stats.min_value + self.stats.gcd.get() * self.bit_unpacker.get(doc, &self.data) + self.stats.min_value + self.stats.gcd.get() * self.unpacker().get(doc, &self.data) + } + + fn get_range(&self, start: u64, output: &mut [u64]) { + debug_assert!(start <= u64::from(self.stats.num_rows)); + debug_assert!(output.len() as u64 <= u64::from(self.stats.num_rows) - start); + if NUM_BITS == 0 { + output.fill(self.stats.min_value); + return; + } + self.unpacker().get_range(start as u32, &self.data, output); + let skip_processing = self.stats.gcd.get() == 1 && self.stats.min_value == 0; + if !skip_processing { + for val in output { + *val = self.stats.min_value + self.stats.gcd.get() * *val; + } + } } #[inline] fn min_value(&self) -> u64 { @@ -139,11 +166,82 @@ impl ColumnCodec for BitpackedCodec { } } +/// Specialize widths with a meaningful decoding speedup; other widths share one decoder. +pub(super) fn load( + bytes: OwnedBytes, +) -> io::Result>> { + let reader = BitpackedCodec::load(bytes)?; + macro_rules! specialize { + ($($bits:literal),* $(,)?) => { + match reader.bit_unpacker.bit_width() { + $( + $bits => super::map_column_values::<_, T>(BitpackedReader::<$bits> { + data: reader.data, + bit_unpacker: reader.bit_unpacker, + stats: reader.stats, + }), + )* + _ => super::map_column_values::<_, T>(reader), + } + }; + } + Ok(specialize!(8, 16, 20, 24, 32, 64,)) +} + #[cfg(test)] mod tests { use super::*; use crate::column_values::u64_based::tests::create_and_validate; + #[test] + fn test_specialized_bit_widths() { + for bits in (0..=56).chain(std::iter::once(64)) { + let mask = u64::MAX.checked_shr(64 - bits as u32).unwrap_or(0); + for gcd in [1, 3] { + if gcd != 1 && mask > (u64::MAX - 7) / gcd { + continue; + } + let min_value = if bits == 64 { 0 } else { 7 }; + let vals: Vec = (0..257) + .map(|i| { + let packed = match i { + 0 => 0, + 1 => mask.min(1), + 2 => mask, + _ => (i as u64).wrapping_mul(0x9e3779b97f4a7c15) & mask, + }; + min_value + gcd * packed + }) + .collect(); + let mut stats = super::super::StatsCollector::default(); + for &val in &vals { + stats.collect(val); + } + let mut buffer = Vec::new(); + BitpackedCodecEstimator + .serialize(&stats.stats(), &mut vals.iter().copied(), &mut buffer) + .unwrap(); + let data = OwnedBytes::new(buffer); + let reader = load::(data.clone()).unwrap(); + let signed_reader = load::(data).unwrap(); + for start in 0..=vals.len() { + let mut output = vec![0; vals.len() - start]; + reader.get_range(start as u64, &mut output); + assert_eq!(output, vals[start..]); + for (i, &val) in output.iter().enumerate() { + assert_eq!(reader.get_val((start + i) as u32), val); + } + let mut signed_output = vec![0; output.len()]; + signed_reader.get_range(start as u64, &mut signed_output); + assert_eq!( + signed_output, + output.into_iter().map(i64::from_u64).collect::>() + ); + } + } + } + } + #[test] fn test_with_codec_data_sets_simple() { create_and_validate::(&[4, 3, 12], "name"); diff --git a/columnar/src/column_values/u64_based/mod.rs b/columnar/src/column_values/u64_based/mod.rs index aa2d9818b..313cd2251 100644 --- a/columnar/src/column_values/u64_based/mod.rs +++ b/columnar/src/column_values/u64_based/mod.rs @@ -114,7 +114,7 @@ impl CodecType { bytes: OwnedBytes, ) -> io::Result>> { match self { - CodecType::Bitpacked => load_specific_codec::(bytes), + CodecType::Bitpacked => bitpacked::load::(bytes), CodecType::Linear => load_specific_codec::(bytes), CodecType::BlockwiseLinear => load_specific_codec::(bytes), } @@ -124,12 +124,16 @@ impl CodecType { fn load_specific_codec( bytes: OwnedBytes, ) -> io::Result>> { - let reader = C::load(bytes)?; - let reader_typed = monotonic_map_column( + Ok(map_column_values::<_, T>(C::load(bytes)?)) +} + +fn map_column_values( + reader: C, +) -> Arc> { + monotonic_map_column( reader, StrictlyMonotonicMappingInverter::from(StrictlyMonotonicMappingToInternal::::new()), - ); - Ok(Arc::new(reader_typed)) + ) } impl CodecType { diff --git a/columnar/src/dynamic_column.rs b/columnar/src/dynamic_column.rs index 58f689ebd..eadcbe736 100644 --- a/columnar/src/dynamic_column.rs +++ b/columnar/src/dynamic_column.rs @@ -1,5 +1,4 @@ use std::net::Ipv6Addr; -use std::sync::Arc; use std::{fmt, io}; use common::file_slice::FileSlice; @@ -124,11 +123,11 @@ impl DynamicColumn { match self { DynamicColumn::I64(column) => Some(DynamicColumn::F64(Column { index: column.index, - values: Arc::new(monotonic_map_column(column.values, MapI64ToF64)), + values: monotonic_map_column(column.values, MapI64ToF64), })), DynamicColumn::U64(column) => Some(DynamicColumn::F64(Column { index: column.index, - values: Arc::new(monotonic_map_column(column.values, MapU64ToF64)), + values: monotonic_map_column(column.values, MapU64ToF64), })), DynamicColumn::F64(_) => Some(self), _ => None, @@ -142,7 +141,7 @@ impl DynamicColumn { } Some(DynamicColumn::I64(Column { index: column.index, - values: Arc::new(monotonic_map_column(column.values, MapU64ToI64)), + values: monotonic_map_column(column.values, MapU64ToI64), })) } DynamicColumn::I64(_) => Some(self), @@ -157,7 +156,7 @@ impl DynamicColumn { } Some(DynamicColumn::U64(Column { index: column.index, - values: Arc::new(monotonic_map_column(column.values, MapI64ToU64)), + values: monotonic_map_column(column.values, MapI64ToU64), })) } DynamicColumn::U64(_) => Some(self), From 4ede8d64f287e7dfcd7c2eb68ee1b6ab35012ced Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Tue, 15 Sep 2026 18:42:48 +0200 Subject: [PATCH 29/49] Optimize low-bit-width block decoding Decode eight values per load for 64-value blocks and specialize bit widths 1 through 7 to enable constant folding in the hot path. --- bitpacker/src/bitpacker.rs | 23 +++++++++++++++---- .../src/column_values/u64_based/bitpacked.rs | 2 +- 2 files changed, 19 insertions(+), 6 deletions(-) diff --git a/bitpacker/src/bitpacker.rs b/bitpacker/src/bitpacker.rs index ce18c3415..c79a77653 100644 --- a/bitpacker/src/bitpacker.rs +++ b/bitpacker/src/bitpacker.rs @@ -138,9 +138,10 @@ impl BitUnpacker { } else { 0 }; + let output_len = output.len(); let (fast, tail) = output.split_at_mut(fast_len); let load = |bit_addr: usize| { - // SAFETY: only called for values in fast, whose eight-byte loads fit in data. + // SAFETY: only called for values in `fast`, whose eight-byte loads fit in `data`. let packed = unsafe { data.as_ptr() .add(bit_addr >> 3) @@ -150,8 +151,20 @@ impl BitUnpacker { u64::from_le(packed) >> (bit_addr & 7) }; let mut bit_addr = start_idx * self.num_bits; - let mut chunks = fast.chunks_exact_mut(4); - for chunk in &mut chunks { + // This is a special case for 64 values of 1-8 bits, where we can decode 8 values at a + // time, which is the fastest possible. + if output_len == 64 && fast_len == 64 && self.num_bits <= 8 { + for chunk in fast.as_chunks_mut::<8>().0 { + let packed = load(bit_addr); + for (i, out) in chunk.iter_mut().enumerate() { + *out = (packed >> (i * self.num_bits)) & self.mask; + } + bit_addr += 8 * self.num_bits; + } + return; + } + let (chunks, remainder) = fast.as_chunks_mut::<4>(); + for chunk in chunks { // Four values plus at most seven leading bits fit in one load. // At 16 bits, values are byte-aligned, so there are no leading bits. if self.num_bits <= 14 || self.num_bits == 16 { @@ -162,7 +175,7 @@ impl BitUnpacker { } else if self.num_bits <= 28 || self.num_bits == 32 { // Two values plus at most seven leading bits fit in one load. // At 32 bits, values are byte-aligned, so there are no leading bits. - for (pair_idx, pair) in chunk.chunks_exact_mut(2).enumerate() { + for (pair_idx, pair) in chunk.as_chunks_mut::<2>().0.iter_mut().enumerate() { let packed = load(bit_addr + pair_idx * 2 * self.num_bits); pair[0] = packed & self.mask; pair[1] = (packed >> self.num_bits) & self.mask; @@ -174,7 +187,7 @@ impl BitUnpacker { } bit_addr += 4 * self.num_bits; } - for out in chunks.into_remainder() { + for out in remainder { *out = load(bit_addr) & self.mask; bit_addr += self.num_bits; } diff --git a/columnar/src/column_values/u64_based/bitpacked.rs b/columnar/src/column_values/u64_based/bitpacked.rs index 25bec8063..19083355b 100644 --- a/columnar/src/column_values/u64_based/bitpacked.rs +++ b/columnar/src/column_values/u64_based/bitpacked.rs @@ -185,7 +185,7 @@ pub(super) fn load( } }; } - Ok(specialize!(8, 16, 20, 24, 32, 64,)) + Ok(specialize!(1, 2, 3, 4, 5, 6, 7, 8, 16, 20, 24, 32, 64,)) } #[cfg(test)] From 21a69132175a225f24ec73f862d66ab409368d27 Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Wed, 16 Sep 2026 14:31:05 +0200 Subject: [PATCH 30/49] Simplify bitpacked range decoding Use scalar decoding for ranges overlapping the final partial load. Clarify the 64-value aggregation block specialization and name the generic decoding chunk size. Apply nightly formatting to the touched benchmark and columnar code. --- benches/agg_bench.rs | 12 ++---------- bitpacker/src/bitpacker.rs | 39 +++++++++++++++++++------------------- columnar/src/column/mod.rs | 6 ++---- 3 files changed, 23 insertions(+), 34 deletions(-) diff --git a/benches/agg_bench.rs b/benches/agg_bench.rs index aba0adf5b..35a6d7096 100644 --- a/benches/agg_bench.rs +++ b/benches/agg_bench.rs @@ -103,11 +103,7 @@ fn terms_on_field(field: &str) -> AggregationRequest { }) } -fn bench_agg( - runner: &mut BenchRunner, - searcher: &Searcher, - execute_filtered: AggregationExecutor, -) { +fn bench_agg(runner: &mut BenchRunner, searcher: &Searcher, execute_filtered: AggregationExecutor) { let multi_terms_vs_nested = vec![ benchmark_config!(nested_terms_status_and_zipf_1000), benchmark_config!(multi_terms_status_and_zipf_1000), @@ -894,11 +890,7 @@ fn execute_agg_filtered_on_single_term(searcher: &Searcher, agg_req: Aggregation execute_agg_with_query(searcher, agg_req, &filter_query); } -fn execute_agg_with_query( - searcher: &Searcher, - agg_req: AggregationRequest, - query: &dyn Query, -) { +fn execute_agg_with_query(searcher: &Searcher, agg_req: AggregationRequest, query: &dyn Query) { let agg_req: Aggregations = serde_json::from_value(agg_req).unwrap(); let collector = get_collector(agg_req); diff --git a/bitpacker/src/bitpacker.rs b/bitpacker/src/bitpacker.rs index c79a77653..dbedd47bf 100644 --- a/bitpacker/src/bitpacker.rs +++ b/bitpacker/src/bitpacker.rs @@ -130,18 +130,18 @@ impl BitUnpacker { "Requested range is out of bounds" ); - // Only the last few values may need a partial eight-byte load. - let fast_len = if data.len() >= 8 { - let last_full_load_bit = (data.len() - 8) * 8 + 7; - let full_values = last_full_load_bit / self.num_bits + 1; - (full_values.max(start_idx) - start_idx).min(output.len()) - } else { - 0 - }; + // Fall back for ranges overlapping the end, where an eight-byte load would be partial. + let last_bit_addr = (start_idx + output.len() - 1) * self.num_bits; + if (last_bit_addr >> 3) + 8 > data.len() { + for (offset, out) in output.iter_mut().enumerate() { + *out = self.get((start_idx + offset) as u32, data); + } + return; + } + let output_len = output.len(); - let (fast, tail) = output.split_at_mut(fast_len); let load = |bit_addr: usize| { - // SAFETY: only called for values in `fast`, whose eight-byte loads fit in `data`. + // SAFETY: the range-end check above guarantees that this load fits in `data`. let packed = unsafe { data.as_ptr() .add(bit_addr >> 3) @@ -151,10 +151,12 @@ impl BitUnpacker { u64::from_le(packed) >> (bit_addr & 7) }; let mut bit_addr = start_idx * self.num_bits; - // This is a special case for 64 values of 1-8 bits, where we can decode 8 values at a - // time, which is the fastest possible. - if output_len == 64 && fast_len == 64 && self.num_bits <= 8 { - for chunk in fast.as_chunks_mut::<8>().0 { + // Tantivy's `COLLECT_BLOCK_BUFFER_LEN` is 64, so optimize its common full-block case by + // decoding eight 1-8 bit values per load. Keep this literal in sync with that constant. + if output_len == 64 && self.num_bits <= 8 { + let (chunks, remainder) = output.as_chunks_mut::<8>(); + debug_assert!(remainder.is_empty()); + for chunk in chunks { let packed = load(bit_addr); for (i, out) in chunk.iter_mut().enumerate() { *out = (packed >> (i * self.num_bits)) & self.mask; @@ -163,7 +165,8 @@ impl BitUnpacker { } return; } - let (chunks, remainder) = fast.as_chunks_mut::<4>(); + const VALUES_PER_CHUNK: usize = 4; + let (chunks, remainder) = output.as_chunks_mut::(); for chunk in chunks { // Four values plus at most seven leading bits fit in one load. // At 16 bits, values are byte-aligned, so there are no leading bits. @@ -185,16 +188,12 @@ impl BitUnpacker { *out = load(bit_addr + i * self.num_bits) & self.mask; } } - bit_addr += 4 * self.num_bits; + bit_addr += VALUES_PER_CHUNK * self.num_bits; } for out in remainder { *out = load(bit_addr) & self.mask; bit_addr += self.num_bits; } - for out in tail { - *out = Self::get_slow_path(self.mask, bit_addr >> 3, (bit_addr & 7) as u32, data); - bit_addr += self.num_bits; - } } // Pass the mask by value so specialized callers don't need to materialize a diff --git a/columnar/src/column/mod.rs b/columnar/src/column/mod.rs index 92a08ad33..11c62bf5d 100644 --- a/columnar/src/column/mod.rs +++ b/columnar/src/column/mod.rs @@ -46,10 +46,8 @@ impl Column { impl Column { pub fn to_u64_monotonic(self) -> Column { - let values = monotonic_map_column( - self.values, - StrictlyMonotonicMappingToInternal::::new(), - ); + let values = + monotonic_map_column(self.values, StrictlyMonotonicMappingToInternal::::new()); Column { index: self.index, values, From f5fb40950a1f2aa8ea4788620f9995ee39aeaeaf Mon Sep 17 00:00:00 2001 From: PSeitz Date: Wed, 23 Sep 2026 12:26:51 +0200 Subject: [PATCH 31/49] Update bitpacker/src/bitpacker.rs Co-authored-by: Paul Masurel --- bitpacker/src/bitpacker.rs | 2 ++ 1 file changed, 2 insertions(+) diff --git a/bitpacker/src/bitpacker.rs b/bitpacker/src/bitpacker.rs index dbedd47bf..baf4120b1 100644 --- a/bitpacker/src/bitpacker.rs +++ b/bitpacker/src/bitpacker.rs @@ -132,6 +132,8 @@ impl BitUnpacker { // Fall back for ranges overlapping the end, where an eight-byte load would be partial. let last_bit_addr = (start_idx + output.len() - 1) * self.num_bits; + // The optimization happening below requires reading full 8 bytes word. + // We check that data is long enough to allow for the last read, and if not, fall back for the following safe implementation if (last_bit_addr >> 3) + 8 > data.len() { for (offset, out) in output.iter_mut().enumerate() { *out = self.get((start_idx + offset) as u32, data); From 6023163e37a81b360e692dc82520f33170794647 Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Wed, 23 Sep 2026 13:51:45 +0200 Subject: [PATCH 32/49] Extract bit unpacking load helper with explicit safety contract Replace the closure with an inline(always) function taking the data slice and bit address, making its inputs and safety requirements explicit. --- bitpacker/src/bitpacker.rs | 24 ++++++++++++++++-------- 1 file changed, 16 insertions(+), 8 deletions(-) diff --git a/bitpacker/src/bitpacker.rs b/bitpacker/src/bitpacker.rs index baf4120b1..d2844c148 100644 --- a/bitpacker/src/bitpacker.rs +++ b/bitpacker/src/bitpacker.rs @@ -142,8 +142,11 @@ impl BitUnpacker { } let output_len = output.len(); - let load = |bit_addr: usize| { - // SAFETY: the range-end check above guarantees that this load fits in `data`. + /// # Safety + /// Eight bytes starting at `bit_addr >> 3` must fit in `data`. + #[inline(always)] + unsafe fn load(data: &[u8], bit_addr: usize) -> u64 { + // SAFETY: the caller guarantees that this load fits in `data`. let packed = unsafe { data.as_ptr() .add(bit_addr >> 3) @@ -151,7 +154,7 @@ impl BitUnpacker { .read_unaligned() }; u64::from_le(packed) >> (bit_addr & 7) - }; + } let mut bit_addr = start_idx * self.num_bits; // Tantivy's `COLLECT_BLOCK_BUFFER_LEN` is 64, so optimize its common full-block case by // decoding eight 1-8 bit values per load. Keep this literal in sync with that constant. @@ -159,7 +162,8 @@ impl BitUnpacker { let (chunks, remainder) = output.as_chunks_mut::<8>(); debug_assert!(remainder.is_empty()); for chunk in chunks { - let packed = load(bit_addr); + // SAFETY: the range-end check above guarantees that this load fits in `data`. + let packed = unsafe { load(data, bit_addr) }; for (i, out) in chunk.iter_mut().enumerate() { *out = (packed >> (i * self.num_bits)) & self.mask; } @@ -173,7 +177,8 @@ impl BitUnpacker { // Four values plus at most seven leading bits fit in one load. // At 16 bits, values are byte-aligned, so there are no leading bits. if self.num_bits <= 14 || self.num_bits == 16 { - let packed = load(bit_addr); + // SAFETY: the range-end check above guarantees that this load fits in `data`. + let packed = unsafe { load(data, bit_addr) }; for (i, out) in chunk.iter_mut().enumerate() { *out = (packed >> (i * self.num_bits)) & self.mask; } @@ -181,19 +186,22 @@ impl BitUnpacker { // Two values plus at most seven leading bits fit in one load. // At 32 bits, values are byte-aligned, so there are no leading bits. for (pair_idx, pair) in chunk.as_chunks_mut::<2>().0.iter_mut().enumerate() { - let packed = load(bit_addr + pair_idx * 2 * self.num_bits); + // SAFETY: the range-end check above guarantees that this load fits in `data`. + let packed = unsafe { load(data, bit_addr + pair_idx * 2 * self.num_bits) }; pair[0] = packed & self.mask; pair[1] = (packed >> self.num_bits) & self.mask; } } else { for (i, out) in chunk.iter_mut().enumerate() { - *out = load(bit_addr + i * self.num_bits) & self.mask; + // SAFETY: the range-end check above guarantees that this load fits in `data`. + *out = unsafe { load(data, bit_addr + i * self.num_bits) } & self.mask; } } bit_addr += VALUES_PER_CHUNK * self.num_bits; } for out in remainder { - *out = load(bit_addr) & self.mask; + // SAFETY: the range-end check above guarantees that this load fits in `data`. + *out = unsafe { load(data, bit_addr) } & self.mask; bit_addr += self.num_bits; } } From c321b170a8b1b5c65083375b49e16f0a62e39955 Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Wed, 23 Sep 2026 14:07:05 +0200 Subject: [PATCH 33/49] Clarify chunk size and packed value type in bit unpacking fast path --- bitpacker/src/bitpacker.rs | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/bitpacker/src/bitpacker.rs b/bitpacker/src/bitpacker.rs index d2844c148..48aa13b4a 100644 --- a/bitpacker/src/bitpacker.rs +++ b/bitpacker/src/bitpacker.rs @@ -159,15 +159,16 @@ impl BitUnpacker { // Tantivy's `COLLECT_BLOCK_BUFFER_LEN` is 64, so optimize its common full-block case by // decoding eight 1-8 bit values per load. Keep this literal in sync with that constant. if output_len == 64 && self.num_bits <= 8 { - let (chunks, remainder) = output.as_chunks_mut::<8>(); + const VALUES_PER_CHUNK: usize = 8; + let (chunks, remainder) = output.as_chunks_mut::(); debug_assert!(remainder.is_empty()); for chunk in chunks { // SAFETY: the range-end check above guarantees that this load fits in `data`. - let packed = unsafe { load(data, bit_addr) }; + let packed: u64 = unsafe { load(data, bit_addr) }; for (i, out) in chunk.iter_mut().enumerate() { *out = (packed >> (i * self.num_bits)) & self.mask; } - bit_addr += 8 * self.num_bits; + bit_addr += VALUES_PER_CHUNK * self.num_bits; } return; } From 08d5836b8f24666c6abc8814e3ee12c4078d7229 Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Wed, 23 Sep 2026 14:15:31 +0200 Subject: [PATCH 34/49] Format with nightly rustfmt --- bitpacker/src/bitpacker.rs | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/bitpacker/src/bitpacker.rs b/bitpacker/src/bitpacker.rs index 48aa13b4a..dc192eeda 100644 --- a/bitpacker/src/bitpacker.rs +++ b/bitpacker/src/bitpacker.rs @@ -133,7 +133,8 @@ impl BitUnpacker { // Fall back for ranges overlapping the end, where an eight-byte load would be partial. let last_bit_addr = (start_idx + output.len() - 1) * self.num_bits; // The optimization happening below requires reading full 8 bytes word. - // We check that data is long enough to allow for the last read, and if not, fall back for the following safe implementation + // We check that data is long enough to allow for the last read, and if not, fall back for + // the following safe implementation if (last_bit_addr >> 3) + 8 > data.len() { for (offset, out) in output.iter_mut().enumerate() { *out = self.get((start_idx + offset) as u32, data); From 95dfcd3ada0c0e8f941f9ece772bd95ad7c8e89a Mon Sep 17 00:00:00 2001 From: trinity-1686a Date: Wed, 23 Sep 2026 12:11:10 +0200 Subject: [PATCH 35/49] add bench for automaton sstable streaming --- sstable/benches/stream_bench.rs | 89 ++++++++++++++++++++++++++++++++- 1 file changed, 88 insertions(+), 1 deletion(-) diff --git a/sstable/benches/stream_bench.rs b/sstable/benches/stream_bench.rs index 70dcdd8e3..f8235ceea 100644 --- a/sstable/benches/stream_bench.rs +++ b/sstable/benches/stream_bench.rs @@ -1,13 +1,68 @@ use std::collections::BTreeSet; +use std::hint::black_box; use std::io; use common::file_slice::FileSlice; use criterion::{Criterion, criterion_group, criterion_main}; use rand::rngs::StdRng; use rand::{Rng, SeedableRng}; +use tantivy_fst::Automaton; use tantivy_sstable::{Dictionary, MonotonicU64SSTable}; const CHARSET: &[u8] = b"abcdefghij"; +const AUTOMATON_PREFIX: &[u8] = b"ab"; +const NUM_AUTOMATON_MATCHES: usize = 1_017; + +// Matches `prefix.*`, but only implement can_match/will_always_match if configured to +// +// this allow comparing effects of optimisations depending on these functions +struct HintedPrefixAutomaton<'a> { + prefix: &'a [u8], + can_match_hint: bool, + always_match_hint: bool, +} + +impl<'a> HintedPrefixAutomaton<'a> { + fn new(prefix: &'a [u8], can_match_hint: bool, always_match_hint: bool) -> Self { + Self { + prefix, + can_match_hint, + always_match_hint, + } + } +} + +impl Automaton for HintedPrefixAutomaton<'_> { + type State = Option; + + fn start(&self) -> Self::State { + Some(0) + } + + fn is_match(&self, state: &Self::State) -> bool { + *state == Some(self.prefix.len()) + } + + fn can_match(&self, state: &Self::State) -> bool { + !self.can_match_hint || state.is_some() + } + + fn will_always_match(&self, state: &Self::State) -> bool { + self.always_match_hint && self.is_match(state) + } + + fn accept(&self, state: &Self::State, byte: u8) -> Self::State { + let Some(pos) = *state else { return None }; + if pos == self.prefix.len() { + return Some(pos); + } + if self.prefix[pos] == byte { + Some(pos + 1) + } else { + None + } + } +} fn generate_key(rng: &mut impl Rng) -> String { let len = rng.random_range(3..12); @@ -56,6 +111,26 @@ fn stream_bench( count } +fn automaton_bench( + dictionary: &Dictionary, + can_match_hint: bool, + always_match_hint: bool, +) -> usize { + let mut stream = dictionary + .search(HintedPrefixAutomaton::new( + AUTOMATON_PREFIX, + black_box(can_match_hint), + black_box(always_match_hint), + )) + .into_stream() + .unwrap(); + let mut count = 0; + while stream.advance() { + count += 1; + } + count +} + pub fn criterion_benchmark(c: &mut Criterion) { let dict = prepare_sstable().unwrap(); c.bench_function("short_scan_init", |b| { @@ -63,7 +138,7 @@ pub fn criterion_benchmark(c: &mut Criterion) { }); c.bench_function("short_scan_init_and_scan", |b| { b.iter(|| { - assert_eq!(stream_bench(&dict, b"fa", b"faz", true), 971); + assert_eq!(stream_bench(&dict, b"fa", b"faz", true), 1051); }) }); c.bench_function("full_scan_init_and_scan_full_with_bound", |b| { @@ -81,6 +156,18 @@ pub fn criterion_benchmark(c: &mut Criterion) { count }) }); + c.bench_function("full_scan_prefix_automaton_no_hints", |b| { + b.iter(|| assert_eq!(automaton_bench(&dict, false, false), NUM_AUTOMATON_MATCHES)) + }); + c.bench_function("full_scan_prefix_automaton_can_match_hint_only", |b| { + b.iter(|| assert_eq!(automaton_bench(&dict, true, false), NUM_AUTOMATON_MATCHES)) + }); + c.bench_function("full_scan_prefix_automaton_always_match_hint_only", |b| { + b.iter(|| assert_eq!(automaton_bench(&dict, false, true), NUM_AUTOMATON_MATCHES)) + }); + c.bench_function("full_scan_prefix_automaton_both_hints", |b| { + b.iter(|| assert_eq!(automaton_bench(&dict, true, true), NUM_AUTOMATON_MATCHES)) + }); } criterion_group!(benches, criterion_benchmark); From aac95014b3a2bd203cbcc6329e3f2a4fc78b2989 Mon Sep 17 00:00:00 2001 From: trinity-1686a Date: Wed, 23 Sep 2026 12:16:16 +0200 Subject: [PATCH 36/49] remove lower_bound check from streamer hotloop --- sstable/src/streamer.rs | 100 ++++++++++++++++++++++++++++++---------- 1 file changed, 75 insertions(+), 25 deletions(-) diff --git a/sstable/src/streamer.rs b/sstable/src/streamer.rs index 9203b3d0a..4ed7522f1 100644 --- a/sstable/src/streamer.rs +++ b/sstable/src/streamer.rs @@ -130,6 +130,7 @@ where delta_reader, key: Vec::new(), term_ord: first_term.checked_sub(1), + lower_bound_reached: self.lower == Bound::Unbounded, lower_bound: self.lower, upper_bound: self.upper, _lifetime: std::marker::PhantomData, @@ -176,6 +177,7 @@ where upper_bound: Bound>, // this field is used to please the type-interface of a dictionary in tantivy _lifetime: std::marker::PhantomData<&'a ()>, + lower_bound_reached: bool, } impl Streamer<'_, TSSTable, AlwaysMatch> @@ -188,6 +190,7 @@ where TSSTable: SSTable delta_reader: DeltaReader::empty(), key: Vec::new(), term_ord: None, + lower_bound_reached: true, lower_bound: Bound::Unbounded, upper_bound: Bound::Unbounded, _lifetime: std::marker::PhantomData, @@ -201,40 +204,84 @@ where A::State: Clone, TSSTable: SSTable, { - /// Advance position the stream on the next item. - /// Before the first call to `.advance()`, the stream - /// is an uninitialized state. - pub fn advance(&mut self) -> bool { - while self.delta_reader.advance().unwrap() { - // An automaton prunes whole blocks, so the ordinal is not simply the previous one - // plus one: on entering a new slice it jumps to that slice's first term ordinal. - // Counting alone would report a term's position among the blocks actually scanned. - self.term_ord = Some(match self.delta_reader.take_first_ordinal() { - Some(first_ordinal) => first_ordinal, - None => self - .term_ord - .map(|term_ord| term_ord + 1u64) - .unwrap_or(0u64), - }); + #[inline(always)] + fn advance_delta_reader(&mut self) -> bool { + if !self.delta_reader.advance().unwrap() { + return false; + } + // An automaton prunes whole blocks, so the ordinal is not simply the previous one + // plus one: on entering a new slice it jumps to that slice's first term ordinal. + // Counting alone would report a term's position among the blocks actually scanned. + self.term_ord = Some(match self.delta_reader.take_first_ordinal() { + Some(first_ordinal) => first_ordinal, + None => self + .term_ord + .map(|term_ord| term_ord + 1u64) + .unwrap_or(0u64), + }); + true + } + + /// Make progress up to the lower bound + /// + /// Returns whether the reader was positioned on a key matching the lower bound. + /// If false, the delta_reader has been exhausted without finding such a key. + fn initialize(&mut self) -> bool { + debug_assert!(!self.lower_bound_reached); + while self.advance_delta_reader() { let common_prefix_len = self.delta_reader.common_prefix_len(); - self.states.truncate(common_prefix_len + 1); self.key.truncate(common_prefix_len); - let mut state: A::State = self.states.last().unwrap().clone(); - for &b in self.delta_reader.suffix() { - state = self.automaton.accept(&state, b); - self.states.push(state.clone()); - } self.key.extend_from_slice(self.delta_reader.suffix()); + let match_lower_bound = match &self.lower_bound { Bound::Unbounded => true, Bound::Included(lower_bound_key) => lower_bound_key[..] <= self.key[..], Bound::Excluded(lower_bound_key) => lower_bound_key[..] < self.key[..], }; - if !match_lower_bound { - continue; + if match_lower_bound { + let mut state: A::State = self.states.last().unwrap().clone(); + for b in &self.key { + state = self.automaton.accept(&state, *b); + self.states.push(state.clone()); + } + self.lower_bound_reached = true; + return true; } - // We match the lower key once. All subsequent keys will pass that bar. - self.lower_bound = Bound::Unbounded; + } + self.lower_bound_reached = true; + false + } + + /// Advance position the stream on the next item. + /// Before the first call to `.advance()`, the stream + /// is an uninitialized state. + pub fn advance(&mut self) -> bool { + let skip_first_delta_advance = if self.lower_bound_reached { + false + } else { + if !self.initialize() { + // no key higher than lower-bound at all + return false; + } + true + }; + if !skip_first_delta_advance { + if !self.advance_delta_reader() { + return false; + } + } + loop { + let common_prefix_len = self.delta_reader.common_prefix_len(); + self.states.truncate(common_prefix_len + 1); + let mut state: A::State = self.states.last().unwrap().clone(); + for &b in self.delta_reader.suffix() { + state = self.automaton.accept(&state, b); + self.states.push(state.clone()); + } + + self.key.truncate(common_prefix_len); + self.key.extend_from_slice(self.delta_reader.suffix()); + let match_upper_bound = match &self.upper_bound { Bound::Unbounded => true, Bound::Included(upper_bound_key) => upper_bound_key[..] >= self.key[..], @@ -246,6 +293,9 @@ where if self.automaton.is_match(&state) { return true; } + if !self.advance_delta_reader() { + break; + } } false } From bec1b08bdf9cc76e17915237f4f6863524307c0b Mon Sep 17 00:00:00 2001 From: trinity-1686a Date: Wed, 23 Sep 2026 15:06:36 +0200 Subject: [PATCH 37/49] use delta-based prefix matcher it does less slice comparisons overall, which improves throughput --- sstable/src/delta.rs | 83 ++++++++++++++++++++++++++++++++++- sstable/src/dictionary.rs | 41 +++++------------ sstable/src/streamer.rs | 92 +++++++++++++++++++++++++++++---------- 3 files changed, 160 insertions(+), 56 deletions(-) diff --git a/sstable/src/delta.rs b/sstable/src/delta.rs index 97d868e4e..df43a2a18 100644 --- a/sstable/src/delta.rs +++ b/sstable/src/delta.rs @@ -1,3 +1,4 @@ +use std::cmp::Ordering; use std::io::{self, BufWriter, Write}; use std::ops::Range; @@ -12,6 +13,64 @@ const FOUR_BIT_LIMITS: usize = 1 << 4; const VINT_MODE: u8 = 1u8; const BLOCK_LEN: usize = 4_000; +/// Incrementally compares delta-encoded keys with a fixed target key. +pub(crate) struct DeltaKeyComparator { + num_matching_bytes: usize, +} + +impl DeltaKeyComparator { + pub(crate) fn new() -> Self { + DeltaKeyComparator { + num_matching_bytes: 0, + } + } + + #[inline(always)] + pub(crate) fn compare( + &mut self, + target: &[u8], + common_prefix_len: usize, + suffix: &[u8], + ) -> Ordering { + match common_prefix_len.cmp(&self.num_matching_bytes) { + // popped bytes already matched => too far + Ordering::Less => return Ordering::Greater, + Ordering::Equal => (), + // the ok prefix is less than current entry prefix => continue to next element + Ordering::Greater => return Ordering::Less, + } + + for (key_byte, target_byte) in suffix.iter().zip(&target[self.num_matching_bytes..]) { + match key_byte.cmp(target_byte) { + Ordering::Equal => self.num_matching_bytes += 1, + ordering => return ordering, + } + } + + (common_prefix_len + suffix.len()).cmp(&target.len()) + } + + #[inline(always)] + pub(crate) fn compare_across_blocks( + &mut self, + target: &[u8], + common_prefix_len: usize, + suffix: &[u8], + ) -> Ordering { + // blocks are independent. On each new block we get a common_prefix_len=0 entry. + // reset our state with it + if common_prefix_len == 0 { + self.num_matching_bytes = target + .iter() + .zip(suffix) + .take_while(|(target_byte, key_byte)| target_byte == key_byte) + .count(); + return suffix.cmp(target); + } + self.compare(target, common_prefix_len, suffix) + } +} + pub struct DeltaWriter where W: io::Write { @@ -241,7 +300,9 @@ where TValueReader: value::ValueReader #[cfg(test)] mod tests { - use super::DeltaReader; + use std::cmp::Ordering; + + use super::{DeltaKeyComparator, DeltaReader}; use crate::value::U64MonotonicValueReader; #[test] @@ -249,4 +310,24 @@ mod tests { let mut delta_reader: DeltaReader = DeltaReader::empty(); assert!(!delta_reader.advance().unwrap()); } + + #[test] + fn test_delta_key_comparator_across_block_reset() { + let mut comparator = DeltaKeyComparator::new(); + let target = b"bba"; + + assert_eq!( + comparator.compare_across_blocks(target, 0, b"baaaaa"), + Ordering::Less + ); + assert_eq!( + comparator.compare_across_blocks(target, 2, b"baaa"), + Ordering::Less + ); + // A zero-length common prefix marks a block reset, so the suffix is a complete key. + assert_eq!( + comparator.compare_across_blocks(target, 0, b"bbbaaa"), + Ordering::Greater + ); + } } diff --git a/sstable/src/dictionary.rs b/sstable/src/dictionary.rs index 5de411467..69b57053f 100644 --- a/sstable/src/dictionary.rs +++ b/sstable/src/dictionary.rs @@ -14,6 +14,7 @@ use itertools::Itertools; use tantivy_fst::Automaton; use tantivy_fst::automaton::AlwaysMatch; +use crate::delta::DeltaKeyComparator; use crate::streamer::{Streamer, StreamerBuilder}; use crate::{BlockAddr, DeltaReader, Reader, SSTable, SSTableIndex, TermOrdinal, VoidSSTable}; @@ -356,41 +357,19 @@ impl Dictionary { ) -> io::Result { let mut term_ord = 0; let key_bytes = key.as_ref(); - let mut ok_bytes = 0; + let mut key_comparator = DeltaKeyComparator::new(); while sstable_delta_reader.advance()? { - let prefix_len = sstable_delta_reader.common_prefix_len(); - let suffix = sstable_delta_reader.suffix(); - - match prefix_len.cmp(&ok_bytes) { - Ordering::Less => return Ok(TermOrdHit::Next(term_ord)), /* popped bytes already matched => too far */ - Ordering::Equal => (), - Ordering::Greater => { - // the ok prefix is less than current entry prefix => continue to next elem + match key_comparator.compare( + key_bytes, + sstable_delta_reader.common_prefix_len(), + sstable_delta_reader.suffix(), + ) { + Ordering::Less => { term_ord += 1; - continue; } + Ordering::Equal => return Ok(TermOrdHit::Exact(term_ord)), + Ordering::Greater => return Ok(TermOrdHit::Next(term_ord)), } - - // we have ok_bytes byte of common prefix, check if this key adds more - for (key_byte, suffix_byte) in key_bytes[ok_bytes..].iter().zip(suffix) { - match suffix_byte.cmp(key_byte) { - Ordering::Less => break, // byte too small - Ordering::Equal => ok_bytes += 1, // new matching - // byte - Ordering::Greater => return Ok(TermOrdHit::Next(term_ord)), // too far - } - } - - if ok_bytes == key_bytes.len() { - if prefix_len + suffix.len() == ok_bytes { - return Ok(TermOrdHit::Exact(term_ord)); - } else { - // current key is a prefix of current element, not a match - return Ok(TermOrdHit::Next(term_ord)); - } - } - - term_ord += 1; } Ok(TermOrdHit::Next(term_ord)) diff --git a/sstable/src/streamer.rs b/sstable/src/streamer.rs index 4ed7522f1..c9716af8a 100644 --- a/sstable/src/streamer.rs +++ b/sstable/src/streamer.rs @@ -1,9 +1,11 @@ +use std::cmp::Ordering; use std::io; use std::ops::Bound; use tantivy_fst::Automaton; use tantivy_fst::automaton::AlwaysMatch; +use crate::delta::DeltaKeyComparator; use crate::dictionary::Dictionary; use crate::{DeltaReader, SSTable, TermOrdinal}; @@ -30,6 +32,22 @@ fn bound_as_byte_slice(bound: &Bound>) -> Bound<&[u8]> { } } +#[inline(always)] +fn matches_upper_bound( + comparator: &mut DeltaKeyComparator, + upper_bound: &Bound>, + common_prefix_len: usize, + suffix: &[u8], +) -> bool { + let (upper_bound_key, inclusive) = match upper_bound { + Bound::Unbounded => return true, + Bound::Included(upper_bound_key) => (upper_bound_key, true), + Bound::Excluded(upper_bound_key) => (upper_bound_key, false), + }; + let ordering = comparator.compare_across_blocks(upper_bound_key, common_prefix_len, suffix); + ordering == Ordering::Less || inclusive && ordering == Ordering::Equal +} + impl<'a, TSSTable, A> StreamerBuilder<'a, TSSTable, A> where A: Automaton, @@ -133,6 +151,7 @@ where lower_bound_reached: self.lower == Bound::Unbounded, lower_bound: self.lower, upper_bound: self.upper, + upper_bound_comparator: DeltaKeyComparator::new(), _lifetime: std::marker::PhantomData, }) } @@ -175,6 +194,7 @@ where term_ord: Option, lower_bound: Bound>, upper_bound: Bound>, + upper_bound_comparator: DeltaKeyComparator, // this field is used to please the type-interface of a dictionary in tantivy _lifetime: std::marker::PhantomData<&'a ()>, lower_bound_reached: bool, @@ -193,6 +213,7 @@ where TSSTable: SSTable lower_bound_reached: true, lower_bound: Bound::Unbounded, upper_bound: Bound::Unbounded, + upper_bound_comparator: DeltaKeyComparator::new(), _lifetime: std::marker::PhantomData, } } @@ -228,17 +249,27 @@ where /// If false, the delta_reader has been exhausted without finding such a key. fn initialize(&mut self) -> bool { debug_assert!(!self.lower_bound_reached); + let mut lower_bound_comparator = DeltaKeyComparator::new(); while self.advance_delta_reader() { let common_prefix_len = self.delta_reader.common_prefix_len(); - self.key.truncate(common_prefix_len); - self.key.extend_from_slice(self.delta_reader.suffix()); - - let match_lower_bound = match &self.lower_bound { - Bound::Unbounded => true, - Bound::Included(lower_bound_key) => lower_bound_key[..] <= self.key[..], - Bound::Excluded(lower_bound_key) => lower_bound_key[..] < self.key[..], + let suffix = self.delta_reader.suffix(); + let (lower_bound_key, inclusive) = match &self.lower_bound { + Bound::Unbounded => unreachable!("unbounded streamers do not need initialization"), + Bound::Included(lower_bound_key) => (lower_bound_key, true), + Bound::Excluded(lower_bound_key) => (lower_bound_key, false), }; + let ordering = lower_bound_comparator.compare_across_blocks( + lower_bound_key, + common_prefix_len, + suffix, + ); + let match_lower_bound = + ordering == Ordering::Greater || inclusive && ordering == Ordering::Equal; if match_lower_bound { + self.key.clear(); + self.key + .extend_from_slice(&lower_bound_key[..common_prefix_len]); + self.key.extend_from_slice(suffix); let mut state: A::State = self.states.last().unwrap().clone(); for b in &self.key { state = self.automaton.accept(&state, *b); @@ -256,23 +287,34 @@ where /// Before the first call to `.advance()`, the stream /// is an uninitialized state. pub fn advance(&mut self) -> bool { - let skip_first_delta_advance = if self.lower_bound_reached { - false - } else { + if !self.lower_bound_reached { if !self.initialize() { // no key higher than lower-bound at all return false; } - true - }; - if !skip_first_delta_advance { - if !self.advance_delta_reader() { + if !matches_upper_bound( + &mut self.upper_bound_comparator, + &self.upper_bound, + 0, + &self.key, + ) { return false; } + if self.automaton.is_match(self.states.last().unwrap()) { + return true; + } } - loop { + + while self.advance_delta_reader() { let common_prefix_len = self.delta_reader.common_prefix_len(); self.states.truncate(common_prefix_len + 1); + + // TODO we could use const-generics to remove this bit when the automaton is an always + // match one + // TODO we could detect when we reach an always match state, and no longer check state + // until we truncate enough (e.g. for the regex `my_prefix.*`, or `.*my_infix.*`) + // TODO we could detect when we reach a !can_match, and skip both state and key + // computation until we truncate that can_t_match out of our state let mut state: A::State = self.states.last().unwrap().clone(); for &b in self.delta_reader.suffix() { state = self.automaton.accept(&state, b); @@ -282,20 +324,22 @@ where self.key.truncate(common_prefix_len); self.key.extend_from_slice(self.delta_reader.suffix()); - let match_upper_bound = match &self.upper_bound { - Bound::Unbounded => true, - Bound::Included(upper_bound_key) => upper_bound_key[..] >= self.key[..], - Bound::Excluded(upper_bound_key) => upper_bound_key[..] > self.key[..], - }; - if !match_upper_bound { + // TODO there is an idea where we only look at the upper bound when our delta_reader + // reached the last block (if we pruned blocks beforehand (do we always?) we cannot + // find that key before that block) + // TODO we could use const-generics to remove this branch from the loop when no + // upper bound is used + if !matches_upper_bound( + &mut self.upper_bound_comparator, + &self.upper_bound, + common_prefix_len, + self.delta_reader.suffix(), + ) { return false; } if self.automaton.is_match(&state) { return true; } - if !self.advance_delta_reader() { - break; - } } false } From 2bb52f1460c2c39499ed34447b8273ba0e24361e Mon Sep 17 00:00:00 2001 From: trinity-1686a Date: Wed, 23 Sep 2026 15:48:23 +0200 Subject: [PATCH 38/49] hide some branches behind const generics --- sstable/src/streamer.rs | 44 +++++++++++++++++++++++++++-------------- 1 file changed, 29 insertions(+), 15 deletions(-) diff --git a/sstable/src/streamer.rs b/sstable/src/streamer.rs index c9716af8a..b92b958a9 100644 --- a/sstable/src/streamer.rs +++ b/sstable/src/streamer.rs @@ -206,7 +206,7 @@ where TSSTable: SSTable pub fn empty() -> Self { Streamer { automaton: AlwaysMatch, - states: Vec::new(), + states: vec![AlwaysMatch.start()], delta_reader: DeltaReader::empty(), key: Vec::new(), term_ord: None, @@ -305,20 +305,34 @@ where } } + // (always_match, no_bound) + match ( + self.automaton + .will_always_match(&self.states.first().unwrap()), + self.upper_bound == Bound::Unbounded, + ) { + (true, true) => self.advance_inner::(), + (true, false) => self.advance_inner::(), + (false, true) => self.advance_inner::(), + (false, false) => self.advance_inner::(), + } + } + + fn advance_inner(&mut self) -> bool { while self.advance_delta_reader() { let common_prefix_len = self.delta_reader.common_prefix_len(); self.states.truncate(common_prefix_len + 1); - // TODO we could use const-generics to remove this bit when the automaton is an always - // match one // TODO we could detect when we reach an always match state, and no longer check state // until we truncate enough (e.g. for the regex `my_prefix.*`, or `.*my_infix.*`) // TODO we could detect when we reach a !can_match, and skip both state and key // computation until we truncate that can_t_match out of our state let mut state: A::State = self.states.last().unwrap().clone(); - for &b in self.delta_reader.suffix() { - state = self.automaton.accept(&state, b); - self.states.push(state.clone()); + if !AUTOMATON_MATCH { + for &b in self.delta_reader.suffix() { + state = self.automaton.accept(&state, b); + self.states.push(state.clone()); + } } self.key.truncate(common_prefix_len); @@ -327,17 +341,17 @@ where // TODO there is an idea where we only look at the upper bound when our delta_reader // reached the last block (if we pruned blocks beforehand (do we always?) we cannot // find that key before that block) - // TODO we could use const-generics to remove this branch from the loop when no - // upper bound is used - if !matches_upper_bound( - &mut self.upper_bound_comparator, - &self.upper_bound, - common_prefix_len, - self.delta_reader.suffix(), - ) { + if !NO_BOUND + && !matches_upper_bound( + &mut self.upper_bound_comparator, + &self.upper_bound, + common_prefix_len, + self.delta_reader.suffix(), + ) + { return false; } - if self.automaton.is_match(&state) { + if AUTOMATON_MATCH || self.automaton.is_match(&state) { return true; } } From 6c311b311dd190df265e0771d8a8c8e07361873e Mon Sep 17 00:00:00 2001 From: trinity-1686a Date: Thu, 24 Sep 2026 19:09:25 +0200 Subject: [PATCH 39/49] use will_always_match hint --- sstable/src/streamer.rs | 119 +++++++++++++++++++++++++++------------- 1 file changed, 82 insertions(+), 37 deletions(-) diff --git a/sstable/src/streamer.rs b/sstable/src/streamer.rs index b92b958a9..9e5d11541 100644 --- a/sstable/src/streamer.rs +++ b/sstable/src/streamer.rs @@ -142,9 +142,16 @@ where Bound::Unbounded => 0, }; + let always_match_at = if self.automaton.will_always_match(&start_state) { + Some(0) + } else { + None + }; + Ok(Streamer { automaton: self.automaton, states: vec![start_state], + always_match_at, delta_reader, key: Vec::new(), term_ord: first_term.checked_sub(1), @@ -198,6 +205,7 @@ where // this field is used to please the type-interface of a dictionary in tantivy _lifetime: std::marker::PhantomData<&'a ()>, lower_bound_reached: bool, + always_match_at: Option, } impl Streamer<'_, TSSTable, AlwaysMatch> @@ -207,6 +215,7 @@ where TSSTable: SSTable Streamer { automaton: AlwaysMatch, states: vec![AlwaysMatch.start()], + always_match_at: Some(0), delta_reader: DeltaReader::empty(), key: Vec::new(), term_ord: None, @@ -305,57 +314,93 @@ where } } - // (always_match, no_bound) match ( + // we could check always_match_at == Some(0), but this actually gets + // inlined into `true` with AlwaysMatch, which is even faster self.automaton .will_always_match(&self.states.first().unwrap()), self.upper_bound == Bound::Unbounded, ) { - (true, true) => self.advance_inner::(), - (true, false) => self.advance_inner::(), - (false, true) => self.advance_inner::(), - (false, false) => self.advance_inner::(), + (true, true) => self.advance_always_match::(), + (true, false) => self.advance_always_match::(), + (false, true) => self.advance_with_automaton::(), + (false, false) => self.advance_with_automaton::(), } } - fn advance_inner(&mut self) -> bool { - while self.advance_delta_reader() { - let common_prefix_len = self.delta_reader.common_prefix_len(); - self.states.truncate(common_prefix_len + 1); + fn advance_always_match(&mut self) -> bool { + if !self.advance_delta_reader() { + return false; + } + self.reconstruct_key_and_check_upper_bound::() + } - // TODO we could detect when we reach an always match state, and no longer check state - // until we truncate enough (e.g. for the regex `my_prefix.*`, or `.*my_infix.*`) - // TODO we could detect when we reach a !can_match, and skip both state and key - // computation until we truncate that can_t_match out of our state - let mut state: A::State = self.states.last().unwrap().clone(); - if !AUTOMATON_MATCH { - for &b in self.delta_reader.suffix() { - state = self.automaton.accept(&state, b); - self.states.push(state.clone()); - } - } - - self.key.truncate(common_prefix_len); - self.key.extend_from_slice(self.delta_reader.suffix()); - - // TODO there is an idea where we only look at the upper bound when our delta_reader - // reached the last block (if we pruned blocks beforehand (do we always?) we cannot - // find that key before that block) - if !NO_BOUND - && !matches_upper_bound( - &mut self.upper_bound_comparator, - &self.upper_bound, - common_prefix_len, - self.delta_reader.suffix(), - ) - { + fn advance_with_automaton(&mut self) -> bool { + // fast path, check if prefix always match and we can skip Vec management + if let Some(always_match_at) = self.always_match_at.take() { + if !self.advance_delta_reader() { return false; } - if AUTOMATON_MATCH || self.automaton.is_match(&state) { + let common_prefix_len = self.delta_reader.common_prefix_len(); + if always_match_at <= common_prefix_len { + self.always_match_at = Some(always_match_at); + return self.reconstruct_key_and_check_upper_bound::(); + } + } else if !self.advance_delta_reader() { + return false; + } + + loop { + let common_prefix_len = self.delta_reader.common_prefix_len(); + self.states.truncate(common_prefix_len + 1); + // TODO we could detect when we reach a !can_match, and skip both state and key + // computation until we truncate that can_t_match out of our state. it's already + // done at the block layer, so not as important + let mut state: A::State = self.states.last().unwrap().clone(); + for &b in self.delta_reader.suffix() { + state = self.automaton.accept(&state, b); + self.states.push(state.clone()); + } + let matches = self.automaton.is_match(&state); + if matches { + self.always_match_at = self + .states + .iter() + .enumerate() + .rev() + .take_while(|(_i, state)| self.automaton.will_always_match(state)) + .last() + .map(|(i, _state)| i); + } + + if !self.reconstruct_key_and_check_upper_bound::() { + return false; + } + if matches { return true; } + if !self.advance_delta_reader() { + return false; + } } - false + } + + #[inline(always)] + fn reconstruct_key_and_check_upper_bound(&mut self) -> bool { + let common_prefix_len = self.delta_reader.common_prefix_len(); + self.key.truncate(common_prefix_len); + self.key.extend_from_slice(self.delta_reader.suffix()); + + // TODO there is an idea where we only look at the upper bound when our delta_reader + // reached the last block (if we pruned blocks beforehand (do we always?) we cannot + // find that key before that block) + NO_BOUND + || matches_upper_bound( + &mut self.upper_bound_comparator, + &self.upper_bound, + common_prefix_len, + self.delta_reader.suffix(), + ) } /// Returns the `TermOrdinal` of the given term. From b9125aad55e8dd3b4e4d6ef76981cc155a8c5397 Mon Sep 17 00:00:00 2001 From: palmoni5 Date: Mon, 28 Sep 2026 12:37:50 +0300 Subject: [PATCH 40/49] Compile RegexPhraseQuery regexes once per weight (#3135) RegexPhraseWeight::phrase_scorer compiled every phrase term's regex again for each segment. Determinizing a wide pattern (e.g. an alternation of typo or morphology expansions) costs milliseconds to tens of milliseconds, so on a multi-segment index the query was dominated by recompiling the same automata. The regexes are now compiled in RegexPhraseQuery::regex_phrase_weight and shared through Arc. An invalid pattern is now reported when the weight is built, instead of when the first segment is scored. --- src/query/phrase_query/regex_phrase_query.rs | 18 ++++++++++++- src/query/phrase_query/regex_phrase_weight.rs | 25 +++++++++++++------ 2 files changed, 35 insertions(+), 8 deletions(-) diff --git a/src/query/phrase_query/regex_phrase_query.rs b/src/query/phrase_query/regex_phrase_query.rs index 98e07d399..6b9e8bb34 100644 --- a/src/query/phrase_query/regex_phrase_query.rs +++ b/src/query/phrase_query/regex_phrase_query.rs @@ -1,3 +1,7 @@ +use std::sync::Arc; + +use tantivy_fst::Regex; + use super::regex_phrase_weight::RegexPhraseWeight; use crate::query::bm25::Bm25Weight; use crate::query::{EnableScoring, Query, Weight}; @@ -149,9 +153,21 @@ impl RegexPhraseQuery { } => Some(Bm25Weight::for_terms(statistics_provider, &terms)?), EnableScoring::Disabled { .. } => None, }; + // Compiled once here rather than per segment: determinizing a large + // pattern can dominate the cost of the query. + let phrase_terms = self + .phrase_terms + .iter() + .map(|(offset, term)| { + let regex = Regex::new(term).map_err(|e| { + crate::TantivyError::InvalidArgument(format!("Invalid regex: {e}")) + })?; + Ok((*offset, Arc::new(regex))) + }) + .collect::>>()?; let weight = RegexPhraseWeight::new( self.field, - self.phrase_terms.clone(), + phrase_terms, bm25_weight_opt, self.max_expansions, self.slop, diff --git a/src/query/phrase_query/regex_phrase_weight.rs b/src/query/phrase_query/regex_phrase_weight.rs index 9cefc555a..e7aa1ab12 100644 --- a/src/query/phrase_query/regex_phrase_weight.rs +++ b/src/query/phrase_query/regex_phrase_weight.rs @@ -20,7 +20,7 @@ type UnionType = SimpleUnion>; /// See RegexPhraseWeight::get_union_from_term_infos for some design decisions. pub struct RegexPhraseWeight { field: Field, - phrase_terms: Vec<(usize, String)>, + phrase_terms: Vec<(usize, Arc)>, similarity_weight_opt: Option, slop: u32, max_expansions: u32, @@ -31,7 +31,7 @@ impl RegexPhraseWeight { /// If `similarity_weight_opt` is None, then scoring is disabled pub fn new( field: Field, - phrase_terms: Vec<(usize, String)>, + phrase_terms: Vec<(usize, Arc)>, similarity_weight_opt: Option, max_expansions: u32, slop: u32, @@ -67,12 +67,9 @@ impl RegexPhraseWeight { let mut posting_lists = Vec::new(); let inverted_index = reader.inverted_index(self.field)?; let mut num_terms = 0; - for &(offset, ref term) in &self.phrase_terms { - let regex = Regex::new(term) - .map_err(|e| crate::TantivyError::InvalidArgument(format!("Invalid regex: {e}")))?; - + for &(offset, ref regex) in &self.phrase_terms { let automaton: AutomatonWeight = - AutomatonWeight::new(self.field, Arc::new(regex)); + AutomatonWeight::new(self.field, Arc::clone(regex)); let term_infos = automaton.get_match_term_infos(reader)?; // If term_infos is empty, the phrase can not match any documents. if term_infos.is_empty() { @@ -351,6 +348,20 @@ mod tests { } } + #[test] + pub fn test_phrase_regex_invalid_pattern_fails_at_weight() -> crate::Result<()> { + let index = create_index(&["a b"])?; + let text_field = index.schema().get_field("text").unwrap(); + let searcher = index.reader()?.searcher(); + let phrase_query = RegexPhraseQuery::new(text_field, vec!["a".into(), "(".into()]); + let enable_scoring = EnableScoring::enabled_from_searcher(&searcher); + assert!(matches!( + phrase_query.regex_phrase_weight(enable_scoring), + Err(crate::TantivyError::InvalidArgument(_)) + )); + Ok(()) + } + #[test] pub fn test_phrase_count() -> crate::Result<()> { let index = create_index(&["a c", "a a b d a b c", " a b"])?; From 047464cf92e5a31d02a696f5158e45f7d34c67eb Mon Sep 17 00:00:00 2001 From: Paul Masurel Date: Mon, 28 Sep 2026 18:46:37 +0200 Subject: [PATCH 41/49] 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), } From 1d294ea6dc5cb4bef3725605a49d8b4519a3cbf7 Mon Sep 17 00:00:00 2001 From: palmoni5 Date: Tue, 29 Sep 2026 10:19:07 +0300 Subject: [PATCH 42/49] Cache RegexPhraseQuery's compiled regexes in the query (#3137) Following up on #3135, the regexes are compiled on first use and kept in the query, so a query reused across searchers, or cloned, determinizes each pattern once. RegexPhraseQuery::regexes exposes them for callers that inspect the patterns before searching, e.g. to count a phrase's per-segment expansions against max_expansions, instead of compiling them a second time. --- src/query/phrase_query/regex_phrase_query.rs | 48 +++++++++++++++---- src/query/phrase_query/regex_phrase_weight.rs | 20 ++++++++ 2 files changed, 59 insertions(+), 9 deletions(-) diff --git a/src/query/phrase_query/regex_phrase_query.rs b/src/query/phrase_query/regex_phrase_query.rs index 6b9e8bb34..1310871ca 100644 --- a/src/query/phrase_query/regex_phrase_query.rs +++ b/src/query/phrase_query/regex_phrase_query.rs @@ -1,5 +1,7 @@ +use std::fmt; use std::sync::Arc; +use once_cell::sync::OnceCell; use tantivy_fst::Regex; use super::regex_phrase_weight::RegexPhraseWeight; @@ -29,6 +31,21 @@ pub struct RegexPhraseQuery { phrase_terms: Vec<(usize, String)>, slop: u32, max_expansions: u32, + regexes: CompiledRegexes, +} + +/// The compiled `phrase_terms`, built on first use. Its `Debug` omits the automata. +#[derive(Clone, Default)] +struct CompiledRegexes(OnceCell>>); + +impl fmt::Debug for CompiledRegexes { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(if self.0.get().is_some() { + "CompiledRegexes(compiled)" + } else { + "CompiledRegexes(pending)" + }) + } } /// Transform a wildcard query to a regex string. @@ -75,6 +92,7 @@ impl RegexPhraseQuery { phrase_terms: terms, slop, max_expansions: 1 << 14, + regexes: CompiledRegexes::default(), } } @@ -115,6 +133,24 @@ impl RegexPhraseQuery { .collect::>() } + /// The compiled regex of each phrase term, in offset order. + /// + /// Compiled on first call and cached, so a query that is reused across searchers, + /// or inspected before searching, determinizes each pattern once. + pub fn regexes(&self) -> crate::Result<&[Arc]> { + let regexes = self.regexes.0.get_or_try_init(|| { + self.phrase_terms + .iter() + .map(|(_, term)| { + Regex::new(term).map(Arc::new).map_err(|e| { + crate::TantivyError::InvalidArgument(format!("Invalid regex: {e}")) + }) + }) + .collect::>>() + })?; + Ok(regexes) + } + /// Returns the [`RegexPhraseWeight`] for the given phrase query given a specific `searcher`. /// /// This function is the same as [`Query::weight()`] except it returns @@ -153,18 +189,12 @@ impl RegexPhraseQuery { } => Some(Bm25Weight::for_terms(statistics_provider, &terms)?), EnableScoring::Disabled { .. } => None, }; - // Compiled once here rather than per segment: determinizing a large - // pattern can dominate the cost of the query. let phrase_terms = self .phrase_terms .iter() - .map(|(offset, term)| { - let regex = Regex::new(term).map_err(|e| { - crate::TantivyError::InvalidArgument(format!("Invalid regex: {e}")) - })?; - Ok((*offset, Arc::new(regex))) - }) - .collect::>>()?; + .map(|(offset, _)| *offset) + .zip(self.regexes()?.iter().cloned()) + .collect(); let weight = RegexPhraseWeight::new( self.field, phrase_terms, diff --git a/src/query/phrase_query/regex_phrase_weight.rs b/src/query/phrase_query/regex_phrase_weight.rs index e7aa1ab12..42018f832 100644 --- a/src/query/phrase_query/regex_phrase_weight.rs +++ b/src/query/phrase_query/regex_phrase_weight.rs @@ -362,6 +362,26 @@ mod tests { Ok(()) } + #[test] + pub fn test_phrase_regexes_are_compiled_once_and_shared() -> crate::Result<()> { + let index = create_index(&["a b"])?; + let text_field = index.schema().get_field("text").unwrap(); + let searcher = index.reader()?.searcher(); + let phrase_query = RegexPhraseQuery::new(text_field, vec!["a.*".into(), "b".into()]); + let first = phrase_query.regexes()?.as_ptr(); + assert_eq!(phrase_query.regexes()?.as_ptr(), first); + + let enable_scoring = EnableScoring::enabled_from_searcher(&searcher); + let _weight = phrase_query.regex_phrase_weight(enable_scoring)?; + let clone = phrase_query.clone(); + let _clone_weight = clone.regex_phrase_weight(enable_scoring)?; + // The query, its clone and both weights hold the same automata. + for regex in phrase_query.regexes()? { + assert_eq!(std::sync::Arc::strong_count(regex), 4); + } + Ok(()) + } + #[test] pub fn test_phrase_count() -> crate::Result<()> { let index = create_index(&["a c", "a a b d a b c", " a b"])?; From 88a17da8e979ac683c393b8511152c01a5b1e065 Mon Sep 17 00:00:00 2001 From: palmoni5 Date: Tue, 29 Sep 2026 10:28:29 +0300 Subject: [PATCH 43/49] Expose the fragment byte range on Snippet (#3134) Snippet only kept a copy of the fragment text, so callers that need to extend the fragment or map it back to the source (e.g. to keep punctuation that the tokenizer leaves outside the last token) had to search for it, which picks the wrong occurrence when the same text appears earlier. --- src/snippet/mod.rs | 43 +++++++++++++++++++++++++++++++++++++++++-- 1 file changed, 41 insertions(+), 2 deletions(-) diff --git a/src/snippet/mod.rs b/src/snippet/mod.rs index ee61b534a..b097dacd8 100644 --- a/src/snippet/mod.rs +++ b/src/snippet/mod.rs @@ -115,6 +115,7 @@ impl FragmentCandidate { #[derive(Debug)] pub struct Snippet { fragment: String, + fragment_range: Range, highlighted: Vec>, snippet_prefix: String, snippet_postfix: String, @@ -122,9 +123,10 @@ pub struct Snippet { impl Snippet { /// Create a new `Snippet`. - fn new(fragment: &str, highlighted: Vec>) -> Self { + fn new(fragment: &str, fragment_range: Range, highlighted: Vec>) -> Self { Self { fragment: fragment.to_string(), + fragment_range, highlighted, snippet_prefix: DEFAULT_SNIPPET_PREFIX.to_string(), snippet_postfix: DEFAULT_SNIPPET_POSTFIX.to_string(), @@ -135,6 +137,7 @@ impl Snippet { pub fn empty() -> Snippet { Snippet { fragment: String::new(), + fragment_range: 0..0, highlighted: Vec::new(), snippet_prefix: String::new(), snippet_postfix: String::new(), @@ -169,6 +172,15 @@ impl Snippet { &self.fragment } + /// Returns the byte range of the fragment within the text the snippet was + /// generated from. + /// + /// For [`SnippetGenerator::snippet_from_doc`], that text is the field's + /// values joined by a single space and trimmed. + pub fn fragment_range(&self) -> Range { + self.fragment_range.clone() + } + /// Returns a list of highlighted positions from the `Snippet`. pub fn highlighted(&self) -> &[Range] { &self.highlighted @@ -250,7 +262,11 @@ fn select_best_fragment_combination(fragments: &[FragmentCandidate], text: &str) .iter() .map(|item| item.start - fragment.start_offset..item.end - fragment.start_offset) .collect(); - Snippet::new(fragment_text, highlighted) + Snippet::new( + fragment_text, + fragment.start_offset..fragment.stop_offset, + highlighted, + ) } else { // When there are no fragments to chose from, // for now create an empty snippet. @@ -596,6 +612,7 @@ Survey in 2016, 2017, and 2018."#; let snippet = select_best_fragment_combination(&fragments[..], text); assert_eq!(snippet.fragment, "c d"); + assert_eq!(snippet.fragment_range(), 4..7); assert_eq!(snippet.to_html(), "c d"); } @@ -619,9 +636,30 @@ Survey in 2016, 2017, and 2018."#; let snippet = select_best_fragment_combination(&fragments[..], text); assert_eq!(snippet.fragment, "e f"); + assert_eq!(snippet.fragment_range(), 8..11); assert_eq!(snippet.to_html(), "e f"); } + #[test] + fn test_snippet_fragment_range_with_repeated_text() { + // "a b" also occurs earlier, straddling a fragment boundary, so a + // text search for the fragment would locate the wrong occurrence. + let text = "x a b y a b"; + + let mut terms = BTreeMap::new(); + terms.insert(String::from("a"), 1.0); + terms.insert(String::from("b"), 1.0); + + let fragments = + search_fragments(&mut From::from(SimpleTokenizer::default()), text, &terms, 3); + + let snippet = select_best_fragment_combination(&fragments[..], text); + assert_eq!(snippet.fragment, "a b"); + assert_eq!(snippet.fragment_range(), 8..11); + assert_eq!(&text[snippet.fragment_range()], snippet.fragment()); + assert_ne!(text.find(snippet.fragment()), Some(8)); + } + #[test] fn test_snippet_with_second_fragment_has_the_highest_score() { let text = "a b c d e f g"; @@ -675,6 +713,7 @@ Survey in 2016, 2017, and 2018."#; let snippet = select_best_fragment_combination(&fragments[..], text); assert_eq!(snippet.fragment, ""); + assert_eq!(snippet.fragment_range(), 0..0); assert_eq!(snippet.to_html(), ""); assert!(snippet.is_empty()); } From 93280f56fbc0f9752918157d1e2bf94b70cf8090 Mon Sep 17 00:00:00 2001 From: omerb-vega Date: Tue, 29 Sep 2026 10:31:00 +0300 Subject: [PATCH 44/49] Map multivalued doc ranks to doc ids with a select cursor (#3133) * columnar: bench row id to doc id conversion on multivalued columns * columnar: map multivalued doc ranks to doc ids with a select cursor MultiValueIndexV2::select_batch_in_place ended with one OptionalIndex::select call per matched doc, each locating its block from scratch. The ranks are sorted and deduplicated at that point, so OptionalIndex::select_batch converts them in one sequential pass with a select cursor. Every Column::get_docids_for_value_range call on a multivalued column goes through here, e.g. fast field range queries. --- columnar/Cargo.toml | 4 ++ columnar/benches/bench_multivalue_docids.rs | 68 +++++++++++++++++++ .../src/column_index/multivalued_index.rs | 4 +- 3 files changed, 73 insertions(+), 3 deletions(-) create mode 100644 columnar/benches/bench_multivalue_docids.rs diff --git a/columnar/Cargo.toml b/columnar/Cargo.toml index 10b49a24d..a2e44978b 100644 --- a/columnar/Cargo.toml +++ b/columnar/Cargo.toml @@ -57,5 +57,9 @@ harness = false name = "bench_optional_index" harness = false +[[bench]] +name = "bench_multivalue_docids" +harness = false + [features] zstd-compression = ["sstable/zstd-compression"] diff --git a/columnar/benches/bench_multivalue_docids.rs b/columnar/benches/bench_multivalue_docids.rs new file mode 100644 index 000000000..e6b649feb --- /dev/null +++ b/columnar/benches/bench_multivalue_docids.rs @@ -0,0 +1,68 @@ +use binggan::{InputGroup, black_box}; +use tantivy_columnar::{Column, ColumnarReader, ColumnarWriter}; + +const NUM_DOCS: u32 = 1_000_000; +const NUM_DISTINCT_VALUES: u64 = 1000; + +/// A multivalued column where `fill_percent` of the docs hold 1 to 4 values, +/// spread over `NUM_DISTINCT_VALUES` distinct values. +fn generate_multivalued_column(fill_percent: u32) -> Column { + let mut columnar_writer = ColumnarWriter::default(); + for doc in 0..NUM_DOCS { + if doc % 100 >= fill_percent { + continue; + } + let num_values = 1 + doc % 4; + for value_idx in 0..num_values { + let value = (doc as u64 * 7 + value_idx as u64) % NUM_DISTINCT_VALUES; + columnar_writer.record_numerical(doc, "field", value); + } + } + let mut buffer: Vec = Vec::new(); + columnar_writer + .serialize(NUM_DOCS, None, &mut buffer) + .unwrap(); + let reader = ColumnarReader::open(buffer).unwrap(); + reader.read_columns("field").unwrap()[0] + .open_u64_lenient() + .unwrap() + .unwrap() +} + +fn main() { + let inputs: Vec<(String, Column)> = [100, 50, 10] + .into_iter() + .map(|fill_percent| { + ( + format!("multi 1-4 values, {fill_percent}% docs"), + generate_multivalued_column(fill_percent), + ) + }) + .collect(); + let mut group: InputGroup = InputGroup::new_with_inputs(inputs); + + group.register("docids_all_values", |column: &Column| { + let mut doc_ids = Vec::new(); + column.get_docids_for_value_range(0..=u64::MAX, 0..NUM_DOCS, &mut doc_ids); + black_box(doc_ids); + }); + group.register("docids_1pct_values", |column: &Column| { + let mut doc_ids = Vec::new(); + column.get_docids_for_value_range(0..=9, 0..NUM_DOCS, &mut doc_ids); + black_box(doc_ids); + }); + // The block-wise fetch of a range query's doc set. + group.register("docids_all_values_blocks_of_1024", |column: &Column| { + let mut doc_ids = Vec::new(); + let mut num_docs = 0; + for block_start in (0..NUM_DOCS).step_by(1024) { + let block_end = (block_start + 1024).min(NUM_DOCS); + doc_ids.clear(); + column.get_docids_for_value_range(0..=u64::MAX, block_start..block_end, &mut doc_ids); + num_docs += doc_ids.len(); + } + black_box(num_docs); + }); + + group.run(); +} diff --git a/columnar/src/column_index/multivalued_index.rs b/columnar/src/column_index/multivalued_index.rs index ad7efd363..338f422c0 100644 --- a/columnar/src/column_index/multivalued_index.rs +++ b/columnar/src/column_index/multivalued_index.rs @@ -338,9 +338,7 @@ impl MultiValueIndexV2 { } ranks.truncate(write_doc_pos); - for rank in ranks.iter_mut() { - *rank = self.optional_index.select(*rank); - } + self.optional_index.select_batch(&mut ranks[..]); } } From 9e1b0de23a631dbe11d9e9f3ecad0a52e11c984a Mon Sep 17 00:00:00 2001 From: palmoni5 Date: Tue, 29 Sep 2026 10:50:24 +0300 Subject: [PATCH 45/49] Expose the Levenshtein automaton of FuzzyTermQuery (#3136) * Expose the Levenshtein automaton of FuzzyTermQuery Callers that need the terms a fuzzy query expands to (e.g. to highlight them) had to copy the private DfaWrapper and depend on the same levenshtein_automata version as tantivy, or their expansion could silently diverge from the query's. FuzzyTermQuery::automaton returns the automaton the query matches with, built from the same cached LevenshteinAutomatonBuilder. * Export DfaWrapper outside of tests The re-export was still behind #[cfg(test)], so FuzzyTermQuery::automaton returned a type other crates could not name. A doctest now uses it from outside the crate. --- src/query/fuzzy_query.rs | 59 ++++++++++++++++++++++++++++++++++------ src/query/mod.rs | 4 +-- 2 files changed, 52 insertions(+), 11 deletions(-) diff --git a/src/query/fuzzy_query.rs b/src/query/fuzzy_query.rs index a0634b96b..a3f8faa43 100644 --- a/src/query/fuzzy_query.rs +++ b/src/query/fuzzy_query.rs @@ -6,7 +6,11 @@ use crate::query::{AutomatonWeight, EnableScoring, Query, Weight}; use crate::schema::{Term, Type}; use crate::TantivyError::InvalidArgument; -pub(crate) struct DfaWrapper(pub DFA); +/// The Levenshtein automaton a [`FuzzyTermQuery`] matches terms with. +/// +/// Obtained from [`FuzzyTermQuery::automaton`], e.g. to stream the term dictionary +/// and collect the terms the query expands to. +pub struct DfaWrapper(pub(crate) DFA); impl Automaton for DfaWrapper { type State = u32; @@ -109,7 +113,21 @@ impl FuzzyTermQuery { } } - fn specialized_weight(&self) -> crate::Result> { + /// Returns the automaton this query matches terms with. + /// + /// For a JSON term, it matches the term's text only, not its JSON path. + /// + /// ```rust + /// use tantivy::query::{DfaWrapper, FuzzyTermQuery}; + /// use tantivy::schema::{Schema, TEXT}; + /// use tantivy::Term; + /// + /// let mut schema_builder = Schema::builder(); + /// let title = schema_builder.add_text_field("title", TEXT); + /// let query = FuzzyTermQuery::new(Term::from_field_text(title, "diary"), 1, true); + /// let _automaton: DfaWrapper = query.automaton().unwrap(); + /// ``` + pub fn automaton(&self) -> crate::Result { static AUTOMATON_BUILDER: [[OnceCell; 2]; 3] = [ [OnceCell::new(), OnceCell::new()], [OnceCell::new(), OnceCell::new()], @@ -158,18 +176,19 @@ impl FuzzyTermQuery { } else { automaton_builder.build_dfa(term_text) }; + Ok(DfaWrapper(automaton)) + } - if let Some((json_path_bytes, _)) = term_value.as_json() { + fn specialized_weight(&self) -> crate::Result> { + let automaton = self.automaton()?; + if let Some((json_path_bytes, _)) = self.term.value().as_json() { Ok(AutomatonWeight::new_for_json_path( self.term.field(), - DfaWrapper(automaton), + automaton, json_path_bytes, )) } else { - Ok(AutomatonWeight::new( - self.term.field(), - DfaWrapper(automaton), - )) + Ok(AutomatonWeight::new(self.term.field(), automaton)) } } } @@ -189,6 +208,30 @@ mod test { use crate::schema::{Schema, STORED, TEXT}; use crate::{assert_nearly_equals, Index, IndexWriter, TantivyDocument, Term}; + #[test] + pub fn test_fuzzy_automaton_streams_expanded_terms() -> crate::Result<()> { + let mut schema_builder = Schema::builder(); + let title = schema_builder.add_text_field("title", TEXT); + let index = Index::create_in_ram(schema_builder.build()); + let mut index_writer: IndexWriter = index.writer_for_tests()?; + index_writer.add_document(doc!(title => "The Diary of a Dairy Cow"))?; + index_writer.add_document(doc!(title => "A Daily Log"))?; + index_writer.commit()?; + let searcher = index.reader()?.searcher(); + + let query = FuzzyTermQuery::new(Term::from_field_text(title, "diary"), 1, true); + let automaton = query.automaton()?; + let inverted_index = searcher.segment_reader(0).inverted_index(title)?; + let mut stream = inverted_index.terms().search(automaton).into_stream()?; + let mut terms = Vec::new(); + while stream.advance() { + terms.push(String::from_utf8(stream.key().to_vec()).unwrap()); + } + assert_eq!(terms, vec!["dairy", "diary"]); + assert_eq!(searcher.search(&query, &Count)?, 1); + Ok(()) + } + #[test] pub fn test_fuzzy_json_path() -> crate::Result<()> { // # Defining the schema diff --git a/src/query/mod.rs b/src/query/mod.rs index a1189c793..dd168e03b 100644 --- a/src/query/mod.rs +++ b/src/query/mod.rs @@ -51,9 +51,7 @@ pub use self::empty_query::{EmptyQuery, EmptyScorer, EmptyWeight}; pub use self::exclude::{Exclude, ExclusionSet}; pub use self::exist_query::ExistsQuery; pub use self::explanation::Explanation; -#[cfg(test)] -pub(crate) use self::fuzzy_query::DfaWrapper; -pub use self::fuzzy_query::FuzzyTermQuery; +pub use self::fuzzy_query::{DfaWrapper, FuzzyTermQuery}; pub use self::intersection::{intersect_scorers, Intersection}; pub use self::more_like_this::{MoreLikeThisQuery, MoreLikeThisQueryBuilder}; pub use self::phrase_prefix_query::PhrasePrefixQuery; From 637803de99e9d87cd0ff3e66de3ee207c74e73c2 Mon Sep 17 00:00:00 2001 From: palmoni5 Date: Tue, 29 Sep 2026 12:45:00 +0300 Subject: [PATCH 46/49] Add RegexPhraseQuery::from_regexes (#3138) RegexQuery can be built from a compiled Regex, but RegexPhraseQuery always compiled its patterns with Regex::new and its default DFA state limit. from_regexes takes the compiled regexes directly, so a caller can build them with its own limits, or reuse ones it already compiled. --- src/query/phrase_query/regex_phrase_query.rs | 30 ++++++++++++++++++ src/query/phrase_query/regex_phrase_weight.rs | 31 +++++++++++++++++++ 2 files changed, 61 insertions(+) diff --git a/src/query/phrase_query/regex_phrase_query.rs b/src/query/phrase_query/regex_phrase_query.rs index 1310871ca..67ca03499 100644 --- a/src/query/phrase_query/regex_phrase_query.rs +++ b/src/query/phrase_query/regex_phrase_query.rs @@ -96,6 +96,36 @@ impl RegexPhraseQuery { } } + /// Creates a new `RegexPhraseQuery` from already compiled regexes, e.g. built with a + /// non-default state limit. + /// + /// Each term is `(offset, pattern, regex)`, where `regex` is the compilation of + /// `pattern`; the pattern is what [`RegexPhraseQuery::phrase_terms`] returns. + pub fn from_regexes( + field: Field, + mut terms: Vec<(usize, String, Arc)>, + slop: u32, + ) -> RegexPhraseQuery { + assert!( + terms.len() > 1, + "A phrase query is required to have strictly more than one term." + ); + terms.sort_by_key(|&(offset, _, _)| offset); + let (phrase_terms, regexes): (Vec<_>, Vec<_>) = terms + .into_iter() + .map(|(offset, pattern, regex)| ((offset, pattern), regex)) + .unzip(); + let compiled = OnceCell::new(); + let _ = compiled.set(regexes); + RegexPhraseQuery { + field, + phrase_terms, + slop, + max_expansions: 1 << 14, + regexes: CompiledRegexes(compiled), + } + } + /// Slop allowed for the phrase. /// /// The query will match if its terms are separated by `slop` terms at most. diff --git a/src/query/phrase_query/regex_phrase_weight.rs b/src/query/phrase_query/regex_phrase_weight.rs index 42018f832..bbaddbc79 100644 --- a/src/query/phrase_query/regex_phrase_weight.rs +++ b/src/query/phrase_query/regex_phrase_weight.rs @@ -382,6 +382,37 @@ mod tests { Ok(()) } + #[test] + pub fn test_phrase_from_regexes_uses_the_given_automata() -> crate::Result<()> { + use std::sync::Arc; + + use tantivy_fst::Regex; + + use crate::collector::Count; + + let index = create_index(&["a b", "aa b", "b a", "a c"])?; + let text_field = index.schema().get_field("text").unwrap(); + let searcher = index.reader()?.searcher(); + let regex_a = Arc::new(Regex::new("a.*").unwrap()); + let regex_b = Arc::new(Regex::new("b").unwrap()); + let from_regexes = RegexPhraseQuery::from_regexes( + text_field, + vec![ + (1, "b".into(), regex_b.clone()), + (0, "a.*".into(), regex_a.clone()), + ], + 0, + ); + let regexes = from_regexes.regexes()?; + assert!(Arc::ptr_eq(®exes[0], ®ex_a)); + assert!(Arc::ptr_eq(®exes[1], ®ex_b)); + + let from_patterns = RegexPhraseQuery::new(text_field, vec!["a.*".into(), "b".into()]); + assert_eq!(searcher.search(&from_regexes, &Count)?, 2); + assert_eq!(searcher.search(&from_patterns, &Count)?, 2); + Ok(()) + } + #[test] pub fn test_phrase_count() -> crate::Result<()> { let index = create_index(&["a c", "a a b d a b c", " a b"])?; From 171340677f6653054c63aa91d7887d610084548d Mon Sep 17 00:00:00 2001 From: trinity-1686a Date: Tue, 29 Sep 2026 11:55:58 +0200 Subject: [PATCH 47/49] address cr --- sstable/src/delta.rs | 10 ++++------ sstable/src/lib.rs | 2 +- 2 files changed, 5 insertions(+), 7 deletions(-) diff --git a/sstable/src/delta.rs b/sstable/src/delta.rs index df43a2a18..986700fef 100644 --- a/sstable/src/delta.rs +++ b/sstable/src/delta.rs @@ -60,12 +60,10 @@ impl DeltaKeyComparator { // blocks are independent. On each new block we get a common_prefix_len=0 entry. // reset our state with it if common_prefix_len == 0 { - self.num_matching_bytes = target - .iter() - .zip(suffix) - .take_while(|(target_byte, key_byte)| target_byte == key_byte) - .count(); - return suffix.cmp(target); + let num_matching_bytes = crate::common_prefix_len(target, suffix); + self.num_matching_bytes = num_matching_bytes; + // cannot panicm at worth we might compare empty slices if num_matching_bytes==len() + return suffix[num_matching_bytes..].cmp(&target[num_matching_bytes..]); } self.compare(target, common_prefix_len, suffix) } diff --git a/sstable/src/lib.rs b/sstable/src/lib.rs index 1f6bd14c7..d925ad540 100644 --- a/sstable/src/lib.rs +++ b/sstable/src/lib.rs @@ -70,7 +70,7 @@ const SSTABLE_VERSION: u32 = 3; /// Given two byte string returns the length of /// the longest common prefix. -fn common_prefix_len(left: &[u8], right: &[u8]) -> usize { +pub(crate) fn common_prefix_len(left: &[u8], right: &[u8]) -> usize { left.iter() .cloned() .zip(right.iter().cloned()) From d5e2842575141b81661fd93d0214347be0fb7e59 Mon Sep 17 00:00:00 2001 From: Walter Woodall Date: Tue, 29 Sep 2026 16:55:20 -0700 Subject: [PATCH 48/49] fix: stream positions during sorted segment merges (#3116) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Sorted segment merges copied every document's positions for a term into a vector before sorting by mapped document ID. Common terms with many positions therefore consumed memory proportional to documents × positions, outside the indexing writer budget. Stream shuffled postings through a min-heap with one cursor per input segment and one reusable positions buffer. The mapping preserves each segment's document order; deleted/filtered postings are skipped. The stacked path continues to stream directly. --- benches/merge_segments.rs | 122 ++++++++-- src/index/inverted_index_plugin.rs | 59 ++--- src/indexer/merger_sorted_index_test.rs | 97 +++++++- src/postings/merger.rs | 305 ++++++++++++++++++++++++ src/postings/mod.rs | 2 + 5 files changed, 534 insertions(+), 51 deletions(-) create mode 100644 src/postings/merger.rs diff --git a/benches/merge_segments.rs b/benches/merge_segments.rs index 7e911f7ec..89e06fe84 100644 --- a/benches/merge_segments.rs +++ b/benches/merge_segments.rs @@ -10,7 +10,8 @@ use std::io::{self, Write}; use std::path::{Path, PathBuf}; use std::sync::{Arc, RwLock}; -use binggan::{black_box, BenchRunner}; +use binggan::plugins::PeakMemAllocPlugin; +use binggan::{black_box, BenchRunner, PeakMemAlloc, INSTRUMENTED_SYSTEM}; use rand::prelude::*; use rand::rngs::StdRng; use rand::SeedableRng; @@ -20,8 +21,11 @@ use tantivy::directory::{ WritePtr, }; use tantivy::indexer::{merge_filtered_segments, NoMergePolicy}; -use tantivy::schema::{Schema, TEXT}; -use tantivy::{doc, HasLen, Index, IndexSettings, Segment}; +use tantivy::schema::{Schema, FAST, TEXT}; +use tantivy::{doc, HasLen, Index, IndexSettings, IndexSortByField, Order, Segment}; + +#[global_allocator] +static GLOBAL: &PeakMemAlloc = &INSTRUMENTED_SYSTEM; #[derive(Clone, Default, Debug)] struct NullDirectory { @@ -196,6 +200,91 @@ fn build_index( } } +/// Like [`build_index`], sorted by a `u64` key that interleaves docs across segments. +fn build_sorted_index( + num_segments: usize, + docs_per_segment: usize, + tokens_per_doc: usize, + vocab_size: usize, +) -> MergeScenario { + let mut schema_builder = Schema::builder(); + let body = schema_builder.add_text_field("body", TEXT); + let sort = schema_builder.add_u64_field("sort", FAST); + let schema = schema_builder.build(); + let index = Index::builder() + .schema(schema) + .settings(IndexSettings { + sort_by_field: Some(IndexSortByField { + field: "sort".into(), + order: Order::Asc, + }), + ..Default::default() + }) + .create_in_ram() + .unwrap(); + + assert!(vocab_size > 0); + let total_tokens = num_segments * docs_per_segment * tokens_per_doc; + let use_unique_terms = vocab_size >= total_tokens; + let mut rng = StdRng::from_seed([7u8; 32]); + let mut next_token_id: u64 = 0; + + { + let mut writer = index.writer_with_num_threads(1, 256_000_000).unwrap(); + writer.set_merge_policy(Box::new(NoMergePolicy)); + for segment in 0..num_segments { + for row in 0..docs_per_segment { + let mut tokens = Vec::with_capacity(tokens_per_doc); + for _ in 0..tokens_per_doc { + let token_id = if use_unique_terms { + let id = next_token_id; + next_token_id += 1; + id + } else { + rng.random_range(0..vocab_size as u64) + }; + tokens.push(format!("term_{token_id}")); + } + let sort_key = (row * num_segments + segment) as u64; + writer + .add_document(doc!(body => tokens.join(" "), sort => sort_key)) + .unwrap(); + } + writer.commit().unwrap(); + } + } + + let segments = index.searchable_segments().unwrap(); + let settings = index.settings().clone(); + let label = format!( + "segments={}, docs/seg={}, tokens/doc={}, vocab={}", + num_segments, docs_per_segment, tokens_per_doc, vocab_size + ); + + MergeScenario { + index, + segments, + settings, + label, + } +} + +fn bench_merge(runner: &mut BenchRunner, group_name: String, scenario: MergeScenario) { + let mut group = runner.new_group(); + group.set_name(group_name); + let segments = scenario.segments.clone(); + let settings = scenario.settings.clone(); + group.register("merge", move |_| { + let output_dir = NullDirectory::default(); + let filter_doc_ids = vec![None; segments.len()]; + let merged_index = + merge_filtered_segments(&segments, settings.clone(), filter_doc_ids, output_dir) + .unwrap(); + black_box(merged_index); + }); + group.run(); +} + fn main() { let scenarios = vec![ build_index(8, 50_000, 12, 8), @@ -205,20 +294,21 @@ fn main() { ]; let mut runner = BenchRunner::new(); + runner.add_plugin(PeakMemAllocPlugin::new(GLOBAL)); for scenario in scenarios { - let mut group = runner.new_group(); - group.set_name(format!("merge_segments inv_index — {}", scenario.label)); - let segments = scenario.segments.clone(); - let settings = scenario.settings.clone(); - group.register("merge", move |_| { - let output_dir = NullDirectory::default(); - let filter_doc_ids = vec![None; segments.len()]; - let merged_index = - merge_filtered_segments(&segments, settings.clone(), filter_doc_ids, output_dir) - .unwrap(); - black_box(merged_index); - }); + let name = format!("merge_segments inv_index — {}", scenario.label); + bench_merge(&mut runner, name, scenario); + } - group.run(); + let sorted_scenarios = vec![ + build_sorted_index(8, 50_000, 12, 8), + build_sorted_index(16, 50_000, 12, 8), + build_sorted_index(16, 100_000, 12, 8), + build_sorted_index(8, 50_000, 8, 8 * 50_000 * 8), + ]; + + for scenario in sorted_scenarios { + let name = format!("merge_segments sorted inv_index — {}", scenario.label); + bench_merge(&mut runner, name, scenario); } } diff --git a/src/index/inverted_index_plugin.rs b/src/index/inverted_index_plugin.rs index bc9c55e58..6ca50296b 100644 --- a/src/index/inverted_index_plugin.rs +++ b/src/index/inverted_index_plugin.rs @@ -17,7 +17,7 @@ use measure_time::debug_time; use tokenizer_api::BoxTokenStream; use crate::directory::{CompositeFile, Directory}; -use crate::docset::{DocSet, TERMINATED}; +use crate::docset::DocSet; use crate::error::DataCorruption; use crate::fieldnorm::{FieldNormReader, FieldNormReaders, FieldNormsSerializer, FieldNormsWriter}; use crate::index::{Segment, SegmentComponent, SegmentReader}; @@ -27,7 +27,8 @@ use crate::json_utils::{index_json_value, IndexingPositionsPerPath}; use crate::plugin::{PluginMergeContext, PluginWriter, PluginWriterContext, SegmentPlugin}; use crate::postings::{ compute_table_memory_size, serialize_postings, IndexingContext, IndexingPosition, - InvertedIndexSerializer, PerFieldPostingsWriter, Postings, PostingsWriter, SegmentPostings, + InvertedIndexSerializer, PerFieldPostingsWriter, Postings, PostingsMerger, PostingsWriter, + SegmentPostings, }; use crate::schema::document::{Document, Value}; use crate::schema::{Field, FieldType, Schema, DATE_TIME_PRECISION_INDEXED}; @@ -595,7 +596,7 @@ fn write_postings_for_field( ); let mut segment_postings_containing_the_term: Vec<(usize, SegmentPostings)> = vec![]; - let mut doc_id_and_positions = vec![]; + let mut merger = PostingsMerger::new(&merged_doc_id_map); while merged_terms.advance() { segment_postings_containing_the_term.clear(); @@ -645,43 +646,35 @@ fn write_postings_for_field( field_serializer.new_term(term_bytes, total_doc_freq, has_term_freq)?; - for (segment_ord, mut segment_postings) in segment_postings_containing_the_term.drain(..) { - let old_to_new_doc_id = &merged_doc_id_map[segment_ord]; - - let mut doc = segment_postings.doc(); - while doc != TERMINATED { - if let Some(remapped_doc_id) = old_to_new_doc_id[doc as usize] { + if doc_id_mapping.is_trivial() { + for (segment_ord, mut postings) in segment_postings_containing_the_term.drain(..) { + let mapping = &merged_doc_id_map[segment_ord]; + while let Some(doc) = crate::postings::next_mapped_doc(&mut postings, mapping) { let term_freq = if has_term_freq { - segment_postings.positions(&mut positions_buffer); - segment_postings.term_freq() + postings.positions(&mut positions_buffer); + postings.term_freq() } else { positions_buffer.clear(); - 0u32 + 0 }; - - if !doc_id_mapping.is_trivial() { - doc_id_and_positions.push(( - remapped_doc_id, - term_freq, - positions_buffer.to_vec(), - )); - } else { - let delta_positions = delta_computer.compute_delta(&positions_buffer); - field_serializer.write_doc(remapped_doc_id, term_freq, delta_positions); - } + let delta_positions = delta_computer.compute_delta(&positions_buffer); + field_serializer.write_doc(doc, term_freq, delta_positions); + postings.advance(); } - - doc = segment_postings.advance(); } - } - if !doc_id_mapping.is_trivial() { - doc_id_and_positions.sort_unstable_by_key(|&(doc_id, _, _)| doc_id); - - for (doc_id, term_freq, positions) in &doc_id_and_positions { - let delta_positions = delta_computer.compute_delta(positions); - field_serializer.write_doc(*doc_id, *term_freq, delta_positions); + } else { + merger.reset(segment_postings_containing_the_term.drain(..)); + while merger.advance() { + let term_freq = if has_term_freq { + merger.positions(&mut positions_buffer); + merger.term_freq() + } else { + positions_buffer.clear(); + 0 + }; + let delta_positions = delta_computer.compute_delta(&positions_buffer); + field_serializer.write_doc(merger.doc(), term_freq, delta_positions); } - doc_id_and_positions.clear(); } field_serializer.close_term()?; } diff --git a/src/indexer/merger_sorted_index_test.rs b/src/indexer/merger_sorted_index_test.rs index 2f230b8df..e717709d4 100644 --- a/src/indexer/merger_sorted_index_test.rs +++ b/src/indexer/merger_sorted_index_test.rs @@ -7,15 +7,16 @@ mod tests { use crate::collector::TopDocs; use crate::fastfield::AliveBitSet; use crate::index::Index; + use crate::indexer::NoMergePolicy; use crate::postings::Postings; use crate::query::QueryParser; use crate::schema::{ self, BytesOptions, Facet, FacetOptions, IndexRecordOption, NumericOptions, - TextFieldIndexing, TextOptions, Value, FAST, STRING, + TextFieldIndexing, TextOptions, Value, FAST, INDEXED, STRING, }; use crate::{ DocAddress, DocSet, IndexSettings, IndexSortByField, IndexWriter, Order, TantivyDocument, - Term, + Term, TERMINATED, }; fn create_test_index_posting_list_issue(index_settings: Option) -> Index { @@ -920,6 +921,98 @@ mod tests { } } + #[test] + fn test_merge_sorted_index_postings_with_deletes_and_missing_sort_keys() -> crate::Result<()> { + for order in [Order::Asc, Order::Desc] { + for record in [ + IndexRecordOption::Basic, + IndexRecordOption::WithFreqs, + IndexRecordOption::WithFreqsAndPositions, + ] { + let mut schema_builder = schema::Schema::builder(); + let id_field = schema_builder.add_u64_field("id", FAST | INDEXED); + let sort_field = schema_builder.add_u64_field("sort", FAST); + let text = schema_builder.add_text_field( + "text", + TextOptions::default().set_indexing_options( + TextFieldIndexing::default() + .set_tokenizer("default") + .set_index_option(record), + ), + ); + let index = Index::builder() + .schema(schema_builder.build()) + .settings(IndexSettings { + sort_by_field: Some(IndexSortByField { + field: "sort".into(), + order, + }), + ..Default::default() + }) + .create_in_ram()?; + let mut writer: IndexWriter = index.writer_with_num_threads(1, 15_000_000)?; + writer.set_merge_policy(Box::new(NoMergePolicy)); + for segment in 0..3u64 { + for row in 0..50u64 { + let id = row * 3 + segment; + let mut document = + doc!(id_field=>id, text=>"common anchor ".repeat((id%7+1) as usize)); + if id % 5 != 0 { + document.add_u64(sort_field, id % 11); + } + writer.add_document(document)?; + } + writer.commit()?; + } + for id in (0..150u64).step_by(7) { + writer.delete_term(Term::from_field_u64(id_field, id)); + } + writer.commit()?; + let segment_ids = index.searchable_segment_ids()?; + assert_eq!(segment_ids.len(), 3); + writer.merge(&segment_ids).wait()?; + + let searcher = index.reader()?.searcher(); + assert_eq!(searcher.segment_readers().len(), 1); + let segment = searcher.segment_reader(0); + let ids = segment.fast_fields().u64("id")?; + let inverted = segment.inverted_index(text)?; + for (token, offset) in [("common", 0), ("anchor", 1)] { + let mut postings = inverted + .read_postings(&Term::from_field_text(text, token), record)? + .unwrap(); + let mut seen = std::collections::BTreeSet::new(); + let mut positions = Vec::new(); + let mut previous_doc = None; + while postings.doc() != TERMINATED { + let doc = postings.doc(); + if let Some(previous) = previous_doc { + assert!(doc > previous); + } + previous_doc = Some(doc); + let id = ids.first(doc).unwrap(); + assert_ne!(id % 7, 0); + assert!(seen.insert(id)); + let repetitions = (id % 7 + 1) as u32; + if record != IndexRecordOption::Basic { + assert_eq!(postings.term_freq(), repetitions); + } + if record == IndexRecordOption::WithFreqsAndPositions { + postings.positions(&mut positions); + assert_eq!( + positions, + (0..repetitions).map(|p| 2 * p + offset).collect::>() + ); + } + postings.advance(); + } + assert_eq!(seen, (0..150u64).filter(|id| id % 7 != 0).collect()); + } + } + } + Ok(()) + } + // #[test] // fn test_merge_sorted_index_asc() { // let index = create_test_index( diff --git a/src/postings/merger.rs b/src/postings/merger.rs new file mode 100644 index 000000000..a49254a43 --- /dev/null +++ b/src/postings/merger.rs @@ -0,0 +1,305 @@ +//! Merge per-segment postings into increasing mapped document id order. + +use std::cmp::Reverse; +use std::collections::binary_heap::PeekMut; +use std::collections::BinaryHeap; + +use crate::docset::{DocSet, TERMINATED}; +use crate::postings::{Postings, SegmentPostings}; +use crate::DocId; + +/// Skip to the next posting whose document is present in `mapping`. +pub(crate) fn next_mapped_doc( + postings: &mut SegmentPostings, + mapping: &[Option], +) -> Option { + while postings.doc() != TERMINATED { + if let Some(doc) = mapping[postings.doc() as usize] { + return Some(doc); + } + postings.advance(); + } + None +} + +struct MappedPostings<'a> { + postings: SegmentPostings, + mapping: &'a [Option], +} + +impl MappedPostings<'_> { + fn advance(&mut self) -> Option { + self.postings.advance(); + next_mapped_doc(&mut self.postings, self.mapping) + } +} + +/// Streams postings from several segments in increasing mapped document id. +/// +/// Create once per field and call [`Self::reset`] for each term. +pub(crate) struct PostingsMerger<'a> { + doc_id_map: &'a [Vec>], + cursors: Vec>, + /// Min-heap of `(current mapped doc, index into cursors)`. + heap: BinaryHeap>, + primed: bool, +} + +impl<'a> PostingsMerger<'a> { + /// `doc_id_map[segment_ord][local_doc]` is the mapped doc, or `None` if dropped. + pub(crate) fn new(doc_id_map: &'a [Vec>]) -> Self { + Self { + doc_id_map, + cursors: Vec::new(), + heap: BinaryHeap::new(), + primed: false, + } + } + + /// Start merging a new term from `(segment_ord, postings)` pairs. + pub(crate) fn reset(&mut self, segments: impl IntoIterator) { + self.cursors.clear(); + self.heap.clear(); + self.primed = false; + let doc_id_map = self.doc_id_map; + for (segment_ord, mut postings) in segments { + let mapping = &doc_id_map[segment_ord][..]; + if let Some(doc) = next_mapped_doc(&mut postings, mapping) { + self.heap.push(Reverse((doc, self.cursors.len()))); + self.cursors.push(MappedPostings { postings, mapping }); + } + } + } + + /// Advance to the next document. Must return `true` before the accessors are called. + pub(crate) fn advance(&mut self) -> bool { + if !self.primed { + self.primed = true; + return !self.heap.is_empty(); + } + let previous = { + let Some(mut top) = self.heap.peek_mut() else { + return false; + }; + let Reverse((previous, cursor_ord)) = *top; + match self.cursors[cursor_ord].advance() { + Some(doc) => { + debug_assert!( + doc > previous, + "merge mapping must preserve per-segment order" + ); + *top = Reverse((doc, cursor_ord)); + } + None => { + PeekMut::pop(top); + } + } + previous + }; + if let Some(Reverse((next, _))) = self.heap.peek() { + debug_assert!( + *next > previous, + "mapped doc ids must be strictly increasing" + ); + } + !self.heap.is_empty() + } + + fn current(&self) -> (DocId, usize) { + self.heap.peek().expect("advance() returned true").0 + } + + pub(crate) fn doc(&self) -> DocId { + self.current().0 + } + + pub(crate) fn term_freq(&self) -> u32 { + self.cursors[self.current().1].postings.term_freq() + } + + pub(crate) fn positions(&mut self, output: &mut Vec) { + let cursor_ord = self.current().1; + self.cursors[cursor_ord].postings.positions(output); + } +} + +#[cfg(test)] +mod tests { + use super::PostingsMerger; + use crate::postings::SegmentPostings; + use crate::schema::{IndexRecordOption, Schema, TEXT}; + use crate::{DocId, Index, IndexWriter, Term}; + + fn collect(merger: &mut PostingsMerger<'_>) -> Vec<(DocId, u32, Vec)> { + let mut docs = Vec::new(); + let mut positions = vec![7, 7, 7]; + while merger.advance() { + merger.positions(&mut positions); + docs.push((merger.doc(), merger.term_freq(), positions.clone())); + } + assert!(!merger.advance()); + docs + } + + fn without_positions(docs: &[(DocId, u32)]) -> Vec<(DocId, u32, Vec)> { + docs.iter() + .map(|&(doc, tf)| (doc, tf, Vec::new())) + .collect() + } + + /// Postings with positions for each token, one segment per inner slice. + fn postings_with_positions( + segments: &[&[&str]], + tokens: &[&str], + ) -> crate::Result>> { + let mut schema_builder = Schema::builder(); + let text = schema_builder.add_text_field("text", TEXT); + let schema = schema_builder.build(); + let mut readers = Vec::new(); + for docs in segments { + let index = Index::create_in_ram(schema.clone()); + let mut writer: IndexWriter = index.writer_for_tests()?; + for body in *docs { + writer.add_document(doc!(text => *body))?; + } + writer.commit()?; + let searcher = index.reader()?.searcher(); + assert_eq!(searcher.segment_readers().len(), 1); + readers.push(searcher.segment_reader(0).inverted_index(text)?); + } + let mut terms = Vec::new(); + for token in tokens { + let term = Term::from_field_text(text, token); + let mut postings = Vec::new(); + for (segment_ord, reader) in readers.iter().enumerate() { + if let Some(segment_postings) = + reader.read_postings(&term, IndexRecordOption::WithFreqsAndPositions)? + { + postings.push((segment_ord, segment_postings)); + } + } + terms.push(postings); + } + Ok(terms) + } + + #[test] + fn test_merges_segments_skipping_deletes() { + // Segment 2 is fully deleted and segment 3 has no postings. + let doc_id_map = vec![ + vec![Some(1), None, Some(3)], + vec![Some(0), Some(2)], + vec![None, None], + Vec::new(), + ]; + let segments = [ + ( + 0, + SegmentPostings::create_from_docs_and_tfs(&[(0, 1), (1, 2), (2, 3)], None), + ), + ( + 1, + SegmentPostings::create_from_docs_and_tfs(&[(0, 4), (1, 5)], None), + ), + ( + 2, + SegmentPostings::create_from_docs_and_tfs(&[(0, 9), (1, 9)], None), + ), + (3, SegmentPostings::empty()), + ]; + let mut merger = PostingsMerger::new(&doc_id_map); + merger.reset(segments); + assert_eq!( + collect(&mut merger), + without_positions(&[(0, 4), (1, 1), (2, 5), (3, 3)]) + ); + } + + #[test] + fn test_empty_term_yields_nothing() { + let doc_id_map = vec![vec![Some(0)]]; + let mut merger = PostingsMerger::new(&doc_id_map); + merger.reset([(0, SegmentPostings::empty())]); + assert_eq!(collect(&mut merger), Vec::new()); + } + + #[test] + fn test_crosses_postings_block_boundary() { + const N: u32 = 200; + let seg0: Vec<(u32, u32)> = (0..N).map(|doc| (doc, doc + 1)).collect(); + let seg1: Vec<(u32, u32)> = (0..N).map(|doc| (doc, 1_000 + doc)).collect(); + let map0: Vec> = (0..N) + .map(|doc| if doc % 10 == 0 { None } else { Some(doc * 2) }) + .collect(); + let map1: Vec> = (0..N).map(|doc| Some(doc * 2 + 1)).collect(); + let doc_id_map = vec![map0, map1]; + let mut merger = PostingsMerger::new(&doc_id_map); + merger.reset([ + (0, SegmentPostings::create_from_docs_and_tfs(&seg0, None)), + (1, SegmentPostings::create_from_docs_and_tfs(&seg1, None)), + ]); + + let mut expected = Vec::new(); + for doc in 0..N { + if doc % 10 != 0 { + expected.push((doc * 2, doc + 1)); + } + expected.push((doc * 2 + 1, 1_000 + doc)); + } + expected.sort_unstable(); + assert_eq!(collect(&mut merger), without_positions(&expected)); + } + + #[test] + fn test_positions_follow_their_document() -> crate::Result<()> { + let segments: [&[&str]; 2] = [&["a b a", "b", "b a b a a"], &["a", "b b a", "a b", "b"]]; + let doc_id_map = vec![ + vec![Some(1), None, Some(4)], + vec![Some(0), Some(2), None, Some(3)], + ]; + let mut terms = postings_with_positions(&segments, &["a", "b"])?.into_iter(); + let mut merger = PostingsMerger::new(&doc_id_map); + + merger.reset(terms.next().unwrap()); + assert_eq!( + collect(&mut merger), + vec![ + (0, 1, vec![0]), + (1, 2, vec![0, 2]), + (2, 1, vec![2]), + (4, 3, vec![1, 3, 4]), + ] + ); + + merger.reset(terms.next().unwrap()); + assert_eq!( + collect(&mut merger), + vec![ + (1, 1, vec![1]), + (2, 2, vec![0, 1]), + (3, 1, vec![0]), + (4, 2, vec![0, 2]), + ] + ); + Ok(()) + } + + #[test] + fn test_reset_discards_unfinished_term() -> crate::Result<()> { + let segments: [&[&str]; 2] = [&["a b", "a"], &["b a", "a b", "a"]]; + let doc_id_map = vec![vec![Some(0), Some(2)], vec![Some(1), Some(3), Some(4)]]; + let mut terms = postings_with_positions(&segments, &["a", "b"])?.into_iter(); + let mut merger = PostingsMerger::new(&doc_id_map); + + merger.reset(terms.next().unwrap()); + assert!(merger.advance()); + assert_eq!(merger.doc(), 0); + + merger.reset(terms.next().unwrap()); + assert_eq!( + collect(&mut merger), + vec![(0, 1, vec![1]), (1, 1, vec![0]), (3, 1, vec![1])] + ); + Ok(()) + } +} diff --git a/src/postings/mod.rs b/src/postings/mod.rs index ea512230c..61f83e62f 100644 --- a/src/postings/mod.rs +++ b/src/postings/mod.rs @@ -9,6 +9,7 @@ pub(crate) mod compression; mod indexing_context; mod json_postings_writer; mod loaded_postings; +mod merger; mod per_field_postings_writer; mod postings; mod postings_writer; @@ -20,6 +21,7 @@ mod skip; mod term_info; pub(crate) use loaded_postings::LoadedPostings; +pub(crate) use merger::{next_mapped_doc, PostingsMerger}; pub(crate) use stacker::compute_table_memory_size; pub use self::block_segment_postings::BlockSegmentPostings; From 1f9e49da6b67393fd9fad3ae49f05a083bdd7575 Mon Sep 17 00:00:00 2001 From: Paul Masurel Date: Wed, 30 Sep 2026 19:12:20 +0200 Subject: [PATCH 49/49] Abstracting columnar from aggregation. (#3112) * Changing the way aggregation access their value. They now get values via a ValueSource abstraction. The aggregation collector also gets the possibility to register ValueSourceProvider describing value columns that are computed on the fly. Finally, segment aggregation that require a full column now manipulates a Arc directly. * CR comments * Clippy * Fixing regression --------- Co-authored-by: Paul Masurel --- columnar/src/column_values/mod.rs | 5 + jitexpr/src/types.rs | 6 +- src/aggregation/accessor_helpers.rs | 76 +++-- src/aggregation/agg_data.rs | 145 ++++++++-- src/aggregation/bucket/histogram/histogram.rs | 58 ++-- src/aggregation/bucket/multi_terms/mod.rs | 2 +- src/aggregation/bucket/range.rs | 22 +- .../term_agg/flattened_term_histogram.rs | 146 ++++++---- src/aggregation/bucket/term_agg/mod.rs | 42 ++- src/aggregation/metric/cardinality/mod.rs | 18 +- .../metric/cardinality/numeric_collector.rs | 19 +- .../metric/cardinality/str_collector.rs | 10 +- src/aggregation/metric/extended_stats.rs | 18 +- src/aggregation/metric/mod.rs | 7 +- src/aggregation/metric/percentiles.rs | 14 +- src/aggregation/metric/stats.rs | 51 +++- src/aggregation/mod.rs | 21 +- .../{ => value_source}/block_accessor.rs | 270 +++++++++--------- src/aggregation/value_source/mod.rs | 126 ++++++++ src/aggregation/value_source/tests.rs | 117 ++++++++ .../value_source/value_source_registry.rs | 72 +++++ src/index/inverted_index_plugin.rs | 6 +- 22 files changed, 905 insertions(+), 346 deletions(-) rename src/aggregation/{ => value_source}/block_accessor.rs (74%) create mode 100644 src/aggregation/value_source/mod.rs create mode 100644 src/aggregation/value_source/tests.rs create mode 100644 src/aggregation/value_source/value_source_registry.rs diff --git a/columnar/src/column_values/mod.rs b/columnar/src/column_values/mod.rs index 64bc69b25..486f2e385 100644 --- a/columnar/src/column_values/mod.rs +++ b/columnar/src/column_values/mod.rs @@ -203,6 +203,11 @@ impl ColumnValues for Arc]) { self.as_ref().get_vals_opt(indexes, output) diff --git a/jitexpr/src/types.rs b/jitexpr/src/types.rs index 71e309739..e767a0ce5 100644 --- a/jitexpr/src/types.rs +++ b/jitexpr/src/types.rs @@ -70,14 +70,16 @@ impl Eq for SafeF64 {} impl Ord for SafeF64 { #[inline(always)] fn cmp(&self, other: &SafeF64) -> Ordering { - self.partial_cmp(&other).unwrap() + self.0 + .partial_cmp(&other.0) + .expect("SafeF64 excludes NaN, so values are totally ordered") } } impl PartialOrd for SafeF64 { #[inline(always)] fn partial_cmp(&self, other: &SafeF64) -> Option { - self.0.partial_cmp(&other.0) + Some(self.cmp(other)) } } diff --git a/src/aggregation/accessor_helpers.rs b/src/aggregation/accessor_helpers.rs index fa51041e4..9d050d1b9 100644 --- a/src/aggregation/accessor_helpers.rs +++ b/src/aggregation/accessor_helpers.rs @@ -1,10 +1,12 @@ //! This will enhance the request tree with access to the fastfield and metadata. use std::io; +use std::sync::Arc; -use columnar::{Column, ColumnType}; +use columnar::{Column, ColumnType, DynamicColumn, DynamicColumnHandle}; -use crate::aggregation::{f64_to_fastfield_u64, Key}; +use crate::aggregation::value_source::ValueSource; +use crate::aggregation::{f64_to_fastfield_u64, Key, ValueSourceRegistry}; use crate::index::SegmentReader; /// Get the missing value as internal u64 representation @@ -55,14 +57,38 @@ pub(crate) fn get_numeric_or_date_column_types() -> &'static [ColumnType] { ] } -/// Get fast field reader or empty as default. -pub(crate) fn get_ff_reader( +fn resolve_registered_source( reader: &SegmentReader, + value_sources: &ValueSourceRegistry, + field_name: &str, + allowed_column_types_opt: Option<&[ColumnType]>, +) -> crate::Result>> { + let Some(provider) = value_sources.get(field_name) else { + return Ok(None); + }; + let source = provider.for_segment(reader)?; + let column_type = source.column_type(); + if let Some(allowed_column_types) = allowed_column_types_opt { + if !allowed_column_types.contains(&column_type) { + return Ok(None); + } + } + Ok(Some(source)) +} + +pub(crate) fn get_value_source( + reader: &SegmentReader, + value_sources: &ValueSourceRegistry, field_name: &str, allowed_column_types: Option<&[ColumnType]>, -) -> crate::Result<(columnar::Column, ColumnType)> { +) -> crate::Result> { + if let Some(registered) = + resolve_registered_source(reader, value_sources, field_name, allowed_column_types)? + { + return Ok(registered); + } let ff_fields = reader.fast_fields(); - let ff_field_with_type = ff_fields + let (column, column_type) = ff_fields .u64_lenient_for_type(allowed_column_types, field_name)? .unwrap_or_else(|| { ( @@ -70,36 +96,52 @@ pub(crate) fn get_ff_reader( ColumnType::U64, ) }); - Ok(ff_field_with_type) + // The empty-column shim stays physical on purpose: several fast paths check + // `as_column()` and would otherwise degrade for a merely absent field. + Ok(Arc::new((column, column_type))) } pub(crate) fn get_dynamic_columns( reader: &SegmentReader, field_name: &str, ) -> crate::Result> { - let ff_fields = reader.fast_fields().dynamic_column_handles(field_name)?; - let cols = ff_fields + let dyn_col_handles: Vec = + reader.fast_fields().dynamic_column_handles(field_name)?; + let dyn_cols: Vec = dyn_col_handles .iter() - .map(|h| h.open()) + .map(DynamicColumnHandle::open) .collect::>()?; - assert!(!ff_fields.is_empty(), "field {field_name} not found"); - Ok(cols) + assert!(!dyn_cols.is_empty(), "field {field_name} not found"); + Ok(dyn_cols) } -/// Get all fast field reader or empty as default. +/// Get all block_value_sources or empty as default. /// /// Is guaranteed to return at least one column. -pub(crate) fn get_all_ff_reader_or_empty( +pub(crate) fn get_all_value_sources( reader: &SegmentReader, + value_sources: &ValueSourceRegistry, field_name: &str, allowed_column_types: Option<&[ColumnType]>, fallback_type: ColumnType, -) -> crate::Result, ColumnType)>> { +) -> crate::Result>> { + // A registered source shadows the physical type fan-out entirely. + if let Some(registered) = + resolve_registered_source(reader, value_sources, field_name, allowed_column_types)? + { + return Ok(vec![registered]); + } let ff_fields = reader.fast_fields(); - let mut ff_field_with_type = + let mut ff_field_with_type: Vec<(Column, ColumnType)> = ff_fields.u64_lenient_for_type_all(allowed_column_types, field_name)?; if ff_field_with_type.is_empty() { ff_field_with_type.push((Column::build_empty_column(reader.num_docs()), fallback_type)); } - Ok(ff_field_with_type) + Ok(ff_field_with_type + .into_iter() + .map(|(column, column_type)| { + let source: Arc = Arc::new((column, column_type)); + source + }) + .collect()) } diff --git a/src/aggregation/agg_data.rs b/src/aggregation/agg_data.rs index 16a08c0e4..0213e002e 100644 --- a/src/aggregation/agg_data.rs +++ b/src/aggregation/agg_data.rs @@ -7,8 +7,8 @@ use serde::Serialize; use tantivy_fst::Regex; use crate::aggregation::accessor_helpers::{ - get_all_ff_reader_or_empty, get_dynamic_columns, get_ff_reader, get_missing_val_as_u64_lenient, - get_numeric_or_date_column_types, + get_all_value_sources, get_dynamic_columns, get_missing_val_as_u64_lenient, + get_numeric_or_date_column_types, get_value_source, }; use crate::aggregation::agg_req::{Aggregation, AggregationVariants, Aggregations}; use crate::aggregation::bucket::{ @@ -28,7 +28,10 @@ use crate::aggregation::metric::{ use crate::aggregation::segment_agg_result::{ GenericSegmentAggregationResultsCollector, SegmentAggregationCollector, }; -use crate::aggregation::{f64_to_fastfield_u64, AggContextParams, ColumnBlockAccessor, Key}; +use crate::aggregation::{ + f64_to_fastfield_u64, AggContextParams, ColumnBlockAccessor, Key, ValueSource, + ValueSourceRegistry, +}; use crate::{SegmentOrdinal, SegmentReader}; #[derive(Default)] @@ -303,7 +306,6 @@ pub(crate) fn build_segment_agg_collector( let req_data = req.get_metric_req_data(node.idx_in_req_data); Ok(Box::new( SegmentPercentilesCollector::from_req_and_validate( - req_data.field_type, req_data.missing_u64, req_data.accessor.clone(), node.idx_in_req_data, @@ -403,6 +405,46 @@ pub(crate) fn build_aggregations_data_from_req( Ok(data) } +/// Resolves the substitute value used for documents that have none. +/// +/// Only the `Str` arms of [`get_missing_val_as_u64_lenient`] read the column's max value — they +/// place the sentinel one past the last real term ordinal — and that bound exists only for a +/// materialized column. A computed text source therefore cannot support `missing`; for numeric +/// types the argument is ignored, so any value will do. +fn missing_value_for_source( + accessor: &dyn ValueSource, + missing: &Key, + field_name: &str, +) -> crate::Result> { + let column_type = accessor.column_type(); + let column_max_value = match accessor.as_column() { + Some(column) => column.max_value(), + None if column_type == ColumnType::Str => { + return Err(crate::TantivyError::InvalidArgument(format!( + "`missing` is not supported for the computed text value source `{field_name}`" + ))); + } + None => 0, + }; + get_missing_val_as_u64_lenient(column_type, column_max_value, missing, field_name) +} + +/// Extracts the materialized column, rejecting a computed source. +/// +/// For the aggregations that read values through per-document random access or +/// `ColumnIndex::has_value`, neither of which `ValueSource` can express. +fn require_physical_column( + source: &dyn ValueSource, + field_name: &str, + agg_kind: &str, +) -> crate::Result> { + source.as_column().cloned().ok_or_else(|| { + crate::TantivyError::InvalidArgument(format!( + "{agg_kind} does not support the computed value source `{field_name}`" + )) + }) +} + fn build_nodes( agg_name: &str, req: &Aggregation, @@ -412,16 +454,17 @@ fn build_nodes( is_top_level: bool, ) -> crate::Result> { use AggregationVariants::*; + let value_sources = &data.context.value_sources; match &req.agg { Range(range_req) => { - let (accessor, field_type) = get_ff_reader( + let accessor = get_value_source( reader, + value_sources, &range_req.field, Some(get_numeric_or_date_column_types()), )?; let idx_in_req_data = data.push_range_req_data(RangeAggReqData { accessor, - field_type, name: agg_name.to_string(), req: range_req.clone(), is_top_level, @@ -434,14 +477,14 @@ fn build_nodes( }]) } Histogram(histo_req) => { - let (accessor, field_type) = get_ff_reader( + let accessor = get_value_source( reader, + value_sources, &histo_req.field, Some(get_numeric_or_date_column_types()), )?; let idx_in_req_data = data.push_histogram_req_data(HistogramAggReqData { accessor, - field_type, name: agg_name.to_string(), req: histo_req.clone(), is_date_histogram: false, @@ -459,14 +502,17 @@ fn build_nodes( }]) } DateHistogram(date_req) => { - let (accessor, field_type) = - get_ff_reader(reader, &date_req.field, Some(&[ColumnType::DateTime]))?; + let accessor = get_value_source( + reader, + value_sources, + &date_req.field, + Some(&[ColumnType::DateTime]), + )?; // Convert to histogram request, normalize to ns precision let mut histo_req = date_req.to_histogram_req()?; histo_req.normalize_date_time(); let idx_in_req_data = data.push_histogram_req_data(HistogramAggReqData { accessor, - field_type, name: agg_name.to_string(), req: histo_req, is_date_histogram: true, @@ -543,10 +589,10 @@ fn build_nodes( )) } }; - let (accessor, field_type) = get_ff_reader(reader, field, allowed_column_types)?; + let accessor = get_value_source(reader, value_sources, field, allowed_column_types)?; + let field_type = accessor.column_type(); let idx_in_req_data = data.push_metric_req_data(MetricAggReqData { accessor, - field_type, name: agg_name.to_string(), collecting_for, missing: *missing, @@ -566,14 +612,15 @@ fn build_nodes( // Percentiles handled as Metric as well AggregationVariants::Percentiles(percentiles_req) => { percentiles_req.validate()?; - let (accessor, field_type) = get_ff_reader( + let accessor = get_value_source( reader, + value_sources, percentiles_req.field_name(), Some(get_numeric_or_date_column_types()), )?; + let field_type = accessor.column_type(); let idx_in_req_data = data.push_metric_req_data(MetricAggReqData { accessor, - field_type, name: agg_name.to_string(), collecting_for: StatsType::Percentiles, missing: percentiles_req.missing, @@ -598,7 +645,18 @@ fn build_nodes( let accessors: Vec<(Column, ColumnType)> = top_hits .field_names() .iter() - .map(|field| get_ff_reader(reader, field, Some(get_numeric_or_date_column_types()))) + .map(|field| { + let source = get_value_source( + reader, + value_sources, + field, + Some(get_numeric_or_date_column_types()), + )?; + // Sort fields are read one document at a time via `values_for_doc`, which has + // no block equivalent. + let column = require_physical_column(&*source, field, "top_hits")?; + Ok((column, source.column_type())) + }) .collect::>()?; let value_accessors = top_hits @@ -715,11 +773,21 @@ fn build_multi_terms_nodes( )); } + let value_sources = data.context.value_sources.clone(); let mut accessors_by_field = Vec::with_capacity(req.terms.len()); for field_def in &req.terms { let field_name = &field_def.field; let str_dict_column = reader.fast_fields().str(field_name)?; - let columns = get_term_agg_accessors(reader, field_name, &field_def.missing, true)?; + // multi_terms resolves missing values through `ColumnIndex::has_value` per document, and + // exposes its columns on a public struct, so it stays physical-only. + let columns = + get_term_agg_accessors(reader, &value_sources, field_name, &field_def.missing, true)? + .into_iter() + .map(|source| { + let column = require_physical_column(&*source, field_name, "multi_terms")?; + Ok((column, source.column_type())) + }) + .collect::>>()?; if let Some((_, column_type)) = columns .iter() @@ -933,10 +1001,11 @@ fn build_children( fn get_term_agg_accessors( reader: &SegmentReader, + value_sources: &ValueSourceRegistry, field_name: &str, missing: &Option, include_bytes: bool, -) -> crate::Result, ColumnType)>> { +) -> crate::Result>> { // `terms` and `multi_terms` both explicitly reject `Bytes` columns downstream, which needs // to actually see them as a real column (rather than the empty shim below) to do so. // `cardinality` has no such rejection: it would hash raw `Bytes` term ordinals as if they @@ -967,14 +1036,15 @@ fn get_term_agg_accessors( }) .unwrap_or(ColumnType::U64); - let column_and_types = get_all_ff_reader_or_empty( + let sources = get_all_value_sources( reader, + value_sources, field_name, Some(&allowed_column_types), fallback_type, )?; - Ok(column_and_types) + Ok(sources) } enum TermsOrCardinalityRequest { @@ -1005,14 +1075,16 @@ fn build_terms_or_cardinality_nodes( let mut nodes = Vec::new(); let str_dict_column = reader.fast_fields().str(field_name)?; + let value_sources = data.context.value_sources.clone(); let include_bytes = matches!(req, TermsOrCardinalityRequest::Terms(_)); - let column_and_types = get_term_agg_accessors(reader, field_name, missing, include_bytes)?; + let sources = + get_term_agg_accessors(reader, &value_sources, field_name, missing, include_bytes)?; // Special handling when missing + multi column or incompatible type on text/date. - let missing_and_more_than_one_col = column_and_types.len() > 1 && missing.is_some(); - let text_on_non_text_col = column_and_types.len() == 1 - && column_and_types[0].1 != ColumnType::Str + let missing_and_more_than_one_col = sources.len() > 1 && missing.is_some(); + let text_on_non_text_col = sources.len() == 1 + && sources[0].column_type() != ColumnType::Str && matches!(missing, Some(Key::Str(_))); let use_special_missing_agg = missing_and_more_than_one_col || text_on_non_text_col; @@ -1029,9 +1101,21 @@ fn build_terms_or_cardinality_nodes( Key::U64(_) => ColumnType::U64, }) .unwrap_or(ColumnType::U64); - let all_accessors = get_all_ff_reader_or_empty(reader, field_name, None, fallback_type)? - .into_iter() - .collect::>(); + // This path inspects `ColumnIndex::has_value` per document per accessor to decide which + // documents are missing across several typed columns. There is no way to ask that through + // `ValueSource`, so it stays physical-only. + let all_accessors = + get_all_value_sources(reader, &value_sources, field_name, None, fallback_type)? + .into_iter() + .map(|source| { + let column = require_physical_column( + &*source, + field_name, + "terms with `missing` across multiple column types", + )?; + Ok((column, source.column_type())) + }) + .collect::>>()?; // This case only happens when we have term aggregation, or we fail let req = req.as_terms().cloned().ok_or_else(|| { crate::TantivyError::InvalidArgument( @@ -1054,11 +1138,12 @@ fn build_terms_or_cardinality_nodes( } // Add one node per accessor - for (accessor, column_type) in column_and_types { + for accessor in sources { + let column_type = accessor.column_type(); let missing_value_for_accessor = if use_special_missing_agg { None } else if let Some(m) = missing.as_ref() { - get_missing_val_as_u64_lenient(column_type, accessor.max_value(), m, field_name)? + missing_value_for_source(&*accessor, m, field_name)? } else { None }; @@ -1085,7 +1170,6 @@ fn build_terms_or_cardinality_nodes( }; let idx_in_req_data = data.push_term_req_data(TermsAggReqData { accessor, - column_type, str_dict_column: str_dict_column.clone(), missing_value_for_accessor, name: agg_name.to_string(), @@ -1109,7 +1193,6 @@ fn build_terms_or_cardinality_nodes( }; let idx_in_req_data = data.push_cardinality_req_data(CardinalityAggReqData { accessor, - column_type, str_dict_column: str_dict_column_for_req, missing_value_for_accessor, name: agg_name.to_string(), diff --git a/src/aggregation/bucket/histogram/histogram.rs b/src/aggregation/bucket/histogram/histogram.rs index 196713730..80079503d 100644 --- a/src/aggregation/bucket/histogram/histogram.rs +++ b/src/aggregation/bucket/histogram/histogram.rs @@ -1,6 +1,7 @@ use std::cmp::Ordering; +use std::sync::Arc; -use columnar::{Column, ColumnType}; +use columnar::ColumnType; use rustc_hash::FxHashMap; use serde::{Deserialize, Serialize}; use tantivy_bitpacker::minmax; @@ -24,9 +25,7 @@ use crate::TantivyError; #[derive(Debug, Clone)] pub(crate) struct HistogramAggReqData { /// The column accessor to access the fast field values. - pub(crate) accessor: Column, - /// The field type of the fast field. - pub(crate) field_type: ColumnType, + pub(crate) accessor: Arc, /// The name of the aggregation. pub(crate) name: String, /// The histogram aggregation request. @@ -447,6 +446,7 @@ pub struct SegmentHistogramCollector { parent_buckets: Vec>, sub_agg: Option, req_data: HistogramAggReqData, + column_type: ColumnType, bucket_id_provider: BucketIdProvider, /// Theoretical bucket range derived from the column min/max, if dense `Vec` storage is /// viable. `None` keeps every parent bucket in the sparse hash map. @@ -491,11 +491,11 @@ impl SegmentAggregationCollector for SegmentHistogramCollector< agg_data .column_block_accessor - .fetch_block(docs, &req.accessor); + .fetch_block(docs, &*req.accessor); // special path for nested buckets if let Some(sub_agg) = &mut self.sub_agg { for (doc, val) in agg_data.column_block_accessor.iter_docid_vals(docs) { - let val = f64_from_fastfield_u64(val, req.field_type); + let val = f64_from_fastfield_u64(val, self.column_type); if bounds.contains(val) { let bucket = store.get_or_create( get_bucket_pos(val), @@ -508,7 +508,7 @@ impl SegmentAggregationCollector for SegmentHistogramCollector< } } else { for val in agg_data.column_block_accessor.iter_vals() { - let val = f64_from_fastfield_u64(val, req.field_type); + let val = f64_from_fastfield_u64(val, self.column_type); if bounds.contains(val) { let bucket = store.get_or_create( get_bucket_pos(val), @@ -586,7 +586,7 @@ impl SegmentHistogramCollector { } buckets.sort_unstable_by(|b1, b2| b1.key.total_cmp(&b2.key)); - let is_date_agg = self.req_data.field_type == ColumnType::DateTime; + let is_date_agg = self.req_data.accessor.column_type() == ColumnType::DateTime; Ok(IntermediateBucketResult::Histogram { buckets, is_date_agg, @@ -609,18 +609,18 @@ impl SegmentHistogramCollector { .limits .add_memory_consumed(req_data.get_memory_consumption() as u64)?; let dense_range = compute_dense_range( - &req_data.accessor, - req_data.field_type, + &*req_data.accessor, req_data.req.interval, req_data.offset, req_data.bounds, ); let sub_agg = sub_agg.map(BufferedSubAggs::new); - + let column_type = req_data.accessor.column_type(); Ok(Self { parent_buckets: Default::default(), sub_agg, req_data, + column_type, bucket_id_provider: BucketIdProvider::default(), dense_range, }) @@ -661,10 +661,12 @@ impl SegmentHistogramCollector<()> { HistogramBuckets::Dense { base_pos, buckets } }) .collect(); + let column_type = req_data.accessor.column_type(); Self { parent_buckets, sub_agg: None, req_data, + column_type, bucket_id_provider: BucketIdProvider::default(), dense_range: None, } @@ -675,7 +677,8 @@ impl SegmentHistogramCollector<()> { /// `histogram` on a date column) and resolves `bounds`/`offset` from the request. fn normalize_histogram_req(req_data: &mut HistogramAggReqData) -> crate::Result<()> { req_data.req.validate()?; - if req_data.field_type == ColumnType::DateTime && !req_data.is_date_histogram { + let field_type = req_data.accessor.column_type(); + if field_type == ColumnType::DateTime && !req_data.is_date_histogram { req_data.req.normalize_date_time(); } req_data.bounds = req_data.req.hard_bounds.unwrap_or(HistogramBounds { @@ -690,13 +693,15 @@ fn normalize_histogram_req(req_data: &mut HistogramAggReqData) -> crate::Result< // emission reads `req.hard_bounds` directly (see `get_req_min_max`), and `hard_bounds` only // ever clips that range, so a wider-than-data bound leaves the result unchanged. if req_data.req.hard_bounds.is_some() { - let col_min = f64_from_fastfield_u64(req_data.accessor.min_value(), req_data.field_type); - let col_max = f64_from_fastfield_u64(req_data.accessor.max_value(), req_data.field_type); - if col_min >= req_data.bounds.min && col_max <= req_data.bounds.max { - req_data.bounds = HistogramBounds { - min: f64::MIN, - max: f64::MAX, - }; + if let Some((min_value, max_value)) = req_data.accessor.bounds() { + let col_min = f64_from_fastfield_u64(min_value, field_type); + let col_max = f64_from_fastfield_u64(max_value, field_type); + if col_min >= req_data.bounds.min && col_max <= req_data.bounds.max { + req_data.bounds = HistogramBounds { + min: f64::MIN, + max: f64::MAX, + }; + } } } Ok(()) @@ -712,8 +717,7 @@ pub(crate) fn prepare_histogram_dense_range( let mut req_data = agg_data.per_request.histogram_req_data[node.idx_in_req_data].clone(); normalize_histogram_req(&mut req_data)?; let dense_range = compute_dense_range( - &req_data.accessor, - req_data.field_type, + &*req_data.accessor, req_data.req.interval, req_data.offset, req_data.bounds, @@ -752,15 +756,19 @@ pub(crate) fn get_bucket_pos_f64(val: f64, interval: f64, offset: f64) -> f64 { /// /// The column min/max bound every value the collector can see, so a `Vec` sized to this range can /// be indexed by `bucket_pos - base_pos` without any out-of-bounds check on the hot path. +/// +/// Returns `None` for a computed source: there is no global range to size the `Vec` from, so the +/// histogram keeps its sparse map. The result is identical, just without the dense fast path. fn compute_dense_range( - accessor: &Column, - field_type: ColumnType, + accessor: &dyn ValueSource, interval: f64, offset: f64, bounds: HistogramBounds, ) -> Option { - let col_min = f64_from_fastfield_u64(accessor.min_value(), field_type); - let col_max = f64_from_fastfield_u64(accessor.max_value(), field_type); + let (min_value, max_value) = accessor.bounds()?; + let field_type = accessor.column_type(); + let col_min = f64_from_fastfield_u64(min_value, field_type); + let col_max = f64_from_fastfield_u64(max_value, field_type); let lo = col_min.max(bounds.min); let hi = col_max.min(bounds.max); if lo > hi { diff --git a/src/aggregation/bucket/multi_terms/mod.rs b/src/aggregation/bucket/multi_terms/mod.rs index d8467961b..0b2d53c23 100644 --- a/src/aggregation/bucket/multi_terms/mod.rs +++ b/src/aggregation/bucket/multi_terms/mod.rs @@ -229,7 +229,7 @@ fn fetch_field_block( let missing_value = block_missing_value(missing); block_accessor.fetch_block_with_missing_unique_per_doc( docs, - &field.column, + &(&field.column, field.column_type), missing_value, true, ); diff --git a/src/aggregation/bucket/range.rs b/src/aggregation/bucket/range.rs index 5857d5a42..e2b581367 100644 --- a/src/aggregation/bucket/range.rs +++ b/src/aggregation/bucket/range.rs @@ -1,7 +1,8 @@ use std::fmt::Debug; use std::ops::Range; +use std::sync::Arc; -use columnar::{Column, ColumnType}; +use columnar::ColumnType; use rustc_hash::FxHashMap; use serde::{Deserialize, Serialize}; @@ -18,6 +19,7 @@ use crate::aggregation::intermediate_agg_result::{ IntermediateRangeBucketEntry, IntermediateRangeBucketResult, }; use crate::aggregation::segment_agg_result::{BucketIdProvider, SegmentAggregationCollector}; +use crate::aggregation::value_source::ValueSource; use crate::aggregation::*; use crate::TantivyError; @@ -26,9 +28,7 @@ use crate::TantivyError; #[derive(Debug, Clone)] pub(crate) struct RangeAggReqData { /// The column accessor to access the fast field values. - pub(crate) accessor: Column, - /// The type of the fast field. - pub(crate) field_type: ColumnType, + pub(crate) accessor: Arc, /// The range aggregation request. pub(crate) req: RangeAggregation, /// The name of the aggregation. @@ -161,7 +161,6 @@ pub struct SegmentRangeCollector { /// The buckets containing the aggregation data. /// One for each ParentBucketId parent_buckets: Vec>, - column_type: ColumnType, pub(crate) req_data: RangeAggReqData, sub_agg: Option>, /// Here things get a bit weird. We need to assign unique bucket ids across all @@ -184,7 +183,7 @@ impl Debug for SegmentRangeCollector { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("SegmentRangeCollector") .field("parent_buckets_len", &self.parent_buckets.len()) - .field("column_type", &self.column_type) + .field("column_type", &self.req_data.accessor.column_type()) .field("name", &self.req_data.name) .field("has_sub_agg", &self.sub_agg.is_some()) .finish() @@ -239,7 +238,7 @@ impl SegmentAggregationCollector for SegmentRangeCollector { parent_bucket_id: BucketId, ) -> crate::Result<()> { self.prepare_max_bucket(parent_bucket_id, agg_data)?; - let field_type = self.column_type; + let field_type = self.req_data.accessor.column_type(); let name = self.req_data.name.to_string(); let buckets = std::mem::take(&mut self.parent_buckets[parent_bucket_id as usize]); @@ -264,7 +263,7 @@ impl SegmentAggregationCollector for SegmentRangeCollector { let bucket = IntermediateBucketResult::Range(IntermediateRangeBucketResult { buckets, - column_type: Some(self.column_type), + column_type: Some(field_type), }); results.push(name, IntermediateAggregationResult::Bucket(bucket))?; @@ -281,7 +280,7 @@ impl SegmentAggregationCollector for SegmentRangeCollector { ) -> crate::Result<()> { agg_data .column_block_accessor - .fetch_block(docs, &self.req_data.accessor); + .fetch_block(docs, &*self.req_data.accessor); let buckets = &mut self.parent_buckets[parent_bucket_id as usize]; @@ -343,7 +342,6 @@ pub(crate) fn build_segment_range_collector( .context .limits .add_memory_consumed(req_data.get_memory_consumption() as u64)?; - let field_type = req_data.field_type; // TODO: A better metric instead of is_top_level would be the number of buckets expected. // E.g. If range agg is not top level, but the parent is a bucket agg with less than 10 buckets, @@ -359,7 +357,6 @@ pub(crate) fn build_segment_range_collector( if is_low_card { Ok(Box::new(SegmentRangeCollector:: { sub_agg: sub_agg.map(LowCardBufferedSubAggs::new), - column_type: field_type, req_data, parent_buckets: Vec::new(), bucket_id_provider: BucketIdProvider::default(), @@ -368,7 +365,6 @@ pub(crate) fn build_segment_range_collector( } else { Ok(Box::new(SegmentRangeCollector:: { sub_agg: sub_agg.map(BufferedSubAggs::new), - column_type: field_type, req_data, parent_buckets: Vec::new(), bucket_id_provider: BucketIdProvider::default(), @@ -379,8 +375,8 @@ pub(crate) fn build_segment_range_collector( impl SegmentRangeCollector { pub(crate) fn create_new_buckets(&mut self) -> crate::Result> { - let field_type = self.column_type; let req_data = &self.req_data; + let field_type = req_data.accessor.column_type(); // The range input on the request is f64. // We need to convert to u64 ranges, because we read the values as u64. // The mapping from the conversion is monotonic so ordering is preserved. diff --git a/src/aggregation/bucket/term_agg/flattened_term_histogram.rs b/src/aggregation/bucket/term_agg/flattened_term_histogram.rs index 675171448..3f741539b 100644 --- a/src/aggregation/bucket/term_agg/flattened_term_histogram.rs +++ b/src/aggregation/bucket/term_agg/flattened_term_histogram.rs @@ -5,8 +5,9 @@ //! [`maybe_build_flattened_collector`] for the conditions under which it is used. use std::fmt::Debug; +use std::sync::Arc; -use columnar::{Column, ColumnType}; +use columnar::{Cardinality, ColumnType, ColumnValues}; use super::{ Bucket, SegmentTermCollector, TermsAggReqData, VecTermBuckets, MAX_NUM_BUCKETS_FOR_COUNT_LANES, @@ -14,7 +15,7 @@ use super::{ }; use crate::aggregation::agg_data::{AggKind, AggRefNode, AggregationsSegmentCtx}; use crate::aggregation::bucket::{ - get_bucket_pos_f64, prepare_histogram_dense_range, HistogramAggReqData, + get_bucket_pos_f64, prepare_histogram_dense_range, DenseRange, HistogramAggReqData, SegmentHistogramCollector, }; use crate::aggregation::buffered_sub_aggs::LowCardSubAggBuffer; @@ -22,7 +23,8 @@ use crate::aggregation::intermediate_agg_result::{ IntermediateAggregationResult, IntermediateAggregationResults, }; use crate::aggregation::segment_agg_result::{BucketIdProvider, SegmentAggregationCollector}; -use crate::aggregation::{f64_from_fastfield_u64, BucketId, ColumnBlockAccessor}; +use crate::aggregation::value_source::ColumnBlockAccessor; +use crate::aggregation::{f64_from_fastfield_u64, BucketId, ValueSource}; /// Maximum number of physical counters in the flattened flat grid. Above this the grid would be too /// large/cache-unfriendly, so we fall back to the general buffered path. Count lanes are included @@ -39,7 +41,7 @@ const SINGLE_COUNT_LANE: usize = 1; const NUM_SMALL_LINEAR_BUCKETS: usize = 4; const NUM_LARGE_LINEAR_BUCKETS: usize = 8; -trait BucketResolver: Debug + 'static { +trait BucketResolver: 'static { /// Fetches the histogram values needed for this block. Resolvers that do not inspect the /// histogram column (notably [`FlattenedSingleBucketResolver`]) leave this as a no-op. fn prepare_block(&mut self, docs: &[crate::DocId]); @@ -81,21 +83,11 @@ fn increment_grid_count( /// Resolver for a histogram whose entire value range maps to one bucket. It deliberately owns no /// block accessor: collecting this shape does not read or decode the histogram column at all. -#[derive(Debug)] +#[derive(Default)] struct SingleBucketResolver { next_count_lane: usize, } -impl SingleBucketResolver { - fn new(hist_req_data: &HistogramAggReqData) -> Self { - assert!( - hist_req_data.accessor.get_cardinality().is_full(), - "SingleBucketResolver requires a full histogram column" - ); - Self { next_count_lane: 0 } - } -} - impl BucketResolver for SingleBucketResolver { #[inline] fn prepare_block(&mut self, _docs: &[crate::DocId]) {} @@ -129,11 +121,10 @@ impl BucketResolver for SingleBucketResolver { /// The general resolver. It preserves the existing field conversion and floating-point bucket /// calculation for histograms that do not use a specialized resolver. -#[derive(Debug)] struct ComputedBucketResolver { hist_block: ColumnBlockAccessor, next_count_lane: usize, - accessor: Column, + column_values: Arc, field_type: ColumnType, interval: f64, offset: f64, @@ -143,16 +134,17 @@ struct ComputedBucketResolver { } impl ComputedBucketResolver { - fn new(hist_req_data: &HistogramAggReqData, base_pos: i64, num_buckets: usize) -> Self { - assert!( - hist_req_data.accessor.get_cardinality().is_full(), - "ComputedBucketResolver requires a full histogram column" - ); + fn new( + hist_req_data: &HistogramAggReqData, + hist_values: Arc, + base_pos: i64, + num_buckets: usize, + ) -> Self { Self { hist_block: ColumnBlockAccessor::default(), next_count_lane: 0, - accessor: hist_req_data.accessor.clone(), - field_type: hist_req_data.field_type, + column_values: hist_values, + field_type: hist_req_data.accessor.column_type(), interval: hist_req_data.req.interval, offset: hist_req_data.offset, base_pos, @@ -166,7 +158,7 @@ impl BucketResolver for ComputedBucketResolver { #[inline] fn prepare_block(&mut self, docs: &[crate::DocId]) { self.hist_block - .fetch_full_column_block(docs, &self.accessor); + .fetch_full_column_block(docs, &*self.column_values); } #[inline] @@ -223,32 +215,30 @@ impl BucketResolver for ComputedBucketResolver { /// Resolver for a small histogram grid. Bucket starts are precomputed in monotonic fast-field /// `u64` space, then scanned linearly. `NUM_BUCKETS` is fixed so the optimizer can unroll the scan. -#[derive(Debug)] struct LinearBucketResolver { hist_block: ColumnBlockAccessor, next_count_lane: usize, - accessor: Column, + column_values: Arc, boundaries: [u64; NUM_BUCKETS], num_buckets: usize, } impl LinearBucketResolver { + /// Returns `None` only when no padding sentinel exists (see below), in which case the caller + /// falls back to [`ComputedBucketResolver`]. fn new( hist_req_data: &HistogramAggReqData, + column_values: Arc, base_pos: i64, num_time_buckets: usize, ) -> Option { assert!(num_time_buckets > 1 && num_time_buckets <= NUM_BUCKETS); - assert!( - hist_req_data.accessor.get_cardinality().is_full(), - "LinearBucketResolver requires a full histogram column" - ); - let max_encoded_value = hist_req_data.accessor.max_value(); + let max_encoded_value = column_values.max_value(); // Padding must compare false for every column value. There is no such `u64` sentinel when // the column contains `u64::MAX`, so that edge case uses the computed resolver instead. let padding = max_encoded_value.checked_add(1)?; let mut boundaries = [padding; NUM_BUCKETS]; - let mut bucket_start = hist_req_data.accessor.min_value(); + let mut bucket_start = column_values.min_value(); for bucket in 1..num_time_buckets { bucket_start = first_encoded_value_for_bucket( bucket_start, @@ -262,7 +252,7 @@ impl LinearBucketResolver { Some(Self { hist_block: ColumnBlockAccessor::default(), next_count_lane: 0, - accessor: hist_req_data.accessor.clone(), + column_values, boundaries, num_buckets: num_time_buckets, }) @@ -282,7 +272,7 @@ impl BucketResolver for LinearBucketResolver u64 { + let field_type = hist_req_data.accessor.column_type(); while encoded_lower_bound < encoded_upper_bound { let encoded_midpoint = encoded_lower_bound + (encoded_upper_bound - encoded_lower_bound) / 2; - let val = f64_from_fastfield_u64(encoded_midpoint, hist_req_data.field_type); + let val = f64_from_fastfield_u64(encoded_midpoint, field_type); let bucket = (get_bucket_pos_f64(val, hist_req_data.req.interval, hist_req_data.offset) as i64 - base_pos) as usize; @@ -359,7 +350,6 @@ fn first_encoded_value_for_bucket( /// At result time the flat grid is expanded back into the regular term map + histogram storage and /// handed to the shared intermediate-result builders, so cross-segment merging is identical to the /// general path. -#[derive(Debug)] struct FlattenedTermHistogramCollector { /// Per-term count of docs *outside* `hard_bounds` (still in `doc_count`, but in no bucket). /// Per-term total = this + the term's `counts` row-sum; left empty when there are no hard @@ -375,7 +365,8 @@ struct FlattenedTermHistogramCollector { /// `bucket_pos` mapped to time-bucket index 0. base_pos: i64, terms_req_data: TermsAggReqData, - /// The (cloned, normalized) histogram request: its column + interval/offset/bounds. + /// The terms full column's values + terms_values: Arc, hist_req_data: HistogramAggReqData, /// Private term block accessor. The bucket resolver owns a histogram block accessor when it /// needs one; the single-bucket resolver deliberately does not. @@ -385,6 +376,15 @@ struct FlattenedTermHistogramCollector { all_docs_in_bounds: bool, } +impl Debug for FlattenedTermHistogramCollector { + fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result { + f.debug_struct("FlattenedTermHistogramCollector") + .field("base_pos", &self.base_pos) + .field("all_docs_in_bounds", &self.all_docs_in_bounds) + .finish_non_exhaustive() + } +} + impl SegmentAggregationCollector for FlattenedTermHistogramCollector { @@ -455,7 +455,7 @@ impl SegmentAggregationCollector // The term column is always needed. The resolver fetches the histogram column only when // bucket selection depends on its values; `SingleBucketResolver` makes this a no-op. self.term_block - .fetch_full_column_block(docs, &self.terms_req_data.accessor); + .fetch_full_column_block(docs, &*self.terms_values); self.bucket_resolver.prepare_block(docs); // Keep separate bounded and unbounded entry points so the common path has no bounds branch, @@ -527,25 +527,26 @@ pub(super) fn maybe_build_flattened_collector( let fuseable = is_top_level // TODO: We can easily support this && terms_req_data.allowed_term_ids.is_none() - && terms_req_data.accessor.get_cardinality().is_full() - // The flat counters are `u32`, bumped once per value, so no count can exceed the column's - // value count. (Essentially always true here: the column is full, so its value count - // equals the doc count, and `DocId` is `u32`.) - && terms_req_data.accessor.values.num_vals() < u32::MAX && node.children.len() == 1 && matches!( node.children[0].kind, AggKind::Histogram | AggKind::DateHistogram ) - && node.children[0].children.is_empty() - && agg_data.per_request.histogram_req_data[node.children[0].idx_in_req_data] - .accessor - .get_cardinality() - .is_full(); + && node.children[0].children.is_empty(); if !fuseable { return Ok(None); } + // Check fullness once, here, and hand the bare values array downstream. + let Some(terms_values) = try_get_column_full_values(&*terms_req_data.accessor) else { + return Ok(None); + }; + // The flat counters are `u32`, bumped once per value, so no count can exceed the + // column's value count. (Essentially always true here: the column is full, so its + // value count equals the doc count, and `DocId` is `u32`.) + if terms_values.num_vals() == u32::MAX { + return Ok(None); + } // Clone + normalize the histogram request and get its dense bucket range; only take the // flattened path when the physical counter grid is small enough. Very small logical grids use // multiple counters per cell; larger grids retain scalar cells to avoid paying for lanes when @@ -554,6 +555,10 @@ pub(super) fn maybe_build_flattened_collector( else { return Ok(None); }; + let Some(hist_values) = try_get_column_full_values(&*hist_req_data.accessor) else { + return Ok(None); + }; + let num_terms = col_max_val.saturating_add(1) as usize; let num_grid_cells = num_terms.saturating_mul(range.len); let use_count_lanes = num_grid_cells <= MAX_NUM_BUCKETS_FOR_COUNT_LANES; @@ -571,40 +576,55 @@ pub(super) fn maybe_build_flattened_collector( agg_data, terms_req_data, hist_req_data, + terms_values, + hist_values, num_terms, - range.len, - range.base_pos, + range, )? } else { build_flattened_collector::( agg_data, terms_req_data, hist_req_data, + terms_values, + hist_values, num_terms, - range.len, - range.base_pos, + range, )? }; Ok(Some(collector)) } +fn try_get_column_full_values(value_source: &dyn ValueSource) -> Option> { + let column = value_source.as_column()?; + if column.get_cardinality() == Cardinality::Full { + Some(column.values.clone()) + } else { + None + } +} + fn build_flattened_collector( agg_data: &mut AggregationsSegmentCtx, terms_req_data: &TermsAggReqData, hist_req_data: HistogramAggReqData, + terms_values: Arc, + hist_values: Arc, num_terms: usize, - num_time_buckets: usize, - base_pos: i64, + range: DenseRange, ) -> crate::Result> { + let num_time_buckets = range.len; + let base_pos = range.base_pos; const { assert!(LANES > 0, "a flattened grid needs at least one count lane") }; let all_docs_in_bounds = hist_req_data.bounds.min == f64::MIN && hist_req_data.bounds.max == f64::MAX; if all_docs_in_bounds && num_time_buckets == 1 { - let resolver = SingleBucketResolver::new(&hist_req_data); + let resolver = SingleBucketResolver::default(); return build_flattened_collector_with_resolver::( agg_data, terms_req_data, + terms_values, hist_req_data, num_terms, base_pos, @@ -614,12 +634,14 @@ fn build_flattened_collector( if all_docs_in_bounds && num_time_buckets <= NUM_SMALL_LINEAR_BUCKETS { if let Some(resolver) = LinearBucketResolver::::new( &hist_req_data, + hist_values.clone(), base_pos, num_time_buckets, ) { return build_flattened_collector_with_resolver::<_, LANES>( agg_data, terms_req_data, + terms_values, hist_req_data, num_terms, base_pos, @@ -629,12 +651,14 @@ fn build_flattened_collector( } else if all_docs_in_bounds && num_time_buckets <= NUM_LARGE_LINEAR_BUCKETS { if let Some(resolver) = LinearBucketResolver::::new( &hist_req_data, + hist_values.clone(), base_pos, num_time_buckets, ) { return build_flattened_collector_with_resolver::<_, LANES>( agg_data, terms_req_data, + terms_values, hist_req_data, num_terms, base_pos, @@ -643,10 +667,16 @@ fn build_flattened_collector( } } - let resolver = ComputedBucketResolver::new(&hist_req_data, base_pos, num_time_buckets); + let resolver = ComputedBucketResolver::new( + &hist_req_data, + hist_values.clone(), + base_pos, + num_time_buckets, + ); build_flattened_collector_with_resolver::<_, LANES>( agg_data, terms_req_data, + terms_values, hist_req_data, num_terms, base_pos, @@ -657,6 +687,7 @@ fn build_flattened_collector( fn build_flattened_collector_with_resolver( agg_data: &mut AggregationsSegmentCtx, terms_req_data: &TermsAggReqData, + terms_values: Arc>, hist_req_data: HistogramAggReqData, num_terms: usize, base_pos: i64, @@ -686,6 +717,7 @@ fn build_flattened_collector_with_resolver, - /// The type of the column. - pub(crate) column_type: ColumnType, + pub(crate) accessor: Arc, /// The string dictionary column if the field is of type text. pub(crate) str_dict_column: Option, /// The missing value as u64 value. @@ -394,7 +393,7 @@ pub(crate) fn build_segment_term_collector( node: &AggRefNode, ) -> crate::Result> { let terms_req_data = req_data.get_term_req_data(node.idx_in_req_data).clone(); - let column_type = terms_req_data.column_type; + let column_type = terms_req_data.accessor.column_type(); if column_type == ColumnType::Bytes { return Err(TantivyError::InvalidArgument(format!( @@ -423,7 +422,11 @@ pub(crate) fn build_segment_term_collector( // Let's see if we can use a vec to aggregate our data // instead of a hashmap. - let col_max_value = terms_req_data.accessor.max_value(); + let col_max_value = terms_req_data + .accessor + .as_column() + .map(|col| col.max_value()) + .unwrap_or(u64::MAX); let max_column_val: u64 = col_max_value.max(terms_req_data.missing_value_for_accessor.unwrap_or(0u64)); @@ -1064,7 +1067,7 @@ impl SegmentAggregationCollector .column_block_accessor .fetch_block_with_missing_unique_per_doc( docs, - &req_data.accessor, + &*req_data.accessor, req_data.missing_value_for_accessor, false, ); @@ -1331,7 +1334,8 @@ where let mut out: Vec<(IntermediateKey, IntermediateTermBucketEntry)> = Vec::with_capacity(entries.len()); - if term_req.column_type == ColumnType::Str { + let column_type = term_req.accessor.column_type(); + if column_type == ColumnType::Str { let fallback_dict = Dictionary::empty(); let term_dict = term_req .str_dict_column @@ -1418,7 +1422,7 @@ where } out.extend(dict); - } else if term_req.column_type == ColumnType::DateTime { + } else if column_type == ColumnType::DateTime { for (val, doc_count) in entries { let intermediate_entry = into_intermediate_bucket_entry( doc_count, @@ -1429,7 +1433,7 @@ where let date = format_date(val)?; out.push((IntermediateKey::Str(date), intermediate_entry)); } - } else if term_req.column_type == ColumnType::Bool { + } else if column_type == ColumnType::Bool { for (val, doc_count) in entries { let intermediate_entry = into_intermediate_bucket_entry( doc_count, @@ -1439,9 +1443,17 @@ where let val = bool::from_u64(val); out.push((IntermediateKey::Bool(val), intermediate_entry)); } - } else if term_req.column_type == ColumnType::IpAddr { + } else if column_type == ColumnType::IpAddr { let compact_space_accessor = term_req .accessor + .as_column() + .ok_or_else(|| { + TantivyError::AggregationError( + crate::aggregation::AggregationError::InternalError( + "IpAddr term keys require a physical column".to_string(), + ), + ) + })? .values .clone() .downcast_arc::() @@ -1464,7 +1476,7 @@ where let val = Ipv6Addr::from_u128(val); out.push((IntermediateKey::IpAddr(val), intermediate_entry)); } - } else if term_req.column_type == ColumnType::F64 { + } else if column_type == ColumnType::F64 { // -0.0 and +0.0 both normalize to I64(0). Sort by normalized key to merge their // buckets below. Other distinct f64 encodings, including NaNs, remain distinct: // NaNs are not normalized and IntermediateKey compares them using total_cmp. @@ -1497,13 +1509,13 @@ where reborrow_opt_collector(&mut sub_agg_collector), agg_data, )?; - let key_val: NumericalValue = match term_req.column_type { + let key_val: NumericalValue = match column_type { ColumnType::U64 => val.into(), ColumnType::I64 => i64::from_u64(val).into(), _ => { return Err(TantivyError::SchemaError(format!( "unknown key type: {}", - term_req.column_type + column_type ))) } }; diff --git a/src/aggregation/metric/cardinality/mod.rs b/src/aggregation/metric/cardinality/mod.rs index 84d9759d3..8440cd948 100644 --- a/src/aggregation/metric/cardinality/mod.rs +++ b/src/aggregation/metric/cardinality/mod.rs @@ -16,8 +16,9 @@ mod str_collector; mod term_ord_accumulator; use std::hash::Hash; +use std::sync::Arc; -use columnar::{Column, ColumnType, StrColumn}; +use columnar::{ColumnType, StrColumn}; use common::BitSet; use datasketches::hll::{Coupon, HllSketch, HllType, HllUnion}; pub(crate) use numeric_collector::SegmentNumericCardinalityCollector; @@ -27,6 +28,7 @@ pub(crate) use term_ord_accumulator::{TermOrdSet, BITSET_MAX_TERM_ORD}; use crate::aggregation::agg_data::{AggRefNode, AggregationsSegmentCtx}; use crate::aggregation::segment_agg_result::SegmentAggregationCollector; +use crate::aggregation::value_source::ValueSource; use crate::aggregation::*; /// Log2 of the number of registers for the HLL sketch. @@ -100,9 +102,7 @@ pub struct CardinalityAggregationReq { /// cardinality aggregation on a segment. pub(crate) struct CardinalityAggReqData { /// The column accessor to access the fast field values. - pub(crate) accessor: Column, - /// The column_type of the field. - pub(crate) column_type: ColumnType, + pub(crate) accessor: Arc, /// The string dictionary column if the field is of type string. pub(crate) str_dict_column: Option, /// The missing value normalized to the internal u64 representation of the field type. @@ -224,9 +224,8 @@ pub(crate) fn build_segment_cardinality_collector( node: &AggRefNode, ) -> crate::Result> { let req_data = req.get_cardinality_req_data(node.idx_in_req_data); - if req_data.column_type != ColumnType::Str { + if req_data.accessor.column_type() != ColumnType::Str { return Ok(Box::new(SegmentNumericCardinalityCollector::from_req( - req_data.column_type, node.idx_in_req_data, req_data.accessor.clone(), req_data.missing_value_for_accessor, @@ -237,7 +236,12 @@ pub(crate) fn build_segment_cardinality_collector( // number of terms. // * small (< BITSET_MAX_TERM_ORD): `BitSet`, pre-allocated. // * large: `TermOrdSet` (sparse HashSet that promotes to a paged bitset). - let max_term_ord_inclusive = req_data.accessor.max_value(); + let Some(column) = req_data.accessor.as_column() else { + return Err(crate::TantivyError::InvalidArgument( + "cardinality over str virtual columns is not supported yet".to_string(), + )); + }; + let max_term_ord_inclusive = column.max_value(); if max_term_ord_inclusive < BITSET_MAX_TERM_ORD { Ok(Box::new( SegmentStrCardinalityCollector::::from_req( diff --git a/src/aggregation/metric/cardinality/numeric_collector.rs b/src/aggregation/metric/cardinality/numeric_collector.rs index 072f922a6..d6b3fa7bc 100644 --- a/src/aggregation/metric/cardinality/numeric_collector.rs +++ b/src/aggregation/metric/cardinality/numeric_collector.rs @@ -10,7 +10,7 @@ use std::fmt::Debug; use std::sync::Arc; use columnar::column_values::CompactSpaceU64Accessor; -use columnar::{Column, ColumnType}; +use columnar::ColumnType; use super::CardinalityCollector; use crate::aggregation::agg_data::AggregationsSegmentCtx; @@ -18,6 +18,7 @@ use crate::aggregation::intermediate_agg_result::{ IntermediateAggregationResult, IntermediateAggregationResults, IntermediateMetricResult, }; use crate::aggregation::segment_agg_result::SegmentAggregationCollector; +use crate::aggregation::value_source::ValueSource; use crate::aggregation::*; use crate::TantivyError; @@ -33,7 +34,7 @@ pub(crate) struct SegmentNumericCardinalityCollector { buckets: Vec>, accessor_idx: usize, /// The column accessor to access the fast field values. - accessor: Column, + accessor: Arc, /// The column_type of the field. column_type: ColumnType, /// Set iff `column_type == ColumnType::IpAddr`. Resolved once at @@ -58,14 +59,22 @@ impl Debug for SegmentNumericCardinalityCollector { impl SegmentNumericCardinalityCollector { pub fn from_req( - column_type: ColumnType, accessor_idx: usize, - accessor: Column, + accessor: Arc, missing_value_for_accessor: Option, ) -> crate::Result { + let column_type = accessor.column_type(); assert_ne!(column_type, ColumnType::Str); let compact_space_accessor = if column_type == ColumnType::IpAddr { let compact_space_accessor = accessor + .as_column() + .ok_or_else(|| { + TantivyError::AggregationError( + crate::aggregation::AggregationError::InternalError( + "IpAddr cardinality requires a physical column".to_string(), + ), + ) + })? .values .clone() .downcast_arc::() @@ -127,7 +136,7 @@ impl SegmentAggregationCollector for SegmentNumericCardinalityCollector { ) -> crate::Result<()> { agg_data.column_block_accessor.fetch_block_with_missing( docs, - &self.accessor, + &*self.accessor, self.missing_value_for_accessor, ); let cardinality = self.buckets[parent_bucket_id as usize] diff --git a/src/aggregation/metric/cardinality/str_collector.rs b/src/aggregation/metric/cardinality/str_collector.rs index 578e029f1..b2ec7d218 100644 --- a/src/aggregation/metric/cardinality/str_collector.rs +++ b/src/aggregation/metric/cardinality/str_collector.rs @@ -10,8 +10,9 @@ use std::fmt::Debug; use std::io; +use std::sync::Arc; -use columnar::{Column, ColumnType, Dictionary}; +use columnar::{ColumnType, Dictionary}; use datasketches::hll::Coupon; use rustc_hash::{FxBuildHasher, FxHashMap, FxHashSet}; @@ -22,6 +23,7 @@ use crate::aggregation::intermediate_agg_result::{ IntermediateAggregationResult, IntermediateAggregationResults, IntermediateMetricResult, }; use crate::aggregation::segment_agg_result::SegmentAggregationCollector; +use crate::aggregation::value_source::ValueSource; use crate::aggregation::*; /// A CouponCache is here to cache the mapping term ordinal -> coupon (see above). @@ -99,7 +101,7 @@ pub(crate) struct SegmentStrCardinalityCollector { buckets: Vec>, accessor_idx: usize, /// The column accessor to access the fast field values (term ordinals). - accessor: Column, + accessor: Arc, /// The missing value normalized to the internal u64 representation of the field type. missing_value_for_accessor: Option, /// Lazily built at finalization time, shared by every bucket. @@ -209,7 +211,7 @@ fn append_to_sketch( impl SegmentStrCardinalityCollector { pub fn from_req( accessor_idx: usize, - accessor: Column, + accessor: Arc, missing_value_for_accessor: Option, max_term_ord_inclusive: u64, ) -> Self { @@ -282,7 +284,7 @@ impl SegmentAggregationCollector ) -> crate::Result<()> { agg_data.column_block_accessor.fetch_block_with_missing( docs, - &self.accessor, + &*self.accessor, self.missing_value_for_accessor, ); let Some(term_ords) = self.buckets[parent_bucket_id as usize].as_mut() else { diff --git a/src/aggregation/metric/extended_stats.rs b/src/aggregation/metric/extended_stats.rs index 1e625d5de..3747a0921 100644 --- a/src/aggregation/metric/extended_stats.rs +++ b/src/aggregation/metric/extended_stats.rs @@ -1,5 +1,6 @@ use std::fmt::Debug; use std::mem; +use std::sync::Arc; use serde::{Deserialize, Serialize}; @@ -321,8 +322,7 @@ impl IntermediateExtendedStats { pub(crate) struct SegmentExtendedStatsCollector { name: String, missing: Option, - field_type: ColumnType, - accessor: columnar::Column, + accessor: Arc, buckets: Vec, sigma: Option, } @@ -331,10 +331,9 @@ impl SegmentExtendedStatsCollector { pub fn from_req(req: &MetricAggReqData, sigma: Option) -> Self { let missing = req .missing - .and_then(|val| f64_to_fastfield_u64(val, &req.field_type)); + .and_then(|val| f64_to_fastfield_u64(val, &req.accessor.column_type())); Self { name: req.name.clone(), - field_type: req.field_type, accessor: req.accessor.clone(), missing, buckets: vec![IntermediateExtendedStats::with_sigma(sigma); 16], @@ -373,11 +372,14 @@ impl SegmentAggregationCollector for SegmentExtendedStatsCollector { ) -> crate::Result<()> { let mut extended_stats = self.buckets[parent_bucket_id as usize].clone(); - agg_data - .column_block_accessor - .fetch_block_with_missing(docs, &self.accessor, self.missing); + agg_data.column_block_accessor.fetch_block_with_missing( + docs, + &*self.accessor, + self.missing, + ); + let field_type = self.accessor.column_type(); for val in agg_data.column_block_accessor.iter_vals() { - let val1 = f64_from_fastfield_u64(val, self.field_type); + let val1 = f64_from_fastfield_u64(val, field_type); extended_stats.collect(val1); } diff --git a/src/aggregation/metric/mod.rs b/src/aggregation/metric/mod.rs index 56ca63f01..16197a7a2 100644 --- a/src/aggregation/metric/mod.rs +++ b/src/aggregation/metric/mod.rs @@ -28,10 +28,10 @@ mod sum; mod top_hits; use std::collections::HashMap; +use std::sync::Arc; pub use average::*; pub use cardinality::*; -use columnar::{Column, ColumnType}; pub use count::*; pub use extended_stats::*; pub use max::*; @@ -43,6 +43,7 @@ pub use stats::*; pub use sum::*; pub use top_hits::*; +use crate::aggregation::value_source::ValueSource; use crate::schema::OwnedValue; /// Contains all information required by metric aggregations like avg, min, max, sum, stats, @@ -51,12 +52,10 @@ use crate::schema::OwnedValue; pub(crate) struct MetricAggReqData { /// True if the field is of number or date type. pub(crate) is_number_or_date_type: bool, - /// The type of the field. - pub(crate) field_type: ColumnType, /// The missing value normalized to the internal u64 representation of the field type. pub(crate) missing_u64: Option, /// The column accessor to access the fast field values. - pub(crate) accessor: Column, + pub(crate) accessor: Arc, /// Used when converting to intermediate result pub(crate) collecting_for: StatsType, /// The missing value diff --git a/src/aggregation/metric/percentiles.rs b/src/aggregation/metric/percentiles.rs index 4d213c927..e2817617e 100644 --- a/src/aggregation/metric/percentiles.rs +++ b/src/aggregation/metric/percentiles.rs @@ -1,4 +1,5 @@ use std::fmt::Debug; +use std::sync::Arc; use serde::{Deserialize, Serialize}; @@ -134,12 +135,10 @@ impl PercentilesAggregationReq { pub(crate) struct SegmentPercentilesCollector { pub(crate) buckets: Vec, pub(crate) accessor_idx: usize, - /// The type of the field. - pub field_type: ColumnType, /// The missing value normalized to the internal u64 representation of the field type. pub missing_u64: Option, /// The column accessor to access the fast field values. - pub(crate) accessor: Column, + pub(crate) accessor: Arc, } #[derive(Clone, Serialize, Deserialize)] @@ -250,14 +249,12 @@ impl PercentilesCollector { impl SegmentPercentilesCollector { pub fn from_req_and_validate( - field_type: ColumnType, missing_u64: Option, - accessor: Column, + accessor: Arc, accessor_idx: usize, ) -> Self { Self { buckets: Vec::with_capacity(64), - field_type, missing_u64, accessor, accessor_idx, @@ -299,12 +296,13 @@ impl SegmentAggregationCollector for SegmentPercentilesCollector { let percentiles = &mut self.buckets[parent_bucket_id as usize]; agg_data.column_block_accessor.fetch_block_with_missing( docs, - &self.accessor, + &*self.accessor, self.missing_u64, ); + let field_type = self.accessor.column_type(); for val in agg_data.column_block_accessor.iter_vals() { - let val1 = f64_from_fastfield_u64(val, self.field_type); + let val1 = f64_from_fastfield_u64(val, field_type); percentiles.collect(val1); } diff --git a/src/aggregation/metric/stats.rs b/src/aggregation/metric/stats.rs index 06a39f6d1..30ec18e33 100644 --- a/src/aggregation/metric/stats.rs +++ b/src/aggregation/metric/stats.rs @@ -1,4 +1,5 @@ use std::fmt::Debug; +use std::sync::Arc; use columnar::{Column, ColumnType}; use serde::{Deserialize, Serialize}; @@ -205,6 +206,10 @@ fn create_collector( collecting_for: req.collecting_for, is_number_or_date_type: req.is_number_or_date_type, missing_u64: req.missing_u64, + column_opt: req + .accessor + .as_column() + .map(|column| (column.clone(), req.accessor.column_type())), accessor: req.accessor.clone(), buckets: vec![IntermediateStats::default()], }) @@ -214,7 +219,7 @@ fn create_collector( pub(crate) fn build_segment_stats_collector( req: &MetricAggReqData, ) -> crate::Result> { - match req.field_type { + match req.accessor.column_type() { ColumnType::I64 => Ok(create_collector::<{ ColumnType::I64 as u8 }>(req)), ColumnType::U64 => Ok(create_collector::<{ ColumnType::U64 as u8 }>(req)), ColumnType::F64 => Ok(create_collector::<{ ColumnType::F64 as u8 }>(req)), @@ -230,7 +235,13 @@ pub(crate) fn build_segment_stats_collector( #[derive(Clone, Debug)] pub(crate) struct SegmentStatsCollector { pub(crate) missing_u64: Option, - pub(crate) accessor: Column, + pub(crate) accessor: Arc, + /// The physical column backing `accessor`, if any, resolved once at construction. + /// + /// `collect` is called once per bucket, often with a single doc (e.g. under a + /// high-cardinality terms agg), so a per-call virtual dispatch through `accessor` is + /// measurable. + pub(crate) column_opt: Option<(Column, ColumnType)>, pub(crate) is_number_or_date_type: bool, pub(crate) buckets: Vec, pub(crate) name: String, @@ -290,21 +301,31 @@ impl SegmentAggregationCollector // skips the block accessor's buffers entirely. // Only valid without a missing value: `values_for_doc` yields nothing for a doc without a // value, so the substitute would be silently dropped. + // Also only valid over a materialized column: `values_for_doc` is per-document random + // access, which a computed source cannot offer. Those fall through to the block path + // below, which is semantically identical. // TODO: remove once we fetch all values for all bucket ids in one go - if docs.len() == 1 && self.missing_u64.is_none() { - collect_stats::( - &mut self.buckets[parent_bucket_id as usize], - self.accessor.values_for_doc(docs[0]), - self.is_number_or_date_type, - )?; - - return Ok(()); + if let Some(column_source) = &self.column_opt { + if docs.len() == 1 && self.missing_u64.is_none() { + collect_stats::( + &mut self.buckets[parent_bucket_id as usize], + column_source.0.values_for_doc(docs[0]), + self.is_number_or_date_type, + )?; + return Ok(()); + } + agg_data.column_block_accessor.fetch_block_with_missing( + docs, + column_source, + self.missing_u64, + ); + } else { + agg_data.column_block_accessor.fetch_block_with_missing( + docs, + &*self.accessor, + self.missing_u64, + ); } - agg_data.column_block_accessor.fetch_block_with_missing( - docs, - &self.accessor, - self.missing_u64, - ); collect_stats::( &mut self.buckets[parent_bucket_id as usize], agg_data.column_block_accessor.iter_vals(), diff --git a/src/aggregation/mod.rs b/src/aggregation/mod.rs index 39da6f45e..81b03551c 100644 --- a/src/aggregation/mod.rs +++ b/src/aggregation/mod.rs @@ -132,7 +132,6 @@ mod agg_data; mod agg_limits; pub mod agg_req; pub mod agg_result; -mod block_accessor; pub mod bucket; pub(crate) mod buffered_sub_aggs; mod collector; @@ -142,10 +141,13 @@ pub mod intermediate_agg_result; pub mod metric; mod segment_agg_result; +mod value_source; use std::cmp::Ordering; use std::fmt::Display; +use std::sync::Arc; -pub(crate) use block_accessor::ColumnBlockAccessor; +pub(crate) use value_source::ColumnBlockAccessor; +pub use value_source::{ValueSource, ValueSourceProvider, ValueSourceRegistry}; #[cfg(test)] mod agg_tests; @@ -184,18 +186,31 @@ pub type BucketId = u32; /// This struct holds shared resources needed during aggregation execution: /// - `limits`: Memory and bucket limits for the aggregation /// - `tokenizers`: TokenizerManager for parsing query strings in filter aggregations +/// - `value_sources`: Named computed columns that aggregations may read instead of a fast field #[derive(Clone, Default)] pub struct AggContextParams { /// Aggregation limits (memory and bucket count) pub limits: AggregationLimitsGuard, /// Tokenizer manager for query string parsing pub tokenizers: TokenizerManager, + /// Computed columns registered by name, resolved in preference to a fast field. + pub value_sources: Arc, } impl AggContextParams { /// Create new aggregation context parameters pub fn new(limits: AggregationLimitsGuard, tokenizers: TokenizerManager) -> Self { - Self { limits, tokenizers } + Self { + limits, + tokenizers, + value_sources: Arc::new(ValueSourceRegistry::default()), + } + } + + /// Attaches named computed columns, which aggregation requests address as field names. + pub fn with_value_sources(mut self, value_sources: Arc) -> Self { + self.value_sources = value_sources; + self } } diff --git a/src/aggregation/block_accessor.rs b/src/aggregation/value_source/block_accessor.rs similarity index 74% rename from src/aggregation/block_accessor.rs rename to src/aggregation/value_source/block_accessor.rs index 2eb592f1b..ccc6b9951 100644 --- a/src/aggregation/block_accessor.rs +++ b/src/aggregation/value_source/block_accessor.rs @@ -1,26 +1,11 @@ use std::cmp::Ordering; -use columnar::{Cardinality, Column, RowId}; +use columnar::{Cardinality, ColumnValues, RowId}; +use crate::aggregation::value_source::ValueSource; use crate::DocId; -/// A source of values for a block of documents. -/// -/// Implementations replace the contents of `values` and, for non-full sources, `docids`. The -/// returned cardinality describes how the two buffers are aligned. Full sources must return one -/// value per input document in the same order. Optional and multivalued sources must populate -/// `docids` with one document id per value. -pub(crate) trait BlockValueSource { - fn load_block( - &self, - docs: &[DocId], - values: &mut Vec, - docids: &mut Vec, - row_ids: &mut Vec, - ) -> Cardinality; -} - -/// Buffers the values associated with a block of documents loaded from a [`BlockValueSource`]. +/// Buffers the values associated with a block of documents loaded from a [`ValueSource`]. /// /// Regardless of their original types, values are loaded in their `u64` representation using the /// associated monotonic mapping. @@ -29,6 +14,7 @@ pub(crate) struct ColumnBlockAccessor { /// Values loaded for the latest document block, in monotonic `u64` representation. val_cache: Vec, /// Document ID corresponding to each value in `val_cache` for a non-full source. + /// For full sources, this is likely to be empty. /// /// A document can occur more than once for a multivalued source. For a full source this buffer /// is ignored because `val_cache` is aligned directly with the requested document block. @@ -37,63 +23,56 @@ pub(crate) struct ColumnBlockAccessor { missing_docids_cache: Vec, /// Scratch buffer available to sources for translating document IDs into value row IDs. row_id_cache: Vec, - /// Cardinality reported by the source that loaded the latest block. - /// For the moment this is reporting the cardinality of the full column, not - /// something specific to the block. + /// Cardinality here is describes the relationship with the loaded doc_id_cache and val_cache. + /// + /// Cheaply hints the cardinality of the given block. + /// + /// It is to be read as a "lower-bound" hint. + /// For instance, a block with one value per doc could have a cardinality property + /// set to full, optional or multivalued (both are technically true). For physical column + /// for instance, we just set cardinality to the column cardinality (although individual + /// blocks could have a stricter cardinality). + /// + /// See also [`Self::has_one_value_per_doc`] if you need a stricter notion of + /// cardinality. cardinality: Cardinality, } -impl BlockValueSource for Column { - #[inline] - fn load_block( - &self, - docs: &[DocId], - values: &mut Vec, - docids: &mut Vec, - row_ids: &mut Vec, - ) -> Cardinality { - let cardinality = self.index.get_cardinality(); - if cardinality.is_full() { - load_full_column_values(docs, self, values); - } else { - docids.clear(); - row_ids.clear(); - self.row_ids_for_docs(docs, docids, row_ids); - values.resize(row_ids.len(), 0u64); - self.values.get_vals(row_ids, values); - } - cardinality - } -} - impl ColumnBlockAccessor { #[inline] - pub(crate) fn fetch_block(&mut self, docs: &[DocId], source: &impl BlockValueSource) { + pub(crate) fn fetch_block(&mut self, docs: &[DocId], source: &S) { self.cardinality = source.load_block( docs, &mut self.val_cache, &mut self.docid_cache, &mut self.row_id_cache, ); + debug_assert!( + !self.cardinality.is_full() || self.val_cache.len() == docs.len(), + "a Full source must return exactly one value per input doc" + ); } - /// Fetches a physical column known to be full without querying its cardinality. + /// Fetches a block from a column known to be full (hence we pass the ColumnValue Object + /// directly). /// - /// This direct-column-only entry point is reserved for specialized collectors whose - /// construction already proved the column is full. + /// docs needs to be strictly increasing. #[inline] - pub(crate) fn fetch_full_column_block(&mut self, docs: &[DocId], accessor: &Column) { - debug_assert!(accessor.index.get_cardinality().is_full()); - load_full_column_values(docs, accessor, &mut self.val_cache); + pub(crate) fn fetch_full_column_block( + &mut self, + docs: &[DocId], + column_values: &dyn ColumnValues, + ) { + super::load_full_column_values(docs, column_values, &mut self.val_cache); self.cardinality = Cardinality::Full; } /// Fetches a block and appends `missing_opt` for documents without a value. #[inline] - pub(crate) fn fetch_block_with_missing( + pub(crate) fn fetch_block_with_missing( &mut self, docs: &[DocId], - source: &impl BlockValueSource, + source: &S, missing_opt: Option, ) { self.fetch_block_with_missing_ordered(docs, source, missing_opt, false) @@ -103,10 +82,10 @@ impl ColumnBlockAccessor { /// true, the missing entries are inserted in document order instead of appended as a second /// run. #[inline] - pub(crate) fn fetch_block_with_missing_ordered( + pub(crate) fn fetch_block_with_missing_ordered( &mut self, docs: &[DocId], - source: &impl BlockValueSource, + source: &S, missing_opt: Option, ordered: bool, ) { @@ -176,7 +155,7 @@ impl ColumnBlockAccessor { pub(crate) fn fetch_block_with_missing_unique_per_doc( &mut self, docs: &[DocId], - source: &impl BlockValueSource, + source: &dyn ValueSource, missing: Option, ordered: bool, ) { @@ -288,37 +267,6 @@ impl ColumnBlockAccessor { } } -#[inline] -fn load_full_column_values(docs: &[DocId], accessor: &Column, values: &mut Vec) { - // Skip the resize when already the right length (common case: fixed-size blocks). - if values.len() != docs.len() { - values.resize(docs.len(), 0u64); - } - // When the docs form a contiguous ascending run we can fetch the values as a single range. - // This lets codecs (e.g. bitpacked) bulk-decode the slice instead of gathering value-by-value. - if is_contiguous(docs) { - accessor.values.get_range(docs[0] as u64, values); - } else { - accessor.values.get_vals(docs, values); - } -} - -/// Returns true if `docs` is a contiguous ascending run `[d, d + 1, ..., d + n - 1]`. -/// -/// Assumes `docs` is sorted ascending and free of duplicates (the invariant for the -/// doc blocks passed to `fetch_block`), so comparing the endpoints is sufficient. -#[inline] -fn is_contiguous(docs: &[u32]) -> bool { - let (Some(&first), Some(&last)) = (docs.first(), docs.last()) else { - return false; - }; - debug_assert!( - docs.windows(2).all(|w| w[0] < w[1]), - "fetch_block requires docs sorted ascending without duplicates" - ); - (last - first) as usize + 1 == docs.len() -} - /// Given two sorted lists of docids `docs` and `hits`, hits is a subset of `docs`. /// Write in the output Vec all of the docs that are not in `hits`. /// @@ -356,14 +304,23 @@ fn find_missing_docs(docs: &[u32], hits: &[u32], output: &mut Vec) { #[cfg(test)] #[allow(clippy::field_reassign_with_default)] mod tests { + use std::sync::Arc; + + use columnar::{Column, ColumnType}; + use super::*; + #[derive(Debug)] struct TestValueSource { cardinality: Cardinality, entries: Vec<(DocId, u64)>, } - impl BlockValueSource for TestValueSource { + impl ValueSource for TestValueSource { + fn column_type(&self) -> ColumnType { + ColumnType::U64 + } + fn load_block( &self, docs: &[DocId], @@ -385,6 +342,75 @@ mod tests { } } + #[test] + fn test_fetch_block_accepts_trait_object() { + let docs = [2, 4, 8]; + let source = TestValueSource { + cardinality: Cardinality::Full, + entries: vec![(2, 20), (4, 40), (8, 80)], + }; + let dyn_source: &dyn ValueSource = &source; + let mut accessor = ColumnBlockAccessor::default(); + + accessor.fetch_block(&docs, dyn_source); + + assert!(accessor.has_one_value_per_doc(&docs)); + assert_eq!( + accessor.iter_docid_vals(&docs).collect::>(), + [(2, 20), (4, 40), (8, 80)] + ); + } + + fn full_column(vals: &[u64]) -> Column { + use columnar::column_index::ColumnIndex; + use columnar::column_values::{ + serialize_and_load_u64_based_column_values, ALL_U64_CODEC_TYPES, + }; + Column { + index: ColumnIndex::Full, + values: serialize_and_load_u64_based_column_values::(&vals, &ALL_U64_CODEC_TYPES), + } + } + + #[test] + fn test_as_column_distinguishes_the_two_kinds() { + let column: Arc = Arc::new((full_column(&[5, 6, 7]), ColumnType::U64)); + assert!(column.as_column().is_some()); + assert_eq!(column.bounds(), Some((5, 7))); + + let computed: Arc = Arc::new(TestValueSource { + cardinality: Cardinality::Full, + entries: vec![(0, 1)], + }); + assert!(computed.as_column().is_none()); + // No global view of a computed source, so no bounds and no bounds-driven fast paths. + assert_eq!(computed.bounds(), None); + } + + #[test] + fn test_fetch_source_block_with_missing_on_a_computed_source() { + // A computed source reports what it produced: docs 0 and 2 have no value, so they are + // absent from `docids` and the source is `Optional`. + let docs = [0, 1, 2, 3]; + let computed: Arc = Arc::new(TestValueSource { + cardinality: Cardinality::Optional, + entries: vec![(1, 11), (3, 33)], + }); + let mut accessor = ColumnBlockAccessor::default(); + + accessor.fetch_block(&docs, &*computed); + assert!(!accessor.has_one_value_per_doc(&docs)); + assert_eq!( + accessor.iter_docid_vals(&docs).collect::>(), + [(1, 11), (3, 33)] + ); + + accessor.fetch_block_with_missing(&docs, &*computed, Some(99)); + let mut pairs = accessor.iter_docid_vals(&docs).collect::>(); + pairs.sort_unstable(); + assert_eq!(pairs, [(0, 99), (1, 11), (2, 99), (3, 33)]); + } + #[test] fn test_find_missing_docs() { let docs: Vec = vec![1, 2, 3, 4, 5, 6, 7, 8, 9, 10]; @@ -392,30 +418,25 @@ mod tests { let mut missing_docs: Vec = Vec::new(); find_missing_docs(&docs, &hits, &mut missing_docs); - assert_eq!(missing_docs, vec![1, 3, 5, 7, 9]); + assert_eq!(missing_docs, [1, 3, 5, 7, 9]); } #[test] fn test_find_missing_docs_empty() { let docs: Vec = Vec::new(); let hits: Vec = vec![2, 4, 6, 8, 10]; - let mut missing_docs: Vec = Vec::new(); - find_missing_docs(&docs, &hits, &mut missing_docs); - - assert_eq!(missing_docs, Vec::::new()); + assert_eq!(missing_docs, [0u32; 0]); } #[test] fn test_find_missing_docs_all_missing() { - let docs: Vec = vec![1, 2, 3, 4, 5]; - let hits: Vec = Vec::new(); - + let docs: &[u32] = &[1, 2, 3, 4, 5]; + let hits: &[u32] = &[]; let mut missing_docs: Vec = vec![10]; - find_missing_docs(&docs, &hits, &mut missing_docs); - - assert_eq!(missing_docs, vec![1, 2, 3, 4, 5]); + find_missing_docs(docs, hits, &mut missing_docs); + assert_eq!(&missing_docs, &[1u32, 2, 3, 4, 5]); } #[test] @@ -426,13 +447,11 @@ mod tests { entries: vec![(2, 20), (4, 40), (8, 80)], }; let mut accessor = ColumnBlockAccessor::default(); - accessor.fetch_block(&docs, &source); - assert!(accessor.has_one_value_per_doc(&docs)); assert_eq!( accessor.iter_docid_vals(&docs).collect::>(), - vec![(2, 20), (4, 40), (8, 80)] + [(2, 20), (4, 40), (8, 80)] ); } @@ -450,7 +469,7 @@ mod tests { assert!(accessor.has_one_value_per_doc(&docs)); assert_eq!( accessor.iter_docid_vals(&docs).collect::>(), - vec![(0, 99), (1, 10), (2, 99), (4, 40)] + [(0, 99), (1, 10), (2, 99), (4, 40)] ); } @@ -468,7 +487,7 @@ mod tests { assert!(!accessor.has_one_value_per_doc(&docs)); assert_eq!( accessor.iter_docid_vals(&docs).collect::>(), - vec![(0, 1), (0, 3), (1, 5)] + [(0, 1), (0, 3), (1, 5)] ); } @@ -489,15 +508,20 @@ mod tests { let docs = [0, 1, 2, 4, 7, 8]; let mut accessor = ColumnBlockAccessor::default(); - accessor.fetch_block_with_missing_ordered(&docs, &column, Some(99), true); + accessor.fetch_block_with_missing_ordered( + &docs, + &(&column, ColumnType::U64), + Some(99), + true, + ); assert_eq!( accessor.iter_vals().collect::>(), - vec![99, 10, 99, 40, 70, 99] + [99, 10, 99, 40, 70, 99] ); assert_eq!( accessor.iter_docid_vals(&docs).collect::>(), - vec![(0, 99), (1, 10), (2, 99), (4, 40), (7, 70), (8, 99)] + [(0, 99), (1, 10), (2, 99), (4, 40), (7, 70), (8, 99)] ); } @@ -507,8 +531,8 @@ mod tests { accessor.docid_cache = vec![0, 0, 2, 3]; accessor.val_cache = vec![10, 10, 10, 10]; accessor.dedup_docid_val_pairs(); - assert_eq!(accessor.docid_cache, vec![0, 2, 3]); - assert_eq!(accessor.val_cache, vec![10, 10, 10]); + assert_eq!(accessor.docid_cache, [0, 2, 3]); + assert_eq!(accessor.val_cache, [10, 10, 10]); } #[test] @@ -518,8 +542,8 @@ mod tests { accessor.docid_cache = vec![0, 0, 0]; accessor.val_cache = vec![1, 2, 1]; accessor.dedup_docid_val_pairs(); - assert_eq!(accessor.docid_cache, vec![0, 0]); - assert_eq!(accessor.val_cache, vec![1, 2]); + assert_eq!(accessor.docid_cache, [0, 0]); + assert_eq!(accessor.val_cache, [1, 2]); } #[test] @@ -529,8 +553,8 @@ mod tests { accessor.docid_cache = vec![0, 0, 0, 1, 1]; accessor.val_cache = vec![3, 1, 3, 5, 5]; accessor.dedup_docid_val_pairs(); - assert_eq!(accessor.docid_cache, vec![0, 0, 1]); - assert_eq!(accessor.val_cache, vec![1, 3, 5]); + assert_eq!(accessor.docid_cache, [0, 0, 1]); + assert_eq!(accessor.val_cache, [1, 3, 5]); } #[test] @@ -539,8 +563,8 @@ mod tests { accessor.docid_cache = vec![0, 0, 1]; accessor.val_cache = vec![1, 2, 3]; accessor.dedup_docid_val_pairs(); - assert_eq!(accessor.docid_cache, vec![0, 0, 1]); - assert_eq!(accessor.val_cache, vec![1, 2, 3]); + assert_eq!(accessor.docid_cache, [0, 0, 1]); + assert_eq!(accessor.val_cache, [1, 2, 3]); } #[test] @@ -549,18 +573,8 @@ mod tests { accessor.docid_cache = vec![0]; accessor.val_cache = vec![1]; accessor.dedup_docid_val_pairs(); - assert_eq!(accessor.docid_cache, vec![0]); - assert_eq!(accessor.val_cache, vec![1]); - } - - #[test] - fn test_is_contiguous() { - assert!(!is_contiguous(&[])); - assert!(is_contiguous(&[5])); - assert!(is_contiguous(&[5, 6, 7, 8])); - assert!(is_contiguous(&[0, 1, 2])); - assert!(!is_contiguous(&[5, 7, 8])); - assert!(!is_contiguous(&[0, 1, 3])); + assert_eq!(accessor.docid_cache, [0]); + assert_eq!(accessor.val_cache, [1]); } #[test] @@ -579,7 +593,7 @@ mod tests { }; let check = |accessor: &mut ColumnBlockAccessor, docs: &[u32]| { - accessor.fetch_block(docs, &column); + accessor.fetch_block(docs, &(&column, ColumnType::U64)); let got: Vec<(u32, u64)> = accessor.iter_docid_vals(docs).collect(); let expected: Vec<(u32, u64)> = docs.iter().map(|&d| (d, vals[d as usize])).collect(); assert_eq!(got, expected); diff --git a/src/aggregation/value_source/mod.rs b/src/aggregation/value_source/mod.rs new file mode 100644 index 000000000..12fec6a41 --- /dev/null +++ b/src/aggregation/value_source/mod.rs @@ -0,0 +1,126 @@ +mod block_accessor; +mod value_source_registry; + +#[cfg(test)] +pub(crate) mod tests; + +use std::borrow::Borrow; + +pub(crate) use block_accessor::ColumnBlockAccessor; +use columnar::{Cardinality, Column, ColumnType, ColumnValues, RowId}; +pub use value_source_registry::{ValueSourceProvider, ValueSourceRegistry}; + +use crate::DocId; + +/// A source of values for a block of documents. +pub trait ValueSource: std::fmt::Debug { + /// Logical type of the encoded values returned by this source. + /// + /// Numeric values use the corresponding monotonic `u64` mapping; string/bytes + /// values are dictionary ordinals and IP addresses use the compact space ord. + fn column_type(&self) -> ColumnType; + + /// Loads the values for `docs` into `values`. + /// + /// Precondition: `docs` has to be strictly increasing. + /// + /// The output buffers are reused across blocks: on entry, `values` and `docids` + /// hold stale data from a previous call. Implementations must clear their + /// content (not append to it). + /// + /// On return, depending on the returned `Cardinality`: + /// - `Full`: `values.len() == docs.len()` and `values[i]` is the value of `docs[i]`. `docids` + /// is left unspecified and must not be read by the caller. + /// - `Optional` / `Multivalued`: `docids.len() == values.len()` and `values[i]` is a value of + /// `docids[i]`. `docids` only contains docs from `docs`. A doc is repeated once per value. + /// + /// `row_ids` is scratch the implementation may use freely. + fn load_block( + &self, + docs: &[DocId], + values: &mut Vec, + docids: &mut Vec, + row_ids: &mut Vec, + ) -> Cardinality; + + /// Returns the physical column, if this source is backed by one. + fn as_column(&self) -> Option<&Column> { + None + } + + /// Global value bounds, for fast paths that need to size or clamp something up front. + fn bounds(&self) -> Option<(u64, u64)> { + let column = self.as_column()?; + Some((column.min_value(), column.max_value())) + } +} + +// Lenient columns have erased their logical type; the tuple retains it alongside the values. +impl> + std::fmt::Debug> ValueSource for (ColumnRef, ColumnType) { + #[inline] + fn column_type(&self) -> ColumnType { + self.1 + } + + #[inline] + fn load_block( + &self, + docs: &[DocId], + values: &mut Vec, + docids: &mut Vec, + row_ids: &mut Vec, + ) -> Cardinality { + let column = self.0.borrow(); + let cardinality = column.index.get_cardinality(); + if cardinality.is_full() { + load_full_column_values(docs, &*column.values, values); + } else { + docids.clear(); + row_ids.clear(); + column.row_ids_for_docs(docs, docids, row_ids); + values.resize(row_ids.len(), 0u64); + column.values.get_vals(row_ids, values); + } + cardinality + } + + #[inline] + fn as_column(&self) -> Option<&Column> { + Some(self.0.borrow()) + } +} + +/// `docs` has to be sorted ascending and free of duplicates. +#[inline] +fn load_full_column_values( + docs: &[DocId], + column_values: &dyn ColumnValues, + values: &mut Vec, +) { + // Skip the resize when already the right length (common case: fixed-size blocks). + if values.len() != docs.len() { + values.resize(docs.len(), 0u64); + } + // When the docs form a contiguous ascending run we can fetch the values as a single range. + // This lets codecs (e.g. bitpacked) bulk-decode the slice instead of gathering value-by-value. + if is_contiguous(docs) { + column_values.get_range(docs[0] as u64, values); + } else { + column_values.get_vals(docs, values); + } +} + +/// Returns true if `docs` is a contiguous ascending run `[d, d + 1, ..., d + n - 1]`. +/// +/// `docs` has to be sorted ascending and free of duplicates. +#[inline] +fn is_contiguous(docs: &[u32]) -> bool { + let (Some(&first), Some(&last)) = (docs.first(), docs.last()) else { + return false; + }; + debug_assert!( + docs.windows(2).all(|w| w[0] < w[1]), + "fetch_block requires docs sorted ascending without duplicates" + ); + (last - first) as usize + 1 == docs.len() +} diff --git a/src/aggregation/value_source/tests.rs b/src/aggregation/value_source/tests.rs new file mode 100644 index 000000000..b2a9c3b4f --- /dev/null +++ b/src/aggregation/value_source/tests.rs @@ -0,0 +1,117 @@ +use std::sync::Arc; + +use columnar::ColumnType; + +use super::*; +use crate::SegmentReader; + +#[derive(Debug)] +pub(crate) struct Constant(u64); + +impl ValueSource for Constant { + fn column_type(&self) -> ColumnType { + ColumnType::U64 + } + + fn load_block( + &self, + docs: &[DocId], + values: &mut Vec, + _docids: &mut Vec, + _row_ids: &mut Vec, + ) -> Cardinality { + values.clear(); + values.resize(docs.len(), self.0); + Cardinality::Full + } +} + +pub(crate) struct ConstantProvider(pub u64); + +impl ValueSourceProvider for ConstantProvider { + fn for_segment(&self, _reader: &SegmentReader) -> crate::Result> { + Ok(Arc::new(Constant(self.0))) + } +} +fn index_with_scores(scores: &[u64]) -> crate::Index { + use crate::schema::{Schema, FAST}; + let mut builder = Schema::builder(); + let score = builder.add_u64_field("score", FAST); + let index = crate::Index::create_in_ram(builder.build()); + let mut writer = index.writer_for_tests().unwrap(); + for &value in scores { + writer.add_document(crate::doc!(score => value)).unwrap(); + } + writer.commit().unwrap(); + index +} + +fn run_agg(index: &crate::Index, aggs: serde_json::Value) -> serde_json::Value { + let mut registry = ValueSourceRegistry::default(); + registry.register("computed", Arc::new(ConstantProvider(1u64))); + run_agg_with_registry(index, aggs, registry) +} + +fn run_agg_with_registry( + index: &crate::Index, + aggs: serde_json::Value, + registry: ValueSourceRegistry, +) -> serde_json::Value { + use crate::aggregation::agg_req::Aggregations; + use crate::aggregation::{AggContextParams, AggregationCollector}; + use crate::query::AllQuery; + + let context = AggContextParams::default().with_value_sources(Arc::new(registry)); + let aggs: Aggregations = serde_json::from_value(aggs).unwrap(); + let collector = AggregationCollector::from_aggs(aggs, context); + let searcher = index.reader().unwrap().searcher(); + let result = searcher.search(&AllQuery, &collector).unwrap(); + serde_json::to_value(result).unwrap() +} + +#[test] +fn test_metric_over_registered_source() { + let index = index_with_scores(&[10, 20, 30, 40]); + let result = run_agg( + &index, + serde_json::json!({ "s": { "stats": { "field": "computed" } } }), + ); + // Every document contributes exactly 1. + assert_eq!(result["s"]["count"], 4); + assert_eq!(result["s"]["sum"], 4.0); + assert_eq!(result["s"]["avg"], 1.0); + assert_eq!(result["s"]["min"], 1.0); + assert_eq!(result["s"]["max"], 1.0); +} + +#[test] +fn test_registered_source_as_sub_aggregation_of_terms() { + let index = index_with_scores(&[7, 7, 7, 9]); + let result = run_agg( + &index, + serde_json::json!({ + "by_score": { + "terms": { "field": "score" }, + "aggs": { "s": { "sum": { "field": "computed" } } } + } + }), + ); + let buckets = result["by_score"]["buckets"].as_array().unwrap(); + assert_eq!(buckets.len(), 2); + assert_eq!(buckets[0]["key"], 7.0); + assert_eq!(buckets[0]["doc_count"], 3); + assert_eq!(buckets[0]["s"]["value"], 3.0); + assert_eq!(buckets[1]["key"], 9.0); + assert_eq!(buckets[1]["doc_count"], 1); + assert_eq!(buckets[1]["s"]["value"], 1.0); +} + +#[test] +fn test_is_contiguous() { + assert!(!is_contiguous(&[])); + assert!(is_contiguous(&[5])); + assert!(is_contiguous(&[5, 6, 7, 8])); + assert!(is_contiguous(&[0, 1, 2])); + assert!(!is_contiguous(&[5, 7, 8])); + assert!(!is_contiguous(&[0, 1, 3])); +} diff --git a/src/aggregation/value_source/value_source_registry.rs b/src/aggregation/value_source/value_source_registry.rs new file mode 100644 index 000000000..f9888a4e6 --- /dev/null +++ b/src/aggregation/value_source/value_source_registry.rs @@ -0,0 +1,72 @@ +//! Registration of named, computed value sources. + +use std::collections::HashMap; +use std::sync::Arc; + +use super::ValueSource; +use crate::SegmentReader; + +/// Creates a value source for each segment. +pub trait ValueSourceProvider: Send + Sync + 'static { + /// Binds this definition to a single segment. + fn for_segment(&self, reader: &SegmentReader) -> crate::Result>; +} + +/// Named computed sources available to an aggregation request. +#[derive(Clone, Default)] +pub struct ValueSourceRegistry { + providers: HashMap>, +} + +impl ValueSourceRegistry { + /// Registers `provider` under `name`, which aggregation requests then use as a field name. + /// + /// Inserting the same name several times results in an override. + pub fn register(&mut self, name: &str, provider: Arc) { + let name = name.to_string(); + self.providers.insert(name, provider); + } + + #[inline] + pub(crate) fn get(&self, name: &str) -> Option<&Arc> { + self.providers.get(name) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::aggregation::value_source::tests::ConstantProvider; + use crate::schema::Schema; + + #[test] + fn test_register_then_get() { + let mut registry = ValueSourceRegistry::default(); + registry.register("computed", Arc::new(ConstantProvider(1))); + assert!(registry.get("computed").is_some()); + assert!(registry.get("absent").is_none()); + } + + #[test] + fn test_register_overrides() { + let mut registry = ValueSourceRegistry::default(); + registry.register("computed", Arc::new(ConstantProvider(1))); + registry.register("computed", Arc::new(ConstantProvider(2))); + let index = crate::Index::create_in_ram(Schema::builder().build()); + let mut writer = index.writer_for_tests().unwrap(); + writer.add_document(crate::doc!()).unwrap(); + writer.commit().unwrap(); + let searcher = index.reader().unwrap().searcher(); + let value_source_provider = registry.get("computed").unwrap(); + let value_source = value_source_provider + .for_segment(searcher.segment_reader(0u32)) + .unwrap(); + let mut values = Vec::new(); + let mut doc_ids = Vec::new(); + let mut row_ids = Vec::new(); + let docs = &[1u32]; + value_source.load_block(docs, &mut values, &mut doc_ids, &mut row_ids); + assert!(doc_ids.is_empty()); + assert_eq!(&values, &[2u64]); + } +} diff --git a/src/index/inverted_index_plugin.rs b/src/index/inverted_index_plugin.rs index 6ca50296b..500052439 100644 --- a/src/index/inverted_index_plugin.rs +++ b/src/index/inverted_index_plugin.rs @@ -705,14 +705,14 @@ fn write_postings_merge( Ok(()) } -#[cfg(not(feature = "compare_hash_only"))] #[cfg(test)] mod tests { - use super::compute_initial_table_size; - + #[cfg(not(feature = "compare_hash_only"))] #[test] fn test_hashmap_size() { + use super::compute_initial_table_size; + assert_eq!(compute_initial_table_size(100_000).unwrap(), 1 << 12); assert_eq!(compute_initial_table_size(1_000_000).unwrap(), 1 << 15); assert_eq!(compute_initial_table_size(15_000_000).unwrap(), 1 << 19);