refactor(proxy): move auto-router usage query into Prisma client

This commit is contained in:
moe-berri 2026-09-25 10:36:16 -07:00
parent bdbbc94fe3
commit 2aa73e6453
4 changed files with 107 additions and 74 deletions

View file

@ -1,29 +1,28 @@
from collections.abc import Mapping, Sequence
from datetime import date, datetime, time, timedelta
from typing import Annotated, Final, Protocol
from fastapi import APIRouter, Depends, HTTPException, Query
from pydantic import BaseModel, TypeAdapter
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY
from litellm.proxy._types import UserAPIKeyAuth, user_api_key_has_admin_view
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.management_endpoints.common_utils import require_caller_user_id_for_non_admin
from litellm.types.management_endpoints.auto_router_endpoints import AutoRouterUsage
router: Final = APIRouter()
class AutoRouterUsage(BaseModel):
model: str
router_name: str | None
router_type: str | None
tier: str | None
requests: int
spend: float
class RoutingUsageDatabase(Protocol):
async def query_raw(self, query: str, *args: object) -> Sequence[Mapping[str, object]]: ...
async def get_auto_router_usage(
self,
*,
start_time: datetime,
end_time: datetime,
user_id: str | None = None,
api_key: str | None = None,
router_name: str | None = None,
router_type: str | None = None,
destination_model: str | None = None,
) -> tuple[AutoRouterUsage, ...]: ...
def routing_usage_database() -> RoutingUsageDatabase:
@ -31,35 +30,9 @@ def routing_usage_database() -> RoutingUsageDatabase:
if prisma_client is None:
raise HTTPException(status_code=503, detail="Database is not connected")
return prisma_client.db # pyright: ignore[reportReturnType] # wrapper forwards query_raw through __getattr__
return prisma_client
ROUTING_USAGE_SQL: Final = f"""
WITH requests AS (
SELECT model, spend,
CASE WHEN jsonb_typeof(metadata->'routing_decision') = 'object'
AND metadata->'routing_decision' <> '{{}}'::jsonb
THEN COALESCE(NULLIF(metadata#>>'{{routing_decision,router_model_name}}', ''), NULLIF(model_group, ''))
END AS router_name,
NULLIF(metadata#>>'{{routing_decision,router_type}}', '') AS router_type,
NULLIF(metadata#>>'{{routing_decision,tier}}', '') AS tier
FROM "LiteLLM_SpendLogs"
WHERE "startTime" >= $1::timestamp AND "startTime" < $2::timestamp
AND ($3::text IS NULL OR "user" = $3)
AND ($4::text IS NULL OR api_key = $4)
AND NULLIF(metadata->>'{INTERNAL_CALL_ORIGIN_METADATA_KEY}', '') IS NULL
AND ($7::text IS NULL OR model = $7)
)
SELECT model, router_name, router_type, tier, COUNT(*)::int AS requests,
COALESCE(SUM(spend), 0)::float8 AS spend
FROM requests
WHERE ($5::text IS NULL OR router_name = $5)
AND ($6::text IS NULL OR router_type = $6)
GROUP BY model, router_name, router_type, tier
ORDER BY spend DESC, model, router_name, tier
"""
_USAGE_ROWS: Final = TypeAdapter(tuple[AutoRouterUsage, ...])
MAX_ROUTING_USAGE_DAYS: Final = 93
@ -100,14 +73,12 @@ async def get_auto_router_usage(
)
if user_id is not None and user_id != scoped_user:
raise HTTPException(status_code=403, detail="Cannot view another user's routing usage")
rows: Final = await db.query_raw(
ROUTING_USAGE_SQL,
datetime.combine(start_date, time.min).isoformat(),
datetime.combine(end_date + timedelta(days=1), time.min).isoformat(),
scoped_user,
api_key,
router_name,
router_type,
destination_model,
return await db.get_auto_router_usage(
start_time=datetime.combine(start_date, time.min),
end_time=datetime.combine(end_date + timedelta(days=1), time.min),
user_id=scoped_user,
api_key=api_key,
router_name=router_name,
router_type=router_type,
destination_model=destination_model,
)
return _USAGE_ROWS.validate_python(rows)

View file

@ -51,6 +51,7 @@ from typing_extensions import ReadOnly, TypedDict
from litellm import _custom_logger_compatible_callbacks_literal
from litellm.constants import (
DEFAULT_MODEL_CREATED_AT_TIME,
INTERNAL_CALL_ORIGIN_METADATA_KEY,
LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL,
MAX_TEAM_LIST_LIMIT,
PROXY_REJECTED_BEFORE_ROUTING_KEY,
@ -239,6 +240,7 @@ from litellm.router_utils.common_utils import resolve_model_group_alias
from litellm.secret_managers.main import str_to_bool
from litellm.types.integrations.slack_alerting import DEFAULT_ALERT_TYPES
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.management_endpoints.auto_router_endpoints import AutoRouterUsage
from litellm.types.mcp import (
MCPDuringCallResponseObject,
MCPPreCallRequestObject,
@ -284,6 +286,33 @@ class _RelTuplesRow(TypedDict):
reltuples: ReadOnly[int]
_ROUTING_USAGE_SQL: Final = f"""
WITH requests AS (
SELECT model, spend,
CASE WHEN jsonb_typeof(metadata->'routing_decision') = 'object'
AND metadata->'routing_decision' <> '{{}}'::jsonb
THEN COALESCE(NULLIF(metadata#>>'{{routing_decision,router_model_name}}', ''), NULLIF(model_group, ''))
END AS router_name,
NULLIF(metadata#>>'{{routing_decision,router_type}}', '') AS router_type,
NULLIF(metadata#>>'{{routing_decision,tier}}', '') AS tier
FROM "LiteLLM_SpendLogs"
WHERE "startTime" >= $1::timestamp AND "startTime" < $2::timestamp
AND ($3::text IS NULL OR "user" = $3)
AND ($4::text IS NULL OR api_key = $4)
AND NULLIF(metadata->>'{INTERNAL_CALL_ORIGIN_METADATA_KEY}', '') IS NULL
AND ($7::text IS NULL OR model = $7)
)
SELECT model, router_name, router_type, tier, COUNT(*)::int AS requests,
COALESCE(SUM(spend), 0)::float8 AS spend
FROM requests
WHERE ($5::text IS NULL OR router_name = $5)
AND ($6::text IS NULL OR router_type = $6)
GROUP BY model, router_name, router_type, tier
ORDER BY spend DESC, model, router_name, tier
"""
_AUTO_ROUTER_USAGE_ROWS: Final = TypeAdapter(tuple[AutoRouterUsage, ...])
_VIEW_SETUP_POLL_INTERVAL_SECONDS: Final = 5.0
_VIEW_SETUP_DEADLINE_SECONDS: Final = 15 * 60.0
_VIEW_SETUP_GATE_TABLE: Final = '"LiteLLM_SpendLogs"'
@ -4774,6 +4803,29 @@ class PrismaClient:
raise e
async def get_auto_router_usage(
self,
*,
start_time: datetime,
end_time: datetime,
user_id: str | None = None,
api_key: str | None = None,
router_name: str | None = None,
router_type: str | None = None,
destination_model: str | None = None,
) -> tuple[AutoRouterUsage, ...]:
rows: Final[object] = await self.db.query_raw( # pyright: ignore[reportAny] # wrappers forward query_raw through __getattr__; validate rows below
_ROUTING_USAGE_SQL,
start_time.isoformat(),
end_time.isoformat(),
user_id,
api_key,
router_name,
router_type,
destination_model,
)
return _AUTO_ROUTER_USAGE_ROWS.validate_python(rows)
async def _query_first_with_cached_plan_fallback(self, sql_query: str, *args) -> dict | None:
"""
Execute a query, recovering once from PostgreSQL's "cached plan must not

View file

@ -210,6 +210,15 @@ class AutoRouterCacheStats(BaseModel):
ttl_1h_turns: int = Field(description="Turns whose cache write used the one-hour TTL")
class AutoRouterUsage(BaseModel):
model: str
router_name: str | None
router_type: str | None
tier: str | None
requests: int
spend: float
class AutoRouterBenchmarkTotals(BaseModel):
"""Session-shape and savings aggregates over auto-routed traffic in the window."""

View file

@ -1,4 +1,4 @@
from datetime import date
from datetime import date, datetime
from typing import Final
from unittest.mock import AsyncMock
@ -11,12 +11,12 @@ from litellm.proxy.management_endpoints.auto_router_usage import get_auto_router
class Database:
def __init__(self) -> None:
self.query_raw = AsyncMock(return_value=[])
self.get_auto_router_usage = AsyncMock(return_value=())
@pytest.mark.asyncio
@pytest.mark.parametrize("role", [LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY])
async def test_admin_filters_reach_query_as_parameters(role: LitellmUserRoles) -> None:
async def test_admin_filters_reach_prisma_client(role: LitellmUserRoles) -> None:
db = Database()
await get_auto_router_usage(
date(2026, 1, 1),
@ -28,14 +28,14 @@ async def test_admin_filters_reach_query_as_parameters(role: LitellmUserRoles) -
user_id="owner",
api_key="key-hash",
)
assert db.query_raw.call_args.args[1:] == (
"2026-01-01T00:00:00",
"2026-01-03T00:00:00",
"owner",
"key-hash",
"router'quoted",
"complexity",
None,
db.get_auto_router_usage.assert_awaited_once_with(
start_time=datetime(2026, 1, 1),
end_time=datetime(2026, 1, 3),
user_id="owner",
api_key="key-hash",
router_name="router'quoted",
router_type="complexity",
destination_model=None,
)
@ -49,14 +49,14 @@ async def test_non_admin_is_scoped_to_own_user_when_filter_omitted() -> None:
db,
destination_model="model-a",
)
assert db.query_raw.call_args.args[1:] == (
"2026-01-01T00:00:00",
"2026-01-02T00:00:00",
"own-user",
None,
None,
None,
"model-a",
db.get_auto_router_usage.assert_awaited_once_with(
start_time=datetime(2026, 1, 1),
end_time=datetime(2026, 1, 2),
user_id="own-user",
api_key=None,
router_name=None,
router_type=None,
destination_model="model-a",
)
@ -74,7 +74,7 @@ async def test_non_admin_cannot_read_other_users_or_unbound_service_account(call
user_id="another-user",
)
assert error.value.status_code == 403
db.query_raw.assert_not_called()
db.get_auto_router_usage.assert_not_called()
@pytest.mark.asyncio
@ -91,7 +91,7 @@ async def test_query_requires_one_specific_model_or_router(model: str | None, ro
router_name=router_name,
)
assert error.value.status_code == 400
db.query_raw.assert_not_called()
db.get_auto_router_usage.assert_not_called()
@pytest.mark.asyncio
@ -103,7 +103,7 @@ async def test_invalid_or_overlong_ranges_never_query_spend_logs(end: date) -> N
date(2026, 1, 1), end, UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), db, destination_model="fast"
)
assert error.value.status_code == 400
db.query_raw.assert_not_called()
db.get_auto_router_usage.assert_not_called()
@pytest.mark.asyncio
@ -113,7 +113,8 @@ async def test_maximum_range_includes_its_last_day() -> None:
date(2026, 1, 1), date(2026, 4, 3), UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), db,
destination_model="fast",
)
assert db.query_raw.call_args.args[1:3] == ("2026-01-01T00:00:00", "2026-04-04T00:00:00")
assert db.get_auto_router_usage.call_args.kwargs["start_time"] == datetime(2026, 1, 1)
assert db.get_auto_router_usage.call_args.kwargs["end_time"] == datetime(2026, 4, 4)
@pytest.mark.asyncio
@ -125,4 +126,4 @@ async def test_router_type_cannot_filter_a_destination_model() -> None:
destination_model="fast", router_type="complexity",
)
assert error.value.status_code == 400
db.query_raw.assert_not_called()
db.get_auto_router_usage.assert_not_called()