diff --git a/backend/main.py b/backend/main.py index 292ece48e7d..e0cef90c979 100644 --- a/backend/main.py +++ b/backend/main.py @@ -8,9 +8,13 @@ Run with: uvicorn backend.main:app --host 0.0.0.0 --port 4001 """ +from collections.abc import AsyncGenerator, Mapping from contextlib import asynccontextmanager +from typing import Final -from fastapi.routing import Mount +from starlette.applications import Starlette +from starlette.routing import Mount +from starlette.types import Lifespan # See gateway/main.py for why we assemble DATABASE_URL(s) here before # importing proxy_server. @@ -43,14 +47,16 @@ def _is_backend_route(route) -> bool: # See gateway/main.py for why the trim runs inside the lifespan instead of at # module scope. -_proxy_lifespan = app.router.lifespan_context +_proxy_lifespan: Final = app.router.lifespan_context @asynccontextmanager -async def _backend_lifespan(app_): - async with _proxy_lifespan(app_): +async def _backend_lifespan( + app_: Starlette, lifespan: Lifespan[Starlette] = _proxy_lifespan +) -> AsyncGenerator[Mapping[str, object], None]: + async with lifespan(app_) as state: app_.router.routes = [r for r in app_.router.routes if _is_backend_route(r)] - yield + yield state if state is not None else {} app.router.lifespan_context = _backend_lifespan diff --git a/gateway/main.py b/gateway/main.py index 61b885b27e4..fb4ae830808 100644 --- a/gateway/main.py +++ b/gateway/main.py @@ -9,9 +9,13 @@ Run with: uvicorn gateway.main:app --host 0.0.0.0 --port 4000 """ +from collections.abc import AsyncGenerator, Mapping from contextlib import asynccontextmanager +from typing import Final -from fastapi.routing import Mount +from starlette.applications import Starlette +from starlette.routing import Mount +from starlette.types import Lifespan # Assemble DATABASE_URL (+ DATABASE_URL_READ_REPLICA) from the discrete # DATABASE_* env vars before proxy_server imports spin up Prisma. Handles @@ -54,14 +58,16 @@ def _is_gateway_route(route) -> bool: # register routes. A module-load filter would miss routes added during # startup; running inside the lifespan, after the inner __aenter__, catches # them while still completing before uvicorn opens the listener. -_proxy_lifespan = app.router.lifespan_context +_proxy_lifespan: Final = app.router.lifespan_context @asynccontextmanager -async def _gateway_lifespan(app_): - async with _proxy_lifespan(app_): +async def _gateway_lifespan( + app_: Starlette, lifespan: Lifespan[Starlette] = _proxy_lifespan +) -> AsyncGenerator[Mapping[str, object], None]: + async with lifespan(app_) as state: app_.router.routes = [r for r in app_.router.routes if _is_gateway_route(r)] - yield + yield state if state is not None else {} app.router.lifespan_context = _gateway_lifespan diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index 0b470cf7bda..0119ce4430e 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -530,11 +530,11 @@ def _register_all_on_startup(inner: "Lifespan[FastAPI]", features: tuple[LazyFea (config pass-through endpoints), so the table is put back in lazy mode's order once it is up.""" @asynccontextmanager - async def lifespan(app: "FastAPI") -> AsyncGenerator[None]: + async def lifespan(app: "FastAPI") -> AsyncGenerator[Mapping[str, object]]: register_all_features(app, features) - async with inner(app): + async with inner(app) as state: _restore_registry_order(app, features) - yield + yield state if state is not None else {} return lifespan diff --git a/tests/unit/proxy/test_component_allowlists.py b/tests/unit/proxy/test_component_allowlists.py index 3641a2d9be9..1a210fcb445 100644 --- a/tests/unit/proxy/test_component_allowlists.py +++ b/tests/unit/proxy/test_component_allowlists.py @@ -26,7 +26,18 @@ RDS IAM token when ``IAM_TOKEN_DB_AUTH`` is set). import json import os import sys -from typing import Final +from collections.abc import AsyncGenerator, Mapping +from contextlib import asynccontextmanager +from functools import partial +from typing import Final, Literal + +import pytest +from starlette.applications import Starlette +from starlette.requests import Request +from starlette.responses import JSONResponse +from starlette.routing import Mount, Route +from starlette.testclient import TestClient +from starlette.types import Lifespan # Importing ``litellm.proxy.proxy_server`` runs its module-level setup, which # reads ``DATABASE_URL`` (Prisma) and ``LITELLM_MASTER_KEY``. Tier-zero CI @@ -43,7 +54,6 @@ _PRE_EXISTING_ENV = {key: os.environ.get(key) for key in _THROWAWAY_ENV} for _key, _value in _THROWAWAY_ENV.items(): os.environ.setdefault(_key, _value) -from fastapi.routing import Mount from prometheus_client import make_asgi_app # gateway/ and backend/ live at the repo root, not inside litellm/. @@ -53,6 +63,7 @@ if _REPO_ROOT not in sys.path: from backend.routes.allowlist import BACKEND_MOUNT_PATHS from gateway.routes.allowlist import GATEWAY_MOUNT_PATHS +from litellm.proxy._lazy_features import LazyFeature, attach_lazy_features from litellm.proxy.proxy_server import app from tests.test_litellm_rust.support.child_interpreter import run_child_interpreter @@ -74,7 +85,10 @@ _DB_ENV_KEYS = ( ) _PRE_DB_ENV = {_key: os.environ.pop(_key, None) for _key in _DB_ENV_KEYS} _PRE_COMPONENT_LIFESPAN = app.router.lifespan_context -from gateway.main import _is_gateway_route +from gateway.main import _gateway_lifespan, _is_gateway_route + +app.router.lifespan_context = _PRE_COMPONENT_LIFESPAN +from backend.main import _backend_lifespan app.router.lifespan_context = _PRE_COMPONENT_LIFESPAN for _key, _previous in _PRE_DB_ENV.items(): @@ -85,7 +99,7 @@ for _key, _previous in _PRE_DB_ENV.items(): _COVERAGE_PROBE: Final = """ import json, os, sys sys.path.insert(0, os.environ["LITELLM_COMPONENT_ALLOWLIST_REPO_ROOT"]) -from fastapi.routing import Mount +from starlette.routing import Mount from backend.routes.allowlist import BACKEND_EXACT_PATHS, BACKEND_PATH_PREFIXES from gateway.routes.allowlist import GATEWAY_EXACT_PATHS, GATEWAY_PATH_PREFIXES from litellm.proxy._lazy_features import loaded_lazy_modules @@ -112,6 +126,101 @@ json.dump({ """ +@pytest.mark.parametrize( + "component_lifespan", (None, _gateway_lifespan, _backend_lifespan), ids=("proxy", "gateway", "backend") +) +@pytest.mark.parametrize("eager", (False, True), ids=("lazy", "eager")) +@pytest.mark.parametrize("state_kind", ("enabled", "disabled", "stateless")) +def test_composed_lifespan_preserves_request_state_and_teardown( + monkeypatch: pytest.MonkeyPatch, + component_lifespan: Lifespan[Starlette] | None, + eager: bool, + state_kind: Literal["enabled", "disabled", "stateless"], +) -> None: + monkeypatch.setenv("LITELLM_DISABLE_LAZY_ROUTES", str(eager).lower()) + receiver: Final = object() + resource: Final = object() + state: Final[Mapping[str, object]] = { + "tracing_receiver": receiver if state_kind == "enabled" else None, + "other_resource": resource, + } + events: Final[list[str]] = [] # mutable-ok: observe startup, requests and teardown across the ASGI boundary + + async def trace_state(request: Request) -> JSONResponse: + events.append("request") + assert events[0] == "startup" and "shutdown" not in events + assert getattr(request.state, "other_resource", None) is (resource if state_kind != "stateless" else None) + assert getattr(request.state, "tracing_receiver", None) is (receiver if state_kind == "enabled" else None) + return JSONResponse({"keys": sorted(request.scope["state"])}) + + def register_trace_route(application: Starlette, module: object) -> None: + application.router.routes.append(Route("/v1/traces", trace_state)) + + @asynccontextmanager + async def stateful_lifespan(application: Starlette) -> AsyncGenerator[Mapping[str, object], None]: + events.append("startup") + application.router.routes.append(Route("/not-a-component-route", trace_state)) + try: + yield state + finally: + events.append("shutdown") + + @asynccontextmanager + async def stateless_lifespan(application: Starlette) -> AsyncGenerator[None, None]: + async with stateful_lifespan(application): + yield + + application: Final = type(app)(lifespan=stateless_lifespan if state_kind == "stateless" else stateful_lifespan) + feature: Final = LazyFeature("traces", __name__, ("/v1/traces",), register_fn=register_trace_route) + attach_lazy_features(application, (feature,)) + if component_lifespan is not None: + application.router.lifespan_context = partial(component_lifespan, lifespan=application.router.lifespan_context) + + with TestClient(application) as client: + response: Final = client.get("/v1/traces") + assert response.status_code == 200, response.text + assert response.json() == {"keys": [] if state_kind == "stateless" else sorted(state)} + filtered: Final = client.get("/not-a-component-route") + assert filtered.status_code == (200 if component_lifespan is None else 404), filtered.text + assert events == (["startup", "request", "request"] if component_lifespan is None else ["startup", "request"]) + assert events == ( + ["startup", "request", "request", "shutdown"] if component_lifespan is None else ["startup", "request", "shutdown"] + ) + + +@pytest.mark.parametrize( + "component_lifespan", (None, _gateway_lifespan, _backend_lifespan), ids=("proxy", "gateway", "backend") +) +@pytest.mark.parametrize("eager", (False, True), ids=("lazy", "eager")) +@pytest.mark.parametrize("phase", ("startup", "shutdown")) +def test_composed_lifespan_propagates_lifecycle_failures( + monkeypatch: pytest.MonkeyPatch, component_lifespan: Lifespan[Starlette] | None, eager: bool, phase: str +) -> None: + monkeypatch.setenv("LITELLM_DISABLE_LAZY_ROUTES", str(eager).lower()) + failure: Final = RuntimeError(f"{phase} failed") + events: Final[list[str]] = [] # mutable-ok: observe lifecycle events across the ASGI boundary + + @asynccontextmanager + async def inner_lifespan(application: Starlette) -> AsyncGenerator[Mapping[str, object], None]: + events.append("startup") + if phase == "startup": + raise failure + yield {} + events.append("shutdown") + raise failure + + application: Final = type(app)(lifespan=inner_lifespan) + attach_lazy_features(application, ()) + if component_lifespan is not None: + application.router.lifespan_context = partial(component_lifespan, lifespan=application.router.lifespan_context) + + with pytest.raises(RuntimeError) as caught: + with TestClient(application): + events.append("serving") + assert caught.value is failure + assert events == (["startup"] if phase == "startup" else ["startup", "serving", "shutdown"]) + + def test_gateway_plus_backend_covers_full_app(): """Every route on the proxy app must be served by gateway or backend.