Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
48 changes: 32 additions & 16 deletions src/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,11 @@
/health until the server (and model) is ready, and only then start the RunPod
serverless job loop so no job is pulled before the backend can serve it.

Before launching vLLM, a metadata-only pre-flight (model_preflight.py) checks
that the configured model is actually fetchable from the Hugging Face Hub, so
a typo'd MODEL_NAME or a gated model without HF_TOKEN fails in seconds instead
of at the end of the download.

If vLLM dies during startup for a reason a restart cannot fix (CUDA OOM, a
MAX_MODEL_LEN the GPU cannot hold, a bad flag, a gated model), the worker stays
up and answers every job with the cause instead of crash-looping; see
Expand All @@ -24,6 +29,7 @@
import urllib.error
import urllib.request

import model_preflight
import startup_errors
from args_builder import build_vllm_args
from download_model import LOCAL_MODEL_ARGS_PATH
Expand Down Expand Up @@ -154,22 +160,32 @@ def main() -> None:
for sig in (signal.SIGTERM, signal.SIGINT):
signal.signal(sig, _forward_signal)

vllm_process = start_vllm()
startup_error = None
try:
wait_for_vllm(vllm_process)
except RuntimeError as e:
logging.error("%s", e)
stop_vllm(vllm_process)
startup_error = startup_errors.classify("".join(recent_output), model=os.getenv("MODEL_NAME"))
if startup_error is None:
# Nothing recognisable, so let the platform restart us: a failed
# download or a bad host is worth another attempt.
sys.exit(1)
# A restart cannot fix this one. Exiting would crash-loop the worker,
# paying for the download and the load on every attempt and showing the
# user a traceback instead of a cause, so stay up and answer jobs with it.
logging.error("vLLM cannot start on this configuration; answering jobs with the cause: %s", startup_error)
# Ask the HF Hub whether the model is fetchable before paying for the
# download: a typo'd MODEL_NAME, a gated model without HF_TOKEN, or a bad
# MODEL_REVISION fails here in seconds instead of at the end of the cold
# start. Only definitive answers fail the boot; see model_preflight.py.
startup_error = model_preflight.check_model_access()
if startup_error:
logging.error(
"Model pre-flight failed; answering jobs with the cause instead of starting vLLM: %s",
startup_error,
)
else:
vllm_process = start_vllm()
try:
wait_for_vllm(vllm_process)
except RuntimeError as e:
logging.error("%s", e)
stop_vllm(vllm_process)
startup_error = startup_errors.classify("".join(recent_output), model=os.getenv("MODEL_NAME"))
if startup_error is None:
# Nothing recognisable, so let the platform restart us: a failed
# download or a bad host is worth another attempt.
sys.exit(1)
# A restart cannot fix this one. Exiting would crash-loop the worker,
# paying for the download and the load on every attempt and showing the
# user a traceback instead of a cause, so stay up and answer jobs with it.
logging.error("vLLM cannot start on this configuration; answering jobs with the cause: %s", startup_error)

# Import here (not at module import time) so the RunPod SDK and handler
# start only after the backend is confirmed healthy (or confirmed dead).
Expand Down
93 changes: 93 additions & 0 deletions src/model_preflight.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,93 @@
"""Fail "the model cannot be fetched" in seconds instead of after the cold start.

When `vllm serve` downloads the weights itself (no baked-in model), a typo'd
MODEL_NAME, a gated model without HF_TOKEN, or a bad MODEL_REVISION only
surfaces once the download attempt fails at the end of a ~20-minute cold
start. One metadata call to the Hugging Face Hub answers the same question in
seconds, before vLLM is even launched, so main.py asks it first and answers
jobs with the cause (never crash-looping, same rule as startup_errors.py).

Only a definitive answer from the Hub fails the boot. A timeout, a 5xx, or any
other ambiguity lets the boot proceed: vLLM makes its own download attempt and
startup_errors.classify still catches the crash, so a transient network blip
is never misreported as "model not found".
"""

import logging
import os
from typing import Mapping, Optional

from huggingface_hub import HfApi
from huggingface_hub.errors import (
GatedRepoError,
HfHubHTTPError,
RepositoryNotFoundError,
RevisionNotFoundError,
)

import startup_errors

# One metadata request; generous enough for a slow Hub, tiny next to the
# STARTUP_TIMEOUT it protects.
PREFLIGHT_TIMEOUT = float(os.getenv("MODEL_PREFLIGHT_TIMEOUT", "15"))

TRUE_VALUES = {"true", "1", "yes", "on"}


def _is_offline(env: Mapping[str, str]) -> bool:
"""Weights are already on disk (Option 2 builds); the Hub must not be hit."""
return any(
str(env.get(name, "")).strip().lower() in TRUE_VALUES
for name in ("HF_HUB_OFFLINE", "TRANSFORMERS_OFFLINE")
)


def check_model_access(env: Optional[Mapping[str, str]] = None) -> Optional[str]:
"""The startup_error for an unfetchable model, or None to proceed with boot.

Validates the same values args_builder.py hands to `vllm serve`
(MODEL_NAME, MODEL_REVISION, HF_TOKEN) with a metadata-only Hub call —
no weights are downloaded.
"""
env = os.environ if env is None else env

