mirror of
https://github.com/quickwit-oss/tantivy.git
synced 2026-10-06 11:52:40 +00:00
faster term agg merges
This commit is contained in:
@@ -1605,7 +1605,8 @@ fn test_percentile_order_segment_level() -> crate::Result<()> {
|
||||
assert!(
|
||||
buckets
|
||||
.entries
|
||||
.contains_key(&IntermediateKey::Str("b".to_string())),
|
||||
.iter()
|
||||
.any(|(key, _)| key == &IntermediateKey::Str("b".to_string())),
|
||||
"\"b\" (higher p50) should survive, not \"a\""
|
||||
);
|
||||
assert!(
|
||||
@@ -1681,7 +1682,8 @@ fn test_percentile_order_prune_intermediate() -> crate::Result<()> {
|
||||
assert!(
|
||||
buckets
|
||||
.entries
|
||||
.contains_key(&IntermediateKey::Str("b".to_string())),
|
||||
.iter()
|
||||
.any(|(key, _)| key == &IntermediateKey::Str("b".to_string())),
|
||||
"\"b\" (higher p50) should survive, not \"a\""
|
||||
);
|
||||
|
||||
|
||||
@@ -1334,15 +1334,15 @@ mod tests {
|
||||
"doc_count": 2,
|
||||
"brands": {
|
||||
"buckets": [
|
||||
{
|
||||
"key": "samsung",
|
||||
"doc_count": 1,
|
||||
"avg_price": { "value": 799.0 }
|
||||
},
|
||||
{
|
||||
"key": "apple",
|
||||
"doc_count": 1,
|
||||
"avg_price": { "value": 999.0 }
|
||||
},
|
||||
{
|
||||
"key": "samsung",
|
||||
"doc_count": 1,
|
||||
"avg_price": { "value": 799.0 }
|
||||
}
|
||||
],
|
||||
"sum_other_doc_count": 0,
|
||||
|
||||
@@ -1325,8 +1325,11 @@ where
|
||||
let (term_doc_count_before_cutoff, sum_other_doc_count) =
|
||||
cut_off_buckets(&mut entries, segment_size, total_doc_count);
|
||||
|
||||
let mut dict: FxHashMap<IntermediateKey, IntermediateTermBucketEntry> = Default::default();
|
||||
dict.reserve(entries.len());
|
||||
// Collected entries, in unspecified order. The cross-segment merge keys entries by their
|
||||
// own dictionary and the final result is sorted by the requested order, so segment-side
|
||||
// order does not matter here.
|
||||
let mut out: Vec<(IntermediateKey, IntermediateTermBucketEntry)> =
|
||||
Vec::with_capacity(entries.len());
|
||||
|
||||
if term_req.column_type == ColumnType::Str {
|
||||
let fallback_dict = Dictionary::empty();
|
||||
@@ -1336,6 +1339,14 @@ where
|
||||
.map(|el| el.dictionary())
|
||||
.unwrap_or_else(|| &fallback_dict);
|
||||
|
||||
// Collect into a map to dedup by key, then flush into `out`. Two cases need it: a real
|
||||
// term may equal the `missing` placeholder, and the min_doc_count==0 fill must skip
|
||||
// already-collected terms. A single-segment query returns this result directly (the
|
||||
// cross-segment merge that would otherwise dedup never runs), so duplicate keys here
|
||||
// would reach the final result unmerged.
|
||||
let mut dict: FxHashMap<IntermediateKey, IntermediateTermBucketEntry> =
|
||||
FxHashMap::with_capacity_and_hasher(entries.len(), Default::default());
|
||||
|
||||
if let Some((intermediate_key, bucket)) = extract_missing_value(&mut entries, term_req)
|
||||
{
|
||||
let intermediate_entry = into_intermediate_bucket_entry(
|
||||
@@ -1405,6 +1416,8 @@ where
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
out.extend(dict);
|
||||
} else if term_req.column_type == ColumnType::DateTime {
|
||||
for (val, doc_count) in entries {
|
||||
let intermediate_entry = into_intermediate_bucket_entry(
|
||||
@@ -1414,7 +1427,7 @@ where
|
||||
)?;
|
||||
let val = i64::from_u64(val);
|
||||
let date = format_date(val)?;
|
||||
dict.insert(IntermediateKey::Str(date), intermediate_entry);
|
||||
out.push((IntermediateKey::Str(date), intermediate_entry));
|
||||
}
|
||||
} else if term_req.column_type == ColumnType::Bool {
|
||||
for (val, doc_count) in entries {
|
||||
@@ -1424,7 +1437,7 @@ where
|
||||
agg_data,
|
||||
)?;
|
||||
let val = bool::from_u64(val);
|
||||
dict.insert(IntermediateKey::Bool(val), intermediate_entry);
|
||||
out.push((IntermediateKey::Bool(val), intermediate_entry));
|
||||
}
|
||||
} else if term_req.column_type == ColumnType::IpAddr {
|
||||
let compact_space_accessor = term_req
|
||||
@@ -1449,7 +1462,7 @@ where
|
||||
)?;
|
||||
let val: u128 = compact_space_accessor.compact_to_u128(val as u32);
|
||||
let val = Ipv6Addr::from_u128(val);
|
||||
dict.insert(IntermediateKey::IpAddr(val), intermediate_entry);
|
||||
out.push((IntermediateKey::IpAddr(val), intermediate_entry));
|
||||
}
|
||||
} else {
|
||||
for (key_val_u64, doc_count) in entries {
|
||||
@@ -1470,15 +1483,16 @@ where
|
||||
}
|
||||
};
|
||||
let key = IntermediateKey::from(key_val.normalize());
|
||||
dict.insert(key, intermediate_entry);
|
||||
out.push((key, intermediate_entry));
|
||||
}
|
||||
};
|
||||
|
||||
Ok(IntermediateBucketResult::Terms {
|
||||
buckets: IntermediateTermBucketResult {
|
||||
entries: dict,
|
||||
entries: out,
|
||||
sum_other_doc_count,
|
||||
doc_count_error_upper_bound: term_doc_count_before_cutoff,
|
||||
..Default::default()
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
use columnar::{Column, ColumnType};
|
||||
use rustc_hash::FxHashMap;
|
||||
|
||||
use crate::aggregation::agg_data::{
|
||||
build_segment_agg_collectors, AggRefNode, AggregationsSegmentCtx,
|
||||
@@ -8,7 +7,7 @@ use crate::aggregation::bucket::term_agg::TermsAggregation;
|
||||
use crate::aggregation::buffered_sub_aggs::{BufferedSubAggs, HighCardBufferedSubAggs};
|
||||
use crate::aggregation::intermediate_agg_result::{
|
||||
IntermediateAggregationResult, IntermediateAggregationResults, IntermediateBucketResult,
|
||||
IntermediateKey, IntermediateTermBucketEntry, IntermediateTermBucketResult,
|
||||
IntermediateTermBucketEntry, IntermediateTermBucketResult,
|
||||
};
|
||||
use crate::aggregation::segment_agg_result::{BucketIdProvider, SegmentAggregationCollector};
|
||||
use crate::aggregation::BucketId;
|
||||
@@ -93,9 +92,6 @@ impl SegmentAggregationCollector for TermMissingAgg {
|
||||
.as_ref()
|
||||
.expect("TermMissingAgg collector, but no missing found in agg req")
|
||||
.clone();
|
||||
let mut entries: FxHashMap<IntermediateKey, IntermediateTermBucketEntry> =
|
||||
Default::default();
|
||||
|
||||
let missing_count = &self.missing_count_per_bucket[parent_bucket_id as usize];
|
||||
let mut missing_entry = IntermediateTermBucketEntry {
|
||||
doc_count: missing_count.missing_count as u64,
|
||||
@@ -108,13 +104,14 @@ impl SegmentAggregationCollector for TermMissingAgg {
|
||||
.add_intermediate_aggregation_result(agg_data, &mut res, missing_count.bucket_id)?;
|
||||
missing_entry.sub_aggregation = res;
|
||||
}
|
||||
entries.insert(missing.into(), missing_entry);
|
||||
let entries = vec![(missing.into(), missing_entry)];
|
||||
|
||||
let bucket = IntermediateBucketResult::Terms {
|
||||
buckets: IntermediateTermBucketResult {
|
||||
entries,
|
||||
sum_other_doc_count: 0,
|
||||
doc_count_error_upper_bound: 0,
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
|
||||
|
||||
@@ -787,10 +787,7 @@ impl IntermediateBucketResult {
|
||||
buckets: term_res_right,
|
||||
},
|
||||
) => {
|
||||
merge_maps(&mut term_res_left.entries, term_res_right.entries)?;
|
||||
term_res_left.sum_other_doc_count += term_res_right.sum_other_doc_count;
|
||||
term_res_left.doc_count_error_upper_bound +=
|
||||
term_res_right.doc_count_error_upper_bound;
|
||||
term_res_left.merge(term_res_right)?;
|
||||
}
|
||||
|
||||
(
|
||||
@@ -888,17 +885,71 @@ pub struct IntermediateRangeBucketResult {
|
||||
pub(crate) column_type: Option<ColumnType>,
|
||||
}
|
||||
|
||||
#[derive(Default, Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
#[derive(Default, Clone, Debug, Serialize, Deserialize)]
|
||||
/// Term aggregation including error counts
|
||||
pub struct IntermediateTermBucketResult {
|
||||
pub(crate) entries: FxHashMap<IntermediateKey, IntermediateTermBucketEntry>,
|
||||
/// Bucket entries as a flat `Vec`, in unspecified order. The final result is sorted by the
|
||||
/// requested order (see `into_final_result`).
|
||||
pub(crate) entries: Vec<(IntermediateKey, IntermediateTermBucketEntry)>,
|
||||
pub(crate) sum_other_doc_count: u64,
|
||||
pub(crate) doc_count_error_upper_bound: u64,
|
||||
/// Transient `key -> position in entries` index, populated lazily on the first merge so that
|
||||
/// folding many segment results stays O(total entries) instead of rebuilding a dictionary on
|
||||
/// every pairwise merge. Not serialized; rebuilt on demand.
|
||||
#[serde(skip)]
|
||||
pub(crate) key_to_idx_in_entries: FxHashMap<IntermediateKey, u32>,
|
||||
}
|
||||
|
||||
// `key_to_idx_in_entries` is a transient acceleration structure and is excluded from equality.
|
||||
impl PartialEq for IntermediateTermBucketResult {
|
||||
fn eq(&self, other: &Self) -> bool {
|
||||
self.entries == other.entries
|
||||
&& self.sum_other_doc_count == other.sum_other_doc_count
|
||||
&& self.doc_count_error_upper_bound == other.doc_count_error_upper_bound
|
||||
}
|
||||
}
|
||||
|
||||
impl IntermediateTermBucketResult {
|
||||
/// Returns a reference to the map of bucket entries keyed by [`IntermediateKey`].
|
||||
pub fn entries(&self) -> &FxHashMap<IntermediateKey, IntermediateTermBucketEntry> {
|
||||
/// Merge `other` into `self` in place.
|
||||
///
|
||||
/// The first merge builds `self.key_to_idx_in_entries` (`key -> position`) once; subsequent
|
||||
/// merges reuse it, so
|
||||
/// each entry is hashed once and the accumulator grows incrementally. This keeps both the
|
||||
/// per-segment fold and the cross-node `merge_fruits` fold linear in the total number of
|
||||
/// entries. Output order is first-seen: existing keys merge in place, new keys are appended.
|
||||
fn merge(&mut self, other: IntermediateTermBucketResult) -> crate::Result<()> {
|
||||
self.sum_other_doc_count += other.sum_other_doc_count;
|
||||
self.doc_count_error_upper_bound += other.doc_count_error_upper_bound;
|
||||
if other.entries.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
if self.entries.is_empty() {
|
||||
self.entries = other.entries;
|
||||
return Ok(());
|
||||
}
|
||||
if self.key_to_idx_in_entries.is_empty() {
|
||||
self.key_to_idx_in_entries
|
||||
.reserve(self.entries.len() + other.entries.len());
|
||||
for (idx, (key, _)) in self.entries.iter().enumerate() {
|
||||
self.key_to_idx_in_entries.insert(key.clone(), idx as u32);
|
||||
}
|
||||
}
|
||||
for (key, bucket) in other.entries {
|
||||
if let Some(&idx_in_entries) = self.key_to_idx_in_entries.get(&key) {
|
||||
let entry = &mut self.entries[idx_in_entries as usize].1;
|
||||
entry.doc_count += bucket.doc_count;
|
||||
entry.sub_aggregation.merge_fruits(bucket.sub_aggregation)?;
|
||||
} else {
|
||||
self.key_to_idx_in_entries
|
||||
.insert(key.clone(), self.entries.len() as u32);
|
||||
self.entries.push((key, bucket));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Returns the bucket entries. Order is unspecified.
|
||||
pub fn entries(&self) -> &[(IntermediateKey, IntermediateTermBucketEntry)] {
|
||||
&self.entries
|
||||
}
|
||||
|
||||
@@ -955,10 +1006,26 @@ impl IntermediateTermBucketResult {
|
||||
});
|
||||
}
|
||||
OrderTarget::Count => {
|
||||
// Tie-break equal counts by key ascending so the output is deterministic
|
||||
// regardless of the merge's (insertion-order) output.
|
||||
let key_tie = |left: &BucketEntry, right: &BucketEntry| {
|
||||
left.key
|
||||
.partial_cmp(&right.key)
|
||||
.expect("expected type string, which is always sortable")
|
||||
};
|
||||
if req.order.order == Order::Desc {
|
||||
buckets.sort_unstable_by_key(|bucket| std::cmp::Reverse(bucket.doc_count()));
|
||||
buckets.sort_unstable_by(|left, right| {
|
||||
right
|
||||
.doc_count()
|
||||
.cmp(&left.doc_count())
|
||||
.then_with(|| key_tie(left, right))
|
||||
});
|
||||
} else {
|
||||
buckets.sort_unstable_by_key(|bucket| bucket.doc_count());
|
||||
buckets.sort_unstable_by(|left, right| {
|
||||
left.doc_count()
|
||||
.cmp(&right.doc_count())
|
||||
.then_with(|| key_tie(left, right))
|
||||
});
|
||||
}
|
||||
}
|
||||
OrderTarget::SubAggregation(name) => {
|
||||
@@ -1013,17 +1080,19 @@ impl IntermediateTermBucketResult {
|
||||
mode: PruneMode,
|
||||
) -> crate::Result<()> {
|
||||
let req_internal = TermsAggregationInternal::from_req(req);
|
||||
// Pruning changes entry positions, invalidating the transient merge index.
|
||||
self.key_to_idx_in_entries.clear();
|
||||
let size = if mode == PruneMode::Final {
|
||||
let min_doc_count = req_internal.min_doc_count;
|
||||
self.entries.retain(|_, e| e.doc_count >= min_doc_count);
|
||||
self.entries
|
||||
.retain(|(_, entry)| entry.doc_count >= min_doc_count);
|
||||
req_internal.size as usize
|
||||
} else {
|
||||
req_internal.segment_size as usize
|
||||
};
|
||||
|
||||
if self.entries.len() > size {
|
||||
let mut entries: Vec<(IntermediateKey, IntermediateTermBucketEntry)> =
|
||||
self.entries.drain().collect();
|
||||
let mut entries = std::mem::take(&mut self.entries);
|
||||
|
||||
match &req_internal.order.target {
|
||||
OrderTarget::SubAggregation(sub_agg_path) => {
|
||||
@@ -1087,10 +1156,10 @@ impl IntermediateTermBucketResult {
|
||||
self.doc_count_error_upper_bound += cutoff_doc_count;
|
||||
}
|
||||
entries.truncate(size);
|
||||
self.entries = entries.into_iter().collect();
|
||||
self.entries = entries;
|
||||
}
|
||||
|
||||
for entry in self.entries.values_mut() {
|
||||
for (_, entry) in &mut self.entries {
|
||||
entry
|
||||
.sub_aggregation
|
||||
.prune_intermediate_results(sub_aggregation_req, mode)?;
|
||||
@@ -1533,9 +1602,8 @@ mod tests {
|
||||
);
|
||||
}
|
||||
let mut term_result = IntermediateTermBucketResult {
|
||||
entries: buckets,
|
||||
sum_other_doc_count: 0,
|
||||
doc_count_error_upper_bound: 0,
|
||||
entries: buckets.into_iter().collect(),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let req: TermsAggregation =
|
||||
@@ -1548,10 +1616,12 @@ mod tests {
|
||||
assert_eq!(term_result.entries.len(), 2);
|
||||
assert!(term_result
|
||||
.entries
|
||||
.contains_key(&IntermediateKey::Str("c".to_string())));
|
||||
.iter()
|
||||
.any(|(key, _)| key == &IntermediateKey::Str("c".to_string())));
|
||||
assert!(term_result
|
||||
.entries
|
||||
.contains_key(&IntermediateKey::Str("e".to_string())));
|
||||
.iter()
|
||||
.any(|(key, _)| key == &IntermediateKey::Str("e".to_string())));
|
||||
assert_eq!(term_result.sum_other_doc_count, 10 + 5 + 1);
|
||||
// final-size cutoff doesn't contribute to error bound
|
||||
assert_eq!(term_result.doc_count_error_upper_bound, 0);
|
||||
@@ -1573,9 +1643,8 @@ mod tests {
|
||||
);
|
||||
}
|
||||
let mut term_result = IntermediateTermBucketResult {
|
||||
entries: buckets,
|
||||
sum_other_doc_count: 0,
|
||||
doc_count_error_upper_bound: 0,
|
||||
entries: buckets.into_iter().collect(),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let req: TermsAggregation =
|
||||
@@ -1588,7 +1657,8 @@ mod tests {
|
||||
assert_eq!(term_result.entries.len(), 4);
|
||||
assert!(!term_result
|
||||
.entries
|
||||
.contains_key(&IntermediateKey::Str("d".to_string())));
|
||||
.iter()
|
||||
.any(|(key, _)| key == &IntermediateKey::Str("d".to_string())));
|
||||
assert_eq!(term_result.sum_other_doc_count, 1);
|
||||
assert_eq!(term_result.doc_count_error_upper_bound, 1);
|
||||
}
|
||||
@@ -1611,9 +1681,8 @@ mod tests {
|
||||
"my_terms".to_string(),
|
||||
IntermediateAggregationResult::Bucket(IntermediateBucketResult::Terms {
|
||||
buckets: IntermediateTermBucketResult {
|
||||
entries: buckets,
|
||||
sum_other_doc_count: 0,
|
||||
doc_count_error_upper_bound: 0,
|
||||
entries: buckets.into_iter().collect(),
|
||||
..Default::default()
|
||||
},
|
||||
}),
|
||||
);
|
||||
@@ -1634,7 +1703,8 @@ mod tests {
|
||||
assert_eq!(buckets.entries.len(), 1);
|
||||
assert!(buckets
|
||||
.entries
|
||||
.contains_key(&IntermediateKey::Str("x".to_string())));
|
||||
.iter()
|
||||
.any(|(key, _)| key == &IntermediateKey::Str("x".to_string())));
|
||||
assert_eq!(buckets.sum_other_doc_count, 60); // y(50) + z(10)
|
||||
}
|
||||
|
||||
@@ -1654,9 +1724,8 @@ mod tests {
|
||||
);
|
||||
}
|
||||
let mut term_result = IntermediateTermBucketResult {
|
||||
entries: buckets,
|
||||
sum_other_doc_count: 0,
|
||||
doc_count_error_upper_bound: 0,
|
||||
entries: buckets.into_iter().collect(),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let req: TermsAggregation =
|
||||
@@ -1670,10 +1739,12 @@ mod tests {
|
||||
assert_eq!(term_result.entries.len(), 2);
|
||||
assert!(term_result
|
||||
.entries
|
||||
.contains_key(&IntermediateKey::Str("a".to_string())));
|
||||
.iter()
|
||||
.any(|(key, _)| key == &IntermediateKey::Str("a".to_string())));
|
||||
assert!(term_result
|
||||
.entries
|
||||
.contains_key(&IntermediateKey::Str("b".to_string())));
|
||||
.iter()
|
||||
.any(|(key, _)| key == &IntermediateKey::Str("b".to_string())));
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
Reference in New Issue
Block a user