Skip to content

Commit b54e39a

Browse files
committed
finalize pwmj drain fix
1 parent c66d9fc commit b54e39a

2 files changed

Lines changed: 147 additions & 54 deletions

File tree

‎datafusion/physical-plan/src/joins/piecewise_merge_join/classic_join.rs‎

Lines changed: 30 additions & 54 deletions
Original file line numberDiff line numberDiff line change
@@ -278,14 +278,7 @@ impl ClassicPWMJStream {
278278
}
279279

280280
if !self.batch_process_state.continue_process {
281-
self.batch_process_state
282-
.output_batches
283-
.finish_buffered_batch()?;
284-
if let Some(batch) = self
285-
.batch_process_state
286-
.output_batches
287-
.next_completed_batch()
288-
{
281+
if let Some(batch) = self.batch_process_state.next_drained_batch()? {
289282
return Ok(StatefulStreamResult::Ready(Some(batch)));
290283
}
291284

@@ -294,7 +287,7 @@ impl ClassicPWMJStream {
294287
}
295288

296289
// Produce more work
297-
let batch = resolve_classic_join(
290+
let scan_batch = resolve_classic_join(
298291
buffered_side,
299292
stream_batch,
300293
&self.schema,
@@ -305,23 +298,18 @@ impl ClassicPWMJStream {
305298
)?;
306299

307300
if !self.batch_process_state.continue_process {
308-
// A flush can queue multiple batches, so transition only after draining.
309-
self.batch_process_state
310-
.output_batches
311-
.finish_buffered_batch()?;
312-
if let Some(batch) = self
313-
.batch_process_state
314-
.output_batches
315-
.next_completed_batch()
316-
{
301+
// The queue can hold several completed batches, so transition only
302+
// after draining all of them.
303+
if let Some(batch) = self.batch_process_state.next_drained_batch()? {
317304
return Ok(StatefulStreamResult::Ready(Some(batch)));
318305
}
319306

307+
// Fully drained; emit the scan's empty tail batch and move on.
320308
self.state = PiecewiseMergeJoinStreamState::FetchStreamBatch;
321-
return Ok(StatefulStreamResult::Ready(Some(batch)));
309+
return Ok(StatefulStreamResult::Ready(Some(scan_batch)));
322310
}
323311

324-
Ok(StatefulStreamResult::Ready(Some(batch)))
312+
Ok(StatefulStreamResult::Ready(Some(scan_batch)))
325313
}
326314

327315
// Process remaining unmatched rows
@@ -335,22 +323,7 @@ impl ClassicPWMJStream {
335323
}
336324

337325
if !self.batch_process_state.continue_process {
338-
if let Some(batch) = self
339-
.batch_process_state
340-
.output_batches
341-
.next_completed_batch()
342-
{
343-
return Ok(StatefulStreamResult::Ready(Some(batch)));
344-
}
345-
346-
self.batch_process_state
347-
.output_batches
348-
.finish_buffered_batch()?;
349-
if let Some(batch) = self
350-
.batch_process_state
351-
.output_batches
352-
.next_completed_batch()
353-
{
326+
if let Some(batch) = self.batch_process_state.next_drained_batch()? {
354327
return Ok(StatefulStreamResult::Ready(Some(batch)));
355328
}
356329

@@ -386,27 +359,11 @@ impl ClassicPWMJStream {
386359
self.batch_process_state.output_batches.push_batch(batch)?;
387360

388361
self.batch_process_state.continue_process = false;
389-
if let Some(batch) = self
390-
.batch_process_state
391-
.output_batches
392-
.next_completed_batch()
393-
{
394-
return Ok(StatefulStreamResult::Ready(Some(batch)));
395-
}
396-
397-
self.batch_process_state
398-
.output_batches
399-
.finish_buffered_batch()?;
400-
if let Some(batch) = self
401-
.batch_process_state
402-
.output_batches
403-
.next_completed_batch()
404-
{
362+
if let Some(batch) = self.batch_process_state.next_drained_batch()? {
405363
return Ok(StatefulStreamResult::Ready(Some(batch)));
406364
}
407365

408366
self.state = PiecewiseMergeJoinStreamState::Completed;
409-
self.batch_process_state.reset();
410367
Ok(StatefulStreamResult::Continue)
411368
}
412369
}
@@ -451,6 +408,17 @@ impl BatchProcessState {
451408
self.continue_process = true;
452409
self.processed_null_count = false;
453410
}
411+
412+
// Pops the next completed batch, flushing the partial remainder once the
413+
// queue is empty. `None` guarantees the coalescer holds no pending rows,
414+
// so the stream may leave its current state without losing output.
415+
fn next_drained_batch(&mut self) -> Result<Option<RecordBatch>> {
416+
if let Some(batch) = self.output_batches.next_completed_batch() {
417+
return Ok(Some(batch));
418+
}
419+
self.output_batches.finish_buffered_batch()?;
420+
Ok(self.output_batches.next_completed_batch())
421+
}
454422
}
455423

456424
impl Stream for ClassicPWMJStream {
@@ -836,7 +804,9 @@ mod tests {
836804

837805
#[tokio::test]
838806
async fn join_right_unmatched_rows_exceeding_batch_size() -> Result<()> {
839-
// 100 < {1, 2, 3} is false, making every streamed row unmatched.
807+
// 100 < {1, 2, 3} is false, making every streamed row unmatched; with
808+
// batch_size 2 the scan finishes with more than one completed batch
809+
// still queued, and all of them must be emitted before the stream ends.
840810
let left = build_table(("a1", &vec![0]), ("b1", &vec![100]), ("c1", &vec![0]));
841811
let right = build_table(
842812
("a2", &vec![10, 20, 30]),
@@ -878,6 +848,9 @@ mod tests {
878848

879849
#[tokio::test]
880850
async fn join_left_unmatched_rows_exact_batch_multiple() -> Result<()> {
851+
// Four unmatched buffered rows with batch_size 2 flush with no
852+
// remainder (an exact multiple); the unmatched pass must terminate
853+
// once the queue drains instead of recomputing the same rows.
881854
let left = build_table(
882855
("a1", &vec![10, 20, 30, 40]),
883856
("b1", &vec![100, 101, 102, 103]),
@@ -896,6 +869,9 @@ mod tests {
896869
let mut stream = join.execute(0, task_ctx)?;
897870

898871
let mut batches = Vec::with_capacity(2);
872+
// Bounded polls (ignoring the empty batch the scan phase emits): a
873+
// non-terminating stream fails the count assertion below instead of
874+
// hanging the test.
899875
for _ in 0..3 {
900876
match stream.next().await.transpose()? {
901877
Some(batch) if batch.num_rows() > 0 => batches.push(batch),

‎datafusion/sqllogictest/test_files/pwmj.slt‎

Lines changed: 117 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -342,5 +342,122 @@ ORDER BY 1,2;
342342
1 3
343343
2 3
344344

345+
# The queries below shrink the output batch size so the join finishes with
346+
# several completed output batches still queued; every queued batch must be
347+
# emitted before the operator changes state.
348+
349+
statement ok
350+
CREATE TABLE drain_t1 (t1_v INT);
351+
352+
statement ok
353+
CREATE TABLE drain_t2 (t2_v INT);
354+
355+
# No `t1_v < t2_v` pair matches, so every row of both tables is unmatched.
356+
statement ok
357+
INSERT INTO drain_t1 VALUES (100), (101), (102), (103);
358+
359+
statement ok
360+
INSERT INTO drain_t2 VALUES (1), (2), (3);
361+
362+
# Shrink the batch size only for query execution so each table above stays a
363+
# single input batch larger than one output batch.
364+
statement ok
365+
set datafusion.execution.batch_size = 2;
366+
367+
# Right join: every unmatched streamed row must be emitted null-extended.
368+
query II
369+
SELECT t1.t1_v, t2.t2_v
370+
FROM drain_t1 t1
371+
RIGHT JOIN drain_t2 t2
372+
ON t1.t1_v < t2.t2_v
373+
ORDER BY 2;
374+
----
375+
NULL 1
376+
NULL 2
377+
NULL 3
378+
379+
query TT
380+
EXPLAIN
381+
SELECT t1.t1_v, t2.t2_v
382+
FROM drain_t1 t1
383+
RIGHT JOIN drain_t2 t2
384+
ON t1.t1_v < t2.t2_v
385+
ORDER BY 2;
386+
----
387+
logical_plan
388+
01)Sort: t2.t2_v ASC NULLS LAST
389+
02)--Right Join: Filter: t1.t1_v < t2.t2_v
390+
03)----SubqueryAlias: t1
391+
04)------TableScan: drain_t1 projection=[t1_v]
392+
05)----SubqueryAlias: t2
393+
06)------TableScan: drain_t2 projection=[t2_v]
394+
physical_plan
395+
01)SortPreservingMergeExec: [t2_v@1 ASC NULLS LAST]
396+
02)--SortExec: expr=[t2_v@1 ASC NULLS LAST], preserve_partitioning=[true]
397+
03)----PiecewiseMergeJoin: operator=Lt, join_type=Right, on=(t1_v < t2_v)
398+
04)------SortExec: expr=[t1_v@0 DESC], preserve_partitioning=[false]
399+
05)--------DataSourceExec: partitions=1, partition_sizes=[1]
400+
06)------RepartitionExec: partitioning=RoundRobinBatch(4), input_partitions=1
401+
07)--------DataSourceExec: partitions=1, partition_sizes=[1]
402+
403+
# Left join: the four unmatched buffered rows are an exact multiple of the
404+
# batch size. The LIMIT keeps the query bounded even if the unmatched pass
405+
# fails to terminate, so a regression fails with extra rows instead of
406+
# hanging the runner.
407+
query II rowsort
408+
SELECT t1.t1_v, t2.t2_v
409+
FROM drain_t1 t1
410+
LEFT JOIN drain_t2 t2
411+
ON t1.t1_v < t2.t2_v
412+
LIMIT 10;
413+
----
414+
100 NULL
415+
101 NULL
416+
102 NULL
417+
103 NULL
418+
419+
query TT
420+
EXPLAIN
421+
SELECT t1.t1_v, t2.t2_v
422+
FROM drain_t1 t1
423+
LEFT JOIN drain_t2 t2
424+
ON t1.t1_v < t2.t2_v
425+
LIMIT 10;
426+
----
427+
logical_plan
428+
01)Limit: skip=0, fetch=10
429+
02)--Left Join: Filter: t1.t1_v < t2.t2_v
430+
03)----SubqueryAlias: t1
431+
04)------Limit: skip=0, fetch=10
432+
05)--------TableScan: drain_t1 projection=[t1_v], fetch=10
433+
06)----SubqueryAlias: t2
434+
07)------TableScan: drain_t2 projection=[t2_v]
435+
physical_plan
436+
01)CoalescePartitionsExec: fetch=10
437+
02)--PiecewiseMergeJoin: operator=Lt, join_type=Left, on=(t1_v < t2_v)
438+
03)----SortExec: TopK(fetch=10), expr=[t1_v@0 DESC], preserve_partitioning=[false]
439+
04)------DataSourceExec: partitions=1, partition_sizes=[1]
440+
05)----RepartitionExec: partitioning=RoundRobinBatch(4), input_partitions=1
441+
06)------DataSourceExec: partitions=1, partition_sizes=[1]
442+
443+
# Full join: both unmatched sides must drain across the state transitions.
444+
query II rowsort
445+
SELECT t1.t1_v, t2.t2_v
446+
FROM drain_t1 t1
447+
FULL JOIN drain_t2 t2
448+
ON t1.t1_v < t2.t2_v
449+
LIMIT 20;
450+
----
451+
100 NULL
452+
101 NULL
453+
102 NULL
454+
103 NULL
455+
NULL 1
456+
NULL 2
457+
NULL 3
458+
459+
statement ok
460+
reset datafusion.execution.batch_size;
461+
345462
statement ok
346463
set datafusion.optimizer.enable_piecewise_merge_join = false;

0 commit comments

Comments
 (0)