@@ -3,6 +3,7 @@ mod tests {
33 use datafusion:: arrow:: util:: pretty:: pretty_format_batches;
44 use datafusion:: error:: DataFusionError ;
55 use datafusion:: execution:: SessionState ;
6+ use datafusion:: prelude:: SessionContext ;
67 use datafusion:: physical_plan:: execute_stream;
78 use datafusion_distributed:: test_utils:: localhost:: start_localhost_context;
89 use datafusion_distributed:: test_utils:: test_work_unit_feed:: {
@@ -612,9 +613,8 @@ mod tests {
612613
613614 #[ tokio:: test]
614615 async fn broadcast_join_over_feeds ( ) -> Result < ( ) , Box < dyn std:: error:: Error > > {
615- let ( plan, results) = run_query (
616+ let ( plan, results) = run_query_with_setup (
616617 r#"
617- SET distributed.broadcast_joins=true;
618618 SELECT
619619 a.tag as a_tag, a.task as a_task, a.partition as a_partition, a.letter,
620620 b.tag as b_tag, b.task as b_task, b.partition as b_partition
@@ -623,6 +623,7 @@ mod tests {
623623 ON a.letter = b.letter
624624 ORDER BY a_task, a_partition, a.letter, b_task, b_partition
625625 "# ,
626+ |ctx : & mut SessionContext | ctx. set_distributed_broadcast_joins ( true ) ,
626627 )
627628 . await ?;
628629
@@ -803,9 +804,8 @@ mod tests {
803804
804805 #[ tokio:: test]
805806 async fn nested_union_budget_exceeds_children_sum ( ) -> Result < ( ) , Box < dyn std:: error:: Error > > {
806- let ( plan, results) = run_query (
807+ let ( plan, results) = run_query_with_setup (
807808 r#"
808- SET distributed.broadcast_joins = true;
809809 SELECT b.tag, a.tag
810810 FROM test_work_unit('big', 4, 'rows(1)', 'rows(1)', 'rows(1)', 'rows(1)') b
811811 INNER JOIN (
@@ -815,6 +815,7 @@ mod tests {
815815 ) a ON a.letter = b.letter
816816 ORDER BY a.tag, b.tag
817817 "# ,
818+ |ctx : & mut SessionContext | ctx. set_distributed_broadcast_joins ( true ) ,
818819 )
819820 . await ?;
820821
@@ -860,11 +861,19 @@ mod tests {
860861 }
861862
862863 async fn run_query ( sql : & str ) -> Result < ( String , String ) , DataFusionError > {
864+ run_query_with_setup ( sql, |_| Ok ( ( ) ) ) . await
865+ }
866+
867+ async fn run_query_with_setup (
868+ sql : & str ,
869+ setup : impl FnOnce ( & mut SessionContext ) -> Result < ( ) , DataFusionError > ,
870+ ) -> Result < ( String , String ) , DataFusionError > {
863871 let ( mut ctx, _guard, _) = start_localhost_context ( 3 , build_state) . await ;
864872 ctx. set_distributed_work_unit_feed ( |p : & RowGeneratorExec | Some ( & p. feed ) ) ;
865873 ctx. set_distributed_user_codec ( TestWorkUnitFeedExecCodec ) ;
866874 ctx. set_distributed_task_estimator ( TestWorkUnitFeedTaskEstimator ) ;
867875 ctx. register_udtf ( "test_work_unit" , Arc :: new ( TestWorkUnitFeedFunction ) ) ;
876+ setup ( & mut ctx) ?;
868877
869878 let mut df_opt = None ;
870879 for sql in sql. split ( ";" ) {
0 commit comments