diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260801000000_add_auto_router_session_rollup/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260801000000_add_auto_router_session_rollup/migration.sql new file mode 100644 index 00000000000..e30748c83f0 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260801000000_add_auto_router_session_rollup/migration.sql @@ -0,0 +1,37 @@ +-- CreateTable +CREATE TABLE IF NOT EXISTS "LiteLLM_AutoRouterSession" ( + "session_id" TEXT NOT NULL, + "model_group" TEXT NOT NULL, + "router_kind" TEXT NOT NULL, + "baseline_model" TEXT, + "first_turn_at" TIMESTAMP(3) NOT NULL, + "last_turn_at" TIMESTAMP(3) NOT NULL, + "turns" INTEGER NOT NULL DEFAULT 0, + "turns_with_usage" INTEGER NOT NULL DEFAULT 0, + "total_tokens" BIGINT NOT NULL DEFAULT 0, + "ephemeral_5m_tokens" BIGINT NOT NULL DEFAULT 0, + "ephemeral_1h_tokens" BIGINT NOT NULL DEFAULT 0, + "spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0, + "baseline_spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0, + "same_model_turns" INTEGER NOT NULL DEFAULT 0, + "same_model_hits" INTEGER NOT NULL DEFAULT 0, + "first_visit_turns" INTEGER NOT NULL DEFAULT 0, + "first_visit_hits" INTEGER NOT NULL DEFAULT 0, + "return_turns" INTEGER NOT NULL DEFAULT 0, + "return_hits" INTEGER NOT NULL DEFAULT 0, + "stale_return_misses" INTEGER NOT NULL DEFAULT 0, + "savable_return_misses" INTEGER NOT NULL DEFAULT 0, + "rescued_spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0, + "replay_spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0, + "last_model" TEXT, + "model_state" JSONB NOT NULL DEFAULT '{}', + "updated_at" TIMESTAMP(3) NOT NULL, + + CONSTRAINT "LiteLLM_AutoRouterSession_pkey" PRIMARY KEY ("session_id","model_group") +); + +-- CreateIndex +CREATE INDEX IF NOT EXISTS "idx_auto_router_session_group_activity" ON "LiteLLM_AutoRouterSession"("model_group", "last_turn_at"); + +-- CreateIndex +CREATE INDEX IF NOT EXISTS "idx_auto_router_session_last_turn" ON "LiteLLM_AutoRouterSession"("last_turn_at"); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 0d7fa8692c8..502541f36cd 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -1391,6 +1391,46 @@ model LiteLLM_AdaptiveRouterSession { @@index([last_activity_at], map: "idx_adaptive_router_session_activity") } +// Per-(session, auto-router) rollup behind the auto-router benchmarks dashboard. +model LiteLLM_AutoRouterSession { + session_id String + model_group String + router_kind String + baseline_model String? + + first_turn_at DateTime + last_turn_at DateTime + turns Int @default(0) + turns_with_usage Int @default(0) + + total_tokens BigInt @default(0) + ephemeral_5m_tokens BigInt @default(0) + ephemeral_1h_tokens BigInt @default(0) + + spend Float @default(0.0) + baseline_spend Float @default(0.0) + + same_model_turns Int @default(0) + same_model_hits Int @default(0) + first_visit_turns Int @default(0) + first_visit_hits Int @default(0) + return_turns Int @default(0) + return_hits Int @default(0) + stale_return_misses Int @default(0) + savable_return_misses Int @default(0) + + rescued_spend Float @default(0.0) + replay_spend Float @default(0.0) + + last_model String? + model_state Json @default("{}") + updated_at DateTime @default(now()) @updatedAt + + @@id([session_id, model_group]) + @@index([model_group, last_turn_at], map: "idx_auto_router_session_group_activity") + @@index([last_turn_at], map: "idx_auto_router_session_last_turn") +} + // --------------------------------------------------------------------------- // Workflow Run Tracking // diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index b24ed4b8282..18cb1ed7c3e 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -57,6 +57,11 @@ from litellm.proxy.db.db_transaction_queue.tool_discovery_queue import ( ToolDiscoveryQueue, ) from litellm.proxy.route_llm_request import ROUTE_ENDPOINT_MAPPING +from litellm.proxy.spend_tracking.auto_router_session_queue import AutoRouterSessionQueue +from litellm.proxy.spend_tracking.auto_router_sessions import ( + auto_router_group_kinds, + turn_from_spend_payload, +) from litellm.proxy.spend_tracking.compression_savings import ( extract_compression_saved_tokens, ) @@ -115,6 +120,16 @@ def _extract_cache_creation_tokens(usage_obj: dict) -> int: return int(details.get("cache_write_tokens", 0) or details.get("cache_creation_tokens", 0) or 0) +def _parse_start_time(value: object) -> datetime | None: + if not isinstance(value, str): + return None + try: + return datetime.fromisoformat(value) + except ValueError: + verbose_proxy_logger.debug("auto_router_sessions: unparseable startTime %s", value) + return None + + class DBSpendUpdateWriter: """ Module responsible for @@ -138,6 +153,7 @@ class DBSpendUpdateWriter: self.daily_agent_spend_update_queue = DailySpendUpdateQueue() self.daily_org_spend_update_queue = DailySpendUpdateQueue() self.daily_tag_spend_update_queue = DailySpendUpdateQueue() + self.auto_router_session_queue = AutoRouterSessionQueue() async def update_database( # LiteLLM management object fields @@ -195,6 +211,8 @@ class DBSpendUpdateWriter: if team_id is not None and team_id != "": payload["team_id"] = team_id + await self._record_auto_router_turn(payload=payload, prisma_client=prisma_client) + if disable_spend_logs is False: await self._insert_spend_log_to_db( payload=payload, @@ -247,6 +265,86 @@ class DBSpendUpdateWriter: end_user_id, ) + async def _record_auto_router_turn( + self, + payload: SpendLogsPayload, + prisma_client: PrismaClient | None, + ) -> None: + """Fold one auto-routed turn into its session rollup. + + Hooked here rather than beside the daily transactions because this is the + one place a request passes through exactly once; the daily path runs per + entity type and would count every turn six times over. + + Independent of ``disable_spend_logs``: the rollup is what the benchmarks + dashboard reads, so turning off per-request logging must not also turn off + the aggregate that replaced it. + + Never raises. A dashboard rollup is not worth failing spend tracking over, + and this runs before the spend log insert, so anything escaping here would + cost the deployment money it could not account for. + """ + try: + await self._record_auto_router_turn_unsafe(payload=payload, prisma_client=prisma_client) + except Exception as e: # noqa: BLE001 # see docstring: this must never break the spend path + verbose_proxy_logger.warning("auto_router_sessions: skipped a turn (%s: %s)", type(e).__name__, e) + + async def _record_auto_router_turn_unsafe( + self, + payload: SpendLogsPayload, + prisma_client: PrismaClient | None, + ) -> None: + from litellm.proxy.proxy_server import llm_router + + session_id = payload.get("session_id") + model_group = payload.get("model_group") + model = payload.get("model") + start_time = payload.get("startTime") + if prisma_client is None or llm_router is None or not session_id or not model_group or not model: + return + router_kind = auto_router_group_kinds(llm_router).get(model_group) + if router_kind is None: + return + started_at = start_time if isinstance(start_time, datetime) else _parse_start_time(start_time) + if started_at is None: + return + + _metadata: SpendLogsMetadata = json.loads(payload["metadata"]) + usage_obj = _metadata.get("usage_object", {}) or {} # mutable-ok: empty fallback for an absent usage payload + cache_read_tokens = _extract_cache_read_tokens(usage_obj) + savings_spend = compute_savings_spend( + model=model, + custom_llm_provider=payload.get("custom_llm_provider", None), + compression_saved_tokens=extract_compression_saved_tokens(_metadata), + cache_read_input_tokens=cache_read_tokens, + routing_decision=_metadata.get("routing_decision"), + model_id=payload.get("model_id"), + llm_router=_get_llm_router, + usage_object=usage_obj, + cost_breakdown=_metadata.get("cost_breakdown"), + ) + await self.auto_router_session_queue.record_turn( + key=(session_id, model_group), + router_kind=router_kind, + # The same setting savings.py priced this turn against, stored on the row so + # the dashboard names the baseline the numbers were actually computed with + # rather than whatever the config says by the time someone opens the tab. + baseline_model=litellm.autorouter_savings_baseline_model, + turn=turn_from_spend_payload( + model=model, + started_at=started_at, + prompt_tokens=payload.get("prompt_tokens", 0) or 0, + completion_tokens=payload.get("completion_tokens", 0) or 0, + total_tokens=payload.get("total_tokens", 0) or 0, + spend=payload.get("spend", 0.0) or 0.0, + autorouter_savings=savings_spend.autorouter, + cache_read_tokens=cache_read_tokens, + cache_creation_tokens=_extract_cache_creation_tokens(usage_obj), + usage_object=usage_obj, + ), + prisma_client=prisma_client, + ) + async def _enqueue_tool_usage_transaction( self, payload: SpendLogsPayload, @@ -797,6 +895,11 @@ class DBSpendUpdateWriter: else: - Regular flow of this method """ + # Flushed outside the Redis buffer on purpose: session rollup writes are + # atomic increments, so two pods committing the same session compose + # correctly without being funnelled through a single elected writer. + await self.auto_router_session_queue.flush(prisma_client=prisma_client) + if RedisUpdateBuffer._should_commit_spend_updates_to_redis(): await self._commit_spend_updates_to_db_with_redis( prisma_client=prisma_client, diff --git a/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py b/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py index d7c70bdb20c..30537a41004 100644 --- a/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py +++ b/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py @@ -184,6 +184,19 @@ class SpendLogCleanup: time_column="start_time", ) + async def _delete_old_auto_router_sessions(self, prisma_client: PrismaClient, cutoff_date: datetime) -> int: + # Auto-router session rollups summarize spend logs, so they expire on the + # same cutoff. A row is keyed on when the session was last active, not + # when it started, so a long-running conversation is not pruned out from + # under itself while it is still being written to. + return await self._delete_old_rows_batched( + prisma_client, + cutoff_date, + table_name="LiteLLM_AutoRouterSession", + key_columns=("session_id", "model_group"), + time_column="last_turn_at", + ) + async def cleanup_old_spend_logs(self, prisma_client: PrismaClient) -> None: """ Main cleanup function. Deletes old spend logs in batches. @@ -243,6 +256,9 @@ class SpendLogCleanup: index_deleted = await self._delete_old_tool_index_rows(prisma_client, cutoff_date) verbose_proxy_logger.info(f"Deleted {index_deleted} expired tool index rows") + sessions_deleted = await self._delete_old_auto_router_sessions(prisma_client, cutoff_date) + verbose_proxy_logger.info(f"Deleted {sessions_deleted} expired auto-router session rollups") + except Exception as e: # .exception() captures the traceback; str(e) alone on a Prisma/DB # timeout is often empty and gives operators no signal to diagnose. diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 2ff39164d80..35e690009ae 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -20,6 +20,7 @@ from datetime import datetime, timedelta, timezone from types import UnionType from typing import ( TYPE_CHECKING, + Annotated, Any, Literal, Optional, @@ -16391,6 +16392,63 @@ async def get_adaptive_router_state( return {"routers": snapshots} +@router.get( + "/auto_router/benchmarks", + tags=["auto_router"], # mutable-ok: HTTPException takes a dict detail + dependencies=[Depends(user_api_key_auth)], # mutable-ok: FastAPI takes a list of dependencies +) +async def get_auto_router_benchmarks( + start_date: str, + end_date: str, + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], +): + """Session-level benchmarks for every configured auto-router. + + Admin-only. For each auto-router alias, reports turns per session, session + length, tokens per session, the dollar savings of the routed mix against the + counterfactual baseline, and how the provider prompt cache behaved. + + Reads the per-session rollup, never the per-request spend logs. + ``start_date`` / ``end_date`` are ``YYYY-MM-DD``; the window is clamped to + the most recent ``BENCHMARKS_MAX_WINDOW_DAYS`` days and the response echoes + the window actually served. Sessions are counted whole when they were active + in the window. Returns 404 when no auto-router is configured. + """ + from litellm.proxy.spend_tracking.auto_router_benchmarks import compute_benchmarks + from litellm.proxy.spend_tracking.auto_router_sessions import auto_router_group_kinds + + if not _user_has_admin_view(user_api_key_dict): + raise HTTPException( + status_code=403, + detail={ # mutable-ok: HTTPException takes a dict detail + "error": CommonProxyErrors.not_allowed_access.value + }, + ) + if llm_router is None: + raise HTTPException( + status_code=404, + detail={ # mutable-ok: HTTPException takes a dict detail + "error": "No auto_router is configured on this proxy." + }, + ) + group_kinds = auto_router_group_kinds(llm_router) + if not group_kinds: + raise HTTPException( + status_code=404, + detail={ # mutable-ok: HTTPException takes a dict detail + "error": "No auto_router is configured on this proxy." + }, + ) + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={ # mutable-ok: HTTPException takes a dict detail + "error": CommonProxyErrors.db_not_connected_error.value + }, + ) + return await compute_benchmarks(prisma_client, group_kinds, start_date, end_date) + + @router.get("/routes", dependencies=[Depends(user_api_key_auth)]) async def get_routes(): """ diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 0d7fa8692c8..502541f36cd 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -1391,6 +1391,46 @@ model LiteLLM_AdaptiveRouterSession { @@index([last_activity_at], map: "idx_adaptive_router_session_activity") } +// Per-(session, auto-router) rollup behind the auto-router benchmarks dashboard. +model LiteLLM_AutoRouterSession { + session_id String + model_group String + router_kind String + baseline_model String? + + first_turn_at DateTime + last_turn_at DateTime + turns Int @default(0) + turns_with_usage Int @default(0) + + total_tokens BigInt @default(0) + ephemeral_5m_tokens BigInt @default(0) + ephemeral_1h_tokens BigInt @default(0) + + spend Float @default(0.0) + baseline_spend Float @default(0.0) + + same_model_turns Int @default(0) + same_model_hits Int @default(0) + first_visit_turns Int @default(0) + first_visit_hits Int @default(0) + return_turns Int @default(0) + return_hits Int @default(0) + stale_return_misses Int @default(0) + savable_return_misses Int @default(0) + + rescued_spend Float @default(0.0) + replay_spend Float @default(0.0) + + last_model String? + model_state Json @default("{}") + updated_at DateTime @default(now()) @updatedAt + + @@id([session_id, model_group]) + @@index([model_group, last_turn_at], map: "idx_auto_router_session_group_activity") + @@index([last_turn_at], map: "idx_auto_router_session_last_turn") +} + // --------------------------------------------------------------------------- // Workflow Run Tracking // diff --git a/litellm/proxy/spend_tracking/auto_router_benchmarks.py b/litellm/proxy/spend_tracking/auto_router_benchmarks.py new file mode 100644 index 00000000000..91557869297 --- /dev/null +++ b/litellm/proxy/spend_tracking/auto_router_benchmarks.py @@ -0,0 +1,281 @@ +"""Read side of the auto-router benchmarks dashboard. + +Answers the customer question "what is the auto-router actually buying me": how +many turns a routed session runs, how long it lasts, how many tokens it burns, +how much cheaper the routed mix is than a single-model baseline, and how the +provider prompt cache behaves as the router moves a session between tiers. + +Every one of those was folded when the turn happened (see +``auto_router_sessions``), so this module only sums pre-folded per-session rows. +It never reads ``LiteLLM_SpendLogs``: the sequential facts behind the cache +numbers cannot be recovered from per-request rows without window functions over +the whole window, which is what this replaced. + +One aggregate query covers every auto-router, rather than four per router. +""" + +from collections.abc import Mapping +from datetime import datetime, timedelta, timezone +from types import MappingProxyType +from typing import TYPE_CHECKING, NamedTuple + +from pydantic import BaseModel, TypeAdapter + +from litellm.proxy.spend_tracking.auto_router_sessions import PROMPT_CACHE_TTL_SECONDS + +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient + +BENCHMARKS_MAX_WINDOW_DAYS = 30 + + +class AutoRouterCacheBenchmark(BaseModel): + """Provider prompt-cache behaviour for one auto-router. + + Sourced from the request's own ``cache_read_input_tokens``, which is the + provider's prefix cache. LiteLLM's ``cache_hit`` column is a different + mechanism entirely and reads false on turns the provider served from cache. + + The three turn buckets are mutually exclusive and exhaustive: every turn is + either the router staying put, arriving somewhere new, or coming back to a + tier this session already used, and they sum to ``turns``. + """ + + ttl_seconds: int + usage_coverage_pct: float + hit_rate_pct: float + turns: int + hits: int + same_model_turns: int + same_model_hits: int + first_visit_turns: int + first_visit_hits: int + return_turns: int + return_hits: int + same_model_hit_rate_pct: float + first_visit_hit_rate_pct: float + return_hit_rate_pct: float + stale_miss_share_pct: float + warming_savable_miss_pct: float + warming_break_even_pct: float + stale_return_misses: int + savable_return_misses: int + warming_rescued_spend: float + warming_replay_spend: float + warming_net_spend: float + + +class AutoRouterGroupBenchmark(BaseModel): + model_group: str + router_kind: str + baseline_model: str | None + sessions: int + turns: int + avg_turns_per_session: float + avg_session_length_seconds: float + total_tokens: int + avg_tokens_per_session: float + actual_spend: float + baseline_spend: float + savings: float + savings_pct: float + cache: AutoRouterCacheBenchmark | None + + +class AutoRouterBenchmarksResponse(BaseModel): + start_date: str + end_date: str + groups: tuple[AutoRouterGroupBenchmark, ...] + + +WARMING_BREAK_EVEN_PCT: Mapping[int, float] = MappingProxyType( + {300: 9.0, 3600: 5.0} # mutable-ok: a JSON object is a dict by definition +) # mutable-ok: frozen by MappingProxyType on this line + + +class _GroupRow(BaseModel): + """One folded auto-router, straight out of the aggregate.""" + + model_group: str + baseline_model: str | None + sessions: int + turns: int + total_session_seconds: float + total_tokens: int + actual_spend: float + baseline_spend: float + turns_with_usage: int + ephemeral_5m_tokens: int + ephemeral_1h_tokens: int + same_model_turns: int + same_model_hits: int + first_visit_turns: int + first_visit_hits: int + return_turns: int + return_hits: int + stale_return_misses: int + savable_return_misses: int + rescued_spend: float + replay_spend: float + + +_GROUP_ROWS = TypeAdapter(tuple[_GroupRow, ...]) + + +class _Window(NamedTuple): + start: str + end: str + + +_GROUP_SQL = """ +SELECT + model_group, + MAX(baseline_model) AS baseline_model, + COUNT(*)::bigint AS sessions, + COALESCE(SUM(turns), 0)::bigint AS turns, + COALESCE(SUM(EXTRACT(EPOCH FROM (last_turn_at - first_turn_at))), 0) AS total_session_seconds, + COALESCE(SUM(total_tokens), 0)::bigint AS total_tokens, + COALESCE(SUM(spend), 0.0) AS actual_spend, + COALESCE(SUM(baseline_spend), 0.0) AS baseline_spend, + COALESCE(SUM(turns_with_usage), 0)::bigint AS turns_with_usage, + COALESCE(SUM(ephemeral_5m_tokens), 0)::bigint AS ephemeral_5m_tokens, + COALESCE(SUM(ephemeral_1h_tokens), 0)::bigint AS ephemeral_1h_tokens, + COALESCE(SUM(same_model_turns), 0)::bigint AS same_model_turns, + COALESCE(SUM(same_model_hits), 0)::bigint AS same_model_hits, + COALESCE(SUM(first_visit_turns), 0)::bigint AS first_visit_turns, + COALESCE(SUM(first_visit_hits), 0)::bigint AS first_visit_hits, + COALESCE(SUM(return_turns), 0)::bigint AS return_turns, + COALESCE(SUM(return_hits), 0)::bigint AS return_hits, + COALESCE(SUM(stale_return_misses), 0)::bigint AS stale_return_misses, + COALESCE(SUM(savable_return_misses), 0)::bigint AS savable_return_misses, + COALESCE(SUM(rescued_spend), 0.0) AS rescued_spend, + COALESCE(SUM(replay_spend), 0.0) AS replay_spend +FROM "LiteLLM_AutoRouterSession" +WHERE model_group = ANY($1::text[]) + AND last_turn_at >= ($2::timestamptz AT TIME ZONE 'UTC') + AND first_turn_at < (($3::timestamptz + INTERVAL '1 day') AT TIME ZONE 'UTC') +GROUP BY model_group +""" + + +def clamp_window(start_date: str, end_date: str) -> _Window: + """Parse the range and enforce ``start >= end - BENCHMARKS_MAX_WINDOW_DAYS``. + + The returned start reflects the window actually served, which the response + echoes so the dashboard can label what it is showing rather than what it + asked for. + """ + start = datetime.fromisoformat(start_date).replace(tzinfo=timezone.utc) + end = datetime.fromisoformat(end_date).replace(tzinfo=timezone.utc) + floor = (end - timedelta(days=BENCHMARKS_MAX_WINDOW_DAYS)).replace(hour=0, minute=0, second=0, microsecond=0) + clamped = max(start, floor) + return _Window(start=clamped.date().isoformat(), end=end.date().isoformat()) + + +def _rate_pct(part: int, whole: int) -> float: + return (100.0 * part / whole) if whole else 0.0 + + +def summarize_cache(row: _GroupRow) -> AutoRouterCacheBenchmark | None: + """Fold one group's cache counters into the dashboard's cache view. + + ``stale_miss_share_pct`` narrows return-to-tier misses to those whose tier + had gone idle past the TTL; the rest missed because the prefix changed, which + keeping caches warm cannot fix. + + ``warming_savable_miss_pct`` narrows further and divides by every cache miss, + so it reads as the share of all misses a refresher could actually have + prevented. A miss qualifies only when the session returned to a tier it had + already used, that tier had gone idle past the TTL, and it came back within + two TTLs. The last bound is what one refresh fired just under the TTL can + bridge; a tier idle longer needs a replay per elapsed TTL, and since every + idle session pays those replays whether or not it returns, bridging further + costs more than the write it avoids. + """ + if row.turns_with_usage == 0: + return None + ttl_seconds = ( + PROMPT_CACHE_TTL_SECONDS["1h"] + if row.ephemeral_1h_tokens > 0 and row.ephemeral_1h_tokens >= row.ephemeral_5m_tokens + else PROMPT_CACHE_TTL_SECONDS["5m"] + ) + hits = row.same_model_hits + row.first_visit_hits + row.return_hits + bucketed_turns = row.same_model_turns + row.first_visit_turns + row.return_turns + return_misses = row.return_turns - row.return_hits + return AutoRouterCacheBenchmark( + ttl_seconds=ttl_seconds, + usage_coverage_pct=_rate_pct(row.turns_with_usage, row.turns), + hit_rate_pct=_rate_pct(hits, bucketed_turns), + turns=bucketed_turns, + hits=hits, + same_model_turns=row.same_model_turns, + same_model_hits=row.same_model_hits, + first_visit_turns=row.first_visit_turns, + first_visit_hits=row.first_visit_hits, + return_turns=row.return_turns, + return_hits=row.return_hits, + same_model_hit_rate_pct=_rate_pct(row.same_model_hits, row.same_model_turns), + first_visit_hit_rate_pct=_rate_pct(row.first_visit_hits, row.first_visit_turns), + return_hit_rate_pct=_rate_pct(row.return_hits, row.return_turns), + stale_miss_share_pct=_rate_pct(row.stale_return_misses, return_misses), + warming_savable_miss_pct=_rate_pct(row.savable_return_misses, bucketed_turns - hits), + warming_break_even_pct=WARMING_BREAK_EVEN_PCT[ttl_seconds], + stale_return_misses=row.stale_return_misses, + savable_return_misses=row.savable_return_misses, + warming_rescued_spend=row.rescued_spend, + warming_replay_spend=row.replay_spend, + warming_net_spend=row.rescued_spend - row.replay_spend, + ) + + +def summarize_group(row: _GroupRow, router_kind: str) -> AutoRouterGroupBenchmark | None: + """Fold one group's session rows into its benchmark. + + ``savings`` keeps its sign. A router that thrashes the prompt cache can cost + more than the baseline it is measured against, and an operator needs to be + able to see that rather than have it floored to zero. + """ + if row.sessions == 0: + return None + savings = row.baseline_spend - row.actual_spend + return AutoRouterGroupBenchmark( + model_group=row.model_group, + router_kind=router_kind, + baseline_model=row.baseline_model, + sessions=row.sessions, + turns=row.turns, + avg_turns_per_session=row.turns / row.sessions, + avg_session_length_seconds=row.total_session_seconds / row.sessions, + total_tokens=row.total_tokens, + avg_tokens_per_session=row.total_tokens / row.sessions, + actual_spend=row.actual_spend, + baseline_spend=row.baseline_spend, + savings=savings, + savings_pct=(100.0 * savings / row.baseline_spend) if row.baseline_spend > 0 else 0.0, + cache=summarize_cache(row), + ) + + +async def compute_benchmarks( + prisma_client: "PrismaClient", + group_kinds: Mapping[str, str], + start_date: str, + end_date: str, +) -> AutoRouterBenchmarksResponse: + """Aggregate the session rollup for every configured auto-router.""" + window = clamp_window(start_date, end_date) + raw = await prisma_client.db.query_raw( + _GROUP_SQL, + list(group_kinds.keys()), # mutable-ok: query_raw binds a list for the text[] parameter + window.start, + window.end, + ) + summarized = ( + summarize_group(row, group_kinds.get(row.model_group, "auto_router")) + for row in _GROUP_ROWS.validate_python(raw) + ) + return AutoRouterBenchmarksResponse( + start_date=window.start, + end_date=window.end, + groups=tuple(group for group in summarized if group is not None), + ) diff --git a/litellm/proxy/spend_tracking/auto_router_session_queue.py b/litellm/proxy/spend_tracking/auto_router_session_queue.py new file mode 100644 index 00000000000..52db6c78800 --- /dev/null +++ b/litellm/proxy/spend_tracking/auto_router_session_queue.py @@ -0,0 +1,242 @@ +"""In-memory aggregation and durable flush for auto-router session rollups. + +Follows ``AdaptiveRouterUpdateQueue``: the logging path only folds into memory, +and a background task drains the aggregate into Postgres with atomic increment +upserts, so two pods writing the same session compose rather than overwrite. + +The one departure is that this queue also caches the session state the fold reads +from. A pod that has never seen a session loads its row once and classifies from +memory thereafter, which is what keeps a session correct across a restart or a +move between pods without paying a read per turn. +""" + +import asyncio +from collections import OrderedDict +from dataclasses import dataclass, replace +from datetime import datetime, timezone +from functools import lru_cache +from typing import TYPE_CHECKING + +from litellm._logging import verbose_proxy_logger +from litellm.proxy.spend_tracking.auto_router_sessions import ( + EMPTY_SESSION_STATE, + SessionState, + TurnDelta, + TurnFacts, + counters_of, + fold_turn, + merge_deltas, + state_column, + state_from_row, +) +from litellm.repositories.table_repositories import AutoRouterSessionRepository + +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient + +SessionKey = tuple[str, str] + +DEFAULT_MAX_TRACKED_SESSIONS = 10_000 + + +@dataclass(frozen=True, slots=True) +class _Pending: + router_kind: str + baseline_model: str | None + first_turn_at: float + last_turn_at: float + delta: TurnDelta + + +@lru_cache(maxsize=1) +def _warn_pending_full(cap: int) -> None: + verbose_proxy_logger.warning( + "auto_router_sessions: %d sessions staged for the next flush; new sessions are not being recorded " + "until it drains. Benchmarks will undercount until then", + cap, + ) + + +def _merge_pending(earlier: _Pending, later: _Pending) -> _Pending: + """Fold two staged batches for one session, oldest first.""" + return _Pending( + router_kind=later.router_kind, + baseline_model=later.baseline_model or earlier.baseline_model, + first_turn_at=min(earlier.first_turn_at, later.first_turn_at), + last_turn_at=max(earlier.last_turn_at, later.last_turn_at), + delta=merge_deltas(earlier.delta, later.delta), + ) + + +def _epoch_to_datetime(value: float) -> datetime: + return datetime.fromtimestamp(value, tz=timezone.utc) + + +class AutoRouterSessionQueue: + """Folds auto-routed turns in memory and flushes them to the session rollup.""" + + def __init__(self, max_tracked_sessions: int = DEFAULT_MAX_TRACKED_SESSIONS) -> None: + self._pending: dict[SessionKey, _Pending] = {} # mutable-ok: drained and replaced wholesale on flush + self._state: OrderedDict[SessionKey, SessionState] = OrderedDict() # mutable-ok: bounded LRU cache + self._lock = asyncio.Lock() + self._max_tracked_sessions = max_tracked_sessions + + async def record_turn( + self, + key: SessionKey, + router_kind: str, + baseline_model: str | None, + turn: TurnFacts, + prisma_client: "PrismaClient", + ) -> None: + """Classify one turn against its session and stage the increments. + + The session id is caller-controlled, so the staged aggregate is capped: + past the cap a session that is already staged keeps accumulating, but a + new one is dropped rather than admitted. Without that bound a caller + sending a fresh id per request grows the aggregate without limit between + flushes, and benchmark rows are not worth an out-of-memory kill. + """ + loaded = await self._session_state(key, prisma_client) + async with self._lock: + current = self._pending.get(key) + if current is None and len(self._pending) >= self._max_tracked_sessions: + _warn_pending_full(self._max_tracked_sessions) + return + cached = self._state.get(key) + delta = fold_turn(cached if cached is not None else loaded, turn) + self._remember(key, delta.state) + self._pending[key] = ( + _Pending( + router_kind=router_kind, + baseline_model=baseline_model, + first_turn_at=turn.started_at, + last_turn_at=turn.started_at, + delta=delta, + ) + if current is None + else replace( + current, + baseline_model=baseline_model or current.baseline_model, + first_turn_at=min(current.first_turn_at, turn.started_at), + last_turn_at=max(current.last_turn_at, turn.started_at), + delta=merge_deltas(current.delta, delta), + ) + ) + + async def _session_state(self, key: SessionKey, prisma_client: "PrismaClient") -> SessionState: + """The session's state, from memory when this pod has seen it before. + + Loading outside the lock keeps a slow read from stalling every other + session's fold; a concurrent loader for the same key at worst repeats the + read, since both resolve to the same stored row. + """ + async with self._lock: + cached = self._state.get(key) + if cached is not None: + self._state.move_to_end(key) + return cached + return await self._load_state(key, prisma_client) + + async def _load_state(self, key: SessionKey, prisma_client: "PrismaClient") -> SessionState: + session_id, model_group = key + try: + row = await AutoRouterSessionRepository(prisma_client).table.find_unique( + where={ # mutable-ok: prisma's write API takes dict payloads + "session_id_model_group": { # mutable-ok: a JSON object is a dict by definition + "session_id": session_id, + "model_group": model_group, + } + } + ) + except Exception as e: # noqa: BLE001 # a read fault must not fail the spend write + verbose_proxy_logger.warning( + "auto_router_sessions: could not load session state for %s (%s); treating as a new session", key, e + ) + return EMPTY_SESSION_STATE + if row is None: + return EMPTY_SESSION_STATE + return state_from_row(row.last_model, row.last_turn_at, row.model_state) + + def _remember(self, key: SessionKey, state: SessionState) -> None: + """Cache the session's next state, evicting the least recently used. + + An evicted session is not lost; its next turn reloads the row it was + already flushed to, which costs one read and classifies identically. + """ + self._state[key] = state + self._state.move_to_end(key) + while len(self._state) > self._max_tracked_sessions: + self._state.popitem(last=False) + + async def flush(self, prisma_client: "PrismaClient") -> int: + """Drain the aggregate into the session rollup. Returns rows written. + + A session whose write fails is staged again rather than dropped. Draining + first and swallowing the error would lose that interval's turns, tokens + and spend permanently on any transient database fault, and because the + upsert is atomic a failure means nothing landed, so replaying it cannot + double-count. + """ + async with self._lock: + batch = self._pending + self._pending = {} # mutable-ok: fresh aggregate for the next interval + + failed = { # mutable-ok: built once from the writes that did not land + key: batch[key] for key in sorted(batch.keys()) if not await self._write(key, batch[key], prisma_client) + } + if failed: + verbose_proxy_logger.warning( + "auto_router_sessions: %d of %d session writes failed; re-staging them for the next flush", + len(failed), + len(batch), + ) + async with self._lock: + for key, pending in failed.items(): + current = self._pending.get(key) + # The retried batch is older than anything staged since, so it + # merges underneath it and the state of the newer one wins. + self._pending[key] = pending if current is None else _merge_pending(pending, current) + return len(batch) - len(failed) + + async def _write(self, key: SessionKey, pending: _Pending, prisma_client: "PrismaClient") -> bool: + session_id, model_group = key + counters = counters_of(pending.delta) + shared = { # mutable-ok: prisma's write API takes dict payloads + "last_turn_at": _epoch_to_datetime(pending.last_turn_at), + "last_model": pending.delta.state.last_model, + "model_state": state_column(pending.delta.state), + "baseline_model": pending.baseline_model, + } + try: + await AutoRouterSessionRepository( + prisma_client + ).table.upsert( + where={ # mutable-ok: prisma's write API takes dict payloads + "session_id_model_group": { # mutable-ok: a JSON object is a dict by definition + "session_id": session_id, + "model_group": model_group, + } + }, + data={ # mutable-ok: prisma's write API takes dict payloads + "create": { # mutable-ok: prisma's write API takes dict payloads + "session_id": session_id, + "model_group": model_group, + "router_kind": pending.router_kind, + "first_turn_at": _epoch_to_datetime(pending.first_turn_at), + **shared, + **counters, + }, + "update": { # mutable-ok: prisma's write API takes dict payloads + **{ # mutable-ok: a JSON object is a dict by definition + field: {"increment": value} # mutable-ok: a JSON object is a dict by definition + for field, value in counters.items() # mutable-ok: spread into the prisma payload immediately below + }, # mutable-ok: spread into the prisma payload immediately below + **shared, + }, + }, + ) + except Exception as e: # noqa: BLE001 # one session's write must not drop the rest of the batch + verbose_proxy_logger.exception("auto_router_sessions: failed to flush session %s (%s)", key, e) + return False + return True diff --git a/litellm/proxy/spend_tracking/auto_router_sessions.py b/litellm/proxy/spend_tracking/auto_router_sessions.py new file mode 100644 index 00000000000..8f7f7e7f946 --- /dev/null +++ b/litellm/proxy/spend_tracking/auto_router_sessions.py @@ -0,0 +1,450 @@ +"""Per-session rollup of auto-router traffic, folded one turn at a time. + +The benchmarks dashboard used to answer every question by scanning +``LiteLLM_SpendLogs`` at read time, deriving each turn's meaning from window +functions over the per-request rows: which model the previous turn used, how long +a tier had been idle, how big the prefix was last time. Those are sequential +facts, and the request that produces them already knows all of them. This module +computes them once, when the turn happens, and folds the answer into a durable +per-(session, auto-router) row. + +The fold is pure. ``fold_turn`` takes the session's prior state and one turn's +facts and returns the increments plus the next state, with no I/O and no clock, +so every rate and dollar formula is testable in isolation. + +Dollars are computed here rather than at read time because the rates belong to +the model that served the turn, and a rollup row has already summed across +models. This is the same reason ``savings.py`` prices in the spend writer. +""" + +from collections.abc import Mapping +from dataclasses import dataclass, fields +from datetime import datetime, timezone +from types import MappingProxyType +from typing import TYPE_CHECKING, Literal, Protocol + +from pydantic import BaseModel, TypeAdapter, ValidationError + +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.router_utils.auto_router_model_naming import classify_strategy_router_model + +if TYPE_CHECKING: + from litellm.router import Router + +PROMPT_CACHE_TTL_SECONDS: Mapping[str, int] = MappingProxyType( + {"5m": 300, "1h": 3600} # mutable-ok: a JSON object is a dict by definition +) # mutable-ok: frozen by MappingProxyType on this line + +TurnBucket = Literal["same_model", "first_visit", "return"] + + +class RateLookup(Protocol): + """Per-token ``(cache_read, cache_write)`` prices for a model at a TTL.""" + + def __call__(self, model: str, ttl_seconds: int) -> tuple[float, float]: ... + + +@dataclass(frozen=True, slots=True) +class ModelMark: + """What a session remembers about a model it has already been served on. + + ``provisioned_replay_spend`` is the replay this model is currently charged + for on the assumption that the turn which set it was the session's last on + that model. A refresher fires once per idle window whether or not the caller + ever comes back, so every use has to carry that charge until the session + proves it returned inside the TTL, at which point the charge is withdrawn. + """ + + last_used_at: float + provisioned_replay_spend: float + + +@dataclass(frozen=True, slots=True) +class SessionState: + """The prior turns of one session on one auto-router, compressed. + + Everything ``fold_turn`` needs to classify the next turn, and nothing else; + this is what the ``model_state`` column round-trips. + """ + + last_model: str | None + last_turn_at: float + model_marks: Mapping[str, ModelMark] + + +EMPTY_SESSION_STATE = SessionState( + last_model=None, + last_turn_at=0.0, + model_marks=MappingProxyType({}), # mutable-ok: frozen by MappingProxyType on this line +) + + +@dataclass(frozen=True, slots=True) +class TurnFacts: + """One auto-routed request, as the spend writer sees it. + + ``autorouter_savings`` arrives already computed by + ``savings.compute_autorouter_savings`` so that the benchmarks tab and the + usage tab cannot report different savings for the same traffic; the + counterfactual baseline is reconstructed from it rather than priced again. + """ + + model: str + started_at: float + prompt_tokens: int + completion_tokens: int + total_tokens: int + cache_read_tokens: int + cache_creation_tokens: int + ephemeral_5m_tokens: int + ephemeral_1h_tokens: int + spend: float + autorouter_savings: float + has_usage: bool + + +@dataclass(frozen=True, slots=True) +class TurnDelta: + """Increments one turn contributes to its session row, plus the next state. + + Every field but ``state`` is additive, so the flusher can hand them straight + to an atomic increment upsert; that is what lets two pods writing the same + session compose rather than overwrite each other. + + The counters are declared here and nowhere else. ``COUNTER_FIELDS`` derives + from this declaration and the merge, the flush payload and the read query all + build off it, so a metric added here reaches the database and the dashboard + without a second edit. Enumerating the same names in five places is how a + rollup field ends up written but never read. + """ + + state: SessionState + turns: int = 0 + turns_with_usage: int = 0 + total_tokens: int = 0 + ephemeral_5m_tokens: int = 0 + ephemeral_1h_tokens: int = 0 + spend: float = 0.0 + baseline_spend: float = 0.0 + same_model_turns: int = 0 + same_model_hits: int = 0 + first_visit_turns: int = 0 + first_visit_hits: int = 0 + return_turns: int = 0 + return_hits: int = 0 + stale_return_misses: int = 0 + savable_return_misses: int = 0 + rescued_spend: float = 0.0 + replay_spend: float = 0.0 + + +COUNTER_FIELDS: tuple[str, ...] = tuple(f.name for f in fields(TurnDelta) if f.name != "state") + + +def counters_of(delta: TurnDelta) -> Mapping[str, float]: + """The additive part of a delta, keyed the way the rollup columns are named.""" + return {name: getattr(delta, name) for name in COUNTER_FIELDS} # mutable-ok: a fresh per-call payload + + +def turn_ttl_seconds(turn: TurnFacts) -> int: + """The prompt-cache TTL this turn was written under. + + Read from the turn's own ``cache_creation`` split rather than guessed for + the window, so a deployment mixing both TTLs is scored per request instead of + having one regime imposed on all of it. Absent any ephemeral breakdown the + provider default of five minutes applies; treating no evidence as the one + hour tier would silently move every staleness verdict. + """ + if turn.ephemeral_1h_tokens > 0 and turn.ephemeral_1h_tokens >= turn.ephemeral_5m_tokens: + return PROMPT_CACHE_TTL_SECONDS["1h"] + return PROMPT_CACHE_TTL_SECONDS["5m"] + + +def cache_rates(model: str, ttl_seconds: int) -> tuple[float, float]: + """``(cache_read, cache_write)`` per-token costs for a model at a TTL. + + Tries the name as given and then bare, because spend rows carry models + provider-prefixed while the cost map often keys them bare, and a single + lookup would silently price the turn at zero. Falls open to zero rates, which + surfaces as no warming economics rather than a raised error inside the spend + writer. + """ + for candidate in _pricing_candidates(model): + try: + info = litellm.get_model_info(model=candidate) + except Exception as e: # noqa: BLE001 # get_model_info raises bare Exception for unmapped models + verbose_proxy_logger.debug("auto_router_sessions: no model info for %s (%s)", candidate, e) + continue + read = float(info.get("cache_read_input_token_cost") or 0.0) + write_5m = float(info.get("cache_creation_input_token_cost") or 0.0) + write_1h = float(info.get("cache_creation_input_token_cost_above_1hr") or 0.0) + return read, (write_1h or write_5m) if ttl_seconds >= PROMPT_CACHE_TTL_SECONDS["1h"] else write_5m + return 0.0, 0.0 + + +def _pricing_candidates(model: str) -> tuple[str, ...]: + stripped = model.split("/", 1)[1] if "/" in model else model + return tuple(dict.fromkeys((model, stripped))) + + +def _bucket(state: SessionState, turn: TurnFacts) -> TurnBucket: + """Which of three mutually exclusive things the router did on this turn. + + The session's opening turn is a first visit to whatever tier served it, which + is what makes the three buckets exhaustive: every turn lands in exactly one, + and they sum to the turn count. The previous split left a session's first + turn in none of them, so the bucket totals silently disagreed with the + headline. + """ + if state.last_model is None: + return "first_visit" + if turn.model == state.last_model: + return "same_model" + return "return" if turn.model in state.model_marks else "first_visit" + + +def fold_turn(state: SessionState, turn: TurnFacts, rates: RateLookup = cache_rates) -> TurnDelta: + """Fold one turn into its session, returning the increments and next state. + + A turn that arrives out of order still contributes its tokens and dollars, + because those are order-free sums, but it is left out of the classification + and of the state: reordering it in would rewrite what "the previous model" + means for turns already folded, and a late arrival is far likelier than a + genuine reversal of a caller's own sequential turns. + + ``rates`` is injected so the fold can be exercised against fixed prices + rather than whatever the cost map happens to say today. + """ + baseline_spend = turn.spend + turn.autorouter_savings + if turn.started_at < state.last_turn_at: + return _unclassified(turn, baseline_spend, state) + + ttl = turn_ttl_seconds(turn) + read_rate, write_rate = rates(turn.model, ttl) + bucket = _bucket(state, turn) + hit = turn.cache_read_tokens > 0 + mark = state.model_marks.get(turn.model) + idle = turn.started_at - mark.last_used_at if mark is not None else 0.0 + + stale = bucket == "return" and not hit and idle > ttl + savable = stale and idle <= 2 * ttl + rescued_spend = turn.cache_creation_tokens * max(write_rate - read_rate, 0.0) if savable else 0.0 + + abandon_spend = (turn.cache_read_tokens + turn.cache_creation_tokens) * read_rate + withdrawn = mark.provisioned_replay_spend if mark is not None and idle <= ttl else 0.0 + + return TurnDelta( + turns=1, + turns_with_usage=1 if turn.has_usage else 0, + total_tokens=turn.total_tokens, + ephemeral_5m_tokens=turn.ephemeral_5m_tokens, + ephemeral_1h_tokens=turn.ephemeral_1h_tokens, + spend=turn.spend, + baseline_spend=baseline_spend, + same_model_turns=1 if bucket == "same_model" else 0, + same_model_hits=1 if bucket == "same_model" and hit else 0, + first_visit_turns=1 if bucket == "first_visit" else 0, + first_visit_hits=1 if bucket == "first_visit" and hit else 0, + return_turns=1 if bucket == "return" else 0, + return_hits=1 if bucket == "return" and hit else 0, + stale_return_misses=1 if stale else 0, + savable_return_misses=1 if savable else 0, + rescued_spend=rescued_spend, + replay_spend=abandon_spend - withdrawn, + state=SessionState( + last_model=turn.model, + last_turn_at=turn.started_at, + model_marks=MappingProxyType( + { # mutable-ok: a JSON object is a dict by definition + **state.model_marks, + turn.model: ModelMark(last_used_at=turn.started_at, provisioned_replay_spend=abandon_spend), + } + ), + ), + ) + + +def _unclassified(turn: TurnFacts, baseline_spend: float, state: SessionState) -> TurnDelta: + return TurnDelta( + state=state, + turns=1, + turns_with_usage=1 if turn.has_usage else 0, + total_tokens=turn.total_tokens, + ephemeral_5m_tokens=turn.ephemeral_5m_tokens, + ephemeral_1h_tokens=turn.ephemeral_1h_tokens, + spend=turn.spend, + baseline_spend=baseline_spend, + ) + + +def auto_router_group_kinds(router: "Router") -> Mapping[str, str]: + """Public alias to router kind, for every auto-router on the proxy. + + ``model_name`` is what a caller sends and what spend rows record, while the + ``litellm_params.model`` string carries the ``auto_router/...`` discriminator + that says it is one. Filtering turns by this mapping is the same filter the + dashboard has always used, and it is load-bearing: the auto-router's own + classifier sub-calls share the session but carry the judge model's group, so + keying on the alias yields one entry per routed turn with no classifier noise. + + Derived per call rather than cached because the router gains and loses + deployments while it runs. + """ + return MappingProxyType( + { # mutable-ok: a JSON object is a dict by definition + str(entry["model_name"]): kind + for entry in (router.model_list or []) # mutable-ok: a JSON object is a dict by definition + if (model := _entry_model(entry)) is not None + and (kind := classify_strategy_router_model(model)) is not None + } + ) + + +def _entry_model(entry: Mapping[str, object]) -> str | None: + params = entry.get("litellm_params") + if not isinstance(params, Mapping): + return None + model = params.get("model") + return model if isinstance(model, str) else None + + +def _ephemeral_split(usage_object: Mapping[str, object]) -> tuple[int, int]: + """``(5m, 1h)`` cache-creation tokens, when the provider breaks them out.""" + creation = usage_object.get("cache_creation") + if not isinstance(creation, Mapping): + return 0, 0 + return ( + int(creation.get("ephemeral_5m_input_tokens") or 0), + int(creation.get("ephemeral_1h_input_tokens") or 0), + ) + + +def turn_from_spend_payload( + model: str, + started_at: datetime, + prompt_tokens: int, + completion_tokens: int, + total_tokens: int, + spend: float, + autorouter_savings: float, + cache_read_tokens: int, + cache_creation_tokens: int, + usage_object: Mapping[str, object], +) -> TurnFacts: + """One spend log payload as the fold sees it. + + The cache token counts arrive already extracted because the spend writer owns + those readers and they have to agree with what the daily rows recorded for the + same request. + """ + ephemeral_5m, ephemeral_1h = _ephemeral_split(usage_object) + return TurnFacts( + model=model, + started_at=as_epoch(started_at), + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens, + total_tokens=total_tokens, + cache_read_tokens=cache_read_tokens, + cache_creation_tokens=cache_creation_tokens, + ephemeral_5m_tokens=ephemeral_5m, + ephemeral_1h_tokens=ephemeral_1h, + spend=spend, + autorouter_savings=autorouter_savings, + has_usage=_reports_cache_usage(usage_object), + ) + + +def _reports_cache_usage(usage_object: Mapping[str, object]) -> bool: + """Whether this turn's usage payload says anything about the prompt cache. + + Coverage answers "can we see cache behaviour here at all", so it counts the + presence of a cache field rather than a non-zero one; a turn that genuinely + read nothing is a miss, not a gap in reporting. A low figure means response + logging is off, which is why it is surfaced next to the hit rate. + """ + return "cache_read_input_tokens" in usage_object or bool(usage_object.get("prompt_tokens_details")) + + +def as_epoch(value: datetime) -> float: + """Seconds since the epoch, treating a naive timestamp as UTC. + + Spend rows are written in UTC but reach here either naive or aware depending + on the driver, and mixing the two silently shifts every idle-time comparison + by the local offset. + """ + return (value if value.tzinfo is not None else value.replace(tzinfo=timezone.utc)).timestamp() + + +class _StoredMark(BaseModel): + last_used_at: float + provisioned_replay_spend: float + + +_STORED_MARKS = TypeAdapter(dict[str, _StoredMark]) + + +def state_to_json(state: SessionState) -> Mapping[str, Mapping[str, float]]: + """``model_marks`` as the ``model_state`` column stores it.""" + return { # mutable-ok: a JSON object is a dict by definition + model: { # mutable-ok: a JSON object is a dict by definition + "last_used_at": mark.last_used_at, + "provisioned_replay_spend": mark.provisioned_replay_spend, + } + for model, mark in state.model_marks.items() + } + + +def state_column(state: SessionState) -> object: + """``model_state`` wrapped the way prisma requires for a Json column. + + Model names contain a slash, and prisma-client-py inlines Json into a GraphQL + document where an unquoted key containing one is a parse error, so a plain + dict fails the whole write. Every writer goes through here so that cannot be + rediscovered one call site at a time. + """ + import prisma + + return prisma.Json(state_to_json(state)) + + +def state_from_row(last_model: str | None, last_turn_at: datetime | None, model_state: object) -> SessionState: + """Rebuild a session's state from its row. + + A row whose ``model_state`` cannot be parsed is treated as a session with no + history rather than raising: the counters it already carries stay correct and + the next turn simply reads as a first visit, which beats failing the spend + write over a state blob. + """ + try: + marks = _STORED_MARKS.validate_python(model_state or {}) # mutable-ok: empty fallback for an absent mapping + except ValidationError as e: + verbose_proxy_logger.warning("auto_router_sessions: unreadable model_state, session history reset (%s)", e) + return EMPTY_SESSION_STATE + return SessionState( + last_model=last_model, + last_turn_at=as_epoch(last_turn_at) if last_turn_at is not None else 0.0, + model_marks=MappingProxyType( + { # mutable-ok: a JSON object is a dict by definition + model: ModelMark( + last_used_at=stored.last_used_at, provisioned_replay_spend=stored.provisioned_replay_spend + ) + for model, stored in marks.items() + } + ), + ) + + +def merge_deltas(left: TurnDelta, right: TurnDelta) -> TurnDelta: + """Combine two folds of the same session so a flush writes one row once. + + Counters add and the later state wins, which is what the database would have + done had the two turns flushed separately; folding them in memory first just + spares the round trip. + """ + return TurnDelta( + state=right.state, + **{ # mutable-ok: a JSON object is a dict by definition + name: getattr(left, name) + getattr(right, name) for name in COUNTER_FIELDS + }, # mutable-ok: a JSON object is a dict by definition + ) diff --git a/litellm/repositories/__init__.py b/litellm/repositories/__init__.py index 1fc3d8dadaf..8b01942425b 100644 --- a/litellm/repositories/__init__.py +++ b/litellm/repositories/__init__.py @@ -17,6 +17,7 @@ from litellm.repositories.table_repositories import ( AdaptiveRouterStateRepository, AgentsRepository, AuditLogRepository, + AutoRouterSessionRepository, CacheConfigRepository, ClaudeCodePluginRepository, ConfigOverridesRepository, @@ -73,6 +74,7 @@ __all__ = [ "AdaptiveRouterStateRepository", "AgentsRepository", "AuditLogRepository", + "AutoRouterSessionRepository", "BudgetRepository", "CacheConfigRepository", "ClaudeCodePluginRepository", diff --git a/litellm/repositories/table_repositories.py b/litellm/repositories/table_repositories.py index af8be986831..67a5348fb8e 100644 --- a/litellm/repositories/table_repositories.py +++ b/litellm/repositories/table_repositories.py @@ -224,3 +224,7 @@ class AuditLogRepository(PrismaTableRepository): class AdaptiveRouterSessionRepository(PrismaTableRepository): table_name = "litellm_adaptiveroutersession" + + +class AutoRouterSessionRepository(PrismaTableRepository): + table_name = "litellm_autoroutersession" diff --git a/schema.prisma b/schema.prisma index 0d7fa8692c8..502541f36cd 100644 --- a/schema.prisma +++ b/schema.prisma @@ -1391,6 +1391,46 @@ model LiteLLM_AdaptiveRouterSession { @@index([last_activity_at], map: "idx_adaptive_router_session_activity") } +// Per-(session, auto-router) rollup behind the auto-router benchmarks dashboard. +model LiteLLM_AutoRouterSession { + session_id String + model_group String + router_kind String + baseline_model String? + + first_turn_at DateTime + last_turn_at DateTime + turns Int @default(0) + turns_with_usage Int @default(0) + + total_tokens BigInt @default(0) + ephemeral_5m_tokens BigInt @default(0) + ephemeral_1h_tokens BigInt @default(0) + + spend Float @default(0.0) + baseline_spend Float @default(0.0) + + same_model_turns Int @default(0) + same_model_hits Int @default(0) + first_visit_turns Int @default(0) + first_visit_hits Int @default(0) + return_turns Int @default(0) + return_hits Int @default(0) + stale_return_misses Int @default(0) + savable_return_misses Int @default(0) + + rescued_spend Float @default(0.0) + replay_spend Float @default(0.0) + + last_model String? + model_state Json @default("{}") + updated_at DateTime @default(now()) @updatedAt + + @@id([session_id, model_group]) + @@index([model_group, last_turn_at], map: "idx_auto_router_session_group_activity") + @@index([last_turn_at], map: "idx_auto_router_session_last_turn") +} + // --------------------------------------------------------------------------- // Workflow Run Tracking // diff --git a/tests/test_litellm/proxy/spend_tracking/test_auto_router_benchmarks.py b/tests/test_litellm/proxy/spend_tracking/test_auto_router_benchmarks.py new file mode 100644 index 00000000000..fe2dae72c61 --- /dev/null +++ b/tests/test_litellm/proxy/spend_tracking/test_auto_router_benchmarks.py @@ -0,0 +1,250 @@ +from datetime import datetime, timedelta, timezone + +import pytest + +from litellm.proxy.spend_tracking.auto_router_benchmarks import ( + BENCHMARKS_MAX_WINDOW_DAYS, + clamp_window, + compute_benchmarks, +) + +GROUP_KINDS = {"claude-auto": "semantic"} + + +class _FakeTable: + """Records every upsert so a test can assert on what would be written.""" + + def __init__(self): + self.upserts = [] + + async def upsert(self, where, data): + self.upserts.append((where, data)) + + +class _FakeDb: + """Returns rows the way prisma really does: a list of plain dicts.""" + + def __init__(self, rows, table=None): + self._rows = rows + self.queries = [] + self.litellm_autoroutersession = table or _FakeTable() + + async def query_raw(self, sql, *args): + self.queries.append((sql, args)) + return self._rows + + +class _FakePrisma: + def __init__(self, rows, table=None): + self.db = _FakeDb(rows, table) + + +def group_row(**overrides): + row = { + "model_group": "claude-auto", + "baseline_model": "anthropic/claude-opus-4-8", + "sessions": 10, + "turns": 100, + "total_session_seconds": 36000.0, + "total_tokens": 1_000_000, + "actual_spend": 10.0, + "baseline_spend": 100.0, + "turns_with_usage": 100, + "ephemeral_5m_tokens": 0, + "ephemeral_1h_tokens": 5000, + "same_model_turns": 60, + "same_model_hits": 57, + "first_visit_turns": 10, + "first_visit_hits": 2, + "return_turns": 30, + "return_hits": 24, + "stale_return_misses": 4, + "savable_return_misses": 2, + "rescued_spend": 6.76, + "replay_spend": 3.91, + } + row.update(overrides) + return row + + +async def benchmarks_for(**overrides): + prisma = _FakePrisma([group_row(**overrides)]) + return await compute_benchmarks(prisma, GROUP_KINDS, "2026-07-02", "2026-08-01") + + +class TestWindowClamping: + def test_a_wider_request_is_clamped_to_the_maximum_window(self): + window = clamp_window("2020-01-01", "2026-08-01") + expected = (datetime(2026, 8, 1, tzinfo=timezone.utc) - timedelta(days=BENCHMARKS_MAX_WINDOW_DAYS)).date() + assert window.start == expected.isoformat() + + def test_a_narrower_request_is_served_as_asked(self): + assert clamp_window("2026-07-25", "2026-08-01").start == "2026-07-25" + + def test_the_response_echoes_the_window_actually_served(self): + window = clamp_window("2020-01-01", "2026-08-01") + assert window.end == "2026-08-01" + + +@pytest.mark.asyncio +class TestSessionShape: + async def test_turns_per_session_divides_turns_by_sessions(self): + result = await benchmarks_for() + assert result.groups[0].avg_turns_per_session == pytest.approx(10.0) + + async def test_session_length_averages_the_summed_durations(self): + result = await benchmarks_for() + assert result.groups[0].avg_session_length_seconds == pytest.approx(3600.0) + + async def test_tokens_per_session_divides_tokens_by_sessions(self): + result = await benchmarks_for() + assert result.groups[0].avg_tokens_per_session == pytest.approx(100_000.0) + + async def test_a_group_with_no_sessions_is_omitted_rather_than_zeroed(self): + result = await benchmarks_for(sessions=0) + assert result.groups == () + + +@pytest.mark.asyncio +class TestSavings: + async def test_savings_is_baseline_minus_actual(self): + result = await benchmarks_for() + assert result.groups[0].savings == pytest.approx(90.0) + assert result.groups[0].savings_pct == pytest.approx(90.0) + + async def test_a_route_that_cost_more_than_the_baseline_reports_a_loss(self): + """Signed on purpose: a cache-thrashing router must not read as zero.""" + result = await benchmarks_for(actual_spend=120.0, baseline_spend=100.0) + assert result.groups[0].savings == pytest.approx(-20.0) + assert result.groups[0].savings_pct == pytest.approx(-20.0) + + async def test_an_unpriced_baseline_reports_no_percentage_instead_of_dividing_by_zero(self): + result = await benchmarks_for(baseline_spend=0.0) + assert result.groups[0].savings_pct == 0.0 + + +@pytest.mark.asyncio +class TestCacheBuckets: + async def test_the_three_buckets_sum_to_the_reported_turn_count(self): + cache = (await benchmarks_for()).groups[0].cache + assert cache is not None + assert cache.same_model_turns + cache.first_visit_turns + cache.return_turns == cache.turns + + async def test_the_headline_rate_is_weighted_by_turns_not_an_average_of_buckets(self): + """57+2+24 hits over 100 turns is 83%, not the 61% mean of the three rates.""" + cache = (await benchmarks_for()).groups[0].cache + assert cache is not None + assert cache.hit_rate_pct == pytest.approx(83.0) + + async def test_each_bucket_reports_its_own_hit_rate(self): + cache = (await benchmarks_for()).groups[0].cache + assert cache is not None + assert cache.same_model_hit_rate_pct == pytest.approx(95.0) + assert cache.first_visit_hit_rate_pct == pytest.approx(20.0) + assert cache.return_hit_rate_pct == pytest.approx(80.0) + + async def test_stale_share_is_measured_against_return_misses_only(self): + cache = (await benchmarks_for()).groups[0].cache + assert cache is not None + assert cache.stale_miss_share_pct == pytest.approx(100.0 * 4 / 6) + + async def test_savable_share_is_measured_against_every_miss(self): + cache = (await benchmarks_for()).groups[0].cache + assert cache is not None + assert cache.warming_savable_miss_pct == pytest.approx(100.0 * 2 / 17) + + async def test_cache_is_omitted_when_nothing_reported_usage(self): + result = await benchmarks_for(turns_with_usage=0) + assert result.groups[0].cache is None + + async def test_coverage_is_the_share_of_turns_that_reported_usage(self): + cache = (await benchmarks_for(turns_with_usage=50)).groups[0].cache + assert cache is not None + assert cache.usage_coverage_pct == pytest.approx(50.0) + + +@pytest.mark.asyncio +class TestWarmingEstimate: + async def test_net_is_rescued_less_replays(self): + cache = (await benchmarks_for()).groups[0].cache + assert cache is not None + assert cache.warming_net_spend == pytest.approx(6.76 - 3.91) + + async def test_break_even_follows_the_ttl_in_use(self): + one_hour = (await benchmarks_for()).groups[0].cache + five_min = (await benchmarks_for(ephemeral_1h_tokens=0, ephemeral_5m_tokens=5000)).groups[0].cache + assert one_hour is not None and five_min is not None + assert one_hour.ttl_seconds == 3600 + assert one_hour.warming_break_even_pct == 5.0 + assert five_min.ttl_seconds == 300 + assert five_min.warming_break_even_pct == 9.0 + + async def test_no_ephemeral_evidence_reads_as_the_five_minute_tier(self): + cache = (await benchmarks_for(ephemeral_1h_tokens=0, ephemeral_5m_tokens=0)).groups[0].cache + assert cache is not None + assert cache.ttl_seconds == 300 + + +@pytest.mark.asyncio +class TestReadPathSource: + async def test_the_dashboard_query_never_touches_the_spend_logs(self): + prisma = _FakePrisma([group_row()]) + await compute_benchmarks(prisma, GROUP_KINDS, "2026-07-02", "2026-08-01") + sql = prisma.db.queries[0][0] + assert "LiteLLM_SpendLogs" not in sql + assert "LiteLLM_AutoRouterSession" in sql + + async def test_one_query_covers_every_configured_auto_router(self): + prisma = _FakePrisma([group_row(), group_row(model_group="claude-router-2")]) + result = await compute_benchmarks( + prisma, {"claude-auto": "semantic", "claude-router-2": "complexity"}, "2026-07-02", "2026-08-01" + ) + assert len(prisma.db.queries) == 1 + assert {g.model_group for g in result.groups} == {"claude-auto", "claude-router-2"} + + async def test_each_group_is_labelled_with_its_router_kind(self): + prisma = _FakePrisma([group_row(model_group="claude-router-2")]) + result = await compute_benchmarks(prisma, {"claude-router-2": "complexity"}, "2026-07-02", "2026-08-01") + assert result.groups[0].router_kind == "complexity" + + +class TestEveryCounterSurvivesToTheDashboard: + """A counter declared on TurnDelta must reach the row and come back out. + + The failure this guards is silent: a metric gets written on every request, + the read query never selects it, and the card shows zero forever with a green + diff and passing tests. Both ends are checked against the one declaration. + """ + + def test_the_read_query_aggregates_every_counter(self): + from litellm.proxy.spend_tracking.auto_router_benchmarks import _GROUP_SQL + from litellm.proxy.spend_tracking.auto_router_sessions import COUNTER_FIELDS + + missing = [name for name in COUNTER_FIELDS if f"SUM({name})" not in _GROUP_SQL] + assert missing == [], f"counters written but never read: {missing}" + + def test_the_response_row_carries_every_counter(self): + from litellm.proxy.spend_tracking.auto_router_benchmarks import _GroupRow + from litellm.proxy.spend_tracking.auto_router_sessions import COUNTER_FIELDS + + # `spend` is the one deliberate rename; the row calls it actual_spend + expected = {"actual_spend" if name == "spend" else name for name in COUNTER_FIELDS} + assert expected <= set(_GroupRow.model_fields) + + def test_the_flush_payload_carries_every_counter(self): + from litellm.proxy.spend_tracking.auto_router_sessions import ( + COUNTER_FIELDS, + EMPTY_SESSION_STATE, + TurnDelta, + counters_of, + ) + + assert set(counters_of(TurnDelta(state=EMPTY_SESSION_STATE))) == set(COUNTER_FIELDS) + + def test_counter_fields_is_derived_not_hand_listed(self): + """Adding a field to TurnDelta must extend COUNTER_FIELDS with no other edit.""" + from dataclasses import fields + + from litellm.proxy.spend_tracking.auto_router_sessions import COUNTER_FIELDS, TurnDelta + + assert set(COUNTER_FIELDS) == {f.name for f in fields(TurnDelta)} - {"state"} diff --git a/tests/test_litellm/proxy/spend_tracking/test_auto_router_sessions.py b/tests/test_litellm/proxy/spend_tracking/test_auto_router_sessions.py new file mode 100644 index 00000000000..c48f741846d --- /dev/null +++ b/tests/test_litellm/proxy/spend_tracking/test_auto_router_sessions.py @@ -0,0 +1,434 @@ +from dataclasses import replace +from datetime import datetime, timezone + +import pytest + +from litellm.proxy.spend_tracking.auto_router_sessions import ( + EMPTY_SESSION_STATE, + PROMPT_CACHE_TTL_SECONDS, + SessionState, + TurnFacts, + fold_turn, + merge_deltas, + state_from_row, + state_to_json, + turn_ttl_seconds, +) + +MODEL_A = "anthropic/claude-haiku-4-5" +MODEL_B = "anthropic/claude-sonnet-4-5" + +READ_RATE = 1e-6 +WRITE_RATE = 5e-6 +FIVE_MIN = PROMPT_CACHE_TTL_SECONDS["5m"] +ONE_HOUR = PROMPT_CACHE_TTL_SECONDS["1h"] + + +def rates(model: str, ttl_seconds: int) -> tuple[float, float]: + """Fixed prices so dollar assertions do not move with the cost map.""" + return READ_RATE, WRITE_RATE + + +def turn( + model: str = MODEL_A, + at: float = 0.0, + read: int = 0, + created: int = 0, + spend: float = 0.01, + savings: float = 0.0, + ephemeral_5m: int | None = None, + ephemeral_1h: int = 0, +) -> TurnFacts: + return TurnFacts( + model=model, + started_at=at, + prompt_tokens=1000, + completion_tokens=100, + total_tokens=1100, + cache_read_tokens=read, + cache_creation_tokens=created, + ephemeral_5m_tokens=created if ephemeral_5m is None else ephemeral_5m, + ephemeral_1h_tokens=ephemeral_1h, + spend=spend, + autorouter_savings=savings, + has_usage=True, + ) + + +def fold_all(turns: tuple[TurnFacts, ...]): + """Fold a whole session, returning every delta in order plus the final state.""" + state = EMPTY_SESSION_STATE + deltas = [] + for one in turns: + delta = fold_turn(state, one, rates=rates) + deltas.append(delta) + state = delta.state + return tuple(deltas), state + + +def buckets(delta) -> int: + return delta.same_model_turns + delta.first_visit_turns + delta.return_turns + + +class TestBucketExhaustiveness: + def test_every_turn_lands_in_exactly_one_bucket(self): + deltas, _ = fold_all( + ( + turn(MODEL_A, at=0, created=5000), + turn(MODEL_A, at=60, read=5000), + turn(MODEL_B, at=120, created=5000), + turn(MODEL_A, at=180, read=5000), + turn(MODEL_B, at=240, read=5000), + ) + ) + assert [buckets(d) for d in deltas] == [1, 1, 1, 1, 1] + assert sum(buckets(d) for d in deltas) == sum(d.turns for d in deltas) + + def test_the_opening_turn_of_a_session_is_a_first_visit(self): + """The old split left it in no bucket, so bucket totals undercounted.""" + delta = fold_turn(EMPTY_SESSION_STATE, turn(MODEL_A, at=0, created=5000), rates=rates) + assert delta.first_visit_turns == 1 + assert buckets(delta) == delta.turns == 1 + + def test_arriving_at_an_unused_model_is_a_first_visit_not_a_return(self): + deltas, _ = fold_all((turn(MODEL_A, at=0), turn(MODEL_B, at=60))) + assert deltas[1].first_visit_turns == 1 + assert deltas[1].return_turns == 0 + + def test_staying_on_the_same_model_is_not_a_return(self): + deltas, _ = fold_all((turn(MODEL_A, at=0), turn(MODEL_A, at=60))) + assert deltas[1].same_model_turns == 1 + assert deltas[1].return_turns == 0 + + +class TestHits: + def test_a_turn_that_read_from_cache_is_a_hit_in_its_own_bucket(self): + deltas, _ = fold_all((turn(MODEL_A, at=0, created=5000), turn(MODEL_A, at=60, read=5000))) + assert deltas[1].same_model_hits == 1 + assert deltas[0].first_visit_hits == 0 + + def test_a_hit_is_read_tokens_not_a_provider_flag(self): + delta = fold_turn(EMPTY_SESSION_STATE, turn(MODEL_A, at=0, read=1, created=0), rates=rates) + assert delta.first_visit_hits == 1 + + +class TestStaleAndSavableReturns: + def _return_after(self, idle: float): + state = fold_turn(EMPTY_SESSION_STATE, turn(MODEL_A, at=0, created=5000), rates=rates).state + state = fold_turn(state, turn(MODEL_B, at=1, created=5000), rates=rates).state + return fold_turn(state, turn(MODEL_A, at=idle, created=5000), rates=rates) + + def test_return_inside_the_ttl_is_neither_stale_nor_savable(self): + delta = self._return_after(FIVE_MIN - 10) + assert delta.return_turns == 1 + assert delta.stale_return_misses == 0 + assert delta.savable_return_misses == 0 + assert delta.rescued_spend == 0.0 + + def test_return_past_the_ttl_within_two_ttls_is_savable_and_rescues_the_write(self): + delta = self._return_after(FIVE_MIN + 10) + assert delta.stale_return_misses == 1 + assert delta.savable_return_misses == 1 + assert delta.rescued_spend == pytest.approx(5000 * (WRITE_RATE - READ_RATE)) + + def test_return_past_two_ttls_is_stale_but_not_savable(self): + delta = self._return_after(2 * FIVE_MIN + 10) + assert delta.stale_return_misses == 1 + assert delta.savable_return_misses == 0 + assert delta.rescued_spend == 0.0 + + def test_a_return_that_hit_is_never_counted_as_a_miss(self): + state = fold_turn(EMPTY_SESSION_STATE, turn(MODEL_A, at=0, created=5000), rates=rates).state + state = fold_turn(state, turn(MODEL_B, at=1, created=5000), rates=rates).state + delta = fold_turn(state, turn(MODEL_A, at=FIVE_MIN + 10, read=5000), rates=rates) + assert delta.return_hits == 1 + assert delta.stale_return_misses == 0 + assert delta.savable_return_misses == 0 + + +class TestWarmingReplayEconomics: + def test_replay_is_withdrawn_when_the_session_returns_inside_the_ttl(self): + """No refresher would have fired, so the provisional charge comes back off.""" + deltas, _ = fold_all((turn(MODEL_A, at=0, created=5000), turn(MODEL_A, at=60, read=5000))) + assert deltas[0].replay_spend == pytest.approx(5000 * READ_RATE) + assert deltas[1].replay_spend == pytest.approx(0.0) + + def test_replay_is_kept_when_the_session_stays_away_past_the_ttl(self): + deltas, _ = fold_all( + (turn(MODEL_A, at=0, created=5000), turn(MODEL_A, at=FIVE_MIN + 10, created=5000)) + ) + assert sum(d.replay_spend for d in deltas) == pytest.approx(2 * 5000 * READ_RATE) + + def test_total_replay_is_each_bridged_gap_plus_one_final_abandon(self): + """Reproduces what the window-function query summed, without the window.""" + deltas, _ = fold_all( + ( + turn(MODEL_A, at=0, created=5000), + turn(MODEL_A, at=60, read=5000), + turn(MODEL_A, at=1400, created=5000), + ) + ) + bridged_gap = 5000 * READ_RATE + final_abandon = 5000 * READ_RATE + assert sum(d.replay_spend for d in deltas) == pytest.approx(bridged_gap + final_abandon) + + def test_each_model_a_session_touches_carries_its_own_abandon_charge(self): + deltas, _ = fold_all((turn(MODEL_A, at=0, created=5000), turn(MODEL_B, at=60, created=3000))) + assert sum(d.replay_spend for d in deltas) == pytest.approx((5000 + 3000) * READ_RATE) + + def test_prefix_for_replay_counts_read_and_written_tokens(self): + delta = fold_turn(EMPTY_SESSION_STATE, turn(MODEL_A, at=0, read=2000, created=3000), rates=rates) + assert delta.replay_spend == pytest.approx(5000 * READ_RATE) + + +class TestTtlSelection: + def test_defaults_to_five_minutes_without_ephemeral_evidence(self): + """Both counters zero must not read as the one hour tier.""" + assert turn_ttl_seconds(turn(MODEL_A, created=0, ephemeral_5m=0, ephemeral_1h=0)) == FIVE_MIN + + def test_one_hour_when_the_turn_wrote_mostly_to_the_long_cache(self): + assert turn_ttl_seconds(turn(MODEL_A, ephemeral_5m=100, ephemeral_1h=5000)) == ONE_HOUR + + def test_five_minutes_when_the_turn_wrote_mostly_to_the_short_cache(self): + assert turn_ttl_seconds(turn(MODEL_A, ephemeral_5m=5000, ephemeral_1h=100)) == FIVE_MIN + + def test_staleness_follows_the_turns_own_ttl(self): + state = fold_turn(EMPTY_SESSION_STATE, turn(MODEL_A, at=0, created=5000), rates=rates).state + state = fold_turn(state, turn(MODEL_B, at=1, created=5000), rates=rates).state + inside_the_hour = fold_turn( + state, turn(MODEL_A, at=1800, created=5000, ephemeral_5m=0, ephemeral_1h=5000), rates=rates + ) + assert inside_the_hour.stale_return_misses == 0 + + +class TestOutOfOrderTurns: + def test_a_late_turn_keeps_its_spend_but_not_its_classification(self): + first = fold_turn(EMPTY_SESSION_STATE, turn(MODEL_A, at=100, created=5000, spend=0.02), rates=rates) + late = fold_turn(first.state, turn(MODEL_B, at=50, created=5000, spend=0.03), rates=rates) + assert late.turns == 1 + assert late.spend == 0.03 + assert buckets(late) == 0 + assert late.replay_spend == 0.0 + + def test_a_late_turn_does_not_rewrite_the_sessions_state(self): + first = fold_turn(EMPTY_SESSION_STATE, turn(MODEL_A, at=100), rates=rates) + late = fold_turn(first.state, turn(MODEL_B, at=50), rates=rates) + assert late.state is first.state + assert late.state.last_model == MODEL_A + + +class TestBaselineSpend: + def test_baseline_is_what_was_paid_plus_what_routing_saved(self): + delta = fold_turn(EMPTY_SESSION_STATE, turn(MODEL_A, spend=0.01, savings=0.09), rates=rates) + assert delta.baseline_spend == pytest.approx(0.10) + + def test_a_route_that_lost_money_reports_a_baseline_below_actual_spend(self): + """Savings are signed, so a cache-thrashing route must stay visible as a loss.""" + delta = fold_turn(EMPTY_SESSION_STATE, turn(MODEL_A, spend=0.05, savings=-0.02), rates=rates) + assert delta.baseline_spend == pytest.approx(0.03) + assert delta.baseline_spend < delta.spend + + +class TestCoverage: + def test_a_turn_without_cache_reporting_still_counts_as_a_turn(self): + """Coverage separates "logging is off" from "the cache was cold".""" + facts = replace(turn(MODEL_A, at=0), has_usage=False) + delta = fold_turn(EMPTY_SESSION_STATE, facts, rates=rates) + assert delta.turns == 1 + assert delta.turns_with_usage == 0 + + def test_a_turn_that_reported_usage_counts_toward_coverage(self): + delta = fold_turn(EMPTY_SESSION_STATE, turn(MODEL_A, at=0, created=5000), rates=rates) + assert delta.turns_with_usage == 1 + + +class TestStateRoundTrip: + def test_state_survives_a_trip_through_the_row(self): + _, state = fold_all((turn(MODEL_A, at=0, created=5000), turn(MODEL_B, at=60, created=3000))) + restored = state_from_row( + state.last_model, + datetime.fromtimestamp(state.last_turn_at, tz=timezone.utc), + state_to_json(state), + ) + assert restored.last_model == state.last_model + assert restored.last_turn_at == pytest.approx(state.last_turn_at) + assert set(restored.model_marks) == set(state.model_marks) + for model, mark in state.model_marks.items(): + assert restored.model_marks[model].last_used_at == pytest.approx(mark.last_used_at) + assert restored.model_marks[model].provisioned_replay_spend == pytest.approx( + mark.provisioned_replay_spend + ) + + def test_a_restored_session_classifies_the_next_turn_the_same_way(self): + """This is the property that makes a pod hop or a restart harmless.""" + _, state = fold_all((turn(MODEL_A, at=0, created=5000), turn(MODEL_B, at=60, created=5000))) + restored = state_from_row( + state.last_model, datetime.fromtimestamp(state.last_turn_at, tz=timezone.utc), state_to_json(state) + ) + next_turn = turn(MODEL_A, at=FIVE_MIN + 100, created=5000) + assert fold_turn(restored, next_turn, rates=rates) == fold_turn(state, next_turn, rates=rates) + + def test_an_unreadable_state_blob_resets_history_instead_of_raising(self): + restored = state_from_row("m", datetime.now(timezone.utc), {"bad": "shape"}) + assert restored is EMPTY_SESSION_STATE + + def test_a_naive_timestamp_is_read_as_utc(self): + naive = datetime(2026, 8, 1, 12, 0, 0) + aware = datetime(2026, 8, 1, 12, 0, 0, tzinfo=timezone.utc) + assert state_from_row(None, naive, {}).last_turn_at == state_from_row(None, aware, {}).last_turn_at + + +class TestMergeDeltas: + def test_counters_add_and_the_later_state_wins(self): + first = fold_turn(EMPTY_SESSION_STATE, turn(MODEL_A, at=0, created=5000), rates=rates) + second = fold_turn(first.state, turn(MODEL_A, at=60, read=5000), rates=rates) + merged = merge_deltas(first, second) + assert merged.turns == 2 + assert merged.first_visit_turns == 1 + assert merged.same_model_turns == 1 + assert merged.spend == pytest.approx(first.spend + second.spend) + assert merged.replay_spend == pytest.approx(first.replay_spend + second.replay_spend) + assert merged.state == second.state + + def test_merging_a_whole_session_matches_folding_it_turn_by_turn(self): + deltas, _ = fold_all( + ( + turn(MODEL_A, at=0, created=5000), + turn(MODEL_A, at=60, read=5000), + turn(MODEL_B, at=400, created=5000), + turn(MODEL_A, at=800, created=5000), + ) + ) + merged = deltas[0] + for delta in deltas[1:]: + merged = merge_deltas(merged, delta) + assert merged.turns == 4 + assert merged.same_model_turns + merged.first_visit_turns + merged.return_turns == 4 + assert merged.replay_spend == pytest.approx(sum(d.replay_spend for d in deltas)) + + +class TestEmptySessionState: + def test_a_fresh_session_has_no_model_history(self): + assert EMPTY_SESSION_STATE.last_model is None + assert EMPTY_SESSION_STATE.model_marks == {} + + def test_folding_never_mutates_the_state_it_was_given(self): + before = state_to_json(EMPTY_SESSION_STATE) + fold_turn(EMPTY_SESSION_STATE, turn(MODEL_A, at=0, created=5000), rates=rates) + assert state_to_json(EMPTY_SESSION_STATE) == before + + def test_session_state_is_hashable_free_of_shared_mutation(self): + _, state = fold_all((turn(MODEL_A, at=0, created=5000),)) + snapshot = dict(state.model_marks) + fold_turn(state, turn(MODEL_B, at=60, created=5000), rates=rates) + assert dict(state.model_marks) == snapshot + + +def test_session_state_type_is_frozen(): + _, state = fold_all((turn(MODEL_A, at=0),)) + with pytest.raises(Exception): + state.last_model = "other" # pyright: ignore[reportAttributeAccessIssue] # asserting frozen + + +def test_state_from_row_without_a_timestamp_starts_at_the_epoch(): + assert state_from_row(None, None, {}) == SessionState( + last_model=None, last_turn_at=0.0, model_marks={} + ) + + +class _RecordingTable: + """A session-rollup table that can be told to fail.""" + + def __init__(self, fail: bool = False): + self.fail = fail + self.upserts: list = [] + + async def find_unique(self, where): + return None + + async def upsert(self, where, data): + if self.fail: + raise RuntimeError("transient database fault") + self.upserts.append((where, data)) + + +class _RecordingPrisma: + def __init__(self, table): + self.db = type("_Db", (), {"litellm_autoroutersession": table})() + + +def _turn_at(at: float, model: str = MODEL_A) -> TurnFacts: + return turn(model, at=at, created=5000) + + +@pytest.mark.asyncio +class TestFlushDurability: + """A transient write fault must not silently delete an interval of traffic.""" + + async def test_a_failed_write_is_restaged_rather_than_dropped(self): + from litellm.proxy.spend_tracking.auto_router_session_queue import AutoRouterSessionQueue + + table = _RecordingTable(fail=True) + prisma = _RecordingPrisma(table) + queue = AutoRouterSessionQueue() + await queue.record_turn(("s1", "g"), "complexity", None, _turn_at(0), prisma) + + assert await queue.flush(prisma) == 0 + + table.fail = False + assert await queue.flush(prisma) == 1 + assert table.upserts[0][1]["create"]["turns"] == 1 + + async def test_a_restaged_batch_merges_under_turns_staged_since(self): + from litellm.proxy.spend_tracking.auto_router_session_queue import AutoRouterSessionQueue + + table = _RecordingTable(fail=True) + prisma = _RecordingPrisma(table) + queue = AutoRouterSessionQueue() + await queue.record_turn(("s1", "g"), "complexity", None, _turn_at(0), prisma) + await queue.flush(prisma) + + table.fail = False + await queue.record_turn(("s1", "g"), "complexity", None, _turn_at(60), prisma) + assert await queue.flush(prisma) == 1 + # Both turns land, once each + assert table.upserts[0][1]["create"]["turns"] == 2 + + async def test_a_successful_flush_stages_nothing_back(self): + from litellm.proxy.spend_tracking.auto_router_session_queue import AutoRouterSessionQueue + + prisma = _RecordingPrisma(_RecordingTable()) + queue = AutoRouterSessionQueue() + await queue.record_turn(("s1", "g"), "complexity", None, _turn_at(0), prisma) + assert await queue.flush(prisma) == 1 + assert await queue.flush(prisma) == 0 + + +@pytest.mark.asyncio +class TestPendingIsBounded: + """`session_id` is caller-controlled, so the staged aggregate needs a ceiling.""" + + async def test_new_sessions_are_refused_once_the_aggregate_is_full(self): + from litellm.proxy.spend_tracking.auto_router_session_queue import AutoRouterSessionQueue + + prisma = _RecordingPrisma(_RecordingTable()) + queue = AutoRouterSessionQueue(max_tracked_sessions=2) + for i in range(5): + await queue.record_turn((f"s{i}", "g"), "complexity", None, _turn_at(0), prisma) + + assert await queue.flush(prisma) == 2 + + async def test_a_session_already_staged_keeps_accumulating_at_the_cap(self): + """Refusing new keys must not stall the conversations already in flight.""" + from litellm.proxy.spend_tracking.auto_router_session_queue import AutoRouterSessionQueue + + table = _RecordingTable() + prisma = _RecordingPrisma(table) + queue = AutoRouterSessionQueue(max_tracked_sessions=1) + await queue.record_turn(("s1", "g"), "complexity", None, _turn_at(0), prisma) + await queue.record_turn(("s2", "g"), "complexity", None, _turn_at(0), prisma) + await queue.record_turn(("s1", "g"), "complexity", None, _turn_at(60), prisma) + + await queue.flush(prisma) + assert len(table.upserts) == 1 + assert table.upserts[0][1]["create"]["turns"] == 2 diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 1ea4113b821..84eceab0630 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -760,6 +760,36 @@ export interface paths { patch?: never; trace?: never; }; + "/auto_router/benchmarks": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** + * Get Auto Router Benchmarks + * @description Session-level benchmarks for every configured auto-router. + * + * Admin-only. For each auto-router alias, reports turns per session, session + * length, tokens per session, the dollar savings of the routed mix against the + * counterfactual baseline, and how the provider prompt cache behaved. + * + * Reads the per-session rollup, never the per-request spend logs. + * ``start_date`` / ``end_date`` are ``YYYY-MM-DD``; the window is clamped to + * the most recent ``BENCHMARKS_MAX_WINDOW_DAYS`` days and the response echoes + * the window actually served. Sessions are counted whole when they were active + * in the window. Returns 404 when no auto-router is configured. + */ + get: operations["get_auto_router_benchmarks_auto_router_benchmarks_get"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/azure/{endpoint}": { parameters: { query?: never; @@ -35515,6 +35545,38 @@ export interface operations { }; }; }; + get_auto_router_benchmarks_auto_router_benchmarks_get: { + parameters: { + query: { + start_date: string; + end_date: string; + }; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; azure_proxy_route_azure__endpoint__get: { parameters: { query?: never;