@@ -2442,4 +2442,185 @@ mod tests {
24422442
24432443 Ok ( ( ) )
24442444 }
2445+
2446+ #[ tokio:: test]
2447+ async fn test_skip_aggregation_after_first_batch ( ) -> Result < ( ) > {
2448+ let schema = Arc :: new ( Schema :: new ( vec ! [
2449+ Field :: new( "key" , DataType :: Int32 , true ) ,
2450+ Field :: new( "val" , DataType :: Int32 , true ) ,
2451+ ] ) ) ;
2452+
2453+ let group_by =
2454+ PhysicalGroupBy :: new_single ( vec ! [ ( col( "key" , & schema) ?, "key" . to_string( ) ) ] ) ;
2455+
2456+ let aggr_expr: Vec < Arc < dyn AggregateExpr > > = vec ! [ create_aggregate_expr(
2457+ & count_udaf( ) ,
2458+ & [ col( "val" , & schema) ?] ,
2459+ & [ datafusion_expr:: col( "val" ) ] ,
2460+ & [ ] ,
2461+ & [ ] ,
2462+ & schema,
2463+ "COUNT(val)" ,
2464+ false ,
2465+ false ,
2466+ false ,
2467+ ) ?] ;
2468+
2469+ let input_data = vec ! [
2470+ RecordBatch :: try_new(
2471+ Arc :: clone( & schema) ,
2472+ vec![
2473+ Arc :: new( Int32Array :: from( vec![ 1 , 2 , 3 ] ) ) ,
2474+ Arc :: new( Int32Array :: from( vec![ 0 , 0 , 0 ] ) ) ,
2475+ ] ,
2476+ )
2477+ . unwrap( ) ,
2478+ RecordBatch :: try_new(
2479+ Arc :: clone( & schema) ,
2480+ vec![
2481+ Arc :: new( Int32Array :: from( vec![ 2 , 3 , 4 ] ) ) ,
2482+ Arc :: new( Int32Array :: from( vec![ 0 , 0 , 0 ] ) ) ,
2483+ ] ,
2484+ )
2485+ . unwrap( ) ,
2486+ ] ;
2487+
2488+ let input = Arc :: new ( MemoryExec :: try_new (
2489+ & [ input_data] ,
2490+ Arc :: clone ( & schema) ,
2491+ None ,
2492+ ) ?) ;
2493+ let aggregate_exec = Arc :: new ( AggregateExec :: try_new (
2494+ AggregateMode :: Partial ,
2495+ group_by,
2496+ aggr_expr,
2497+ vec ! [ None ] ,
2498+ Arc :: clone ( & input) as Arc < dyn ExecutionPlan > ,
2499+ schema,
2500+ ) ?) ;
2501+
2502+ let mut session_config = SessionConfig :: default ( ) ;
2503+ session_config = session_config. set (
2504+ "datafusion.execution.skip_partial_aggregation_probe_rows_threshold" ,
2505+ ScalarValue :: Int64 ( Some ( 2 ) ) ,
2506+ ) ;
2507+ session_config = session_config. set (
2508+ "datafusion.execution.skip_partial_aggregation_probe_ratio_threshold" ,
2509+ ScalarValue :: Float64 ( Some ( 0.1 ) ) ,
2510+ ) ;
2511+
2512+ let ctx = TaskContext :: default ( ) . with_session_config ( session_config) ;
2513+ let output = collect ( aggregate_exec. execute ( 0 , Arc :: new ( ctx) ) ?) . await ?;
2514+
2515+ let expected = [
2516+ "+-----+-------------------+" ,
2517+ "| key | COUNT(val)[count] |" ,
2518+ "+-----+-------------------+" ,
2519+ "| 1 | 1 |" ,
2520+ "| 2 | 1 |" ,
2521+ "| 3 | 1 |" ,
2522+ "| 2 | 1 |" ,
2523+ "| 3 | 1 |" ,
2524+ "| 4 | 1 |" ,
2525+ "+-----+-------------------+" ,
2526+ ] ;
2527+ assert_batches_eq ! ( expected, & output) ;
2528+
2529+ Ok ( ( ) )
2530+ }
2531+
2532+ #[ tokio:: test]
2533+ async fn test_skip_aggregation_after_threshold ( ) -> Result < ( ) > {
2534+ let schema = Arc :: new ( Schema :: new ( vec ! [
2535+ Field :: new( "key" , DataType :: Int32 , true ) ,
2536+ Field :: new( "val" , DataType :: Int32 , true ) ,
2537+ ] ) ) ;
2538+
2539+ let group_by =
2540+ PhysicalGroupBy :: new_single ( vec ! [ ( col( "key" , & schema) ?, "key" . to_string( ) ) ] ) ;
2541+
2542+ let aggr_expr: Vec < Arc < dyn AggregateExpr > > = vec ! [ create_aggregate_expr(
2543+ & count_udaf( ) ,
2544+ & [ col( "val" , & schema) ?] ,
2545+ & [ datafusion_expr:: col( "val" ) ] ,
2546+ & [ ] ,
2547+ & [ ] ,
2548+ & schema,
2549+ "COUNT(val)" ,
2550+ false ,
2551+ false ,
2552+ false ,
2553+ ) ?] ;
2554+
2555+ let input_data = vec ! [
2556+ RecordBatch :: try_new(
2557+ Arc :: clone( & schema) ,
2558+ vec![
2559+ Arc :: new( Int32Array :: from( vec![ 1 , 2 , 3 ] ) ) ,
2560+ Arc :: new( Int32Array :: from( vec![ 0 , 0 , 0 ] ) ) ,
2561+ ] ,
2562+ )
2563+ . unwrap( ) ,
2564+ RecordBatch :: try_new(
2565+ Arc :: clone( & schema) ,
2566+ vec![
2567+ Arc :: new( Int32Array :: from( vec![ 2 , 3 , 4 ] ) ) ,
2568+ Arc :: new( Int32Array :: from( vec![ 0 , 0 , 0 ] ) ) ,
2569+ ] ,
2570+ )
2571+ . unwrap( ) ,
2572+ RecordBatch :: try_new(
2573+ Arc :: clone( & schema) ,
2574+ vec![
2575+ Arc :: new( Int32Array :: from( vec![ 2 , 3 , 4 ] ) ) ,
2576+ Arc :: new( Int32Array :: from( vec![ 0 , 0 , 0 ] ) ) ,
2577+ ] ,
2578+ )
2579+ . unwrap( ) ,
2580+ ] ;
2581+
2582+ let input = Arc :: new ( MemoryExec :: try_new (
2583+ & [ input_data] ,
2584+ Arc :: clone ( & schema) ,
2585+ None ,
2586+ ) ?) ;
2587+ let aggregate_exec = Arc :: new ( AggregateExec :: try_new (
2588+ AggregateMode :: Partial ,
2589+ group_by,
2590+ aggr_expr,
2591+ vec ! [ None ] ,
2592+ Arc :: clone ( & input) as Arc < dyn ExecutionPlan > ,
2593+ schema,
2594+ ) ?) ;
2595+
2596+ let mut session_config = SessionConfig :: default ( ) ;
2597+ session_config = session_config. set (
2598+ "datafusion.execution.skip_partial_aggregation_probe_rows_threshold" ,
2599+ ScalarValue :: Int64 ( Some ( 5 ) ) ,
2600+ ) ;
2601+ session_config = session_config. set (
2602+ "datafusion.execution.skip_partial_aggregation_probe_ratio_threshold" ,
2603+ ScalarValue :: Float64 ( Some ( 0.1 ) ) ,
2604+ ) ;
2605+
2606+ let ctx = TaskContext :: default ( ) . with_session_config ( session_config) ;
2607+ let output = collect ( aggregate_exec. execute ( 0 , Arc :: new ( ctx) ) ?) . await ?;
2608+
2609+ let expected = [
2610+ "+-----+-------------------+" ,
2611+ "| key | COUNT(val)[count] |" ,
2612+ "+-----+-------------------+" ,
2613+ "| 1 | 1 |" ,
2614+ "| 2 | 2 |" ,
2615+ "| 3 | 2 |" ,
2616+ "| 4 | 1 |" ,
2617+ "| 2 | 1 |" ,
2618+ "| 3 | 1 |" ,
2619+ "| 4 | 1 |" ,
2620+ "+-----+-------------------+" ,
2621+ ] ;
2622+ assert_batches_eq ! ( expected, & output) ;
2623+
2624+ Ok ( ( ) )
2625+ }
24452626}
0 commit comments