Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
72 changes: 57 additions & 15 deletions src/modelinfo/parsers/huggingface.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import concurrent.futures
import json
import os
import re
import struct
import urllib.error
import urllib.parse
Expand All @@ -24,12 +25,20 @@ def _get_hf_endpoint() -> str:
return endpoint


def _hub_file_url(repo_id: str, filename: str) -> str:
repo_path = urllib.parse.quote(repo_id, safe="/")
file_path = urllib.parse.quote(filename, safe="/")
return f"{_get_hf_endpoint()}/{repo_path}/resolve/main/{file_path}"


def _get_hf_token() -> str | None:
token = os.environ.get("HF_TOKEN")
if token:
return token

cache_path = os.path.expanduser("~/.cache/huggingface/token")
cache_home = os.environ.get("XDG_CACHE_HOME", "~/.cache")
hf_home = os.environ.get("HF_HOME", os.path.join(cache_home, "huggingface"))
cache_path = os.path.expanduser(os.environ.get("HF_TOKEN_PATH", os.path.join(hf_home, "token")))
if os.path.exists(cache_path):
try:
with open(cache_path, "r", encoding="utf-8") as f:
Expand All @@ -53,8 +62,7 @@ def _make_request(
limit: int | None = None,
timeout: float = 10.0,
) -> bytes:
if headers is None:
headers = {}
headers = dict(headers) if headers is not None else {}

