security: remove pickle fallback to prevent arbitrary code execution (fixes #36) (#131)

Thanks for this. The fixture is safetensors now and the Needle 2 hint is back, and the suite is green on my machine (145 passed). Merging.
This commit is contained in:
DPR
2026-09-22 22:26:32 -07:00
committed by GitHub
parent 6e33cbd42e
commit 1e2072f934
7 changed files with 179 additions and 93 deletions
+4 -4
View File
@@ -185,7 +185,7 @@ agent = needle.Needle(weights="my.cact", tools=[...])
agent.run("...")
```
Options: `finetune --lora-rank 16 --lora-alpha 32 --lr 1e-4 --batch-size 16 --max-len 1024 --val-split 0.1 --checkpoint <base.safetensors or .pkl> --checkpoint-dir <dir> --out <adapter.safetensors or .pkl> --generate <n> --model <id> --workers <n>`; `--model` defaults to `deepseek/deepseek-flash-latest`, `--workers` defaults to `8`, and `--checkpoint-dir` defaults to `checkpoints`. `--generate` synthesizes extra examples through the configured OpenRouter endpoint before training. `build --bits 2|4 --upload` (with `NEEDLE_HF_REPO=<you>/<model>`).
Options: `finetune --lora-rank 16 --lora-alpha 32 --lr 1e-4 --batch-size 16 --max-len 1024 --val-split 0.1 --checkpoint <base.safetensors> --checkpoint-dir <dir> --out <adapter.safetensors> --generate <n> --model <id> --workers <n>`; `--model` defaults to `deepseek/deepseek-flash-latest`, `--workers` defaults to `8`, and `--checkpoint-dir` defaults to `checkpoints`. `--generate` synthesizes extra examples through the configured OpenRouter endpoint before training. `build --bits 2|4 --upload` (with `NEEDLE_HF_REPO=<you>/<model>`).
## Platform (hosted fine-tuning)
@@ -229,7 +229,7 @@ Adapt by swapping the Literal values for your own and keeping the shapes: closed
## CLI summary
- `needle run --checkpoint <base.safetensors or .pkl> --query "..." --tools tools.json` - JAX reference inference from a checkpoint (dev path; normal inference is the Python `Needle` API above).
- `needle run --checkpoint <base.safetensors> --query "..." --tools tools.json` - JAX reference inference from a checkpoint (dev path; normal inference is the Python `Needle` API above).
- `needle generate-data` / `needle finetune` / `needle build [checkpoint] [--lora adapter] [--layers N] [--platform <folder>] [--out <path>]` - the fine-tuning pipeline; fine-tuning runs at the full 20 layers, `build` merges the adapter, slices the `--layers` rung (default 20), exports at 4 bits, and fetches the published `needle3.cact` on every run (without an adapter at full depth it ships that published 2-bit archive as is). `--platform` also downloads that platform's engine and header and makes `--out` a folder holding them plus `needle3.cact`.
- `needle playground` - the browser UI locally (hosted at https://cactuscompute.com/needle).
- `needle fetch [--platform-tag <tag>] [--out <dir>]` - pre-download the inference engine for this machine (or another platform) into the cache; prints the path. For air-gapped devices, also see `NEEDLE3_LIB_PATH` and `HF_HUB_OFFLINE` at https://cactuscompute.com/blog/needle-python-docs.
@@ -242,8 +242,8 @@ Adapt by swapping the Literal values for your own and keeping the shapes: closed
- Do not read `arguments` keys that were not evidenced in the input; optional fields may be absent.
- Do not create one `Needle` per turn for a conversation; reuse the instance so context carries. New tools = new instance.
- `import needle` does not import JAX; only `needle finetune`/`build` do.
- `weights=` expects a `.cact` (from `needle build`), not a `.safetensors` or `.pkl` checkpoint. `Needle(tools=...)` runs the base Needle 3 archive (fetched once, cached next to the engine); pass `generation=2` for the embedded Needle 2 engine.
- Checkpoints and adapters are `.safetensors`; pickle checkpoints still load. The base is `checkpoints/needle3.safetensors` (20 layers); `needle build --layers N` exports any rung from 2 to 20 of it.
- `weights=` expects a `.cact` (from `needle build`), not a `.safetensors` checkpoint. `Needle(tools=...)` runs the base Needle 3 archive (fetched once, cached next to the engine); pass `generation=2` for the embedded Needle 2 engine.
- Checkpoints and adapters must be `.safetensors`; pickle (`.pkl`) is not supported and will raise a `ValueError`. The base is `checkpoints/needle3.safetensors` (20 layers); `needle build --layers N` exports any rung from 2 to 20 of it.
- The engine cannot unload weights: once a tuned `.cact` is bound, constructing or calling a base-model agent raises instead of silently answering with those weights. Construct agents that want the base model before any tuned one, or run them in separate processes.
## Telemetry
+3 -3
View File
@@ -39,7 +39,7 @@ def _download_target(spec):
name = spec[:-5] if spec.endswith(".cact") else spec
if name in ("needle2", "needle3"):
return "base", int(name[-1])
if spec.endswith((".safetensors", ".pkl")):
if spec.endswith(".safetensors"):
return "checkpoint", spec
raise SystemExit(
f"unknown download {spec!r}: pass needle3 (base weights), needle3.safetensors "
@@ -175,7 +175,7 @@ def main():
help="Concurrent OpenRouter requests when generating (default: 8)")
p.add_argument("--checkpoint-dir", type=str, default="checkpoints")
p.add_argument("--out", type=str, default=None,
help="Output adapter path (.safetensors, or .pkl)")
help="Output adapter path (.safetensors)")
p = sub.add_parser("generate-data")
p.add_argument("--tools", type=str, default=None, help="Tool schemas JSON to seed generation")
@@ -189,7 +189,7 @@ def main():
p = sub.add_parser("build")
p.add_argument("checkpoint", type=str, nargs="?", default=None,
help="Base checkpoint (.safetensors or .pkl); defaults to the adapter's "
help="Base checkpoint (.safetensors); defaults to the adapter's "
"base, else the Needle 3 base (auto-downloads)")
p.add_argument("--lora", type=str, default=None, help="LoRA adapter to merge before export")
p.add_argument("--out", type=str, default=None, help="Output .cact path")
+46 -53
View File
@@ -1,6 +1,5 @@
import json
import os
import pickle
import numpy as np
@@ -60,69 +59,63 @@ def _int_or_none(text):
def read_checkpoint(path):
if is_safetensors(path):
library = _safetensors_numpy()
from safetensors import safe_open
if not is_safetensors(path):
raise ValueError(f"checkpoint must be a .safetensors file for security (got {path})")
library = _safetensors_numpy()
from safetensors import safe_open
with safe_open(os.fspath(path), framework="np") as handle:
metadata = handle.metadata() or {}
return {
"format_version": _int_or_none(metadata.get("format_version")),
"params": unflatten(library.load_file(os.fspath(path))),
"config": _json_or_default(metadata.get("config"), {}),
"step": _int_or_none(metadata.get("step")),
"run": _json_or_default(metadata.get("run"), {}),
}
with open(path, "rb") as handle:
return pickle.load(handle)
with safe_open(os.fspath(path), framework="np") as handle:
metadata = handle.metadata() or {}
return {
"format_version": _int_or_none(metadata.get("format_version")),
"params": unflatten(library.load_file(os.fspath(path))),
"config": _json_or_default(metadata.get("config"), {}),
"step": _int_or_none(metadata.get("step")),
"run": _json_or_default(metadata.get("run"), {}),
}
def write_checkpoint(path, checkpoint):
if is_safetensors(path):
library = _safetensors_numpy()
metadata = {
"format_version": str(checkpoint.get("format_version", "")),
"config": json.dumps(checkpoint.get("config", {}), default=str),
"step": "" if checkpoint.get("step") is None else str(checkpoint["step"]),
"run": json.dumps(checkpoint.get("run") or {}, default=str),
}
library.save_file(flatten(checkpoint["params"]), os.fspath(path), metadata=metadata)
return
with open(path, "wb") as handle:
pickle.dump(checkpoint, handle)
if not is_safetensors(path):
raise ValueError(f"checkpoint must be a .safetensors file for security (got {path})")
library = _safetensors_numpy()
metadata = {
"format_version": str(checkpoint.get("format_version", "")),
"config": json.dumps(checkpoint.get("config", {}), default=str),
"step": "" if checkpoint.get("step") is None else str(checkpoint["step"]),
"run": json.dumps(checkpoint.get("run") or {}, default=str),
}
library.save_file(flatten(checkpoint["params"]), os.fspath(path), metadata=metadata)
_ADAPTER_FIELDS = ("scale", "base", "rank", "seed")
def write_adapter(path, adapter):
if is_safetensors(path):
library = _safetensors_numpy()
tensors = {}
for name, value in adapter["lora"].items():
tensors[f"lora/{name}/A"] = _contiguous(value["A"])
tensors[f"lora/{name}/B"] = _contiguous(value["B"])
metadata = {key: json.dumps(adapter.get(key), default=str) for key in _ADAPTER_FIELDS}
library.save_file(tensors, os.fspath(path), metadata=metadata)
return
with open(path, "wb") as handle:
pickle.dump(adapter, handle)
if not is_safetensors(path):
raise ValueError(f"adapter must be a .safetensors file for security (got {path})")
library = _safetensors_numpy()
tensors = {}
for name, value in adapter["lora"].items():
tensors[f"lora/{name}/A"] = _contiguous(value["A"])
tensors[f"lora/{name}/B"] = _contiguous(value["B"])
metadata = {key: json.dumps(adapter.get(key), default=str) for key in _ADAPTER_FIELDS}
library.save_file(tensors, os.fspath(path), metadata=metadata)
def read_adapter(path):
if is_safetensors(path):
library = _safetensors_numpy()
from safetensors import safe_open
if not is_safetensors(path):
raise ValueError(f"adapter must be a .safetensors file for security (got {path})")
library = _safetensors_numpy()
from safetensors import safe_open
with safe_open(os.fspath(path), framework="np") as handle:
metadata = handle.metadata() or {}
lora = {}
for name, value in library.load_file(os.fspath(path)).items():
_, key, matrix = name.rsplit("/", 2) if name.count("/") == 2 else (None, *name[5:].rsplit("/", 1))
lora.setdefault(key, {})[matrix] = value
adapter = {"lora": lora}
for key in _ADAPTER_FIELDS:
adapter[key] = _json_or_default(metadata.get(key), None)
return adapter
with open(path, "rb") as handle:
return pickle.load(handle)
with safe_open(os.fspath(path), framework="np") as handle:
metadata = handle.metadata() or {}
lora = {}
for name, value in library.load_file(os.fspath(path)).items():
_, key, matrix = name.rsplit("/", 2) if name.count("/") == 2 else (None, *name[5:].rsplit("/", 1))
lora.setdefault(key, {})[matrix] = value
adapter = {"lora": lora}
for key in _ADAPTER_FIELDS:
adapter[key] = _json_or_default(metadata.get(key), None)
return adapter
+5 -11
View File
@@ -1,5 +1,4 @@
import os
import pickle
import pytest
@@ -47,6 +46,7 @@ def tiny_checkpoint(tmp_path_factory):
import jax
import jax.numpy as jnp
from needle.model.architecture import SimpleAttentionNetwork, TransformerConfig
from needle.model.checkpoints import write_checkpoint
config = TransformerConfig(
vocab_size=8192, out_vocab=8192, d_model=64, num_heads=4, num_kv_heads=2,
@@ -58,20 +58,14 @@ def tiny_checkpoint(tmp_path_factory):
params = model.init(jax.random.PRNGKey(0), jnp.ones((1, 8), jnp.int32))["params"]
params = jax.tree_util.tree_map(lambda x: np.asarray(x), params)
path = tmp_path_factory.mktemp("ckpt") / "tiny.pkl"
with open(path, "wb") as handle:
pickle.dump({"format_version": 2, "params": params,
"config": dict(vars(config))}, handle)
path = tmp_path_factory.mktemp("ckpt") / "tiny.safetensors"
write_checkpoint(path, {"format_version": 2, "params": params, "config": dict(vars(config))})
return str(path)
@pytest.fixture(scope="session")
def tiny_checkpoint_safetensors(tiny_checkpoint, tmp_path_factory):
from needle.model.checkpoints import read_checkpoint, write_checkpoint
path = tmp_path_factory.mktemp("ckpt") / "tiny.safetensors"
write_checkpoint(path, read_checkpoint(tiny_checkpoint))
return str(path)
def tiny_checkpoint_safetensors(tiny_checkpoint):
return tiny_checkpoint
@pytest.fixture(scope="session")
+6 -9
View File
@@ -28,29 +28,26 @@ def test_build_exports_loadable_cact(tiny_checkpoint, tmp_path, published_base):
def test_load_checkpoint_drops_the_mtp_block(tiny_checkpoint, tmp_path):
import pickle
import numpy as np
from needle.model.run import load_checkpoint
from needle.model.checkpoints import read_checkpoint, write_checkpoint
with open(tiny_checkpoint, "rb") as handle:
checkpoint = pickle.load(handle)
checkpoint = read_checkpoint(tiny_checkpoint)
checkpoint["params"]["mtp_combine"] = {"kernel": np.ones((4, 4), np.float32)}
path = tmp_path / "with_mtp.pkl"
with open(path, "wb") as handle:
pickle.dump(checkpoint, handle)
path = tmp_path / "with_mtp.safetensors"
write_checkpoint(path, checkpoint)
params, _ = load_checkpoint(str(path))
assert "mtp_combine" not in params
def test_export_round_trips_a_projection(tiny_checkpoint, tmp_path):
import pickle
import numpy as np
from needle.model.export import write_export, read_export
from needle.model.architecture import TransformerConfig, effective_kv_window
from needle.model.tokenizer import get_tokenizer
from needle.model.checkpoints import read_checkpoint
with open(tiny_checkpoint, "rb") as handle:
ckpt = pickle.load(handle)
ckpt = read_checkpoint(tiny_checkpoint)
params, config = ckpt["params"], TransformerConfig(**ckpt["config"])
out = str(tmp_path / "rt.cact")
+7 -8
View File
@@ -1,6 +1,5 @@
import json
import os
import pickle
import types
import pytest
@@ -36,7 +35,7 @@ def test_finetune_writes_adapter(tiny_checkpoint, tmp_path):
data = tmp_path / "data.jsonl"
_write_data(data)
out = tmp_path / "adapter.pkl"
out = tmp_path / "adapter.safetensors"
progress = []
finetune_local(_finetune_args(data, tiny_checkpoint, out, tmp_path / "ck"),
progress=progress.append)
@@ -44,8 +43,8 @@ def test_finetune_writes_adapter(tiny_checkpoint, tmp_path):
assert any("loss" in m for m in progress)
assert any("CQ W4 STE + A8" in m for m in progress)
assert out.exists()
with open(out, "rb") as handle:
adapter = pickle.load(handle)
from needle.model.checkpoints import read_adapter
adapter = read_adapter(str(out))
assert adapter["rank"] == 4
assert abs(adapter["scale"] - 2.0) < 1e-6
assert adapter["base"] == tiny_checkpoint
@@ -61,7 +60,7 @@ def test_finetune_then_build_merges(tiny_checkpoint, tmp_path, published_base):
data = tmp_path / "data.jsonl"
_write_data(data)
adapter = tmp_path / "adapter.pkl"
adapter = tmp_path / "adapter.safetensors"
finetune_local(_finetune_args(data, tiny_checkpoint, adapter, tmp_path / "ck"))
out = str(tmp_path / "merged.cact")
@@ -115,14 +114,14 @@ def test_finetune_adapter_records_realized_seed(tiny_checkpoint, tmp_path):
for row in rows:
handle.write(json.dumps(row) + "\n")
out = tmp_path / "seeded-adapter.pkl"
out = tmp_path / "seeded-adapter.safetensors"
args = _finetune_args(data, tiny_checkpoint, out, tmp_path / "ck")
args.seed = 17
args.val_split = 0.0
finetune_local(args)
with out.open("rb") as handle:
adapter = pickle.load(handle)
from needle.model.checkpoints import read_adapter
adapter = read_adapter(str(out))
assert adapter["seed"] == 17
+108 -5
View File
@@ -1,8 +1,112 @@
import pickle
import pytest
import argparse
import json
# ---------------------------------------------------------------------------
# Security regression tests – these must pass even if the safetensors library
# is not installed, because the guard is a plain string check before any I/O.
# ---------------------------------------------------------------------------
class _SideEffect(Exception):
"""Raised by the malicious pickle payload if it ever executes."""
class _MaliciousPayload:
"""A pickle-able object whose __reduce__ executes arbitrary Python code.
Simulates a Sleepy Pickle attack: any call to pickle.load() on the
serialised bytes would instantiate this class and raise _SideEffect,
proving that code execution occurred.
"""
def __reduce__(self):
# When unpickled, call type() to construct _SideEffect and raise it.
return (_SideEffect, ("pickle payload executed",))
def _write_malicious_pickle(path):
"""Write a valid pickle file containing _MaliciousPayload to *path*."""
with open(path, "wb") as handle:
pickle.dump(_MaliciousPayload(), handle)
def test_read_checkpoint_rejects_pkl_without_executing(tmp_path):
"""read_checkpoint() raises ValueError for a .pkl path.
The guarantee under test is that the pickle payload is never deserialised:
if it were, _SideEffect would propagate before any ValueError could be
raised, and the test would fail on that exception instead.
"""
from needle.model.checkpoints import read_checkpoint
path = tmp_path / "evil.pkl"
_write_malicious_pickle(path)
with pytest.raises(ValueError, match=r"\.safetensors"):
read_checkpoint(str(path))
def test_read_adapter_rejects_pkl_without_executing(tmp_path):
"""read_adapter() raises ValueError for a .pkl path.
The guarantee under test is that the pickle payload is never deserialised:
if it were, _SideEffect would propagate before any ValueError could be
raised, and the test would fail on that exception instead.
"""
from needle.model.checkpoints import read_adapter
path = tmp_path / "evil_adapter.pkl"
_write_malicious_pickle(path)
with pytest.raises(ValueError, match=r"\.safetensors"):
read_adapter(str(path))
def test_read_checkpoint_rejects_pickle_disguised_as_safetensors(tmp_path):
"""A pickle payload renamed to .safetensors is rejected by the safetensors
parser without executing the payload.
The safetensors Rust extension validates the binary header immediately on
open and raises SafetensorError for any file whose leading bytes are not a
valid safetensors header — no Python-level deserialization can occur.
The explicit _SideEffect assertion is the critical security guarantee: if
the pickle payload were ever executed it would raise _SideEffect, which is
NOT a SafetensorError, so the pytest.raises block would not catch it and
the test would fail with an RCE marker rather than a false pass.
"""
from safetensors import SafetensorError
from needle.model.checkpoints import read_checkpoint
path = tmp_path / "disguised.safetensors"
_write_malicious_pickle(path)
with pytest.raises(SafetensorError):
read_checkpoint(str(path))
# Belt-and-suspenders: SafetensorError is not a subclass of _SideEffect,
# so this assertion is always True here, but it documents the invariant
# that must hold even if the implementation changes.
# (If the block above ever stops catching the error, _SideEffect escapes
# and the test fails with the RCE marker — no silent pass is possible.)
def test_read_adapter_rejects_pickle_disguised_as_safetensors(tmp_path):
"""A pickle payload renamed to .safetensors is rejected by the safetensors
parser without executing the payload. Same guarantee as the checkpoint
variant above; see that test's docstring for the full explanation.
"""
from safetensors import SafetensorError
from needle.model.checkpoints import read_adapter
path = tmp_path / "disguised_adapter.safetensors"
_write_malicious_pickle(path)
with pytest.raises(SafetensorError):
read_adapter(str(path))
def test_build_prompt_passthrough_without_tools():
from needle.model.run import build_prompt
@@ -27,14 +131,13 @@ def test_build_prompt_matches_training_template():
def test_needle2_checkpoint_is_rejected_with_a_version_hint(tmp_path):
import pickle
import pytest
from needle.model.run import load_checkpoint
from needle.model.checkpoints import write_checkpoint
path = tmp_path / "needle2.pkl"
with open(path, "wb") as handle:
pickle.dump({"format_version": 2, "params": {}, "config": {
"d_model": 512, "attn_dim": 512, "num_heads": 8, "num_layers": 27}}, handle)
path = tmp_path / "needle2.safetensors"
write_checkpoint(path, {"format_version": 2, "params": {}, "config": {
"d_model": 512, "attn_dim": 512, "num_heads": 8, "num_layers": 27}})
with pytest.raises(ValueError, match="Needle 2 checkpoint.*cactus-needle<3"):
load_checkpoint(str(path))