Skip to content

[SPARK-60113][CORE][SQL] Wait for the pipelined Python UDF writer to exit before task completion cleanup - #59332

Open
viirya wants to merge 2 commits into
apache:masterfrom
viirya:SPARK-60113-wait-pipelined-writer
Open

viirya wants to merge 2 commits into
apache:masterfrom
viirya:SPARK-60113-wait-pipelined-writer

Conversation

@viirya

@viirya viirya commented Oct 10, 2026 •

Copy link
Copy Markdown
Member

What changes were proposed in this pull request?

In pipelined Python UDF mode (spark.python.udf.pipelined.enabled), the listener that stops the writer thread now waits until the writer has actually exited, and it runs before every task completion listener.

  1. Add TaskContext.addTaskPreCompletionListener (private[spark]). TaskContextImpl runs these listeners before every listener added with addTaskCompletionListener, regardless of when or from which thread that listener was added: invokeListeners always takes the next listener from the pre-completion stack first. The default implementation in TaskContext falls back to addTaskCompletionListener, and BarrierTaskContext delegates.

  2. The writer is now wrapped in BasePythonRunner.PipelinedWriterTask, which counts down a latch in a finally around the whole writer. The listener interrupts the writer with cancel(true) as before, and then waits on that latch instead of calling Future.get(). If the listener runs before the pool has started the writer, an AtomicBoolean claim makes sure the writer is skipped, so the listener does not wait on a latch that would never be counted down.

  3. The listener is added with addTaskPreCompletionListener.

  4. The wait has no time bound, the same as the WriterThread.join() that the writer thread used before SPARK-44705 (added for SPARK-33277). While it waits, it logs a warning every 10 seconds. If the listener thread is interrupted (for example, by a task kill), it keeps waiting and does not restore the interrupt flag, so that the completion listeners that run next are not affected; the kill is already recorded in the TaskContext.

  5. A fatal error of the writer (e.g. OutOfMemoryError) is now also passed to the reader through writer.setException with a socket shutdownOutput(), and it is logged; previously the thread pool's FutureTask kept it and the reader waited for the Python worker.

Why are the changes needed?

The existing listener does:

writerFuture.cancel(true)
try { writerFuture.get() } catch { case _: CancellationException | ... => }

Its comment says this waits for the writer to exit before the listeners that run after it (registered earlier, LIFO) free the memory backing the input rows. But once cancel(true) succeeds, FutureTask.get() throws CancellationException right away and does not wait for the interrupted runnable to return. So the writer can still be serializing a row, or adding it to a HybridRowQueue page, when the queue's close listener frees that page. That is a use-after-free, and with off-heap memory it can crash the executor. A HybridQueue.createNewQueue racing with close() can also add a page that nobody closes until cleanUpAllAllocatedMemory.

A task completes while the writer is still running when the output is not fully consumed, for example with a LIMIT or when the task fails.

Ordering the writer listener in the completion listener stack is not enough either. Listeners run in reverse order of registration, and the writer thread itself adds completion listeners lazily while it pulls the upstream iterator, e.g. ParquetReaderCallback.initIfNotAlready of the DSv2 Parquet reader (which closes the vectorized reader and its off-heap column vectors) or the UnsafeExternalSorter that ExternalAppendOnlyUnsafeRowArray creates once it spills. Those would run before the writer listener, which is why the writer is stopped from a pre-completion listener.

The wait cannot be bounded safely: if the listener gave up, the listeners that run next would still free the memory, and throwing from the listener does not stop them either. Even a writer blocked in upstream I/O that ignores interrupts (e.g. a slow JDBC fetch) adds the row it gets next to the HybridRowQueue, whose InMemoryRowQueue writes through the page's base object captured at construction, so a freed page would be corrupted even on-heap (HeapMemoryAllocator pools pages for reuse). The price is that task completion is held for as long as such I/O blocks. Usually the writer exits soon after the interrupt: channel.write fails with ClosedByInterruptException, and otherwise the writer loop checks the interrupt flag between writeNextInputToStream calls, each of which can pull and serialize a whole batch (e.g. for the Arrow runners).

The pipelined mode was added in 4.3.0 by SPARK-56642, so branch-4.3 and branch-4.x have the same issue.

Does this PR introduce any user-facing change?

No. This fixes a bug in an opt-in mode that is not released yet (4.3.0).

How was this patch tested?

New tests in BasePythonRunnerSuite:

  • Task completion stays blocked until a running writer that ignores interrupts exits. Both a cleanup listener added on the task thread before the writer starts and one added by the writer thread itself see that the writer has exited. The tests first wait until task completion has reached the writer listener.
  • An interrupt of the thread running the listeners does not stop the wait and is not passed to the completion listeners that run next.
  • A writer stopped before it starts never runs.
  • Stopping waits for a writer that was claimed by run() on a thread whose interrupt flag was already set.
  • A fatal error of the writer is logged.

New tests in TaskContextSuite for the pre-completion listener ordering (including one added from another thread while completion listeners run) and failure handling.

Before the fix, a test with the old listener code showed markTaskCompleted returning while the writer was still running. With the old cancel-then-get() behavior put back, the blocking tests fail, and with the writer listener added as an ordinary completion listener, the cleanup listener added by the writer thread runs while the writer is still running.

Also ran pyspark.sql.tests.pandas.test_pipelined_udf.

Was this patch authored or co-authored using generative AI tooling?

Generated-by: Claude Code (Claude Opus 5.5)

This pull request and its description were written by Isaac.

…efore task completion cleanup

In pipelined Python UDF mode, the task completion listener that stops the writer called
`writerFuture.cancel(true)` and then `writerFuture.get()`, intending to wait for the writer to
exit before the listeners that run after it free the memory backing the input rows. Once
`cancel(true)` succeeds, `FutureTask.get()` throws `CancellationException` immediately without
waiting for the interrupted runnable, so the writer could still be serializing a row or adding
it to a `HybridRowQueue` page when that page was freed: a use-after-free that can crash the
executor with off-heap memory.

Wrap the writer in `PipelinedWriterTask`, which counts down a latch when the writer returns,
and have the listener interrupt the writer and wait on that latch. If the listener gets there
before the writer starts, the writer is skipped instead. The wait has no time bound, like the
`WriterThread.join()` this mode replaced (SPARK-33277): giving up would allow the
use-after-free. It logs a warning every 10 seconds while it waits, and an interrupt of the
listener thread is restored after the wait.

Co-authored-by: Claude Code <noreply@anthropic.com>

@dongjoon-hyun dongjoon-hyun left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thank you for working on this, @viirya. I left 8 inline comments. Here is a summary.

  1. (PythonRunner.scala L203) The wait only covers listeners registered before startPipelinedWriter. Listeners that the writer thread itself registers lazily while pulling upstream input (e.g. ParquetReaderCallback.initIfNotAlready of the DSv2 Parquet reader, or UnsafeExternalSorter's cleanupResources() when ExternalAppendOnlyUnsafeRowArray switches to the spillable sorter) are pushed later, so they run first under LIFO and can free memory while the writer is still running and not even interrupted yet.
  2. (L239) The unbounded wait can block task completion indefinitely when the writer is stuck in upstream I/O that ignores interrupts (e.g. JDBC ResultSet.next(), HTTP/S3 stream reads). The task reaper is disabled by default and only handles killed tasks, not a normal LIMIT completion.
  3. (L259) Restoring the interrupt flag means all remaining (earlier-registered) cleanup listeners run with the flag set, so NIO/lock-based cleanup in them can fail.
  4. (BasePythonRunnerSuite.scala L94) The blocking assertions can pass even when the completion thread has not reached the listener yet, so the test may not catch a regression on a slow CI machine.
  5. (BasePythonRunnerSuite.scala L141) The race where the pool thread enters FutureTask.run and wins the claim after cancel(true) is not tested.
  6. (L240) The hand-written interrupt-preserving loop can be replaced with Guava Uninterruptibles.awaitUninterruptibly, which is already used in core.
  7. (L204) A fatal error thrown by the writer is captured by the FutureTask and never logged or propagated.
  8. (L231) The comment says the writer checks the interrupt flag between rows, but it checks only between writeNextInputToStream calls (a whole batch for Arrow runners).

private[python] def startPipelinedWriter(writer: Runnable, context: TaskContext): Unit = {
val writerTask = new PipelinedWriterTask(writer)
val writerFuture = pipelinedWriterThreadPool.submit(writerTask)
context.addTaskCompletionListener[Unit] { _ =>

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This orders the writer listener only against the listeners registered before this line. In pipelined mode, the upstream iterator is first pulled on the writer pool thread, so any completion listener that upstream operators register lazily is pushed after this one and runs before it under LIFO, while the writer is still running and not even interrupted yet. For example:

  • DSv2 Parquet: ParquetReaderCallback.initIfNotAlready registers closeCurrent() when the first file is opened (ParquetPartitionReaderFactory.scala L391), which closes the vectorized reader (off-heap column vectors when spark.sql.columnVector.offheap.enabled=true).
  • ExternalAppendOnlyUnsafeRowArray (SMJ / Window upstream) creates an UnsafeExternalSorter lazily once the threshold is exceeded, and its constructor registers cleanupResources().

In both cases, a LIMIT that completes the task frees the memory under the running writer, which is the same use-after-free that this PR is trying to prevent. The new tests register the cleanup listener on the task thread before the writer starts, so they don't cover this. Could we stop and wait for the writer before any completion listener runs (e.g. a pre-completion hook in TaskContextImpl), instead of relying on the position in the LIFO stack?

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Good catch, confirmed. DataSourceRDD.compute creates the reader inside flatMap, so ParquetReaderCallback.initIfNotAlready runs on the writer thread and its closeCurrent() lands above the writer listener in the stack.

Added TaskContext.addTaskPreCompletionListener (private[spark], default falls back to addTaskCompletionListener). TaskContextImpl.invokeListeners now always takes the next listener from the pre-completion stack first, so these run before every completion listener, including ones added later or from other threads, and also when a pre-completion listener is added while another thread is already running the completion listeners. The writer listener uses it. The test now adds a cleanup listener from the writer thread itself, and it fails if the writer listener is switched back to an ordinary completion listener.

As far as I can tell, the WriterThread listener before SPARK-44705 had the same gap, since it was also registered before the writer thread started pulling upstream.

val writerTask = new PipelinedWriterTask(writer)
val writerFuture = pipelinedWriterThreadPool.submit(writerTask)
context.addTaskCompletionListener[Unit] { _ =>
writerTask.stopAndAwaitExit(writerFuture, taskIdentifier(context))

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Since the listener no longer looks at the Future, a fatal Throwable escaping PipelinedWriterRunnable.run (NonFatal excludes VirtualMachineError, LinkageError, etc.) is stored in the FutureTask and silently dropped. In that case the writer also doesn't call shutdownOutput() or writer.setException(...), so the Python worker keeps waiting for input and the reader waits for output with no root cause recorded. Since the listener now knows the writer has exited, could we log the cause here (e.g. future.get() after the latch, ignoring CancellationException)?

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed in two places. PipelinedWriterRunnable now also calls writer.setException(t) and shutdownOutput() for fatal errors before rethrowing them, so the reader fails with the cause instead of waiting for Python. PipelinedWriterTask.run logs anything that escapes the writer. A future.get() after the latch would not be enough here: once cancel(true) has succeeded, the future's outcome is no longer available.

def stopAndAwaitExit(future: Future[_], taskName: String): Unit = {
// Interrupts the writer thread if the writer is running. This unblocks channel.write
// (the JDK closes the channel and throws ClosedByInterruptException), and the writer
// loop checks the interrupt flag between rows, so the wait below is normally short.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

nit. checks the interrupt flag between rows is not accurate. The loop checks it only between writeNextInputToStream calls. For Arrow runners, one call serializes a whole batch (up to arrowMaxRecordsPerBatch rows) and pulls from upstream, and ColumnarArrowEvalPythonEvaluatorFactory Path 2 copies all rows of an input ColumnarBatch into the HybridRowQueue in one step without checking the interrupt. The same wording is in the PR description. Since the wait below is unbounded, it would be good to describe it precisely.

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed the comment: it now says that the interrupt is seen only between writeNextInputToStream calls, each of which can pull and serialize a whole batch. Will fix the PR description as well.

}
// Wait without a time bound: giving up would let the listeners that run next free memory
// the writer may still be reading, and failing this listener would not stop them either.
// A killed task that stays stuck here is handled by the task reaper, if enabled.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This unbounded wait is a behavior change for writers blocked in upstream I/O that does not respond to interrupts, e.g. JDBC ResultSet.next() on a slow fetch, or an HTTP/S3 stream read with a long (or no) socket timeout. When a LIMIT is satisfied, the task completes normally, cancel(true) cannot unblock the read, and markTaskCompleted now waits until that read returns, holding the task slot and only logging a warning every 10 seconds. Previously the listener returned immediately.

Also, the task reaper only applies to killed tasks and spark.task.reaper.enabled is false by default, so nothing recovers the normal completion case. In these examples the rows are on-heap, so the wait doesn't buy any safety. Could we mention this trade-off explicitly, or consider limiting the wait to the cases where the input memory can actually be freed under the writer?

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The trade-off is real, and I documented it in the comment. I still kept the wait unbounded, because in these examples the wait does buy safety. In BatchEvalPython / ArrowEvalPython, the row the writer gets next from the blocked upstream read goes to queue.add before the loop checks the interrupt again. InMemoryRowQueue writes through the page's base object that it captured at construction. So even for on-heap memory, a freed page (at least the page size, which HeapMemoryAllocator pools and hands to other consumers) would be silently corrupted, and off-heap it would crash. I couldn't find a reliable way to tell when giving up would be safe. The WriterThread.join() before SPARK-44705 was also unbounded, for every Python UDF.

// Wait without a time bound: giving up would let the listeners that run next free memory
// the writer may still be reading, and failing this listener would not stop them either.
// A killed task that stays stuck here is handled by the task reaper, if enabled.
val startNs = System.nanoTime()

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

nit. The interrupted flag, the outer try/finally, and the inner catch InterruptedException re-implement Guava's Uninterruptibles.awaitUninterruptibly(CountDownLatch, long, TimeUnit), which also restores the caller's interrupt status before returning. core already uses Uninterruptibles (e.g. AsyncEventQueue.scala). This could be simplified to something like:

while (!Uninterruptibles.awaitUninterruptibly(
    exited, pipelinedWriterExitWarnIntervalMs, TimeUnit.MILLISECONDS)) {
  logWarning(...)
}

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Following your next comment, the wait no longer restores the interrupt status, while Uninterruptibles.awaitUninterruptibly always restores it. So I kept the loop, which is now just a catch-and-continue without the outer try/finally.

}
}
} finally {
if (interrupted) {

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Restoring the interrupt flag here means that every listener registered earlier (which runs after this one) runs with the flag set, e.g. worker.stop() (L560), queue.close() of the HybridRowQueue (which deletes spill files), and the reader/allocator closes. Since this listener can now take much longer, a task-kill interrupt is more likely to land during this wait and be carried over. Any blocking NIO channel operation or lock wait in those listeners would then fail with ClosedByInterruptException / InterruptedException, and the cleanup would be skipped or logged as a listener failure. Previously, the old listener didn't wait at all. Is this intended? If so, could we add a comment about the impact on the following listeners? The new test asserts the propagation, so it would lock in this behavior.

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Not intended. Changed it to not restore the flag, with a comment: the kill is already recorded in the TaskContext, and the old listener also swallowed the interrupt (get() threw InterruptedException). The test now asserts that a later completion listener does not see the flag.

completion.start()
try {
// Task completion must stay blocked while the writer is still running.
completion.join(500)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

completion.isAlive and exitedBeforeCleanup.isEmpty are also true when the completion thread has not reached the writer listener yet. On a slow CI machine where the thread isn't scheduled within 500 ms, this test would pass even with the old cancel(true) + get() listener. How about registering a listener after startPipelinedWriter (which runs first under LIFO) that counts down a latch, and waiting on it before checking isAlive? The same applies to completion.join(200) + interrupt() at L121, which doesn't verify that the interrupt arrives during the wait.

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed. The tests add a pre-completion listener after startPipelinedWriter, which runs right before the writer listener and counts down a latch. The tests wait for that latch before checking isAlive, and the interrupt test also interrupts only after it.

val writer = new UninterruptibleWriter
val writerTask = new BasePythonRunner.PipelinedWriterTask(writer)
// A future the pool has not started yet: stopping must not wait for it.
val future = new FutureTask[Unit](writerTask, ())

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This FutureTask is never run by a pool thread, so the test covers only the case where FutureTask.run sees the cancelled state and never calls PipelinedWriterTask.run. The real race isn't covered: the pool thread has already entered FutureTask.run, cancel(true) interrupts it, and run() still wins the claim and starts the writer with the interrupt flag already set. In that case the listener must wait for the latch, so a test for it would protect against a future change that breaks the countdown on that path.

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Added "stopping waits for a writer that started after the interrupt": the writer is claimed by run() on a thread whose interrupt flag is already set, and stopAndAwaitExit must block until it exits.

- Add TaskContext.addTaskPreCompletionListener, whose listeners run before every task
  completion listener, and stop the pipelined writer from one. The writer thread can add
  completion listeners lazily while it pulls the upstream iterator (e.g. the DSv2 Parquet
  reader's close callback), and those ran before the writer listener.
- Do not restore the interrupt status after the wait, so the completion listeners that run
  next are not affected by it.
- Propagate fatal writer errors to the reader and log them.
- Document the trade-off of the unbounded wait and fix the comment on where the writer sees
  the interrupt.
- Make the tests wait until task completion reaches the writer listener, and cover a cleanup
  listener added by the writer thread and a writer that starts after the interrupt.

Co-authored-by: Claude Code <noreply@anthropic.com>
@viirya viirya changed the title [SPARK-60113][SQL] Wait for the pipelined Python UDF writer to exit before task completion cleanup [SPARK-60113][CORE][SQL] Wait for the pipelined Python UDF writer to exit before task completion cleanup Oct 10, 2026

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants