Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_add_milvus_grpc_transport

# Conflicts:
#	litellm/proxy/common_utils/http_parsing_utils.py
This commit is contained in:
Yujong Lee 2026-09-04 18:46:26 -07:00
commit e9a014090a
13 changed files with 818 additions and 104 deletions

View file

@ -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:

View file

@ -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]

View file

@ -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,17 +44,19 @@ def _is_json_content_type(content_type: str) -> bool:
return _normalize_media_type(content_type) == "application/json"
def _type_args(annotation: object) -> tuple[object, ...]:
return tuple(get_args(annotation))
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: object = annotation # rebind-ok: type qualifiers may be nested to arbitrary depth
while get_origin(unwrapped) in (Annotated, NotRequired, ReadOnly, Required):
unwrapped = _type_args(unwrapped)[0] # rebind-ok: peel one qualifier per iteration
unwrapped: Final = _unqualified(annotation)
candidates: Final = (
tuple(arg for arg in _type_args(unwrapped) if arg is not type(None))
tuple(arg for arg in get_args(unwrapped) if arg is not type(None))
if get_origin(unwrapped) in (Union, UnionType)
else (unwrapped,)
)

View file

@ -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: 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)
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,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),
),
rollover_caps=rollover_caps,
cache_keys=(
@ -682,6 +700,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)),
),
)

View file

@ -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:"):

View file

@ -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
@ -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
@ -2580,6 +2581,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 +2616,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

View file

@ -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).
]

View file

@ -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")]

View file

@ -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):

View file

@ -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, include_extras=True))) == {
"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):

View file

@ -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
@ -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: Final = _make_counter_invalidation_job(monkeypatch)
budget: Final = _budget_row(budget_id="budget-1")
mock_prisma_client.data["budget"] = [budget]
test_enduser: Final = 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: 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
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)
@ -3028,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."""

View file

@ -9,6 +9,7 @@ from __future__ import annotations
from datetime import datetime, timedelta, timezone
from types import SimpleNamespace
from typing import Final
import pytest
@ -18,7 +19,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 +48,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 +252,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: Final = _FakePrismaClient(end_user_row=SimpleNamespace(user_id="customer-42", spend=0.0))
result: Final = 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: 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")
== 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: 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 == []
@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: 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 == []

View file

@ -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
@ -222,16 +223,19 @@ 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."""
fake_cache = _make_spend_counter_cache(redis_get_value=2.0)
@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: 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))
result = await ps.get_current_spend(
counter_key="spend:end_user:e1",
result: Final = await ps.get_current_spend(
counter_key=counter_key,
fallback_spend=20.0,
max_budget=10.0,
)
@ -241,6 +245,72 @@ async def test_get_current_spend_floors_end_user_tag_against_fallback(monkeypatc
fake_cache.redis_cache.async_set_max.assert_not_called()
def _make_prisma_with_end_user_row(spend: float | None):
prisma: Final = 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: Final = _make_spend_counter_cache(redis_get_value=0.0)
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
prisma: Final = _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: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: 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: Final = 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: 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: Final = await ps.get_current_spend(
counter_key="spend:end_user:customer-42",
fallback_spend=20.0,
max_budget=10.0,
)
assert result == 20.0
fake_cache.redis_cache.async_set_max.assert_not_called()
@pytest.mark.asyncio
async def test_get_current_spend_floors_window_against_spend_logs(monkeypatch):
"""Per-window counters have no DB row but aggregate from spend logs. A