model = (env.get("MODEL_NAME") or "").strip()
if not model:
# MODEL / VLLM_CONFIG_FILE deploys name the model elsewhere; let vLLM
# resolve those itself.
return None
if _is_offline(env):
return None
if os.path.exists(model) or "://" in model:
# A local path or a non-HF source (e.g. the Run:ai model streamer);
# not the Hub's question to answer.
return None

revision = (env.get("MODEL_REVISION") or "").strip() or None
token = (env.get("HF_TOKEN") or "").strip() or None

try:
api = HfApi(token=token)
info = api.model_info(model, revision=revision, timeout=PREFLIGHT_TIMEOUT)
if getattr(info, "gated", False):
# Gated repos serve their metadata publicly; only a file-access
# check answers "may this worker download the weights?".
api.auth_check(model)
except GatedRepoError:
return startup_errors.gated_message(model)
except RevisionNotFoundError:
return startup_errors.revision_not_found_message(model, revision or "main")
except RepositoryNotFoundError:
return startup_errors.not_found_message(model)
except HfHubHTTPError as e:
status = getattr(getattr(e, "response", None), "status_code", None)
if status in (401, 403):
return startup_errors.gated_message(model)
logging.warning("Model pre-flight got HTTP %s from the HF Hub; proceeding with boot: %s", status, e)
return None
except Exception as e: # DNS failure, timeout, TLS trouble, ...
logging.warning("Model pre-flight could not reach the HF Hub; proceeding with boot: %s", e)
return None

logging.info("Model pre-flight: %s is fetchable", model if not revision else f"{model}@{revision}")
return None
39 changes: 29 additions & 10 deletions src/startup_errors.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,33 @@
_NO_SPACE = re.compile(r"No space left on device|ENOSPC|errno 28", re.I)


# The model-access messages are shared with model_preflight.py, which detects
# the same failures before the cold start; whether the user hits the fast path
# or the post-crash path, the wording must be identical.
def gated_message(named: str) -> str:
return (
f"{named} is gated or private on Hugging Face and the worker was not "
f"allowed to download it. Set HF_TOKEN to a token whose account has "
f"accepted the model's license, or pick a public model."
)


def not_found_message(named: str) -> str:
return (
f"{named} was not found on Hugging Face. Check MODEL_NAME for typos "
f"(it must be the full `org/repo` id) and MODEL_REVISION if set. Private "
f"repositories also return not-found until HF_TOKEN grants access."
)


def revision_not_found_message(named: str, revision: str) -> str:
return (
f"MODEL_REVISION={revision} does not exist for {named} on Hugging Face. "
f"Check the model page for the branch, tag, or commit hash, or unset "
f"MODEL_REVISION to use the default branch."
)


def classify(output: str, model: Optional[str] = None) -> Optional[str]:
"""One actionable message for a known fatal failure, or None to let it retry."""
named = model or "The model"
Expand Down Expand Up @@ -89,18 +116,10 @@ def classify(output: str, model: Optional[str] = None) -> Optional[str]:
)

if _GATED.search(output):
return (
f"{named} is gated or private on Hugging Face and the worker was not "
f"allowed to download it. Set HF_TOKEN to a token whose account has "
f"accepted the model's license, or pick a public model."
)
return gated_message(named)

if _NOT_FOUND.search(output):
return (
f"{named} was not found on Hugging Face. Check MODEL_NAME for typos "
f"(it must be the full `org/repo` id) and MODEL_REVISION if set. Private "
f"repositories also return not-found until HF_TOKEN grants access."
)
return not_found_message(named)

