Skip to content

Commit 751b81c

Browse files
committed
use enum PruneMode instead of bool
1 parent 886322f commit 751b81c

2 files changed

Lines changed: 39 additions & 39 deletions

File tree

src/aggregation/agg_tests.rs

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@ use crate::aggregation::agg_result::AggregationResults;
55
use crate::aggregation::collector::AggregationCollector;
66
use crate::aggregation::intermediate_agg_result::{
77
IntermediateAggregationResult, IntermediateAggregationResults, IntermediateBucketResult,
8-
IntermediateKey,
8+
IntermediateKey, PruneMode,
99
};
1010
use crate::aggregation::tests::{get_test_index_2_segments, get_test_index_from_values_and_terms};
1111
use crate::aggregation::DistributedAggregationCollector;
@@ -1590,7 +1590,7 @@ fn test_percentile_order_prune_intermediate() -> crate::Result<()> {
15901590
);
15911591
}
15921592

1593-
intermediate.prune_intermediate_results(&agg_req, false)?;
1593+
intermediate.prune_intermediate_results(&agg_req, PruneMode::Final)?;
15941594

15951595
let IntermediateAggregationResult::Bucket(IntermediateBucketResult::Terms { buckets }) =
15961596
intermediate.aggs_res.get("my_terms").unwrap()

src/aggregation/intermediate_agg_result.rs

Lines changed: 37 additions & 37 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,16 @@ use crate::aggregation::bucket::TermsAggregationInternal;
3232
use crate::aggregation::metric::CardinalityCollector;
3333
use crate::TantivyError;
3434

35+
/// Controls which size limit is applied when pruning intermediate aggregation results.
36+
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
37+
pub enum PruneMode {
38+
/// Use the same rules for pruning as the per-segment pruning, notably using `segment_size`.
39+
Intermediate,
40+
/// Use the same rules for pruning as what happen when creating normal results.
41+
/// Uses `size`, and possibly apply other filtering such as `min_doc_count`.
42+
Final,
43+
}
44+
3545
/// Contains the intermediate aggregation result, which is optimized to be merged with other
3646
/// intermediate results.
3747
///
@@ -224,17 +234,14 @@ impl IntermediateAggregationResults {
224234
}
225235

