mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(proxy): preserve state through composed lifespans (#44214)
This commit is contained in:
parent
2c9a971171
commit
3ae491a06c
4 changed files with 138 additions and 17 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue