From daced81f20a6b98c01beecf66fa997310d90bfb4 Mon Sep 17 00:00:00 2001 From: amasen02 Date: Fri, 4 Sep 2026 15:37:53 +0530 Subject: [PATCH 01/11] fix(proxy): invalidate end-user spend counter and cache on budget reset (#39726) Signed-off-by: amasen02 --- .../proxy/common_utils/reset_budget_job.py | 24 ++++++++++++++++- .../common_utils/test_reset_budget_job.py | 26 +++++++++++++++++++ 2 files changed, 49 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index 47f69732e95..12ba75aea24 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -38,6 +38,7 @@ from litellm.proxy.common_utils.timezone_utils import ( get_budget_reset_settings, ) from litellm.proxy.common_utils.user_api_key_cache import ( + end_user_cache_key, model_access_group_cache_key, model_access_group_spend_counter_key, tag_cache_key, @@ -177,6 +178,21 @@ def _model_access_group_cache_keys(row: _ModelAccessGroupRow) -> tuple[str, ...] return (model_access_group_cache_key(row.access_group_name),) +def _enduser_counter_key(row: _EndUserRow) -> str: + return f"spend:end_user:{row.user_id}" + + +def _enduser_cache_keys(row: _EndUserRow) -> tuple[str, ...]: + return (end_user_cache_key(row.user_id),) + + +def _enduser_carried_spend(row: _EndUserRow, caps: Mapping[str, float]) -> float: + if not caps: + return 0.0 + effective_budget_id = row.budget_id or litellm.max_end_user_budget_id + return _carried_spend(row.spend, caps.get(effective_budget_id) if effective_budget_id is not None else None) + + def _budget_link_where( budget_ids: Sequence[str], extra: Mapping[str, object] = MappingProxyType({}), @@ -650,6 +666,7 @@ class ResetBudgetJob: if _rollover_enabled() else {} # mutable-ok: empty sentinel immediately frozen by MappingProxyType ) + endusers: Final[tuple[_EndUserRow, ...]] = await self._collect_endusers_to_reset(budget_ids) return _BudgetCascade( budgets=tuple(budgets_to_reset), budget_ids=budget_ids, @@ -661,7 +678,7 @@ class ResetBudgetJob: for b in budgets_to_reset if b.budget_id is not None and b.budget_duration is not None ), - endusers=await self._collect_endusers_to_reset(budget_ids), + endusers=endusers, counter_resets=( *( (_team_membership_counter_key(row), _row_carried_spend(row, rollover_caps)) @@ -674,6 +691,10 @@ class ResetBudgetJob: (_model_access_group_counter_key(row), _row_carried_spend(row, rollover_caps)) for row in model_access_groups ), + *( + (_enduser_counter_key(row), _enduser_carried_spend(row, rollover_caps)) + for row in endusers + ), ), rollover_caps=rollover_caps, cache_keys=( @@ -682,6 +703,7 @@ class ResetBudgetJob: *(key for row in orgs for key in _org_cache_keys(row)), *(key for row in tags for key in _tag_cache_keys(row)), *(key for row in model_access_groups for key in _model_access_group_cache_keys(row)), + *(key for row in endusers for key in _enduser_cache_keys(row)), ), ) diff --git a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py index 03b05bd9d87..e6bbfd8c7d4 100644 --- a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py +++ b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py @@ -1495,6 +1495,32 @@ def test_budget_table_reset_invalidates_every_tag_not_just_the_first(reset_budge assert deleted == {"tag:tenant-a", "tag:tenant-b", "tag:tenant-c"} +def test_budget_table_reset_invalidates_enduser_counter_and_cache(reset_budget_job, mock_prisma_client, monkeypatch): + """When an end user's budget resets, its Redis spend counter is zeroed and its management cache is evicted.""" + counter_cache = _make_counter_invalidation_job(monkeypatch) + budget = _budget_row(budget_id="budget-1") + mock_prisma_client.data["budget"] = [budget] + test_enduser = type( + "LiteLLM_EndUserTable", + (), + { + "spend": 20.0, + "litellm_budget_table": budget, + "budget_id": "budget-1", + "user_id": "customer-42", + }, + ) + mock_prisma_client.data["enduser"] = [test_enduser] + + asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table()) + + counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:end_user:customer-42", value=0.0, ttl=60) + counter_cache.redis_cache.async_set_cache.assert_any_await(key="spend:end_user:customer-42", value=0.0, ttl=60) + deleted = {call.kwargs.get("key") for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list} + assert "end_user_id:customer-42" in deleted + + + def test_budget_table_reset_commits_even_when_cache_eviction_fails(reset_budget_job, mock_prisma_client, monkeypatch): """Eviction runs after the commit, so a broken cache cannot undo the write.""" counter_cache = _make_counter_invalidation_job(monkeypatch) From 3623aecc6419ed5442bea4efd435cdca3246101a Mon Sep 17 00:00:00 2001 From: amasen02 Date: Fri, 4 Sep 2026 16:33:57 +0530 Subject: [PATCH 02/11] style(proxy): add Final type annotations to enduser budget reset variables --- litellm/proxy/common_utils/reset_budget_job.py | 2 +- .../proxy/common_utils/test_reset_budget_job.py | 10 +++++----- 2 files changed, 6 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index 12ba75aea24..4fb544cbb15 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -189,7 +189,7 @@ def _enduser_cache_keys(row: _EndUserRow) -> tuple[str, ...]: def _enduser_carried_spend(row: _EndUserRow, caps: Mapping[str, float]) -> float: if not caps: return 0.0 - effective_budget_id = row.budget_id or litellm.max_end_user_budget_id + effective_budget_id: Final[str | None] = row.budget_id or litellm.max_end_user_budget_id return _carried_spend(row.spend, caps.get(effective_budget_id) if effective_budget_id is not None else None) diff --git a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py index e6bbfd8c7d4..bc9926a314f 100644 --- a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py +++ b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py @@ -4,7 +4,7 @@ import sys import types from datetime import datetime, timedelta, timezone from datetime import time as dt_time -from typing import Any, Dict, List +from typing import Any, Dict, Final, List from unittest.mock import AsyncMock, MagicMock import httpx @@ -1497,10 +1497,10 @@ def test_budget_table_reset_invalidates_every_tag_not_just_the_first(reset_budge def test_budget_table_reset_invalidates_enduser_counter_and_cache(reset_budget_job, mock_prisma_client, monkeypatch): """When an end user's budget resets, its Redis spend counter is zeroed and its management cache is evicted.""" - counter_cache = _make_counter_invalidation_job(monkeypatch) - budget = _budget_row(budget_id="budget-1") + counter_cache: Final = _make_counter_invalidation_job(monkeypatch) + budget: Final = _budget_row(budget_id="budget-1") mock_prisma_client.data["budget"] = [budget] - test_enduser = type( + test_enduser: Final = type( "LiteLLM_EndUserTable", (), { @@ -1516,7 +1516,7 @@ def test_budget_table_reset_invalidates_enduser_counter_and_cache(reset_budget_j counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:end_user:customer-42", value=0.0, ttl=60) counter_cache.redis_cache.async_set_cache.assert_any_await(key="spend:end_user:customer-42", value=0.0, ttl=60) - deleted = {call.kwargs.get("key") for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list} + deleted: Final = {call.kwargs.get("key") for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list} assert "end_user_id:customer-42" in deleted From 9c8594c7b8c4c948abb85507e7611762b0b0f70f Mon Sep 17 00:00:00 2001 From: amasen02 Date: Fri, 4 Sep 2026 16:37:58 +0530 Subject: [PATCH 03/11] style(proxy): format reset_budget_job with ruff --- litellm/proxy/common_utils/reset_budget_job.py | 5 +---- 1 file changed, 1 insertion(+), 4 deletions(-) diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index 4fb544cbb15..f2648c8466e 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -691,10 +691,7 @@ class ResetBudgetJob: (_model_access_group_counter_key(row), _row_carried_spend(row, rollover_caps)) for row in model_access_groups ), - *( - (_enduser_counter_key(row), _enduser_carried_spend(row, rollover_caps)) - for row in endusers - ), + *((_enduser_counter_key(row), _enduser_carried_spend(row, rollover_caps)) for row in endusers), ), rollover_caps=rollover_caps, cache_keys=( From 2042364fc2976ea735ab3d8c77dd4f4b27df3b84 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Fri, 4 Sep 2026 11:50:28 -0700 Subject: [PATCH 04/11] fix(proxy): strip every TypedDict qualifier before numeric form-field detection _numeric_form_type only peeled a single ReadOnly layer, so a field still wrapped in Required/NotRequired was read as non-numeric and dropped from the mapping. Which qualifiers survive get_type_hints varies by interpreter version and by include_extras, so on Python 3.10 NotRequired[ReadOnly[int]] reached the check intact and the field was silently skipped, which is what turns the mapped test red on the 3.10 leg only. Peel Required/NotRequired/ReadOnly/Annotated in any order and nesting instead. The one production caller feeds a schema with no qualifiers, so the resulting mapping is unchanged on every interpreter in the matrix, but a field written the house-convention way stops being dropped. --- .../proxy/common_utils/http_parsing_utils.py | 16 +++++++++++++--- .../common_utils/test_http_parsing_utils.py | 18 ++++++++++++++++++ 2 files changed, 31 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py index 552d1ea434f..a396e543e94 100644 --- a/litellm/proxy/common_utils/http_parsing_utils.py +++ b/litellm/proxy/common_utils/http_parsing_utils.py @@ -2,11 +2,11 @@ import json import re from collections.abc import Collection, Mapping from types import MappingProxyType, UnionType -from typing import Any, Final, Union, get_args, get_origin +from typing import Annotated, Any, Final, Union, get_args, get_origin import orjson from fastapi import Request, UploadFile, status -from typing_extensions import ReadOnly +from typing_extensions import NotRequired, ReadOnly, Required from litellm._logging import verbose_proxy_logger from litellm.constants import MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB @@ -18,6 +18,8 @@ from litellm.types.router import Deployment _FORM_CONTENT_TYPES: Final[frozenset[str]] = frozenset({"application/x-www-form-urlencoded", "multipart/form-data"}) +_ANNOTATION_QUALIFIERS: Final[frozenset[object]] = frozenset({Annotated, NotRequired, ReadOnly, Required}) + def _normalize_media_type(content_type: str) -> str: """Return the bare media type per RFC 7231: strip params, trim, lowercase.""" @@ -42,9 +44,17 @@ def _is_json_content_type(content_type: str) -> bool: return _normalize_media_type(content_type) == "application/json" +def _unqualified(annotation: object) -> object: + """Which qualifiers ``get_type_hints`` already stripped varies by interpreter version, so peel them all.""" + if get_origin(annotation) not in _ANNOTATION_QUALIFIERS: + return annotation + qualified: Final[tuple[object, ...]] = get_args(annotation) + return _unqualified(qualified[0]) + + def _numeric_form_type(annotation: object) -> type[int] | type[float] | None: """The scalar to parse an ``int``/``float``-typed field as, else ``None``.""" - unwrapped: Final = get_args(annotation)[0] if get_origin(annotation) is ReadOnly else annotation + unwrapped: Final = _unqualified(annotation) candidates: Final = ( tuple(arg for arg in get_args(unwrapped) if arg is not type(None)) if get_origin(unwrapped) in (Union, UnionType) diff --git a/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py b/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py index fcfb9342176..011571a37e0 100644 --- a/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py +++ b/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py @@ -1053,6 +1053,8 @@ class TestNumericFormFields: read_only: ReadOnly[int | None] not_required: NotRequired[ReadOnly[int]] required: Required[ReadOnly[Annotated[float, "meta"]]] + read_only_not_required: ReadOnly[NotRequired[int]] + read_only_required: ReadOnly[Required[float]] assert dict(numeric_form_fields(get_type_hints(Schema))) == { "plain": int, @@ -1061,6 +1063,22 @@ class TestNumericFormFields: "read_only": int, "not_required": int, "required": float, + "read_only_not_required": int, + "read_only_required": float, + } + + def test_qualifiers_are_unwrapped_when_get_type_hints_keeps_extras(self): + from typing_extensions import Annotated, NotRequired, ReadOnly, Required, TypedDict + + class Schema(TypedDict, total=False): + annotated: ReadOnly[Annotated[int, "meta"]] + not_required: NotRequired[ReadOnly[int]] + required: Required[ReadOnly[Annotated[float, "meta"]]] + + assert dict(numeric_form_fields(get_type_hints(Schema, include_extras=True))) == { + "annotated": int, + "not_required": int, + "required": float, } def test_non_scalar_and_bool_fields_are_skipped(self): From 976f8625f34c9c0eb7ac5e5976493dde2c4cd997 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 4 Sep 2026 16:42:29 -0700 Subject: [PATCH 05/11] test(proxy): cover default-tier end-user counter reset with rollover --- .../common_utils/test_reset_budget_job.py | 32 +++++++++++++++++++ 1 file changed, 32 insertions(+) diff --git a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py index bc9926a314f..56c0efb41d2 100644 --- a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py +++ b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py @@ -3054,6 +3054,38 @@ def test_budget_cascade_carries_enduser_overage_when_rollover_enabled( } in enduser_writes +def test_budget_cascade_carries_default_tier_enduser_counter_when_rollover_enabled( + rollover_enabled, reset_budget_job, mock_prisma_client, monkeypatch +): + """An end user on the default budget (no budget_id on its row) 5 over the cap + keeps a counter of 5 in the next window and loses its cached object.""" + import litellm + + counter_cache: Final = _make_counter_invalidation_job(monkeypatch) + monkeypatch.setattr(litellm, "max_end_user_budget_id", "default-enduser-budget") + mock_prisma_client.data["budget"] = [ + _budget_row(budget_id="default-enduser-budget", budget_duration="1d", max_budget=10.0) + ] + implicit_enduser: Final = type( + "EndUserRow", + (), + { + "spend": 15.0, + "user_id": "enduser-implicit", + "budget_id": None, + "model_dump": lambda self=None: {"spend": 15.0, "user_id": "enduser-implicit", "budget_id": None, "blocked": False}, + }, + ) + mock_prisma_client.db.litellm_endusertable.set_find_many_results([implicit_enduser]) + + asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table()) + + counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:end_user:enduser-implicit", value=5.0, ttl=60) + counter_cache.redis_cache.async_set_cache.assert_any_await(key="spend:end_user:enduser-implicit", value=5.0, ttl=60) + deleted: Final = {call.kwargs.get("key") for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list} + assert "end_user_id:enduser-implicit" in deleted + + def _replay_spend_writes(writes, spend): """Apply the queued update_many statements in order, the way the DB transaction executes them, and return the row's final spend.""" From 89086db28233514c3cc07333fcffdcd974cc8573 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 4 Sep 2026 17:39:46 -0700 Subject: [PATCH 06/11] fix(proxy): floor end-user budget checks on the DB row after a reset The reset job evicts the cached end-user object only from its own worker's in-memory cache (plus Redis), so every other uvicorn worker and replica keeps the pre-reset spend for up to user_api_key_cache_ttl (60s by default). Those workers pass that stale spend as fallback_spend, and since the authoritative floor read returned None for spend:end_user: keys, get_current_spend handed the stale value straight back and the end user kept getting 429 after the rollover on every worker but the one that ran the reset. The floor read now consults LiteLLM_EndUserTable.spend for end-user counters, the same way keys, teams, users, and orgs already read their rows. It runs only when the shared counter sits below the cached spend (a reset or a Redis restart) and stays behind the existing 5s in-process marker, so the normal request path still does no DB read. Cold end-user counters keep seeding from the cached object rather than the row, so from_db is unchanged for them. --- litellm/proxy/db/spend_counter_reseed.py | 25 +++++- litellm/proxy/proxy_server.py | 46 +++++++---- .../proxy/db/test_spend_counter_reseed.py | 69 +++++++++++++++- .../proxy/proxy_server/test_spend_counters.py | 79 +++++++++++++++++-- 4 files changed, 196 insertions(+), 23 deletions(-) diff --git a/litellm/proxy/db/spend_counter_reseed.py b/litellm/proxy/db/spend_counter_reseed.py index 7b3c261036e..a38b8a47dbd 100644 --- a/litellm/proxy/db/spend_counter_reseed.py +++ b/litellm/proxy/db/spend_counter_reseed.py @@ -26,6 +26,7 @@ from litellm.proxy._types import Litellm_EntityType from litellm.repositories.organization_repository import OrganizationRepository from litellm.repositories.table_repositories import ( BudgetWindowSpendRepository, + EndUserRepository, SpendLogsRepository, TeamMembershipRepository, ) @@ -36,6 +37,8 @@ from litellm.repositories.verification_token_repository import ( ) if TYPE_CHECKING: + from prisma.types import LiteLLM_EndUserTableWhereUniqueInput + from litellm.caching.dual_cache import DualCache from litellm.proxy.utils import PrismaClient @@ -47,6 +50,8 @@ _WINDOW_SPEND_ENTITY_TYPES: Final[Mapping[str, str]] = MappingProxyType( } ) +END_USER_COUNTER_PREFIX: Final = "spend:end_user:" + _WINDOW_SPEND_LOG_FIELDS: Final[Mapping[str, str]] = MappingProxyType( { "Key": "api_key", @@ -74,6 +79,10 @@ class SpendCounterReseed: End-user and tag spend counters intentionally do not reseed here. Their auth paths already load the corresponding objects via get_end_user_object() and get_tag_objects_batch(); callers pass those values as fallback_spend. + end_user_from_db is the one end-user read, used only as the budget floor when + a counter sits below that cached spend: a worker that did not run the budget + reset still caches the pre-reset end-user object, and LiteLLM_EndUserTable + is the row the reset zeroed. """ _locks: ClassVar["OrderedDict[str, asyncio.Lock]"] = OrderedDict() @@ -129,7 +138,7 @@ class SpendCounterReseed: elif counter_key.startswith("spend:user:"): user_id = counter_key[len("spend:user:") :] row = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id}) - elif counter_key.startswith("spend:end_user:") or counter_key.startswith("spend:tag:"): + elif counter_key.startswith(END_USER_COUNTER_PREFIX) or counter_key.startswith("spend:tag:"): return None elif counter_key.startswith("spend:org:"): org_id: Final = counter_key[len("spend:org:") :] @@ -143,6 +152,20 @@ class SpendCounterReseed: return None return float(getattr(row, "spend", 0.0) or 0.0) + @staticmethod + async def end_user_from_db(prisma_client: Optional["PrismaClient"], counter_key: str) -> float | None: + if prisma_client is None or not counter_key.startswith(END_USER_COUNTER_PREFIX): + return None + where: Final[LiteLLM_EndUserTableWhereUniqueInput] = {"user_id": counter_key[len(END_USER_COUNTER_PREFIX) :]} + try: + row: Final = await EndUserRepository(prisma_client).table.find_unique(where=where) + except Exception: # noqa: BLE001 # a failed floor read falls back to the cached spend, like from_db + verbose_proxy_logger.exception("SpendCounterReseed.end_user_from_db: failed for %s", counter_key) + return None + if row is None: + return None + return float(row.spend or 0.0) + @staticmethod def _is_key_or_team_window_counter(counter_key: str) -> bool: for prefix in ("spend:key:", "spend:team:"): diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 1f39a78e12a..47c0811d903 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -423,7 +423,7 @@ from litellm.proxy.db.proxy_worker_heartbeat import ( PROXY_WORKER_HEARTBEAT_INTERVAL_SECONDS, ProxyWorkerHeartbeat, ) -from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed +from litellm.proxy.db.spend_counter_reseed import END_USER_COUNTER_PREFIX, SpendCounterReseed from litellm.proxy.discovery_endpoints import ui_discovery_endpoints_router from litellm.proxy.fine_tuning_endpoints.endpoints import router as fine_tuning_router from litellm.proxy.fine_tuning_endpoints.endpoints import set_fine_tuning_config @@ -2580,6 +2580,29 @@ async def reseed_spend_counter_from_db(counter_key: str) -> None: await _repair_stale_spend_counter(counter_key=counter_key, db_spend=db_spend) +async def _floor_spend_from_db( + counter_key: str, + window_entity_type: str | None, + window_entity_id: str | None, + window_duration: str | None, + window_start: datetime | None, +) -> float | None: + if counter_key.startswith(END_USER_COUNTER_PREFIX): + return await SpendCounterReseed.end_user_from_db(prisma_client=prisma_client, counter_key=counter_key) + entity_spend: Final = await SpendCounterReseed.from_db(prisma_client=prisma_client, counter_key=counter_key) + if entity_spend is not None: + return entity_spend + if window_entity_type is None or window_entity_id is None or window_start is None: + return None + return await SpendCounterReseed.window_from_db( + prisma_client=prisma_client, + entity_type=window_entity_type, + entity_id=window_entity_id, + window_duration=window_duration, + window_start=window_start, + ) + + async def _authoritative_floor_spend( counter_key: str, window_entity_type: str | None = None, @@ -2592,20 +2615,13 @@ async def _authoritative_floor_spend( if cached is not None: return float(cached) - db_spend = await SpendCounterReseed.from_db(prisma_client=prisma_client, counter_key=counter_key) - if ( - db_spend is None - and window_entity_type is not None - and window_entity_id is not None - and window_start is not None - ): - db_spend = await SpendCounterReseed.window_from_db( - prisma_client=prisma_client, - entity_type=window_entity_type, - entity_id=window_entity_id, - window_duration=window_duration, - window_start=window_start, - ) + db_spend: Final = await _floor_spend_from_db( + counter_key=counter_key, + window_entity_type=window_entity_type, + window_entity_id=window_entity_id, + window_duration=window_duration, + window_start=window_start, + ) if db_spend is None: return None diff --git a/tests/test_litellm/proxy/db/test_spend_counter_reseed.py b/tests/test_litellm/proxy/db/test_spend_counter_reseed.py index 816f9ae72f4..3bd6d93328d 100644 --- a/tests/test_litellm/proxy/db/test_spend_counter_reseed.py +++ b/tests/test_litellm/proxy/db/test_spend_counter_reseed.py @@ -18,7 +18,7 @@ from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed WINDOW_START = datetime(2026, 8, 1, tzinfo=timezone.utc) -class _FakeWindowSpendTable: +class _FakeFindUniqueTable: def __init__(self, row: SimpleNamespace | None, error: Exception | None = None) -> None: self._row = row self._error = error @@ -47,10 +47,13 @@ class _FakePrismaClient: row: SimpleNamespace | None = None, spend_logs_total: float = 0.0, error: Exception | None = None, + end_user_row: SimpleNamespace | None = None, + end_user_error: Exception | None = None, ) -> None: self.db = SimpleNamespace( - litellm_budgetwindowspend=_FakeWindowSpendTable(row=row, error=error), + litellm_budgetwindowspend=_FakeFindUniqueTable(row=row, error=error), litellm_spendlogs=_FakeSpendLogsTable(total=spend_logs_total), + litellm_endusertable=_FakeFindUniqueTable(row=end_user_row, error=end_user_error), ) @@ -248,3 +251,65 @@ async def test_coalesced_window_seeds_a_cold_counter_from_the_row(): assert result == 4.5 assert cache.in_memory_cache.get_cache(key=counter_key) == 4.5 assert prisma.db.litellm_spendlogs.call_count == 0 + + +@pytest.mark.asyncio +async def test_end_user_from_db_reads_the_end_user_row_by_user_id(): + prisma = _FakePrismaClient(end_user_row=SimpleNamespace(user_id="customer-42", spend=0.0)) + + result = await SpendCounterReseed.end_user_from_db( + prisma_client=prisma, counter_key="spend:end_user:customer-42" + ) + + assert result == 0.0 + assert prisma.db.litellm_endusertable.where_clauses == [{"user_id": "customer-42"}] + + +@pytest.mark.asyncio +async def test_end_user_from_db_returns_the_recorded_spend(): + prisma = _FakePrismaClient(end_user_row=SimpleNamespace(user_id="customer-42", spend=12.5)) + + assert ( + await SpendCounterReseed.end_user_from_db(prisma_client=prisma, counter_key="spend:end_user:customer-42") + == 12.5 + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("counter_key", ["spend:key:hashed", "spend:team:t1", "spend:tag:t1"]) +async def test_end_user_from_db_ignores_other_counter_kinds_without_touching_the_db(counter_key): + prisma = _FakePrismaClient(end_user_row=SimpleNamespace(user_id="x", spend=5.0)) + + assert await SpendCounterReseed.end_user_from_db(prisma_client=prisma, counter_key=counter_key) is None + assert prisma.db.litellm_endusertable.where_clauses == [] + + +@pytest.mark.asyncio +async def test_end_user_from_db_returns_none_without_a_row_a_client_or_on_db_error(): + assert ( + await SpendCounterReseed.end_user_from_db(prisma_client=None, counter_key="spend:end_user:customer-42") + is None + ) + assert ( + await SpendCounterReseed.end_user_from_db( + prisma_client=_FakePrismaClient(end_user_row=None), counter_key="spend:end_user:customer-42" + ) + is None + ) + assert ( + await SpendCounterReseed.end_user_from_db( + prisma_client=_FakePrismaClient(end_user_error=RuntimeError("db down")), + counter_key="spend:end_user:customer-42", + ) + is None + ) + + +@pytest.mark.asyncio +async def test_from_db_still_never_reads_the_end_user_row(): + """A cold end-user counter keeps seeding from the cached end-user object the auth + path already loaded; the row is read only as the budget floor.""" + prisma = _FakePrismaClient(end_user_row=SimpleNamespace(user_id="customer-42", spend=5.0)) + + assert await SpendCounterReseed.from_db(prisma_client=prisma, counter_key="spend:end_user:customer-42") is None + assert prisma.db.litellm_endusertable.where_clauses == [] diff --git a/tests/test_litellm/proxy/proxy_server/test_spend_counters.py b/tests/test_litellm/proxy/proxy_server/test_spend_counters.py index fb3de990deb..7cc1390fcbd 100644 --- a/tests/test_litellm/proxy/proxy_server/test_spend_counters.py +++ b/tests/test_litellm/proxy/proxy_server/test_spend_counters.py @@ -223,21 +223,90 @@ async def test_get_current_spend_floor_caches_db_read(monkeypatch): @pytest.mark.asyncio async def test_get_current_spend_floors_end_user_tag_against_fallback(monkeypatch): - """End-user and tag counters have no DB row (from_db returns None). When the - counter is stale-low, enforcement falls back to the caller's recorded spend - (loaded fresh in auth) instead of trusting the stale counter.""" + """Tag counters have no DB row (from_db returns None), and an end-user counter has + none to read without a DB client. When such a counter is stale-low, enforcement + falls back to the caller's recorded spend (loaded fresh in auth) instead of + trusting the stale counter.""" fake_cache = _make_spend_counter_cache(redis_get_value=2.0) monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + monkeypatch.setattr(ps, "prisma_client", None) monkeypatch.setattr(ps.SpendCounterReseed, "from_db", AsyncMock(return_value=None)) + for counter_key in ("spend:end_user:e1", "spend:tag:t1"): + result = await ps.get_current_spend( + counter_key=counter_key, + fallback_spend=20.0, + max_budget=10.0, + ) + + assert result == 20.0 + # no DB row to repair against, so the shared counter is left untouched + fake_cache.redis_cache.async_set_max.assert_not_called() + + +def _make_prisma_with_end_user_row(spend: float | None): + prisma = MagicMock() + prisma.db.litellm_endusertable.find_unique = AsyncMock( + return_value=None if spend is None else MagicMock(spend=spend) + ) + return prisma + + +@pytest.mark.asyncio +async def test_get_current_spend_end_user_floor_admits_after_a_reset_on_a_stale_worker(monkeypatch): + """The reset job zeroes LiteLLM_EndUserTable.spend and the shared counter, but it + evicts the cached end-user object only on the worker that ran the reset. Every + other worker still passes the pre-reset spend as fallback_spend, and that stale + copy must not out-vote the reset row.""" + fake_cache = _make_spend_counter_cache(redis_get_value=0.0) + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + prisma = _make_prisma_with_end_user_row(spend=0.0) + monkeypatch.setattr(ps, "prisma_client", prisma) + result = await ps.get_current_spend( - counter_key="spend:end_user:e1", + counter_key="spend:end_user:customer-42", + fallback_spend=0.000032, + max_budget=0.00003, + fallback_authoritative=True, + ) + + assert result == 0.0 + prisma.db.litellm_endusertable.find_unique.assert_awaited_once_with(where={"user_id": "customer-42"}) + fake_cache.redis_cache.async_set_max.assert_not_called() + + +@pytest.mark.asyncio +async def test_get_current_spend_end_user_floor_repairs_a_stale_low_counter(monkeypatch): + """After a Redis restart the end-user counter can sit below the recorded spend; + the row wins and the shared counter is raised so other workers stop admitting on + the stale value.""" + fake_cache = _make_spend_counter_cache(redis_get_value=2.0) + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + monkeypatch.setattr(ps, "prisma_client", _make_prisma_with_end_user_row(spend=12.0)) + + result = await ps.get_current_spend( + counter_key="spend:end_user:customer-42", + fallback_spend=12.0, + max_budget=10.0, + ) + + assert result == 12.0 + fake_cache.redis_cache.async_set_max.assert_awaited_once_with(key="spend:end_user:customer-42", value=12.0) + + +@pytest.mark.asyncio +async def test_get_current_spend_end_user_without_a_row_keeps_the_cached_spend(monkeypatch): + fake_cache = _make_spend_counter_cache(redis_get_value=0.0) + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + monkeypatch.setattr(ps, "prisma_client", _make_prisma_with_end_user_row(spend=None)) + + result = await ps.get_current_spend( + counter_key="spend:end_user:customer-42", fallback_spend=20.0, max_budget=10.0, ) assert result == 20.0 - # no DB row to repair against, so the shared counter is left untouched fake_cache.redis_cache.async_set_max.assert_not_called() From b4fd63f621cd8dd98755944aeca0e9dfaa78daa1 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 4 Sep 2026 17:53:45 -0700 Subject: [PATCH 07/11] chore(proxy): annotate the new spend-counter test locals and correct the floor comments --- litellm/proxy/proxy_server.py | 7 ++++--- .../proxy/db/test_spend_counter_reseed.py | 11 ++++++----- .../proxy/proxy_server/test_spend_counters.py | 17 +++++++++-------- 3 files changed, 19 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 47c0811d903..9ad3c910f58 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -2477,7 +2477,8 @@ async def get_current_spend( authoritative source depends on the counter: primary key/team/user/org counters read the DB row; per-window counters (``window_start`` supplied) read the maintained window-spend row and only aggregate spend logs when - that row is missing or stale; end-user/tag counters have no DB row, so the caller's + that row is missing or stale; end-user counters read ``LiteLLM_EndUserTable``, the + row the budget reset zeroes; tag counters have no DB row, so the caller's ``fallback_spend`` (loaded fresh in auth) is authoritative. The DB read is skipped for healthy primary counters (counter at or above recorded spend) and cached in-process for a few seconds, so a persistently stale counter @@ -2511,8 +2512,8 @@ async def get_current_spend( await _repair_stale_spend_counter(counter_key=counter_key, db_spend=authoritative) return authoritative elif fallback_spend > current: - # end-user / tag counters have no DB row; fallback_spend is the - # authoritative recorded value loaded in auth. + # nothing to read (tag counters, an end user without a row or a DB client, a + # failed read); fallback_spend is the authoritative recorded value loaded in auth. return fallback_spend # Opt-in hard guarantee: when the spend backing this admit decision came diff --git a/tests/test_litellm/proxy/db/test_spend_counter_reseed.py b/tests/test_litellm/proxy/db/test_spend_counter_reseed.py index 3bd6d93328d..8cb3fc665eb 100644 --- a/tests/test_litellm/proxy/db/test_spend_counter_reseed.py +++ b/tests/test_litellm/proxy/db/test_spend_counter_reseed.py @@ -9,6 +9,7 @@ from __future__ import annotations from datetime import datetime, timedelta, timezone from types import SimpleNamespace +from typing import Final import pytest @@ -255,9 +256,9 @@ async def test_coalesced_window_seeds_a_cold_counter_from_the_row(): @pytest.mark.asyncio async def test_end_user_from_db_reads_the_end_user_row_by_user_id(): - prisma = _FakePrismaClient(end_user_row=SimpleNamespace(user_id="customer-42", spend=0.0)) + prisma: Final = _FakePrismaClient(end_user_row=SimpleNamespace(user_id="customer-42", spend=0.0)) - result = await SpendCounterReseed.end_user_from_db( + result: Final = await SpendCounterReseed.end_user_from_db( prisma_client=prisma, counter_key="spend:end_user:customer-42" ) @@ -267,7 +268,7 @@ async def test_end_user_from_db_reads_the_end_user_row_by_user_id(): @pytest.mark.asyncio async def test_end_user_from_db_returns_the_recorded_spend(): - prisma = _FakePrismaClient(end_user_row=SimpleNamespace(user_id="customer-42", spend=12.5)) + prisma: Final = _FakePrismaClient(end_user_row=SimpleNamespace(user_id="customer-42", spend=12.5)) assert ( await SpendCounterReseed.end_user_from_db(prisma_client=prisma, counter_key="spend:end_user:customer-42") @@ -278,7 +279,7 @@ async def test_end_user_from_db_returns_the_recorded_spend(): @pytest.mark.asyncio @pytest.mark.parametrize("counter_key", ["spend:key:hashed", "spend:team:t1", "spend:tag:t1"]) async def test_end_user_from_db_ignores_other_counter_kinds_without_touching_the_db(counter_key): - prisma = _FakePrismaClient(end_user_row=SimpleNamespace(user_id="x", spend=5.0)) + prisma: Final = _FakePrismaClient(end_user_row=SimpleNamespace(user_id="x", spend=5.0)) assert await SpendCounterReseed.end_user_from_db(prisma_client=prisma, counter_key=counter_key) is None assert prisma.db.litellm_endusertable.where_clauses == [] @@ -309,7 +310,7 @@ async def test_end_user_from_db_returns_none_without_a_row_a_client_or_on_db_err async def test_from_db_still_never_reads_the_end_user_row(): """A cold end-user counter keeps seeding from the cached end-user object the auth path already loaded; the row is read only as the budget floor.""" - prisma = _FakePrismaClient(end_user_row=SimpleNamespace(user_id="customer-42", spend=5.0)) + prisma: Final = _FakePrismaClient(end_user_row=SimpleNamespace(user_id="customer-42", spend=5.0)) assert await SpendCounterReseed.from_db(prisma_client=prisma, counter_key="spend:end_user:customer-42") is None assert prisma.db.litellm_endusertable.where_clauses == [] diff --git a/tests/test_litellm/proxy/proxy_server/test_spend_counters.py b/tests/test_litellm/proxy/proxy_server/test_spend_counters.py index 7cc1390fcbd..ef6f8120c82 100644 --- a/tests/test_litellm/proxy/proxy_server/test_spend_counters.py +++ b/tests/test_litellm/proxy/proxy_server/test_spend_counters.py @@ -22,6 +22,7 @@ from __future__ import annotations import asyncio from datetime import datetime +from typing import Final from unittest.mock import AsyncMock, MagicMock import pytest @@ -233,7 +234,7 @@ async def test_get_current_spend_floors_end_user_tag_against_fallback(monkeypatc monkeypatch.setattr(ps.SpendCounterReseed, "from_db", AsyncMock(return_value=None)) for counter_key in ("spend:end_user:e1", "spend:tag:t1"): - result = await ps.get_current_spend( + result: Final = await ps.get_current_spend( counter_key=counter_key, fallback_spend=20.0, max_budget=10.0, @@ -245,7 +246,7 @@ async def test_get_current_spend_floors_end_user_tag_against_fallback(monkeypatc def _make_prisma_with_end_user_row(spend: float | None): - prisma = MagicMock() + prisma: Final = MagicMock() prisma.db.litellm_endusertable.find_unique = AsyncMock( return_value=None if spend is None else MagicMock(spend=spend) ) @@ -258,9 +259,9 @@ async def test_get_current_spend_end_user_floor_admits_after_a_reset_on_a_stale_ evicts the cached end-user object only on the worker that ran the reset. Every other worker still passes the pre-reset spend as fallback_spend, and that stale copy must not out-vote the reset row.""" - fake_cache = _make_spend_counter_cache(redis_get_value=0.0) + fake_cache: Final = _make_spend_counter_cache(redis_get_value=0.0) monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) - prisma = _make_prisma_with_end_user_row(spend=0.0) + prisma: Final = _make_prisma_with_end_user_row(spend=0.0) monkeypatch.setattr(ps, "prisma_client", prisma) result = await ps.get_current_spend( @@ -280,11 +281,11 @@ async def test_get_current_spend_end_user_floor_repairs_a_stale_low_counter(monk """After a Redis restart the end-user counter can sit below the recorded spend; the row wins and the shared counter is raised so other workers stop admitting on the stale value.""" - fake_cache = _make_spend_counter_cache(redis_get_value=2.0) + fake_cache: Final = _make_spend_counter_cache(redis_get_value=2.0) monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) monkeypatch.setattr(ps, "prisma_client", _make_prisma_with_end_user_row(spend=12.0)) - result = await ps.get_current_spend( + result: Final = await ps.get_current_spend( counter_key="spend:end_user:customer-42", fallback_spend=12.0, max_budget=10.0, @@ -296,11 +297,11 @@ async def test_get_current_spend_end_user_floor_repairs_a_stale_low_counter(monk @pytest.mark.asyncio async def test_get_current_spend_end_user_without_a_row_keeps_the_cached_spend(monkeypatch): - fake_cache = _make_spend_counter_cache(redis_get_value=0.0) + fake_cache: Final = _make_spend_counter_cache(redis_get_value=0.0) monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) monkeypatch.setattr(ps, "prisma_client", _make_prisma_with_end_user_row(spend=None)) - result = await ps.get_current_spend( + result: Final = await ps.get_current_spend( counter_key="spend:end_user:customer-42", fallback_spend=20.0, max_budget=10.0, From 4ffd2ffb25017c861a350f24a27b6441ceeb2fd1 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 4 Sep 2026 18:00:40 -0700 Subject: [PATCH 08/11] test(proxy): parametrize the stale end-user counter case so no Final local sits in a loop --- .../proxy/proxy_server/test_spend_counters.py | 18 +++++++++--------- 1 file changed, 9 insertions(+), 9 deletions(-) diff --git a/tests/test_litellm/proxy/proxy_server/test_spend_counters.py b/tests/test_litellm/proxy/proxy_server/test_spend_counters.py index ef6f8120c82..86e97a334df 100644 --- a/tests/test_litellm/proxy/proxy_server/test_spend_counters.py +++ b/tests/test_litellm/proxy/proxy_server/test_spend_counters.py @@ -223,24 +223,24 @@ async def test_get_current_spend_floor_caches_db_read(monkeypatch): @pytest.mark.asyncio -async def test_get_current_spend_floors_end_user_tag_against_fallback(monkeypatch): +@pytest.mark.parametrize("counter_key", ("spend:end_user:e1", "spend:tag:t1")) +async def test_get_current_spend_floors_end_user_tag_against_fallback(monkeypatch, counter_key): """Tag counters have no DB row (from_db returns None), and an end-user counter has none to read without a DB client. When such a counter is stale-low, enforcement falls back to the caller's recorded spend (loaded fresh in auth) instead of trusting the stale counter.""" - fake_cache = _make_spend_counter_cache(redis_get_value=2.0) + fake_cache: Final = _make_spend_counter_cache(redis_get_value=2.0) monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) monkeypatch.setattr(ps, "prisma_client", None) monkeypatch.setattr(ps.SpendCounterReseed, "from_db", AsyncMock(return_value=None)) - for counter_key in ("spend:end_user:e1", "spend:tag:t1"): - result: Final = await ps.get_current_spend( - counter_key=counter_key, - fallback_spend=20.0, - max_budget=10.0, - ) + result: Final = await ps.get_current_spend( + counter_key=counter_key, + fallback_spend=20.0, + max_budget=10.0, + ) - assert result == 20.0 + assert result == 20.0 # no DB row to repair against, so the shared counter is left untouched fake_cache.redis_cache.async_set_max.assert_not_called() From cd113c3a2e46bd8d468ea48219d861c107e06260 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 4 Sep 2026 18:13:51 -0700 Subject: [PATCH 09/11] ci: allowlist the bounded _unqualified qualifier peel in the recursion detector --- tests/code_coverage_tests/recursive_detector.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/code_coverage_tests/recursive_detector.py b/tests/code_coverage_tests/recursive_detector.py index 0578dc60119..e9f87ba6cae 100644 --- a/tests/code_coverage_tests/recursive_detector.py +++ b/tests/code_coverage_tests/recursive_detector.py @@ -66,6 +66,7 @@ IGNORE_FUNCTIONS = [ "_json_safe", # max depth set (_MAX_DEPTH) plus a seen-ids cycle guard for self-referential input. "_redact_agent_params_tree", # max depth set (default 10), same shape as _redact_sensitive_litellm_params. "_restore_redacted_nested_value", # max depth set (default 10), mirrors _redact_agent_params_tree on the write side. + "_unqualified", # bounded by the qualifier depth of a static TypedDict annotation (Annotated, Required/NotRequired, ReadOnly around one type, no cycles possible). ] From 853fed824e841ce3e3952f137b6a3d577ced07d4 Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Fri, 4 Sep 2026 18:40:35 -0700 Subject: [PATCH 10/11] fix(bedrock): stop sending toolConfig tool definitions to guardrails on passthrough converse (#39281) Bedrock passthrough Converse routes flattened every non-empty string under toolConfig.tools into the guardrail INPUT texts, so tool names, tool descriptions and JSON-schema strings (object, property names, titles, type names, enum values) each arrived as a separate guardrail item. A request whose only prompt was one benign user message could be blocked outright because a denied term appeared in an app-authored tool definition. Tool definitions are now excluded from the extracted texts, matching every other guardrail translation handler, which carries tool definitions in the structured tools input rather than in texts. Caller content stays scanned: message text, toolUse.input, toolResult content and json, and additionalModelRequestFields are unchanged. Resolves LIT-5797 --- .../guardrail_translation/handler.py | 27 ++++-- .../guardrail_translation/test_handler.py | 97 ++++++++++++++----- 2 files changed, 90 insertions(+), 34 deletions(-) diff --git a/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py b/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py index 9d35a87855e..b87f6196e51 100644 --- a/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py +++ b/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py @@ -89,11 +89,24 @@ def _extract_converse_texts( top-level ``text`` blocks this scans the arbitrary-JSON fields a caller can hide prompt content in -- ``toolUse.input`` and ``toolResult.content[].json`` (alongside ``toolResult.content[].text``) -- - as well as the request-level fields still forwarded to Bedrock that a caller - can route blocked content through: ``toolConfig.tools`` (tool names, - descriptions and input schemas) and ``additionalModelRequestFields``. Tool - message blocks are skipped when tool messages are excluded, but tool - definitions are always scanned to match the chat-completions guardrail path. + as well as ``additionalModelRequestFields``, a free-form model-parameter bag + with no schema that a caller can route blocked content through. + + ``toolConfig.tools`` is deliberately NOT scanned. Tool definitions are + app-authored config, so their names, descriptions and JSON-schema strings + ("object", property names, titles, type names, enum values) would each reach + the guardrail as a separate INPUT item, producing false positives and + inflating guardrail usage for a request whose only prompt is one user + message. No other guardrail translation handler puts tool definitions in + ``texts``; the chat and messages handlers carry them in the structured + ``tools`` input instead, which this handler does not populate because a + Bedrock ``toolSpec`` is not the OpenAI tool shape those consumers expect. + + ``additionalModelRequestFields`` is treated differently on purpose. Bedrock + gives ``toolConfig.tools`` a fixed schema whose contents are tool metadata by + contract, while ``additionalModelRequestFields`` is free-form and defined by + the target model, so what it carries cannot be classified without knowing + that model. Scanning it stays the fail-closed default. """ holders: Final[list[_StringHolder]] = [] @@ -121,10 +134,6 @@ def _extract_converse_texts( _collect_block_text(inner, holders) _collect_strings(inner.get("json"), holders) - tool_config: Final = body.get("toolConfig") - if isinstance(tool_config, dict): - _collect_strings(tool_config.get("tools"), holders) - _collect_strings(body.get("additionalModelRequestFields"), holders) texts: Final = [container[key] for container, key in holders] diff --git a/tests/test_litellm/llms/bedrock/passthrough/guardrail_translation/test_handler.py b/tests/test_litellm/llms/bedrock/passthrough/guardrail_translation/test_handler.py index 744ed50dbcb..dee8366ce2d 100644 --- a/tests/test_litellm/llms/bedrock/passthrough/guardrail_translation/test_handler.py +++ b/tests/test_litellm/llms/bedrock/passthrough/guardrail_translation/test_handler.py @@ -170,22 +170,27 @@ class TestExtractConverseTexts: texts, _ = _extract_converse_texts(body, skip_system=False, skip_tool=False) assert texts == [] - def test_extracts_tool_config_description_and_schema(self): + def test_tool_config_definitions_not_extracted(self): + """Tool definitions are app-authored config, so nothing under + toolConfig.tools reaches the guardrail as input content.""" body = { - "messages": [{"role": "user", "content": [{"text": "hi"}]}], + "messages": [ + {"role": "user", "content": [{"text": "How much lag is there in my data?"}]} + ], "toolConfig": { "tools": [ { "toolSpec": { "name": "lookup", - "description": "blocked tool description", + "description": "tool description", "inputSchema": { "json": { "type": "object", "properties": { - "q": { + "agent_name": { "type": "string", - "description": "blocked schema description", + "title": "Agent Name", + "enum": ["alpha", "beta", "gamma"], } }, } @@ -196,20 +201,56 @@ class TestExtractConverseTexts: }, } texts, _ = _extract_converse_texts(body, skip_system=False, skip_tool=False) - assert "blocked tool description" in texts - assert "blocked schema description" in texts + assert texts == ["How much lag is there in my data?"] - def test_tool_config_scanned_even_when_tool_messages_skipped(self): + def test_every_tool_definition_excluded_not_just_the_first(self): + """A per-tool scan that only skipped tools[0] would still leak the rest.""" body = { "messages": [{"role": "user", "content": [{"text": "hi"}]}], "toolConfig": { "tools": [ - {"toolSpec": {"name": "fn", "description": "blocked description"}} + {"toolSpec": {"name": "first", "description": "first description"}}, + {"toolSpec": {"name": "second", "description": "second description"}}, + {"toolSpec": {"name": "third", "description": "third description"}}, + ] + }, + } + texts, _ = _extract_converse_texts(body, skip_system=False, skip_tool=False) + assert texts == ["hi"] + + def test_tool_config_definitions_not_extracted_when_tool_messages_skipped(self): + body = { + "messages": [{"role": "user", "content": [{"text": "hi"}]}], + "toolConfig": { + "tools": [ + {"toolSpec": {"name": "fn", "description": "tool description"}} ] }, } texts, _ = _extract_converse_texts(body, skip_system=False, skip_tool=True) - assert "blocked description" in texts + assert texts == ["hi"] + + def test_tool_use_input_still_extracted_alongside_tool_config(self): + """Only tool DEFINITIONS are excluded; caller content inside a toolUse + block is still scanned.""" + body = { + "messages": [ + { + "role": "user", + "content": [ + {"text": "hi"}, + {"toolUse": {"toolUseId": "t1", "name": "fn", "input": {"q": "user secret"}}}, + ], + } + ], + "toolConfig": { + "tools": [ + {"toolSpec": {"name": "fn", "description": "tool description"}} + ] + }, + } + texts, _ = _extract_converse_texts(body, skip_system=False, skip_tool=False) + assert texts == ["hi", "user secret"] def test_extracts_additional_model_request_fields(self): body = { @@ -437,9 +478,9 @@ class TestBedrockPassthroughGuardrailHandlerInput: assert "blocked content" in sent_texts @pytest.mark.asyncio - async def test_tool_config_description_scanned_and_masked(self): - """Blocked text hidden in toolConfig.tools[].toolSpec.description is still - forwarded to Bedrock, so the guardrail must see it and mask it in place.""" + async def test_tool_config_definitions_not_sent_and_left_untouched(self): + """Tool definitions never reach the guardrail, and the body forwarded to + Bedrock keeps them byte for byte.""" handler = BedrockPassthroughGuardrailHandler() data = _converse_data() data["data"]["toolConfig"] = { @@ -453,36 +494,42 @@ class TestBedrockPassthroughGuardrailHandlerInput: } ] } - guardrail = _make_guardrail( - {"texts": ["You are helpful.", "Hello world", "lookup", "[REDACTED]", "object"]} - ) + original_tool_config = copy.deepcopy(data["data"]["toolConfig"]) + guardrail = _make_guardrail({"texts": ["[REDACTED]", "[REDACTED]"]}) result = await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) sent_texts = guardrail.apply_guardrail.call_args.kwargs["inputs"]["texts"] - assert "email john@example.com" in sent_texts - tool_spec = result["data"]["toolConfig"]["tools"][0]["toolSpec"] - assert tool_spec["description"] == "[REDACTED]" + assert sent_texts == ["You are helpful.", "Hello world"] + assert result["data"]["toolConfig"] == original_tool_config @pytest.mark.asyncio - async def test_tool_config_description_blocking_propagates(self): - """A blocking guardrail must reject content hidden in a tool description.""" + async def test_blocking_guardrail_not_triggered_by_tool_description(self): + """LIT-5797: a request whose only prompt is a benign user message must not + be blocked because a denied term appears in a tool definition.""" handler = BedrockPassthroughGuardrailHandler() data = _converse_data() data["data"]["toolConfig"] = { "tools": [{"toolSpec": {"name": "fn", "description": "blocked content"}}] } + + async def _block_on_denied_term(**kwargs): + texts = kwargs["inputs"]["texts"] + if any("blocked content" in text for text in texts): + raise GuardrailBlocked("Blocked") + return {"texts": texts} + guardrail = MagicMock() guardrail.guardrail_name = "block-guard" guardrail.skip_system_message_in_guardrail = False guardrail.skip_tool_message_in_guardrail = False - guardrail.apply_guardrail = AsyncMock(side_effect=GuardrailBlocked("Blocked")) + guardrail.apply_guardrail = AsyncMock(side_effect=_block_on_denied_term) - with pytest.raises(GuardrailBlocked): - await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + result = await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) sent_texts = guardrail.apply_guardrail.call_args.kwargs["inputs"]["texts"] - assert "blocked content" in sent_texts + assert "blocked content" not in sent_texts + assert result["data"]["toolConfig"]["tools"][0]["toolSpec"]["description"] == "blocked content" @pytest.mark.asyncio async def test_additional_model_request_fields_scanned_and_masked(self): From 2c3c7dd1a60f80595350671c7f654dfb2c16933d Mon Sep 17 00:00:00 2001 From: moe-berri Date: Fri, 4 Sep 2026 18:41:47 -0700 Subject: [PATCH 11/11] feat(shadow_eval): judge tool-call turns instead of dropping or erroring on them (#39818) * fix(shadow_eval): tell a tool-call shadow reply apart from an empty one Both arrive at the attempt row as the same 'shadow router returned an empty response', because _chat_final_text returns empty for a tool-final turn by design and for a reply that genuinely carried no text. Those are different things: an arm that chose a tool where the real model wrote prose is a divergence a text judge cannot score, and the sampling side already drops the real arm's tool-final turns for exactly that reason, so the shadow side reads as a fault where the real side reads as a filter. A job that is almost all 'empty response' gives no way to tell a tool-happy arm from a broken one. The error now names which of the two happened, and carries the finish_reason and the routed model so the row says what the arm was doing. Every varying part sits behind the first semicolon: operators read these by grouping on the error text, and interpolating the model into the leading sentence would make each row its own group. The outcome stays 'error'. Whether a tool-call reply should instead be its own non-judged outcome, excluded from the loss rate the way the real arm's tool-final turns already are, needs the four aggregation predicates that spell judged as outcome != 'error' rewritten, and a decision on how to surface the new bucket. That is a separate change. * fix(shadow_eval): read the tool name of a custom tool call A custom tool call carries its name under custom.name with no function key, so every one of them reported as tool=unnamed. * feat(shadow_eval): judge tool calls instead of dropping the turn A turn where either arm called a tool was discarded before it could be compared: the real arm's at sampling, the shadow arm's as an error row. On agentic traffic that is most of the traffic, so a job set to sample 10% was sampling 10% of the prose-only slice. Tool calls now serialize to text on every surface and are judged like any other response, and the judge is told a tool call is not a defect so it scores the choice rather than the shape. * feat(shadow_eval): show the judge what tools were available Both arms were offered the same tools, but the judge only ever saw the chosen call in isolation, with no way to tell whether a better tool existed or the arguments matched what the tool expects. Threads the request's tool definitions (name and description only) into the judge prompt, capped and omitted entirely on turns that offered none. * fix(shadow_eval): read a custom tool definition's name from custom, not function A chat-completions custom tool definition nests name and description under custom, mirroring how a custom tool call nests them (openai.types.chat. ChatCompletionCustomToolParam). Reading only function rendered every one as unnamed, telling the judge nothing about what it was. --- litellm/integrations/shadow_eval_logger.py | 163 ++++++++-- .../integrations/test_shadow_eval_logger.py | 287 +++++++++++++++++- 2 files changed, 416 insertions(+), 34 deletions(-) diff --git a/litellm/integrations/shadow_eval_logger.py b/litellm/integrations/shadow_eval_logger.py index 27da785331a..b28d21ba1ca 100644 --- a/litellm/integrations/shadow_eval_logger.py +++ b/litellm/integrations/shadow_eval_logger.py @@ -165,28 +165,91 @@ def _chat_request_from_responses( ) -def _chat_final_text(response_obj: object) -> str: - """The assistant's text, or empty when the turn carries tool calls: only text-final - turns produce a judgeable A/B comparison.""" +def _chat_choice(response_obj: object) -> object | None: + """The response's first choice, from a payload mapping or a duck-typed ModelResponse.""" try: - message: Final = ( - response_obj["choices"][0]["message"] - if isinstance(response_obj, Mapping) - else response_obj.choices[0].message # pyright: ignore[reportAttributeAccessIssue] # duck-typed ModelResponse - ) + if isinstance(response_obj, Mapping): + return response_obj["choices"][0] + return response_obj.choices[0] # pyright: ignore[reportAttributeAccessIssue] # duck-typed ModelResponse except (AttributeError, KeyError, IndexError, TypeError): + return None + + +def _field_reader(obj: object) -> Callable[[str], object]: + return obj.get if isinstance(obj, Mapping) else lambda key: getattr(obj, key, None) + + +def _chat_message_reader(response_obj: object) -> Callable[[str], object] | None: + """Field access over the assistant message of a chat response, or None for a payload + with no readable message.""" + choice: Final = _chat_choice(response_obj) + if choice is None: + return None + message: Final = _field_reader(choice)("message") + return _field_reader(message) if message is not None else None + + +def _chat_final_text(response_obj: object) -> str: + """The turn's judgeable text: prose, or every tool call serialized alongside it as + `[tool call] name(arguments)` when the assistant chose to act instead of, or as well + as, answering directly. A tool call is a real turn, not a gap, so this is what both + the real arm's sampling decision and the shadow arm's reply compare against.""" + read: Final = _chat_message_reader(response_obj) + if read is None: return "" - read: Final = message.get if isinstance(message, Mapping) else lambda key: getattr(message, key, None) - if read("tool_calls") or read("function_call"): - return "" - return extract_text_from_content(read("content")) + prose: Final = extract_text_from_content(read("content")) + if not (read("tool_calls") or read("function_call")): + return prose + serialized: Final = _serialize_tool_calls(read) + return f"{prose} {serialized}".strip() if prose else serialized + + +def _chat_finish_reason(response_obj: object) -> str: + choice: Final = _chat_choice(response_obj) + raw: Final = _field_reader(choice)("finish_reason") if choice is not None else None + return str(raw) if raw else "unknown" + + +_RESPONSES_TOOL_CALL_TYPES: Final = frozenset(("function_call", "custom_tool_call")) + + +def _tool_calls_list(read: Callable[[str], object]) -> tuple[object, ...]: + calls: Final = read("tool_calls") + listed: Final = tuple(calls) if isinstance(calls, Sequence) and not isinstance(calls, str) else () + single: Final = read("function_call") + return listed if listed else ((single,) if single is not None else ()) + + +def _tool_call_invocation(call: object) -> str: + """One tool call as `name(arguments)`. Custom tool calls name themselves and carry their + arguments under `custom` rather than `function`.""" + read_call: Final = _field_reader(call) + payload: Final = read_call("function") or read_call("custom") or call + read_payload: Final = _field_reader(payload) + name: Final = read_payload("name") + arguments: Final = read_payload("arguments") or read_payload("input") or "" + return f"{name or 'unnamed'}({arguments})" + + +def _serialize_tool_calls(read: Callable[[str], object]) -> str: + """Every tool call in a reply as text a judge built for prose can still read.""" + return ", ".join(f"[tool call] {_tool_call_invocation(call)}" for call in _tool_calls_list(read)) + + +def _shadow_empty_reply_error(response_obj: object, routed_model: str) -> str: + """Why a shadow reply yielded no judgeable text at all: no prose, and no tool call to + serialize either. The stable sentence comes first and every varying part after the + semicolon, so grouping rows by error still yields one row per cause.""" + detail: Final = f"finish_reason={_chat_finish_reason(response_obj)}, model={routed_model or 'unknown'}" + return f"shadow router returned an empty response; {detail}" def _responses_final_text(response_obj: object) -> str: - """The turn's aggregated output text, or empty when the turn carries tool calls. A - dict-shaped payload is validated into the owner type first, because ``output_text`` - is a derived property rather than a serialized field, so it never exists on a dict; - a dict the owner type rejects is unjudgeable and skipped.""" + """The turn's judgeable text: the aggregated output plus any tool call serialized + alongside it, the same way the chat surface renders one. A dict-shaped payload is + validated into the owner type first, because ``output_text`` is a derived property + rather than a serialized field, so it never exists on a dict; a dict the owner type + rejects is unjudgeable and skipped.""" from litellm.types.llms.openai import ResponsesAPIResponse try: @@ -199,11 +262,16 @@ def _responses_final_text(response_obj: object) -> str: if not isinstance(output, Sequence): return "" items: Final = tuple(item.model_dump() if isinstance(item, BaseModel) else item for item in output) - if any( - not isinstance(item, Mapping) or item.get("type") in ("function_call", "custom_tool_call") for item in items - ): + if any(not isinstance(item, Mapping) for item in items): return "" - return str(getattr(response, "output_text", "") or "") + calls: Final = tuple( + item for item in items if isinstance(item, Mapping) and item.get("type") in _RESPONSES_TOOL_CALL_TYPES + ) + prose: Final = str(getattr(response, "output_text", "") or "") + if not calls: + return prose + serialized: Final = ", ".join(f"[tool call] {_tool_call_invocation(call)}" for call in calls) + return f"{prose} {serialized}".strip() if prose else serialized class _SurfaceOps: @@ -273,8 +341,8 @@ def _judgeable_sample( response_obj: object, ) -> tuple[tuple[Mapping[str, object], ...], Mapping[str, object], str] | None: """The normalized chat conversation, the forwardable generation params, and the - judgeable final text; None when this request's shapes cannot be sampled (tool-final - turn, empty text, or a shape the owner transformations reject).""" + judgeable final text; None when this request's shapes cannot be sampled (no text and no + tool call to serialize, or a shape the owner transformations reject).""" try: request: Final = ops.chat_request(kwargs, model_parameters) items: Final = _MESSAGE_ITEMS_ADAPTER.validate_python(request.get("messages")) @@ -307,6 +375,11 @@ PAIRWISE_JUDGE_SYSTEM_PROMPT: Final = """You are an impartial quality judge comp The responses are labeled A and B in random order. You do not know which system produced which. +A response may be prose, or a tool call shown as `[tool call] name(arguments)` if the +assistant chose to act instead of answering directly. A tool call is not a defect: judge +whether calling that tool was the right response to the conversation, the same as you +would judge prose. + Criteria: correctness, completeness, clarity, conciseness. Return ONLY valid JSON in this exact format, no other text: @@ -376,14 +449,37 @@ def _unmask_preference(raw_preference: str, real_is_a: bool) -> str: return "tie" -def _judge_user_prompt(conversation: str, response_a: str, response_b: str) -> str: +_MAX_JUDGE_TOOL_DEFS_CHARS: Final = 2_000 + + +def _tool_definitions_text(tools: object) -> str: + """The tools available to both arms, name and description only: enough for the judge + to tell whether the chosen tool, and not some other one, was the right call, without + forwarding parameter schemas it does not need to score that.""" + if not isinstance(tools, Sequence) or isinstance(tools, str): + return "" + entries: Final = tuple( + _field_reader(t)("function") or _field_reader(t)("custom") or t for t in tools if not isinstance(t, str) + ) + lines: Final = tuple( + f"- {_field_reader(e)('name') or 'unnamed'}: {_field_reader(e)('description') or 'no description'}" + for e in entries + ) + if not lines: + return "" + return ("Tools available to both responses:\n" + "\n".join(lines))[:_MAX_JUDGE_TOOL_DEFS_CHARS] + + +def _judge_user_prompt(conversation: str, response_a: str, response_b: str, tool_definitions: str = "") -> str: """The judge prompt under one total character budget: each response is capped, and - the conversation tail gets whatever budget the responses left over.""" + the conversation tail gets whatever budget the responses and tool definitions left + over.""" a: Final = response_a[:_MAX_JUDGE_RESPONSE_CHARS] b: Final = response_b[:_MAX_JUDGE_RESPONSE_CHARS] - conversation_budget: Final = _MAX_JUDGE_PROMPT_CHARS - len(a) - len(b) + prefix: Final = f"{tool_definitions}\n\n" if tool_definitions else "" + conversation_budget: Final = _MAX_JUDGE_PROMPT_CHARS - len(a) - len(b) - len(prefix) return ( - f"Conversation:\n{conversation[-conversation_budget:]}\n\n" + f"{prefix}Conversation:\n{conversation[-conversation_budget:]}\n\n" f"Response A:\n{a}\n\n" f"Response B:\n{b}\n\n" "Which response is better?" @@ -942,6 +1038,7 @@ class ShadowEvalLogger(CustomLogger): messages=messages, real_text=real_text, shadow_text=shadow.text, + tools=shadow_params.get("tools"), parent_metadata=parent_metadata, ) if isinstance(verdict, _CallFailure): @@ -1080,15 +1177,18 @@ class ShadowEvalLogger(CustomLogger): classifier_cost=_decision_classifier_cost(shadow_metadata), ) text: Final = _chat_final_text(response) + routed_model: Final = str( + getattr(response, "model", None) or _routing_decision(shadow_metadata).get("routed_model") or "" + ) if not text: return _CallFailure( - "shadow router returned an empty response", + _shadow_empty_reply_error(response, routed_model), cost=_call_cost(response), classifier_cost=_decision_classifier_cost(shadow_metadata), ) return _ShadowResponse( text=text, - model=str(getattr(response, "model", None) or _routing_decision(shadow_metadata).get("routed_model") or ""), + model=routed_model, tier=_routed_tier(shadow_metadata), cost=_call_cost(response), classifier_cost=_decision_classifier_cost(shadow_metadata), @@ -1100,9 +1200,12 @@ class ShadowEvalLogger(CustomLogger): messages: Sequence[Mapping[str, object]], real_text: str, shadow_text: str, + tools: object, parent_metadata: Mapping[str, object], ) -> "_JudgeVerdict | _CallFailure": - """Blind pairwise judge with A/B labels randomized to cancel position bias.""" + """Blind pairwise judge with A/B labels randomized to cancel position bias. Both + arms were offered the same tools, so the judge is shown their definitions too: a + tool call is only assessable against what else was available to call instead.""" real_is_a: Final = random.random() < 0.5 response_a: Final = real_text if real_is_a else shadow_text response_b: Final = shadow_text if real_is_a else real_text @@ -1117,7 +1220,7 @@ class ShadowEvalLogger(CustomLogger): {"role": "system", "content": PAIRWISE_JUDGE_SYSTEM_PROMPT}, # mutable-ok: SDK message { "role": "user", - "content": _judge_user_prompt(conversation, response_a, response_b), + "content": _judge_user_prompt(conversation, response_a, response_b, _tool_definitions_text(tools)), }, # mutable-ok: SDK message ] try: diff --git a/tests/test_litellm/integrations/test_shadow_eval_logger.py b/tests/test_litellm/integrations/test_shadow_eval_logger.py index 5628d69de26..f273a285d49 100644 --- a/tests/test_litellm/integrations/test_shadow_eval_logger.py +++ b/tests/test_litellm/integrations/test_shadow_eval_logger.py @@ -24,7 +24,13 @@ from litellm.integrations.shadow_eval_logger import ( _unmask_preference, ) from litellm.types.guardrails import GuardrailEventHooks -from litellm.types.utils import SHADOW_EVAL_JUDGE_CALL_ORIGIN, SHADOW_EVAL_ROUTER_CALL_ORIGIN, ModelResponse +from litellm.types.utils import ( + SHADOW_EVAL_JUDGE_CALL_ORIGIN, + SHADOW_EVAL_ROUTER_CALL_ORIGIN, + ChatCompletionCustomToolCallPayload, + ChatCompletionMessageCustomToolCall, + ModelResponse, +) def _job(**overrides) -> ActiveShadowEvalJob: @@ -120,6 +126,39 @@ def _router( return router +def _shadow_reply_router(message, finish_reason="stop", routed_model="cheap-model"): + """A router whose shadow arm answers with a caller-supplied message, so a reply that + yields no judgeable text can be posed as the two different things it can be: an arm + that chose a tool, or an arm that returned nothing.""" + router = MagicMock() + router.model_group_alias = {} + router.get_model_list = MagicMock(return_value=[{"litellm_params": {"model": "openai/gpt-4o-mini"}}]) + + async def acompletion(**kwargs): + if kwargs["metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) != SHADOW_EVAL_ROUTER_CALL_ORIGIN: + return {"choices": [{"message": {"content": '{"preference": "A", "confidence": 0.9}'}}]} + kwargs["metadata"]["routing_decision"] = {"tier_label": "SIMPLE", "routed_model": routed_model} + return {"choices": [{"message": message, "finish_reason": finish_reason}]} + + router.acompletion = MagicMock(side_effect=acompletion) + return router + + +TOOL_CALL_MESSAGE = { + "content": None, + "tool_calls": [{"id": "c1", "type": "function", "function": {"name": "Read", "arguments": "{}"}}], +} + +CUSTOM_TOOL_CALL_MESSAGE = { + "content": None, + "tool_calls": [ + ChatCompletionMessageCustomToolCall( + id="c2", custom=ChatCompletionCustomToolCallPayload(name="exec_sql", input="select 1") + ) + ], +} + + def _spend_counter(store=None): """In-memory stand-in for the proxy's cross-pod spend counter: reads take the max of the counter and the caller's fallback, exactly like get_current_spend does for a key @@ -368,7 +407,13 @@ class TestSurfaceNormalization: ], ids=["tool-final-chat-turn", "tool-final-responses-turn"], ) - async def test_unjudgeable_turns_are_skipped_without_consuming_budget(self, response_mutation, kwargs_mutation): + async def test_a_tool_final_turn_is_sampled_and_serialized_for_the_judge( + self, response_mutation, kwargs_mutation + ): + """A turn where the real model called a tool used to be dropped before sampling, on + every surface. On agentic traffic that is most of the traffic, so a job set to + sample 10% was really sampling 10% of the prose-only slice and calling it 10% of + the key. The turn is sampled like any other and the call is serialized as text.""" from litellm.types.llms.openai import ResponsesAPIResponse hook_kwargs = _success_kwargs(**({"call_type": "acompletion"} | kwargs_mutation)) @@ -406,6 +451,38 @@ class TestSurfaceNormalization: prisma, router = await self._drive(hook_kwargs, response) + judge_prompt = next( + call.kwargs["messages"][-1]["content"] + for call in router.acompletion.call_args_list + if call.kwargs["metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) != SHADOW_EVAL_ROUTER_CALL_ORIGIN + ) + assert "[tool call] f({})" in judge_prompt + prisma.db.litellm_shadowevalattempt.create.assert_called_once() + + @pytest.mark.parametrize( + "response_mutation,kwargs_mutation", + [ + ("chat-no-content", {}), + ("responses-no-output", {"call_type": "aresponses"}), + ], + ids=["empty-chat-turn", "empty-responses-turn"], + ) + async def test_turns_with_nothing_to_compare_are_skipped_without_consuming_budget( + self, response_mutation, kwargs_mutation + ): + """No prose and no tool call leaves the judge nothing to score, so the turn is + still skipped rather than billed.""" + from litellm.types.llms.openai import ResponsesAPIResponse + + hook_kwargs = _success_kwargs(**({"call_type": "acompletion"} | kwargs_mutation)) + if response_mutation == "chat-no-content": + response = {"choices": [{"message": {"content": ""}}]} + else: + hook_kwargs["messages"] = "do the thing" + response = ResponsesAPIResponse.model_validate(RESPONSES_API_RESPONSE | {"output": []}) + + prisma, router = await self._drive(hook_kwargs, response) + router.acompletion.assert_not_called() prisma.db.litellm_shadowevalattempt.create.assert_not_called() @@ -1134,6 +1211,206 @@ class TestShadowPipeline: assert row["shadow_cost"] == 0.007 assert logger._test_counter["spend:shadow_eval:job-1"] == 0.007 + async def _no_text_error(self, router) -> str: + prisma = _prisma() + await _logger(router=router, prisma=prisma)._run_shadow_eval( + job=_job(), + request_id="req-1", + messages=({"role": "user", "content": "hi"},), + real_text="real answer", + real_model="claude-opus", + real_cost=0.0, + real_classifier_cost=0.0, + real_cache_hit=False, + control_tier=None, + shadow_params={}, + parent_metadata={}, + ) + row = prisma.db.litellm_shadowevalattempt.create.call_args.kwargs["data"] + assert row["outcome"] == "error" + return row["error"] + + async def _judged_shadow_row(self, router: MagicMock, shadow_params: dict | None = None) -> dict: + prisma = _prisma() + await _logger(router=router, prisma=prisma)._run_shadow_eval( + job=_job(), + request_id="req-1", + messages=({"role": "user", "content": "hi"},), + real_text="real answer", + real_model="claude-opus", + real_cost=0.0, + real_classifier_cost=0.0, + real_cache_hit=False, + control_tier=None, + shadow_params=shadow_params or {}, + parent_metadata={}, + ) + return prisma.db.litellm_shadowevalattempt.create.call_args.kwargs["data"] + + async def test_a_tool_call_shadow_reply_is_judged_rather_than_discarded(self): + """An arm that calls a tool where the real model wrote prose has answered, it just + answered by acting. Dropping that turn threw away the comparison the job exists to + make, and on agentic traffic it threw away most of them, so the tool call is + serialized into text and judged like any other response.""" + row = await self._judged_shadow_row(_shadow_reply_router(TOOL_CALL_MESSAGE, finish_reason="tool_calls")) + + assert row["outcome"] != "error" + assert row["error"] is None + assert row["confidence"] == 0.9 + + async def test_a_tool_call_reaches_the_judge_as_readable_text(self): + """The judge only ever sees strings, so a tool call has to arrive as its name and + arguments. A serialization that dropped either would ask the judge to score a + response it cannot tell apart from any other tool call.""" + router = _shadow_reply_router(TOOL_CALL_MESSAGE, finish_reason="tool_calls") + await self._judged_shadow_row(router) + + judge_prompt = next( + call.kwargs["messages"][-1]["content"] + for call in router.acompletion.call_args_list + if call.kwargs["metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) != SHADOW_EVAL_ROUTER_CALL_ORIGIN + ) + + assert "[tool call] Read({})" in judge_prompt + + async def test_the_judge_sees_what_tools_were_available(self): + """Scoring whether a tool call was the right response needs to know what else the + arm could have called instead. Without the tool list, the judge can score the + arguments but not whether Read, specifically, was the correct choice.""" + router = _shadow_reply_router(TOOL_CALL_MESSAGE, finish_reason="tool_calls") + tools = [ + {"type": "function", "function": {"name": "Read", "description": "read a file from disk"}}, + {"type": "function", "function": {"name": "Bash", "description": "run a shell command"}}, + ] + await self._judged_shadow_row(router, shadow_params={"tools": tools}) + + judge_prompt = next( + call.kwargs["messages"][-1]["content"] + for call in router.acompletion.call_args_list + if call.kwargs["metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) != SHADOW_EVAL_ROUTER_CALL_ORIGIN + ) + + assert "Read: read a file from disk" in judge_prompt + assert "Bash: run a shell command" in judge_prompt + + async def test_a_custom_tool_definition_is_named_for_the_judge(self): + """A custom tool definition nests name and description under `custom`, not + `function`, so reading only `function` renders every one of them as unnamed and + tells the judge nothing about what the arm could have called.""" + from openai.types.chat import ChatCompletionCustomToolParam + + router = _shadow_reply_router(TOOL_CALL_MESSAGE, finish_reason="tool_calls") + tools = [ + ChatCompletionCustomToolParam( + type="custom", + custom={"name": "exec_sql", "description": "run a read-only sql query"}, + ) + ] + await self._judged_shadow_row(router, shadow_params={"tools": tools}) + + judge_prompt = next( + call.kwargs["messages"][-1]["content"] + for call in router.acompletion.call_args_list + if call.kwargs["metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) != SHADOW_EVAL_ROUTER_CALL_ORIGIN + ) + + assert "exec_sql: run a read-only sql query" in judge_prompt + assert "unnamed" not in judge_prompt + + @pytest.mark.parametrize("shadow_params", [{}, {"tools": []}], ids=["omitted", "empty-list"]) + async def test_no_tool_definitions_section_when_the_turn_offered_no_tools(self, shadow_params): + """Padding every judge prompt with an empty tools section wastes budget on the + turns, still the majority, that never offered one, whether tools was left out of + the request entirely or sent as an empty list.""" + router = _shadow_reply_router({"content": "hello"}, finish_reason="stop") + await self._judged_shadow_row(router, shadow_params=shadow_params) + + judge_prompt = next( + call.kwargs["messages"][-1]["content"] + for call in router.acompletion.call_args_list + if call.kwargs["metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) != SHADOW_EVAL_ROUTER_CALL_ORIGIN + ) + + assert "Tools available" not in judge_prompt + + async def test_a_custom_tool_call_serializes_its_name_and_input(self): + """Custom tool calls carry no `function` key: name and arguments live under + `custom`, so reading only `function` serializes every one of them as unnamed.""" + router = _shadow_reply_router(CUSTOM_TOOL_CALL_MESSAGE, finish_reason="tool_calls") + await self._judged_shadow_row(router) + + judge_prompt = next( + call.kwargs["messages"][-1]["content"] + for call in router.acompletion.call_args_list + if call.kwargs["metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) != SHADOW_EVAL_ROUTER_CALL_ORIGIN + ) + + assert "[tool call] exec_sql(select 1)" in judge_prompt + + async def test_the_judge_is_told_a_tool_call_is_not_a_defect(self): + """The judge scores on completeness and clarity. Handed a tool call with no + instruction, it marks it down for not reading like an answer, which would bias + every verdict against a tool-calling arm on exactly the traffic that calls tools.""" + router = _shadow_reply_router(TOOL_CALL_MESSAGE, finish_reason="tool_calls") + await self._judged_shadow_row(router) + + system_prompt = next( + call.kwargs["messages"][0]["content"] + for call in router.acompletion.call_args_list + if call.kwargs["metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) != SHADOW_EVAL_ROUTER_CALL_ORIGIN + ) + + assert "tool call" in system_prompt + assert "not a defect" in system_prompt + + async def test_prose_written_alongside_a_tool_call_survives_into_the_verdict(self): + """Some providers write a sentence before acting. Serializing only the call would + hide half of what the arm actually said from the judge.""" + router = _shadow_reply_router( + {"content": "Let me look that up.", "tool_calls": TOOL_CALL_MESSAGE["tool_calls"]}, + finish_reason="tool_calls", + ) + await self._judged_shadow_row(router) + + judge_prompt = next( + call.kwargs["messages"][-1]["content"] + for call in router.acompletion.call_args_list + if call.kwargs["metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) != SHADOW_EVAL_ROUTER_CALL_ORIGIN + ) + + assert "Let me look that up. [tool call] Read({})" in judge_prompt + + async def test_an_empty_shadow_reply_names_the_finish_reason_and_the_routed_model(self): + """A reply that really carried no text is diagnosable only if the row says what + the arm was doing when it produced none: a truncated turn and a model that answers + with nothing are different faults with different fixes.""" + error = await self._no_text_error( + _shadow_reply_router({"content": ""}, finish_reason="length", routed_model="some-model") + ) + + assert "empty response" in error + assert "finish_reason=length" in error + assert "model=some-model" in error + + async def test_no_text_errors_stay_groupable_across_models_and_finish_reasons(self): + """Operators read these rows by grouping on the error text, which is how a job's + failures collapse to a handful of causes. Every varying part therefore has to sit + behind the first semicolon, or each row becomes its own group and the count that + made the problem visible stops existing.""" + first = await self._no_text_error( + _shadow_reply_router({"content": None}, finish_reason="length", routed_model="model-a") + ) + second = await self._no_text_error( + _shadow_reply_router( + {"content": ""}, + finish_reason="stop", + routed_model="model-b", + ) + ) + + assert first != second + assert first.split(";")[0] == second.split(";")[0] + async def test_a_pipeline_error_after_the_shadow_call_keeps_its_billed_cost(self, monkeypatch: pytest.MonkeyPatch): """An unexpected error between the billed shadow call and the attempt write must still record the shadow cost, or the per-key dollar gate undercounts forever.""" @@ -1691,11 +1968,13 @@ class TestSamplingFunnel: prisma.db.litellm_shadowevalattempt.create.assert_not_awaited() async def test_an_unjudgeable_sampled_request_counts_unjudgeable(self): + """A tool call still serializes into judgeable text; a turn with neither prose nor + a tool call to serialize is the one case left with nothing to compare.""" prisma = _prisma() logger = _logger(router=_router(), prisma=prisma, jobs=(_job(),)) - tool_final = {"choices": [{"message": {"content": None, "tool_calls": [{"type": "function", "function": {}}]}}]} + empty = {"choices": [{"message": {"content": None}}]} - await logger.async_log_success_event(_success_kwargs(), tool_final, None, None) + await logger.async_log_success_event(_success_kwargs(), empty, None, None) await _drain(logger) assert logger._test_funnel == [("job-1", "unjudgeable")]