if _UNSUPPORTED_ARCH.search(output):
return (
Expand Down
3 changes: 3 additions & 0 deletions tests/requirements.txt
Original file line number Diff line number Diff line change
@@ -1,4 +1,7 @@
# Tests are CPU-only and never import vllm/torch. aiohttp is needed because
# src/handler.py imports it at module load (the proxy code is not exercised).
# huggingface-hub is needed because src/model_preflight.py imports its error
# types at module load (the Hub is never hit; HfApi is mocked).
pytest>=8,<10
aiohttp
huggingface-hub>=0.26
154 changes: 154 additions & 0 deletions tests/test_model_preflight.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,154 @@
"""Which model configurations fail fast at boot, and which are let through."""

import sys
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import MagicMock, patch

import pytest
from huggingface_hub.errors import (
GatedRepoError,
HfHubHTTPError,
RepositoryNotFoundError,
RevisionNotFoundError,
)

sys.path.insert(0, str(Path(__file__).resolve().parent.parent / "src"))

import model_preflight # noqa: E402
import startup_errors # noqa: E402
from model_preflight import check_model_access # noqa: E402


def hub_that_raises(error, gated=False):
"""A patched HfApi whose model_info raises `error` (or passes if None)."""
api = MagicMock()
if error is not None:
api.model_info.side_effect = error
else:
api.model_info.return_value = SimpleNamespace(gated=gated)
return patch.object(model_preflight, "HfApi", return_value=api), api


def hub_error(cls, message: str, status: int = 404):
"""Build a huggingface_hub error (they require a response since hub 1.x)."""
response = SimpleNamespace(status_code=status, headers={}, request=None)
return cls(message, response=response)


def http_error(status: int) -> HfHubHTTPError:
return hub_error(HfHubHTTPError, f"HTTP {status}", status)


class TestDefinitiveFailures:
def test_nonexistent_repo_fails_fast_with_the_not_found_message(self):
patcher, _ = hub_that_raises(hub_error(RepositoryNotFoundError, "404 Client Error"))
with patcher:
message = check_model_access({"MODEL_NAME": "definitely-not-a-real-org/nope-123"})

assert "definitely-not-a-real-org/nope-123 was not found" in message
assert "MODEL_NAME" in message

def test_gated_model_without_a_token_names_hf_token(self):
# Gated repos serve model_info publicly; the failure surfaces on the
# follow-up auth_check, exactly as on the real Hub.
patcher, api = hub_that_raises(None, gated="manual")
api.auth_check.side_effect = hub_error(GatedRepoError, "401 Client Error", 401)
with patcher:
message = check_model_access({"MODEL_NAME": "meta-llama/Llama-3.1-8B-Instruct"})

assert "gated or private" in message
assert "HF_TOKEN" in message

def test_gated_model_with_an_accepted_token_proceeds(self):
patcher, api = hub_that_raises(None, gated="manual")
with patcher:
assert check_model_access({"MODEL_NAME": "org/gated", "HF_TOKEN": "hf_ok"}) is None
api.auth_check.assert_called_once_with("org/gated")

def test_bad_revision_gets_its_own_message(self):
patcher, _ = hub_that_raises(hub_error(RevisionNotFoundError, "404 Client Error"))
with patcher:
message = check_model_access({"MODEL_NAME": "org/model", "MODEL_REVISION": "v9"})

assert "MODEL_REVISION=v9" in message
assert "org/model" in message

@pytest.mark.parametrize("status", [401, 403])
def test_plain_auth_rejections_read_as_gated(self, status):
patcher, _ = hub_that_raises(http_error(status))
with patcher:
assert "HF_TOKEN" in check_model_access({"MODEL_NAME": "org/model"})

def test_wording_matches_the_post_crash_classifier(self):
# The user must read the same sentence whether the failure is caught
# here in seconds or by startup_errors.classify after a crash.
patcher, api = hub_that_raises(None, gated="auto")
api.auth_check.side_effect = hub_error(GatedRepoError, "401", 401)
with patcher:
fast = check_model_access({"MODEL_NAME": "org/model"})
gated_crash = "huggingface_hub.errors.GatedRepoError: 401 Client Error."

assert fast == startup_errors.classify(gated_crash, model="org/model")


class TestAmbiguityLetsBootProceed:
@pytest.mark.parametrize(
"error",
[
ConnectionError("Connection reset by peer"),
TimeoutError("timed out"),
http_error(500),
http_error(503),
http_error(429),
OSError("Temporary failure in name resolution"),
],
)
def test_transient_errors_are_never_reported_as_not_found(self, error):
# A network blip must get a retry (vLLM makes its own attempt), not a
# fatal "model not found" the user cannot act on.
patcher, _ = hub_that_raises(error)
with patcher:
assert check_model_access({"MODEL_NAME": "org/model"}) is None


class TestValidModel:
def test_fetchable_model_passes_and_boot_proceeds(self):
patcher, api = hub_that_raises(None)
with patcher:
assert check_model_access({"MODEL_NAME": "Qwen/Qwen2.5-0.5B-Instruct"}) is None
api.model_info.assert_called_once()
api.auth_check.assert_not_called() # not gated: one metadata call is enough

def test_checks_the_same_values_vllm_will_use(self):
# args_builder maps MODEL_NAME -> --model, MODEL_REVISION -> --revision,
# and vLLM reads HF_TOKEN; the pre-flight must validate those, no others.
patcher, api = hub_that_raises(None)
with patcher as hf_api:
check_model_access(
{"MODEL_NAME": "org/model", "MODEL_REVISION": "v2", "HF_TOKEN": "hf_secret"}
)

hf_api.assert_called_once_with(token="hf_secret")
args, kwargs = api.model_info.call_args
assert args == ("org/model",)
assert kwargs["revision"] == "v2"


class TestSkippedConfigurations:
@pytest.mark.parametrize(
"env",
[
{}, # no MODEL_NAME: config-file / MODEL deploys resolve elsewhere
{"MODEL_NAME": "org/model", "HF_HUB_OFFLINE": "1"},
{"MODEL_NAME": "org/model", "TRANSFORMERS_OFFLINE": "1"},
{"MODEL_NAME": "s3://bucket/model"}, # non-HF source
{"MODEL_NAME": str(Path(__file__).parent)}, # local path
{"MODEL_NAME": "", "VLLM_CONFIG_FILE": "/config.yaml"},
],
)
def test_never_touches_the_hub(self, env):
patcher, api = hub_that_raises(None)
with patcher:
assert check_model_access(env) is None
api.model_info.assert_not_called()
Loading