token = _get_hf_token()
if token:
Expand All @@ -63,6 +71,18 @@ def _make_request(
req = urllib.request.Request(url, headers=headers)
try:
with urllib.request.urlopen(req, timeout=timeout) as response:
requested_range = req.get_header("Range", "")
if (requested_range.startswith("bytes=")
and requested_range.split("=", 1)[1].split("-", 1)[0] != "0"
and getattr(response, "status", None) == 200):
raise ValueError("Server ignored the requested byte range; refusing data from the wrong offset.")
content_range = getattr(response, "headers", {}).get("Content-Range")
if requested_range.startswith("bytes=") and getattr(response, "status", None) == 206 and content_range:
match = re.fullmatch(r"bytes (\d+)-(\d+)/(?:\d+|\*)", content_range)
expected_start = int(requested_range.split("=", 1)[1].split("-", 1)[0])
if (match is None or int(match[1]) != expected_start
or int(match[2]) < int(match[1])):
raise ValueError("Server returned data for an invalid or unexpected byte range.")
if limit is not None:
return response.read(limit)
return response.read()
Expand All @@ -74,7 +94,7 @@ def _make_request(
raise

def _fetch_safetensors_header(repo_id: str, filename: str, timeout: float = 10.0) -> Dict[str, Any]:
url = f"{_get_hf_endpoint()}/{repo_id}/resolve/main/{filename}"
url = _hub_file_url(repo_id, filename)

# 1. Fetch the first 500KB in a single roundtrip
headers = {"Range": "bytes=0-500000"}
Expand All @@ -90,6 +110,8 @@ def _fetch_safetensors_header(repo_id: str, filename: str, timeout: float = 10.0
raise ValueError(f"File {filename} is too small to contain a SafeTensors header.")

header_size = struct.unpack("<Q", chunk[:8])[0]
if header_size > 100 * 1024 * 1024:
raise ValueError(f"Header length ({header_size} bytes) exceeds maximum allowed size.")

# 2. Slice locally if it fits
if 8 + header_size <= len(chunk):
Expand All @@ -99,10 +121,16 @@ def _fetch_safetensors_header(repo_id: str, filename: str, timeout: float = 10.0
headers = {"Range": f"bytes=8-{8+header_size-1}"}
json_bytes = _make_request(url, headers=headers, limit=header_size, timeout=timeout)

return json.loads(json_bytes)
if len(json_bytes) != header_size:
raise ValueError(f"SafeTensors header in {filename} is truncated: expected {header_size} bytes, got {len(json_bytes)}.")

header = json.loads(json_bytes)
if not isinstance(header, dict):
raise ValueError(f"SafeTensors header in {filename} must be a JSON object.")
return header

def _get_remote_file_size_fallback(repo_id: str, filename: str, timeout: float = 10.0) -> float:
req = urllib.request.Request(f"{_get_hf_endpoint()}/{repo_id}/resolve/main/{filename}", method="HEAD")
req = urllib.request.Request(_hub_file_url(repo_id, filename), method="HEAD")
token = _get_hf_token()
if token:
req.add_header("Authorization", f"Bearer {token}")
Expand All @@ -115,16 +143,22 @@ def _get_remote_file_size_fallback(repo_id: str, filename: str, timeout: float =

class RemoteFileStream:
def __init__(self, url: str, chunk_size: int = 1024*1024, timeout: float = 10.0):
if chunk_size <= 0:
raise ValueError("chunk_size must be positive")
self.url = url
self.chunk_size = chunk_size
self.timeout = timeout
self.buffer = b""
self.position = 0

def read(self, size: int = -1) -> bytes:
if size == 0:
return b""
if size == -1:
raise NotImplementedError("Unlimited remote read is not supported.")

if size < -1:
raise ValueError("read size must be nonnegative or -1")
end_pos = self.position + size
if end_pos > 50 * 1024 * 1024:
raise ValueError("Remote header read limit exceeded (50MB). File might be invalid or too large.")
Expand Down Expand Up @@ -157,11 +191,14 @@ def read(self, size: int = -1) -> bytes:

def seek(self, offset: int, whence: int = 0) -> int:
if whence == 0:
self.position = offset
position = offset
elif whence == 1:
self.position += offset
position = self.position + offset
else:
raise NotImplementedError("Seek from end is not supported.")
if position < 0:
raise ValueError("negative seek position")
self.position = position
return self.position

def tell(self) -> int:
Expand All @@ -172,7 +209,7 @@ def close(self) -> None:


def _fetch_remote_gguf_single(real_repo_id: str, filename: str, fallback_size: float | None, timeout: float) -> Tuple[Dict[str, Any], float]:
url = f"{_get_hf_endpoint()}/{real_repo_id}/resolve/main/{filename}"
url = _hub_file_url(real_repo_id, filename)
stream = RemoteFileStream(url, timeout=timeout)
from modelinfo.parsers.gguf import parse_gguf_header
tensors = parse_gguf_header(stream)
Expand All @@ -191,7 +228,7 @@ def _fetch_remote_gguf_group(real_repo_id: str, gguf_files: List[Dict[str, Any]]
header_target = gguf_files[0]

header_file = header_target["filename"]
url = f"{_get_hf_endpoint()}/{real_repo_id}/resolve/main/{header_file}"
url = _hub_file_url(real_repo_id, header_file)
stream = RemoteFileStream(url, timeout=timeout)
from modelinfo.parsers.gguf import parse_gguf_header
tensors = parse_gguf_header(stream)
Expand All @@ -212,10 +249,15 @@ def _fetch_remote_gguf_group(real_repo_id: str, gguf_files: List[Dict[str, Any]]
return tensors


def _fetch_shards_concurrently(real_repo_id: str, unique_shards: List[str], timeout: float) -> Tuple[Dict[str, Any], int]:
def _fetch_shards_concurrently(real_repo_id: str, unique_shards: List[str], timeout: float, weight_map: Dict[str, str] | None = None) -> Tuple[Dict[str, Any], int]:
def fetch_shard(shard: str):
try:
header = _fetch_safetensors_header(real_repo_id, shard, timeout=timeout)
if weight_map is not None:
assigned = [name for name, filename in weight_map.items() if filename == shard]
if any(name not in header for name in assigned):
raise ValueError(f"Indexed tensor missing from shard {shard!r}")
header = {name: header[name] for name in assigned}
return shard, header, None
except Exception as e:
return shard, {}, e
Expand All @@ -241,7 +283,7 @@ def _fetch_remote_safetensors_sharded(
fetch_tensors: bool,
timeout: float
) -> Tuple[Dict[str, Any], float]:
index_url = f"{_get_hf_endpoint()}/{real_repo_id}/resolve/main/model.safetensors.index.json"
index_url = _hub_file_url(real_repo_id, "model.safetensors.index.json")
index_data = json.loads(_make_request(index_url, timeout=timeout).decode("utf-8"))

weight_map = index_data.get("weight_map", {})
Expand All @@ -261,7 +303,7 @@ def _fetch_remote_safetensors_sharded(
"total_size": total_size
}
else:
tensors, missing_shards = _fetch_shards_concurrently(real_repo_id, unique_shards, timeout)
tensors, missing_shards = _fetch_shards_concurrently(real_repo_id, unique_shards, timeout, weight_map)
tensors["__metadata__"] = {
"missing_shards": missing_shards,
"total_shards": len(unique_shards),
Expand All @@ -273,7 +315,7 @@ def _fetch_remote_safetensors_sharded(

def _fetch_remote_safetensors_single(real_repo_id: str, timeout: float) -> Tuple[Dict[str, Any], float]:
total_size = 0.0
req = urllib.request.Request(f"{_get_hf_endpoint()}/{real_repo_id}/resolve/main/model.safetensors", method="HEAD")
req = urllib.request.Request(_hub_file_url(real_repo_id, "model.safetensors"), method="HEAD")
token = _get_hf_token()
if token:
req.add_header("Authorization", f"Bearer {token}")
Expand Down Expand Up @@ -315,7 +357,7 @@ def fetch_huggingface_repo(repo_id: str, fetch_tensors: bool = False, timeout: f

config = None
if "config.json" in filenames:
config_url = f"{_get_hf_endpoint()}/{real_repo_id}/resolve/main/config.json"
config_url = _hub_file_url(real_repo_id, "config.json")
config = json.loads(_make_request(config_url, timeout=timeout).decode("utf-8"))

# Find GGUF siblings
Expand Down
5 changes: 4 additions & 1 deletion src/modelinfo/parsers/safetensors.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,10 @@ def _read_single_header(path: str) -> dict[str, Any]:
if len(json_bytes) != header_length:
raise EOFError("Invalid SafeTensors file: Unexpected end of file while reading JSON header.")

return json.loads(json_bytes)
header = json.loads(json_bytes)
if not isinstance(header, dict):
raise ValueError(f"SafeTensors header in {path} must be a JSON object.")
return header

def parse_safetensors_header(path: str) -> dict[str, Any]:
dir_path = os.path.dirname(path)
Expand Down
16 changes: 16 additions & 0 deletions tests/test_hf_content_range.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,16 @@
import io
from unittest.mock import patch

import pytest

from modelinfo.parsers.huggingface import _make_request


@pytest.mark.parametrize('content_range', ['bytes 0-7/16', 'bytes 9-16/32', 'invalid'])
def test_partial_response_rejects_wrong_range(content_range):
response = io.BytesIO(b'bad-data')
response.status = 206
response.headers = {'Content-Range': content_range}
with patch('modelinfo.parsers.huggingface._get_hf_token', return_value=None), patch('modelinfo.parsers.huggingface.urllib.request.urlopen', return_value=response):
with pytest.raises(ValueError, match='range'):
_make_request('https://example.test/file', {'Range': 'bytes=8-15'}, limit=8)
24 changes: 24 additions & 0 deletions tests/test_hf_filename_encoding.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,24 @@
import struct

import pytest

from modelinfo.parsers import huggingface as hf


@pytest.mark.parametrize(("filename", "suffix"), [
("sub dir/model.safetensors", "sub%20dir/model.safetensors"),
("model#1.safetensors", "model%231.safetensors"),
("model?x.safetensors", "model%3Fx.safetensors"),
("model%20.safetensors", "model%2520.safetensors"),
])
def test_file_names_are_encoded_as_path_components(monkeypatch, filename, suffix):
monkeypatch.setenv("HF_ENDPOINT", "https://hub.example")
urls = []

def request(url, **kwargs):
urls.append(url)
return struct.pack("<Q", 2) + b"{}"

monkeypatch.setattr(hf, "_make_request", request)
assert hf._fetch_safetensors_header("org/model", filename) == {}
assert urls == ["https://hub.example/org/model/resolve/main/" + suffix]
18 changes: 18 additions & 0 deletions tests/test_hf_header_limit.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,18 @@
import struct

import pytest

from modelinfo.parsers import huggingface as hf


def test_remote_safetensors_rejects_oversized_header_before_second_request(monkeypatch):
calls = []

def request(*args, **kwargs):
calls.append(kwargs.get("limit"))
return struct.pack("<Q", 100 * 1024 * 1024 + 1)

monkeypatch.setattr(hf, "_make_request", request)
with pytest.raises(ValueError, match="maximum"):
hf._fetch_safetensors_header("org/model", "model.safetensors")
assert len(calls) == 1
33 changes: 33 additions & 0 deletions tests/test_hf_header_object.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,33 @@
import json
import struct

import pytest

from modelinfo.parsers import huggingface as hf


@pytest.mark.parametrize('header', [None, [], [1], 'weights', 42, True])
def test_remote_safetensors_header_requires_object(monkeypatch, header):
payload = json.dumps(header).encode()
monkeypatch.setattr(hf, '_make_request', lambda *a, **kw: struct.pack('<Q', len(payload)) + payload)
with pytest.raises(ValueError, match='JSON object'):
hf._fetch_safetensors_header('owner/model', 'model.safetensors')


def test_invalid_header_shard_is_recorded_as_missing(monkeypatch):
payload = b'[]'
monkeypatch.setattr(hf, '_make_request', lambda *a, **kw: struct.pack('<Q', len(payload)) + payload)
tensors, missing = hf._fetch_shards_concurrently('owner/model', ['model.safetensors'], 1)
assert tensors == {}
assert missing == 1


@pytest.mark.parametrize('header', [None, [], [1], 'weights', 42, True])
def test_local_safetensors_header_requires_object(tmp_path, header):
from modelinfo.parsers.safetensors import _read_single_header

payload = json.dumps(header).encode()
path = tmp_path / 'model.safetensors'
path.write_bytes(struct.pack('<Q', len(payload)) + payload)
with pytest.raises(ValueError, match='JSON object'):
_read_single_header(str(path))
19 changes: 19 additions & 0 deletions tests/test_hf_index_authority.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,19 @@
import json
from unittest.mock import patch

from modelinfo.parsers.huggingface import _fetch_remote_safetensors_sharded


def test_remote_shards_include_only_tensors_assigned_by_index():
index = {'weight_map': {'weight': 'one.safetensors'}, 'metadata': {'total_size': 8}}
header = {'weight': {'shape': [2], 'dtype': 'F32'}, 'extra': {'shape': [100], 'dtype': 'F32'}}
with patch('modelinfo.parsers.huggingface._make_request', return_value=json.dumps(index).encode()), patch('modelinfo.parsers.huggingface._fetch_safetensors_header', return_value=header):
tensors, _ = _fetch_remote_safetensors_sharded('org/model', None, True, 10)
assert set(tensors) == {'weight', '__metadata__'}


def test_remote_shard_missing_an_indexed_tensor_is_incomplete():
index = {'weight_map': {'weight': 'one.safetensors'}}
with patch('modelinfo.parsers.huggingface._make_request', return_value=json.dumps(index).encode()), patch('modelinfo.parsers.huggingface._fetch_safetensors_header', return_value={}):
tensors, _ = _fetch_remote_safetensors_sharded('org/model', None, True, 10)
assert tensors['__metadata__']['missing_shards'] == 1
22 changes: 22 additions & 0 deletions tests/test_hf_range_response.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,22 @@
import io

import pytest

from modelinfo.parsers import huggingface as hf


def test_nonzero_range_rejects_full_response(monkeypatch):
monkeypatch.setattr(hf, "_get_hf_token", lambda: None)
response = io.BytesIO(b"wrong prefix")
response.status = 200
monkeypatch.setattr(hf.urllib.request, "urlopen", lambda *a, **k: response)
with pytest.raises(ValueError, match="range"):
hf._make_request("https://hub.example/file", {"Range": "bytes=8-15"}, limit=8)


def test_nonzero_range_accepts_partial_response(monkeypatch):
monkeypatch.setattr(hf, "_get_hf_token", lambda: None)
response = io.BytesIO(b"expected")
response.status = 206
monkeypatch.setattr(hf.urllib.request, "urlopen", lambda *a, **k: response)
assert hf._make_request("https://hub.example/file", {"Range": "bytes=8-15"}, limit=8) == b"expected"
23 changes: 23 additions & 0 deletions tests/test_hf_token_paths.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,23 @@
import pytest

from modelinfo.parsers.huggingface import _get_hf_token


@pytest.mark.parametrize('variable,relative', [
('HF_TOKEN_PATH', 'custom-token'),
('HF_HOME', 'custom-home/token'),
('XDG_CACHE_HOME', 'custom-cache/huggingface/token'),
])
def test_custom_huggingface_token_cache(monkeypatch, tmp_path, variable, relative):
for name in ('HF_TOKEN', 'HF_TOKEN_PATH', 'HF_HOME', 'XDG_CACHE_HOME'):
monkeypatch.delenv(name, raising=False)
monkeypatch.setenv('HOME', str(tmp_path))
monkeypatch.setenv('USERPROFILE', str(tmp_path))
token_file = tmp_path / relative
token_file.parent.mkdir(parents=True, exist_ok=True)
token_file.write_text('fixture-token\n')
value = token_file if variable == 'HF_TOKEN_PATH' else (token_file.parent if variable == 'HF_HOME' else token_file.parent.parent)
monkeypatch.setenv(variable, str(value))
assert _get_hf_token() == 'fixture-token'
monkeypatch.setenv('HF_TOKEN', 'environment-fixture')
assert _get_hf_token() == 'environment-fixture'
14 changes: 14 additions & 0 deletions tests/test_hf_truncated_header.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,14 @@
import struct
from unittest.mock import patch

import pytest

from modelinfo.parsers.huggingface import _fetch_safetensors_header


@pytest.mark.parametrize('declared_size', [64, 600000])
def test_remote_header_rejects_truncation_even_when_json_is_complete(declared_size):
first_chunk = struct.pack('<Q', declared_size) + b'{}'
with patch('modelinfo.parsers.huggingface._make_request', side_effect=[first_chunk, b'{}']):
with pytest.raises(ValueError, match='truncated'):
_fetch_safetensors_header('org/model', 'model.safetensors')
Loading
Loading