diff --git a/routstr/core/main.py b/routstr/core/main.py index 4defd7cb..54004699 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -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()) diff --git a/routstr/proxy.py b/routstr/proxy.py index a2cf9492..7c0cf908 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -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 - _upstreams = await init_upstreams() - logger.info(f"Initialized {len(_upstreams)} upstream providers") - await refresh_model_maps() + 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 ) diff --git a/tests/unit/test_models_warming_up.py b/tests/unit/test_models_warming_up.py new file mode 100644 index 00000000..99f286e7 --- /dev/null +++ b/tests/unit/test_models_warming_up.py @@ -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()