Skip to content
Open
66 changes: 38 additions & 28 deletions src/modelinfo/parsers/gguf.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,36 +10,46 @@
29: "IQ1_M", 30: "BF16", 31: "Q4_0_4_4", 32: "Q4_0_4_8", 33: "Q4_0_8_8",
}

def _read_exact(f: Any, size: int) -> bytes:
data = bytearray()
while len(data) < size:
chunk = f.read(size - len(data))
if not chunk:
raise EOFError("Unexpected end of GGUF header")
data.extend(chunk)
return bytes(data)


def _read_gguf_value(f: Any, val_type: int) -> Any:
if val_type == 0:
return struct.unpack("<B", f.read(1))[0]
return struct.unpack("<B", _read_exact(f, 1))[0]
elif val_type == 1:
return struct.unpack("<b", f.read(1))[0]
return struct.unpack("<b", _read_exact(f, 1))[0]
elif val_type == 2:
return struct.unpack("<H", f.read(2))[0]
return struct.unpack("<H", _read_exact(f, 2))[0]
elif val_type == 3:
return struct.unpack("<h", f.read(2))[0]
return struct.unpack("<h", _read_exact(f, 2))[0]
elif val_type == 4:
return struct.unpack("<I", f.read(4))[0]
return struct.unpack("<I", _read_exact(f, 4))[0]
elif val_type == 5:
return struct.unpack("<i", f.read(4))[0]
return struct.unpack("<i", _read_exact(f, 4))[0]
elif val_type == 6:
return struct.unpack("<f", f.read(4))[0]
return struct.unpack("<f", _read_exact(f, 4))[0]
elif val_type == 7:
return struct.unpack("<?", f.read(1))[0]
return struct.unpack("<?", _read_exact(f, 1))[0]
elif val_type == 8:
slen = struct.unpack("<Q", f.read(8))[0]
return f.read(slen).decode("utf-8")
slen = struct.unpack("<Q", _read_exact(f, 8))[0]
return _read_exact(f, slen).decode("utf-8")
elif val_type == 9:
arr_type = struct.unpack("<I", f.read(4))[0]
arr_len = struct.unpack("<Q", f.read(8))[0]
arr_type = struct.unpack("<I", _read_exact(f, 4))[0]
arr_len = struct.unpack("<Q", _read_exact(f, 8))[0]
return [_read_gguf_value(f, arr_type) for _ in range(arr_len)]
elif val_type == 10:
return struct.unpack("<Q", f.read(8))[0]
return struct.unpack("<Q", _read_exact(f, 8))[0]
elif val_type == 11:
return struct.unpack("<q", f.read(8))[0]
return struct.unpack("<q", _read_exact(f, 8))[0]
elif val_type == 12:
return struct.unpack("<d", f.read(8))[0]
return struct.unpack("<d", _read_exact(f, 8))[0]
else:
raise ValueError(f"Unknown GGUF value type: {val_type}")

Expand All @@ -55,37 +65,37 @@ def parse_gguf_header(path_or_file: str | Any) -> Dict[str, Any]:

def _parse_gguf_header_from_stream(f: Any) -> Dict[str, Any]:
tensors: Dict[str, Any] = {}
magic = f.read(4)
magic = _read_exact(f, 4)
if magic != b"GGUF":
raise ValueError("Invalid GGUF file: Magic bytes missing.")

version = struct.unpack("<I", f.read(4))[0]
version = struct.unpack("<I", _read_exact(f, 4))[0]
if version < 2:
raise ValueError(f"Unsupported GGUF version: {version}")

tensor_count = struct.unpack("<Q", f.read(8))[0]
kv_count = struct.unpack("<Q", f.read(8))[0]
tensor_count = struct.unpack("<Q", _read_exact(f, 8))[0]
kv_count = struct.unpack("<Q", _read_exact(f, 8))[0]

metadata = {}
for _ in range(kv_count):
key_len = struct.unpack("<Q", f.read(8))[0]
key_name = f.read(key_len).decode("utf-8")
val_type = struct.unpack("<I", f.read(4))[0]
key_len = struct.unpack("<Q", _read_exact(f, 8))[0]
key_name = _read_exact(f, key_len).decode("utf-8")
val_type = struct.unpack("<I", _read_exact(f, 4))[0]
metadata[key_name] = _read_gguf_value(f, val_type)

