diff --git a/CHANGELOG.md b/CHANGELOG.md index 1861e91c3..52f429ffe 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,6 +9,7 @@ metadata and the backend fallback mirror it. ## [Unreleased] **Highlights** +- Validate http response ranges before writing segments (#2451) — thanks @rudycelekli! - Ask VoiceStudio Agent adds chat, harness selection, feature presets, read-only planning and autopilot app actions without a source checkout (#2407) - Home credits contributors with over 10 commits in three responsive rows of round avatars stacked from right to left with an All contributors link, with GitHub and X links on Palash's hover card (#2407) - The VoiceStudio.sh Open Source title opens a website preview below the clicked item, within the right content area, with navigation and external-browser controls (#2407) diff --git a/backend/services/segmented_download.py b/backend/services/segmented_download.py index fedfbea72..d27bc947a 100644 --- a/backend/services/segmented_download.py +++ b/backend/services/segmented_download.py @@ -21,6 +21,7 @@ from __future__ import annotations import asyncio import json import os +import re from typing import Callable, Optional import httpx @@ -170,6 +171,12 @@ async def segmented_download( headers = {**_auth_headers(final_url, token), "Range": f"bytes={start}-{end}"} async with client.stream("GET", final_url, headers=headers) as r: r.raise_for_status() + match = re.fullmatch(r"bytes\s+([0-9]+)-([0-9]+)/([0-9]+|\*)", + r.headers.get("content-range", "").strip(), re.IGNORECASE) + if (r.status_code != 206 or match is None + or (int(match[1]), int(match[2])) != (start, end) + or (match[3] != "*" and int(match[3]) != size)): + raise ValueError(f"invalid response range for bytes {start}-{end}/{size}") got = 0 with open(part, "r+b") as fh: fh.seek(start) diff --git a/docs/install/troubleshooting.md b/docs/install/troubleshooting.md index 1bc1ef012..4270f9118 100644 --- a/docs/install/troubleshooting.md +++ b/docs/install/troubleshooting.md @@ -683,6 +683,8 @@ order: - macOS/Linux: `export HF_ENDPOINT=https://hf-mirror.com` - Windows (PowerShell): `[Environment]::SetEnvironmentVariable("HF_ENDPOINT","https://hf-mirror.com","User")` +A segmented download refuses a response whose status or Content-Range does not match the requested bytes and file size. An invalid response is not published as the model file; retry through a server or mirror that supports correct byte ranges. + **Manual fallback** (if downloads keep failing), pull the weights yourself into the same cache, then relaunch: diff --git a/tests/test_download_range_contract.py b/tests/test_download_range_contract.py new file mode 100644 index 000000000..680e8c382 --- /dev/null +++ b/tests/test_download_range_contract.py @@ -0,0 +1,54 @@ +"""Actual HTTP ranges must prove their identity before a model is published.""" +import asyncio +import threading +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer + +import pytest + + +@pytest.mark.parametrize('mode', ['wrong_start', 'wrong_total', 'ignored_range', 'valid', 'valid_unknown_total']) +def test_wrong_range_response_never_commits_model(tmp_path, mode): + from services.segmented_download import segmented_download + + content = b'A' * (4 * 1024 * 1024) + b'B' * (4 * 1024 * 1024) + + class Handler(BaseHTTPRequestHandler): + def log_message(self, *args): + pass + + def do_HEAD(self): + self.send_response(200) + self.send_header('Content-Length', str(len(content))) + self.send_header('Accept-Ranges', 'bytes') + self.end_headers() + + def do_GET(self): + lo, hi = map(int, self.headers['Range'].removeprefix('bytes=').split('-')) + self.send_response(200 if mode == 'ignored_range' else 206) + advertised_lo = lo + 1 if mode == 'wrong_start' else lo + total = '*' if mode == 'valid_unknown_total' else len(content) + 1 if mode == 'wrong_total' else len(content) + self.send_header('Content-Range', f'bytes {advertised_lo}-{hi}/{total}') + self.send_header('Content-Length', str(hi - lo + 1)) + self.end_headers() + try: + self.wfile.write(content[lo:hi + 1] if mode.startswith('valid') else b'X' * (hi - lo + 1)) + except (BrokenPipeError, ConnectionResetError): + pass # Invalid headers may be rejected before the response body. + + server = ThreadingHTTPServer(('127.0.0.1', 0), Handler) + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + destination = tmp_path / 'model.bin' + try: + call = segmented_download(f'http://127.0.0.1:{server.server_port}/model', str(destination), num_connections=2) + if mode.startswith('valid'): + asyncio.run(call) + assert destination.read_bytes() == content + else: + with pytest.raises(ValueError, match='range'): + asyncio.run(call) + assert not destination.exists() + finally: + server.shutdown() + server.server_close() + thread.join(timeout=5) diff --git a/tests/test_fdl_segmented_download.py b/tests/test_fdl_segmented_download.py index fa4496852..2d5bd6217 100644 --- a/tests/test_fdl_segmented_download.py +++ b/tests/test_fdl_segmented_download.py @@ -39,7 +39,7 @@ def _ranged_handler(payload=PAYLOAD, *, accept_ranges=True, record=None): if rng and accept_ranges: lo, hi = rng.replace("bytes=", "").split("-") lo, hi = int(lo), int(hi) - return httpx.Response(206, content=payload[lo:hi + 1]) + return httpx.Response(206, headers={"Content-Range": f"bytes {lo}-{hi}/{len(payload)}"}, content=payload[lo:hi + 1]) return httpx.Response(200, content=payload) return handler @@ -177,7 +177,7 @@ def test_concurrency_stays_at_num_connections(tmp_path, monkeypatch): except asyncio.TimeoutError: pass lo, hi = request.headers["range"].replace("bytes=", "").split("-") - return httpx.Response(206, content=PAYLOAD[int(lo):int(hi) + 1]) + return httpx.Response(206, headers={"Content-Range": f"bytes {lo}-{hi}/{len(PAYLOAD)}"}, content=PAYLOAD[int(lo):int(hi) + 1]) finally: state["inflight"] -= 1 @@ -215,7 +215,7 @@ def test_dropped_connection_resumes_from_manifest(tmp_path, monkeypatch): raise httpx.RemoteProtocolError("peer closed connection", request=request) lo, hi = request.headers["range"].replace("bytes=", "").split("-") served.append(int(hi) - int(lo) + 1) - return httpx.Response(206, content=PAYLOAD[int(lo):int(hi) + 1]) + return httpx.Response(206, headers={"Content-Range": f"bytes {lo}-{hi}/{len(PAYLOAD)}"}, content=PAYLOAD[int(lo):int(hi) + 1]) with pytest.raises(httpx.RemoteProtocolError): _download(handler, dest, expected_size=len(PAYLOAD), num_connections=2)