apache / apache/datafusion-comet

[EPIC] Memory pool and accounting audit sweep

Open
#5,212 4 comments 0 reactions 2 assignees Claimed by @comphead View on GitHub
area:ffi area:shuffle bug correctness EPIC performance priority:high
Dominant language
Scala
Stars
1.3k
Forks
373
Avg merge
2d 4h
Merged PRs (30d)
198

Description

## What / Why

Comet's memory accounting is spread across four largely independent layers — the native DataFusion `MemoryPool` implementations, the JNI bridge to Spark's `TaskMemoryManager`, the JVM shuffle allocators, and the native shuffle writer's own `MemoryReservation`. Each has been modified in isolation over time, and there are **no unit tests for any of the pools**. This EPIC collects the findings from a sweep across all four layers.

The headline item is that the **default off-heap pool silently caps a task at `pool_size / num_consumers`** (position 1 below), and the second is an **unbounded leak of per-task pool entries that pin JVM objects for the executor's lifetime** (position 2). The rest range from cross-task interference in JVM shuffle spilling down to metrics that always report zero.

Verified against `datafusion-execution` 54.1.0, `arrow-buffer` 59.1.0, and `spark-core_2.12` 3.5.7 sources.

## Scope of the sweep

- `native/core/src/execution/memory_pools/` — `fair_pool.rs`, `unified_pool.rs`, `task_shared.rs`, `config.rs`, `logging_pool.rs`, `mod.rs`
- `native/core/src/execution/jni_api.rs` — pool creation/release, tracing pool registry
- `native/jni-bridge/src/comet_task_memory_manager.rs` and `spark/src/main/java/org/apache/spark/CometTaskMemoryManager.java`
- `spark/src/main/java/org/apache/spark/shuffle/comet/` — the JVM shuffle allocators
- `SpillWriter` / `SpillSorter` / `CometShuffleExternalSorter` / `CometDiskBlockWriter`
- `native/shuffle/src/partitioners/multi_partition.rs` — the native shuffle writer's reservation

## Findings — correctness, high

### 1. `CometFairMemoryPool` compares pool-wide usage against a per-consumer limit

`native/core/src/execution/memory_pools/fair_pool.rs:151-162`

```rust
let limit = self.pool_size.checked_div(num).expect("overflow in checked_div");
let used = state.used; // pool-wide total
if limit < used + additional {
return resources_err!(...);
}
```

`limit` is the *per-consumer* fair share but `used` is the *whole pool's* usage, so a task can never use more than `pool_size / num_consumers` in total. Effective capacity shrinks linearly with every registered consumer. This is the **default off-heap pool type** (`fair_unified`, `CometConf.scala:686`), so it affects the default configuration.

The comment justifying the switch away from `reservation.size()` (`fair_pool.rs:155-157`) is incorrect:

> We use `state.used` instead of `reservation.size()` because DataFusion 53+ calls `pool.try_grow()` before incrementing the reservation's atomic size, so `reservation.size()` would not include prior grows.

`MemoryReservation::try_grow` (datafusion-execution 54.1.0, `memory_pool/mod.rs:471-475`) calls `pool.try_grow(self, capacity)` and *then* `self.size.fetch_add(capacity)`. So `reservation.size()` does include every prior grow — it excludes only the current `additional`, which is exactly why upstream `FairSpillPool` writes `reservation.size() + additional > available` (`memory_pool/pool.rs:249`). The pre-existing code was correct. The corresponding comment on `shrink` (`fair_pool.rs:128-130`) *is* correct, since `MemoryReservation::shrink` decrements before calling the pool; the two cases appear to have been conflated during the DataFusion 53 upgrade (`90633dcc4`).

Secondary divergence: `num` counts every registered consumer, whereas upstream divides by `num_spill` (spillable consumers only), making Comet's limit tighter still.

### 2. `TASK_SHARED_MEMORY_POOLS` entries leak, pinning JVM objects for the executor's lifetime

`memory_pools/task_shared.rs:25`, `memory_pools/mod.rs:57`/`:73`/`:122`, `jni_api.rs:952`

