diff --git a/litellm/proxy/management_endpoints/auto_router_usage.py b/litellm/proxy/management_endpoints/auto_router_usage.py index c4511251e65..ba34abc3b62 100644 --- a/litellm/proxy/management_endpoints/auto_router_usage.py +++ b/litellm/proxy/management_endpoints/auto_router_usage.py @@ -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) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index b8cc30ad8a7..988cae6ba66 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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 diff --git a/litellm/types/management_endpoints/auto_router_endpoints.py b/litellm/types/management_endpoints/auto_router_endpoints.py index e191470ec6e..1128062f8db 100644 --- a/litellm/types/management_endpoints/auto_router_endpoints.py +++ b/litellm/types/management_endpoints/auto_router_endpoints.py @@ -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.""" diff --git a/tests/unit/proxy/management_endpoints/test_auto_router_usage.py b/tests/unit/proxy/management_endpoints/test_auto_router_usage.py index fabd2bef130..dcb7d864fc5 100644 --- a/tests/unit/proxy/management_endpoints/test_auto_router_usage.py +++ b/tests/unit/proxy/management_endpoints/test_auto_router_usage.py @@ -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()