feat: return retryable 503 for unknown models during startup warm-up

This commit is contained in:
9qeklajc
2026-10-02 01:23:29 +02:00
parent 058d8d1e64
commit 0bdfc1b811
3 changed files with 95 additions and 5 deletions
+2 -1
View File
@@ -148,12 +148,13 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
# await ensure_models_bootstrapped()
from ..proxy import get_upstreams
from ..proxy import get_upstreams, mark_warming_up
from ..upstream.helpers import refresh_upstreams_models_periodically
# Provider discovery hits every upstream's /models (plus the OpenRouter
# catalog for unpriced models), so keep it off the readiness path: the
# app serves as soon as the DB is up and fills its model maps after.
mark_warming_up()
bootstrap_task = asyncio.create_task(_bootstrap_providers_and_pricing())
btc_price_task = asyncio.create_task(update_prices_periodically())
+32 -1
View File
@@ -25,6 +25,7 @@ from .core.db import (
)
from .core.error_scope import (
ERROR_SCOPE_HEADER,
ERROR_SCOPE_NODE,
ERROR_SCOPE_UPSTREAM,
UPSTREAM_ERROR_STATUS,
UPSTREAM_UNAVAILABLE,
@@ -62,8 +63,14 @@ from .upstream.request_correction import correct_request, extract_error_message
logger = get_logger(__name__)
MODEL_PATH_HEADER = "x-routstr-model-path"
MODELS_WARMING_UP = "MODELS_WARMING_UP"
WARMING_UP_RETRY_AFTER_SECONDS = 2
proxy_router = APIRouter()
# Startup loads providers in the background, so until that first load finishes
# an unknown model may just not be loaded yet.
_warming_up = False
_upstreams: list[BaseUpstreamProvider] = []
_provider_map: dict[
str, list[tuple[Model, BaseUpstreamProvider]]
@@ -80,10 +87,23 @@ async def _finish_read_transaction(session: AsyncSession) -> None:
async def initialize_upstreams() -> None:
"""Initialize upstream providers from database during application startup."""
global _upstreams
global _upstreams, _warming_up
try:
_upstreams = await init_upstreams()
logger.info(f"Initialized {len(_upstreams)} upstream providers")
await refresh_model_maps()
finally:
_warming_up = False
def mark_warming_up() -> None:
"""Answer unknown models with a retryable 503 until initialize_upstreams() ends."""
global _warming_up
_warming_up = True
def is_warming_up() -> bool:
return _warming_up
async def reinitialize_upstreams() -> None:
@@ -650,6 +670,17 @@ async def _proxy(
candidates = get_candidates(model_id)
if not candidates:
if is_warming_up():
response = create_error_response(
"service_unavailable",
"Models are still loading, retry shortly",
503,
request=request,
code=MODELS_WARMING_UP,
error_scope=ERROR_SCOPE_NODE,
)
response.headers["Retry-After"] = str(WARMING_UP_RETRY_AFTER_SECONDS)
return response
return create_error_response(
"invalid_model", f"Model '{model_id}' not found", 400, request=request
)
+58
View File
@@ -0,0 +1,58 @@
"""Unknown models get a retryable 503 while startup is still loading providers."""
import json
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from routstr import proxy as proxy_module
from routstr.core.error_scope import ERROR_SCOPE_HEADER, ERROR_SCOPE_NODE
from .test_model_path_routing import _make_request, _run_proxy
@pytest.fixture(autouse=True)
def _reset_warming_up(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(proxy_module, "_warming_up", False)
def _chat_request() -> MagicMock:
return _make_request(
{"authorization": "Bearer sk-mpkey"},
json.dumps({"model": "not-loaded-yet"}).encode(),
)
@pytest.mark.asyncio
async def test_unknown_model_while_warming_up_is_retryable_503() -> None:
proxy_module.mark_warming_up()
response = await _run_proxy(_chat_request(), [])
assert response.status_code == 503
assert response.headers["Retry-After"] == "2"
assert response.headers[ERROR_SCOPE_HEADER] == ERROR_SCOPE_NODE
assert json.loads(bytes(response.body))["error"]["code"] == "MODELS_WARMING_UP"
@pytest.mark.asyncio
async def test_unknown_model_after_warm_up_is_invalid_model() -> None:
response = await _run_proxy(_chat_request(), [])
assert response.status_code == 400
assert json.loads(bytes(response.body))["error"]["type"] == "invalid_model"
@pytest.mark.asyncio
async def test_failed_initialization_still_ends_warm_up() -> None:
proxy_module.mark_warming_up()
with (
patch.object(
proxy_module, "init_upstreams", AsyncMock(side_effect=RuntimeError)
),
pytest.raises(RuntimeError),
):
await proxy_module.initialize_upstreams()
assert not proxy_module.is_warming_up()