diff --git a/studio/backend/hub/services/download_lifecycle.py b/studio/backend/hub/services/download_lifecycle.py index 44c39337fbd..649780c4211 100644 --- a/studio/backend/hub/services/download_lifecycle.py +++ b/studio/backend/hub/services/download_lifecycle.py @@ -8,6 +8,7 @@ import signal import subprocess import sys +import time import threading from pathlib import Path from typing import Callable, Optional @@ -206,7 +207,9 @@ def finalize_worker_exit( repo_type: Optional[RepoType] = None, repo_id: Optional[str] = None, transport: Optional[str] = None, -) -> None: + cancel_marker_transport: Optional[str] = None, + defer_error: bool = False, +) -> str: """Block until *proc* exits, then record the job's terminal state in *registry*. Drains and scrubs stderr first, then classifies the exit code. A no-op when the process was already dropped (e.g. superseded). @@ -218,7 +221,7 @@ def finalize_worker_exit( rc = proc.wait() cancel_requested = registry.cancel_requested(key) if not registry.drop_process(key, proc): - return + return "idle" stderr_text = download_registry.scrub_secrets( (stderr_data or b"").decode("utf-8", "replace").strip(), hf_token = hf_token, @@ -226,6 +229,8 @@ def finalize_worker_exit( state = classify_exit(rc, cancel_requested = cancel_requested) if state == "complete": registry.set_job(key, "complete") + if transport == download_registry.TRANSPORT_HTTP: + registry.update_job_transport(key, download_registry.TRANSPORT_HTTP) if stderr_text: if download_manifest.MANIFEST_DEGRADED_MARKER in stderr_text: logger.warning( @@ -258,18 +263,226 @@ def finalize_worker_exit( metadata.variant if metadata is not None and metadata.variant else download_registry.variant_from_key(key), - transport, + cancel_marker_transport or transport, logger = logger, ) else: - registry.set_job( + if not defer_error: + registry.set_job( + key, + "error", + stderr_text or f"worker exited with code {rc}", + ) + logger.error( + f"{log_prefix} failed for {label} (rc={rc}): {stderr_text}", + ) + return state + + +def _set_retry_failure_state( + registry: download_registry.DownloadRegistry, + key: str, + error: str, + *, + repo_type: RepoType, + repo_id: str, + fallback_variant: Optional[str], + fallback_transport: Optional[str], + logger, +) -> str: + state, metadata = registry.set_error_unless_cancelled(key, error) + if state == "cancelled": + download_registry.persist_cancel_marker( + repo_type, + repo_id, + metadata.variant if metadata is not None and metadata.variant else fallback_variant, + metadata.transport + if metadata is not None and metadata.transport + else fallback_transport, + logger = logger, + ) + return state + + +def _try_http_retry( + registry: download_registry.DownloadRegistry, + key: str, + *, + hf_token: Optional[str], + label: str, + log_prefix: str, + logger, + repo_type: RepoType, + repo_id: str, + watch_name: str, +) -> bool: + """Reclaim *key* with HTTP transport and spawn a recovery worker. + + Returns ``True`` when the HTTP worker was successfully registered. + Caller is responsible for ensuring this is only called when: the job is + in ``"error"`` state, the original transport was XET, and HTTP is available. + + Derives variant and blob-hash metadata from the registry entry written by + the original XET claim so callers do not re-construct worker arguments. + Re-queries peer protection hashes at spawn time to reflect any concurrent + sibling changes between the XET failure and this call. + """ + original_metadata = registry.get_job_metadata(key) + if original_metadata is None: + logger.debug("%s XET retry skipped for %s; metadata unavailable", log_prefix, label) + _set_retry_failure_state( + registry, + key, + "XET retry skipped: metadata unavailable", + repo_type = repo_type, + repo_id = repo_id, + fallback_variant = download_registry.variant_from_key(key), + fallback_transport = download_registry.TRANSPORT_XET, + logger = logger, + ) + return False + if original_metadata.transport != download_registry.TRANSPORT_XET: + logger.debug( + "%s XET retry skipped for %s; original transport was %s", + log_prefix, + label, + original_metadata.transport, + ) + _set_retry_failure_state( + registry, + key, + f"XET retry skipped: original transport was {original_metadata.transport}", + repo_type = repo_type, + repo_id = repo_id, + fallback_variant = original_metadata.variant, + fallback_transport = original_metadata.transport, + logger = logger, + ) + return False + variant = original_metadata.variant + blob_hashes = original_metadata.blob_hashes + progress_blob_hashes = original_metadata.progress_blob_hashes + completed_baseline_bytes = ( + download_registry.completed_blob_bytes( + repo_type, + repo_id, + progress_blob_hashes, + ) + if progress_blob_hashes + else 0 + ) + generation = registry.current_generation(key) + registry.release_active_slot(key) + while True: + if registry.cancel_requested(key): + _set_retry_failure_state( + registry, + key, + "HTTP retry cancelled before reclaiming the download slot", + repo_type = repo_type, + repo_id = repo_id, + fallback_variant = variant, + fallback_transport = original_metadata.transport, + logger = logger, + ) + return False + + claimed, conflict_state = registry.claim( key, - "error", - stderr_text or f"worker exited with code {rc}", + download_registry.TRANSPORT_HTTP, + repo_type = repo_type, + repo_id = repo_id, + variant = variant, + blob_hashes = blob_hashes, + progress_blob_hashes = progress_blob_hashes, + completed_baseline_bytes = completed_baseline_bytes, + generation = generation, + replace_active = True, + cancel_marker_transport = original_metadata.transport, + ) + if claimed: + break + if conflict_state == "deleting": + logger.debug( + "%s XET retry claim rejected for %s; repo is being deleted", + log_prefix, + label, + ) + _set_retry_failure_state( + registry, + key, + "HTTP retry could not reclaim the download slot", + repo_type = repo_type, + repo_id = repo_id, + fallback_variant = variant, + fallback_transport = original_metadata.transport, + logger = logger, + ) + return False + logger.debug( + "%s XET retry claim blocked for %s by active sibling state %s; waiting", + log_prefix, + label, + conflict_state, + ) + time.sleep(0.05) + + args: list[str] = ["--repo-id", repo_id] + if repo_type == "dataset": + args.append("--dataset") + elif variant: + args.extend(["--variant", variant]) + + # Re-query at spawn time: sibling state may have changed since XET failed. + peer_hashes = registry.peer_blob_hashes(key) if variant else frozenset() + + logger.warning( + "%s XET worker failed for %s; retrying over HTTP", + log_prefix, + label, + ) + try: + proc = spawn_worker( + args, + hf_token, + use_xet = False, + protected_blob_hashes = peer_hashes or None, ) + except Exception as exc: + scrubbed = download_registry.scrub_secrets(str(exc), hf_token = hf_token) logger.error( - f"{log_prefix} failed for {label} (rc={rc}): {stderr_text}", + "%s HTTP retry spawn failed for %s: %s", + log_prefix, + label, + scrubbed, + ) + registry.update_job_transport(key, original_metadata.transport) + _set_retry_failure_state( + registry, + key, + scrubbed, + repo_type = repo_type, + repo_id = repo_id, + fallback_variant = variant, + fallback_transport = original_metadata.transport, + logger = logger, ) + return False + + return register_worker( + registry, + key, + proc, + hf_token = hf_token, + label = label, + log_prefix = log_prefix, + logger = logger, + repo_type = repo_type, + repo_id = repo_id, + transport = download_registry.TRANSPORT_HTTP, + cancel_marker_transport = original_metadata.transport, + watch_name = watch_name, + ) def kill_and_reap_process( @@ -305,6 +518,7 @@ def register_worker( repo_type: RepoType, repo_id: str, transport: str, + cancel_marker_transport: Optional[str] = None, watch_name: str, ) -> bool: if not registry.register_process(key, proc): @@ -315,7 +529,14 @@ def register_worker( def _watch() -> None: try: - finalize_worker_exit( + can_retry_http = ( + transport == download_registry.TRANSPORT_XET + and download_registry.download_transport_unavailable_reason( + download_registry.TRANSPORT_HTTP + ) + is None + ) + state = finalize_worker_exit( registry, key, proc, @@ -326,7 +547,25 @@ def _watch() -> None: repo_type = repo_type, repo_id = repo_id, transport = transport, + cancel_marker_transport = cancel_marker_transport, + defer_error = can_retry_http, ) + # XET-to-HTTP recovery: when a non-cancelled XET worker fails and + # HTTP is available, attempt one automatic retry over HTTP. The + # transport check is the recursion guard: an HTTP worker that errors + # never satisfies `transport == TRANSPORT_XET`, so it stays terminal. + if can_retry_http and state == "error": + _try_http_retry( + registry, + key, + hf_token = worker_token, + label = label, + log_prefix = log_prefix, + logger = logger, + repo_type = repo_type, + repo_id = repo_id, + watch_name = watch_name, + ) except Exception: # finalize_worker_exit is the only thing that clears running/cancelling; # if it raises, force a terminal state so claim() isn't blocked until restart. @@ -422,8 +661,19 @@ def cancel_worker( return "cancelling" return registry.get_job(key).state # Worker already exited; let its watcher classify the real return code. - # Arming a pending cancel here could mislabel a genuine failure as a cancel. if proc.poll() is not None: + get_metadata = getattr(registry, "get_job_metadata", None) + metadata = get_metadata(key) if get_metadata is not None else None + can_retry_http = ( + metadata is not None + and metadata.transport == download_registry.TRANSPORT_XET + and download_registry.download_transport_unavailable_reason( + download_registry.TRANSPORT_HTTP + ) + is None + ) + if can_retry_http and registry.mark_pending_cancel(key, generation): + return "cancelling" return registry.get_job(key).state if not registry.request_cancel(key, proc, generation): diff --git a/studio/backend/hub/tests/test_download_lifecycle.py b/studio/backend/hub/tests/test_download_lifecycle.py index a4baafa317b..82b1b83ba19 100644 --- a/studio/backend/hub/tests/test_download_lifecycle.py +++ b/studio/backend/hub/tests/test_download_lifecycle.py @@ -1,7 +1,13 @@ # SPDX-License-Identifier: AGPL-3.0-only # Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 +import io +import json +import logging + from hub.services import download_lifecycle +from hub.utils import download_registry +from hub.utils import state_dir def _set_xet_reason(monkeypatch, reason): @@ -12,6 +18,37 @@ def _set_xet_reason(monkeypatch, reason): ) +def _make_proc(rc, stderr = b""): + class _Proc: + pid = 4242 + + def __init__(self): + self._rc = rc + self._waited = False + self.killed = False + self.stderr = io.BytesIO(stderr) + + def poll(self): + return self._rc if self._waited else None + + def wait(self, timeout = None): + self._waited = True + return self._rc + + def kill(self): + self.killed = True + + return _Proc() + + +class _ImmediateThread: + def __init__(self, *, target, **_kwargs): + self._target = target + + def start(self): + self._target() + + def test_resolve_effective_use_xet_keeps_http_when_not_requested(monkeypatch): _set_xet_reason(monkeypatch, "should not be consulted") assert download_lifecycle.resolve_effective_use_xet(False) is False @@ -25,3 +62,1165 @@ def test_resolve_effective_use_xet_keeps_xet_when_available(monkeypatch): def test_resolve_effective_use_xet_downgrades_when_xet_unavailable(monkeypatch): _set_xet_reason(monkeypatch, "Xet transport is unavailable because hf_xet is not installed.") assert download_lifecycle.resolve_effective_use_xet(True) is False + + +def test_download_watcher_retries_xet_failure_over_http_for_model_and_dataset( + monkeypatch, tmp_path +): + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state") + monkeypatch.setattr(download_lifecycle.threading, "Thread", _ImmediateThread) + real_register_worker = download_lifecycle.register_worker + + cases = [ + ( + "model", + "Org/Model", + "Q4_K_M", + ["--repo-id", "Org/Model", "--variant", "Q4_K_M"], + frozenset({"mainhash"}), + frozenset({"mainhash", "mmprojhash"}), + 12, + 18, + ), + ( + "dataset", + "Org/Data", + None, + ["--repo-id", "Org/Data", "--dataset"], + frozenset(), + frozenset(), + 0, + 0, + ), + ] + + for ( + repo_type, + repo_id, + variant, + expected_args, + blob_hashes, + progress_blob_hashes, + baseline_bytes, + retry_baseline_bytes, + ) in cases: + registry = download_registry.DownloadRegistry() + key = ( + download_registry.normalize_job_key(f"{repo_id}::{variant}") + if variant is not None + else download_registry.normalize_repo_key(repo_id) + ) + assert registry.claim( + key, + download_registry.TRANSPORT_XET, + repo_type = repo_type, + repo_id = repo_id, + variant = variant, + blob_hashes = blob_hashes, + progress_blob_hashes = progress_blob_hashes, + completed_baseline_bytes = baseline_bytes, + ) + original_generation = registry.current_generation(key) + proc = _make_proc(1, b"xet failed") + spawned = [] + retry_registers = [] + baseline_calls = [] + + def fake_completed_blob_bytes(*args): + baseline_calls.append(args) + return retry_baseline_bytes + + def fake_spawn_worker( + args, + hf_token, + *, + use_xet, + protected_blob_hashes = None, + ): + assert registry.get_job(key).state == "running" + spawned.append((args, use_xet, protected_blob_hashes)) + return _make_proc(0, b"http retry") + + def fake_register_worker(*_args, **kwargs): + assert registry.get_job_metadata(key).transport == download_registry.TRANSPORT_HTTP + if repo_type == "model" and variant: + sibling_claimed, sibling_state = registry.claim( + "Org/Model::Q5_K_M", + download_registry.TRANSPORT_XET, + repo_type = "model", + repo_id = "Org/Model", + variant = "Q5_K_M", + blob_hashes = frozenset({"q5-main"}), + progress_blob_hashes = frozenset({"q5-main", "mmprojhash"}), + ) + assert sibling_claimed is False + assert sibling_state == "running" + retry_registers.append(kwargs) + return True + + monkeypatch.setattr(download_registry, "completed_blob_bytes", fake_completed_blob_bytes) + monkeypatch.setattr(download_lifecycle, "spawn_worker", fake_spawn_worker) + monkeypatch.setattr(download_lifecycle, "register_worker", fake_register_worker) + + assert real_register_worker( + registry, + key, + proc, + hf_token = None, + label = f"{repo_id}{f' [{variant}]' if variant else ''}", + log_prefix = "Download", + logger = logging.getLogger("test"), + repo_type = repo_type, + repo_id = repo_id, + transport = download_registry.TRANSPORT_XET, + watch_name = f"{repo_type}-watch", + ) + + assert spawned == [(expected_args, False, None)] + assert ( + retry_registers and retry_registers[0]["transport"] == download_registry.TRANSPORT_HTTP + ) + metadata = registry.get_job_metadata(key) + assert metadata is not None + assert metadata.transport == download_registry.TRANSPORT_HTTP + assert metadata.blob_hashes == blob_hashes + assert metadata.progress_blob_hashes == progress_blob_hashes + assert metadata.completed_baseline_bytes == retry_baseline_bytes + assert registry.current_generation(key) == original_generation + assert baseline_calls == ( + [("model", "Org/Model", progress_blob_hashes)] if progress_blob_hashes else [] + ) + assert registry.get_job(key).state == "running" + + +def test_download_watcher_defers_http_retry_while_sibling_xet_variant_is_active( + monkeypatch, tmp_path +): + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state") + monkeypatch.setattr(download_lifecycle.threading, "Thread", _ImmediateThread) + real_register_worker = download_lifecycle.register_worker + registry = download_registry.DownloadRegistry() + key_a = download_registry.normalize_job_key("Org/Model::Q4_K_M") + key_b = download_registry.normalize_job_key("Org/Model::Q5_K_M") + assert registry.claim( + key_a, + download_registry.TRANSPORT_XET, + repo_type = "model", + repo_id = "Org/Model", + variant = "Q4_K_M", + blob_hashes = frozenset({"mainhash-a"}), + progress_blob_hashes = frozenset({"mainhash-a"}), + ) + assert registry.claim( + key_b, + download_registry.TRANSPORT_XET, + repo_type = "model", + repo_id = "Org/Model", + variant = "Q5_K_M", + blob_hashes = frozenset({"mainhash-b"}), + progress_blob_hashes = frozenset({"mainhash-b"}), + ) + proc = _make_proc(1, b"xet failed") + sleep_calls = [] + spawned = [] + + def fake_sleep(seconds): + sleep_calls.append(seconds) + assert registry.get_job(key_a).state == "running" + assert registry.get_job_metadata(key_a).transport == download_registry.TRANSPORT_XET + assert registry.begin_delete("Org/Model", "Q4_K_M") is False + registry.set_job(key_b, "complete") + + def fake_spawn_worker( + args, + hf_token, + *, + use_xet, + protected_blob_hashes = None, + ): + assert use_xet is False + assert registry.get_job_metadata(key_a).transport == download_registry.TRANSPORT_HTTP + spawned.append((args, protected_blob_hashes)) + return _make_proc(0, b"http retry") + + monkeypatch.setattr(download_lifecycle.time, "sleep", fake_sleep) + monkeypatch.setattr(download_lifecycle, "spawn_worker", fake_spawn_worker) + + assert real_register_worker( + registry, + key_a, + proc, + hf_token = None, + label = "Org/Model [Q4_K_M]", + log_prefix = "Download", + logger = logging.getLogger("test"), + repo_type = "model", + repo_id = "Org/Model", + transport = download_registry.TRANSPORT_XET, + watch_name = "model-watch", + ) + + assert sleep_calls == [0.05] + assert spawned == [(["--repo-id", "Org/Model", "--variant", "Q4_K_M"], None)] + metadata = registry.get_job_metadata(key_a) + assert metadata is not None + assert metadata.transport == download_registry.TRANSPORT_HTTP + assert registry.get_job(key_a).state == "complete" + assert registry.get_job(key_b).state == "complete" + + +def test_download_watcher_does_not_deadlock_when_sibling_xet_variants_both_fail( + monkeypatch, tmp_path +): + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state") + monkeypatch.setattr(download_lifecycle.threading, "Thread", _ImmediateThread) + real_register_worker = download_lifecycle.register_worker + registry = download_registry.DownloadRegistry() + key_a = download_registry.normalize_job_key("Org/Model::Q4_K_M") + key_b = download_registry.normalize_job_key("Org/Model::Q5_K_M") + for key, variant, blob_hash in ( + (key_a, "Q4_K_M", "mainhash-a"), + (key_b, "Q5_K_M", "mainhash-b"), + ): + assert registry.claim( + key, + download_registry.TRANSPORT_XET, + repo_type = "model", + repo_id = "Org/Model", + variant = variant, + blob_hashes = frozenset({blob_hash}), + progress_blob_hashes = frozenset({blob_hash}), + ) + + spawned = [] + triggered_b = False + + def fake_spawn_worker( + args, + hf_token, + *, + use_xet, + protected_blob_hashes = None, + ): + assert use_xet is False + spawned.append(args) + return _make_proc(0, b"http retry") + + def fake_sleep(_seconds): + nonlocal triggered_b + assert registry.get_job(key_a).state == "running" + assert registry.get_job_metadata(key_a).transport == download_registry.TRANSPORT_XET + assert not triggered_b + triggered_b = True + assert real_register_worker( + registry, + key_b, + _make_proc(1, b"xet failed b"), + hf_token = None, + label = "Org/Model [Q5_K_M]", + log_prefix = "Download", + logger = logging.getLogger("test"), + repo_type = "model", + repo_id = "Org/Model", + transport = download_registry.TRANSPORT_XET, + watch_name = "model-watch-b", + ) + + monkeypatch.setattr(download_lifecycle.time, "sleep", fake_sleep) + monkeypatch.setattr(download_lifecycle, "spawn_worker", fake_spawn_worker) + + assert real_register_worker( + registry, + key_a, + _make_proc(1, b"xet failed a"), + hf_token = None, + label = "Org/Model [Q4_K_M]", + log_prefix = "Download", + logger = logging.getLogger("test"), + repo_type = "model", + repo_id = "Org/Model", + transport = download_registry.TRANSPORT_XET, + watch_name = "model-watch-a", + ) + + assert triggered_b + assert spawned == [ + ["--repo-id", "Org/Model", "--variant", "Q5_K_M"], + ["--repo-id", "Org/Model", "--variant", "Q4_K_M"], + ] + assert registry.get_job(key_a).state == "complete" + assert registry.get_job(key_b).state == "complete" + + +def test_download_watcher_keeps_http_failure_terminal(monkeypatch, tmp_path): + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state") + monkeypatch.setattr(download_lifecycle.threading, "Thread", _ImmediateThread) + real_register_worker = download_lifecycle.register_worker + registry = download_registry.DownloadRegistry() + key = download_registry.normalize_repo_key("Org/Data") + assert registry.claim( + key, + download_registry.TRANSPORT_HTTP, + repo_type = "dataset", + repo_id = "Org/Data", + ) + proc = _make_proc(1, b"http failed") + + def fake_register_worker(*_args, **_kwargs): + raise AssertionError("HTTP failure should stay terminal") + + monkeypatch.setattr(download_lifecycle, "register_worker", fake_register_worker) + + assert real_register_worker( + registry, + key, + proc, + hf_token = None, + label = "Org/Data", + log_prefix = "Download", + logger = logging.getLogger("test"), + repo_type = "dataset", + repo_id = "Org/Data", + transport = download_registry.TRANSPORT_HTTP, + watch_name = "dataset-watch", + ) + + assert registry.get_job(key).state == "error" + + +def test_http_retry_skip_without_metadata_sets_error(monkeypatch, tmp_path): + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state") + registry = download_registry.DownloadRegistry() + key = download_registry.normalize_job_key("Org/Model::Q4_K_M") + registry.set_job(key, "running") + + def fake_spawn_worker(*_args, **_kwargs): + raise AssertionError("metadata-free retry should not spawn a worker") + + def fake_register_worker(*_args, **_kwargs): + raise AssertionError("metadata-free retry should not register a worker") + + monkeypatch.setattr(download_lifecycle, "spawn_worker", fake_spawn_worker) + monkeypatch.setattr(download_lifecycle, "register_worker", fake_register_worker) + + assert not download_lifecycle._try_http_retry( + registry, + key, + hf_token = None, + label = "Org/Model [Q4_K_M]", + log_prefix = "Download", + logger = logging.getLogger("test"), + repo_type = "model", + repo_id = "Org/Model", + watch_name = "model-watch", + ) + + assert registry.get_job(key).state == "error" + + +def test_http_retry_skip_for_non_xet_transport_sets_error(monkeypatch, tmp_path): + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state") + registry = download_registry.DownloadRegistry() + key = download_registry.normalize_repo_key("Org/Data") + assert registry.claim( + key, + download_registry.TRANSPORT_HTTP, + repo_type = "dataset", + repo_id = "Org/Data", + ) + + def fake_spawn_worker(*_args, **_kwargs): + raise AssertionError("non-XET retry should not spawn a worker") + + def fake_register_worker(*_args, **_kwargs): + raise AssertionError("non-XET retry should not register a worker") + + monkeypatch.setattr(download_lifecycle, "spawn_worker", fake_spawn_worker) + monkeypatch.setattr(download_lifecycle, "register_worker", fake_register_worker) + + assert not download_lifecycle._try_http_retry( + registry, + key, + hf_token = None, + label = "Org/Data", + log_prefix = "Download", + logger = logging.getLogger("test"), + repo_type = "dataset", + repo_id = "Org/Data", + watch_name = "dataset-watch", + ) + + assert registry.get_job(key).state == "error" + + +def test_download_watcher_restores_xet_transport_when_http_retry_spawn_fails(monkeypatch, tmp_path): + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state") + monkeypatch.setattr(download_lifecycle.threading, "Thread", _ImmediateThread) + real_register_worker = download_lifecycle.register_worker + registry = download_registry.DownloadRegistry() + key = download_registry.normalize_job_key("Org/Model::Q4_K_M") + assert registry.claim( + key, + download_registry.TRANSPORT_XET, + repo_type = "model", + repo_id = "Org/Model", + variant = "Q4_K_M", + blob_hashes = frozenset({"mainhash"}), + progress_blob_hashes = frozenset({"mainhash", "mmprojhash"}), + completed_baseline_bytes = 12, + ) + proc = _make_proc(1, b"xet failed") + + def fake_spawn_worker(*_args, **_kwargs): + raise RuntimeError("HTTP retry spawn failed") + + def fake_register_worker(*_args, **_kwargs): + raise AssertionError("HTTP retry should not register after spawn failure") + + monkeypatch.setattr(download_lifecycle, "spawn_worker", fake_spawn_worker) + monkeypatch.setattr(download_lifecycle, "register_worker", fake_register_worker) + + assert real_register_worker( + registry, + key, + proc, + hf_token = None, + label = "Org/Model [Q4_K_M]", + log_prefix = "Download", + logger = logging.getLogger("test"), + repo_type = "model", + repo_id = "Org/Model", + transport = download_registry.TRANSPORT_XET, + watch_name = "model-watch", + ) + + assert registry.get_job(key).state == "error" + metadata = registry.get_job_metadata(key) + assert metadata is not None + assert metadata.transport == download_registry.TRANSPORT_XET + + +def test_download_watcher_preserves_xet_marker_when_http_retry_is_cancelled_before_register( + monkeypatch, tmp_path +): + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state") + monkeypatch.setattr(download_lifecycle.threading, "Thread", _ImmediateThread) + real_register_worker = download_lifecycle.register_worker + registry = download_registry.DownloadRegistry() + key = download_registry.normalize_job_key("Org/Model::Q4_K_M") + assert registry.claim( + key, + download_registry.TRANSPORT_XET, + repo_type = "model", + repo_id = "Org/Model", + variant = "Q4_K_M", + blob_hashes = frozenset({"mainhash"}), + progress_blob_hashes = frozenset({"mainhash"}), + ) + generation = registry.current_generation(key) + proc = _make_proc(1, b"xet failed") + markers = [] + + def fake_spawn_worker(*_args, **_kwargs): + assert registry.get_job_metadata(key).transport == download_registry.TRANSPORT_HTTP + assert registry.mark_pending_cancel(key, generation) + return _make_proc(0, b"http retry") + + def fake_persist_cancel_marker(*args, **_kwargs): + markers.append(args) + + monkeypatch.setattr(download_lifecycle, "spawn_worker", fake_spawn_worker) + monkeypatch.setattr(download_registry, "persist_cancel_marker", fake_persist_cancel_marker) + + assert real_register_worker( + registry, + key, + proc, + hf_token = None, + label = "Org/Model [Q4_K_M]", + log_prefix = "Download", + logger = logging.getLogger("test"), + repo_type = "model", + repo_id = "Org/Model", + transport = download_registry.TRANSPORT_XET, + watch_name = "model-watch", + ) + + assert registry.get_job(key).state == "cancelled" + assert markers == [("model", "Org/Model", "Q4_K_M", download_registry.TRANSPORT_XET)] + assert registry.get_job_metadata(key) is None + + +def test_download_watcher_honors_cancel_before_http_retry_spawn(monkeypatch, tmp_path): + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state") + monkeypatch.setattr(download_lifecycle.threading, "Thread", _ImmediateThread) + real_register_worker = download_lifecycle.register_worker + registry = download_registry.DownloadRegistry() + key = download_registry.normalize_job_key("Org/Model::Q4_K_M") + assert registry.claim( + key, + download_registry.TRANSPORT_XET, + repo_type = "model", + repo_id = "Org/Model", + variant = "Q4_K_M", + blob_hashes = frozenset({"mainhash"}), + progress_blob_hashes = frozenset({"mainhash"}), + ) + generation = registry.current_generation(key) + proc = _make_proc(1, b"xet failed") + proc._waited = True + markers = [] + + def fake_drain_stderr_excerpt(_stream): + assert ( + download_lifecycle.cancel_worker( + registry, + key, + generation = generation, + label = "Org/Model [Q4_K_M]", + logger = logging.getLogger("test"), + ) + == "cancelling" + ) + return b"xet failed" + + def fake_spawn_worker(*_args, **_kwargs): + raise AssertionError("cancelled retry should not spawn a worker") + + def fake_register_worker(*_args, **_kwargs): + raise AssertionError("cancelled retry should not register a worker") + + def fake_persist_cancel_marker(*args, **_kwargs): + markers.append(args) + + monkeypatch.setattr(download_lifecycle, "drain_stderr_excerpt", fake_drain_stderr_excerpt) + monkeypatch.setattr(download_lifecycle, "spawn_worker", fake_spawn_worker) + monkeypatch.setattr(download_lifecycle, "register_worker", fake_register_worker) + monkeypatch.setattr(download_registry, "persist_cancel_marker", fake_persist_cancel_marker) + + assert real_register_worker( + registry, + key, + proc, + hf_token = None, + label = "Org/Model [Q4_K_M]", + log_prefix = "Download", + logger = logging.getLogger("test"), + repo_type = "model", + repo_id = "Org/Model", + transport = download_registry.TRANSPORT_XET, + watch_name = "model-watch", + ) + + assert registry.get_job(key).state == "cancelled" + assert markers == [("model", "Org/Model", "Q4_K_M", download_registry.TRANSPORT_XET)] + + +def test_download_watcher_preserves_xet_marker_when_http_retry_is_cancelled_after_register( + monkeypatch, tmp_path +): + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state") + monkeypatch.setattr(download_lifecycle.threading, "Thread", _ImmediateThread) + real_register_worker = download_lifecycle.register_worker + registry = download_registry.DownloadRegistry() + key = download_registry.normalize_job_key("Org/Model::Q4_K_M") + assert registry.claim( + key, + download_registry.TRANSPORT_XET, + repo_type = "model", + repo_id = "Org/Model", + variant = "Q4_K_M", + blob_hashes = frozenset({"mainhash"}), + progress_blob_hashes = frozenset({"mainhash"}), + ) + generation = registry.current_generation(key) + markers = [] + + class _CancellingProc: + pid = 5252 + + def __init__(self): + self.stderr = io.BytesIO(b"cancelled") + + def wait(self, timeout = None): + proc = registry.get_process(key) + assert proc is self + assert registry.request_cancel(key, self, generation) + return -9 + + def poll(self): + return None + + def kill(self): + raise AssertionError("registered retry worker should exit by cancellation") + + def fake_spawn_worker(*_args, **_kwargs): + assert registry.get_job_metadata(key).transport == download_registry.TRANSPORT_HTTP + return _CancellingProc() + + def fake_persist_cancel_marker(*args, **_kwargs): + markers.append(args) + + monkeypatch.setattr(download_lifecycle, "spawn_worker", fake_spawn_worker) + monkeypatch.setattr(download_registry, "persist_cancel_marker", fake_persist_cancel_marker) + + assert real_register_worker( + registry, + key, + _make_proc(1, b"xet failed"), + hf_token = None, + label = "Org/Model [Q4_K_M]", + log_prefix = "Download", + logger = logging.getLogger("test"), + repo_type = "model", + repo_id = "Org/Model", + transport = download_registry.TRANSPORT_XET, + watch_name = "model-watch", + ) + + assert registry.get_job(key).state == "cancelled" + assert markers == [("model", "Org/Model", "Q4_K_M", download_registry.TRANSPORT_XET)] + assert registry.get_job_metadata(key) is not None + assert registry.get_job_metadata(key).transport == download_registry.TRANSPORT_HTTP + + +def test_http_retry_shutdown_and_breadcrumb_preserve_xet_marker(monkeypatch, tmp_path): + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state") + registry = download_registry.DownloadRegistry() + key = download_registry.normalize_job_key("Org/Model::Q4_K_M") + assert registry.claim( + key, + download_registry.TRANSPORT_HTTP, + repo_type = "model", + repo_id = "Org/Model", + variant = "Q4_K_M", + blob_hashes = frozenset({"mainhash"}), + progress_blob_hashes = frozenset({"mainhash"}), + cancel_marker_transport = download_registry.TRANSPORT_XET, + ) + metadata = registry.get_job_metadata(key) + assert metadata is not None + assert metadata.transport == download_registry.TRANSPORT_HTTP + assert metadata.cancel_marker_transport == download_registry.TRANSPORT_XET + markers = [] + + class _KillableProc: + pid = 6262 + + def __init__(self): + self._rc = None + self.killed = False + + def poll(self): + return self._rc + + def kill(self): + self.killed = True + self._rc = -9 + + def wait(self, timeout = None): + return self._rc + + proc = _KillableProc() + assert registry.register_process(key, proc) + worker_files = list(state_dir.workers_dir().glob("*.json")) + assert len(worker_files) == 1 + payload = json.loads(worker_files[0].read_text(encoding = "utf-8")) + assert payload["transport"] == download_registry.TRANSPORT_HTTP + assert payload["cancel_marker_transport"] == download_registry.TRANSPORT_XET + + def fake_persist_cancel_marker(*args, **_kwargs): + markers.append(args) + + monkeypatch.setattr(download_registry, "persist_cancel_marker", fake_persist_cancel_marker) + + registry.terminate_all("download") + + assert proc.killed is True + assert markers == [("model", "Org/Model", "Q4_K_M", download_registry.TRANSPORT_XET)] + + +def test_download_watcher_preserves_pending_cancel_when_http_retry_claim_fails( + monkeypatch, tmp_path +): + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state") + monkeypatch.setattr(download_lifecycle.threading, "Thread", _ImmediateThread) + real_register_worker = download_lifecycle.register_worker + registry = download_registry.DownloadRegistry() + key = download_registry.normalize_job_key("Org/Model::Q4_K_M") + assert registry.claim( + key, + download_registry.TRANSPORT_XET, + repo_type = "model", + repo_id = "Org/Model", + variant = "Q4_K_M", + blob_hashes = frozenset({"mainhash"}), + progress_blob_hashes = frozenset({"mainhash"}), + ) + generation = registry.current_generation(key) + proc = _make_proc(1, b"xet failed") + markers = [] + original_claim = registry.claim + + def fake_spawn_worker(*_args, **_kwargs): + raise AssertionError("failed retry claim should not spawn a worker") + + def fake_claim(*args, **kwargs): + if kwargs.get("replace_active"): + assert registry.mark_pending_cancel(key, generation) + return False, "cancelling" + return original_claim(*args, **kwargs) + + def fake_register_worker(*_args, **_kwargs): + raise AssertionError("failed retry claim should not register a worker") + + def fake_persist_cancel_marker(*args, **_kwargs): + markers.append(args) + + monkeypatch.setattr(registry, "claim", fake_claim) + monkeypatch.setattr(download_lifecycle, "spawn_worker", fake_spawn_worker) + monkeypatch.setattr(download_lifecycle, "register_worker", fake_register_worker) + monkeypatch.setattr(download_registry, "persist_cancel_marker", fake_persist_cancel_marker) + + assert real_register_worker( + registry, + key, + proc, + hf_token = None, + label = "Org/Model [Q4_K_M]", + log_prefix = "Download", + logger = logging.getLogger("test"), + repo_type = "model", + repo_id = "Org/Model", + transport = download_registry.TRANSPORT_XET, + watch_name = "model-watch", + ) + + assert registry.get_job(key).state == "cancelled" + assert markers == [("model", "Org/Model", "Q4_K_M", download_registry.TRANSPORT_XET)] + + +def test_download_watcher_preserves_pending_cancel_when_http_retry_spawn_fails( + monkeypatch, tmp_path +): + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state") + monkeypatch.setattr(download_lifecycle.threading, "Thread", _ImmediateThread) + real_register_worker = download_lifecycle.register_worker + registry = download_registry.DownloadRegistry() + key = download_registry.normalize_job_key("Org/Model::Q4_K_M") + assert registry.claim( + key, + download_registry.TRANSPORT_XET, + repo_type = "model", + repo_id = "Org/Model", + variant = "Q4_K_M", + blob_hashes = frozenset({"mainhash"}), + progress_blob_hashes = frozenset({"mainhash"}), + ) + generation = registry.current_generation(key) + proc = _make_proc(1, b"xet failed") + markers = [] + + def fake_spawn_worker(*_args, **_kwargs): + assert registry.mark_pending_cancel(key, generation) + raise RuntimeError("HTTP retry spawn failed") + + def fake_register_worker(*_args, **_kwargs): + raise AssertionError("failed retry spawn should not register a worker") + + def fake_persist_cancel_marker(*args, **_kwargs): + markers.append(args) + + monkeypatch.setattr(download_lifecycle, "spawn_worker", fake_spawn_worker) + monkeypatch.setattr(download_lifecycle, "register_worker", fake_register_worker) + monkeypatch.setattr(download_registry, "persist_cancel_marker", fake_persist_cancel_marker) + + assert real_register_worker( + registry, + key, + proc, + hf_token = None, + label = "Org/Model [Q4_K_M]", + log_prefix = "Download", + logger = logging.getLogger("test"), + repo_type = "model", + repo_id = "Org/Model", + transport = download_registry.TRANSPORT_XET, + watch_name = "model-watch", + ) + + assert registry.get_job(key).state == "cancelled" + assert markers == [("model", "Org/Model", "Q4_K_M", download_registry.TRANSPORT_XET)] + + +def test_download_watcher_persists_cancel_without_retry(monkeypatch, tmp_path): + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state") + monkeypatch.setattr(download_lifecycle.threading, "Thread", _ImmediateThread) + real_register_worker = download_lifecycle.register_worker + registry = download_registry.DownloadRegistry() + key = download_registry.normalize_repo_key("Org/Data") + assert registry.claim( + key, + download_registry.TRANSPORT_XET, + repo_type = "dataset", + repo_id = "Org/Data", + ) + proc = _make_proc(130, b"cancelled") + markers = [] + + def fake_persist_cancel_marker(*args, **_kwargs): + markers.append(args) + + def fake_register_worker(*_args, **_kwargs): + raise AssertionError("cancelled workers should not retry") + + monkeypatch.setattr(download_registry, "persist_cancel_marker", fake_persist_cancel_marker) + monkeypatch.setattr(download_lifecycle, "register_worker", fake_register_worker) + + assert real_register_worker( + registry, + key, + proc, + hf_token = None, + label = "Org/Data", + log_prefix = "Download", + logger = logging.getLogger("test"), + repo_type = "dataset", + repo_id = "Org/Data", + transport = download_registry.TRANSPORT_XET, + watch_name = "dataset-watch", + ) + + assert registry.get_job(key).state == "cancelled" + assert markers == [("dataset", "Org/Data", None, download_registry.TRANSPORT_XET)] + + +def test_active_download_refs_include_waiting_http_retry(monkeypatch, tmp_path): + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state") + monkeypatch.setattr(download_lifecycle.threading, "Thread", _ImmediateThread) + real_register_worker = download_lifecycle.register_worker + registry = download_registry.DownloadRegistry() + key_a = download_registry.normalize_job_key("Org/Model::Q4_K_M") + key_b = download_registry.normalize_job_key("Org/Model::Q5_K_M") + assert registry.claim( + key_a, + download_registry.TRANSPORT_XET, + repo_type = "model", + repo_id = "Org/Model", + variant = "Q4_K_M", + blob_hashes = frozenset({"mainhash-a"}), + progress_blob_hashes = frozenset({"mainhash-a"}), + ) + assert registry.claim( + key_b, + download_registry.TRANSPORT_XET, + repo_type = "model", + repo_id = "Org/Model", + variant = "Q5_K_M", + blob_hashes = frozenset({"mainhash-b"}), + progress_blob_hashes = frozenset({"mainhash-b"}), + ) + proc = _make_proc(1, b"xet failed") + waiting_variants = [] + + def fake_sleep(_seconds): + # The Q4_K_M retry is released from the repo guard while it waits for the + # slot; deletion stays blocked, so the listing must still surface it. + assert registry.begin_delete("Org/Model", "Q4_K_M") is False + refs = download_lifecycle.active_download_refs(registry, "Org/Model", with_variant = True) + waiting_variants.append(sorted(ref.variant for ref in refs)) + registry.set_job(key_b, "complete") + + def fake_spawn_worker( + args, + hf_token, + *, + use_xet, + protected_blob_hashes = None, + ): + assert use_xet is False + return _make_proc(0, b"http retry") + + monkeypatch.setattr(download_lifecycle.time, "sleep", fake_sleep) + monkeypatch.setattr(download_lifecycle, "spawn_worker", fake_spawn_worker) + + assert real_register_worker( + registry, + key_a, + proc, + hf_token = None, + label = "Org/Model [Q4_K_M]", + log_prefix = "Download", + logger = logging.getLogger("test"), + repo_type = "model", + repo_id = "Org/Model", + transport = download_registry.TRANSPORT_XET, + watch_name = "model-watch", + ) + + # While the HTTP retry waited for the slot, the listing must include both the + # still-active sibling and the released retry so the frontend can adopt or + # cancel that backend-running retry after a reload. + assert waiting_variants == [["Q4_K_M", "Q5_K_M"]] + assert registry.get_job(key_a).state == "complete" + + +def test_http_retry_bails_when_shutdown_settles_parked_retry(monkeypatch, tmp_path): + # A parked XET->HTTP retry has dropped its worker and active-slot guard, so it + # is absent from terminate_all()'s process snapshot. If shutdown runs while the + # retry waits behind an active sibling, terminate_all() must still settle the + # no-process job so the reclaim loop bails instead of spawning a fresh HTTP + # worker that outlives shutdown cleanup. + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state") + monkeypatch.setattr(download_lifecycle.threading, "Thread", _ImmediateThread) + real_register_worker = download_lifecycle.register_worker + registry = download_registry.DownloadRegistry() + key_a = download_registry.normalize_job_key("Org/Model::Q4_K_M") + key_b = download_registry.normalize_job_key("Org/Model::Q5_K_M") + assert registry.claim( + key_a, + download_registry.TRANSPORT_XET, + repo_type = "model", + repo_id = "Org/Model", + variant = "Q4_K_M", + blob_hashes = frozenset({"mainhash-a"}), + progress_blob_hashes = frozenset({"mainhash-a"}), + ) + assert registry.claim( + key_b, + download_registry.TRANSPORT_XET, + repo_type = "model", + repo_id = "Org/Model", + variant = "Q5_K_M", + blob_hashes = frozenset({"mainhash-b"}), + progress_blob_hashes = frozenset({"mainhash-b"}), + ) + # The sibling holds a live worker so it appears in terminate_all()'s snapshot. + sibling_proc = _make_proc(1, b"sibling xet") + assert registry.register_process(key_b, sibling_proc) + markers = [] + shutdown_calls = 0 + + def fake_persist_cancel_marker(*args, **_kwargs): + markers.append(args) + + def fake_spawn_worker(*_args, **_kwargs): + raise AssertionError("parked retry must not spawn a worker after shutdown") + + def fake_sleep(_seconds): + nonlocal shutdown_calls + shutdown_calls += 1 + # Shutdown fires while key_a's retry is parked in the reclaim wait loop. + registry.terminate_all("download") + # The killed sibling's watcher then clears the repo guard, which without a + # settled retry would let the loop reclaim the slot and spawn a worker. + registry.set_job(key_b, "cancelled") + + monkeypatch.setattr(download_lifecycle.time, "sleep", fake_sleep) + monkeypatch.setattr(download_lifecycle, "spawn_worker", fake_spawn_worker) + monkeypatch.setattr(download_registry, "persist_cancel_marker", fake_persist_cancel_marker) + + assert real_register_worker( + registry, + key_a, + _make_proc(1, b"xet failed a"), + hf_token = None, + label = "Org/Model [Q4_K_M]", + log_prefix = "Download", + logger = logging.getLogger("test"), + repo_type = "model", + repo_id = "Org/Model", + transport = download_registry.TRANSPORT_XET, + watch_name = "model-watch-a", + ) + + assert shutdown_calls == 1 + assert registry.get_job(key_a).state == "cancelled" + # The bailed retry still persists its XET cancel marker / orphan breadcrumb. + assert ("model", "Org/Model", "Q4_K_M", download_registry.TRANSPORT_XET) in markers + + +def test_shutdown_settles_registered_worker_that_exited_with_error(monkeypatch, tmp_path): + # A registered worker that already exited nonzero but whose watcher has not + # yet run is absent from terminate_all()'s live snapshot. The old `key in + # self._processes` skip left it `running`, so the watcher would later enter + # _try_http_retry() and spawn an HTTP worker after shutdown. terminate_all() + # must settle it (cancelling + pending cancel) and persist its cancel marker. + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state") + registry = download_registry.DownloadRegistry() + key = download_registry.normalize_job_key("Org/Model::Q4_K_M") + assert registry.claim( + key, + download_registry.TRANSPORT_XET, + repo_type = "model", + repo_id = "Org/Model", + variant = "Q4_K_M", + blob_hashes = frozenset({"mainhash"}), + progress_blob_hashes = frozenset({"mainhash"}), + ) + exited_proc = _make_proc(1, b"xet failed") + assert registry.register_process(key, exited_proc) + # The worker exited nonzero; drop_process has not run yet, so it stays in the + # process table while poll() already reports the error. + exited_proc.wait() + assert exited_proc.poll() == 1 + + markers = [] + monkeypatch.setattr( + download_registry, + "persist_cancel_marker", + lambda *args, **_kwargs: markers.append(args), + ) + + registry.terminate_all("download") + + # Settled so a later watcher pass sees the intentional stop and bails instead + # of retrying over HTTP, and its XET cancel marker is persisted by shutdown. + assert registry.get_job(key).state == "cancelling" + assert registry.cancel_requested(key) is True + assert ("model", "Org/Model", "Q4_K_M", download_registry.TRANSPORT_XET) in markers + + +def test_shutdown_leaves_cleanly_exited_registered_worker_alone(monkeypatch, tmp_path): + # A registered worker that exited cleanly (rc == 0) completed; its watcher + # will mark it done. terminate_all() must not mark it cancelling or persist a + # stale cancel marker for it. + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state") + registry = download_registry.DownloadRegistry() + key = download_registry.normalize_job_key("Org/Model::Q4_K_M") + assert registry.claim( + key, + download_registry.TRANSPORT_XET, + repo_type = "model", + repo_id = "Org/Model", + variant = "Q4_K_M", + blob_hashes = frozenset({"mainhash"}), + progress_blob_hashes = frozenset({"mainhash"}), + ) + done_proc = _make_proc(0) + assert registry.register_process(key, done_proc) + done_proc.wait() + assert done_proc.poll() == 0 + + markers = [] + monkeypatch.setattr( + download_registry, + "persist_cancel_marker", + lambda *args, **_kwargs: markers.append(args), + ) + + registry.terminate_all("download") + + assert registry.get_job(key).state == "running" + assert registry.cancel_requested(key) is False + assert markers == [] + + +def test_shutdown_persists_marker_for_parked_no_process_retry(monkeypatch, tmp_path): + # A parked XET->HTTP retry has dropped its worker, so it has no live process. + # run_lifespan_shutdown only awaits terminate_all(); if it returns before the + # daemon watcher wakes, terminate_all() itself must have persisted the XET + # cancel marker so the next launch keeps resumable/cancelled state. + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state") + registry = download_registry.DownloadRegistry() + key = download_registry.normalize_job_key("Org/Model::Q4_K_M") + # Retry runs over HTTP but its cancel marker must name the original XET slot. + assert registry.claim( + key, + download_registry.TRANSPORT_HTTP, + repo_type = "model", + repo_id = "Org/Model", + variant = "Q4_K_M", + blob_hashes = frozenset({"mainhash"}), + progress_blob_hashes = frozenset({"mainhash"}), + metadata_transport = download_registry.TRANSPORT_HTTP, + cancel_marker_transport = download_registry.TRANSPORT_XET, + ) + # No register_process: the retry is parked in the reclaim wait loop. + assert registry.get_process(key) is None + + markers = [] + monkeypatch.setattr( + download_registry, + "persist_cancel_marker", + lambda *args, **_kwargs: markers.append(args), + ) + + registry.terminate_all("download") + + # Persisted synchronously by terminate_all, not deferred to the watcher. + assert registry.get_job(key).state == "cancelling" + assert registry.cancel_requested(key) is True + assert ("model", "Org/Model", "Q4_K_M", download_registry.TRANSPORT_XET) in markers + + +def test_shutdown_leaves_registered_http_worker_error_alone(monkeypatch, tmp_path): + # A registered HTTP worker that exited nonzero on its own is a genuine terminal + # download failure, not a shutdown cancel and not retry-capable (only XET can + # spawn a post-shutdown HTTP retry). terminate_all() must NOT settle it or + # persist a cancel marker, or idle_status would report a real error as + # cancelled/resumable after restart. Its error status stays intact. + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state") + registry = download_registry.DownloadRegistry() + key = download_registry.normalize_job_key("Org/Model::Q4_K_M") + assert registry.claim( + key, + download_registry.TRANSPORT_HTTP, + repo_type = "model", + repo_id = "Org/Model", + variant = "Q4_K_M", + blob_hashes = frozenset({"mainhash"}), + progress_blob_hashes = frozenset({"mainhash"}), + metadata_transport = download_registry.TRANSPORT_HTTP, + ) + exited_proc = _make_proc(1, b"http failed") + assert registry.register_process(key, exited_proc) + # The worker exited nonzero; drop_process has not run yet, so it stays in the + # process table while poll() already reports the error. + exited_proc.wait() + assert exited_proc.poll() == 1 + + markers = [] + monkeypatch.setattr( + download_registry, + "persist_cancel_marker", + lambda *args, **_kwargs: markers.append(args), + ) + + registry.terminate_all("download") + + # Left running so the watcher records its real error; no cancel marker. + assert registry.get_job(key).state == "running" + assert registry.cancel_requested(key) is False + assert markers == [] + + +def test_has_active_peer_variant_sees_released_retry_peer(monkeypatch, tmp_path): + # A Q4 XET->HTTP retry peer between release_active_slot() and its reclaim is + # briefly absent from _repo_active while its job stays active and still owns + # the shared companion. Deleting a different variant (Q5) must still see it as + # an active peer so companion deletion is blocked; scanning only _repo_active + # would miss it and let the delete remove the shared mmproj out from under it. + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state") + registry = download_registry.DownloadRegistry() + key_q4 = download_registry.normalize_job_key("Org/Model::Q4_K_M") + assert registry.claim( + key_q4, + download_registry.TRANSPORT_XET, + repo_type = "model", + repo_id = "Org/Model", + variant = "Q4_K_M", + blob_hashes = frozenset({"mainhash-q4"}), + progress_blob_hashes = frozenset({"mainhash-q4"}), + ) + # The retry has released its slot but its job stays active while it waits to + # reclaim over HTTP. + registry.release_active_slot(key_q4) + assert registry.get_job(key_q4).state == "running" + assert key_q4 not in registry._repo_active.get("org/model", set()) + + # A different variant's delete must be blocked by the released Q4 peer. + assert registry.has_active_peer_variant("Org/Model", "Q5_K_M") is True + # Same variant is not a peer: the released-but-active scan still respects the + # target-variant filter. + assert registry.has_active_peer_variant("Org/Model", "Q4_K_M") is False diff --git a/studio/backend/hub/utils/download_registry.py b/studio/backend/hub/utils/download_registry.py index 777d63e1b5b..274038b2926 100644 --- a/studio/backend/hub/utils/download_registry.py +++ b/studio/backend/hub/utils/download_registry.py @@ -45,7 +45,7 @@ import threading import time import weakref -from dataclasses import dataclass, field +from dataclasses import dataclass, field, replace from pathlib import Path from typing import Iterator, Literal, Optional @@ -126,6 +126,9 @@ def write_worker_breadcrumb(key: str, pid: int, metadata: Optional["DownloadMeta "repo_id": metadata.repo_id if metadata is not None else None, "variant": metadata.variant if metadata is not None else None, "transport": metadata.transport if metadata is not None else None, + "cancel_marker_transport": metadata.cancel_marker_transport + if metadata is not None + else None, } tmp = path.with_name(f".{path.name}.tmp-{pid}") try: @@ -305,7 +308,7 @@ def reap_orphan_workers() -> None: data.get("repo_type"), repo_id, data.get("variant"), - data.get("transport"), + data.get("cancel_marker_transport") or data.get("transport"), ) except Exception as exc: logger.debug("Reaper failed for breadcrumb %s: %s", entry, exc) @@ -699,6 +702,7 @@ class DownloadMetadata: repo_id: str variant: Optional[str] transport: Optional[str] + cancel_marker_transport: Optional[str] = None # GGUF variant main/writable hashes, identifying the variant-specific shards # for concurrency decisions. blob_hashes: frozenset[str] = field(default_factory = frozenset) @@ -801,6 +805,7 @@ def __init__(self, max_terminal: int = 64) -> None: self._processes: dict[str, subprocess.Popen] = {} self._repo_active: dict[str, set[str]] = {} self._metadata: dict[str, DownloadMetadata] = {} + self._cancel_marker_transports: dict[str, str] = {} self._pending_cancel: dict[str, Optional[int]] = {} self._generations: dict[str, int] = {} # Monotonic across keys so an evicted then re-claimed key never reuses a @@ -839,6 +844,7 @@ def set_job( if state in TERMINAL_STATES: self._put_terminal_job_locked(key, state, error) self._pending_cancel.pop(key, None) + self._cancel_marker_transports.pop(key, None) repo = _repo_of_key(key) active = self._repo_active.get(repo) if active is not None: @@ -848,6 +854,57 @@ def set_job( else: self._jobs[key] = DownloadState(state, error) + def set_error_unless_cancelled( + self, key: str, error: str + ) -> tuple[JobState, Optional[DownloadMetadata]]: + key = normalize_job_key(key) + with self._lock: + current = self._jobs.get(key, DownloadState("idle")).state + has_pending_cancel = key in self._pending_cancel + pending_generation = self._pending_cancel.get(key) + metadata = self._metadata.get(key) + should_cancel = current == "cancelling" or ( + has_pending_cancel and self._generation_matches_locked(key, pending_generation) + ) + terminal_state: JobState = "cancelled" if should_cancel else "error" + marker_transport = self._cancel_marker_transports.pop(key, None) + if marker_transport is None and metadata is not None: + marker_transport = metadata.cancel_marker_transport + self._put_terminal_job_locked( + key, + terminal_state, + None if should_cancel else error, + ) + self._pending_cancel.pop(key, None) + repo = _repo_of_key(key) + active = self._repo_active.get(repo) + if active is not None: + active.discard(key) + if not active: + self._repo_active.pop(repo, None) + if should_cancel and metadata is not None and marker_transport is not None: + metadata = replace(metadata, transport = marker_transport) + return terminal_state, metadata + + def update_job_transport(self, key: str, transport: str) -> None: + key = normalize_job_key(key) + with self._lock: + metadata = self._metadata.get(key) + if metadata is None or metadata.transport == transport: + return + self._metadata[key] = replace(metadata, transport = transport) + + def release_active_slot(self, key: str) -> None: + key = normalize_job_key(key) + repo = _repo_of_key(key) + with self._lock: + active = self._repo_active.get(repo) + if active is None: + return + active.discard(key) + if not active: + self._repo_active.pop(repo, None) + def get_job(self, key: str) -> DownloadState: key = normalize_job_key(key) with self._lock: @@ -884,6 +941,14 @@ def register_process(self, key: str, proc: subprocess.Popen) -> bool: ): self._put_terminal_job_locked(key, "cancelled") metadata_to_persist = self._metadata.pop(key, None) + marker_transport = self._cancel_marker_transports.pop(key, None) + if marker_transport is None and metadata_to_persist is not None: + marker_transport = metadata_to_persist.cancel_marker_transport + if metadata_to_persist is not None and marker_transport is not None: + metadata_to_persist = replace( + metadata_to_persist, + transport = marker_transport, + ) repo = _repo_of_key(key) active = self._repo_active.get(repo) if active is not None: @@ -963,6 +1028,10 @@ def claim( blob_hashes: Optional[frozenset[str]] = None, progress_blob_hashes: Optional[frozenset[str]] = None, completed_baseline_bytes: int = 0, + generation: Optional[int] = None, + replace_active: bool = False, + metadata_transport: Optional[str] = None, + cancel_marker_transport: Optional[str] = None, ) -> tuple[bool, str]: key = normalize_job_key(key) repo = _repo_of_key(key) @@ -1007,10 +1076,13 @@ def claim( if conflict_state is not None: return False, conflict_state current = self._jobs.get(key, DownloadState("idle")).state - if current in _ACTIVE_STATES: + if current in _ACTIVE_STATES and not replace_active: return False, current - self._generation_seq += 1 - self._generations[key] = self._generation_seq + if generation is None: + self._generation_seq += 1 + self._generations[key] = self._generation_seq + else: + self._generations[key] = generation self._jobs[key] = DownloadState("running") self._repo_active.setdefault(repo, active).add(key) if repo_type and repo_id: @@ -1018,7 +1090,8 @@ def claim( repo_type = repo_type, repo_id = repo_id, variant = variant, - transport = transport, + transport = metadata_transport if metadata_transport is not None else transport, + cancel_marker_transport = cancel_marker_transport, blob_hashes = requested_hashes, progress_blob_hashes = requested_progress_hashes, completed_baseline_bytes = max( @@ -1026,8 +1099,13 @@ def claim( int(completed_baseline_bytes or 0), ), ) + if cancel_marker_transport is not None: + self._cancel_marker_transports[key] = cancel_marker_transport + else: + self._cancel_marker_transports.pop(key, None) else: self._metadata.pop(key, None) + self._cancel_marker_transports.pop(key, None) return True, "running" def adoptable(self, key: str) -> bool: @@ -1053,7 +1131,8 @@ def _delete_blocked_by_active_locked(self, repo_id: str, variant: Optional[str]) download. A variant delete conflicts only with that same variant or a whole-repo download writing the shared snapshot; other quantizations download concurrently and never block it.""" - for key in self._repo_active.get(repo_id, set()): + active_keys = self._repo_active.get(repo_id, set()) + for key in active_keys: job = self._jobs.get(key) if job is None or job.state not in _ACTIVE_STATES: continue @@ -1062,6 +1141,16 @@ def _delete_blocked_by_active_locked(self, repo_id: str, variant: Optional[str]) other_variant = self._active_job_variant_locked(key) if other_variant is None or other_variant == variant: return True + for key, job in self._jobs.items(): + if key in active_keys or _repo_of_key(key) != repo_id: + continue + if job.state not in _ACTIVE_STATES: + continue + if variant is None: + return True + other_variant = self._active_job_variant_locked(key) + if other_variant is None or other_variant == variant: + return True return False def peer_blob_hashes(self, key: str) -> frozenset[str]: @@ -1108,6 +1197,16 @@ def active_job_refs(self, repo_id: Optional[str] = None) -> list[ActiveDownloadR candidate_keys = list(self._repo_active.get(repo_key, set())) else: candidate_keys = [key for active in self._repo_active.values() for key in active] + # An XET->HTTP retry handoff briefly drops its key from _repo_active + # while its job stays active; include those released-but-active jobs + # so the waiting retry still lists and can be adopted or cancelled. + seen = set(candidate_keys) + for key, job in self._jobs.items(): + if key in seen or job.state not in _ACTIVE_STATES: + continue + if repo_key is not None and _repo_of_key(key) != repo_key: + continue + candidate_keys.append(key) refs: list[ActiveDownloadRef] = [] for key in candidate_keys: job = self._jobs.get(key) @@ -1169,12 +1268,25 @@ def has_active_peer_variant(self, repo_id: str, variant: Optional[str]) -> bool: repo_id = normalize_repo_key(repo_id) target = (variant or "").strip().lower() or None with self._lock: - for key in self._repo_active.get(repo_id, set()): + active_keys = self._repo_active.get(repo_id, set()) + for key in active_keys: job = self._jobs.get(key) if job is None or job.state not in _ACTIVE_STATES: continue if self._active_job_variant_locked(key) != target: return True + # An XET->HTTP retry peer between release_active_slot() and its reclaim + # is briefly absent from _repo_active while its job stays active and + # still owns the shared companion; mirror the released-but-active scan + # used by _delete_blocked_by_active_locked so it still blocks companion + # deletion of a different variant. + for key, job in self._jobs.items(): + if key in active_keys or _repo_of_key(key) != repo_id: + continue + if job.state not in _ACTIVE_STATES: + continue + if self._active_job_variant_locked(key) != target: + return True return False def request_cancel( @@ -1198,17 +1310,58 @@ def request_cancel( return True def terminate_all(self, kind: str = "download") -> None: + settled_no_proc: list[Optional[DownloadMetadata]] = [] with self._lock: live = [ (key, proc, self._metadata.get(key)) for key, proc in self._processes.items() if proc.poll() is None ] + live_keys = {key for key, _proc, _metadata in live} # Flag as an intentional stop so the watcher's exit classification # reports them cancelled rather than an OOM/crash once SIGKILL lands. for key, _proc, _metadata in live: if self._jobs.get(key, DownloadState("idle")).state == "running": self._jobs[key] = DownloadState("cancelling") + # Settle active jobs without a live worker too. Two cases: an + # XET->HTTP retry parked in the reclaim wait loop has dropped its + # worker and slot guard, so it is absent from `live`; and a + # registered worker that already exited with an error but whose + # watcher has not yet run would otherwise stay `running` and spawn an + # HTTP retry after this shutdown snapshot. Skip a registered worker + # that exited cleanly (rc == 0): it completed and the watcher will + # mark it done, so marking it cancelling would strand a stale marker. + for key, job in list(self._jobs.items()): + if job.state not in _ACTIVE_STATES or key in live_keys: + continue + proc = self._processes.get(key) + if proc is not None: + if proc.poll() == 0: + continue + # A registered worker that exited nonzero on its own over HTTP + # is a genuine terminal download failure, not a shutdown cancel + # and not retry-capable: leave its error status intact rather + # than persisting a cancel marker that would read as + # cancelled/resumable after restart. Only an exited XET worker + # could still spawn a post-shutdown HTTP retry, so only that + # needs settling here. + metadata = self._metadata.get(key) + if metadata is not None and metadata.transport == TRANSPORT_HTTP: + continue + self._pending_cancel[key] = self._generations.get(key) + self._jobs[key] = DownloadState("cancelling") + settled_no_proc.append(self._metadata.get(key)) + # Persist a cancel marker for each settled no-live-worker job outside the + # lock (mirroring the reaped path) so shutdown records resumable/cancelled + # state even if it returns before the daemon watcher wakes to do so. + for metadata in settled_no_proc: + if metadata is not None: + persist_cancel_marker( + metadata.repo_type, + metadata.repo_id, + metadata.variant, + metadata.cancel_marker_transport or metadata.transport, + ) reaped: list[tuple[str, subprocess.Popen, Optional[DownloadMetadata]]] = [] for key, proc, metadata in live: try: @@ -1222,7 +1375,7 @@ def terminate_all(self, kind: str = "download") -> None: metadata.repo_type, metadata.repo_id, metadata.variant, - metadata.transport, + metadata.cancel_marker_transport or metadata.transport, ) continue reaped.append((key, proc, metadata)) @@ -1242,7 +1395,7 @@ def terminate_all(self, kind: str = "download") -> None: metadata.repo_type, metadata.repo_id, metadata.variant, - metadata.transport, + metadata.cancel_marker_transport or metadata.transport, ) diff --git a/tests/studio/playwright_extra_ui.py b/tests/studio/playwright_extra_ui.py index 209a8a06f1a..32409261f3a 100644 --- a/tests/studio/playwright_extra_ui.py +++ b/tests/studio/playwright_extra_ui.py @@ -413,7 +413,10 @@ def shoot(name: str) -> None: soft_fail(f"chat-only mode should keep /export reachable; url={page.url}") else: unavailable = page.get_by_text(re.compile(r"Export unavailable", re.I)).first - if unavailable.count() == 0: + try: + # The export hardware probe settles asynchronously on slower runners. + unavailable.wait_for(state = "visible", timeout = 8000) + except Exception: soft_fail("chat-only /export did not show the export unavailable gate") else: info("OK chat-only /export rendered the unavailable gate")