@@ -32,6 +32,16 @@ use crate::aggregation::bucket::TermsAggregationInternal;
3232use crate :: aggregation:: metric:: CardinalityCollector ;
3333use 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