Merge remote-tracking branch 'origin/main' into mallets/seqnum-field

This commit is contained in:
Luca Cominardi
2026-10-01 11:40:33 +02:00
101 changed files with 7522 additions and 2428 deletions
+2 -2
View File
@@ -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}}
+5 -1
View File
@@ -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.
+66 -23
View File
@@ -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<std::alloc::System> = &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<BenchmarkConfig>);
@@ -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 {
+106 -16
View File
@@ -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<std::alloc::System> = &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);
}
}
+130 -3
View File
@@ -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::<u64>()
.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::<VALUES_PER_CHUNK>();
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::<VALUES_PER_CHUNK>();
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<u64> = (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! {
+4
View File
@@ -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"]
@@ -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<u8> = 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<Column> = 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();
}
+2 -4
View File
@@ -47,10 +47,8 @@ impl<T: PartialOrd + Default> Column<T> {
impl<T: MonotonicallyMappableToU64> Column<T> {
pub fn to_u64_monotonic(self) -> Column<u64> {
let values = Arc::new(monotonic_map_column(
self.values,
StrictlyMonotonicMappingToInternal::<T>::new(),
));
let values =
monotonic_map_column(self.values, StrictlyMonotonicMappingToInternal::<T>::new());
Column {
index: self.index,
values,
@@ -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[..]);
}
}
+5
View File
@@ -203,6 +203,11 @@ impl<T: Copy + PartialOrd + Debug + 'static> ColumnValues<T> for Arc<dyn ColumnV
self.as_ref().get_val(idx)
}
#[inline(always)]
fn get_vals(&self, indexes: &[u32], output: &mut [T]) {
self.as_ref().get_vals(indexes, output)
}
#[inline(always)]
fn get_vals_opt(&self, indexes: &[u32], output: &mut [Option<T>]) {
self.as_ref().get_vals_opt(indexes, output)
+42 -3
View File
@@ -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<C, T, Input> {
pub fn monotonic_map_column<C, T, Input, Output>(
from_column: C,
monotonic_mapping: T,
) -> impl ColumnValues<Output>
) -> Arc<dyn ColumnValues<Output>>
where
C: ColumnValues<Input> + 'static,
T: StrictlyMonotonicFn<Input, Output> + 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::<Input>() == TypeId::of::<Output>() {
let column: Arc<dyn ColumnValues<Input>> = Arc::new(from_column);
return (&column as &dyn Any)
.downcast_ref::<Arc<dyn ColumnValues<Output>>>()
.unwrap()
.clone();
}
Arc::new(MonotonicMappingColumn {
from_column,
monotonic_mapping,
_phantom: PhantomData,
}
})
}
impl<C, T, Input, Output> ColumnValues<Output> for MonotonicMappingColumn<C, T, Input>
@@ -104,6 +114,35 @@ mod tests {
StrictlyMonotonicMappingInverter, StrictlyMonotonicMappingToInternal,
};
#[test]
fn test_u128_identity() {
let column: Arc<dyn ColumnValues<u128>> = monotonic_map_column(
VecColumn::from(vec![u128::MAX]),
StrictlyMonotonicMappingInverter::from(
StrictlyMonotonicMappingToInternal::<u128>::new(),
),
);
assert_eq!(column.get_val(0), u128::MAX);
}
#[test]
fn test_same_type_non_identity() {
struct Shift;
impl StrictlyMonotonicFn<u64, u64> 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<u64> = (0..100u64).map(|el| el * 10).collect();
@@ -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<External, Internal> {
/// 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<T> From<T> for StrictlyMonotonicMappingInverter<T> {
impl<From, To, T> StrictlyMonotonicFn<To, From> for StrictlyMonotonicMappingInverter<T>
where T: StrictlyMonotonicFn<From, To>
{
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<External: MonotonicallyMappableToU128, T: MonotonicallyMappableToU128>
StrictlyMonotonicFn<External, u128> for StrictlyMonotonicMappingToInternal<T>
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<External: MonotonicallyMappableToU64, T: MonotonicallyMappableToU64>
StrictlyMonotonicFn<External, u64> for StrictlyMonotonicMappingToInternal<T>
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
@@ -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
}
+1 -1
View File
@@ -108,7 +108,7 @@ pub fn open_u128_mapped<T: MonotonicallyMappableToU128 + Debug>(
let reader = CompactSpaceDecompressor::open(bytes)?;
let inverted: StrictlyMonotonicMappingInverter<StrictlyMonotonicMappingToInternal<T>> =
StrictlyMonotonicMappingToInternal::<T>::new().into();
Ok(Arc::new(monotonic_map_column(reader, inverted)))
Ok(monotonic_map_column(reader, inverted))
}
/// Returns the u64 representation of the u128 data.
@@ -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<const NUM_BITS: u8 = { u8::MAX }> {
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<const NUM_BITS: u8> BitpackedReader<NUM_BITS> {
#[inline(always)]
fn unpacker(&self) -> BitUnpacker {
if NUM_BITS == u8::MAX {
self.bit_unpacker
} else {
BitUnpacker::new(NUM_BITS)
}
}
}
impl<const NUM_BITS: u8> ColumnValues for BitpackedReader<NUM_BITS> {
#[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<T: MonotonicallyMappableToU64>(
bytes: OwnedBytes,
) -> io::Result<Arc<dyn ColumnValues<T>>> {
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<u64> = (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::<u64>(data.clone()).unwrap();
let signed_reader = load::<i64>(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::<Vec<_>>()
);
}
}
}
}
#[test]
fn test_with_codec_data_sets_simple() {
create_and_validate::<BitpackedCodec>(&[4, 3, 12], "name");
+9 -5
View File
@@ -114,7 +114,7 @@ impl CodecType {
bytes: OwnedBytes,
) -> io::Result<Arc<dyn ColumnValues<T>>> {
match self {
CodecType::Bitpacked => load_specific_codec::<BitpackedCodec, T>(bytes),
CodecType::Bitpacked => bitpacked::load::<T>(bytes),
CodecType::Linear => load_specific_codec::<LinearCodec, T>(bytes),
CodecType::BlockwiseLinear => load_specific_codec::<BlockwiseLinearCodec, T>(bytes),
}
@@ -124,12 +124,16 @@ impl CodecType {
fn load_specific_codec<C: ColumnCodec, T: MonotonicallyMappableToU64>(
bytes: OwnedBytes,
) -> io::Result<Arc<dyn ColumnValues<T>>> {
let reader = C::load(bytes)?;
let reader_typed = monotonic_map_column(
Ok(map_column_values::<_, T>(C::load(bytes)?))
}
fn map_column_values<C: ColumnValues + 'static, T: MonotonicallyMappableToU64>(
reader: C,
) -> Arc<dyn ColumnValues<T>> {
monotonic_map_column(
reader,
StrictlyMonotonicMappingInverter::from(StrictlyMonotonicMappingToInternal::<T>::new()),
);
Ok(Arc::new(reader_typed))
)
}
impl CodecType {
+4 -5
View File
@@ -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),
+4
View File
@@ -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"
+3 -2
View File
@@ -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<dyn Error>> {
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)?),
],
)?;
+1 -1
View File
@@ -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),
+27 -18
View File
@@ -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<str>),
}
@@ -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<i64> for Literal {
}
}
impl From<f64> 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<f64> 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<Literal, NonFiniteFloat> {
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);
}
}
+5 -1
View File
@@ -1,11 +1,15 @@
mod infer_types;
mod literal;
mod presence;
mod serde;
mod untyped_expr;
pub use infer_types::{InferredTypeSet, TypeError, infer_types, infer_types_with_target};
pub(crate) use infer_types::{infer_type_with_variable_types, infer_types_aux};
pub use literal::Literal;
pub use literal::{Literal, NonFiniteFloat};
pub use presence::{
ConditionSet, VariablePresenceCondition, required_presence, required_presence_for_true,
};
pub(crate) use serde::format_variable_name;
pub use serde::{DeserializeError, deserialize, serialize};
pub use untyped_expr::UntypedExpr;
+902
View File
@@ -0,0 +1,902 @@
//! Necessary conditions on the presence of variables.
//!
//! Most functions return null as soon as one of their arguments is null. An expression can
//! therefore often only produce a value, or only evaluate to `true`, if some of its variables are
//! present. For instance, `(EQ (ADD a 1i64) b)` is null unless both `a` and `b` are present.
//!
//! A caller evaluating a predicate over many documents can use this to skip the documents missing
//! these variables without evaluating the expression.
use std::sync::Arc;
use crate::ast::{Function, Literal, UntypedExpr};
/// A boolean formula over the presence of variables.
///
/// It is meant to be used as a necessary condition: it is implied by some property of an
/// expression (producing a value, or evaluating to `true`), but it does not imply it.
///
/// As much as possible, we try to normalize these object, in order to simplify them.
#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub enum VariablePresenceCondition {
/// Always satisfied.
Always,
/// Never satisfied.
Never,
/// Satisfied when the variable is present.
Present(Arc<str>),
/// Satisfied when all of the conditions are satisfied.
All(ConditionSet),
/// Satisfied when at least one of the conditions is satisfied.
Any(ConditionSet),
}
/// The children of a [`VariablePresenceCondition::All`] or [`VariablePresenceCondition::Any`] node.
///
/// It can only be built through [`VariablePresenceCondition::all`] and
/// [`VariablePresenceCondition::any`], which
/// uphold the following hidden contract, on which the derived `Eq` and `Hash` rely:
/// - children are sorted and distinct, and there are at least two of them;
/// - no child is `Always`, `Never`, or a node of the same kind as the parent.
#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct ConditionSet(Vec<VariablePresenceCondition>);
impl ConditionSet {
/// Returns the children, in canonical order.
pub fn iter(&self) -> impl Iterator<Item = &VariablePresenceCondition> {
self.0.iter()
}
}
impl VariablePresenceCondition {
pub fn all(
conditions: impl IntoIterator<Item = VariablePresenceCondition>,
) -> VariablePresenceCondition {
let mut children: Vec<VariablePresenceCondition> = Vec::new();
for condition in conditions {
match condition {
VariablePresenceCondition::Always => {}
VariablePresenceCondition::Never => return VariablePresenceCondition::Never,
VariablePresenceCondition::All(grand_children) => children.extend(grand_children.0),
condition => children.push(condition),
}
}
children.sort();
children.dedup();
match children.len() {
0 => VariablePresenceCondition::Always,
1 => children.pop().unwrap(),
_ => VariablePresenceCondition::All(ConditionSet(children)),
}
}
pub fn any(
conditions: impl IntoIterator<Item = VariablePresenceCondition>,
) -> VariablePresenceCondition {
let mut children: Vec<VariablePresenceCondition> = Vec::new();
for condition in conditions {
match condition {
VariablePresenceCondition::Never => {}
VariablePresenceCondition::Always => return VariablePresenceCondition::Always,
VariablePresenceCondition::Any(grand_children) => children.extend(grand_children.0),
condition => children.push(condition),
}
}
children.sort();
children.dedup();
match children.len() {
0 => VariablePresenceCondition::Never,
1 => children.pop().unwrap(),
_ => VariablePresenceCondition::Any(ConditionSet(children)),
}
}
/// Evaluates the condition, given the presence of each variable.
#[cfg(test)]
pub fn eval(&self, is_present: &mut impl FnMut(&str) -> bool) -> bool {
match self {
VariablePresenceCondition::Always => true,
VariablePresenceCondition::Never => false,
VariablePresenceCondition::Present(variable_name) => is_present(variable_name),
VariablePresenceCondition::All(conditions) => conditions
.iter()
.all(|condition| condition.eval(&mut *is_present)),
VariablePresenceCondition::Any(conditions) => conditions
.iter()
.any(|condition| condition.eval(&mut *is_present)),
}
}
}
/// Returns a necessary presence condition for `expr` to evaluate to a non-null value.
pub fn required_presence(expr: &UntypedExpr) -> VariablePresenceCondition {
match expr {
UntypedExpr::Literal(_) => VariablePresenceCondition::Always,
UntypedExpr::Variable(variable_name) => {
VariablePresenceCondition::Present(variable_name.clone())
}
UntypedExpr::FnCall { function, args } => required_presence_for_fn_call(*function, args),
}
}
/// Returns a necessary presence condition for `expr` to evaluate to a present `true`.
pub fn required_presence_for_true(expr: &UntypedExpr) -> VariablePresenceCondition {
match expr {
UntypedExpr::Literal(Literal::Bool(true)) => VariablePresenceCondition::Always,
UntypedExpr::Literal(_) => VariablePresenceCondition::Never,
UntypedExpr::Variable(variable_name) => {
VariablePresenceCondition::Present(variable_name.clone())
}
UntypedExpr::FnCall { function, args } => {
required_presence_for_true_for_fn_call(*function, args)
}
}
}
fn required_presence_for_fn_call(
function: Function,
args: &[UntypedExpr],
) -> VariablePresenceCondition {
// null argument as "null in, null out" would make callers skip matching documents.
match function {
// A null argument makes the result null.
//
// AND belongs here: `(AND false none)` is null.
Function::Abs
| Function::Add
| Function::And
| Function::Ceil
| Function::Concat
| Function::Divide
| Function::Eq
| Function::Floor
| Function::Gt
| Function::GtEq
| Function::IntMod
| Function::Left
| Function::Lower
| Function::Lt
| Function::LtEq
| Function::Max
| Function::Min
| Function::Multiply
| Function::Pow
| Function::RegexpExtract
| Function::Right
| Function::Round
| Function::SplitAfter
| Function::SplitBefore
| Function::Sqrt
| Function::Substring
| Function::SubstringCount
| Function::Subtract
| Function::TextJoin
| Function::Trim
| Function::Upper => VariablePresenceCondition::all(args.iter().map(required_presence)),
// OR is null only if all of its arguments are null.
Function::Or => VariablePresenceCondition::any(args.iter().map(required_presence)),
// IF is null if its condition is null. Otherwise it takes the presence of the selected
// branch.
Function::If => {
let [condition, when_true, when_false] = args else {
return VariablePresenceCondition::Always;
};
VariablePresenceCondition::all([
required_presence(condition),
VariablePresenceCondition::any([
required_presence(when_true),
required_presence(when_false),
]),
])
}
// These functions always return a present value.
//
// REGEXP_LIKE returns `false` for a null input.
Function::IsNotNull
| Function::IsNull
| Function::Neq
| Function::Not
| Function::RegexpLike => VariablePresenceCondition::Always,
}
}
fn required_presence_for_true_for_fn_call(
function: Function,
args: &[UntypedExpr],
) -> VariablePresenceCondition {
match function {
Function::And => {
VariablePresenceCondition::all(args.iter().map(required_presence_for_true))
}
Function::Or => VariablePresenceCondition::any(args.iter().map(required_presence_for_true)),
Function::If => {
let [condition, when_true, when_false] = args else {
return VariablePresenceCondition::Always;
};
VariablePresenceCondition::all([
required_presence(condition),
VariablePresenceCondition::any([
required_presence_for_true(when_true),
required_presence_for_true(when_false),
]),
])
}
Function::IsNotNull => {
let [arg] = args else {
return VariablePresenceCondition::Always;
};
required_presence(arg)
}
// REGEXP_LIKE returns `false` for a null input.
Function::RegexpLike => {
let Some(input) = args.first() else {
return VariablePresenceCondition::Always;
};
required_presence(input)
}
// A `true` result is in particular a present result. This fallback is therefore correct
// for any function, including functions added later.
//
// It yields `Always` for NOT, NEQ, and IS_NULL, which are `true` when their
// argument is null.
_ => required_presence_for_fn_call(function, args),
}
}
#[cfg(test)]
mod tests {
use std::collections::HashMap;
use std::collections::hash_map::DefaultHasher;
use std::hash::{Hash, Hasher};
use proptest::prelude::*;
use proptest::strategy::BoxedStrategy;
use super::*;
use crate::ast::{InferredTypeSet, deserialize, infer_types_with_target};
use crate::compile::{StringArena, compile};
use crate::types::{VarType, VariableValue};
fn present(variable_name: &str) -> VariablePresenceCondition {
VariablePresenceCondition::Present(Arc::from(variable_name))
}
fn all(conditions: Vec<VariablePresenceCondition>) -> VariablePresenceCondition {
VariablePresenceCondition::all(conditions)
}
fn any(conditions: Vec<VariablePresenceCondition>) -> VariablePresenceCondition {
VariablePresenceCondition::any(conditions)
}
fn for_true(expr: &str) -> VariablePresenceCondition {
required_presence_for_true(&deserialize(expr).unwrap())
}
fn for_value(expr: &str) -> VariablePresenceCondition {
required_presence(&deserialize(expr).unwrap())
}
#[test]
fn test_all_simplification() {
assert_eq!(
VariablePresenceCondition::all([]),
VariablePresenceCondition::Always
);
assert_eq!(
VariablePresenceCondition::all([VariablePresenceCondition::Always, present("a")]),
present("a")
);
assert_eq!(
VariablePresenceCondition::all([present("a"), VariablePresenceCondition::Never]),
VariablePresenceCondition::Never
);
assert_eq!(
VariablePresenceCondition::all([
present("a"),
all(vec![present("b"), present("a")]),
present("c"),
]),
all(vec![present("a"), present("b"), present("c")])
);
assert_eq!(
VariablePresenceCondition::all([any(vec![present("a"), present("b")]), present("c")]),
all(vec![any(vec![present("a"), present("b")]), present("c")])
);
}
#[test]
fn test_any_simplification() {
assert_eq!(
VariablePresenceCondition::any([]),
VariablePresenceCondition::Never
);
assert_eq!(
VariablePresenceCondition::any([VariablePresenceCondition::Never, present("a")]),
present("a")
);
assert_eq!(
VariablePresenceCondition::any([present("a"), VariablePresenceCondition::Always]),
VariablePresenceCondition::Always
);
assert_eq!(
VariablePresenceCondition::any([present("a"), any(vec![present("b"), present("a")])]),
any(vec![present("a"), present("b")])
);
}
fn hash_of(condition: &VariablePresenceCondition) -> u64 {
let mut hasher = DefaultHasher::new();
condition.hash(&mut hasher);
hasher.finish()
}
fn assert_same(left: VariablePresenceCondition, right: VariablePresenceCondition) {
assert_eq!(left, right);
assert_eq!(hash_of(&left), hash_of(&right));
}
#[test]
fn test_canonical_order() {
let (a, b, c) = (present("a"), present("b"), present("c"));
assert_same(
VariablePresenceCondition::all([a.clone(), b.clone()]),
VariablePresenceCondition::all([b.clone(), a.clone()]),
);
assert_same(
VariablePresenceCondition::any([c.clone(), a.clone(), b.clone()]),
VariablePresenceCondition::any([b.clone(), c.clone(), a.clone()]),
);
assert_ne!(
VariablePresenceCondition::all([a.clone(), b.clone()]),
VariablePresenceCondition::any([a.clone(), b.clone()])
);
let VariablePresenceCondition::All(children) =
VariablePresenceCondition::all([c.clone(), a.clone()])
else {
panic!("expected an All node");
};
assert_eq!(children.iter().collect::<Vec<_>>(), vec![&a, &c]);
}
#[test]
fn test_canonical_grouping_and_repetition() {
let (a, b, c) = (present("a"), present("b"), present("c"));
assert_same(
VariablePresenceCondition::all([
a.clone(),
VariablePresenceCondition::all([b.clone(), c.clone()]),
]),
VariablePresenceCondition::all([
VariablePresenceCondition::all([c.clone(), a.clone()]),
b.clone(),
]),
);
assert_same(
VariablePresenceCondition::all([a.clone(), a.clone()]),
a.clone(),
);
assert_same(
VariablePresenceCondition::any([
VariablePresenceCondition::all([a.clone(), b.clone()]),
VariablePresenceCondition::all([b.clone(), a.clone()]),
]),
VariablePresenceCondition::all([a.clone(), b.clone()]),
);
}
/// A condition tree built without any normalization.
#[derive(Clone, Debug)]
enum RawCondition {
Always,
Never,
Present(usize),
All(Vec<RawCondition>),
Any(Vec<RawCondition>),
}
const RAW_VARIABLES: [&str; 4] = ["a", "b", "c", "d"];
impl RawCondition {
fn eval(&self, present_mask: u32) -> bool {
match self {
RawCondition::Always => true,
RawCondition::Never => false,
RawCondition::Present(variable_ord) => present_mask & (1 << variable_ord) != 0,
RawCondition::All(children) => {
children.iter().all(|child| child.eval(present_mask))
}
RawCondition::Any(children) => {
children.iter().any(|child| child.eval(present_mask))
}
}
}
/// Builds the canonical condition, visiting children in reverse order if `reverse`.
fn build(&self, reverse: bool) -> VariablePresenceCondition {
let build_children = |children: &[RawCondition]| {
let mut built: Vec<VariablePresenceCondition> =
children.iter().map(|child| child.build(reverse)).collect();
if reverse {
built.reverse();
}
built
};
match self {
RawCondition::Always => VariablePresenceCondition::Always,
RawCondition::Never => VariablePresenceCondition::Never,
RawCondition::Present(variable_ord) => present(RAW_VARIABLES[*variable_ord]),
RawCondition::All(children) => {
VariablePresenceCondition::all(build_children(children))
}
RawCondition::Any(children) => {
VariablePresenceCondition::any(build_children(children))
}
}
}
}
fn raw_conditions() -> impl Strategy<Value = RawCondition> {
let leaf = prop_oneof![
1 => Just(RawCondition::Always),
1 => Just(RawCondition::Never),
6 => (0..RAW_VARIABLES.len()).prop_map(RawCondition::Present),
];
leaf.prop_recursive(4, 32, 4, |inner| {
prop_oneof![
prop::collection::vec(inner.clone(), 0..4).prop_map(RawCondition::All),
prop::collection::vec(inner, 0..4).prop_map(RawCondition::Any),
]
})
}
/// Checks the hidden contract of `ConditionSet`, recursively.
fn assert_canonical(condition: &VariablePresenceCondition) {
let (children, is_all) = match condition {
VariablePresenceCondition::All(children) => (children, true),
VariablePresenceCondition::Any(children) => (children, false),
_ => return,
};
let children: Vec<&VariablePresenceCondition> = children.iter().collect();
assert!(children.len() >= 2, "{condition:?}");
assert!(
children.windows(2).all(|pair| pair[0] < pair[1]),
"{condition:?}"
);
for child in &children {
assert_canonical(child);
match (child, is_all) {
(VariablePresenceCondition::Always | VariablePresenceCondition::Never, _) => {
panic!("neutral or absorbing child in {condition:?}")
}
(VariablePresenceCondition::All(_), true)
| (VariablePresenceCondition::Any(_), false) => {
panic!("same-kind child in {condition:?}")
}
_ => {}
}
}
}
proptest! {
#[test]
fn proptest_canonical_form(raw in raw_conditions()) {
let condition = raw.build(false);
assert_canonical(&condition);
let reversed = raw.build(true);
prop_assert_eq!(&condition, &reversed);
prop_assert_eq!(hash_of(&condition), hash_of(&reversed));
for present_mask in 0..(1u32 << RAW_VARIABLES.len()) {
let mut is_present = |variable_name: &str| {
let variable_ord =
RAW_VARIABLES.iter().position(|name| *name == variable_name).unwrap();
present_mask & (1 << variable_ord) != 0
};
prop_assert_eq!(condition.eval(&mut is_present), raw.eval(present_mask));
}
}
}
fn presence_of<'a>(present_names: &'a [&'a str]) -> impl FnMut(&str) -> bool + 'a {
move |variable_name: &str| present_names.contains(&variable_name)
}
#[test]
fn test_eval() {
let condition = all(vec![present("a"), any(vec![present("b"), present("c")])]);
assert!(condition.eval(&mut presence_of(&["a", "c"])));
assert!(!condition.eval(&mut presence_of(&["a"])));
assert!(!condition.eval(&mut presence_of(&["b", "c"])));
assert!(VariablePresenceCondition::Always.eval(&mut presence_of(&[])));
assert!(!VariablePresenceCondition::Never.eval(&mut presence_of(&["a"])));
}
#[test]
fn test_literals() {
assert_eq!(for_true("true"), VariablePresenceCondition::Always);
assert_eq!(for_true("false"), VariablePresenceCondition::Never);
assert_eq!(for_true("none"), VariablePresenceCondition::Never);
assert_eq!(for_value("none"), VariablePresenceCondition::Always);
assert_eq!(for_value("1u64"), VariablePresenceCondition::Always);
// `none` in a strict function is conservatively ignored.
assert_eq!(for_true("(EQ a none)"), present("a"));
}
#[test]
fn test_variable() {
assert_eq!(for_true("flag"), present("flag"));
assert_eq!(for_value("a"), present("a"));
}
#[test]
fn test_strict_functions() {
assert_eq!(
for_true("(EQ (ADD a 1i64) b)"),
all(vec![present("a"), present("b")])
);
assert_eq!(for_true("(GT (ABS a) 3i64)"), present("a"));
assert_eq!(
for_value(r#"(CONCAT "," "true" (UPPER a) (SUBSTRING b 0i64 2i64))"#),
all(vec![present("a"), present("b")])
);
assert_eq!(for_value(r#"(REGEXP_EXTRACT a "(x+)" 1u64)"#), present("a"));
assert_eq!(for_value("(ADD)"), VariablePresenceCondition::Always);
}
#[test]
fn test_and_or() {
assert_eq!(
for_true("(AND (EQ a 1i64) (LT b 2i64) c)"),
all(vec![present("a"), present("b"), present("c")])
);
assert_eq!(
for_true("(OR (EQ a 1i64) (EQ b 2i64))"),
any(vec![present("a"), present("b")])
);
assert_eq!(
for_true("(AND (OR (EQ a 1i64) (EQ b 2i64)) (EQ c 3i64))"),
all(vec![any(vec![present("a"), present("b")]), present("c")])
);
// AND is null as soon as one of its arguments is null.
assert_eq!(for_value("(AND (NOT a) b)"), present("b"));
assert_eq!(
for_value("(OR (EQ a 1i64) (EQ b 2i64))"),
any(vec![present("a"), present("b")])
);
}
#[test]
fn test_null_tolerant_functions() {
assert_eq!(
for_true("(NOT (EQ a 1i64))"),
VariablePresenceCondition::Always
);
assert_eq!(for_true("(NEQ a 1i64)"), VariablePresenceCondition::Always);
assert_eq!(for_true("(IS_NULL a)"), VariablePresenceCondition::Always);
assert_eq!(
for_value("(IS_NOT_NULL a)"),
VariablePresenceCondition::Always
);
assert_eq!(
for_true("(IS_NOT_NULL (ADD a b))"),
all(vec![present("a"), present("b")])
);
assert_eq!(
for_value(r#"(REGEXP_LIKE a "x")"#),
VariablePresenceCondition::Always
);
assert_eq!(for_true(r#"(REGEXP_LIKE a "x")"#), present("a"));
assert_eq!(
for_true("(OR (EQ a 1i64) (IS_NULL b))"),
VariablePresenceCondition::Always
);
assert_eq!(for_true("(AND (EQ a 1i64) (IS_NULL b))"), present("a"));
}
#[test]
fn test_if() {
assert_eq!(
for_value("(IF c a b)"),
all(vec![present("c"), any(vec![present("a"), present("b")])])
);
assert_eq!(
for_true("(IF c (EQ a 1i64) (EQ b 1i64))"),
all(vec![present("c"), any(vec![present("a"), present("b")])])
);
assert_eq!(for_true("(IF c true false)"), present("c"));
assert_eq!(
for_true("(IF c false false)"),
VariablePresenceCondition::Never
);
assert_eq!(for_value("(IF c 1i64 a)"), present("c"));
}
// The property tests below check that the conditions are indeed necessary, by comparing them
// with the compiled expression over random inputs.
const VARIABLES: [(&str, VarType); 7] = [
("b0", VarType::Bool),
("b1", VarType::Bool),
("n0", VarType::I64),
("n1", VarType::I64),
("f0", VarType::F64),
("s0", VarType::Str),
("s1", VarType::Str),
];
/// Values are indexed by variable, in the order of `VARIABLES`. `None` means null.
type Assignment = Vec<Option<u8>>;
fn variable_value(var_type: VarType, value_ord: u8) -> VariableValue<'static> {
let value_ord = value_ord as usize;
match var_type {
VarType::Bool => VariableValue::some([true, false, true, false][value_ord]),
VarType::I64 => VariableValue::some([0i64, 1, -2, 3][value_ord]),
VarType::F64 => VariableValue::some([0.0f64, 1.5, -1.0, 2.0][value_ord]),
VarType::Str => VariableValue::some(["", "a", "ab,a", "ba"][value_ord]),
VarType::U64 | VarType::None => unreachable!(),
}
}
fn assignments() -> impl Strategy<Value = Vec<Assignment>> {
let value = prop_oneof![Just(None), (0u8..4).prop_map(Some)];
prop::collection::vec(prop::collection::vec(value, VARIABLES.len()), 1..16)
}
struct ExprStrategies {
boolean: BoxedStrategy<String>,
number: BoxedStrategy<String>,
string: BoxedStrategy<String>,
}
fn leaves() -> ExprStrategies {
let pick = |choices: &'static [&'static str]| {
prop::sample::select(choices)
.prop_map(str::to_string)
.boxed()
};
ExprStrategies {
boolean: pick(&["b0", "b1", "true", "false", "none"]),
number: pick(&["n0", "n1", "f0", "0i64", "3i64", "-2i64", "1.5f64", "none"]),
string: pick(&["s0", "s1", r#""a""#, r#""""#, "none"]),
}
}
fn unary(arg: &BoxedStrategy<String>, template: &'static str) -> BoxedStrategy<String> {
arg.clone()
.prop_map(move |arg| template.replace("$0", &arg))
.boxed()
}
fn binary(
left: &BoxedStrategy<String>,
right: &BoxedStrategy<String>,
template: &'static str,
) -> BoxedStrategy<String> {
(left.clone(), right.clone())
.prop_map(move |(left, right)| template.replace("$0", &left).replace("$1", &right))
.boxed()
}
fn ternary(
first: &BoxedStrategy<String>,
second: &BoxedStrategy<String>,
third: &BoxedStrategy<String>,
template: &'static str,
) -> BoxedStrategy<String> {
(first.clone(), second.clone(), third.clone())
.prop_map(move |(first, second, third)| {
template
.replace("$0", &first)
.replace("$1", &second)
.replace("$2", &third)
})
.boxed()
}
/// Returns strategies generating well-typed expressions of the given depth.
fn exprs(depth: u32) -> ExprStrategies {
let leaves = leaves();
if depth == 0 {
return leaves;
}
let ExprStrategies {
boolean: b,
number: n,
string: s,
} = exprs(depth - 1);
let any_kind = prop_oneof![b.clone(), n.clone(), s.clone()].boxed();
let boolean = prop::strategy::Union::new(vec![
leaves.boolean,
binary(&b, &b, "(AND $0 $1)"),
ternary(&b, &b, &b, "(AND $0 $1 $2)"),
binary(&b, &b, "(OR $0 $1)"),
ternary(&b, &b, &b, "(OR $0 $1 $2)"),
unary(&b, "(NOT $0)"),
unary(&any_kind, "(IS_NULL $0)"),
unary(&any_kind, "(IS_NOT_NULL $0)"),
binary(&n, &n, "(EQ $0 $1)"),
binary(&s, &s, "(EQ $0 $1)"),
binary(&b, &b, "(EQ $0 $1)"),
binary(&n, &n, "(NEQ $0 $1)"),
binary(&s, &s, "(NEQ $0 $1)"),
binary(&n, &n, "(LT $0 $1)"),
binary(&n, &n, "(LT_EQ $0 $1)"),
binary(&n, &n, "(GT $0 $1)"),
binary(&s, &s, "(GT_EQ $0 $1)"),
unary(&s, r#"(REGEXP_LIKE $0 "a")"#),
ternary(&b, &b, &b, "(IF $0 $1 $2)"),
])
.boxed();
let number = prop::strategy::Union::new(vec![
leaves.number,
unary(&n, "(ADD $0)"),
binary(&n, &n, "(ADD $0 $1)"),
binary(&n, &n, "(SUBTRACT $0 $1)"),
binary(&n, &n, "(MULTIPLY $0 $1)"),
binary(&n, &n, "(DIVIDE $0 $1)"),
binary(&n, &n, "(POW $0 $1)"),
binary(&n, &n, "(INT_MOD $0 $1)"),
binary(&n, &n, "(MIN $0 $1)"),
binary(&n, &n, "(MAX $0 $1)"),
unary(&n, "(ABS $0)"),
unary(&n, "(CEIL $0)"),
unary(&n, "(FLOOR $0)"),
unary(&n, "(SQRT $0)"),
unary(&n, "(ROUND $0)"),
unary(&n, "(ROUND $0 1i64)"),
// SUBSTRING_COUNT is not generated: its native implementation builds a slice from a
// null pointer when the haystack is null, which aborts debug builds.
ternary(&b, &n, &n, "(IF $0 $1 $2)"),
])
.boxed();
let string = prop::strategy::Union::new(vec![
leaves.string,
unary(&s, "(UPPER $0)"),
unary(&s, "(LOWER $0)"),
unary(&s, "(LEFT $0 1i64)"),
unary(&s, "(RIGHT $0 1i64)"),
unary(&s, "(SUBSTRING $0 0i64 1i64)"),
binary(&s, &s, r#"(CONCAT "," "false" $0 $1)"#),
binary(&s, &s, r#"(TEXT_JOIN "," "true" $0 $1)"#),
unary(&s, r#"(TRIM $0 "a" "both")"#),
unary(&s, r#"(SPLIT_AFTER $0 ",")"#),
unary(&s, r#"(SPLIT_BEFORE $0 "," 0i64)"#),
unary(&s, r#"(REGEXP_EXTRACT $0 "(a)b" 1u64)"#),
// IF is not generated for strings: with a null condition, it returns the selected
// branch instead of null, as the string pointer is not cleared.
])
.boxed();
ExprStrategies {
boolean,
number,
string,
}
}
/// Compiles `expr_str`, then checks that `required(expr)` holds for every assignment where
/// `holds(result)` is true.
///
/// Following tantivy's fast field binding, a variable is bound only if its type is accepted by
/// type inference. Unbound variables are null, and therefore absent.
fn check_necessary_condition(
expr_str: &str,
target_type: InferredTypeSet,
assignments: &[Assignment],
required: fn(&UntypedExpr) -> VariablePresenceCondition,
holds: fn(VarType, VariableValue) -> bool,
) -> Result<(), TestCaseError> {
let expr = deserialize(expr_str).unwrap();
let Ok(inferred_types) = infer_types_with_target(&expr, target_type) else {
return Err(TestCaseError::reject("type inference failed"));
};
let mut variable_types: HashMap<&str, VarType> =
HashMap::with_capacity(inferred_types.len());
for (variable_name, accepted_types) in &inferred_types {
let (_, var_type) = VARIABLES
.iter()
.find(|(name, _)| name == variable_name)
.unwrap();
if accepted_types.contains(*var_type) {
variable_types.insert(*variable_name, *var_type);
}
}
// Some expressions trip debug assertions of the compiler, unrelated to presence. For
// instance, `(SQRT (CEIL n0))` asks CEIL for a f64, while it always returns an i64.
let compile_result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
compile(&expr, &variable_types)
}));
let Ok(Ok(compiled_fn)) = compile_result else {
return Err(TestCaseError::reject("compilation failed"));
};
let condition = required(&expr);
let mut string_arena = StringArena::default();
for assignment in assignments {
let variable_ord = |variable_name: &str| {
VARIABLES
.iter()
.position(|(name, _)| *name == variable_name)
.unwrap()
};
let args: Vec<VariableValue> = compiled_fn
.inputs()
.iter()
.map(
|input| match assignment[variable_ord(&input.variable_name)] {
Some(value_ord) => variable_value(input.r#type, value_ord),
None => VariableValue::none(),
},
)
.collect();
// SAFETY: Each slot follows the compiled input order, and uses the input type.
let result = unsafe { compiled_fn.call(&args, &mut string_arena) };
if !holds(compiled_fn.result_type(), result) {
continue;
}
let mut is_present = |variable_name: &str| {
variable_types.contains_key(variable_name)
&& assignment[variable_ord(variable_name)].is_some()
};
prop_assert!(
condition.eval(&mut is_present),
"{expr_str} holds for {assignment:?}, but {condition:?} does not"
);
}
Ok(())
}
fn is_true(result_type: VarType, result: VariableValue) -> bool {
// SAFETY: The union member is selected with the result type.
result_type == VarType::Bool && unsafe { result.as_bool() } == Some(true)
}
fn is_present(result_type: VarType, result: VariableValue) -> bool {
// SAFETY: The union member is selected with the result type.
unsafe {
match result_type {
VarType::Bool => result.as_bool().is_some(),
VarType::F64 => result.as_f64().is_some(),
VarType::U64 => result.as_u64().is_some(),
VarType::I64 => result.as_i64().is_some(),
VarType::Str => result.as_str().is_some(),
VarType::None => false,
}
}
}
proptest! {
// Compiler debug assertions reject a fraction of the generated expressions.
#![proptest_config(ProptestConfig {
max_global_rejects: 1 << 16,
..ProptestConfig::with_cases(512)
})]
#[test]
fn proptest_required_presence_for_true_is_necessary(
expr in exprs(3).boolean,
assignments in assignments(),
) {
check_necessary_condition(
&expr,
InferredTypeSet::BOOLEAN,
&assignments,
required_presence_for_true,
is_true,
)?;
}
#[test]
fn proptest_required_presence_is_necessary(
expr in prop_oneof![exprs(3).boolean, exprs(3).number, exprs(3).string],
assignments in assignments(),
) {
check_necessary_condition(
&expr,
InferredTypeSet::ALL,
&assignments,
required_presence,
is_present,
)?;
}
}
}
+29 -21
View File
@@ -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<Literal> {
}
if let Some(value_str) = atom.strip_suffix("f64") {
let val = value_str.parse::<f64>().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::<f64>().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());
}
}
+306
View File
@@ -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<Arc<CompiledFn>, 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<OnceLock<CompilationResult>>;
/// 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<Mutex<ExprCompilationCacheInner>>,
}
struct ExprCompilationCacheInner {
capacity: usize,
// We use Option here to lazily allocate on the first insertion.
entries: Option<LruCache<ExprCacheKey, CompilationSlot>>,
}
impl ExprCompilationCacheInner {
fn entries(&mut self) -> Option<&mut LruCache<ExprCacheKey, CompilationSlot>> {
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<Arc<CompiledFn>, 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<CompilationSlot> {
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<Arc<CompiledFn>> = 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);
}
}
+7 -7
View File
@@ -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"),
}
+5 -3
View File
@@ -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<cranelift_module::ModuleError>),
Module(#[source] Arc<cranelift_module::ModuleError>),
#[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<cranelift_module::ModuleError> for CompileError {
fn from(error: cranelift_module::ModuleError) -> Self {
CompileError::Module(Box::new(error))
CompileError::Module(Arc::new(error))
}
}
+3 -1
View File
@@ -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) => {
+12 -14
View File
@@ -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<str>),
}
+1 -1
View File
@@ -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),
}
}
+17 -5
View File
@@ -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);
+20 -7
View File
@@ -98,7 +98,7 @@ impl From<IsNullFnCall> 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));
+1 -1
View File
@@ -38,7 +38,7 @@ fn constant_length(expression: &UntypedExpr) -> Result<Option<usize>, 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,
})
+1 -1
View File
@@ -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 {
+1 -1
View File
@@ -39,7 +39,7 @@ fn constant_length(expression: &UntypedExpr) -> Result<Option<usize>, 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,
})
+4 -5
View File
@@ -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,
+1 -1
View File
@@ -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,
})
+1 -1
View File
@@ -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,
})
+5 -2
View File
@@ -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]
+1 -1
View File
@@ -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,
+2
View File
@@ -1,3 +1,5 @@
//! EXPERIMENTAL. The API is likely to change in the near future.
pub mod ast;
pub mod compile;
pub mod types;
+195 -1
View File
@@ -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<Self> {
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<i128>) -> 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<Ordering> {
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<H: Hasher>(&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<VariablePrimitiveOpt> 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<SafeF64> {
[
-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<f64> = 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() {
+20 -11
View File
@@ -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<LenientError>) {
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(),
+3 -2
View File
@@ -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 },
}
+59 -17
View File
@@ -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<Option<Arc<dyn ValueSource>>> {
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<u64>, ColumnType)> {
) -> crate::Result<Arc<dyn ValueSource>> {
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<Vec<columnar::DynamicColumn>> {
let ff_fields = reader.fast_fields().dynamic_column_handles(field_name)?;
let cols = ff_fields
let dyn_col_handles: Vec<DynamicColumnHandle> =
reader.fast_fields().dynamic_column_handles(field_name)?;
let dyn_cols: Vec<DynamicColumn> = dyn_col_handles
.iter()
.map(|h| h.open())
.map(DynamicColumnHandle::open)
.collect::<io::Result<_>>()?;
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<Vec<(columnar::Column<u64>, ColumnType)>> {
) -> crate::Result<Vec<Arc<dyn ValueSource>>> {
// 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<dyn ValueSource> = Arc::new((column, column_type));
source
})
.collect())
}
+131 -83
View File
@@ -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<TermsAggReqData>,
pub(crate) term_req_data: Vec<TermsAggReqData>,
/// HistogramAggReqData contains the request data for a histogram aggregation.
pub histogram_req_data: Vec<HistogramAggReqData>,
pub(crate) histogram_req_data: Vec<HistogramAggReqData>,
/// RangeAggReqData contains the request data for a range aggregation.
pub range_req_data: Vec<RangeAggReqData>,
pub(crate) range_req_data: Vec<RangeAggReqData>,
/// FilterAggReqData contains the request data for a filter aggregation.
pub filter_req_data: Vec<FilterAggReqData>,
pub(crate) filter_req_data: Vec<FilterAggReqData>,
/// Shared by avg, min, max, sum, stats, extended_stats, count
pub stats_metric_req_data: Vec<MetricAggReqData>,
pub(crate) stats_metric_req_data: Vec<MetricAggReqData>,
/// CardinalityAggReqData contains the request data for a cardinality aggregation.
pub cardinality_req_data: Vec<CardinalityAggReqData>,
pub(crate) cardinality_req_data: Vec<CardinalityAggReqData>,
/// TopHitsAggReqData contains the request data for a top_hits aggregation.
pub top_hits_req_data: Vec<TopHitsAggReqData>,
pub(crate) top_hits_req_data: Vec<TopHitsAggReqData>,
/// MissingTermAggReqData contains the request data for a missing term aggregation.
pub missing_term_req_data: Vec<MissingTermAggReqData>,
pub(crate) missing_term_req_data: Vec<MissingTermAggReqData>,
/// CompositeAggReqData contains the request data for a composite aggregation.
pub composite_req_data: Vec<CompositeAggReqData>,
pub(crate) composite_req_data: Vec<CompositeAggReqData>,
/// MultiTermsAggReqData contains the request data for a multi_terms aggregation.
pub multi_terms_req_data: Vec<MultiTermsAggReqData>,
pub(crate) multi_terms_req_data: Vec<MultiTermsAggReqData>,
/// Request tree used to build collectors.
pub agg_tree: Vec<AggRefNode>,
pub(crate) agg_tree: Vec<AggRefNode>,
}
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<dyn SegmentAggregationCollector> =
if is_str && max_term_ord_inclusive < BITSET_MAX_TERM_ORD {
Box::new(SegmentCardinalityCollector::<BitSet>::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::<TermOrdSet>::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<Option<u64>> {
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<Column<u64>> {
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<Vec<AggRefNode>> {
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<u64>, 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::<crate::Result<_>>()?;
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::<crate::Result<Vec<_>>>()?;
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<Key>,
include_bytes: bool,
) -> crate::Result<Vec<(Column<u64>, ColumnType)>> {
) -> crate::Result<Vec<Arc<dyn ValueSource>>> {
// `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::<Vec<_>>();
// 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::<crate::Result<Vec<_>>>()?;
// 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(),
@@ -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<CompositeSourceAccessors>,
pub(crate) composite_accessors: Vec<CompositeSourceAccessors>,
}
impl CompositeAggReqData {
+2 -1
View File
@@ -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::{
+5 -7
View File
@@ -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<DocumentQueryEvaluator>,
pub(crate) evaluator: Rc<DocumentQueryEvaluator>,
/// 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 {
+39 -31
View File
@@ -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<u64>,
/// The field type of the fast field.
pub field_type: ColumnType,
pub(crate) accessor: Arc<dyn ValueSource>,
/// 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<B> {
parent_buckets: Vec<HistogramBuckets<B>>,
sub_agg: Option<HighCardBufferedSubAggs>,
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<B: BucketIdSlot> 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<B: BucketIdSlot> 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<B: BucketIdSlot> SegmentHistogramCollector<B> {
}
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<B: BucketIdSlot> SegmentHistogramCollector<B> {
.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<u64>,
field_type: ColumnType,
accessor: &dyn ValueSource,
interval: f64,
offset: f64,
bounds: HistogramBounds,
) -> Option<DenseRange> {
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 {
+233 -118
View File
@@ -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<MultiTermsFieldAccessor>,
pub(crate) fields: Vec<MultiTermsFieldAccessor>,
/// 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<Option<MultiTermsMissingAccessor>>,
/// Sub-aggregation descriptor (empty when no sub-aggs).
pub sub_aggregations: Aggregations,
pub(crate) missing_accessors: Vec<Option<MultiTermsMissingAccessor>>,
/// 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<I>(&self, keys: &mut [Self::PackingType], field_idx: usize, values: I)
where I: IntoIterator<Item = u64>;
fn unpack(
&self,
key: &Self::PackingType,
req_data: &MultiTermsAggReqData,
) -> crate::Result<Vec<IntermediateKey>>;
/// 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<Vec<IntermediateKey>> {
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<Vec<IntermediateKey>> {
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<Vec<IntermediateKey>, 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<P: MultiTermsPacking, B>(
packing: &P,
entries: &[MultiTermsBucketEntry<P::PackingType, B>],
req_data: &MultiTermsAggReqData,
) -> crate::Result<Vec<Vec<IntermediateKey>>> {
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<MultiTermsBucketEntry> = 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::<Vec<_>>()
.join("|");
let keys: Vec<Key> = 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::<crate::Result<_>>()?;
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<IntermediateKey>, IntermediateTermBucketEntry)| {
key_vec.iter().cloned().map(Key::from).collect::<Vec<_>>()
};
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::<crate::Result<Vec<_>>>()?;
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::<Vec<_>>()
.join("|");
let keys: Vec<Key> = 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(())
}
+13 -17
View File
@@ -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<u64>,
/// The type of the fast field.
pub field_type: ColumnType,
pub(crate) accessor: Arc<dyn ValueSource>,
/// 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<B: SubAggBuffer> {
/// The buckets containing the aggregation data.
/// One for each ParentBucketId
parent_buckets: Vec<Vec<SegmentRangeAndBucketEntry>>,
column_type: ColumnType,
pub(crate) req_data: RangeAggReqData,
sub_agg: Option<BufferedSubAggs<B>>,
/// Here things get a bit weird. We need to assign unique bucket ids across all
@@ -184,7 +183,7 @@ impl<B: SubAggBuffer> Debug for SegmentRangeCollector<B> {
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<B: SubAggBuffer> SegmentAggregationCollector for SegmentRangeCollector<B> {
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<B: SubAggBuffer> SegmentAggregationCollector for SegmentRangeCollector<B> {
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<B: SubAggBuffer> SegmentAggregationCollector for SegmentRangeCollector<B> {
) -> 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::<LowCardSubAggBuffer> {
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::<HighCardSubAggBuffer> {
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<B: SubAggBuffer> SegmentRangeCollector<B> {
pub(crate) fn create_new_buckets(&mut self) -> crate::Result<Vec<SegmentRangeAndBucketEntry>> {
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.
@@ -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<const LANES: usize>(
/// 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<u64>,
column_values: Arc<dyn ColumnValues>,
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<dyn ColumnValues>,
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<const NUM_BUCKETS: usize> {
hist_block: ColumnBlockAccessor,
next_count_lane: usize,
accessor: Column<u64>,
column_values: Arc<dyn ColumnValues>,
boundaries: [u64; NUM_BUCKETS],
num_buckets: usize,
}
impl<const NUM_BUCKETS: usize> LinearBucketResolver<NUM_BUCKETS> {
/// 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<dyn ColumnValues>,
base_pos: i64,
num_time_buckets: usize,
) -> Option<Self> {
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<const NUM_BUCKETS: usize> LinearBucketResolver<NUM_BUCKETS> {
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<const NUM_BUCKETS: usize> BucketResolver for LinearBucketResolver<NUM_BUCKE
#[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]
@@ -328,10 +318,11 @@ fn first_encoded_value_for_bucket(
hist_req_data: &HistogramAggReqData,
base_pos: i64,
) -> 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<R: BucketResolver, const LANES: usize> {
/// 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<R: BucketResolver, const LANES: usize> {
/// `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<dyn ColumnValues>,
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<R: BucketResolver, const LANES: usize> {
all_docs_in_bounds: bool,
}
impl<R: BucketResolver, const LANES: usize> Debug for FlattenedTermHistogramCollector<R, LANES> {
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<R: BucketResolver, const LANES: usize> SegmentAggregationCollector
for FlattenedTermHistogramCollector<R, LANES>
{
@@ -455,7 +455,7 @@ impl<R: BucketResolver, const LANES: usize> 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::<SINGLE_COUNT_LANE>(
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<Arc<dyn ColumnValues>> {
let column = value_source.as_column()?;
if column.get_cardinality() == Cardinality::Full {
Some(column.values.clone())
} else {
None
}
}
fn build_flattened_collector<const LANES: usize>(
agg_data: &mut AggregationsSegmentCtx,
terms_req_data: &TermsAggReqData,
hist_req_data: HistogramAggReqData,
terms_values: Arc<dyn ColumnValues>,
hist_values: Arc<dyn ColumnValues>,
num_terms: usize,
num_time_buckets: usize,
base_pos: i64,
range: DenseRange,
) -> crate::Result<Box<dyn SegmentAggregationCollector>> {
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::<SingleBucketResolver, LANES>(
agg_data,
terms_req_data,
terms_values,
hist_req_data,
num_terms,
base_pos,
@@ -614,12 +634,14 @@ fn build_flattened_collector<const LANES: usize>(
if all_docs_in_bounds && num_time_buckets <= NUM_SMALL_LINEAR_BUCKETS {
if let Some(resolver) = LinearBucketResolver::<NUM_SMALL_LINEAR_BUCKETS>::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<const LANES: usize>(
} else if all_docs_in_bounds && num_time_buckets <= NUM_LARGE_LINEAR_BUCKETS {
if let Some(resolver) = LinearBucketResolver::<NUM_LARGE_LINEAR_BUCKETS>::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<const LANES: usize>(
}
}
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<const LANES: usize>(
fn build_flattened_collector_with_resolver<R: BucketResolver, const LANES: usize>(
agg_data: &mut AggregationsSegmentCtx,
terms_req_data: &TermsAggReqData,
terms_values: Arc<dyn ColumnValues<u64>>,
hist_req_data: HistogramAggReqData,
num_terms: usize,
base_pos: i64,
@@ -686,6 +717,7 @@ fn build_flattened_collector_with_resolver<R: BucketResolver, const LANES: usize
counts,
base_pos,
terms_req_data: terms_req_data.clone(),
terms_values,
hist_req_data,
term_block: ColumnBlockAccessor::default(),
bucket_resolver,
+37 -25
View File
@@ -1,10 +1,11 @@
use std::fmt::Debug;
use std::hash::Hash;
use std::net::Ipv6Addr;
use std::sync::Arc;
use columnar::column_values::CompactSpaceU64Accessor;
use columnar::{
Column, ColumnType, Dictionary, MonotonicallyMappableToU128, MonotonicallyMappableToU64,
ColumnType, Dictionary, MonotonicallyMappableToU128, MonotonicallyMappableToU64,
NumericalValue, StrColumn,
};
use common::{BitSet, TinySet};
@@ -26,7 +27,7 @@ use crate::aggregation::intermediate_agg_result::{
IntermediateKey, IntermediateTermBucketEntry, IntermediateTermBucketResult,
};
use crate::aggregation::segment_agg_result::{BucketIdProvider, SegmentAggregationCollector};
use crate::aggregation::{format_date, BucketId, Key};
use crate::aggregation::{format_date, BucketId, Key, ValueSource};
use crate::error::DataCorruption;
use crate::TantivyError;
@@ -35,25 +36,23 @@ mod flattened_term_histogram;
/// Contains all information required by the SegmentTermCollector to perform the
/// terms aggregation on a segment.
#[derive(Debug, Clone)]
pub struct TermsAggReqData {
pub(crate) struct TermsAggReqData {
/// The column accessor to access the fast field values.
pub accessor: Column<u64>,
/// The type of the column.
pub column_type: ColumnType,
pub(crate) accessor: Arc<dyn ValueSource>,
/// The string dictionary column if the field is of type text.
pub str_dict_column: Option<StrColumn>,
pub(crate) str_dict_column: Option<StrColumn>,
/// The missing value as u64 value.
pub missing_value_for_accessor: Option<u64>,
pub(crate) missing_value_for_accessor: Option<u64>,
/// 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<BitSet>,
pub(crate) allowed_term_ids: Option<BitSet>,
/// 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<Box<dyn SegmentAggregationCollector>> {
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<TermMap: TermAggregationMap, B: SubAggBuffer> 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::<CompactSpaceU64Accessor>()
@@ -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<K> GetDocCount for (K, IntermediateTermBucketEntry) {
fn doc_count(&self) -> u64 {
self.1.doc_count
}
+4 -4
View File
@@ -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<u64>, ColumnType)>,
pub(crate) accessors: Vec<(Column<u64>, 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 {
+110 -9
View File
@@ -1213,19 +1213,20 @@ trait MergeFruits {
fn merge_fruits(&mut self, other: Self) -> crate::Result<()>;
}
fn merge_maps<V: MergeFruits + Clone, T: Eq + PartialEq + Hash>(
fn merge_maps<V: MergeFruits, T: Eq + Hash>(
entries_left: &mut FxHashMap<T, V>,
mut entries_right: FxHashMap<T, V>,
entries_right: FxHashMap<T, V>,
) -> 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;
File diff suppressed because it is too large Load Diff
+614
View File
@@ -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<Key>,
}
/// 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<dyn ValueSource>,
/// The string dictionary column if the field is of type string.
pub(crate) str_dict_column: Option<StrColumn>,
/// The missing value normalized to the internal u64 representation of the field type.
pub(crate) missing_value_for_accessor: Option<u64>,
/// 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::<Self>()
}
}
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<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
let bytes = self.sketch.serialize();
serializer.serialize_bytes(&bytes)
}
}
impl<'de> Deserialize<'de> for CardinalityCollector {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let bytes: Vec<u8> = 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<f64> {
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<u8> {
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<Box<dyn SegmentAggregationCollector>> {
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::<BitSet>::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::<TermOrdSet>::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);
}
}
@@ -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<Option<CardinalityCollector>>,
accessor_idx: usize,
/// The column accessor to access the fast field values.
accessor: Arc<dyn ValueSource>,
/// 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<Arc<CompactSpaceU64Accessor>>,
/// The missing value normalized to the internal u64 representation of the field type.
missing_value_for_accessor: Option<u64>,
}
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<dyn ValueSource>,
missing_value_for_accessor: Option<u64>,
) -> crate::Result<Self> {
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::<CompactSpaceU64Accessor>()
.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<f64> {
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(())
}
}
@@ -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<Coupon>,
missing_coupon_opt: Option<Coupon>,
},
Sparse {
coupon_map: FxHashMap<u64, Coupon>,
missing_coupon_opt: Option<Coupon>,
},
}
impl CouponCache {
fn new(
term_ords: Vec<u64>,
coupons: Vec<Coupon>,
missing_coupon_opt: Option<Coupon>,
) -> 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<Coupon> =
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<u64, Coupon> = 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<S: TermOrdAccumulator> {
/// Buckets are Some(_) until they get consumed by
/// `add_intermediate_aggregation_result`.
buckets: Vec<Option<S>>,
accessor_idx: usize,
/// The column accessor to access the fast field values (term ordinals).
accessor: Arc<dyn ValueSource>,
/// The missing value normalized to the internal u64 representation of the field type.
missing_value_for_accessor: Option<u64>,
/// Lazily built at finalization time, shared by every bucket.
coupon_cache: Option<CouponCache>,
/// Largest term_ord that may be inserted into a bucket, i.e.
/// `accessor.max_value()`.
max_term_ord_inclusive: u64,
}
impl<S: TermOrdAccumulator> Debug for SegmentStrCardinalityCollector<S> {
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<S: TermOrdAccumulator>(
buckets: &[Option<S>],
dictionary: &Dictionary,
missing_value_opt: Option<&Key>,
) -> io::Result<CouponCache> {
// 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<u64> = 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<Coupon> = 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<Coupon> = 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<S: TermOrdAccumulator> SegmentStrCardinalityCollector<S> {
pub fn from_req(
accessor_idx: usize,
accessor: Arc<dyn ValueSource>,
missing_value_for_accessor: Option<u64>,
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<S: TermOrdAccumulator + 'static> SegmentAggregationCollector
for SegmentStrCardinalityCollector<S>
{
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<f64> {
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<String> = (0..32).map(|i| format!("term_{i}")).collect();
let term_refs: Vec<Vec<&str>> = terms.iter().map(|t| vec![t.as_str()]).collect::<Vec<_>>();
// 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);
}
}
@@ -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<u64> -> 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<u64> 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<Option<Box<Page>>>` 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<Option<Box<PagedBitsetPage>>>,
/// 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<Item = u64> + '_ {
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<u64>` 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<u64>` 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<Item = u64>);
/// 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<Item = u64> + '_;
}
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<Item = u64> + '_ {
// `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<Item = u64>) {
for ord in ords {
<Self as TermOrdAccumulator>::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<u64>),
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<Item = u64>) {
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<Item = u64> + '_ {
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<u64> = bitset.iter_sorted().collect();
let mut expected: Vec<u64> = 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 = <TermOrdSet as TermOrdAccumulator>::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<u64> = 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 = <BitSet as TermOrdAccumulator>::new(255);
for ord in [0u64, 1, 63, 64, 65, 128, 200, 200, 0] {
<BitSet as TermOrdAccumulator>::insert(&mut set, ord);
}
assert_eq!(<BitSet as TermOrdAccumulator>::len(&set), 7);
let collected: Vec<u64> = set.iter_ords().collect();
assert_eq!(collected, vec![0, 1, 63, 64, 65, 128, 200]);
}
}
+10 -8
View File
@@ -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<u64>,
field_type: ColumnType,
accessor: columnar::Column<u64>,
accessor: Arc<dyn ValueSource>,
buckets: Vec<IntermediateExtendedStats>,
sigma: Option<f64>,
}
@@ -331,10 +331,9 @@ impl SegmentExtendedStatsCollector {
pub fn from_req(req: &MetricAggReqData, sigma: Option<f64>) -> 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);
}
+9 -10
View File
@@ -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<u64>,
pub(crate) missing_u64: Option<u64>,
/// The column accessor to access the fast field values.
pub accessor: Column<u64>,
pub(crate) accessor: Arc<dyn ValueSource>,
/// Used when converting to intermediate result
pub collecting_for: StatsType,
pub(crate) collecting_for: StatsType,
/// The missing value
pub missing: Option<f64>,
pub(crate) missing: Option<f64>,
/// The name of the aggregation.
pub name: String,
pub(crate) name: String,
}
impl MetricAggReqData {
+6 -8
View File
@@ -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<PercentilesCollector>,
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<u64>,
/// The column accessor to access the fast field values.
pub accessor: Column<u64>,
pub(crate) accessor: Arc<dyn ValueSource>,
}
#[derive(Clone, Serialize, Deserialize)]
@@ -250,14 +249,12 @@ impl PercentilesCollector {
impl SegmentPercentilesCollector {
pub fn from_req_and_validate(
field_type: ColumnType,
missing_u64: Option<u64>,
accessor: Column<u64>,
accessor: Arc<dyn ValueSource>,
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);
}
+36 -15
View File
@@ -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<const TYPE_ID: u8>(
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<const TYPE_ID: u8>(
pub(crate) fn build_segment_stats_collector(
req: &MetricAggReqData,
) -> crate::Result<Box<dyn SegmentAggregationCollector>> {
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<const COLUMN_TYPE_ID: u8> {
pub(crate) missing_u64: Option<u64>,
pub(crate) accessor: Column<u64>,
pub(crate) accessor: Arc<dyn ValueSource>,
/// 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<u64>, ColumnType)>,
pub(crate) is_number_or_date_type: bool,
pub(crate) buckets: Vec<IntermediateStats>,
pub(crate) name: String,
@@ -290,21 +301,31 @@ impl<const COLUMN_TYPE_ID: u8> 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::<COLUMN_TYPE_ID>(
&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::<COLUMN_TYPE_ID>(
&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::<COLUMN_TYPE_ID>(
&mut self.buckets[parent_bucket_id as usize],
agg_data.column_block_accessor.iter_vals(),
+6 -6
View File
@@ -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<u64>, ColumnType)>,
pub(crate) accessors: Vec<(Column<u64>, ColumnType)>,
/// The accessors to access the fast field values for retrieving document fields.
pub value_accessors: HashMap<String, Vec<DynamicColumn>>,
pub(crate) value_accessors: HashMap<String, Vec<DynamicColumn>>,
/// 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 {
+18 -3
View File
@@ -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<ValueSourceRegistry>,
}
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<ValueSourceRegistry>) -> Self {
self.value_sources = value_sources;
self
}
}
@@ -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<u64>,
docids: &mut Vec<DocId>,
row_ids: &mut Vec<RowId>,
) -> 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<u64>,
/// 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<DocId>,
/// Scratch buffer available to sources for translating document IDs into value row IDs.
row_id_cache: Vec<RowId>,
/// 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<u64> {
#[inline]
fn load_block(
&self,
docs: &[DocId],
values: &mut Vec<u64>,
docids: &mut Vec<DocId>,
row_ids: &mut Vec<RowId>,
) -> 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<S: ValueSource + ?Sized>(&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<u64>) {
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<u64>,
) {
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<S: ValueSource + ?Sized>(
&mut self,
docs: &[DocId],
source: &impl BlockValueSource,
source: &S,
missing_opt: Option<u64>,
) {
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<S: ValueSource + ?Sized>(
&mut self,
docs: &[DocId],
source: &impl BlockValueSource,
source: &S,
missing_opt: Option<u64>,
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<u64>,
ordered: bool,
) {
@@ -288,37 +267,6 @@ impl ColumnBlockAccessor {
}
}
#[inline]
fn load_full_column_values(docs: &[DocId], accessor: &Column<u64>, values: &mut Vec<u64>) {
// 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<u32>) {
#[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::<Vec<_>>(),
[(2, 20), (4, 40), (8, 80)]
);
}
fn full_column(vals: &[u64]) -> Column<u64> {
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::<u64>(&vals, &ALL_U64_CODEC_TYPES),
}
}
#[test]
fn test_as_column_distinguishes_the_two_kinds() {
let column: Arc<dyn ValueSource> = 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<dyn ValueSource> = 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<dyn ValueSource> = 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::<Vec<_>>(),
[(1, 11), (3, 33)]
);
accessor.fetch_block_with_missing(&docs, &*computed, Some(99));
let mut pairs = accessor.iter_docid_vals(&docs).collect::<Vec<_>>();
pairs.sort_unstable();
assert_eq!(pairs, [(0, 99), (1, 11), (2, 99), (3, 33)]);
}
#[test]
fn test_find_missing_docs() {
let docs: Vec<u32> = vec![1, 2, 3, 4, 5, 6, 7, 8, 9, 10];
@@ -392,30 +418,25 @@ mod tests {
let mut missing_docs: Vec<u32> = 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<u32> = Vec::new();
let hits: Vec<u32> = vec![2, 4, 6, 8, 10];
let mut missing_docs: Vec<u32> = Vec::new();
find_missing_docs(&docs, &hits, &mut missing_docs);
assert_eq!(missing_docs, Vec::<u32>::new());
assert_eq!(missing_docs, [0u32; 0]);
}
#[test]
fn test_find_missing_docs_all_missing() {
let docs: Vec<u32> = vec![1, 2, 3, 4, 5];
let hits: Vec<u32> = Vec::new();
let docs: &[u32] = &[1, 2, 3, 4, 5];
let hits: &[u32] = &[];
let mut missing_docs: Vec<u32> = 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<_>>(),
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<_>>(),
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<_>>(),
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<_>>(),
vec![99, 10, 99, 40, 70, 99]
[99, 10, 99, 40, 70, 99]
);
assert_eq!(
accessor.iter_docid_vals(&docs).collect::<Vec<_>>(),
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);
+126
View File
@@ -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<u64>,
docids: &mut Vec<DocId>,
row_ids: &mut Vec<RowId>,
) -> Cardinality;
/// Returns the physical column, if this source is backed by one.
fn as_column(&self) -> Option<&Column<u64>> {
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<ColumnRef: Borrow<Column<u64>> + 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<u64>,
docids: &mut Vec<DocId>,
row_ids: &mut Vec<RowId>,
) -> 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<u64>> {
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<u64>,
values: &mut Vec<u64>,
) {
// 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()
}
+117
View File
@@ -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<u64>,
_docids: &mut Vec<DocId>,
_row_ids: &mut Vec<columnar::RowId>,
) -> 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<Arc<dyn ValueSource>> {
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]));
}
@@ -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<Arc<dyn ValueSource>>;
}
/// Named computed sources available to an aggregation request.
#[derive(Clone, Default)]
pub struct ValueSourceRegistry {
providers: HashMap<String, Arc<dyn ValueSourceProvider>>,
}
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<dyn ValueSourceProvider>) {
let name = name.to_string();
self.providers.insert(name, provider);
}
#[inline]
pub(crate) fn get(&self, name: &str) -> Option<&Arc<dyn ValueSourceProvider>> {
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]);
}
}
+1 -7
View File
@@ -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()?;
}
+9 -14
View File
@@ -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<DirectoryLock, TryAcquireLockError> {
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<WritePtr, OpenWriteError>;
+11 -16
View File
@@ -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<PathBuf> = (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![];
+131 -57
View File
@@ -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<AtomicUsize>,
reported_bytes: usize,
unreported_bytes: usize,
}
impl MemoryUsageTracker {
fn new(shared_usage: Arc<AtomicUsize>) -> 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<Vec<u8>>,
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<usize> {
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<PathBuf, FileSlice>,
active_writers: HashSet<PathBuf>,
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<u8>) -> bool {
self.fs.insert(path, FileSlice::from(data)).is_some()
}
fn open_read(&self, path: &Path) -> Result<FileSlice, OpenReadError> {
@@ -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<RwLock<InnerDirectory>>,
active_writer_mem_usage: Arc<AtomicUsize>,
}
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<WritePtr, OpenWriteError> {
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<Vec<u8>, 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);
}
}
+11 -11
View File
@@ -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]);
-1
View File
@@ -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(())
}
+22 -1
View File
@@ -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<Arc<dyn SegmentPlugin>>,
#[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)
}
}
+29 -35
View File
@@ -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);
+5
View File
@@ -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 {
+95 -2
View File
@@ -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<IndexSettings>) -> 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::<Vec<_>>()
);
}
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(
+2
View File
@@ -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};
+305
View File
@@ -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<DocId>],
) -> Option<DocId> {
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<DocId>],
}
impl MappedPostings<'_> {
fn advance(&mut self) -> Option<DocId> {
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<Option<DocId>>],
cursors: Vec<MappedPostings<'a>>,
/// Min-heap of `(current mapped doc, index into cursors)`.
heap: BinaryHeap<Reverse<(DocId, usize)>>,
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<Option<DocId>>]) -> 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<Item = (usize, SegmentPostings)>) {
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<u32>) {
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<u32>)> {
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<u32>)> {
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<Vec<Vec<(usize, SegmentPostings)>>> {
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<Option<DocId>> = (0..N)
.map(|doc| if doc % 10 == 0 { None } else { Some(doc * 2) })
.collect();
let map1: Vec<Option<DocId>> = (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(())
}
}
+2
View File
@@ -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;
-1
View File
@@ -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(())
}
+2 -1
View File
@@ -47,7 +47,8 @@ pub struct AllScorer {
impl AllScorer {
/// Creates a new AllScorer with `max_doc` docs.
pub fn new(max_doc: DocId) -> AllScorer {
AllScorer { doc: 0u32, max_doc }
let doc = if max_doc == 0u32 { TERMINATED } else { 0 };
AllScorer { doc, max_doc }
}
}
@@ -1,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<SegmentF> {
(self.segment_predicate_factory)(segment_reader)
fn doc_predicate(
&self,
segment_reader: &SegmentReader,
) -> crate::Result<ConstOrVariableSegmentPredicate<SegmentF>> {
let predicate = (self.segment_predicate_factory)(segment_reader)?;
Ok(ConstOrVariableSegmentPredicate::Variable {
predicate,
necessary_condition: Box::new(AllScorer::new(segment_reader.max_doc())),
})
}
}
@@ -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<Self, TypeError> {
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<ConstOrVariableSegmentPredicate<JitExprEvalState>> {
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<dyn DocSet> = build_necessary_condition_docset(
&self.required_presence,
&opened_columns,
segment_reader.max_doc(),
);
if necessary_condition.doc() == TERMINATED {
return Ok(ConstOrVariableSegmentPredicate::Const(false));
}
let compiled_fn = segment_reader
.index()
.expr_compilation_cache()
.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<VariableValue> =
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<DynamicColumn> =
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<dyn DocSet> {
match condition {
VariablePresenceCondition::Always => Box::new(AllScorer::new(max_doc)),
VariablePresenceCondition::Never => Box::new(EmptyScorer),
VariablePresenceCondition::Present(variable_name) => {
let Some(column) = columns.get(variable_name.as_ref()) else {
return Box::new(EmptyScorer);
};
let exists_column_index = match column.column_index() {
ColumnIndex::Empty { .. } => return Box::new(EmptyScorer),
ColumnIndex::Full => return Box::new(AllScorer::new(max_doc)),
ColumnIndex::Optional(optional_index) => {
ExistsColumnIndex::Optional(optional_index.clone())
}
ColumnIndex::Multivalued(multivalued_index) => {
ExistsColumnIndex::Multivalued(multivalued_index.clone())
}
};
Box::new(ExistsDocSet::new(exists_column_index))
}
VariablePresenceCondition::All(conditions) => {
let mut doc_sets: Vec<Box<dyn DocSet>> = conditions
.iter()
.map(|condition| build_necessary_condition_docset(condition, columns, max_doc))
.collect();
match doc_sets.len() {
0 => Box::new(AllScorer::new(max_doc)),
1 => doc_sets.pop().unwrap(),
_ => Box::new(Intersection::new(doc_sets, max_doc)),
}
}
VariablePresenceCondition::Any(conditions) => {
let doc_sets: Vec<Box<dyn DocSet>> = conditions
.iter()
.map(|condition| build_necessary_condition_docset(condition, columns, max_doc))
.collect();
Box::new(SimpleUnion::build(doc_sets))
}
}
}
fn open_input_column(
reader: &SegmentReader,
name: &str,
accepted_types: InferredTypeSet,
) -> io::Result<Option<DynamicColumn>> {
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<VarType> {
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<Option<DynamicColumn>>,
// One reusable buffer per string column.
string_inputs: Vec<String>,
// 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<VariableValue<'static>>,
}
/// A wrapper to make sure the variable value buffer is cleared even if the evaluation
/// panicked.
struct ClearOnDrop<'a>(&'a mut Vec<VariableValue<'a>>);
impl<'a> ClearOnDrop<'a> {
fn wrap(input_values: &'a mut Vec<VariableValue<'static>>) -> 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<VariableValue<'_>> =
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<bool> = unsafe { self.compiled.call(inputs_vec.0).as_bool() };
eval_result == Some(true)
}
}
fn fill_input_values<'buffer>(
columns: &[Option<DynamicColumn>],
string_inputs: &'buffer mut [String],
input_values: &mut Vec<VariableValue<'buffer>>,
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<VariableValue> = 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
// );
// }
}
+387 -55
View File
@@ -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<TSegmentDocPredicate> {
doc_predicate: TSegmentDocPredicate,
doc: DocId,
max_doc: DocId,
necessary_condition: Box<dyn DocSet>,
}
impl<TSegmentDocPredicate: SegmentDocPredicate> DocPredicateDocSet<TSegmentDocPredicate> {
/// Creates a `DocPredicateDocSet` positioned on its first matching document.
fn new(doc_predicate: TSegmentDocPredicate, necessary_condition: Box<dyn DocSet>) -> Self {
let first_candidate = necessary_condition.doc();
let mut doc_set = DocPredicateDocSet {
doc_predicate,
necessary_condition,
};
doc_set.find_match(first_candidate);
doc_set
}
/// Creates a `DocPredicateDocSet`, and seeks it to `target`, following
/// [`Weight::scorer_danger`]'s contract.
///
/// Documents before `target` are not evaluated.
fn new_seeked_to(
doc_predicate: TSegmentDocPredicate,
necessary_condition: Box<dyn DocSet>,
target: DocId,
) -> (SeekDangerResult, Self) {
let first_candidate = necessary_condition.doc();
let mut doc_set = DocPredicateDocSet {
doc_predicate,
necessary_condition,
};
if target >= TERMINATED {
if doc_set.necessary_condition.doc() < TERMINATED {
doc_set.necessary_condition.seek(TERMINATED);
}
return (SeekDangerResult::SeekLowerBound(TERMINATED), doc_set);
}
let seek_result = match first_candidate.cmp(&target) {
Ordering::Less => doc_set.seek_danger(target),
Ordering::Equal => doc_set.eval_candidate(target),
Ordering::Greater => SeekDangerResult::SeekLowerBound(first_candidate),
};
(seek_result, doc_set)
}
/// Evaluates the predicate on `candidate`.
///
/// Hidden contract: the necessary condition is positioned on `candidate`.
fn eval_candidate(&mut self, candidate: DocId) -> SeekDangerResult {
if self.doc_predicate.eval(candidate) {
SeekDangerResult::Found
} else {
SeekDangerResult::SeekLowerBound(candidate + 1)
}
}
/// Advances to the first matching document at or after `candidate`.
///
/// Hidden contract: the necessary condition is positioned on `candidate`.
fn find_match(&mut self, mut candidate: DocId) -> DocId {
debug_assert_eq!(candidate, self.necessary_condition.doc());
while candidate != TERMINATED && !self.doc_predicate.eval(candidate) {
candidate = self.necessary_condition.advance();
}
candidate
}
}
impl<TSegmentDocPredicate: SegmentDocPredicate> DocSet
for DocPredicateDocSet<TSegmentDocPredicate>
{
fn advance(&mut self) -> DocId {
if self.doc == TERMINATED {
if self.doc() == TERMINATED {
return TERMINATED;
}
self.find_match(self.doc + 1)
let candidate = self.necessary_condition.advance();
self.find_match(candidate)
}
fn seek(&mut self, target: DocId) -> DocId {
debug_assert!(target >= self.doc);
if self.doc == TERMINATED {
return TERMINATED;
let doc = self.doc();
debug_assert!(target >= doc);
// In a valid state, the current doc is a match (or TERMINATED).
if doc >= target {
return doc;
}
self.find_match(target)
let candidate = self.necessary_condition.seek(target);
self.find_match(candidate)
}
fn seek_danger(&mut self, target: DocId) -> SeekDangerResult {
if target >= self.max_doc {
self.doc = TERMINATED;
return SeekDangerResult::SeekLowerBound(TERMINATED);
}
if self.doc_predicate.eval(target) {
self.doc = target;
SeekDangerResult::Found
} else {
SeekDangerResult::SeekLowerBound(target + 1)
match self.necessary_condition.seek_danger(target) {
SeekDangerResult::Found => self.eval_candidate(target),
// Following `seek_danger`'s contract, we are now in an invalid state, and `doc()` may
// return anything until a subsequent `seek_danger` returns `Found`.
seek_lower_bound @ SeekDangerResult::SeekLowerBound(_) => seek_lower_bound,
}
}
fn doc(&self) -> DocId {
self.doc
self.necessary_condition.doc()
}
fn size_hint(&self) -> u32 {
self.max_doc
self.necessary_condition.size_hint()
}
}
impl<TSegmentDocPredicate: SegmentDocPredicate> DocPredicateDocSet<TSegmentDocPredicate> {
fn find_match(&mut self, mut target: DocId) -> DocId {
loop {
match self.seek_danger(target) {
SeekDangerResult::Found => return target,
SeekDangerResult::SeekLowerBound(next_target) => {
if next_target >= TERMINATED {
return TERMINATED;
}
target = next_target;
}
}
}
fn cost(&self) -> u64 {
// `cost` is the method used to tell how costly it is to consume a DocSet entirely.
//
// This is used in intersection to have cheaper docset "lead" the intersection.
//
// Here, we naturally use a model where we use the cost of the necessary condition
// multiplied by some factor expressing how slow it is to evaluate an expression.
self.necessary_condition.cost() * self.doc_predicate.cost()
}
}
@@ -141,14 +211,27 @@ pub trait DocPredicateBoxable: std::fmt::Debug + 'static + Send + Sync {
impl<TDocPredicate: DocPredicate> DocPredicateBoxable for TDocPredicate {
fn scorer(&self, segment_reader: &SegmentReader, boost: f32) -> crate::Result<Box<dyn Scorer>> {
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<dyn 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(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<TDocPredicate: DocPredicate> DocPredicateBoxable for TDocPredicate {
target: DocId,
boost: f32,
) -> crate::Result<(SeekDangerResult, Box<dyn Scorer>)> {
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<dyn Scorer>;
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<P: SegmentDocPredicate> {
/// Can be emitted to hint that a predicate will be always true or false on a segment.
/// Returning Const instead of a variable is an optimization.
Const(bool),
/// A regular SegmentDocPredicate, evaluated document by document.
Variable {
/// The predicate to evaluate.
predicate: P,
/// The [`DocSet`] of the documents on which `predicate` is evaluated.
///
/// Hidden contract: it must contain every document for which `predicate.eval` returns
/// true. Documents outside of it are never evaluated, and are considered as not
/// matching. Use an [`AllScorer`](crate::query::AllScorer) to evaluate every document of
/// the segment.
///
/// The `DocSet` must be positioned on its first document.
necessary_condition: Box<dyn DocSet>,
},
}
/// A per-query predicate that produces a [`SegmentDocPredicate`] for each
/// segment.
///
@@ -186,19 +299,37 @@ pub trait DocPredicate: Send + Sync + 'static + std::fmt::Debug {
fn doc_predicate(
&self,
segment_reader: &SegmentReader,
) -> crate::Result<Self::SegmentDocPredicate>;
) -> crate::Result<ConstOrVariableSegmentPredicate<Self::SegmentDocPredicate>>;
}
/// 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<AtomicUsize>,
}
impl SegmentDocPredicate for EvenDocIds {
fn eval(&mut self, doc_id: DocId) -> bool {
self.num_evals.fetch_add(1, AtomicOrdering::Relaxed);
doc_id.is_multiple_of(2)
}
}
/// `EvenDocIds`, with a fixed necessary condition.
#[derive(Debug)]
struct EvenWithNecessaryCondition {
necessary_condition: Vec<DocId>,
num_evals: Arc<AtomicUsize>,
}
impl EvenWithNecessaryCondition {
fn new(necessary_condition: Vec<DocId>) -> Self {
EvenWithNecessaryCondition {
necessary_condition,
num_evals: Arc::default(),
}
}
}
impl DocPredicate for EvenWithNecessaryCondition {
type SegmentDocPredicate = EvenDocIds;
fn doc_predicate(
&self,
_segment_reader: &SegmentReader,
) -> crate::Result<ConstOrVariableSegmentPredicate<EvenDocIds>> {
Ok(ConstOrVariableSegmentPredicate::Variable {
predicate: EvenDocIds {
num_evals: self.num_evals.clone(),
},
necessary_condition: Box::new(VecDocSet::from(self.necessary_condition.clone())),
})
}
}
#[test]
fn test_necessary_condition_restricts_evaluations() {
let index = create_index_for_test(10);
let searcher = index.reader().unwrap().searcher();
let predicate = EvenWithNecessaryCondition::new(vec![1, 2, 3, 4, 6, 9]);
let num_evals = predicate.num_evals.clone();
let query: DocPredicateQuery = predicate.into();
assert_eq!(searcher.search(&query, &DocSetCollector).unwrap().len(), 3);
assert_eq!(num_evals.load(AtomicOrdering::Relaxed), 6);
}
#[test]
fn test_necessary_condition_size_hint_and_cost() {
let index = create_index_for_test(10);
let searcher = index.reader().unwrap().searcher();
let query: DocPredicateQuery =
EvenWithNecessaryCondition::new(vec![1, 2, 3, 4, 6, 9]).into();
let weight = query
.weight(EnableScoring::disabled_from_searcher(&searcher))
.unwrap();
let scorer = weight.scorer(searcher.segment_reader(0), 1.0).unwrap();
assert_eq!(scorer.size_hint(), 6);
assert_eq!(scorer.cost(), 600);
// Without a necessary condition, all docs are candidates.
let scorer = even_doc_id_query()
.scorer(searcher.segment_reader(0), 1.0)
.unwrap();
assert_eq!(scorer.size_hint(), 10);
assert_eq!(scorer.cost(), 1000);
}
#[test]
fn test_necessary_condition_scorer_danger() {
let index = create_index_for_test(10);
let searcher = index.reader().unwrap().searcher();
let segment_reader = searcher.segment_reader(0);
let scorer_danger = |necessary_condition: Vec<DocId>, target: DocId| {
let predicate = EvenWithNecessaryCondition::new(necessary_condition);
let num_evals = predicate.num_evals.clone();
let query: DocPredicateQuery = predicate.into();
let (seek_result, scorer) = query.scorer_danger(segment_reader, target, 1.0).unwrap();
(seek_result, scorer, num_evals.load(AtomicOrdering::Relaxed))
};
// The necessary condition starts after the target.
let (seek_result, mut scorer, num_evals) = scorer_danger(vec![4, 6], 1);
assert_eq!(seek_result, SeekDangerResult::SeekLowerBound(4));
assert_eq!(num_evals, 0);
assert_eq!(scorer.seek_danger(4), SeekDangerResult::Found);
assert_eq!(scorer.doc(), 4);
// The target is the first candidate, and matches.
let (seek_result, scorer, _) = scorer_danger(vec![2, 6], 2);
assert_eq!(seek_result, SeekDangerResult::Found);
assert_eq!(scorer.doc(), 2);
// The target is a candidate, but does not match.
let (seek_result, mut scorer, num_evals) = scorer_danger(vec![1, 3, 4], 3);
assert_eq!(seek_result, SeekDangerResult::SeekLowerBound(4));
assert_eq!(num_evals, 1);
assert_eq!(scorer.seek_danger(4), SeekDangerResult::Found);
// The target is not a candidate: it is not evaluated.
let (seek_result, _, num_evals) = scorer_danger(vec![1, 6], 2);
assert_eq!(seek_result, SeekDangerResult::SeekLowerBound(6));
assert_eq!(num_evals, 0);
// No match after the target. The lower bound can stop on a non-matching candidate.
let (seek_result, mut scorer, _) = scorer_danger(vec![1, 3], 2);
assert_eq!(seek_result, SeekDangerResult::SeekLowerBound(3));
assert_eq!(scorer.seek_danger(3), SeekDangerResult::SeekLowerBound(4));
assert_eq!(
scorer.seek_danger(4),
SeekDangerResult::SeekLowerBound(TERMINATED)
);
}
proptest! {
#[test]
fn proptest_necessary_condition_doc_set(
candidates in prop::collection::btree_set(0u32..200, 0..60),
modulo in 1u32..5,
targets in prop::collection::vec(0u32..220, 0..30),
advances in prop::collection::vec(any::<bool>(), 0..30),
) {
let candidates: Vec<DocId> = candidates.into_iter().collect();
let expected: Vec<DocId> = candidates
.iter()
.copied()
.filter(|doc| doc.is_multiple_of(modulo))
.collect();
let new_doc_set = || {
DocPredicateDocSet::new(
move |doc: DocId| doc.is_multiple_of(modulo),
Box::new(VecDocSet::from(candidates.clone())),
)
};
let first_match = |target: DocId| {
expected
.iter()
.copied()
.find(|doc| *doc >= target)
.unwrap_or(TERMINATED)
};
// advance
let mut doc_set = new_doc_set();
let mut matches: Vec<DocId> = Vec::new();
while doc_set.doc() != TERMINATED {
matches.push(doc_set.doc());
doc_set.advance();
}
prop_assert_eq!(&matches, &expected);
// interleaved seek and advance
let mut doc_set = new_doc_set();
for (target, advance) in targets.iter().zip(advances.iter()) {
let target = (*target).max(doc_set.doc());
if *advance && doc_set.doc() != TERMINATED {
let current = doc_set.doc();
prop_assert_eq!(doc_set.advance(), first_match(current + 1));
} else {
prop_assert_eq!(doc_set.seek(target), first_match(target));
}
}
// seek_danger, following its contract: strictly increasing targets, respecting
// the returned lower bounds.
let mut sorted_targets = targets.clone();
sorted_targets.sort_unstable();
sorted_targets.dedup();
let mut doc_set = new_doc_set();
let mut lower_bound = doc_set.doc();
let mut previous_target = None;
for requested_target in sorted_targets {
let target = requested_target.max(lower_bound);
if previous_target.is_some_and(|previous| previous >= target) {
continue;
}
previous_target = Some(target);
let next_match = first_match(target);
match doc_set.seek_danger(target) {
SeekDangerResult::Found => {
prop_assert_eq!(next_match, target);
prop_assert_eq!(doc_set.doc(), target);
}
SeekDangerResult::SeekLowerBound(bound) => {
prop_assert!(next_match != target || target == TERMINATED);
prop_assert!(bound > target || target == TERMINATED);
prop_assert!(bound <= next_match);
lower_bound = bound;
}
}
}
}
}
#[test]
fn test_doc_predicate_query_scorer_danger_target_past_max_doc() {
let index = create_index_for_test(4);
+1 -1
View File
@@ -182,7 +182,7 @@ impl Weight for FastFieldExistsWeight {
}
}
enum ExistsColumnIndex {
pub(crate) enum ExistsColumnIndex {
Optional(OptionalIndex),
Multivalued(MultiValueIndex),
}
+51 -8
View File
@@ -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<AutomatonWeight<DfaWrapper>> {
/// 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<DfaWrapper> {
static AUTOMATON_BUILDER: [[OnceCell<LevenshteinAutomatonBuilder>; 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<AutomatonWeight<DfaWrapper>> {
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
+1 -3
View File
@@ -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;
+77 -1
View File
@@ -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<Vec<Arc<Regex>>>);
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<Regex>)>,
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::<Vec<Term>>()
}
/// 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<Regex>]> {
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::<crate::Result<Vec<_>>>()
})?;
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,
+69 -7
View File
@@ -20,7 +20,7 @@ type UnionType = SimpleUnion<Box<dyn Postings + 'static>>;
/// 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<Regex>)>,
similarity_weight_opt: Option<Bm25Weight>,
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<Regex>)>,
similarity_weight_opt: Option<Bm25Weight>,
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<Regex> =
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(&regexes[0], &regex_a));
assert!(Arc::ptr_eq(&regexes[1], &regex_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"])?;
+26
View File
@@ -1106,6 +1106,7 @@ fn convert_to_query(fuzzy: &FxHashMap<Field, Fuzzy>, 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,
+41 -2
View File
@@ -115,6 +115,7 @@ impl FragmentCandidate {
#[derive(Debug)]
pub struct Snippet {
fragment: String,
fragment_range: Range<usize>,
highlighted: Vec<Range<usize>>,
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<Range<usize>>) -> Self {
fn new(fragment: &str, fragment_range: Range<usize>, highlighted: Vec<Range<usize>>) -> 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<usize> {
self.fragment_range.clone()
}
/// Returns a list of highlighted positions from the `Snippet`.
pub fn highlighted(&self) -> &[Range<usize>] {
&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(), "<b>c</b> 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 <b>f</b>");
}
#[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());
}
+88 -1
View File
@@ -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<usize>;
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<MonotonicU64SSTable>,
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);
+80 -1
View File
@@ -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<W, TValueWriter>
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<U64MonotonicValueReader> = 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
);
}
}
+10 -31
View File
@@ -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<TSSTable: SSTable> Dictionary<TSSTable> {
) -> io::Result<TermOrdHit> {
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))
+1 -1
View File
@@ -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())

Some files were not shown because too many files have changed in this diff Show More