Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
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