tensors["__metadata__"] = metadata

for _ in range(tensor_count):
name_len = struct.unpack("<Q", f.read(8))[0]
name = f.read(name_len).decode("utf-8")
name_len = struct.unpack("<Q", _read_exact(f, 8))[0]
name = _read_exact(f, name_len).decode("utf-8")

n_dims = struct.unpack("<I", f.read(4))[0]
n_dims = struct.unpack("<I", _read_exact(f, 4))[0]
shape = []
for _ in range(n_dims):
shape.append(struct.unpack("<Q", f.read(8))[0])
shape.append(struct.unpack("<Q", _read_exact(f, 8))[0])

t_type = struct.unpack("<I", f.read(4))[0]
f.read(8) # skip offset bytes
t_type = struct.unpack("<I", _read_exact(f, 4))[0]
_read_exact(f, 8) # skip offset bytes

# Strict GGUF tensor type mapping
dtype = GGML_TYPE_MAP.get(t_type, "Unknown")
Expand Down
33 changes: 33 additions & 0 deletions src/modelinfo/parsers/pytorch.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,16 @@ def dummy_rebuild_tensor_v2(
dtype = "I32"
elif storage.name == "LongStorage":
dtype = "I64"
elif storage.name == "DoubleStorage":
dtype = "F64"
elif storage.name == "ShortStorage":
dtype = "I16"
elif storage.name == "CharStorage":
dtype = "I8"
elif storage.name == "ByteStorage":
dtype = "U8"
elif storage.name == "BoolStorage":
dtype = "BOOL"
return {"shape": list(size), "dtype": dtype}


Expand All @@ -46,9 +56,27 @@ class RestrictedUnpickler(pickle.Unpickler):
"BFloat16Storage",
"IntStorage",
"LongStorage",
"DoubleStorage",
"ShortStorage",
"CharStorage",
"ByteStorage",
"BoolStorage",
},
}

def persistent_load(self, saved_id: Any) -> DummyStorage:
# ZIP checkpoints refer to tensor storage via a five-element ID.
# Only reconstruct metadata using our existing dummy allowlist; never
# load storage bytes, import torch, or invoke a pickle-provided callable.
if not isinstance(saved_id, tuple) or len(saved_id) != 5:
raise pickle.UnpicklingError("Invalid storage persistent ID")
kind, storage_type, _, _, _ = saved_id
if (kind != "storage" or not isinstance(storage_type, type)
or not issubclass(storage_type, DummyStorage)
or storage_type is DummyStorage):
raise pickle.UnpicklingError("Unsupported storage persistent ID")
return storage_type()

def find_class(self, module: str, name: str) -> Any:
if module in self.ALLOWED_MODULES and name in self.ALLOWED_MODULES[module]:
if name == "OrderedDict":
Expand Down Expand Up @@ -84,6 +112,11 @@ def parse_pytorch_header(path: str) -> Dict[str, Any]:
data = unpickler.load()

if isinstance(data, dict):
for key in ("state_dict", "model_state_dict"):
candidate = data.get(key)
if isinstance(candidate, dict) and "shape" not in candidate:
data = candidate
break
for k, v in data.items():
if isinstance(v, dict) and "shape" in v:
tensors[k] = v
Expand Down
9 changes: 6 additions & 3 deletions src/modelinfo/parsers/safetensors.py
Original file line number Diff line number Diff line change
Expand Up @@ -63,9 +63,12 @@ def parse_safetensors_header(path: str) -> dict[str, Any]:
total_size += os.path.getsize(shard_path)
try:
shard_header = _read_single_header(shard_path)
for k, v in shard_header.items():
if k != "__metadata__":
tensors[k] = v
for name, assigned_shard in weight_map.items():
if assigned_shard != shard:
continue
if name not in shard_header:
raise ValueError(f"Indexed tensor {name!r} is missing from shard {shard!r}")
tensors[name] = shard_header[name]
except FileNotFoundError:
missing_shards += 1

Expand Down
125 changes: 125 additions & 0 deletions tests/test_binary_parsers.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,125 @@
"""Small generated checkpoint fixtures; no model downloads or torch dependency."""

import json
import struct
import zipfile
from pickle import UnpicklingError

import pytest

from modelinfo.parsers.pytorch import parse_pytorch_header
from modelinfo.parsers.safetensors import parse_safetensors_header


