fix(studio): recover stalled Hub downloads over HTTP (#6858)
* fix(studio): recover stalled Hub downloads over HTTP * fix(studio): preserve retry generation and progress baseline * fix(studio): keep XET retry handoff nonterminal * fix(studio): preserve retry cancellation on claim failure * fix(studio): make retry failure cancellation atomic * fix(studio): close skipped retry state gaps * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Stabilize chat-only export gate detection on Windows * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Retrigger CI on a user-authored head * fix(studio): serialize XET HTTP retry handoff * List XET to HTTP retries that are briefly released from the repo guard as active downloads for PR #6858 * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Settle no-process active downloads on shutdown so a parked XET retry cannot spawn after cleanup for PR #6858 * Settle exited-error and no-process downloads on shutdown and persist their cancel markers for PR #6858 * Keep terminal HTTP failures uncancelled and block companion deletion for released retry peers for PR #6858 * Trim download lifecycle test coverage * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> Co-authored-by: Daniel Han <danielhanchen@gmail.com> Co-authored-by: Etherll <61019402+Etherll@users.noreply.github.com>
This commit is contained in:
parent
1c7bce427e
commit
49f2879cf8
3 changed files with 557 additions and 34 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue