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/Cargo.toml b/Cargo.toml index a4ff435ae..a5b050c78 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" } @@ -84,7 +87,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" @@ -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/benches/agg_bench.rs b/benches/agg_bench.rs index 0ffd4322c..35a6d7096 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,26 +53,45 @@ fn main() { execute_agg_filtered }; runner.set_name(input_name); - bench_agg(&mut runner, &index, execute_filtered); + bench_agg(&mut runner, &searcher, 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 reader = index.reader().unwrap(); + let searcher = reader.searcher(); 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); 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()) + group.register_with_input(benchmark_name, &searcher, move |searcher| { + execute_agg(searcher, agg_req.clone()) }); } group.run(); @@ -82,7 +103,7 @@ 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), @@ -140,6 +161,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), @@ -192,8 +214,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(); @@ -528,6 +550,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 @@ -835,34 +868,32 @@ 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()); } @@ -1058,6 +1089,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()? { @@ -1110,6 +1142,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)); @@ -1184,6 +1217,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) { @@ -1195,6 +1229,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", @@ -1209,6 +1251,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 { 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/bitpacker/src/bitpacker.rs b/bitpacker/src/bitpacker.rs index 1efd2755e..dc192eeda 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,107 @@ 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" + ); + + // 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); + } + return; + } + + let output_len = output.len(); + /// # 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) + .cast::() + .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. + if output_len == 64 && self.num_bits <= 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: u64 = unsafe { load(data, bit_addr) }; + for (i, out) in chunk.iter_mut().enumerate() { + *out = (packed >> (i * self.num_bits)) & self.mask; + } + bit_addr += VALUES_PER_CHUNK * self.num_bits; + } + return; + } + 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. + if self.num_bits <= 14 || self.num_bits == 16 { + // 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; + } + } 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.as_chunks_mut::<2>().0.iter_mut().enumerate() { + // 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() { + // 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 { + // 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; + } + } + + // 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 +219,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 +415,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/Cargo.toml b/columnar/Cargo.toml index 6e07375c5..fb5e3912b 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/mod.rs b/columnar/src/column/mod.rs index 3bc61cba0..5bf7963cb 100644 --- a/columnar/src/column/mod.rs +++ b/columnar/src/column/mod.rs @@ -47,10 +47,8 @@ impl Column { impl Column { pub fn to_u64_monotonic(self) -> Column { - let values = Arc::new(monotonic_map_column( - self.values, - StrictlyMonotonicMappingToInternal::::new(), - )); + let values = + monotonic_map_column(self.values, StrictlyMonotonicMappingToInternal::::new()); Column { index: self.index, values, 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[..]); } } diff --git a/columnar/src/column_values/mod.rs b/columnar/src/column_values/mod.rs index 4911012af..87d8cca32 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/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..19083355b 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!(1, 2, 3, 4, 5, 6, 7, 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 42ef335ad..59ae64e09 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), diff --git a/jitexpr/Cargo.toml b/jitexpr/Cargo.toml index e770cb802..85033708c 100644 --- a/jitexpr/Cargo.toml +++ b/jitexpr/Cargo.toml @@ -12,5 +12,9 @@ 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" + +[dev-dependencies] +proptest = "1.7.0" 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..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; +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/ast/serde.rs b/jitexpr/src/ast/serde.rs index 705ad8d6c..959edff15 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 { @@ -579,19 +591,15 @@ 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, - ] { - let serialized = serialize(&UntypedExpr::literal(value)); + // `-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 { 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/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/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/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/jitexpr/src/types.rs b/jitexpr/src/types.rs index 18a8cb957..e767a0ce5 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,82 @@ pub enum VarType { None, } +/// 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() { + return None; + } + if val == 0.0 { + // True of both zeros; `0.0` is the representative. + Some(SafeF64(0.0)) + } 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 == other.0 + } +} + +// Sound because NaN is excluded, so equality is reflexive. +impl Eq for SafeF64 {} + +impl Ord for SafeF64 { + #[inline(always)] + fn cmp(&self, other: &SafeF64) -> Ordering { + 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 { + Some(self.cmp(other)) + } +} + +/// 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); + } +} + /// The payload of a primitive runtime value. /// /// This union is deliberately untagged. The corresponding @@ -291,7 +371,121 @@ 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 zeros -- which fold together -- 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_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] + 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)); + } #[test] fn test_runtime_value_layouts() { diff --git a/query-grammar/src/query_grammar.rs b/query-grammar/src/query_grammar.rs index aaa7800e8..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 @@ -687,12 +685,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); } @@ -1125,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/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/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 475cef2ea..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::{ @@ -22,14 +22,16 @@ 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, }; -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)] @@ -37,8 +39,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 +117,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 { @@ -286,39 +288,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]; @@ -336,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, @@ -436,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, @@ -445,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, @@ -467,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, @@ -492,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, @@ -576,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, @@ -599,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, @@ -631,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 @@ -690,7 +715,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, @@ -749,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() @@ -827,7 +861,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)?; @@ -968,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 @@ -1002,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 { @@ -1040,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; @@ -1064,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( @@ -1089,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 }; @@ -1120,12 +1170,11 @@ 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(), req: TermsAggregationInternal::from_req(req), - sug_aggregations: sub_aggs.clone(), + sub_aggregations: sub_aggs.clone(), allowed_term_ids, is_top_level, }); @@ -1144,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/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..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; @@ -22,21 +23,19 @@ 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, - /// The field type of the fast field. - pub field_type: ColumnType, + pub(crate) accessor: Arc, /// 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. @@ -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 5b3af5f08..0b2d53c23 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; @@ -158,23 +159,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 { @@ -230,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, ); @@ -252,11 +251,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)] @@ -288,18 +284,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] } } @@ -332,20 +318,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 } } @@ -735,8 +711,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( @@ -749,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, @@ -1097,6 +1073,65 @@ 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() + .map(|m| m.missing_value == value) + .unwrap_or(false) + { + 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 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| { + decoded_keys.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(), + )); + } + for (position, key) in positions.into_iter().zip(decoded_keys) { + keys[position].push(key); + } + } + 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. @@ -1332,86 +1367,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) @@ -1553,6 +1601,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( @@ -1603,7 +1708,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(()) } @@ -2706,7 +2811,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); @@ -2723,23 +2828,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/range.rs b/src/aggregation/bucket/range.rs index a66ea8225..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,23 +19,22 @@ 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; /// 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, - /// The type of the fast field. - pub field_type: ColumnType, + pub(crate) accessor: Arc, /// 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 { @@ -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 column_type: ColumnType, + pub(crate) accessor: Arc, /// 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 { @@ -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 @@ -1390,7 +1394,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; @@ -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 ))) } }; @@ -1554,7 +1566,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 } 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/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; diff --git a/src/aggregation/metric/cardinality.rs b/src/aggregation/metric/cardinality.rs deleted file mode 100644 index ca0a00fb2..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 struct CardinalityAggReqData { - /// The column accessor to access the fast field values. - pub accessor: Column, - /// The column_type of the field. - pub column_type: ColumnType, - /// The string dictionary column if the field is of type string. - pub str_dict_column: Option, - /// The missing value normalized to the internal u64 representation of the field type. - pub missing_value_for_accessor: Option, - /// The name of the aggregation. - pub name: String, - /// The aggregation request. - pub 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..8440cd948 --- /dev/null +++ b/src/aggregation/metric/cardinality/mod.rs @@ -0,0 +1,614 @@ +//! 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 std::sync::Arc; + +use columnar::{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::value_source::ValueSource; +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: 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. + 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) + .expect("LG_K is within the supported range"), + 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).expect("LG_K is within the supported range"); + 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.accessor.column_type() != ColumnType::Str { + return Ok(Box::new(SegmentNumericCardinalityCollector::from_req( + 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 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( + 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..d6b3fa7bc --- /dev/null +++ b/src/aggregation/metric/cardinality/numeric_collector.rs @@ -0,0 +1,264 @@ +//! 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::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::value_source::ValueSource; +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: Arc, + /// 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( + accessor_idx: usize, + 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::() + .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..b2ec7d218 --- /dev/null +++ b/src/aggregation/metric/cardinality/str_collector.rs @@ -0,0 +1,399 @@ +//! 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 std::sync::Arc; + +use columnar::{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::value_source::ValueSource; +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_value(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: 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. + 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_value(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_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_value("__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: Arc, + 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]); + } +} 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 05ff861f2..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,26 +43,25 @@ 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, /// 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, - /// The type of the field. - pub field_type: ColumnType, + pub(crate) is_number_or_date_type: bool, /// 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: Arc, /// 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..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 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/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 { 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/directory/composite_file.rs b/src/directory/composite_file.rs index 93e063880..f3d3d4cd6 100644 --- a/src/directory/composite_file.rs +++ b/src/directory/composite_file.rs @@ -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..9e3a7b14d 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, TerminatingWrite, 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.terminate().map_err(TryAcquireLockError::from)?; Ok(DirectoryLock::from(Box::new(DirectoryLockGuard { directory: directory.box_clone(), path: filepath.to_owned(), @@ -138,10 +139,8 @@ 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`]. + /// 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`]. /// /// However, depending on the directory implementation, /// it might be required to call [`Directory::sync_directory()`] to ensure @@ -150,15 +149,11 @@ pub trait Directory: DirectoryClone + fmt::Debug + Send + Sync + 'static { /// a POSIX filesystem.) /// /// Write operations may be aggressively buffered. - /// The client of this trait is responsible for calling flush + /// The client of this trait is responsible for calling terminate /// 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`] triggering terminate. /// /// The file may not previously exist. fn open_write(&self, path: &Path) -> Result; diff --git a/src/directory/mmap_directory/mod.rs b/src/directory/mmap_directory/mod.rs index 0370a8b9e..1c33f0358 100644 --- a/src/directory/mmap_directory/mod.rs +++ b/src/directory/mmap_directory/mod.rs @@ -319,7 +319,7 @@ impl Drop for ReleaseLockFile { } /// This Write wraps a File, but has the specificity of -/// call `sync_all` on flush. +/// calling `sync_all` on terminate. struct SafeFileWriter(File); impl SafeFileWriter { @@ -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,10 +440,6 @@ 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. @@ -559,10 +555,11 @@ 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() + .terminate() + .unwrap(); let readonlymap = mmap_directory.open_read(&path).unwrap(); assert_eq!(readonlymap.len(), 0); } @@ -578,12 +575,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.terminate().unwrap(); } let mut keep = vec![]; diff --git a/src/directory/ram_directory.rs b/src/directory/ram_directory.rs index 7ea6db9e0..2ce9d39d1 100644 --- a/src/directory/ram_directory.rs +++ b/src/directory/ram_directory.rs @@ -1,6 +1,7 @@ -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}; use std::sync::{Arc, RwLock}; use std::{fmt, result}; @@ -14,6 +15,49 @@ 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; + } + } + + fn release(&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) { + self.release(); + } +} + /// Writer associated with the [`RamDirectory`]. /// /// The Writer just writes a buffer. @@ -21,64 +65,82 @@ struct VecWriter { path: PathBuf, shared_directory: RamDirectory, data: Cursor>, - is_flushed: bool, + memory_usage: MemoryUsageTracker, + is_finished: 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, - is_flushed: true, + memory_usage, + 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 \ - 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 - ) + ); + } + if let Ok(mut fs) = self.shared_directory.fs.write() { + fs.active_writers.remove(&self.path); } } } 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(); + self.memory_usage.add(capacity - previous_capacity); Ok(buf.len()) } 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() + 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); + self.is_finished = true; + self.memory_usage.release(); + Ok(()) } } #[derive(Default)] struct InnerDirectory { fs: HashMap, + active_writers: HashSet, watch_router: WatchCallbackList, } 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 { @@ -96,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 { @@ -117,10 +179,11 @@ 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 terminating the writer. #[derive(Clone, Default)] pub struct RamDirectory { fs: Arc>, + active_writer_mem_usage: Arc, } impl RamDirectory { @@ -129,24 +192,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`]. @@ -198,16 +249,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 { - Ok(BufWriter::new(Box::new(vec_writer))) + 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()); + Ok(BufWriter::new(Box::new(vec_writer))) } fn atomic_read(&self, path: &Path) -> Result, OpenReadError> { @@ -244,7 +293,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] @@ -257,7 +307,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.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); @@ -265,21 +315,45 @@ mod tests { } #[test] - fn test_ram_directory_deep_clone() { + fn test_dropped_writer_releases_path() { 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 writer = dir.open_write(path).unwrap(); + assert!(dir.open_write(path).is_err()); + + drop(writer); + dir.open_write(path).unwrap().terminate().unwrap(); + assert!(dir.exists(path).unwrap()); + } + + #[test] + fn test_active_writer_memory_usage() { + let dir = RamDirectory::default(); + 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(); + assert_eq!(dir.total_mem_usage(), grown_capacity); + assert_eq!(dir.clone().total_mem_usage(), dir.total_mem_usage()); + + let file_len = 1 + MEMORY_USAGE_UPDATE_THRESHOLD + first_capacity + 1; + writer.terminate().unwrap(); + assert_eq!(dir.total_mem_usage(), file_len); + + dir.delete(path).unwrap(); + assert_eq!(dir.total_mem_usage(), 0); } } diff --git a/src/directory/tests.rs b/src/directory/tests.rs index a2c8473ce..dd6d5de6a 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.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); @@ -135,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)?; + let writer = directory.open_write(test_path)?; + assert!(directory.open_write(test_path).is_err()); + writer.terminate()?; assert!(directory.exists(test_path).unwrap()); assert!(directory.open_write(test_path).is_err()); assert!(directory.delete(test_path).is_ok()); @@ -144,13 +145,12 @@ 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()); + 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()); } fn test_directory_delete(directory: &dyn Directory) -> crate::Result<()> { @@ -158,7 +158,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.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/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/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/inverted_index_plugin.rs b/src/index/inverted_index_plugin.rs index 68ddf3e48..19273c926 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()?; } @@ -714,11 +707,12 @@ fn write_postings_merge( #[cfg(test)] mod tests { - use super::compute_initial_table_size; - #[test] #[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); 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/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/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/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; 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/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 ac7192580..acd5fd206 100644 --- a/src/query/doc_predicate_query/function_predicate.rs +++ b/src/query/doc_predicate_query/function_predicate.rs @@ -1,5 +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 @@ -48,8 +50,15 @@ 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> { + 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 new file mode 100644 index 000000000..9698ac3fd --- /dev/null +++ b/src/query/doc_predicate_query/jitexpr_predicate.rs @@ -0,0 +1,836 @@ +use std::collections::HashMap; +use std::io; + +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::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. +/// +/// 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. +/// +/// 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 { + /// 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(); + let required_presence = required_presence_for_true(&expression); + Ok(Self { + expression, + inferred_inputs, + required_presence, + }) + } + + /// 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); + } + + // 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() + .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(); + 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)) + } + } +} + +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 inputs_vec = ClearOnDrop::wrap(&mut self.input_values); + + fill_input_values( + &self.columns_opt, + &mut self.string_inputs, + 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, DocSetCollector}; + use crate::query::doc_predicate_query::DocPredicateQuery; + 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(); + 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 + ); + } + + /// 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() { + // 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..e8c74e106 100644 --- a/src/query/doc_predicate_query/mod.rs +++ b/src/query/doc_predicate_query/mod.rs @@ -1,13 +1,20 @@ +use std::cmp::Ordering; 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. @@ -59,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() } } @@ -141,14 +211,27 @@ 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 { + 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))) + } + } } fn scorer_danger( @@ -157,18 +240,48 @@ 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 { + 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)))) + } + } } } +/// 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), + /// 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 /// segment. /// @@ -186,19 +299,37 @@ 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`]. 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(); @@ -271,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), } 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; diff --git a/src/query/phrase_query/regex_phrase_query.rs b/src/query/phrase_query/regex_phrase_query.rs index 98e07d399..67ca03499 100644 --- a/src/query/phrase_query/regex_phrase_query.rs +++ b/src/query/phrase_query/regex_phrase_query.rs @@ -1,3 +1,9 @@ +use std::fmt; +use std::sync::Arc; + +use once_cell::sync::OnceCell; +use tantivy_fst::Regex; + use super::regex_phrase_weight::RegexPhraseWeight; use crate::query::bm25::Bm25Weight; use crate::query::{EnableScoring, Query, Weight}; @@ -25,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. @@ -71,6 +92,37 @@ impl RegexPhraseQuery { phrase_terms: terms, slop, max_expansions: 1 << 14, + regexes: CompiledRegexes::default(), + } + } + + /// 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), } } @@ -111,6 +163,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 @@ -149,9 +219,15 @@ impl RegexPhraseQuery { } => Some(Bm25Weight::for_terms(statistics_provider, &terms)?), EnableScoring::Disabled { .. } => None, }; + let phrase_terms = self + .phrase_terms + .iter() + .map(|(offset, _)| *offset) + .zip(self.regexes()?.iter().cloned()) + .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..bbaddbc79 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,71 @@ 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_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_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"])?; diff --git a/src/query/query_parser/query_parser.rs b/src/query/query_parser/query_parser.rs index 73416ee2d..4d0185fdb 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,31 @@ 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}"); + prop_assert!(query_parser.parse_query(&query).is_err()); + let (_, errors) = query_parser.parse_query_lenient(&query); + prop_assert!(!errors.is_empty()); + } + + #[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, 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()); } 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); diff --git a/sstable/src/delta.rs b/sstable/src/delta.rs index 97d868e4e..986700fef 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,62 @@ 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 { + 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) + } +} + pub struct DeltaWriter where W: io::Write { @@ -241,7 +298,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 +308,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/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()) diff --git a/sstable/src/streamer.rs b/sstable/src/streamer.rs index 9203b3d0a..9e5d11541 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, @@ -124,14 +142,23 @@ 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), + lower_bound_reached: self.lower == Bound::Unbounded, lower_bound: self.lower, upper_bound: self.upper, + upper_bound_comparator: DeltaKeyComparator::new(), _lifetime: std::marker::PhantomData, }) } @@ -174,8 +201,11 @@ 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, + always_match_at: Option, } impl Streamer<'_, TSSTable, AlwaysMatch> @@ -184,12 +214,15 @@ where TSSTable: SSTable pub fn empty() -> Self { Streamer { automaton: AlwaysMatch, - states: Vec::new(), + states: vec![AlwaysMatch.start()], + always_match_at: Some(0), delta_reader: DeltaReader::empty(), key: Vec::new(), term_ord: None, + lower_bound_reached: true, lower_bound: Bound::Unbounded, upper_bound: Bound::Unbounded, + upper_bound_comparator: DeltaKeyComparator::new(), _lifetime: std::marker::PhantomData, } } @@ -201,53 +234,173 @@ where A::State: Clone, TSSTable: SSTable, { + #[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); + let mut lower_bound_comparator = DeltaKeyComparator::new(); + while self.advance_delta_reader() { + let common_prefix_len = self.delta_reader.common_prefix_len(); + 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); + self.states.push(state.clone()); + } + self.lower_bound_reached = true; + return true; + } + } + 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 { - 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), - }); + if !self.lower_bound_reached { + if !self.initialize() { + // no key higher than lower-bound at all + return false; + } + 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; + } + } + + 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_always_match::(), + (true, false) => self.advance_always_match::(), + (false, true) => self.advance_with_automaton::(), + (false, false) => self.advance_with_automaton::(), + } + } + + fn advance_always_match(&mut self) -> bool { + if !self.advance_delta_reader() { + return false; + } + self.reconstruct_key_and_check_upper_bound::() + } + + 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; + } + 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); - self.key.truncate(common_prefix_len); + // 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()); } - 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; + 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); } - // We match the lower key once. All subsequent keys will pass that bar. - self.lower_bound = Bound::Unbounded; - 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 { + + if !self.reconstruct_key_and_check_upper_bound::() { return false; } - if self.automaton.is_match(&state) { + 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.