def write_safetensors(path, header):
payload = json.dumps(header).encode()
path.write_bytes(struct.pack("<Q", len(payload)) + payload)


def test_safetensors_extracts_header_without_reading_tensor_payload(tmp_path):
path = tmp_path / "model.safetensors"
header = {
"__metadata__": {"format": "pt"},
"weight": {"shape": [2, 3], "dtype": "F16", "data_offsets": [0, 12]},
}
write_safetensors(path, header)
assert parse_safetensors_header(str(path)) == header


@pytest.mark.parametrize("data", [b"", b"1234567", struct.pack("<Q", 20) + b"{}"])
def test_safetensors_rejects_truncated_headers(tmp_path, data):
path = tmp_path / "truncated.safetensors"
path.write_bytes(data)
with pytest.raises(EOFError, match="Unexpected end of file"):
parse_safetensors_header(str(path))


def test_safetensors_rejects_oversized_header_before_reading_it(tmp_path):
path = tmp_path / "oversized.safetensors"
path.write_bytes(struct.pack("<Q", 100 * 1024 * 1024 + 1))
with pytest.raises(ValueError, match="exceeds maximum allowed size"):
parse_safetensors_header(str(path))


def test_safetensors_rejects_malformed_json(tmp_path):
path = tmp_path / "invalid.safetensors"
path.write_bytes(struct.pack("<Q", 1) + b"{")
with pytest.raises(json.JSONDecodeError):
parse_safetensors_header(str(path))


def test_safetensors_index_deduplicates_shards_and_reports_missing_files(tmp_path):
index = tmp_path / "model.safetensors.index.json"
shard = tmp_path / "part1.safetensors"
header = {
"__metadata__": {"format": "pt"},
"a": {"shape": [2], "dtype": "F16"},
"b": {"shape": [3], "dtype": "F32"},
}
write_safetensors(shard, header)
index.write_text(
json.dumps(
{
"weight_map": {
"a": shard.name,
"b": shard.name,
"c": "missing.safetensors",
}
}
)
)
result = parse_safetensors_header(str(index))
assert result["a"] == header["a"]
assert result["b"] == header["b"]
assert "c" not in result
assert result["__metadata__"] == {
"missing_shards": 1,
"total_shards": 2,
"is_sharded": True,
"disk_size": shard.stat().st_size,
}


def test_pytorch_reads_ordered_state_dict_from_archive(tmp_path):
path = tmp_path / "model.pt"
# Fixed protocol-0 OrderedDict fixture: inspect metadata without generating
# executable serialization during the test.
state = b"ccollections\nOrderedDict\np0\n(tRp1\nVweight\np2\n(dp3\nVshape\np4\n(lp5\nI2\naI3\nasVdtype\np6\nVF16\np7\nssVstep\np8\nI12\ns."
with zipfile.ZipFile(path, "w") as archive:
archive.writestr("checkpoint/data.pkl", state)
assert parse_pytorch_header(str(path)) == {
"weight": {"shape": [2, 3], "dtype": "F16"},
"step": {"shape": [], "dtype": "F32"},
}


def test_pytorch_rejects_non_zip_checkpoint(tmp_path):
path = tmp_path / "legacy.pt"
path.write_bytes(b"not a zip")
with pytest.raises(ValueError, match="not a valid zip archive"):
parse_pytorch_header(str(path))


def test_pytorch_rejects_archive_without_metadata(tmp_path):
path = tmp_path / "empty.pt"
with zipfile.ZipFile(path, "w") as archive:
archive.writestr("checkpoint/version", b"3")
with pytest.raises(ValueError, match="Could not find data.pkl"):
parse_pytorch_header(str(path))


def test_pytorch_rejects_unapproved_pickle_global(tmp_path):
path = tmp_path / "forbidden.pt"
# A harmless builtin still must not bypass the explicit unpickler allowlist.
with zipfile.ZipFile(path, "w") as archive:
archive.writestr("checkpoint/data.pkl", b"cbuiltins\nlen\n.")
with pytest.raises(UnpicklingError, match="forbidden"):
parse_pytorch_header(str(path))


def test_pytorch_propagates_truncated_pickle_error(tmp_path):
path = tmp_path / "truncated.pt"
with zipfile.ZipFile(path, "w") as archive:
archive.writestr("checkpoint/data.pkl", b"")
with pytest.raises(EOFError):
parse_pytorch_header(str(path))
35 changes: 35 additions & 0 deletions tests/test_gguf_truncation.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,35 @@
import io
import struct

