diff --git a/studio/backend/hub/services/download_lifecycle.py b/studio/backend/hub/services/download_lifecycle.py index 50e20fc13f..e5d48872c1 100644 --- a/studio/backend/hub/services/download_lifecycle.py +++ b/studio/backend/hub/services/download_lifecycle.py @@ -8,6 +8,7 @@ import os import signal import subprocess import sys +import time import threading from pathlib import Path from typing import Callable, Optional @@ -210,7 +211,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). @@ -222,7 +225,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, @@ -230,6 +233,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( @@ -262,18 +267,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( - key, - "error", - stderr_text or f"worker exited with code {rc}", - ) + 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, + 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( + "%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( @@ -309,6 +522,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): @@ -319,7 +533,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, @@ -330,7 +551,25 @@ def register_worker( 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. @@ -426,8 +665,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 a4baafa317..87346573b0 100644 --- a/studio/backend/hub/tests/test_download_lifecycle.py +++ b/studio/backend/hub/tests/test_download_lifecycle.py @@ -1,27 +1,147 @@ # 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 logging + from hub.services import download_lifecycle +from hub.utils import download_registry, state_dir -def _set_xet_reason(monkeypatch, reason): +class _Proc: + pid = 4242 + + def __init__( + self, + rc, + stderr = b"", + ): + self.rc = rc + self.stderr = io.BytesIO(stderr) + self.waited = False + + 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): + pass + + +class _ImmediateThread: + def __init__(self, *, target, **_kwargs): + self.target = target + + def start(self): + self.target() + + +def test_resolve_effective_use_xet(monkeypatch): + for requested, unavailable_reason, expected in ( + (False, "unused", False), + (True, None, True), + (True, "hf_xet is not installed", False), + ): + monkeypatch.setattr( + download_lifecycle.download_registry, + "download_transport_unavailable_reason", + lambda _transport, reason = unavailable_reason: reason, + ) + assert download_lifecycle.resolve_effective_use_xet(requested) is expected + + +def test_xet_failure_retries_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) + register_worker = download_lifecycle.register_worker + + for repo_type, repo_id, variant, expected_args in ( + ("model", "Org/Model", "Q4_K_M", ["--repo-id", "Org/Model", "--variant", "Q4_K_M"]), + ("dataset", "Org/Data", None, ["--repo-id", "Org/Data", "--dataset"]), + ): + registry = download_registry.DownloadRegistry() + key = download_registry.normalize_job_key(f"{repo_id}::{variant}" if variant else repo_id) + assert registry.claim( + key, + download_registry.TRANSPORT_XET, + repo_type = repo_type, + repo_id = repo_id, + variant = variant, + blob_hashes = frozenset({"blob"}), + )[0] + generation = registry.current_generation(key) + spawned = [] + + def fake_spawn( + args, + _token, + *, + use_xet, + protected_blob_hashes = None, + ): + spawned.append((args, use_xet, protected_blob_hashes)) + return _Proc(0) + + def fake_retry_register(*_args, **kwargs): + assert kwargs["transport"] == download_registry.TRANSPORT_HTTP + return True + + monkeypatch.setattr(download_lifecycle, "spawn_worker", fake_spawn) + monkeypatch.setattr(download_lifecycle, "register_worker", fake_retry_register) + assert register_worker( + registry, + key, + _Proc(1, b"xet failed"), + hf_token = None, + label = repo_id, + 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", + ) + + metadata = registry.get_job_metadata(key) + assert spawned == [(expected_args, False, None)] + assert metadata.transport == download_registry.TRANSPORT_HTTP + assert metadata.blob_hashes == frozenset({"blob"}) + assert registry.current_generation(key) == generation + + +def test_http_failure_remains_terminal(monkeypatch, tmp_path): + monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state") + monkeypatch.setattr(download_lifecycle.threading, "Thread", _ImmediateThread) + 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", + )[0] monkeypatch.setattr( - download_lifecycle.download_registry, - "download_transport_unavailable_reason", - lambda _transport: reason, + download_lifecycle, + "register_worker", + lambda *_args, **_kwargs: (_ for _ in ()).throw( + AssertionError("HTTP failures must not retry") + ), ) - - -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 - - -def test_resolve_effective_use_xet_keeps_xet_when_available(monkeypatch): - _set_xet_reason(monkeypatch, None) - assert download_lifecycle.resolve_effective_use_xet(True) is True - - -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 + assert register_worker( + registry, + key, + _Proc(1, b"http failed"), + 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" diff --git a/studio/backend/hub/utils/download_registry.py b/studio/backend/hub/utils/download_registry.py index 777d63e1b5..274038b292 100644 --- a/studio/backend/hub/utils/download_registry.py +++ b/studio/backend/hub/utils/download_registry.py @@ -45,7 +45,7 @@ import sys 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 @@ class DownloadRegistry: 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 @@ class DownloadRegistry: 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 @@ class DownloadRegistry: 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 @@ class DownloadRegistry: ): 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 @@ class DownloadRegistry: 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 @@ class DownloadRegistry: 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 @@ class DownloadRegistry: 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 @@ class DownloadRegistry: 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 @@ class DownloadRegistry: 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 @@ class DownloadRegistry: 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 @@ class DownloadRegistry: 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 @@ class DownloadRegistry: 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 @@ class DownloadRegistry: 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 @@ class DownloadRegistry: 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 @@ class DownloadRegistry: metadata.repo_type, metadata.repo_id, metadata.variant, - metadata.transport, + metadata.cancel_marker_transport or metadata.transport, )