Skip to content
Merged
61 changes: 44 additions & 17 deletions datafusion/core/tests/fuzz_cases/join_fuzz.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1487,7 +1487,7 @@ fn make_staggered_batches_binary(
//
// What only randomization reaches: `LeftSemi`/`LeftAnti` record matches in a shared
// `AtomicUsize` watermark and emit once, from whichever streamed partition finishes last. The
// streamed side below is spread round-robin over several partitions as one-row batches, so
// streamed side below is spread round-robin over several partitions in small batches, so
// batches arrive in an order no static test pins down, and the counter that gates the final
// pass is seeded from a partition count that deliberately disagrees with the `num_partitions`
// argument. `RightSemi`/`RightAnti` take the mirror path -- every partition emits its own
Expand Down Expand Up @@ -1575,25 +1575,39 @@ fn pwmj_single_exec<K: PwmjFuzzKey>(
.unwrap()
}

/// Rows spread round-robin across `nparts` partitions, one row per batch, with any partition
/// that draws no row left holding a single empty batch. Used for the streamed side throughout,
/// and for the buffered side of the right existence joins, which place no single-partition
/// requirement on it.
/// Rows spread round-robin across `nparts` partitions, with any partition that draws no row
/// left holding a single empty batch. With `stagger_seed`, each partition's rows are cut into
/// random-size batches, empty ones included, by [`stagger_batch_with_seed`]; without it, one
/// row per batch. Used for the streamed side throughout, and for the buffered side of the
/// right existence joins, which place no single-partition requirement on it.
fn pwmj_parts_exec<K: PwmjFuzzKey>(
ids: &[i32],
keys: &[Option<K>],
nparts: usize,
stagger_seed: Option<u64>,
) -> Arc<dyn ExecutionPlan> {
let nparts = nparts.max(1);
let mut partitions: Vec<Vec<RecordBatch>> = vec![Vec::new(); nparts];
let mut part_rows: Vec<(Vec<i32>, Vec<Option<K>>)> = vec![Default::default(); nparts];
for (row, (&id, key)) in ids.iter().zip(keys.iter()).enumerate() {
partitions[row % nparts].push(pwmj_kv_batch(&[id], &[*key]));
}
for p in partitions.iter_mut() {
if p.is_empty() {
p.push(pwmj_kv_batch::<K>(&[], &[]));
}
}
part_rows[row % nparts].0.push(id);
part_rows[row % nparts].1.push(*key);
}
let partitions: Vec<Vec<RecordBatch>> = part_rows
.iter()
.enumerate()
.map(|(part, (ids, keys))| match stagger_seed {
// `stagger_batch_with_seed` panics on a batch with no rows.
_ if ids.is_empty() => vec![pwmj_kv_batch::<K>(&[], &[])],
Some(seed) => {
stagger_batch_with_seed(pwmj_kv_batch(ids, keys), seed + part as u64)
}
None => ids
.iter()
.zip(keys)
.map(|(&id, key)| pwmj_kv_batch(&[id], &[*key]))
.collect(),
})
.collect();
MemorySourceConfig::try_new_exec(&partitions, pwmj_kv_schema::<K>(), None).unwrap()
}

Expand Down Expand Up @@ -1820,6 +1834,9 @@ async fn run_pwmj_fuzz<K: PwmjFuzzKey + std::fmt::Debug>(task_ctx: &Arc<TaskCont
let left_keys = gen_keys(left_len, &mut rng);
let right_ids: Vec<i32> = (0..right_len as i32).collect();
let right_keys = gen_keys(right_len, &mut rng);
// Staggered streamed batches mix rows that can and cannot match, so the classic joins
// set some aside before the sort and map the rest back to their rows.
let stagger_seed = Some(seed);

for op in ops {
for join_type in join_types {
Expand All @@ -1828,7 +1845,7 @@ async fn run_pwmj_fuzz<K: PwmjFuzzKey + std::fmt::Debug>(task_ctx: &Arc<TaskCont
// single-partition shape below never reaches.
let buffered = match join_type {
JoinType::RightSemi | JoinType::RightAnti | JoinType::RightMark => {
pwmj_parts_exec(&left_ids, &left_keys, buffered_nparts)
pwmj_parts_exec(&left_ids, &left_keys, buffered_nparts, None)
}
_ => pwmj_single_exec(&left_ids, &left_keys),
};
Expand All @@ -1840,7 +1857,12 @@ async fn run_pwmj_fuzz<K: PwmjFuzzKey + std::fmt::Debug>(task_ctx: &Arc<TaskCont
let got = pwmj_collect_mark_pairs(
pwmj_plan(
buffered,
pwmj_parts_exec(&right_ids, &right_keys, nparts),
pwmj_parts_exec(
&right_ids,
&right_keys,
nparts,
stagger_seed,
),
op,
join_type,
),
Expand Down Expand Up @@ -1872,7 +1894,7 @@ async fn run_pwmj_fuzz<K: PwmjFuzzKey + std::fmt::Debug>(task_ctx: &Arc<TaskCont
let got = pwmj_collect_id_pairs(
pwmj_plan(
buffered,
pwmj_parts_exec(&right_ids, &right_keys, nparts),
pwmj_parts_exec(&right_ids, &right_keys, nparts, stagger_seed),
op,
join_type,
),
Expand Down Expand Up @@ -2000,7 +2022,12 @@ async fn pwmj_watermark_nulls_duplicates_multi_partition() {
let got = pwmj_collect_id_pairs(
pwmj_plan(
pwmj_single_exec(&left_ids, &left_keys),
pwmj_parts_exec(&case.streamed_ids, &case.streamed_keys, nparts),
pwmj_parts_exec(
&case.streamed_ids,
&case.streamed_keys,
nparts,
None,
),
Operator::Lt,
join_type,
),
Expand Down
Loading
Loading