226236
/// Re-prune intermediate results using the limits from the aggregation request.
227-
///
228-
/// `use_segment_size` controls which size limit is applied: `true` uses `segment_size`,
229-
/// `false` uses `size`.
230237
pub fn prune_intermediate_results(
231238
&mut self,
232239
req: &Aggregations,
233-
use_segment_size: bool,
240+
mode: PruneMode,
234241
) -> crate::Result<()> {
235242
for (key, agg_res) in self.aggs_res.iter_mut() {
236243
if let Some(agg_req) = req.get(key.as_str()) {
237-
agg_res.prune_intermediate_results(agg_req, use_segment_size)?;
244+
agg_res.prune_intermediate_results(agg_req, mode)?;
238245
}
239246
}
240247
Ok(())
@@ -378,11 +385,11 @@ impl IntermediateAggregationResult {
378385
pub(crate) fn prune_intermediate_results(
379386
&mut self,
380387
req: &Aggregation,
381-
use_segment_size: bool,
388+
mode: PruneMode,
382389
) -> crate::Result<()> {
383390
match self {
384391
IntermediateAggregationResult::Bucket(bucket) => {
385-
bucket.prune_intermediate_results(req, use_segment_size)
392+
bucket.prune_intermediate_results(req, mode)
386393
}
387394
IntermediateAggregationResult::Metric(_) => Ok(()),
388395
}
@@ -692,49 +699,43 @@ impl IntermediateBucketResult {
692699
pub(crate) fn prune_intermediate_results(
693700
&mut self,
694701
req: &Aggregation,
695-
use_segment_size: bool,
702+
mode: PruneMode,
696703
) -> crate::Result<()> {
697704
match self {
698705
IntermediateBucketResult::Terms { buckets } => {
699706
let terms_req = req
700707
.agg
701708
.as_term()
702709
.expect("unexpected aggregation, expected term aggregation");
703-
buckets.prune_intermediate_results(
704-
terms_req,
705-
req.sub_aggregation(),
706-
use_segment_size,
707-
)
710+
buckets.prune_intermediate_results(terms_req, req.sub_aggregation(), mode)
708711
}
709712
IntermediateBucketResult::Range(range_res) => {
710713
for entry in range_res.buckets.values_mut() {
711714
entry
712715
.sub_aggregation_res
713-
.prune_intermediate_results(req.sub_aggregation(), use_segment_size)?;
716+
.prune_intermediate_results(req.sub_aggregation(), mode)?;
714717
}
715718
Ok(())
716719
}
717720
IntermediateBucketResult::Histogram { buckets, .. } => {
718721
for entry in buckets.iter_mut() {
719722
entry
720723
.sub_aggregation
721-
.prune_intermediate_results(req.sub_aggregation(), use_segment_size)?;
724+
.prune_intermediate_results(req.sub_aggregation(), mode)?;
722725
}
723726
Ok(())
724727
}
725728
IntermediateBucketResult::Filter {
726729
sub_aggregations, ..
727-
} => {
728-
sub_aggregations.prune_intermediate_results(req.sub_aggregation(), use_segment_size)
729-
}
730+
} => sub_aggregations.prune_intermediate_results(req.sub_aggregation(), mode),
730731
IntermediateBucketResult::Composite { buckets } => {
731-
if !use_segment_size {
732+
if mode == PruneMode::Final {
732733
buckets.trim()?;
733734
}
734735
for entry in buckets.entries.values_mut() {
735736
entry
736737
.sub_aggregation
737-
.prune_intermediate_results(req.sub_aggregation(), use_segment_size)?;
738+
.prune_intermediate_results(req.sub_aggregation(), mode)?;
738739
}
739740
Ok(())
740741
}
@@ -963,19 +964,16 @@ impl IntermediateTermBucketResult {
963964
&mut self,
964965
req: &TermsAggregation,
965966
sub_aggregation_req: &Aggregations,
966-
use_segment_size: bool,
967+
mode: PruneMode,
967968
) -> crate::Result<()> {
968969
let req_internal = TermsAggregationInternal::from_req(req);
969-
let size = if use_segment_size {
970-
req_internal.segment_size as usize
971-
} else {
972-
req_internal.size as usize
973-
};
974-
975-
if !use_segment_size {
970+
let size = if mode == PruneMode::Final {
976971
let min_doc_count = req_internal.min_doc_count;
977972
self.entries.retain(|_, e| e.doc_count >= min_doc_count);
978-
}
973+
req_internal.size as usize
974+
} else {
975+
req_internal.segment_size as usize
976+
};
979977

980978
if self.entries.len() > size {
981979
let mut entries: Vec<(IntermediateKey, IntermediateTermBucketEntry)> =
@@ -1039,7 +1037,7 @@ impl IntermediateTermBucketResult {
10391037
.iter()
10401038
.map(|(_, e)| e.doc_count)
10411039
.sum::<u64>();
1042-
if use_segment_size {
1040+
if mode == PruneMode::Intermediate {
10431041
self.doc_count_error_upper_bound += cutoff_doc_count;
10441042
}
10451043
entries.truncate(size);
@@ -1049,7 +1047,7 @@ impl IntermediateTermBucketResult {
10491047
for entry in self.entries.values_mut() {
10501048
entry
10511049
.sub_aggregation
1052-
.prune_intermediate_results(sub_aggregation_req, use_segment_size)?;
1050+
.prune_intermediate_results(sub_aggregation_req, mode)?;
10531051
}
10541052

10551053
Ok(())
@@ -1497,9 +1495,9 @@ mod tests {
14971495
let req: TermsAggregation =
14981496
serde_json::from_str(r#"{"field": "myfield", "size": 2, "segment_size": 4}"#).unwrap();
14991497

1500-
// use_segment_size=false → keep top 2 by count: c(20), e(15); prune a(10), b(5), d(1)
1498+
// Final mode, keep top 2 by count: c(20), e(15); prune a(10), b(5), d(1)
15011499
term_result
1502-
.prune_intermediate_results(&req, &Default::default(), false)
1500+
.prune_intermediate_results(&req, &Default::default(), PruneMode::Final)
15031501
.unwrap();
15041502
assert_eq!(term_result.entries.len(), 2);
15051503
assert!(term_result
@@ -1537,9 +1535,9 @@ mod tests {
15371535
let req: TermsAggregation =
15381536
serde_json::from_str(r#"{"field": "myfield", "size": 2, "segment_size": 4}"#).unwrap();
15391537

1540-
// use_segment_size=true → keep top 4 by count: c(20), e(15), a(10), b(5); prune d(1)
1538+
// Intermediate mode, keep top 4 by count: c(20), e(15), a(10), b(5); prune d(1)
15411539
term_result
1542-
.prune_intermediate_results(&req, &Default::default(), true)
1540+
.prune_intermediate_results(&req, &Default::default(), PruneMode::Intermediate)
15431541
.unwrap();
15441542
assert_eq!(term_result.entries.len(), 4);
15451543
assert!(!term_result
@@ -1578,7 +1576,9 @@ mod tests {
15781576
serde_json::from_str(r#"{"my_terms": {"terms": {"field": "myfield", "size": 1}}}"#)
15791577
.unwrap();
15801578

1581-
results.prune_intermediate_results(&req, false).unwrap();
1579+
results
1580+
.prune_intermediate_results(&req, PruneMode::Final)
1581+
.unwrap();
15821582

15831583
let IntermediateAggregationResult::Bucket(IntermediateBucketResult::Terms { buckets }) =
15841584
results.aggs_res.get("my_terms").unwrap()
@@ -1619,7 +1619,7 @@ mod tests {
16191619

16201620
// asc key order, size=2 → keep "a" and "b"
16211621
term_result
1622-
.prune_intermediate_results(&req, &Default::default(), false)
1622+
.prune_intermediate_results(&req, &Default::default(), PruneMode::Final)
16231623
.unwrap();
16241624
assert_eq!(term_result.entries.len(), 2);
16251625
assert!(term_result

0 commit comments

Comments
 (0)