@@ -20,7 +20,7 @@ use std::time::SystemTime;
2020
2121use crate :: fuzz_cases:: join_fuzz:: JoinTestType :: { HjSmj , NljHj } ;
2222
23- use arrow:: array:: { ArrayRef , BinaryArray , Int32Array , StringArray } ;
23+ use arrow:: array:: { ArrayRef , BinaryArray , Int32Array } ;
2424use arrow:: compute:: SortOptions ;
2525use arrow:: datatypes:: Schema ;
2626use arrow:: record_batch:: RecordBatch ;
@@ -30,24 +30,22 @@ use datafusion::datasource::memory::MemorySourceConfig;
3030use datafusion:: datasource:: source:: DataSourceExec ;
3131use datafusion:: logical_expr:: { JoinType , Operator } ;
3232use datafusion:: physical_expr:: expressions:: BinaryExpr ;
33+ use datafusion:: physical_plan:: collect;
3334use datafusion:: physical_plan:: expressions:: Column ;
3435use datafusion:: physical_plan:: joins:: utils:: { ColumnIndex , JoinFilter } ;
3536use datafusion:: physical_plan:: joins:: {
3637 HashJoinExec , NestedLoopJoinExec , PartitionMode , SortMergeJoinExec ,
3738} ;
38- use datafusion:: physical_plan:: { ExecutionPlan , collect} ;
3939use datafusion:: prelude:: { SessionConfig , SessionContext } ;
4040use datafusion_common:: { NullEquality , ScalarValue } ;
4141use datafusion_execution:: TaskContext ;
4242use datafusion_execution:: disk_manager:: { DiskManagerBuilder , DiskManagerMode } ;
43- use datafusion_execution:: memory_pool:: FairSpillPool ;
4443use datafusion_execution:: runtime_env:: RuntimeEnvBuilder ;
4544use datafusion_physical_expr:: PhysicalExprRef ;
4645use datafusion_physical_expr:: expressions:: Literal ;
4746
4847use itertools:: Itertools ;
49- use rand:: rngs:: SmallRng ;
50- use rand:: { Rng , SeedableRng } ;
48+ use rand:: Rng ;
5149use test_utils:: stagger_batch_with_seed;
5250
5351// Determines what Fuzz tests needs to run
@@ -1128,131 +1126,133 @@ impl JoinFuzzTestCase {
11281126 }
11291127}
11301128
1131- /// Compare a guaranteed-spilling SMJ against an unlimited-memory hash join
1132- /// for filtered materializing joins.
1129+ /// Fuzz test: compare SMJ (with spilling) against HJ (no spill) for filtered
1130+ /// outer joins under memory pressure. This exercises the deferred filtering +
1131+ /// spill read-back path that unit tests can't easily cover with random data.
11331132#[ tokio:: test]
11341133async fn test_filtered_join_spill_fuzz ( ) {
1135- let join_types = [
1136- JoinType :: Inner ,
1137- JoinType :: Left ,
1138- JoinType :: Right ,
1139- JoinType :: Full ,
1140- ] ;
1141- let input1 = make_spill_join_batches ( 256 , 32 , 512 , 1 ) ;
1142- let input2 = make_spill_join_batches ( 256 , 32 , 512 , 2 ) ;
1143- let schema1 = input1[ 0 ] . schema ( ) ;
1144- let schema2 = input2[ 0 ] . schema ( ) ;
1145- let filter = col_lt_col_filter ( Arc :: clone ( & schema1) , Arc :: clone ( & schema2) ) ;
1146- let on = vec ! [
1147- (
1148- Arc :: new( Column :: new_with_schema( "a" , & schema1) . unwrap( ) ) as _,
1149- Arc :: new( Column :: new_with_schema( "a" , & schema2) . unwrap( ) ) as _,
1150- ) ,
1151- (
1152- Arc :: new( Column :: new_with_schema( "b" , & schema1) . unwrap( ) ) as _,
1153- Arc :: new( Column :: new_with_schema( "b" , & schema2) . unwrap( ) ) as _,
1154- ) ,
1155- ] ;
1134+ let join_types = [ JoinType :: Left , JoinType :: Right , JoinType :: Full ] ;
11561135
11571136 let runtime_spill = RuntimeEnvBuilder :: new ( )
1158- . with_memory_pool ( Arc :: new ( FairSpillPool :: new ( 1024 ) ) )
1137+ . with_memory_limit ( 4096 , 1.0 )
11591138 . with_disk_manager_builder (
11601139 DiskManagerBuilder :: default ( ) . with_mode ( DiskManagerMode :: OsTmpDirectory ) ,
11611140 )
11621141 . build_arc ( )
11631142 . unwrap ( ) ;
11641143
1165- for join_type in join_types {
1166- for batch_size in [ 2 , 50 ] {
1167- let session_config = SessionConfig :: new ( ) . with_batch_size ( batch_size) ;
1144+ for join_type in & join_types {
1145+ for ( left_extra, right_extra) in [ ( true , true ) , ( false , true ) , ( true , false ) ] {
1146+ let input1 = make_staggered_batches_i32 ( 1000 , left_extra) ;
1147+ let input2 = make_staggered_batches_i32 ( 1000 , right_extra) ;
11681148
1169- let left_hj = MemorySourceConfig :: try_new_exec (
1170- std:: slice:: from_ref ( & input1) ,
1171- Arc :: clone ( & schema1) ,
1172- None ,
1173- )
1174- . unwrap ( ) ;
1175- let right_hj = MemorySourceConfig :: try_new_exec (
1176- std:: slice:: from_ref ( & input2) ,
1177- Arc :: clone ( & schema2) ,
1178- None ,
1179- )
1180- . unwrap ( ) ;
1181- let hj = Arc :: new (
1182- HashJoinExec :: try_new (
1183- left_hj,
1184- right_hj,
1185- on. clone ( ) ,
1186- Some ( filter. clone ( ) ) ,
1187- & join_type,
1149+ let schema1 = input1[ 0 ] . schema ( ) ;
1150+ let schema2 = input2[ 0 ] . schema ( ) ;
1151+ let filter = col_lt_col_filter ( schema1. clone ( ) , schema2. clone ( ) ) ;
1152+
1153+ let on = vec ! [
1154+ (
1155+ Arc :: new( Column :: new_with_schema( "a" , & schema1) . unwrap( ) ) as _,
1156+ Arc :: new( Column :: new_with_schema( "a" , & schema2) . unwrap( ) ) as _,
1157+ ) ,
1158+ (
1159+ Arc :: new( Column :: new_with_schema( "b" , & schema1) . unwrap( ) ) as _,
1160+ Arc :: new( Column :: new_with_schema( "b" , & schema2) . unwrap( ) ) as _,
1161+ ) ,
1162+ ] ;
1163+
1164+ for batch_size in [ 2 , 49 , 100 ] {
1165+ let session_config = SessionConfig :: new ( ) . with_batch_size ( batch_size) ;
1166+
1167+ // HJ baseline (no memory limit)
1168+ let left_hj = MemorySourceConfig :: try_new_exec (
1169+ std:: slice:: from_ref ( & input1) ,
1170+ schema1. clone ( ) ,
11881171 None ,
1189- PartitionMode :: Partitioned ,
1190- NullEquality :: NullEqualsNothing ,
1191- false ,
11921172 )
1193- . unwrap ( ) ,
1194- ) ;
1195- let ctx_hj = SessionContext :: new_with_config ( session_config. clone ( ) ) ;
1196- let hj_collected = collect ( hj, ctx_hj. task_ctx ( ) ) . await . unwrap ( ) ;
1173+ . unwrap ( ) ;
1174+ let right_hj = MemorySourceConfig :: try_new_exec (
1175+ std:: slice:: from_ref ( & input2) ,
1176+ schema2. clone ( ) ,
1177+ None ,
1178+ )
1179+ . unwrap ( ) ;
1180+ let hj = Arc :: new (
1181+ HashJoinExec :: try_new (
1182+ left_hj,
1183+ right_hj,
1184+ on. clone ( ) ,
1185+ Some ( filter. clone ( ) ) ,
1186+ join_type,
1187+ None ,
1188+ PartitionMode :: Partitioned ,
1189+ NullEquality :: NullEqualsNothing ,
1190+ false ,
1191+ )
1192+ . unwrap ( ) ,
1193+ ) ;
1194+ let ctx_hj = SessionContext :: new_with_config ( session_config. clone ( ) ) ;
1195+ let hj_collected = collect ( hj, ctx_hj. task_ctx ( ) ) . await . unwrap ( ) ;
11971196
1198- let left_smj = MemorySourceConfig :: try_new_exec (
1199- std:: slice:: from_ref ( & input1) ,
1200- Arc :: clone ( & schema1) ,
1201- None ,
1202- )
1203- . unwrap ( ) ;
1204- let right_smj = MemorySourceConfig :: try_new_exec (
1205- std:: slice:: from_ref ( & input2) ,
1206- Arc :: clone ( & schema2) ,
1207- None ,
1208- )
1209- . unwrap ( ) ;
1210- let smj = Arc :: new (
1211- SortMergeJoinExec :: try_new (
1212- left_smj,
1213- right_smj,
1214- on. clone ( ) ,
1215- Some ( filter. clone ( ) ) ,
1216- join_type,
1217- vec ! [ SortOptions :: default ( ) ; on. len( ) ] ,
1218- NullEquality :: NullEqualsNothing ,
1197+ // SMJ with spilling
1198+ let left_smj = MemorySourceConfig :: try_new_exec (
1199+ std:: slice:: from_ref ( & input1) ,
1200+ schema1. clone ( ) ,
1201+ None ,
12191202 )
1220- . unwrap ( ) ,
1221- ) ;
1222- let task_ctx_spill = Arc :: new (
1223- TaskContext :: default ( )
1224- . with_session_config ( session_config)
1225- . with_runtime ( Arc :: clone ( & runtime_spill) ) ,
1226- ) ;
1227- let smj_collected =
1228- collect ( Arc :: clone ( & smj) as Arc < dyn ExecutionPlan > , task_ctx_spill)
1229- . await
1230- . unwrap ( ) ;
1231-
1232- assert ! (
1233- smj. metrics( ) . unwrap( ) . spill_count( ) . unwrap_or_default( ) > 0 ,
1234- "expected SMJ to spill for {join_type:?} batch_size={batch_size}" ,
1235- ) ;
1203+ . unwrap ( ) ;
1204+ let right_smj = MemorySourceConfig :: try_new_exec (
1205+ std:: slice:: from_ref ( & input2) ,
1206+ schema2. clone ( ) ,
1207+ None ,
1208+ )
1209+ . unwrap ( ) ;
1210+ let smj = Arc :: new (
1211+ SortMergeJoinExec :: try_new (
1212+ left_smj,
1213+ right_smj,
1214+ on. clone ( ) ,
1215+ Some ( filter. clone ( ) ) ,
1216+ * join_type,
1217+ vec ! [ SortOptions :: default ( ) ; on. len( ) ] ,
1218+ NullEquality :: NullEqualsNothing ,
1219+ )
1220+ . unwrap ( ) ,
1221+ ) ;
1222+ let task_ctx_spill = Arc :: new (
1223+ TaskContext :: default ( )
1224+ . with_session_config ( session_config)
1225+ . with_runtime ( Arc :: clone ( & runtime_spill) ) ,
1226+ ) ;
1227+ let smj_collected = collect ( smj, task_ctx_spill) . await . unwrap ( ) ;
12361228
1237- let hj_rows: usize = hj_collected. iter ( ) . map ( |b| b. num_rows ( ) ) . sum ( ) ;
1238- let smj_rows: usize = smj_collected. iter ( ) . map ( |b| b. num_rows ( ) ) . sum ( ) ;
1239- assert_eq ! (
1240- hj_rows, smj_rows,
1241- "row count mismatch for {join_type:?} batch_size={batch_size}: \
1242- HJ={hj_rows} SMJ={smj_rows}",
1243- ) ;
1229+ let hj_rows: usize = hj_collected. iter ( ) . map ( |b| b. num_rows ( ) ) . sum ( ) ;
1230+ let smj_rows: usize = smj_collected. iter ( ) . map ( |b| b. num_rows ( ) ) . sum ( ) ;
12441231
1245- if hj_rows > 0 {
1246- let hj_fmt = pretty_format_batches ( & hj_collected) . unwrap ( ) . to_string ( ) ;
1247- let smj_fmt = pretty_format_batches ( & smj_collected) . unwrap ( ) . to_string ( ) ;
1248- let mut hj_sorted: Vec < & str > = hj_fmt. trim ( ) . lines ( ) . collect ( ) ;
1249- hj_sorted. sort_unstable ( ) ;
1250- let mut smj_sorted: Vec < & str > = smj_fmt. trim ( ) . lines ( ) . collect ( ) ;
1251- smj_sorted. sort_unstable ( ) ;
12521232 assert_eq ! (
1253- hj_sorted, smj_sorted,
1254- "content mismatch for {join_type:?} batch_size={batch_size}" ,
1233+ hj_rows, smj_rows,
1234+ "Row count mismatch for {join_type:?} batch_size={batch_size} \
1235+ left_extra={left_extra} right_extra={right_extra}: \
1236+ HJ={hj_rows} SMJ={smj_rows}"
12551237 ) ;
1238+
1239+ if hj_rows > 0 {
1240+ let hj_fmt =
1241+ pretty_format_batches ( & hj_collected) . unwrap ( ) . to_string ( ) ;
1242+ let smj_fmt =
1243+ pretty_format_batches ( & smj_collected) . unwrap ( ) . to_string ( ) ;
1244+
1245+ let mut hj_sorted: Vec < & str > = hj_fmt. trim ( ) . lines ( ) . collect ( ) ;
1246+ hj_sorted. sort_unstable ( ) ;
1247+ let mut smj_sorted: Vec < & str > = smj_fmt. trim ( ) . lines ( ) . collect ( ) ;
1248+ smj_sorted. sort_unstable ( ) ;
1249+
1250+ assert_eq ! (
1251+ hj_sorted, smj_sorted,
1252+ "Content mismatch for {join_type:?} batch_size={batch_size} \
1253+ left_extra={left_extra} right_extra={right_extra}"
1254+ ) ;
1255+ }
12561256 }
12571257 }
12581258 }
@@ -1347,39 +1347,3 @@ fn make_staggered_batches_binary(
13471347 // preserve your existing randomized partitioning
13481348 stagger_batch_with_seed ( batch, 42 )
13491349}
1350-
1351- /// Sorted, low-cardinality inputs whose wide payloads force SMJ key-group
1352- /// spilling. `(a, b)` is sorted (`b` is constant) and `x` is nullable to
1353- /// exercise filter NULL handling.
1354- fn make_spill_join_batches (
1355- len : usize ,
1356- num_keys : i32 ,
1357- payload_len : usize ,
1358- seed : u64 ,
1359- ) -> Vec < RecordBatch > {
1360- let mut rng = SmallRng :: seed_from_u64 ( seed) ;
1361-
1362- let mut keys: Vec < i32 > = ( 0 ..len) . map ( |_| rng. random_range ( 0 ..num_keys) ) . collect ( ) ;
1363- keys. sort_unstable ( ) ;
1364- let a = Int32Array :: from_iter_values ( keys) ;
1365- let b = Int32Array :: from_iter_values ( std:: iter:: repeat_n ( 0 , len) ) ;
1366- let x = Int32Array :: from_iter ( ( 0 ..len) . map ( |_| {
1367- if rng. random_range ( 0 ..10 ) == 0 {
1368- None
1369- } else {
1370- Some ( rng. random_range ( 0 ..1000 ) )
1371- }
1372- } ) ) ;
1373-
1374- let payload = "a" . repeat ( payload_len) ;
1375- let p = StringArray :: from_iter_values ( std:: iter:: repeat_n ( payload. as_str ( ) , len) ) ;
1376- let batch = RecordBatch :: try_from_iter ( vec ! [
1377- ( "a" , Arc :: new( a) as ArrayRef ) ,
1378- ( "b" , Arc :: new( b) as ArrayRef ) ,
1379- ( "x" , Arc :: new( x) as ArrayRef ) ,
1380- ( "p" , Arc :: new( p) as ArrayRef ) ,
1381- ] )
1382- . unwrap ( ) ;
1383-
1384- stagger_batch_with_seed ( batch, 7 )
1385- }
0 commit comments