faster term agg merges

This commit is contained in:
Pascal Seitz
2026-09-07 18:26:37 +08:00
committed by PSeitz
parent 797543f3ca
commit 44a403d446
5 changed files with 137 additions and 53 deletions
+4 -2
View File
@@ -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\""
);
+5 -5
View File
@@ -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,
+21 -7
View File
@@ -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()
},
})
}
+3 -6
View File
@@ -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()
},
};
+104 -33
View File
@@ -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]