From 44a403d4468519e61a1900f528179fa3bbfe5e88 Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Wed, 12 Aug 2026 19:41:19 +0800 Subject: [PATCH] faster term agg merges --- src/aggregation/agg_tests.rs | 6 +- src/aggregation/bucket/filter.rs | 10 +- src/aggregation/bucket/term_agg/mod.rs | 28 +++-- src/aggregation/bucket/term_missing_agg.rs | 9 +- src/aggregation/intermediate_agg_result.rs | 137 ++++++++++++++++----- 5 files changed, 137 insertions(+), 53 deletions(-) diff --git a/src/aggregation/agg_tests.rs b/src/aggregation/agg_tests.rs index 1c735a9e6..cd818d0f4 100644 --- a/src/aggregation/agg_tests.rs +++ b/src/aggregation/agg_tests.rs @@ -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\"" ); diff --git a/src/aggregation/bucket/filter.rs b/src/aggregation/bucket/filter.rs index 9ee6116e1..656ac7795 100644 --- a/src/aggregation/bucket/filter.rs +++ b/src/aggregation/bucket/filter.rs @@ -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, diff --git a/src/aggregation/bucket/term_agg/mod.rs b/src/aggregation/bucket/term_agg/mod.rs index e9be561ab..e0183636e 100644 --- a/src/aggregation/bucket/term_agg/mod.rs +++ b/src/aggregation/bucket/term_agg/mod.rs @@ -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 = 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 = + 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() }, }) } diff --git a/src/aggregation/bucket/term_missing_agg.rs b/src/aggregation/bucket/term_missing_agg.rs index afdc84741..72dacfd07 100644 --- a/src/aggregation/bucket/term_missing_agg.rs +++ b/src/aggregation/bucket/term_missing_agg.rs @@ -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 = - 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() }, }; diff --git a/src/aggregation/intermediate_agg_result.rs b/src/aggregation/intermediate_agg_result.rs index 1b4101c9c..bd24e9a9f 100644 --- a/src/aggregation/intermediate_agg_result.rs +++ b/src/aggregation/intermediate_agg_result.rs @@ -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, } -#[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, + /// 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, +} + +// `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 { + /// 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]