import pytest

from modelinfo.parsers.gguf import parse_gguf_header


def tensor_header():
return (b'GGUF' + struct.pack('<IQQ', 3, 1, 0)
+ struct.pack('<Q', 6) + b'weight'
+ struct.pack('<IQIQ', 1, 3, 0, 0))


@pytest.mark.parametrize('missing', range(1, 9))
def test_rejects_truncated_final_tensor_offset(missing):
with pytest.raises(EOFError):
parse_gguf_header(io.BytesIO(tensor_header()[:-missing]))


def test_rejects_truncated_final_metadata_string():
payload = (b'GGUF' + struct.pack('<IQQ', 3, 0, 1)
+ struct.pack('<Q', 3) + b'key' + struct.pack('<IQ', 8, 5) + b'ab')
with pytest.raises(EOFError):
parse_gguf_header(io.BytesIO(payload))


class ShortReads(io.BytesIO):
def read(self, size=-1):
return super().read(min(size, 2))


def test_stream_can_return_partial_reads():
assert parse_gguf_header(ShortReads(tensor_header()))['weight'] == {
'shape': [3], 'dtype': 'F32'}
16 changes: 16 additions & 0 deletions tests/test_pytorch_checkpoint_wrapper.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,16 @@
import pickle
import zipfile

import pytest

from modelinfo.parsers.pytorch import parse_pytorch_header


@pytest.mark.parametrize('wrapper', ['state_dict', 'model_state_dict'])
def test_wrapped_state_dict_excludes_training_metadata(tmp_path, wrapper):
weight = {'shape': [2, 3], 'dtype': 'F32'}
data = {wrapper: {'layer.weight': weight}, 'epoch': 10, 'loss': 0.125}
path = tmp_path / 'checkpoint.pt'
with zipfile.ZipFile(path, 'w') as archive:
archive.writestr('checkpoint/data.pkl', pickle.dumps(data, protocol=2))

Check failure on line 15 in tests/test_pytorch_checkpoint_wrapper.py

View check run for this annotation

Codacy Production / Codacy Static Code Analysis

tests/test_pytorch_checkpoint_wrapper.py#L15

Avoid using `pickle`, which is known to lead to code execution vulnerabilities.
assert parse_pytorch_header(str(path)) == {'layer.weight': weight}
31 changes: 31 additions & 0 deletions tests/test_pytorch_storage.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,31 @@
"""Fixed protocol-2 fixture using torch.save's storage persistent ID layout."""
import io
import pickle
import zipfile

import pytest

from modelinfo.parsers.pytorch import RestrictedUnpickler, parse_pytorch_header


@pytest.mark.parametrize('storage,dtype', [('FloatStorage', 'F32'), ('HalfStorage', 'F16'),
('BFloat16Storage', 'BF16'), ('IntStorage', 'I32'), ('LongStorage', 'I64'),
('DoubleStorage', 'F64'), ('ShortStorage', 'I16'), ('CharStorage', 'I8'),
('ByteStorage', 'U8'), ('BoolStorage', 'BOOL')])
def test_tensor_with_persistent_storage_id(tmp_path, storage, dtype):
payload = (b'\x80\x02}X\x06\x00\x00\x00weightctorch._utils\n_rebuild_tensor_v2\n'
b'((X\x07\x00\x00\x00storagectorch\n' + storage.encode() + b'\n'
b'X\x01\x00\x00\x000X\x03\x00\x00\x00cpuK\x06tQK\x00K\x02K\x03\x86'
b'K\x03K\x01\x86\x89ccollections\nOrderedDict\n)RtRs.')
path = tmp_path / 'model.pt'
with zipfile.ZipFile(path, 'w') as archive:
archive.writestr('model/data.pkl', payload)
assert parse_pytorch_header(str(path)) == {'weight': {'shape': [2, 3], 'dtype': dtype}}


@pytest.mark.parametrize('pid', [None, (), ('module',), ('storage', int, '0', 'cpu', 6),
('storage', None, '0', 'cpu', 6)])
def test_invalid_persistent_ids_are_rejected(pid):
unpickler = RestrictedUnpickler(io.BytesIO())
with pytest.raises(pickle.UnpicklingError):
unpickler.persistent_load(pid)
Loading
Loading