`num_plans` is incremented inside `create_memory_pool` (`jni_api.rs:435`), but `createPlan` can still fail afterwards — `local_dirs` decoding (`jni_api.rs:444-451`), `prepare_datafusion_session_context` (`:456`) and the key-unwrapper global ref (`:476`) all use `?`. On the JVM side `plan` is a field initializer (`CometExecIterator.scala:87`) evaluated *before* the task-completion listener is registered (`:139`), so a failed `createPlan` means `releasePlan` is never called and `num_plans` never returns to zero.

Symmetrically, `releasePlan` runs `update_metrics(env, execution_context)?` (`jni_api.rs:950`) *before* `handle_task_shared_pool_release` (`:952`), so a metrics failure also strands the entry — and skips the `Box::from_raw` at `:967`.

Each stranded entry holds an `Arc>` for that task's `CometTaskMemoryManager`, which transitively pins `TaskMemoryManager` and `TaskContext`. Keys are unique task-attempt ids and nothing else ever prunes the map, so it grows without bound.

Fix: move the release ahead of `update_metrics`, and make `create_memory_pool` the last fallible step in `createPlan` (or register the count via a guard).

### 3. `num_plans -= 1` underflow poisons a `std::sync::Mutex` and bricks the executor

`memory_pools/task_shared.rs:51-53`

`per_task_memory_pool.num_plans -= 1` on a `usize` panics on underflow while holding `TASK_SHARED_MEMORY_POOLS.lock().unwrap()`. Because that is a `std::sync::Mutex`, the panic poisons it and every subsequent `create_memory_pool` / release in the executor fails at `.unwrap()`. Use `saturating_sub`, and either `parking_lot::Mutex` (as `fair_pool.rs` already does) or `lock().unwrap_or_else(|e| e.into_inner())`.

### 4. `NativeMemoryConsumer.toString()` throws `UnknownFormatConversionException`

`spark/src/main/java/org/apache/spark/CometTaskMemoryManager.java:117`

```java
return String.format("NativeMemoryConsumer(id=%)", id);
```

`%)` is an invalid conversion specifier — this throws rather than formatting (verified). Reachable whenever Spark stringifies the consumer: `TaskMemoryManager.acquireExecutionMemory` logs `logger.debug("Task {} acquired {} for {}", ..., requestingConsumer)` (Spark 3.5.7 `TaskMemoryManager.java:200`), and `logger.error("error while calling spill() on " + consumerToSpill, e)` at `:248`. Under DEBUG logging on that class the exception propagates out of `acquireMemory`, back through JNI, into `CometUnifiedMemoryPool::try_grow`.

Not reachable via the `showMemoryUsage()` call at `CometTaskMemoryManager.java:74` — but only because of the next finding.

### 5. Comet native memory is invisible to Spark's own accounting

`spark/src/main/java/org/apache/spark/CometTaskMemoryManager.java:62`

```java
long acquired = internal.acquireExecutionMemory(size, nativeMemoryConsumer);
```

Calling `TaskMemoryManager.acquireExecutionMemory` directly bypasses `MemoryConsumer.acquireMemory`, so `nativeMemoryConsumer.used` stays 0 forever. Consequences:

- `showMemoryUsage()` skips the consumer entirely (its `totalMemUsage > 0` guard) and attributes all Comet native memory to the "bytes … not associated with specific consumers" bucket — precisely the diagnostic we call at `CometTaskMemoryManager.java:74` when an acquisition comes back short.
- Spark's spill-victim ordering (`TaskMemoryManager.java:175`, `key = c.getUsed()`) always sorts Comet's consumer as zero.

The class already tracks the correct number in its own `AtomicLong` (`:47`); it just never publishes it where Spark looks.

### 6. `CometDiskBlockWriter.currentWriters` is executor-wide, so spilling steals across tasks

`spark/src/main/java/org/apache/spark/sql/comet/execution/shuffle/CometDiskBlockWriter.java:66`, `:377-406`

