mirror of
https://github.com/quickwit-oss/tantivy.git
synced 2026-10-06 03:42:36 +00:00
Merge remote-tracking branch 'origin/main' into mallets/seqnum-field
This commit is contained in:
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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! {
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
@@ -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[..]);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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");
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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)?),
|
||||
],
|
||||
)?;
|
||||
|
||||
|
||||
@@ -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
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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
@@ -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());
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
@@ -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"),
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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) => {
|
||||
|
||||
@@ -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>),
|
||||
}
|
||||
|
||||
|
||||
@@ -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),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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));
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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() {
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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 },
|
||||
}
|
||||
|
||||
@@ -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
@@ -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 {
|
||||
|
||||
@@ -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::{
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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]);
|
||||
}
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+142
-128
@@ -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);
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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]);
|
||||
}
|
||||
}
|
||||
@@ -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()?;
|
||||
}
|
||||
|
||||
@@ -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>;
|
||||
|
||||
@@ -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
@@ -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
@@ -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]);
|
||||
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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};
|
||||
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
|
||||
@@ -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
|
||||
// );
|
||||
// }
|
||||
}
|
||||
@@ -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);
|
||||
|
||||
@@ -182,7 +182,7 @@ impl Weight for FastFieldExistsWeight {
|
||||
}
|
||||
}
|
||||
|
||||
enum ExistsColumnIndex {
|
||||
pub(crate) enum ExistsColumnIndex {
|
||||
Optional(OptionalIndex),
|
||||
Multivalued(MultiValueIndex),
|
||||
}
|
||||
|
||||
@@ -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
@@ -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;
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(®exes[0], ®ex_a));
|
||||
assert!(Arc::ptr_eq(®exes[1], ®ex_b));
|
||||
|
||||
let from_patterns = RegexPhraseQuery::new(text_field, vec!["a.*".into(), "b".into()]);
|
||||
assert_eq!(searcher.search(&from_regexes, &Count)?, 2);
|
||||
assert_eq!(searcher.search(&from_patterns, &Count)?, 2);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
pub fn test_phrase_count() -> crate::Result<()> {
|
||||
let index = create_index(&["a c", "a a b d a b c", " a b"])?;
|
||||
|
||||
@@ -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
@@ -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());
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
Reference in New Issue
Block a user