This is an automated email from the ASF dual-hosted git repository.
zaynt4606 pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/celeborn.git
The following commit(s) were added to refs/heads/main by this push:
new fd453c074a [CELEBORN-2428] Fix TierWriter buffer accounting when
CompositeByteBuf insertion fails
fd453c074a is described below
commit fd453c074ad0f997e98d7c4546a7412f6e0133d8
Author: Kalvin2077 <[email protected]>
AuthorDate: Tue Sep 15 11:21:12 2026 +0800
[CELEBORN-2428] Fix TierWriter buffer accounting when CompositeByteBuf
insertion fails
### What changes were proposed in this pull request?
Fix `TierWriter` buffer accounting and reference-count handling when
`CompositeByteBuf.addComponent` fails before insertion or during consolidation.
Add regression tests and include the shuffle key in write-error logs.
### Why are the changes needed?
`addComponent` can fail before or after taking ownership of the buffer.
Treating both cases identically can cause incorrect storage counters or buffer
reference-count leaks.
### Does this PR resolve a correctness bug?
- [ ] Yes
### Does this PR introduce _any_ user-facing change?
- [ ] Yes
### How was this patch tested?
- Added tests for pre-insertion OOM, capacity overflow, and post-insertion
consolidation OOM.
- Ran `TierWriterSuite`: 11 tests passed.
- Ran `./build/mvn spotless:check -pl worker`.
Closes #3809 from Kalvin2077/fix/CELEBORN-2428.
Authored-by: Kalvin2077 <[email protected]>
Signed-off-by: zhengtao <[email protected]>
AI-Contributed/Feature: 0/69
AI-Contributed/UT: 0/255
---
.../service/deploy/worker/PushDataHandler.scala | 2 +-
.../service/deploy/worker/storage/TierWriter.scala | 67 +++---
.../deploy/worker/storage/TierWriterSuite.scala | 255 ++++++++++++++++++++-
3 files changed, 293 insertions(+), 31 deletions(-)
diff --git
a/worker/src/main/scala/org/apache/celeborn/service/deploy/worker/PushDataHandler.scala
b/worker/src/main/scala/org/apache/celeborn/service/deploy/worker/PushDataHandler.scala
index 0accb33adb..f6c32e77f3 100644
---
a/worker/src/main/scala/org/apache/celeborn/service/deploy/worker/PushDataHandler.scala
+++
b/worker/src/main/scala/org/apache/celeborn/service/deploy/worker/PushDataHandler.scala
@@ -1546,7 +1546,7 @@ class PushDataHandler(val workerSource: WorkerSource)
extends BaseMessageHandler
workerSource.incCounter(WorkerSource.WRITE_DATA_HARD_SPLIT_COUNT)
result(index) = StatusCode.HARD_SPLIT
} else {
- logError("Exception encountered when write.", e)
+ logError(s"Exception encountered when writing shuffle
$shuffleKey.", e)
workerSource.incCounter(WorkerSource.WRITE_DATA_FAIL_COUNT)
val cause =
if (isPrimary) {
diff --git
a/worker/src/main/scala/org/apache/celeborn/service/deploy/worker/storage/TierWriter.scala
b/worker/src/main/scala/org/apache/celeborn/service/deploy/worker/storage/TierWriter.scala
index a04f1d6761..6a39236fcd 100644
---
a/worker/src/main/scala/org/apache/celeborn/service/deploy/worker/storage/TierWriter.scala
+++
b/worker/src/main/scala/org/apache/celeborn/service/deploy/worker/storage/TierWriter.scala
@@ -24,6 +24,7 @@ import java.util.concurrent.TimeUnit
import java.util.concurrent.atomic.AtomicInteger
import scala.collection.JavaConverters.asScalaBufferConverter
+import scala.util.control.NonFatal
import io.netty.buffer.{ByteBuf, CompositeByteBuf}
import org.apache.hadoop.fs.{FileSystem, FSDataOutputStream}
@@ -88,6 +89,32 @@ abstract class TierWriterBase(
protected def writeInternal(buf: ByteBuf): Unit
+ protected def accountAddedBytes(numBytes: Int): Unit
+
+ protected def appendToFlushBuffer(buf: ByteBuf): Unit = {
+ val numBytes = buf.readableBytes()
+ val writerIndexBefore = flushBuffer.writerIndex()
+ val numComponentsBefore = flushBuffer.numComponents()
+ buf.retain()
+ try {
+ flushBuffer.addComponent(true, buf)
+ } catch {
+ case t: Throwable if t.isInstanceOf[OutOfMemoryError] || NonFatal(t) =>
+ // addComponent can fail after taking ownership while consolidating
components.
+ val componentAdded =
+ flushBuffer.writerIndex() != writerIndexBefore ||
+ flushBuffer.numComponents() != numComponentsBefore
+ if (componentAdded) {
+ accountAddedBytes(numBytes)
+ }
+ logError(
+ s"Failed to add $numBytes bytes to flush buffer for shuffle
$shuffleKey file $filename.",
+ t)
+ throw t
+ }
+ accountAddedBytes(numBytes)
+ }
+
def needEvict(): Boolean
def evict(file: TierWriterBase): Unit
@@ -322,25 +349,15 @@ class MemoryTierWriter(
// Memory file won't produce flush task
override def genFlushTask(finalFlush: Boolean, keepBuffer: Boolean):
FlushTask = null
- override def writeInternal(buf: ByteBuf): Unit = {
- buf.retain()
- val numBytes = buf.readableBytes()
- try {
- flushBuffer.addComponent(true, buf)
- } catch {
- case oom: OutOfMemoryError =>
- // memory tier writer will not flush
- // add the bytes into flusher buffer is flush completed
- metaHandler.afterFlush(numBytes)
- MemoryManager.instance.incrementMemoryFileStorage(numBytes)
- throw oom
- }
+ override protected def accountAddedBytes(numBytes: Int): Unit = {
// memory tier writer will not flush
// add the bytes into flusher buffer is flush completed
metaHandler.afterFlush(numBytes)
MemoryManager.instance().incrementMemoryFileStorage(numBytes)
}
+ override def writeInternal(buf: ByteBuf): Unit = appendToFlushBuffer(buf)
+
override def closeStreams(): Unit = {
try {
flushBuffer.consolidate()
@@ -437,14 +454,10 @@ class LocalTierWriter(
if (flushBufferReadableBytes != 0 && flushBufferReadableBytes + numBytes
>= flusherBufferSize) {
flush(false)
}
- buf.retain()
- try {
- flushBuffer.addComponent(true, buf)
- } catch {
- case oom: OutOfMemoryError =>
- MemoryManager.instance.incrementDiskBuffer(numBytes)
- throw oom
- }
+ appendToFlushBuffer(buf)
+ }
+
+ override protected def accountAddedBytes(numBytes: Int): Unit = {
MemoryManager.instance.incrementDiskBuffer(numBytes)
}
@@ -659,14 +672,10 @@ class DfsTierWriter(
if (flushBufferReadableBytes != 0 && flushBufferReadableBytes + numBytes
>= flusherBufferSize) {
flush(false)
}
- buf.retain()
- try {
- flushBuffer.addComponent(true, buf)
- } catch {
- case oom: OutOfMemoryError =>
- MemoryManager.instance.incrementDiskBuffer(numBytes)
- throw oom
- }
+ appendToFlushBuffer(buf)
+ }
+
+ override protected def accountAddedBytes(numBytes: Int): Unit = {
MemoryManager.instance.incrementDiskBuffer(numBytes)
}
diff --git
a/worker/src/test/scala/org/apache/celeborn/service/deploy/worker/storage/TierWriterSuite.scala
b/worker/src/test/scala/org/apache/celeborn/service/deploy/worker/storage/TierWriterSuite.scala
index ee6903ddf6..ba40bd67f7 100644
---
a/worker/src/test/scala/org/apache/celeborn/service/deploy/worker/storage/TierWriterSuite.scala
+++
b/worker/src/test/scala/org/apache/celeborn/service/deploy/worker/storage/TierWriterSuite.scala
@@ -21,7 +21,7 @@ import java.io.IOException
import java.nio.file.Files
import java.util.concurrent.atomic.{AtomicInteger, AtomicLong}
-import io.netty.buffer.UnpooledByteBufAllocator
+import io.netty.buffer.{AbstractByteBufAllocator, ByteBuf, CompositeByteBuf,
DuplicatedByteBuf, UnpooledByteBufAllocator}
import org.mockito.Mockito
import org.mockito.MockitoSugar.when
import org.scalatest.BeforeAndAfterEach
@@ -37,6 +37,63 @@ import org.apache.celeborn.service.deploy.worker.WorkerSource
import org.apache.celeborn.service.deploy.worker.memory.MemoryManager
class TierWriterSuite extends AnyFunSuite with BeforeAndAfterEach {
+ private class LargeReadableByteBuf(buffer: ByteBuf, virtualSize: Int)
+ extends DuplicatedByteBuf(buffer) {
+ override def capacity(): Int = virtualSize
+
+ override def readableBytes(): Int = virtualSize
+ }
+
+ private class PreInsertionOomByteBuf(buffer: ByteBuf, failure:
OutOfMemoryError)
+ extends DuplicatedByteBuf(buffer) {
+ override def readerIndex(): Int = throw failure
+ }
+
+ private class OomByteBufAllocator(failure: OutOfMemoryError)
+ extends AbstractByteBufAllocator(false) {
+ override protected def newHeapBuffer(initialCapacity: Int, maxCapacity:
Int): ByteBuf =
+ throw failure
+
+ override protected def newDirectBuffer(initialCapacity: Int, maxCapacity:
Int): ByteBuf =
+ throw failure
+
+ override def isDirectBufferPooled(): Boolean = false
+ }
+
+ private def capacityOverflowFlushBuffer(): CompositeByteBuf = {
+ val flushBuffer = new CompositeByteBuf(
+ UnpooledByteBufAllocator.DEFAULT,
+ false,
+ Int.MaxValue)
+ val component = UnpooledByteBufAllocator.DEFAULT.buffer(1,
Int.MaxValue).writeByte(0)
+ flushBuffer.addComponent(
+ true,
+ new LargeReadableByteBuf(component, Int.MaxValue - 512))
+ flushBuffer
+ }
+
+ private def consolidationOomFlushBuffer(failure: OutOfMemoryError):
CompositeByteBuf = {
+ val flushBuffer = new CompositeByteBuf(new OomByteBufAllocator(failure),
false, 1)
+ flushBuffer.addComponent(true, UnpooledByteBufAllocator.DEFAULT.buffer(0,
0))
+ flushBuffer
+ }
+
+ private def restoreCounters(memoryCounterBefore: Long, diskCounterBefore:
Long): Unit = {
+ val memoryManager = MemoryManager.instance()
+ val memoryCounterDelta = memoryManager.getMemoryFileStorageCounter -
memoryCounterBefore
+ if (memoryCounterDelta > 0) {
+
memoryManager.releaseMemoryFileStorage(Math.toIntExact(memoryCounterDelta))
+ } else if (memoryCounterDelta < 0) {
+
memoryManager.incrementMemoryFileStorage(Math.toIntExact(-memoryCounterDelta))
+ }
+ val diskCounterDelta = memoryManager.getDiskBufferCounter.get() -
diskCounterBefore
+ if (diskCounterDelta > 0) {
+ memoryManager.releaseDiskBuffer(Math.toIntExact(diskCounterDelta))
+ } else if (diskCounterDelta < 0) {
+ memoryManager.incrementDiskBuffer(Math.toIntExact(-diskCounterDelta))
+ }
+ }
+
private def prepareMemoryWriter: MemoryTierWriter = {
val celebornConf = new CelebornConf()
@@ -333,4 +390,200 @@ class TierWriterSuite extends AnyFunSuite with
BeforeAndAfterEach {
assert(fileLen == 10240)
assert(localTierWriter.closed === true)
}
+
+ test("memory tier writer should preserve state on pre-insertion OOM") {
+ val memoryTierWriter = prepareMemoryWriter
+ val failure = new OutOfMemoryError("component creation failed")
+ val buf = new PreInsertionOomByteBuf(
+ WriterUtils.generateSparkFormatData(UnpooledByteBufAllocator.DEFAULT, 0),
+ failure)
+ val refCntBeforeWrite = buf.refCnt()
+ val memoryManager = MemoryManager.instance()
+ val memoryCounterBefore = memoryManager.getMemoryFileStorageCounter
+ val diskCounterBefore = memoryManager.getDiskBufferCounter.get()
+ val fileLengthBefore = memoryTierWriter.fileInfo.getFileLength
+ val writerIndexBefore = memoryTierWriter.flushBuffer.writerIndex()
+ val numComponentsBefore = memoryTierWriter.flushBuffer.numComponents()
+ var callerReferenceReleased = false
+
+ try {
+ memoryTierWriter.numPendingWrites.incrementAndGet()
+ val thrown = intercept[OutOfMemoryError](memoryTierWriter.write(buf))
+
+ assert(thrown eq failure)
+ assert(memoryTierWriter.flushBuffer.writerIndex() === writerIndexBefore)
+ assert(memoryTierWriter.flushBuffer.numComponents() ===
numComponentsBefore)
+ assert(memoryManager.getMemoryFileStorageCounter === memoryCounterBefore)
+ assert(memoryManager.getDiskBufferCounter.get() === diskCounterBefore)
+ assert(memoryTierWriter.fileInfo.getFileLength === fileLengthBefore)
+ assert(buf.refCnt() === refCntBeforeWrite)
+ assert(buf.release())
+ callerReferenceReleased = true
+ assert(buf.refCnt() === 0)
+ } finally {
+ memoryTierWriter.destroy(new IOException("test cleanup"))
+ restoreCounters(memoryCounterBefore, diskCounterBefore)
+ if (!callerReferenceReleased && buf.refCnt() > 0) {
+ buf.release(buf.refCnt())
+ }
+ }
+ }
+
+ test("memory tier writer should preserve state on composite buffer capacity
overflow") {
+ val memoryTierWriter = prepareMemoryWriter
+ val buf =
WriterUtils.generateSparkFormatData(UnpooledByteBufAllocator.DEFAULT, 0)
+ val refCntBeforeWrite = buf.refCnt()
+ val memoryManager = MemoryManager.instance()
+ val memoryCounterBefore = memoryManager.getMemoryFileStorageCounter
+ val diskCounterBefore = memoryManager.getDiskBufferCounter.get()
+ val fileLengthBefore = memoryTierWriter.fileInfo.getFileLength
+ val overflowBuffer = capacityOverflowFlushBuffer()
+ val writerIndexBefore = overflowBuffer.writerIndex()
+ val numComponentsBefore = overflowBuffer.numComponents()
+ val originalFlushBuffer = memoryTierWriter.flushBuffer
+ memoryTierWriter.flushBuffer = overflowBuffer
+ originalFlushBuffer.release()
+ var callerReferenceReleased = false
+
+ try {
+ memoryTierWriter.numPendingWrites.incrementAndGet()
+ val thrown =
intercept[IllegalArgumentException](memoryTierWriter.write(buf))
+
+ assert(thrown.getMessage.contains("overflow"))
+ assert(overflowBuffer.writerIndex() === writerIndexBefore)
+ assert(overflowBuffer.numComponents() === numComponentsBefore)
+ assert(memoryManager.getMemoryFileStorageCounter === memoryCounterBefore)
+ assert(memoryTierWriter.fileInfo.getFileLength === fileLengthBefore)
+ assert(buf.refCnt() === refCntBeforeWrite)
+ assert(buf.release())
+ callerReferenceReleased = true
+ assert(buf.refCnt() === 0)
+ } finally {
+ memoryTierWriter.destroy(new IOException("test cleanup"))
+ restoreCounters(memoryCounterBefore, diskCounterBefore)
+ if (!callerReferenceReleased && buf.refCnt() > 0) {
+ buf.release(buf.refCnt())
+ }
+ }
+ }
+
+ test("memory tier writer should account data added before consolidation
OOM") {
+ val memoryTierWriter = prepareMemoryWriter
+ val memoryFileInfo = memoryTierWriter.fileInfo.asInstanceOf[MemoryFileInfo]
+ val buf =
WriterUtils.generateSparkFormatData(UnpooledByteBufAllocator.DEFAULT, 0)
+ val numBytes = buf.readableBytes()
+ val refCntBeforeWrite = buf.refCnt()
+ val memoryManager = MemoryManager.instance()
+ val memoryCounterBefore = memoryManager.getMemoryFileStorageCounter
+ val diskCounterBefore = memoryManager.getDiskBufferCounter.get()
+ val fileLengthBefore = memoryTierWriter.fileInfo.getFileLength
+ val failure = new OutOfMemoryError("consolidation failed")
+ val failingBuffer = consolidationOomFlushBuffer(failure)
+ val numComponentsBefore = failingBuffer.numComponents()
+ val originalFlushBuffer = memoryTierWriter.flushBuffer
+ memoryTierWriter.flushBuffer = failingBuffer
+ originalFlushBuffer.release()
+ var memoryBufferReleased = false
+ var callerReferenceReleased = false
+
+ try {
+ memoryTierWriter.numPendingWrites.incrementAndGet()
+ val thrown = intercept[OutOfMemoryError](memoryTierWriter.write(buf))
+
+ assert(thrown eq failure)
+ assert(failingBuffer.writerIndex() === numBytes)
+ assert(failingBuffer.readableBytes() === numBytes)
+ assert(failingBuffer.numComponents() === numComponentsBefore + 1)
+ assert(memoryManager.getMemoryFileStorageCounter === memoryCounterBefore
+ numBytes)
+ assert(memoryTierWriter.fileInfo.getFileLength === fileLengthBefore +
numBytes)
+ assert(buf.refCnt() === refCntBeforeWrite + 1)
+
+ memoryTierWriter.numPendingWrites.set(0)
+ assert(memoryTierWriter.close() === fileLengthBefore + numBytes)
+ val releasedBytes = memoryFileInfo.releaseMemoryBuffers()
+ memoryBufferReleased = true
+ memoryManager.releaseMemoryFileStorage(releasedBytes)
+ assert(releasedBytes === numBytes)
+ assert(memoryManager.getMemoryFileStorageCounter === memoryCounterBefore)
+ assert(buf.refCnt() === refCntBeforeWrite)
+ assert(buf.release())
+ callerReferenceReleased = true
+ assert(buf.refCnt() === 0)
+ } finally {
+ if (!memoryTierWriter.closed) {
+ memoryTierWriter.destroy(new IOException("test cleanup"))
+ } else if (!memoryBufferReleased && memoryFileInfo.getBuffer.refCnt() >
0) {
+ val releasedBytes = memoryFileInfo.releaseMemoryBuffers()
+ memoryManager.releaseMemoryFileStorage(releasedBytes)
+ }
+ restoreCounters(memoryCounterBefore, diskCounterBefore)
+ if (!callerReferenceReleased && buf.refCnt() > 0) {
+ buf.release(buf.refCnt())
+ }
+ }
+ }
+
+ test("local tier writer should account data added before consolidation OOM")
{
+ val localTierWriter = prepareLocalTierWriter(false)
+ val buf =
WriterUtils.generateSparkFormatData(UnpooledByteBufAllocator.DEFAULT, 0)
+ val numBytes = buf.readableBytes()
+ val refCntBeforeWrite = buf.refCnt()
+ val memoryManager = MemoryManager.instance()
+ val memoryCounterBefore = memoryManager.getMemoryFileStorageCounter
+ val diskCounterBefore = memoryManager.getDiskBufferCounter.get()
+ val failure = new OutOfMemoryError("consolidation failed")
+ val failingBuffer = consolidationOomFlushBuffer(failure)
+ val numComponentsBefore = failingBuffer.numComponents()
+ val originalFlushBuffer = localTierWriter.flushBuffer
+ localTierWriter.flushBuffer = failingBuffer
+ localTierWriter.getFlusher.returnBuffer(originalFlushBuffer, false)
+ var flushBufferReleased = false
+ var callerReferenceReleased = false
+
+ try {
+ localTierWriter.numPendingWrites.incrementAndGet()
+ val thrown = intercept[OutOfMemoryError](localTierWriter.write(buf))
+
+ assert(thrown eq failure)
+ assert(failingBuffer.writerIndex() === numBytes)
+ assert(failingBuffer.readableBytes() === numBytes)
+ assert(failingBuffer.numComponents() === numComponentsBefore + 1)
+ assert(memoryManager.getDiskBufferCounter.get() === diskCounterBefore +
numBytes)
+ assert(buf.refCnt() === refCntBeforeWrite + 1)
+
+ localTierWriter.returnBuffer(false)
+ assert(memoryManager.getDiskBufferCounter.get() === diskCounterBefore)
+ assert(failingBuffer.writerIndex() === 0)
+ assert(failingBuffer.numComponents() === 0)
+ assert(buf.refCnt() === refCntBeforeWrite)
+ val returnedBuffer = localTierWriter.getFlusher.takeBuffer()
+ assert(returnedBuffer eq failingBuffer)
+ assert(returnedBuffer.release())
+ flushBufferReleased = true
+
+ localTierWriter.numPendingWrites.set(0)
+ localTierWriter.close()
+ assert(buf.release())
+ callerReferenceReleased = true
+ assert(buf.refCnt() === 0)
+ } finally {
+ if (!flushBufferReleased) {
+ if (localTierWriter.flushBuffer != null) {
+ localTierWriter.flushBuffer.release()
+ localTierWriter.flushBuffer = null
+ } else if (failingBuffer.refCnt() > 0) {
+ val returnedBuffer = localTierWriter.getFlusher.takeBuffer()
+ returnedBuffer.release()
+ }
+ }
+ localTierWriter.numPendingWrites.set(0)
+ if (!localTierWriter.closed) {
+ localTierWriter.close()
+ }
+ restoreCounters(memoryCounterBefore, diskCounterBefore)
+ if (!callerReferenceReleased && buf.refCnt() > 0) {
+ buf.release(buf.refCnt())
+ }
+ }
+ }
}