Fix/adjust diffusion: round 21 P1+P2 batch for PR #5754

P1 #1 + #2: ``LoadRequest._no_embedded_hf_tokens`` and
``ValidateModelRequest._no_embedded_hf_tokens`` now cover
``gguf_variant`` in addition to ``model_path``. A caller could
pass a variant like ``Q4_K_M-hf_xxxxxxxx`` that flowed into
structured log sinks via the GGUF resolver path; the matching
``DiffusionLoadRequest`` validator already covered every string
field, so this restores parity.

P1 #3: ``/api/inference/unload`` now also matches the llama
``loading_model_identifier`` when picking the GGUF branch. A
pending GGUF download (``is_active`` still False,
``loading_model_identifier`` populated) used to fall through to
the safetensors branch and respond ``status="unloaded"`` while
llama-server kept downloading.

P1 #4 + #5: the final safetensors-handoff sweeps (route-level
``_release_safetensors_chat_for`` and backend
``_release_chat_backend_for_diffusion``) now check ``active_model_name``
and ``loading_models`` WITHOUT the initial ``owned_names`` filter.
A concurrent ``/load`` that landed AFTER the snapshot was
previously ignored, so a chat model that began loading during the
unload window let training / export / GGUF chat / diffusion start
anyway and race the new chat for VRAM.

P2 #6: added ``_preflight_diffusers_subfolder_config`` and
invoked it for GGUF loads with a transformer class
(``effective_base``, ``"transformer"``). A custom base companion
that had ``model_index.json`` but lacked
``transformer/config.json`` previously passed the round 19
preflight, unloaded chat, then failed inside
``from_single_file``.

P2 #7: ``_scrub_validation_obj`` in main.py also scrubs string
dict KEYS. Pydantic ``string_type`` errors surface ``input``
verbatim, and a malformed payload like
``{"repo_id": {"hf_xxxxx": "owner/repo"}}`` would otherwise leak
the token through the 422 response body.

All 85 diffusion-relevant + 35 model-validation tests pass
locally. Existing fakes for ``hf_hub_download`` updated to
accept the new ``subfolder=`` kwarg the round 21 preflight uses.

(P1 #3 cross-workload GPU handoff lock from round 20 is still
deferred; round 21's P1 #4 / #5 raised the sweep-level guarantee,
which closes the most common race without the deadlock risk of
holding a process-wide lock across the entire load.)
This commit is contained in:
Daniel Han-Chen 2026-05-25 10:32:43 +00:00
commit 04bd9b2da5
5 changed files with 146 additions and 34 deletions

View file

