mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
refactor(proxy): move auto-router usage query into Prisma client
This commit is contained in:
parent
bdbbc94fe3
commit
2aa73e6453
4 changed files with 107 additions and 74 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue