Repository navigation
Conversation
…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
left a comment
There was a problem hiding this comment.
Thank you for working on this, @viirya. I left 8 inline comments. Here is a summary.
- (
PythonRunner.scalaL203) The wait only covers listeners registered beforestartPipelinedWriter. Listeners that the writer thread itself registers lazily while pulling upstream input (e.g.ParquetReaderCallback.initIfNotAlreadyof the DSv2 Parquet reader, orUnsafeExternalSorter'scleanupResources()whenExternalAppendOnlyUnsafeRowArrayswitches 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. - (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 normalLIMITcompletion. - (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.
- (
BasePythonRunnerSuite.scalaL94) 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. - (
BasePythonRunnerSuite.scalaL141) The race where the pool thread entersFutureTask.runand wins the claim aftercancel(true)is not tested. - (L240) The hand-written interrupt-preserving loop can be replaced with Guava
Uninterruptibles.awaitUninterruptibly, which is already used incore. - (L204) A fatal error thrown by the writer is captured by the
FutureTaskand never logged or propagated. - (L231) The comment says the writer checks the interrupt flag between rows, but it checks only between
writeNextInputToStreamcalls (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] { _ => |
There was a problem hiding this comment.
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.initIfNotAlreadyregisterscloseCurrent()when the first file is opened (ParquetPartitionReaderFactory.scalaL391), which closes the vectorized reader (off-heap column vectors whenspark.sql.columnVector.offheap.enabled=true). ExternalAppendOnlyUnsafeRowArray(SMJ / Window upstream) creates anUnsafeExternalSorterlazily once the threshold is exceeded, and its constructor registerscleanupResources().
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?
There was a problem hiding this comment.
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)) |
There was a problem hiding this comment.
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)?
There was a problem hiding this comment.
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. |
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
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. |
There was a problem hiding this comment.
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?
There was a problem hiding this comment.
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() |
There was a problem hiding this comment.
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(...)
}There was a problem hiding this comment.
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) { |
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
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) |
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
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, ()) |
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
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>
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.Add
TaskContext.addTaskPreCompletionListener(private[spark]).TaskContextImplruns these listeners before every listener added withaddTaskCompletionListener, regardless of when or from which thread that listener was added:invokeListenersalways takes the next listener from the pre-completion stack first. The default implementation inTaskContextfalls back toaddTaskCompletionListener, andBarrierTaskContextdelegates.The writer is now wrapped in
BasePythonRunner.PipelinedWriterTask, which counts down a latch in afinallyaround the whole writer. The listener interrupts the writer withcancel(true)as before, and then waits on that latch instead of callingFuture.get(). If the listener runs before the pool has started the writer, anAtomicBooleanclaim makes sure the writer is skipped, so the listener does not wait on a latch that would never be counted down.The listener is added with
addTaskPreCompletionListener.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 theTaskContext.A fatal error of the writer (e.g.
OutOfMemoryError) is now also passed to the reader throughwriter.setExceptionwith a socketshutdownOutput(), and it is logged; previously the thread pool'sFutureTaskkept it and the reader waited for the Python worker.Why are the changes needed?
The existing listener does:
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()throwsCancellationExceptionright away and does not wait for the interrupted runnable to return. So the writer can still be serializing a row, or adding it to aHybridRowQueuepage, 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. AHybridQueue.createNewQueueracing withclose()can also add a page that nobody closes untilcleanUpAllAllocatedMemory.A task completes while the writer is still running when the output is not fully consumed, for example with a
LIMITor 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.initIfNotAlreadyof the DSv2 Parquet reader (which closes the vectorized reader and its off-heap column vectors) or theUnsafeExternalSorterthatExternalAppendOnlyUnsafeRowArraycreates 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, whoseInMemoryRowQueuewrites through the page's base object captured at construction, so a freed page would be corrupted even on-heap (HeapMemoryAllocatorpools 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.writefails withClosedByInterruptException, and otherwise the writer loop checks the interrupt flag betweenwriteNextInputToStreamcalls, 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:run()on a thread whose interrupt flag was already set.New tests in
TaskContextSuitefor 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
markTaskCompletedreturning 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.