@ -201,6 +201,60 @@ def _expand_existing_local_path(value: str) -> str:
return value
def _preflight_diffusers_subfolder_config(
repo: str,
subfolder: str,
hf_token: Optional[str],
) -> None:
"""Round 21 P2 #6: also probe ``{subfolder}/config.json``.
The full-repo preflight at ``_preflight_full_diffusers_repo``
only proves ``model_index.json`` exists. For GGUF loads the
follow-up ``from_single_file(..., config=effective_base,
subfolder="transformer")`` still needs a matching
``transformer/config.json`` on the base companion. Without
this second probe a base that has model_index.json but no
transformer config would still unload chat before the load
failed.
"""
if not repo or not subfolder:
return
try:
local = Path(repo).expanduser()
except (OSError, ValueError):
local = None
if local is not None and local.exists():
config_path = local / subfolder / "config.json"
if not config_path.is_file():
raise RuntimeError(
f"Diffusion repo '{_display_repo_id(repo)}' is missing "
f"{subfolder}/config.json."
)
return
if (local is not None and local.is_absolute()) or repo.startswith("~"):
# Local-only path that does not exist -- _preflight_full_diffusers_repo
# already raised for the absent directory, so reaching here means the
# caller is loading a Hub id that just looks like a path. Fall through
# to the network probe.
pass
try:
from huggingface_hub import hf_hub_download as _hf_hub_download
except Exception:
return
try:
_hf_hub_download(
repo_id = repo,
filename = "config.json",
subfolder = subfolder,
token = hf_token,
)
except Exception as exc:
raise RuntimeError(
f"Could not access diffusion repo '{_display_repo_id(repo)}' "
f"{subfolder}/config.json before unloading the current model."
) from exc
def _preflight_full_diffusers_repo(repo: str, hf_token: Optional[str]) -> None:
"""Prove a full diffusers repo is accessible before any unloads.
@ -904,6 +958,21 @@ class DiffusionBackend:
# ``effective_base`` so a bad companion repo is
# caught BEFORE chat / export are released.
_preflight_full_diffusers_repo(effective_base, hf_token)
# Round 21 P2 #6: the GGUF transformer path also
# consumes ``effective_base`` via
# ``from_single_file(config=effective_base,
# subfolder="transformer")``. A base that has
# ``model_index.json`` but lacks
# ``transformer/config.json`` would pass the
# round-19 preflight and only fail AFTER the chat
# unload. Run the subfolder probe too so the
# second cheap failure mode is also caught early.
if gguf_filename and fam.transformer_class:
_preflight_diffusers_subfolder_config(
effective_base,
"transformer",
hf_token,
)
# Round 20 P1 #2: ``diffusers.GGUFQuantizationConfig``
# imports the ``gguf`` package lazily at construction
@ -1536,17 +1605,18 @@ def _release_chat_backend_for_diffusion() -> None:
)
_require_unload(loading)
# Round 19 P1 #2: final sweep using the initial snapshot of
# owned names. Catches races where a name we did not explicitly
# unload (because it appeared in loading_models between the
# snapshot and the unload calls) is still owned after the loop.
remaining_loading = (
set(getattr(backend, "loading_models", set()) or set()) & owned_names
)
# Round 21 P1 #5: final sweep without the owned_names filter.
# A concurrent ``/load`` that appeared AFTER the initial
# snapshot was previously ignored, so a chat model that started
# loading during the diffusion handoff slipped through and
# raced the diffusion allocation for VRAM. Treat ANY surviving
# active / loading entry as a failure so the surrounding
# load_model raises and the caller retries.
remaining_loading = set(getattr(backend, "loading_models", set()) or set())
remaining_active = getattr(backend, "active_model_name", None)
if remaining_loading or (remaining_active in owned_names):
if remaining_loading or remaining_active:
raise RuntimeError(
"The existing safetensors chat model is still active or loading "
"A safetensors chat model is still active or loading "
"after unload; retry before loading a diffusion image model."
)

View file

@ -333,7 +333,15 @@ def _scrub_validation_obj(value):
if isinstance(value, list):
return [_scrub_validation_obj(v) for v in value]
if isinstance(value, dict):
return {k: _scrub_validation_obj(v) for k, v in value.items()}
# Round 21 P2 #7: pydantic surfaces ``input`` for ``string_type``
# validation errors verbatim, including dict KEYS like
# ``{"hf_xxxxx": "owner/repo"}``. Scrub string keys too so the
# token does not leak through the 422 response body.
return {
(_scrub_validation_obj(k) if isinstance(k, str) else k):
_scrub_validation_obj(v)
for k, v in value.items()
}
return value

View file

@ -73,7 +73,12 @@ class LoadRequest(BaseModel):
def _no_identifier_control_chars(cls, v, info):
return _no_control_chars(v, info.field_name)
@field_validator("model_path")
# Round 21 P1 #1: also reject embedded HF tokens in
# ``gguf_variant``. A caller can pass a variant string like
# ``Q4_K_M-hf_xxxxxxxx`` that flows into log sinks via the
# GGUF resolver path; without this only ``model_path`` was
# protected.
@field_validator("model_path", "gguf_variant")
@classmethod
def _no_embedded_hf_tokens(cls, v, info):
return _reject_embedded_hf_token(v, info.field_name)
@ -171,7 +176,9 @@ class ValidateModelRequest(BaseModel):
def _no_identifier_control_chars(cls, v, info):
return _no_control_chars(v, info.field_name)
@field_validator("model_path")
# Round 21 P1 #2: extend embedded-token rejection to
# ``gguf_variant`` here too (mirrors LoadRequest).
@field_validator("model_path", "gguf_variant")
@classmethod
def _no_embedded_hf_tokens(cls, v, info):
return _reject_embedded_hf_token(v, info.field_name)

View file

@ -502,19 +502,20 @@ async def _release_safetensors_chat_for(workload: str) -> None:
)
await _unload_required(loading)
# Round 19 P1 #1: final sweep using the set of names that were
# initially present. Catches races where a model name we did not
# explicitly unload (because it appeared between the snapshot and
# the unload calls) is still in the owned set after the loop.
remaining_loading = (
set(getattr(inf, "loading_models", set()) or set()) & owned_names
)
# Round 21 P1 #4: final sweep without the owned_names filter.
# A concurrent ``/load`` that appeared AFTER the initial
# snapshot was previously ignored here, so a chat model that
# started loading during the unload window let the surrounding
# training / export / GGUF / diffusion start anyway. Treat ANY
# surviving active / loading entry as a failure so the caller
# retries rather than racing the new chat load for VRAM.
remaining_loading = set(getattr(inf, "loading_models", set()) or set())
remaining_active = getattr(inf, "active_model_name", None)
if remaining_loading or (remaining_active in owned_names):
if remaining_loading or remaining_active:
raise HTTPException(
status_code = 503,
detail = (
"The existing safetensors chat model is still active or loading "
"A safetensors chat model is still active or loading "
f"after unload; retry before starting {workload}."
),
)
@ -1667,12 +1668,32 @@ async def unload_model(
try:
# Check if the GGUF backend has this model loaded or is loading it
llama_backend = get_llama_cpp_backend()
if llama_backend.is_active and (
llama_backend.model_identifier == request.model_path
loaded_identifier = getattr(llama_backend, "model_identifier", None)
loading_identifier = getattr(llama_backend, "loading_model_identifier", None)
# Round 21 P1 #3: a GGUF download that has not yet flipped
# ``is_active`` to True (model_identifier still None,
# ``loading_model_identifier`` populated) used to fall
# through to the safetensors branch, which silently
# responded ``status="unloaded"`` while llama-server kept
# downloading. Match on either the loaded OR loading
# identifier so the explicit unload route can actually
# cancel a pending GGUF load.
llama_matches_request = (
loaded_identifier == request.model_path
or loading_identifier == request.model_path
or is_registered_native_path_label(
llama_backend.model_identifier, request.model_path
loaded_identifier, request.model_path
)
or not llama_backend.is_loaded
or is_registered_native_path_label(
loading_identifier, request.model_path
)
)
if (
getattr(llama_backend, "is_active", False)
or loading_identifier
) and (
llama_matches_request
or not getattr(llama_backend, "is_loaded", False)
):
# Round 19 P1 #6: previously this called
# ``llama_backend.unload_model()`` and unconditionally

View file

@ -393,10 +393,15 @@ def _install_fake_diffusers(monkeypatch, *, raise_on_pipeline = False):
monkeypatch.setitem(sys.modules, "diffusers", fake)
# Pretend HF Hub gave us a local file without actually fetching.
# Round 21: accept arbitrary kwargs (round 20 preflight adds
# ``filename="model_index.json"`` and round 21 preflight adds
# ``subfolder="transformer"``) so existing tests that exercise
# the GGUF path do not hit a TypeError from the fake signature.
fake_hub = types.ModuleType("huggingface_hub")
fake_hub.hf_hub_download = (
lambda repo_id, filename, token = None: f"/fake/{repo_id}/{filename}"
)
def _fake_download(repo_id, filename, token = None, subfolder = None, **_kwargs):
sub = f"{subfolder}/" if subfolder else ""
return f"/fake/{repo_id}/{sub}{filename}"
fake_hub.hf_hub_download = _fake_download
monkeypatch.setitem(sys.modules, "huggingface_hub", fake_hub)
# Force CPU dtype so the test does not need CUDA.
@ -1308,12 +1313,13 @@ def test_load_model_accepts_relative_local_dir(monkeypatch, tmp_path):
def _boom(**kwargs):
# Round 20 P1 #1 added a base-repo preflight that downloads
# the diffusers ``model_index.json`` of the auto-picked
# companion repo BEFORE the chat unload. Allow that call
# through (it would otherwise hit the network) but still
# reject any attempt to download the GGUF itself, which is
# what this test guards.
if kwargs.get("filename") == "model_index.json":
return "/tmp/model_index.json"
# companion repo BEFORE the chat unload. Round 21 P2 #6
# added a second preflight for ``transformer/config.json``
# on that same companion. Allow both preflight kinds through
# but still reject any attempt to download the GGUF itself,
# which is what this test guards.
if kwargs.get("filename") in ("model_index.json", "config.json"):
return "/tmp/preflight"
raise AssertionError("hf_hub_download must not run for a local dir")
fake_hub = SimpleNamespace(hf_hub_download = _boom)