`currentWriters` is `static final`, holding writers from *all* concurrent tasks. `ArrowIPCWriter.spill(required)` sorts that global list and spills other tasks' writers, accumulating their freed bytes into the `totalFreed >= required` check. With `CometUnifiedShuffleMemoryAllocator` each task has its own `TaskMemoryManager`, so freeing another task's pages does nothing for *this* task's acquisition — the loop breaks early believing it succeeded, having forced unrelated tasks to spill for no benefit. Should be scoped per task attempt.

### 7. `getActiveMemoryUsage()` reads a `LinkedList` cross-thread without synchronization

`CometDiskBlockWriter.java:273`, `SpillWriter.java:249-258`

`SpillWriter.getMemoryUsage()` documents "Assume this method won't be called on a spilling writer, so we don't need to synchronize it" — but the comparator in position 6 above (`CometDiskBlockWriter.java:387-388`) calls it on writers owned by other threads that may be concurrently inserting or freeing pages, risking `ConcurrentModificationException` or torn sums. `SpillSorter` overrides it with `synchronized` (`SpillSorter.java:154-166`); `ArrowIPCWriter` does not.

### 8. Double `releasePlan` on a freed pointer

`spark/src/main/scala/org/apache/comet/CometExecIterator.scala:225-245`

`closed = true` is set *last*. If `nativeLib.releasePlan(plan)` (`:233`) succeeds but `traceMemoryUsage()` (`:236`) throws, `closed` stays false and the task-completion listener (`:139`) calls `close()` again → a second `releasePlan` on a pointer already freed by `Box::from_raw` (`jni_api.rs:967`). Set `closed = true` before the teardown calls, or use `try`/`finally`.

## Findings — correctness, lower

### 9. Non-volatile cross-thread state in `CometDiskBlockWriter`

`spilling` (`:105`) and `totalWritten` (`:91`) are mutated from other tasks' threads via `doSpill()`. `totalWritten` is written under `synchronized (this)` at `:185` but read and written unguarded at `:245` and `:265`.

### 10. Dead spill-tracking list

`CometDiskBlockWriter.java:69` — `spillingWriters` is never added to, so the loop at `:278` never executes. Either populate it or delete both.

### 11. `allocator.getUsed()` always reports 0 in on-heap mode

`CometBoundedShuffleMemoryAllocator` tracks its own `allocatedMemory` field (`:58`) and never touches `MemoryConsumer.used`, so the tracing calls at `CometUnsafeShuffleWriter.java:228`/`:251` and `CometBypassMergeSortShuffleWriter.java:222`/`:241` log zero. `CometUnifiedShuffleMemoryAllocator` reports correctly (via `allocatePage`/`freePage`), so the same metric means different things in the two modes.

### 12. Peak memory lost across spills

`CometShuffleExternalSorter.java:199` calls `activeSpillSorter.freeMemory()` without `updatePeakMemoryUsed()` first, so `getPeakMemoryUsedBytes()` only ever reflects the high-water mark after the last spill.

### 13. Reservation stranded on the native shuffle spill error path

`native/shuffle/src/partitioners/multi_partition.rs:504-533` — if `partition_writer.write` fails (`:521`), `reservation.free()` (`:528`) and `pinned_buffers.clear()` (`:529`) are skipped even though `partitioned_batches()` already took ownership of the batches and indices. Additionally the comment at `:460-462` ("`try_grow` is evaluated first so the reservation accounts for this batch either way") is wrong — a failed `try_grow` does not grow the reservation.

### 14. Acquired bytes lost on overflow in the unified pool

`memory_pools/unified_pool.rs:150-160` — if `fetch_update` fails, the function returns `Err` without releasing the already-acquired bytes back to Spark. Separately, `used` is incremented by `acquired` while the caller's reservation grows by `additional`. Practically unreachable (requires `usize` overflow), but the asymmetry is trivial to remove.

### 15. On-heap shuffle allocator singleton freezes config and shares the page table

`CometShuffleMemoryAllocator.getInstance` (`:47-53`) returns a process-wide instance whose `totalMemory` is fixed from the first task's `SparkConf` and whose 8192-entry page table (`CometBoundedShuffleMemoryAllocator.PAGE_TABLE_SIZE`) is shared by all concurrent tasks, so "Have already allocated a maximum of 8192 pages" is an executor-wide rather than per-task limit.

### 16. `memory_limit()` not overridden on the Comet pools

Both Comet pools inherit the default `MemoryLimit::Unknown`, even though `CometFairMemoryPool` knows its `pool_size`. Low impact in DataFusion 54 (only `runtime_env.rs:261` display and the arrow-pool wrapper consume it), but trivial to fix.

## Findings — performance

### 17. `CometFairMemoryPool` holds its mutex across the JNI round-trip

`memory_pools/fair_pool.rs:149-183` — `state.lock()` is taken at `:149` and held through `self.acquire(additional)` at `:165`, a JVM call that can block on Spark's `synchronized` `MemoryManager` and may itself trigger spills. Because the pool is *task-shared* across Comet's tokio workers, this serializes the whole task's memory traffic behind a lock held across a JVM call. It is also a non-reentrant `parking_lot::Mutex`, so this becomes a deadlock if any Comet `MemoryConsumer` ever spills back through the pool. `CometUnifiedMemoryPool` does the same job with a lock-free `AtomicUsize`; the fair pool could do the limit check under the lock, drop it, then acquire.

### 18. One JNI round-trip per grow/shrink, with no batching or hysteresis

`memory_pools/unified_pool.rs:64-78` — DataFusion calls `try_grow`/`shrink` per batch in aggregates, joins and sorts; each becomes a JNI call plus contention on Spark's executor-wide `MemoryManager` lock, and every `shrink` releases immediately however small. Acquiring in chunks with slack and releasing with hysteresis is likely the highest-leverage optimization in this area.

### 19. Tracing double-counts task-shared pools and takes a global lock per batch

`jni_api.rs:437-441` wraps each context's pool in its own `LoggingMemoryPool` when `COMET_DEBUG_MEMORY` is set, defeating the `Arc::as_ptr` dedup in `total_reserved_for_thread` (`:182`) and `unregister_and_total` (`:164`), so a task-shared pool is counted once per plan. Separately, `total_reserved_for_thread` takes a global mutex and walks the map on every `executePlan` (`:896`). Both are debug/tracing-only paths.

## Test coverage gap

There are **no unit tests for any memory pool** — no `#[cfg(test)]` module anywhere under `native/core/src/execution/memory_pools/` — and no JVM tests for the shuffle allocators or `CometTaskMemoryManager` (only `SpillSorterSuite.scala` exists in this area).

Position 1 is exactly the kind of regression a handful of pool tests would have caught. A minimal suite for each pool: register N consumers, assert a single consumer can grow to `pool_size / N`, and assert N consumers can collectively reach `pool_size`.

## Suggested sequencing

1. Positions 1, 2, 3, 4, 5 — small, independent, all on the default off-heap path — plus the pool unit tests above.
2. Positions 6, 7, 9, 10 — the `CometDiskBlockWriter` cross-task spilling and its associated data races, likely one change.
3. Positions 8, 11, 12, 13, 14, 15, 16 — independent small fixes, good first issues.
4. Positions 17, 18, 19 — performance work, needs a benchmark to demonstrate the win.

## Not investigated

- Whether native pool sizing and the JVM shuffle allocator's budget should be coordinated. In on-heap mode `memoryLimit` (`CometExecIterator.scala:306`) is the full `COMET_ONHEAP_MEMORY_OVERHEAD` while `CometBoundedShuffleMemoryAllocator.totalMemory` is `COMET_SHUFFLE_JVM_MEMORY_FACTOR` times the same overhead, so the two can jointly exceed the nominal budget.
- Whether `memoryLimitPerTask = memoryLimit * coresPerTask / numCores` (`CometExecIterator.scala:309`) can round to 0 in practice, which would make `GreedyMemoryPool::new(0)` fail every allocation.
- Behaviour when a task runs multiple native plans: the task-shared pool holds a global ref to the *first* plan's `CometTaskMemoryManager`, so later plans' memory is accounted through the first manager. `cometTaskMemoryManager.getUsed` at `CometExecIterator.scala:239` will then over-report for the first iterator and report 0 for the others, which may explain spurious "closed with non-zero memory usage" warnings.

Contributor guide

Open the contributing guide

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.