fix(proxy): preserve state through composed lifespans (#44214)

This commit is contained in:
tin-berri 2026-10-02 10:56:52 -07:00 • committed by GitHub
parent 2c9a971171
commit 3ae491a06c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 138 additions and 17 deletions

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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.