diff --git a/src/modelinfo/parsers/huggingface.py b/src/modelinfo/parsers/huggingface.py index b36dd7f..f5e5cc0 100644 --- a/src/modelinfo/parsers/huggingface.py +++ b/src/modelinfo/parsers/huggingface.py @@ -1,6 +1,7 @@ import concurrent.futures import json import os +import re import struct import urllib.error import urllib.parse @@ -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: @@ -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: @@ -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() @@ -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"} @@ -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(" 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): @@ -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}") @@ -115,6 +143,8 @@ 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 @@ -122,9 +152,13 @@ def __init__(self, url: str, chunk_size: int = 1024*1024, timeout: float = 10.0) 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.") @@ -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: @@ -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) @@ -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) @@ -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 @@ -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", {}) @@ -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), @@ -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}") @@ -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 diff --git a/src/modelinfo/parsers/safetensors.py b/src/modelinfo/parsers/safetensors.py index 2e7d705..b76b438 100644 --- a/src/modelinfo/parsers/safetensors.py +++ b/src/modelinfo/parsers/safetensors.py @@ -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) diff --git a/tests/test_hf_content_range.py b/tests/test_hf_content_range.py new file mode 100644 index 0000000..e5853e7 --- /dev/null +++ b/tests/test_hf_content_range.py @@ -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) diff --git a/tests/test_hf_filename_encoding.py b/tests/test_hf_filename_encoding.py new file mode 100644 index 0000000..dc3aac1 --- /dev/null +++ b/tests/test_hf_filename_encoding.py @@ -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("= len(payload): + raise urllib.error.HTTPError(request.full_url, 416, "range", {}, None) + return response(payload[start : end + 1]) + + monkeypatch.setattr(hf.urllib.request, "urlopen", urlopen) + stream = hf.RemoteFileStream("https://hub.example/file", chunk_size=4) + assert stream.read(3) == b"abc" + assert stream.seek(1) == 1 + assert stream.read(3) == b"bcd" + assert calls == [(0, 3)] + assert stream.read(20) == b"efghij" + assert stream.tell() == 10 + assert stream.read(1) == b"" + + +def test_request_does_not_retain_token_in_reused_headers(monkeypatch): + calls = [] + def urlopen(request, timeout): + calls.append(request) + return response(b"ok") + monkeypatch.setattr(hf.urllib.request, "urlopen", urlopen) + headers = {"Range": "bytes=0-1"} + hf._make_request("https://hub.example/file", headers=headers) + monkeypatch.setattr(hf, "_get_hf_token", lambda: None) + hf._make_request("https://hub.example/file", headers=headers) + assert calls[0].get_header("Authorization") == "Bearer test-token" + assert calls[1].get_header("Authorization") is None + assert headers == {"Range": "bytes=0-1"} diff --git a/tests/test_remote_stream_positions.py b/tests/test_remote_stream_positions.py new file mode 100644 index 0000000..2f158c9 --- /dev/null +++ b/tests/test_remote_stream_positions.py @@ -0,0 +1,28 @@ +import pytest + +from modelinfo.parsers import huggingface as hf + + +def test_negative_seek_does_not_corrupt_position(): + stream = hf.RemoteFileStream("https://hub.example/file") + assert stream.seek(3) == 3 + with pytest.raises(ValueError): + stream.seek(-4, 1) + assert stream.tell() == 3 + with pytest.raises(ValueError): + stream.seek(-1) + assert stream.tell() == 3 + + +def test_invalid_negative_read_does_not_move_cursor(): + stream = hf.RemoteFileStream("https://hub.example/file") + stream.buffer = b"abc" + with pytest.raises(ValueError): + stream.read(-2) + assert stream.tell() == 0 + + +@pytest.mark.parametrize("chunk_size", [0, -1]) +def test_invalid_chunk_size_rejected(chunk_size): + with pytest.raises(ValueError): + hf.RemoteFileStream("https://hub.example/file", chunk_size=chunk_size) diff --git a/tests/test_remote_zero_read.py b/tests/test_remote_zero_read.py new file mode 100644 index 0000000..8154cb7 --- /dev/null +++ b/tests/test_remote_zero_read.py @@ -0,0 +1,12 @@ +from unittest.mock import patch + +from modelinfo.parsers.huggingface import RemoteFileStream + + +def test_zero_length_read_after_seek_never_fetches_or_moves(): + stream = RemoteFileStream('https://example.test/model.gguf') + for position in (100, 60 * 1024 * 1024): + stream.seek(position) + with patch('modelinfo.parsers.huggingface._make_request', side_effect=AssertionError('unexpected network request')): + assert stream.read(0) == b'' + assert stream.tell() == position