merge: main into litellm_hotfix_7072_team_callbacks
Some checks failed
ai-gateway image / ai-gateway release image (push) Has been cancelled
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-15 13:37:45 +00:00
commit acba499e47
42 changed files with 3778 additions and 695 deletions

View file

@ -538,7 +538,7 @@ context_window_fallbacks: Optional[List] = None
content_policy_fallbacks: Optional[List] = None
allowed_fails: int = 3
allow_dynamic_callback_disabling: bool = True
num_retries_per_request: Optional[int] = None # cap on Router retries of one model group; resets per fallback hop
num_retries_per_request: Optional[int] = None # for the request overall (incl. fallbacks + model retries)
####### SECRET MANAGERS #####################
secret_manager_client: Optional[Any] = (
None # list of instantiated key management clients - e.g. azure kv, infisical, etc.

View file

@ -309,8 +309,8 @@ def max_retries_per_request_hit(kwargs: Mapping[str, object], num_retries_per_re
metadata: Final = kwargs.get(get_metadata_variable_name_from_kwargs(kwargs))
if not isinstance(metadata, Mapping):
return False
attempted_retries: Final = metadata.get("attempted_retries")
return type(attempted_retries) is int and 0 < attempted_retries and num_retries_per_request <= attempted_retries
retry_count: Final = metadata.get("request_retry_count")
return type(retry_count) is int and 0 < retry_count and num_retries_per_request <= retry_count
def get_or_create_metadata_bucket(

View file

@ -20,6 +20,7 @@ from itertools import chain, repeat
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Protocol, cast, overload, runtime_checkable
from pydantic import TypeAdapter, ValidationError
from typing_extensions import ReadOnly, TypedDict, assert_never
from litellm._logging import verbose_proxy_logger
@ -104,9 +105,24 @@ class ToolResultBlockTextTarget:
block_idx: int
InputWriteBackTarget = (
MessageContentTarget | ContentBlockTextTarget | ToolResultStringTarget | ToolResultBlockTextTarget
)
@dataclass(frozen=True, slots=True)
class SystemStringTarget:
pass
@dataclass(frozen=True, slots=True)
class SystemBlockTextTarget:
block_idx: int
@dataclass(frozen=True, slots=True)
class ToolUseInputTarget:
msg_idx: int
content_idx: int
MessageTextTarget = MessageContentTarget | ContentBlockTextTarget | ToolResultStringTarget | ToolResultBlockTextTarget
InputWriteBackTarget = SystemStringTarget | SystemBlockTextTarget | MessageTextTarget
def _as_str_mapping(value: Mapping[str, object]) -> Mapping[str, object]:
@ -147,10 +163,17 @@ class ScannedText:
target: InputWriteBackTarget
@dataclass(frozen=True, slots=True)
class ScannedToolCall:
tool_call: ChatCompletionToolCallChunk
target: ToolUseInputTarget
@dataclass(frozen=True, slots=True)
class ExtractedInput:
scanned: tuple[ScannedText, ...]
images: tuple[str, ...]
tool_calls: tuple[ScannedToolCall, ...] = ()
EMPTY_EXTRACTED_INPUT: Final = ExtractedInput(scanned=(), images=())
@ -162,6 +185,74 @@ class _ToolCallShape:
arguments: str
def _is_client_tool_use(block: Mapping[str, object]) -> bool:
return (
block.get("type") == "tool_use"
and isinstance(block.get("id"), str)
and isinstance(block.get("name"), str)
and isinstance(block.get("input"), dict)
)
def _write_back_system_block(system: object, block_idx: int, response: str) -> None:
if not isinstance(system, list):
return
text_blocks: Final = tuple(block for block in system if isinstance(block, dict) and block.get("type") == "text")
if block_idx < len(text_blocks):
text_blocks[block_idx]["text"] = (
response # mutable-ok: guardrails rewrite the caller's request payload in place
)
def _write_back_message_text(message: _WritableMessage, target: MessageTextTarget, response: str) -> None:
content: Final = message.get("content", None)
if content is None:
return
match target:
case MessageContentTarget():
if isinstance(content, str):
message["content"] = response # mutable-ok: guardrails rewrite the caller's request payload in place
case ContentBlockTextTarget(content_idx=content_idx):
if isinstance(content, list):
content[content_idx]["text"] = (
response # mutable-ok: guardrails rewrite the caller's request payload in place
)
case ToolResultStringTarget(content_idx=content_idx):
if isinstance(content, list):
content[content_idx]["content"] = (
response # mutable-ok: guardrails rewrite the caller's request payload in place
)
case ToolResultBlockTextTarget(content_idx=content_idx, block_idx=block_idx):
if isinstance(content, list):
content[content_idx]["content"][block_idx]["text"] = (
response # mutable-ok: guardrails rewrite the caller's request payload in place
)
case _:
assert_never(target)
_TOOL_USE_INPUT_ADAPTER: Final = TypeAdapter(dict[str, object])
def _rewritten_tool_use_input(arguments: str) -> Mapping[str, object] | None:
try:
return _TOOL_USE_INPUT_ADAPTER.validate_json(arguments)
except ValidationError:
return None
def _write_back_tool_use(
message: _WritableMessage, target: ToolUseInputTarget, shape: _ToolCallShape, rewritten_input: Mapping[str, object]
) -> None:
content: Final = message.get("content", None)
block: Final = content[target.content_idx] if isinstance(content, list) else None
if not isinstance(block, dict):
return
block["input"] = rewritten_input # mutable-ok: guardrails rewrite the caller's request payload in place
if shape.name is not None and shape.name != block.get("name"):
block["name"] = shape.name # mutable-ok: guardrails rewrite the caller's request payload in place
@dataclass(frozen=True, slots=True)
class _SSEFieldRewrite:
"""One field of one nested section of a buffered SSE event, rewritten."""
@ -453,9 +544,8 @@ class AnthropicMessagesHandler(BaseTranslation):
skip_tool: Final = effective_skip_tool_message_for_guardrail(guardrail_to_apply)
scan_only_tool_results: Final = effective_scan_only_tool_results_for_guardrail(guardrail_to_apply)
# Exclude only the trusted top-level prompt. In-sequence system entries are untrusted
# and must stay aligned with texts_to_check for positional masking. When the top-level
# prompt is included, the pre-existing count mismatch disables positional masking.
# The top-level prompt is translated on its own below so it can be hoisted in front of
# any mid-turn system entries and scanned first, aligned with that structured position.
translation_source: Final = { # mutable-ok: API message payload
key: value for key, value in data.items() if key != "system"
}
@ -491,7 +581,12 @@ class AnthropicMessagesHandler(BaseTranslation):
]
)
# Step 1: Extract all text content and images
# Step 1: Extract all text content, images, and tool calls
top_level_system_scanned: Final = (
()
if hoisted_system_message is None or scan_only_tool_results
else self._extract_top_level_system_text(hoisted_system_message)
)
extracted: Final = tuple(
self._extract_input_text_and_images(
message=message,
@ -502,17 +597,27 @@ class AnthropicMessagesHandler(BaseTranslation):
)
for msg_idx, message in enumerate(messages)
)
scanned: Final = tuple(item for one_message in extracted for item in one_message.scanned)
scanned: Final = (
*top_level_system_scanned,
*(item for one_message in extracted for item in one_message.scanned),
)
texts_to_check: Final = [item.text for item in scanned] # mutable-ok: GenericGuardrailAPIInputs takes list[str]
images_to_check: Final = [
image for one_message in extracted for image in one_message.images
] # mutable-ok: GenericGuardrailAPIInputs takes list[str]
scanned_tool_calls: Final = tuple(item for one_message in extracted for item in one_message.tool_calls)
tool_calls_to_check: Final = [
item.tool_call for item in scanned_tool_calls
] # mutable-ok: GenericGuardrailAPIInputs takes list[ChatCompletionToolCallChunk]
pre_guardrail_tool_calls: Final = _tool_call_shapes(tool_calls_to_check)
# Step 2: Apply guardrail to all texts in batch
if texts_to_check:
# Step 2: Apply guardrail to all texts and tool calls in batch
if texts_to_check or tool_calls_to_check:
inputs: Final = GenericGuardrailAPIInputs(texts=texts_to_check)
if images_to_check:
inputs["images"] = images_to_check
if tool_calls_to_check:
inputs["tool_calls"] = tool_calls_to_check
if tools_to_check:
inputs["tools"] = tools_to_check
original_structured_messages: Final = structured_messages
@ -573,9 +678,16 @@ class AnthropicMessagesHandler(BaseTranslation):
else:
if guardrailed_texts and len(guardrailed_texts) != len(scanned):
raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name)
self._apply_guardrail_tool_calls_to_input(
messages=messages,
scanned_tool_calls=scanned_tool_calls,
pre_guardrail_tool_calls=pre_guardrail_tool_calls,
returned_tool_calls=guardrailed_inputs.get("tool_calls"),
guardrail_name=guardrail_to_apply.guardrail_name,
)
# Step 3: Map guardrail responses back to original message structure
await self._apply_guardrail_responses_to_input(
messages=messages,
data=data,
responses=guardrailed_texts,
scanned=scanned,
)
@ -601,6 +713,19 @@ class AnthropicMessagesHandler(BaseTranslation):
hoisted: Final = probe.get("messages") or [] # mutable-ok: API message payload
return hoisted[0] if hoisted else None
@staticmethod
def _extract_top_level_system_text(hoisted_system_message: AllMessageValues) -> tuple[ScannedText, ...]:
content: Final = hoisted_system_message.get("content")
if isinstance(content, str):
return (ScannedText(content, SystemStringTarget()),)
if not isinstance(content, list):
return ()
return tuple(
ScannedText(text_str, SystemBlockTextTarget(block_idx))
for block_idx, block in enumerate(content)
if isinstance(block, dict) and isinstance(text_str := block.get("text"), str)
)
@staticmethod
def _openai_system_message_to_anthropic(
message: Mapping[str, object],
@ -855,9 +980,25 @@ class AnthropicMessagesHandler(BaseTranslation):
for content_idx, content_item in enumerate(content)
if isinstance(content_item, dict)
)
tool_use_blocks: Final = (
()
if scan_only_tool_results
else tuple(
(content_idx, content_item)
for content_idx, content_item in enumerate(content)
if isinstance(content_item, dict) and _is_client_tool_use(content_item)
)
)
return ExtractedInput(
scanned=tuple(item for block in blocks for item in block.scanned),
images=tuple(image for block in blocks for image in block.images),
tool_calls=tuple(
ScannedToolCall(
tool_call=AnthropicConfig.convert_tool_use_to_openai_format(content_item, tool_call_idx),
target=ToolUseInputTarget(msg_idx, content_idx),
)
for tool_call_idx, (content_idx, content_item) in enumerate(tool_use_blocks)
),
)
@classmethod
@ -943,43 +1084,59 @@ class AnthropicMessagesHandler(BaseTranslation):
async def _apply_guardrail_responses_to_input(
self,
messages: Sequence[_WritableMessage],
responses: list[str],
data: dict[str, object], # mutable-ok: API message payload
responses: Sequence[str],
scanned: tuple[ScannedText, ...],
) -> None:
"""
Apply guardrail responses back to input messages.
Apply guardrail responses back to the top-level system prompt and the input messages.
"""
raw_messages: Final = data.get("messages")
messages: Final[Sequence[_WritableMessage]] = raw_messages if isinstance(raw_messages, list) else ()
for item, guardrail_response in zip(scanned, responses):
target = item.target
message = messages[target.msg_idx]
content = message.get("content", None)
if content is None:
continue
match target:
case MessageContentTarget():
if isinstance(content, str):
message["content"] = (
guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place
)
case ContentBlockTextTarget(content_idx=content_idx):
if isinstance(content, list):
content[content_idx]["text"] = (
guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place
)
case ToolResultStringTarget(content_idx=content_idx):
if isinstance(content, list):
content[content_idx]["content"] = (
guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place
)
case ToolResultBlockTextTarget(content_idx=content_idx, block_idx=block_idx):
if isinstance(content, list):
content[content_idx]["content"][block_idx]["text"] = (
match item.target:
case SystemStringTarget():
if isinstance(data.get("system"), str):
data["system"] = (
guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place
)
case SystemBlockTextTarget(block_idx=block_idx):
_write_back_system_block(data.get("system"), block_idx, guardrail_response)
case (
MessageContentTarget()
| ContentBlockTextTarget()
| ToolResultStringTarget()
| ToolResultBlockTextTarget() as message_target
):
_write_back_message_text(messages[message_target.msg_idx], message_target, guardrail_response)
case _:
assert_never(target)
assert_never(item.target)
@staticmethod
def _apply_guardrail_tool_calls_to_input(
messages: Sequence[_WritableMessage],
scanned_tool_calls: tuple[ScannedToolCall, ...],
pre_guardrail_tool_calls: tuple[_ToolCallShape, ...],
returned_tool_calls: Sequence[object] | None,
guardrail_name: str | None,
) -> None:
post_guardrail_tool_calls: Final = _tool_call_shapes(
returned_tool_calls
if returned_tool_calls is not None and len(returned_tool_calls) == len(pre_guardrail_tool_calls)
else tuple(item.tool_call for item in scanned_tool_calls)
)
rewritten: Final = tuple(
(item, after, _rewritten_tool_use_input(after.arguments))
for item, before, after in zip(scanned_tool_calls, pre_guardrail_tool_calls, post_guardrail_tool_calls)
if before != after
)
applicable: Final = tuple(
(item, after, rewritten_input) for item, after, rewritten_input in rewritten if rewritten_input is not None
)
if len(applicable) != len(rewritten):
raise unappliable_request_rewrite(guardrail_name)
for item, after, rewritten_input in applicable:
_write_back_tool_use(messages[item.target.msg_idx], item.target, after, rewritten_input)
async def process_output_response(
self,

View file

@ -11,6 +11,7 @@ from concurrent.futures import ThreadPoolExecutor
from datetime import datetime
from functools import partial
from threading import Lock
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, ParamSpec, TypeVar, cast, get_args, overload
import httpx
@ -96,6 +97,77 @@ def _assume_role_params(
)
_SecureTransportBool = TypedDict("_SecureTransportBool", {"aws:SecureTransport": ReadOnly[Literal["true"]]})
class _SecureTransportCondition(TypedDict):
Bool: ReadOnly[_SecureTransportBool]
class _SessionPolicyStatement(TypedDict):
Sid: ReadOnly[str]
Effect: ReadOnly[Literal["Allow"]]
Action: ReadOnly[tuple[str, ...]]
Resource: ReadOnly[Literal["*"]]
Condition: ReadOnly[_SecureTransportCondition]
class WebIdentitySessionPolicy(TypedDict):
Version: ReadOnly[Literal["2012-10-17"]]
Statement: ReadOnly[tuple[_SessionPolicyStatement, ...]]
_WEB_IDENTITY_SESSION_POLICY_ACTIONS: Final[Mapping[str, tuple[str, ...]]] = MappingProxyType(
{
"BedrockLiteLLM": (
"bedrock:InvokeModel",
"bedrock:InvokeModelWithResponseStream",
"bedrock:CountTokens",
"bedrock:Rerank",
"bedrock:Retrieve",
"bedrock:ListKnowledgeBases",
"bedrock:InvokeAgent",
"bedrock:ApplyGuardrail",
"bedrock:GetGuardrail",
"bedrock:ListGuardrails",
),
"BedrockAgentCoreLiteLLM": (
"bedrock-agentcore:InvokeAgentRuntime",
"bedrock-agentcore:InvokeAgentRuntimeForUser",
"bedrock-agentcore:InvokeGateway",
),
"ClaudePlatformLiteLLM": (
"aws-external-anthropic:CreateInference",
"aws-external-anthropic:CreateBatchInference",
"aws-external-anthropic:CancelBatchInference",
"aws-external-anthropic:DeleteBatchInference",
"aws-external-anthropic:CountTokens",
"aws-external-anthropic:Get*",
"aws-external-anthropic:List*",
),
"BedrockMantleLiteLLM": ("bedrock-mantle:CreateInference",),
}
)
_SECURE_TRANSPORT_ONLY: Final = _SecureTransportCondition(Bool=_SecureTransportBool({"aws:SecureTransport": "true"}))
def build_web_identity_session_policy() -> WebIdentitySessionPolicy:
return WebIdentitySessionPolicy(
Version="2012-10-17",
Statement=tuple(
_SessionPolicyStatement(
Sid=sid,
Effect="Allow",
Action=actions,
Resource="*",
Condition=_SECURE_TRANSPORT_ONLY,
)
for sid, actions in _WEB_IDENTITY_SESSION_POLICY_ACTIONS.items()
),
)
class BedrockRequestTarget(BaseModel):
aws_region_name: str
aws_bedrock_runtime_endpoint: str | None
@ -940,60 +1012,12 @@ class BaseAWSLLM(SignsRequestsWithAWS):
# auth only (static creds + IRSA take other code paths).
# https://docs.aws.amazon.com/STS/latest/APIReference/API_AssumeRoleWithWebIdentity.html
# https://boto3.amazonaws.com/v1/documentation/api/latest/reference/services/sts/client/assume_role_with_web_identity.html
bedrock_session_policy: Final = {
"Version": "2012-10-17",
"Statement": [
{
"Sid": "BedrockLiteLLM",
"Effect": "Allow",
"Action": [
"bedrock:InvokeModel",
"bedrock:InvokeModelWithResponseStream",
"bedrock:CountTokens",
"bedrock:ApplyGuardrail",
"bedrock:GetGuardrail",
"bedrock:ListGuardrails",
],
"Resource": "*",
"Condition": {"Bool": {"aws:SecureTransport": "true"}},
},
# Claude Platform on AWS (added by #27678 for the
# ``bedrock/claude_platform/<model>`` route) lives under
# a separate IAM action namespace; without these entries
# the OIDC path 403s on every claude_platform request
# even with a fully permissive identity policy (#30200).
{
"Sid": "ClaudePlatformLiteLLM",
"Effect": "Allow",
"Action": [
"aws-external-anthropic:CreateInference",
"aws-external-anthropic:CreateBatchInference",
"aws-external-anthropic:CancelBatchInference",
"aws-external-anthropic:DeleteBatchInference",
"aws-external-anthropic:CountTokens",
"aws-external-anthropic:Get*",
"aws-external-anthropic:List*",
],
"Resource": "*",
"Condition": {"Bool": {"aws:SecureTransport": "true"}},
},
{
"Sid": "BedrockMantleLiteLLM",
"Effect": "Allow",
"Action": [
"bedrock-mantle:CreateInference",
],
"Resource": "*",
"Condition": {"Bool": {"aws:SecureTransport": "true"}},
},
],
}
assume_role_params: Final = {
"RoleArn": aws_role_name,
"RoleSessionName": aws_session_name,
"WebIdentityToken": oidc_token,
"DurationSeconds": 3600,
"Policy": json.dumps(bedrock_session_policy, separators=(",", ":")),
"Policy": json.dumps(build_web_identity_session_policy(), separators=(",", ":")),
}
# Add ExternalId parameter if provided

View file

@ -115,6 +115,19 @@ def _parse_setup(session_configuration_request: str) -> BidiGenerateContentSetup
return envelope.get("setup", empty_setup)
def _grounding_metadata_from_frame(frame: Mapping[str, object]) -> tuple[Mapping[str, object], ...]:
"""Read ``serverContent.groundingMetadata`` off the frame that carries the turn's usage.
Live reports grounding in the server frames rather than in ``usageMetadata``, and it emits both
on the same frame, so the per-query charge is countable at the point usage is built.
"""
server_content: Final = frame.get("serverContent")
if not isinstance(server_content, Mapping):
return ()
metadata: Final = server_content.get("groundingMetadata")
return (metadata,) if isinstance(metadata, Mapping) else ()
# Google bills Live transcription at an estimated 25 audio tokens/sec of input and
# 175 text tokens/min of output (ai.google.dev/gemini-api/docs/pricing).
GEMINI_LIVE_TRANSCRIBE_AUDIO_TOKENS_PER_SECOND: Final = 25
@ -323,7 +336,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
)
elif key == "input_audio_transcription" and value is not None:
optional_params["inputAudioTranscription"] = {}
elif key == "turn_detection":
elif key == "turn_detection" and value is not None:
value_typed = cast(OpenAIRealtimeTurnDetection, value)
if (
isinstance(value_typed, dict)
@ -1049,6 +1062,11 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
{**cast(dict, message), "usageMetadata": resolved_usage_metadata},
),
)
grounding_metadata: Final = _grounding_metadata_from_frame(message)
if grounding_metadata:
VertexGeminiConfig._set_grounding_usage_counters( # pyright: ignore[reportPrivateUsage] # shared with the chat path; no public alias exists yet
_chat_completion_usage, grounding_metadata
)
else:
_chat_completion_usage = get_empty_usage()

View file

@ -580,8 +580,8 @@ What the command changed is recorded in `~/.litellm/claude_configure_state.json`
`lite configure claude`, `lite login --config-claude`, `lite up` and `lite autoroute up` also install a status line (`~/.litellm/statusline.py`, registered as `statusLine` in `~/.claude/settings.json` unless you already run one) that shows which model the auto-router actually served the last turn and, once the proxy has recorded the session, what the session cost against the router's savings baseline:
```
claude-auto · Routed to: claude-haiku-4-5 -63% vs Claude Opus 5
LiteLLM ████████░░░░░░░░░░░░░░░░ $0.14
Routed to: claude-haiku-4-5 -63% vs Claude Opus 5
claude-auto ████████░░░░░░░░░░░░░░░░ $0.14
Claude Opus 5 ████████████████████████ $0.38
```

View file

@ -28,6 +28,7 @@ import os
import sys
import tempfile
import time
import unicodedata
import urllib.error
import urllib.request
from collections.abc import Callable, Mapping
@ -42,7 +43,6 @@ FETCH_TIMEOUT_SECONDS: Final = 3
BAR_WIDTH: Final = 24
BAR_FULL: Final = "\u2588"
BAR_EMPTY: Final = "\u2591"
SEPARATOR: Final = " \u00b7 "
TRANSCRIPT_SCAN_LIMIT_BYTES: Final = 4 * 1024 * 1024
CLAUDE_BASE_URL_ENV_KEYS: Final = ("ANTHROPIC_BASE_URL",)
CLAUDE_API_KEY_ENV_KEYS: Final = ("ANTHROPIC_AUTH_TOKEN", "ANTHROPIC_API_KEY")
@ -50,7 +50,6 @@ CODEX_BASE_URL_ENV_KEYS: Final = ("OPENAI_BASE_URL",)
CODEX_API_KEY_ENV_KEYS: Final = ("OPENAI_API_KEY",)
CODEX_STOP_EVENT: Final = "Stop"
SYNTHETIC_MODEL: Final = "<synthetic>"
LITELLM_LABEL: Final = "LiteLLM"
RESET: Final = "\033[0m"
BOLD: Final = "\033[1m"
DIM: Final = "\033[90m"
@ -302,31 +301,37 @@ def _bar(fraction: float, color: str, width: int, use_color: bool) -> str:
return f"{color}{BAR_FULL * filled}{DIM}{BAR_EMPTY * (width - filled)}{RESET}"
def _display_width(label: str) -> int:
return sum(
2 if unicodedata.east_asian_width(character) in ("W", "F") else 1
for character in label
if unicodedata.category(character) not in ("Mn", "Me")
)
def render(model: str, session: Session | None, config_dir: Path, use_color: bool, bar_width: int = BAR_WIDTH) -> str:
def paint(code: str, text: str) -> str:
return f"{code}{text}{RESET}" if use_color else text
routed: Final = paint(BOLD, f"Routed to: {model}")
if session is None:
if session is None or session.baseline_model is None or session.baseline_spend <= 0:
return routed
header: Final = f"{session.router_name}{SEPARATOR}{routed}"
if session.baseline_model is None or session.baseline_spend <= 0:
return header
reference: Final = baseline_label(session.baseline_model, config_dir)
pct: Final = (session.baseline_spend - session.spend) / session.baseline_spend * 100
delta: Final = paint(LITELLM_COLOR, f"{'-' if pct >= 0 else '+'}{abs(round(pct))}% vs {reference}")
peak: Final = max(session.spend, session.baseline_spend)
label_width: Final = max(len(LITELLM_LABEL), len(reference))
label_width: Final = max(_display_width(session.router_name), _display_width(reference))
rows: Final = (
(LITELLM_LABEL, session.spend, LITELLM_COLOR),
(session.router_name, session.spend, LITELLM_COLOR),
(reference, session.baseline_spend, BASELINE_COLOR),
)
lines: Final = (
f"{paint(DIM, label.ljust(label_width))} {_bar(amount / peak, color, bar_width, use_color)} "
f"{paint(DIM, label + ' ' * (label_width - _display_width(label)))} "
f"{_bar(amount / peak, color, bar_width, use_color)} "
f"{paint(DIM, f'${amount:.2f}')}"
for label, amount, color in rows
)
return "\n".join((f"{header} {delta}", *lines))
return "\n".join((f"{routed} {delta}", *lines))
def color_enabled(env: Mapping[str, str]) -> bool:

View file

@ -1600,8 +1600,8 @@ class PanwPrismaAirsHandler(CustomGuardrail):
Args:
texts: Flattened text entries from the framework.
messages: Original request messages (request_data["messages"]),
NOT structured_messages (which may have injected system content).
messages: The structured messages the framework flattened into ``texts``,
hoisted top-level system prompt included, so positions line up.
Returns a set of scannable indices, or None on count mismatch or no user/developer
message (safety fallback to existing role-filter behavior).
@ -1788,15 +1788,10 @@ class PanwPrismaAirsHandler(CustomGuardrail):
structured_messages: Final = inputs.get("structured_messages")
if structured_messages:
# For Anthropic /v1/messages: default to latest-user-only scanning.
# Uses request_data["messages"] (original format), NOT structured_messages
# (which has injected system content from adapter translation).
if self._use_latest_user_only(request_data, logging_obj):
original_messages: Final = request_data.get("messages")
if original_messages:
scannable_indices = self._get_latest_user_text_indices(texts, original_messages)
scannable_indices = self._get_latest_user_text_indices(texts, structured_messages)
# Fall through to existing role filtering if:
# - not Anthropic, OR flag explicitly False, OR
# - no original messages, OR
# - latest-user extraction returned None (no user / count mismatch)
if scannable_indices is None:
scannable_indices = self._get_scannable_text_indices(texts, structured_messages)

View file

@ -336,7 +336,7 @@ _CLIENT_PRICING_METADATA_FIELDS: Final = frozenset({"model_info", "standard_logg
# and read by spend logs as fact; a client value has no legitimate meaning and no
# key or team setting keeps it, so the strip is never gated.
_ROUTER_RESERVED_METADATA_FIELDS: Final = frozenset(
{"attempted_fallbacks", "original_model_group", CLIENT_OUTPUT_CEILING_METADATA_KEY}
{"attempted_fallbacks", "original_model_group", "request_retry_count", CLIENT_OUTPUT_CEILING_METADATA_KEY}
)
_ALLOW_CLIENT_PRICING_OVERRIDE_METADATA_KEY: Final = "allow_client_pricing_override"

View file

@ -147,6 +147,7 @@ from litellm.types.proxy.management_endpoints.key_management_endpoints import (
BulkUpdateKeyRequest,
BulkUpdateKeyResponse,
BulkUpdateTeamKeysRequest,
CustomKeyPolicyRequest,
FailedKeyUpdate,
KeySearchWhere,
SuccessfulKeyUpdate,
@ -285,6 +286,7 @@ def _config_table(prisma_client: PrismaClient) -> _ConfigTableActions:
class _CustomKeyHooksModule(Protocol):
user_custom_key_generate: Callable[..., Awaitable[Mapping[str, object]]] | None
user_custom_key_update: Callable[..., Awaitable[Mapping[str, object]]] | None
user_custom_key_policy: Callable[..., Awaitable[Mapping[str, object]]] | None
def _custom_key_generate_hook(
@ -299,6 +301,161 @@ def _custom_key_update_hook(
return hooks.user_custom_key_update
def _custom_key_policy_hook(
hooks: _CustomKeyHooksModule,
) -> Callable[..., Awaitable[Mapping[str, object]]] | None:
return hooks.user_custom_key_policy
async def _enforce_custom_key_update_policy(
hook: Callable[..., Awaitable[Mapping[str, object]]] | None,
data: UpdateKeyRequest,
) -> None:
if hook is None:
return
if not inspect.iscoroutinefunction(hook):
raise ValueError("user_custom_key_update must be a coroutine")
result: Final = await hook(data)
if not result.get("decision", True):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=result.get("message", "Authentication Failed - Custom Auth Rule"),
)
async def _enforce_custom_key_policy(
hook: Callable[..., Awaitable[Mapping[str, object]]] | None,
build_policy_request: Callable[[], CustomKeyPolicyRequest],
) -> None:
if hook is None:
return
if not inspect.iscoroutinefunction(hook):
raise ValueError("user_custom_key_policy must be a coroutine")
result: Final = await hook(build_policy_request())
if not result.get("decision", True):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=result.get("message", "Authentication Failed - Custom Auth Rule"),
)
_KEY_UPDATE_JSON_STRING_COLUMNS: Final = frozenset({"router_settings", "budget_limits"})
_KEY_METADATA_REQUEST_FIELDS: Final = frozenset(
(*LiteLLM_ManagementEndpoint_MetadataFields_Premium, *LiteLLM_ManagementEndpoint_MetadataFields)
)
def _decode_json_string_column(column: str, value: object) -> object:
if column in _KEY_UPDATE_JSON_STRING_COLUMNS and isinstance(value, str):
return json.loads(value)
return value
def _verification_token_from_row(row: Mapping[str, object]) -> LiteLLM_VerificationToken:
org_id: Final = row["organization_id"] if "organization_id" in row else row.get("org_id")
return LiteLLM_VerificationToken.model_validate(MappingProxyType({**row, "org_id": org_id}))
def _effective_key_after_update(
existing_key_row: LiteLLM_VerificationToken,
non_default_values: Mapping[str, object],
) -> LiteLLM_VerificationToken:
overlay: Final = MappingProxyType(
{column: _decode_json_string_column(column, value) for column, value in non_default_values.items()}
)
return _verification_token_from_row(
MappingProxyType({**existing_key_row.model_dump(), **overlay, "object_permission": None})
)
def _update_policy_request(
operation: Literal["update", "regenerate"],
existing_key_row: LiteLLM_VerificationToken,
non_default_values: Mapping[str, object],
request: UpdateKeyRequest | RegenerateKeyRequest,
) -> CustomKeyPolicyRequest:
return CustomKeyPolicyRequest(
operation=operation,
existing_key=_verification_token_from_row(existing_key_row.model_dump()),
effective_key=_effective_key_after_update(
existing_key_row=existing_key_row, non_default_values=non_default_values
),
request=request,
)
def _generate_budget_windows(
budget_limits: Sequence[BudgetLimitEntry] | None,
) -> tuple[Mapping[str, object], ...] | None:
if not budget_limits:
return None
return tuple(
MappingProxyType(
{
**window.model_dump(),
"reset_at": get_budget_reset_time(budget_duration=window.budget_duration).isoformat(),
}
)
for window in budget_limits
)
def _effective_key_for_generate(data: GenerateKeyRequest, now: datetime) -> LiteLLM_VerificationToken:
requested: Final = data.model_dump(exclude_unset=True, exclude_none=True)
metadata_fields: Final = MappingProxyType(
{field: value for field, value in requested.items() if field in _KEY_METADATA_REQUEST_FIELDS}
)
column_fields: Final = MappingProxyType(
{field: value for field, value in requested.items() if field not in _KEY_METADATA_REQUEST_FIELDS}
)
metadata: Final = data.metadata or MappingProxyType({})
folded_metadata: Final = {**metadata, **metadata_fields} # mutable-ok: encrypt_callback_vars needs a dict
columns: Final = handle_key_type(data, {**column_fields}) # mutable-ok: handle_key_type mutates in place
expires: Final = (
now + timedelta(seconds=duration_in_seconds(duration=data.duration)) if data.duration is not None else None
)
budget_reset_at: Final = (
get_budget_reset_time(budget_duration=data.budget_duration) if data.budget_duration is not None else None
)
key_rotation_at: Final = (
now + timedelta(seconds=duration_in_seconds(duration=data.rotation_interval))
if data.auto_rotate and data.rotation_interval
else None
)
return _verification_token_from_row(
MappingProxyType(
{
**columns,
"metadata": encrypt_callback_vars(folded_metadata),
"expires": expires,
"budget_reset_at": budget_reset_at,
"key_rotation_at": key_rotation_at,
"budget_limits": _generate_budget_windows(data.budget_limits),
"object_permission": None,
}
)
)
_EMPTY_DURATION_MEANS_UNCHANGED: Final = frozenset({"duration", "budget_duration"})
def _regenerate_request_as_update_request(key: str, data: RegenerateKeyRequest) -> UpdateKeyRequest | None:
changed_fields: Final = MappingProxyType(
{
field: value
for field, value in data.model_dump(exclude_unset=True).items()
if field in UpdateKeyRequest.model_fields
and field != "key"
and not (field in _EMPTY_DURATION_MEANS_UNCHANGED and value == "")
}
)
if not changed_fields:
return None
return UpdateKeyRequest(key=key, **changed_fields)
class _LegacyDumpable(Protocol):
def dict(self) -> Mapping[str, object]: ...
@ -992,6 +1149,7 @@ async def _common_key_generation_helper(
litellm_changed_by: str | None,
team_table: LiteLLM_TeamTableCachedObj | None,
) -> GenerateKeyResponse:
from litellm.proxy import proxy_server
from litellm.proxy.proxy_server import (
litellm_proxy_admin_name,
llm_router,
@ -1140,6 +1298,16 @@ async def _common_key_generation_helper(
"litellm.proxy.proxy_server.generate_key_fn(): Enterprise key management params not applied - %s", e
)
await _enforce_custom_key_policy(
hook=_custom_key_policy_hook(proxy_server),
build_policy_request=lambda: CustomKeyPolicyRequest(
operation="generate",
existing_key=None,
effective_key=_effective_key_for_generate(data=data, now=datetime.now(timezone.utc)),
request=data,
),
)
# TODO: @ishaan-jaff: Migrate all budget tracking to use LiteLLM_BudgetTable
_budget_id = data.budget_id
if prisma_client is not None and data.soft_budget is not None:
@ -2325,12 +2493,6 @@ async def prepare_key_update_data(
# sentinel for Json? columns, so store the JSON literal null
non_default_values["budget_limits"] = json.dumps(None)
if "object_permission" in non_default_values:
non_default_values = await _handle_update_object_permission(
data_json=non_default_values,
existing_key_row=existing_key_row,
)
_metadata: Final = existing_key_row.metadata or {}
# validate model_max_budget
@ -2351,13 +2513,12 @@ async def prepare_key_update_data(
async def _handle_update_object_permission(
data_json: dict,
existing_key_row: LiteLLM_VerificationToken,
prisma_client: PrismaClient,
) -> dict:
"""
Handle the update of object permission.
"""
from litellm.proxy.proxy_server import prisma_client
"""Persist the requested object permission row and swap it for its id, only after the key policy allowed the write."""
if "object_permission" not in data_json:
return data_json
# Use the common helper to handle the object permission update
object_permission_id: Final = await handle_update_object_permission_common(
data_json=data_json,
existing_object_permission_id=existing_key_row.object_permission_id,
@ -2491,6 +2652,7 @@ async def _process_single_key_update(
llm_router: Router | None,
user_custom_key_update: Callable | None = None,
existing_key_row: LiteLLM_VerificationToken | None = None,
user_custom_key_policy: Callable[..., Awaitable[Mapping[str, object]]] | None = None,
) -> dict[str, object]:
"""
Process a single key update with all validations and checks.
@ -2603,6 +2765,16 @@ async def _process_single_key_update(
data=update_key_request, existing_key_row=existing_key_row, prisma_client=prisma_client, llm_router=llm_router
)
await _enforce_custom_key_policy(
hook=user_custom_key_policy,
build_policy_request=lambda: _update_policy_request(
operation="update",
existing_key_row=existing_key_row,
non_default_values=non_default_values,
request=update_key_request,
),
)
# Update key in database
if prisma_client is None:
raise HTTPException(
@ -2610,7 +2782,12 @@ async def _process_single_key_update(
detail={"error": "Database not connected"},
)
_data: Final = {**non_default_values, "token": update_key_request.key}
update_values: Final = await _handle_update_object_permission(
data_json=non_default_values,
existing_key_row=existing_key_row,
prisma_client=prisma_client,
)
_data: Final = {**update_values, "token": update_key_request.key}
response: Final[Mapping[str, object] | None] = cast( # cast-ok: every update_data branch returns a str-keyed dict
"Mapping[str, object] | None",
await prisma_client.update_data(token=update_key_request.key, data=_data),
@ -3103,19 +3280,7 @@ async def update_key_fn(
user_api_key_cache=user_api_key_cache,
)
# Custom key update hook
custom_key_update_hook: Final[Callable[..., Awaitable[Mapping[str, object]]] | None] = _custom_key_update_hook(
proxy_server
)
if custom_key_update_hook is not None:
if inspect.iscoroutinefunction(custom_key_update_hook):
result: Final = await custom_key_update_hook(data)
else:
raise ValueError("user_custom_key_update must be a coroutine")
decision: Final = result.get("decision", True)
message: Final = result.get("message", "Authentication Failed - Custom Auth Rule")
if not decision:
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=message)
await _enforce_custom_key_update_policy(hook=_custom_key_update_hook(proxy_server), data=data)
# Enforce upperbound key params on update (don't fill defaults)
_enforce_upperbound_key_params(data, fill_defaults=False)
@ -3142,21 +3307,36 @@ async def update_key_fn(
existing_key_alias=existing_key_row.key_alias,
)
await _enforce_custom_key_policy(
hook=_custom_key_policy_hook(proxy_server),
build_policy_request=lambda: _update_policy_request(
operation="update",
existing_key_row=existing_key_row,
non_default_values=non_default_values,
request=data,
),
)
if prisma_client is None:
raise Exception("Not connected to DB!")
update_values: Final = await _handle_update_object_permission(
data_json=non_default_values,
existing_key_row=existing_key_row,
prisma_client=prisma_client,
)
changed_by: Final = user_api_key_dict.user_id or litellm_proxy_admin_name
response: Final = (
await _update_key_row_with_soft_budget(
prisma_client=prisma_client,
key=key,
data=data,
non_default_values=non_default_values,
non_default_values=update_values,
existing_key_row=existing_key_row,
changed_by=changed_by,
)
if "soft_budget" in data.model_fields_set
else await prisma_client.update_data(token=key, data=MappingProxyType({**non_default_values, "token": key}))
else await prisma_client.update_data(token=key, data=MappingProxyType({**update_values, "token": key}))
)
# Delete - key from cache, since it's been updated!
@ -3291,6 +3471,7 @@ async def bulk_update_keys(
)
custom_key_update_hook: Final = _custom_key_update_hook(proxy_server)
custom_key_policy_hook: Final = _custom_key_policy_hook(proxy_server)
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value:
raise HTTPException(
@ -3338,6 +3519,7 @@ async def bulk_update_keys(
proxy_logging_obj=proxy_logging_obj,
llm_router=llm_router,
user_custom_key_update=custom_key_update_hook,
user_custom_key_policy=custom_key_policy_hook,
)
successful_updates.append(
@ -3455,6 +3637,7 @@ async def bulk_update_team_keys(
)
custom_key_update_hook: Final = _custom_key_update_hook(proxy_server)
custom_key_policy_hook: Final = _custom_key_policy_hook(proxy_server)
if prisma_client is None:
raise HTTPException(
@ -3585,6 +3768,7 @@ async def bulk_update_team_keys(
proxy_logging_obj=proxy_logging_obj,
llm_router=llm_router,
user_custom_key_update=custom_key_update_hook,
user_custom_key_policy=custom_key_policy_hook,
existing_key_row=existing_by_token[db_token],
)
@ -5118,6 +5302,7 @@ async def _execute_virtual_key_regeneration(
proxy_logging_obj: ProxyLogging,
) -> GenerateKeyResponse:
"""Generate new token, update DB, invalidate cache, and return response."""
from litellm.proxy import proxy_server
from litellm.proxy.proxy_server import hash_token
# Mirror the /key/update ownership rebind guard. See helper docstring.
@ -5165,6 +5350,9 @@ async def _execute_virtual_key_regeneration(
non_default_values = {}
if data is not None:
update_request: Final = _regenerate_request_as_update_request(key=hashed_api_key, data=data)
if update_request is not None:
await _enforce_custom_key_update_policy(hook=_custom_key_update_hook(proxy_server), data=update_request)
# Enforce upperbound key params on regenerate (don't fill defaults)
_enforce_upperbound_key_params(data, fill_defaults=False)
non_default_values = await prepare_key_update_data(
@ -5175,7 +5363,21 @@ async def _execute_virtual_key_regeneration(
if new_key_alias != key_in_db.key_alias:
_validate_key_alias_format(key_alias=new_key_alias)
verbose_proxy_logger.debug("non_default_values: %s", non_default_values)
update_data.update(non_default_values)
await _enforce_custom_key_policy(
hook=_custom_key_policy_hook(proxy_server),
build_policy_request=lambda: _update_policy_request(
operation="regenerate",
existing_key_row=key_in_db,
non_default_values=non_default_values,
request=data if data is not None else RegenerateKeyRequest(),
),
)
update_values: Final = await _handle_update_object_permission(
data_json=non_default_values,
existing_key_row=key_in_db,
prisma_client=prisma_client,
)
update_data.update(update_values)
jsonified_update_data: Final[Mapping[str, object]] = prisma_client.jsonify_object(data=update_data)
# Snapshot before the token update: the FK cascade rewrites mapping rows to the new hash,
@ -5185,6 +5387,13 @@ async def _execute_virtual_key_regeneration(
prisma_client=prisma_client,
)
await _persist_deleted_verification_tokens(
keys=[key_in_db],
prisma_client=prisma_client,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
)
# If grace period set, insert deprecated key so old key remains valid
await _insert_deprecated_key(
prisma_client=prisma_client,
@ -5484,17 +5693,6 @@ async def regenerate_key_fn(
if litellm_changed_by is not None and not isinstance(litellm_changed_by, str):
litellm_changed_by = None
# Save the old key record to deleted table before regeneration.
# This preserves key_alias and team_id metadata for historical spend records.
# If this fails, abort the regeneration to avoid permanently losing the
# old hash→metadata mapping.
await _persist_deleted_verification_tokens(
keys=[_key_in_db],
prisma_client=prisma_client,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
)
return await _execute_virtual_key_regeneration(
prisma_client=prisma_client,
llm_router=llm_router,

View file

@ -5,18 +5,105 @@ Handles cost tracking and logging for Vertex AI Live API WebSocket passthrough e
Supports different modalities: text, audio, video, and web search.
"""
from collections.abc import Mapping, Sequence
from datetime import datetime
from typing import Any, Final
from itertools import chain, pairwise
from types import MappingProxyType
from typing import Final, Literal, TypeAlias
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.vertex_ai.gemini.grounding_requests import GroundingRequests, calculate_grounding_requests
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.base_passthrough_logging_handler import (
BasePassthroughLoggingHandler,
)
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.openai_passthrough_logging_handler import (
PassThroughEndpointLoggingTypedDict,
)
from litellm.types.utils import LlmProviders, ModelResponse, Usage
from litellm.utils import get_model_info
from litellm.types.utils import (
CompletionTokensDetailsWrapper,
CostBreakdown,
LlmProviders,
ModelResponse,
PromptTokensDetailsWrapper,
Usage,
)
_NO_GROUNDING: Final = GroundingRequests(web_search_requests=None, google_maps_grounding_requests=None)
_AGGREGATED_FIELDS: Final = frozenset(
{
"promptTokenCount",
"candidatesTokenCount",
"totalTokenCount",
"toolUsePromptTokenCount",
"promptTokensDetails",
"candidatesTokensDetails",
}
)
def _detail_entries(raw: object) -> tuple[Mapping[str, object], ...]:
"""Narrow one turn's ``*TokensDetails`` value to the entries that are actually shaped like one."""
return tuple(entry for entry in raw if isinstance(entry, Mapping)) if isinstance(raw, Sequence) else ()
def _grounding_metadata(websocket_messages: Sequence[object]) -> tuple[Mapping[str, object], ...]:
"""Collect every ``serverContent.groundingMetadata`` a session emitted.
Live reports grounding in the server frames, never in ``usageMetadata``, so the per-query
charge has to be counted here rather than derived from the token totals.
"""
return tuple(
metadata
for message in websocket_messages
if isinstance(message, Mapping)
for server_content in (message.get("serverContent"),)
if isinstance(server_content, Mapping)
for metadata in (server_content.get("groundingMetadata"),)
if isinstance(metadata, Mapping)
)
def _turns(websocket_messages: Sequence[object]) -> tuple[tuple[object, ...], ...]:
"""Split a session at every ``usageMetadata`` frame; frames after the last one never got their usage."""
closes: Final = tuple(
index + 1
for index, message in enumerate(websocket_messages)
if isinstance(message, Mapping) and isinstance(message.get("usageMetadata"), dict)
)
return tuple(tuple(websocket_messages[start:end]) for start, end in pairwise((0, *closes)))
def _session_grounding_requests(websocket_messages: Sequence[object]) -> GroundingRequests:
per_turn: Final = tuple(
calculate_grounding_requests(_grounding_metadata(turn)) for turn in _turns(websocket_messages)
)
web_search_requests: Final = sum(requests.web_search_requests or 0 for requests in per_turn)
google_maps_grounding_requests: Final = sum(requests.google_maps_grounding_requests or 0 for requests in per_turn)
return GroundingRequests(
web_search_requests=web_search_requests or None,
google_maps_grounding_requests=google_maps_grounding_requests or None,
)
_SummedField: TypeAlias = Literal[
"input_cost",
"output_cost",
"tool_usage_cost",
"cache_read_cost",
"cache_creation_cost",
"reasoning_cost",
"original_cost",
"discount_amount",
"margin_fixed_amount",
"margin_total_amount",
]
def _summed(breakdowns: Sequence[CostBreakdown], field: _SummedField) -> float | None:
values: Final = tuple(value for breakdown in breakdowns if (value := breakdown.get(field)) is not None)
return sum(values) if values else None
class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler):
@ -48,186 +135,110 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler):
"""Return the LLM provider name."""
return LlmProviders.VERTEX_AI
@staticmethod
def _resolve_detail_counts(
details: Sequence[Mapping[str, object]],
declared_total: object,
) -> tuple[tuple[str, int], ...]:
"""
Pair each of one turn's ``*TokensDetails`` entries with its token count.
Live sometimes names the modality that carries the rest of a turn without a
``tokenCount``, and reading the absent key as zero drops those tokens from the
breakdown, so real audio ends up priced as text. A lone unpriced entry therefore takes
whatever the turn's declared count leaves over. Two or more cannot be told apart, so
they are left out and the cost calculator charges the remainder as text.
"""
priced: Final = tuple(
(str(detail.get("modality", "TEXT")), count)
for detail in details
if isinstance(count := detail.get("tokenCount"), int)
)
unpriced: Final = tuple(
str(detail.get("modality", "TEXT")) for detail in details if not isinstance(detail.get("tokenCount"), int)
)
if len(unpriced) != 1 or not isinstance(declared_total, int):
return priced
residual: Final = declared_total - sum(count for _, count in priced)
return priced if residual <= 0 else (*priced, (unpriced[0], residual))
@staticmethod
def _sum_by_modality(counts: Sequence[tuple[str, int]]) -> Mapping[str, int]:
"""Total the (modality, tokenCount) pairs of one or more turns per modality."""
return MappingProxyType({modality: sum(c for m, c in counts if m == modality) for modality, _ in counts})
@staticmethod
def _merged_modality_totals(
snapshots: Sequence[Mapping[str, object]],
count_key: str,
details_key: str,
) -> Mapping[str, int]:
"""Total every turn's per-modality counts, so the breakdown adds up the way the totals do."""
return VertexAILivePassthroughLoggingHandler._sum_by_modality(
tuple(
chain.from_iterable(
VertexAILivePassthroughLoggingHandler._resolve_detail_counts(
_detail_entries(snapshot.get(details_key)), snapshot.get(count_key)
)
for snapshot in snapshots
)
)
)
@staticmethod
def _extract_usage_metadata_from_websocket_messages(
websocket_messages: list[dict],
websocket_messages: Sequence[object],
) -> dict | None:
"""
Extract and aggregate usage metadata from a list of WebSocket messages.
Live emits one ``usageMetadata`` per turn and Google charges per turn for every token in
the session context window, which is the current turn's tokens plus all accumulated
tokens from previous turns, so the turns add up rather than restating each other. See
the Live API note under https://cloud.google.com/vertex-ai/generative-ai/pricing.
Args:
websocket_messages: List of WebSocket messages from the Live API
Returns:
Dictionary containing aggregated usage metadata, or None if not found
"""
all_usage_metadata: Final = []
snapshots: Final = tuple(
metadata
for message in websocket_messages
if isinstance(message, Mapping)
for metadata in (message.get("usageMetadata"),)
if isinstance(metadata, dict)
)
# Collect all usage metadata messages
for message in websocket_messages:
if isinstance(message, dict) and "usageMetadata" in message:
all_usage_metadata.append(message["usageMetadata"])
if not all_usage_metadata:
if not snapshots:
return None
# If only one usage metadata, return it as-is
if len(all_usage_metadata) == 1:
return all_usage_metadata[0]
# Aggregate multiple usage metadata messages
aggregated: Final[dict[str, Any]] = {
"promptTokenCount": 0,
"candidatesTokenCount": 0,
"totalTokenCount": 0,
"promptTokensDetails": [],
"candidatesTokensDetails": [],
prompt_totals: Final = VertexAILivePassthroughLoggingHandler._merged_modality_totals(
snapshots, "promptTokenCount", "promptTokensDetails"
)
candidate_totals: Final = VertexAILivePassthroughLoggingHandler._merged_modality_totals(
snapshots, "candidatesTokenCount", "candidatesTokensDetails"
)
return {
**{key: value for key, value in snapshots[0].items() if key not in _AGGREGATED_FIELDS},
"promptTokenCount": sum(snapshot.get("promptTokenCount", 0) for snapshot in snapshots),
"candidatesTokenCount": sum(snapshot.get("candidatesTokenCount", 0) for snapshot in snapshots),
"totalTokenCount": sum(snapshot.get("totalTokenCount", 0) for snapshot in snapshots),
"toolUsePromptTokenCount": sum(snapshot.get("toolUsePromptTokenCount", 0) for snapshot in snapshots),
"promptTokensDetails": [
{"modality": modality, "tokenCount": count} for modality, count in prompt_totals.items() if count > 0
],
"candidatesTokensDetails": [
{"modality": modality, "tokenCount": count} for modality, count in candidate_totals.items() if count > 0
],
}
# Aggregate token counts
for usage in all_usage_metadata:
aggregated["promptTokenCount"] += usage.get("promptTokenCount", 0)
aggregated["candidatesTokenCount"] += usage.get("candidatesTokenCount", 0)
aggregated["totalTokenCount"] += usage.get("totalTokenCount", 0)
# Aggregate token details by modality
modality_totals: Final = {}
for usage in all_usage_metadata:
# Process prompt tokens details
for detail in usage.get("promptTokensDetails", []):
modality = detail.get("modality", "TEXT")
token_count = detail.get("tokenCount", 0)
if modality not in modality_totals:
modality_totals[modality] = {"prompt": 0, "candidate": 0}
modality_totals[modality]["prompt"] += token_count
# Process candidate tokens details
for detail in usage.get("candidatesTokensDetails", []):
modality = detail.get("modality", "TEXT")
token_count = detail.get("tokenCount", 0)
if modality not in modality_totals:
modality_totals[modality] = {"prompt": 0, "candidate": 0}
modality_totals[modality]["candidate"] += token_count
# Convert aggregated modality totals back to details format
for modality, totals in modality_totals.items():
if totals["prompt"] > 0:
aggregated["promptTokensDetails"].append({"modality": modality, "tokenCount": totals["prompt"]})
if totals["candidate"] > 0:
aggregated["candidatesTokensDetails"].append({"modality": modality, "tokenCount": totals["candidate"]})
# Add any additional fields from the first usage metadata
first_usage: Final = all_usage_metadata[0]
for key, value in first_usage.items():
if key not in aggregated:
aggregated[key] = value
return aggregated
@staticmethod
def _calculate_live_api_cost(
model: str,
usage_metadata: dict,
custom_llm_provider: str = "vertex_ai",
) -> float:
"""
Calculate cost for Vertex AI Live API based on usage metadata.
Args:
model: The model name (e.g., "gemini-2.0-flash-live-preview-04-09")
usage_metadata: Usage metadata from the Live API response
custom_llm_provider: The LLM provider (default: "vertex_ai")
Returns:
Total cost in USD
"""
try:
# Get model pricing information
model_info: Final = get_model_info(model=model, custom_llm_provider=custom_llm_provider)
verbose_proxy_logger.debug("Vertex AI Live API model info for '%s': %s", model, model_info)
# Check if pricing info is available
if not model_info or not model_info.get("input_cost_per_token"):
verbose_proxy_logger.error("No pricing info found for %s in local model pricing database", model)
return 0.0
total_cost = 0.0
# Extract token counts from usage metadata
prompt_token_count: Final = usage_metadata.get("promptTokenCount", 0)
candidates_token_count: Final = usage_metadata.get("candidatesTokenCount", 0)
# Calculate base text token costs
input_cost_per_token: Final = model_info.get("input_cost_per_token", 0.0)
output_cost_per_token: Final = model_info.get("output_cost_per_token", 0.0)
total_cost += prompt_token_count * input_cost_per_token
total_cost += candidates_token_count * output_cost_per_token
# Handle modality-specific costs if present
prompt_tokens_details: Final = usage_metadata.get("promptTokensDetails", [])
candidates_tokens_details: Final = usage_metadata.get("candidatesTokensDetails", [])
# Process prompt tokens by modality
for detail in prompt_tokens_details:
modality = detail.get("modality", "TEXT")
token_count = detail.get("tokenCount", 0)
if modality == "AUDIO":
audio_cost_per_token = model_info.get("input_cost_per_audio_token", 0.0)
total_cost += token_count * audio_cost_per_token
elif modality == "VIDEO":
# Video tokens are typically per second, but we'll treat as per token for now
video_cost_per_token = model_info.get("input_cost_per_video_per_second", 0.0)
total_cost += token_count * video_cost_per_token
# TEXT tokens are already handled above
# Process candidate tokens by modality
for detail in candidates_tokens_details:
modality = detail.get("modality", "TEXT")
token_count = detail.get("tokenCount", 0)
if modality == "AUDIO":
audio_cost_per_token = model_info.get("output_cost_per_audio_token", 0.0)
total_cost += token_count * audio_cost_per_token
elif modality == "VIDEO":
# Video tokens are typically per second, but we'll treat as per token for now
video_cost_per_token = model_info.get("output_cost_per_video_per_second", 0.0)
total_cost += token_count * video_cost_per_token
# TEXT tokens are already handled above
# Handle web search costs if present
tool_use_prompt_token_count: Final = usage_metadata.get("toolUsePromptTokenCount", 0)
if tool_use_prompt_token_count > 0:
# Web search typically has a fixed cost per request
web_search_cost: Final = model_info.get("web_search_cost_per_request", 0.0)
if isinstance(web_search_cost, (int, float)) and web_search_cost > 0:
total_cost += web_search_cost
else:
# Fallback to token-based pricing for tool use
total_cost += tool_use_prompt_token_count * input_cost_per_token
verbose_proxy_logger.debug(
f"Vertex AI Live API cost calculation - Model: {model}, "
f"Prompt tokens: {prompt_token_count}, "
f"Candidate tokens: {candidates_token_count}, "
f"Total cost: ${total_cost:.6f}"
)
return total_cost
except Exception as e:
verbose_proxy_logger.error("Error calculating Vertex AI Live API cost: %s", e)
return 0.0
@staticmethod
def _create_usage_object_from_metadata(
usage_metadata: dict,
model: str,
grounding_requests: GroundingRequests = _NO_GROUNDING,
) -> Usage:
"""
Create a LiteLLM Usage object from Live API usage metadata.
@ -235,48 +246,124 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler):
Args:
usage_metadata: Usage metadata from the Live API response
model: The model name
grounding_requests: The Search and Maps grounding requests summed over the session's
turns, matching the per-turn charge
Returns:
LiteLLM Usage object
"""
prompt_tokens: Final = usage_metadata.get("promptTokenCount", 0)
completion_tokens: Final = usage_metadata.get("candidatesTokenCount", 0)
total_tokens: Final = usage_metadata.get("totalTokenCount", 0)
prompt_by_modality: Final = VertexAILivePassthroughLoggingHandler._sum_by_modality(
VertexAILivePassthroughLoggingHandler._resolve_detail_counts(
_detail_entries(usage_metadata.get("promptTokensDetails")), usage_metadata.get("promptTokenCount")
)
)
candidates_by_modality: Final = VertexAILivePassthroughLoggingHandler._sum_by_modality(
VertexAILivePassthroughLoggingHandler._resolve_detail_counts(
_detail_entries(usage_metadata.get("candidatesTokensDetails")),
usage_metadata.get("candidatesTokenCount"),
)
)
# Create modality-specific token details if available
prompt_tokens_details: Final = usage_metadata.get("promptTokensDetails", [])
candidates_tokens_details: Final = usage_metadata.get("candidatesTokensDetails", [])
# Extract text tokens from details
text_prompt_tokens = 0
text_completion_tokens = 0
for detail in prompt_tokens_details:
if detail.get("modality") == "TEXT":
text_prompt_tokens = detail.get("tokenCount", 0)
break
for detail in candidates_tokens_details:
if detail.get("modality") == "TEXT":
text_completion_tokens = detail.get("tokenCount", 0)
break
# If no text tokens found in details, use total counts
if text_prompt_tokens == 0:
text_prompt_tokens = prompt_tokens
if text_completion_tokens == 0:
text_completion_tokens = completion_tokens
prompt_tokens: Final = usage_metadata.get("promptTokenCount", 0) or sum(prompt_by_modality.values())
completion_tokens: Final = usage_metadata.get("candidatesTokenCount", 0) or sum(candidates_by_modality.values())
return Usage(
prompt_tokens=text_prompt_tokens,
completion_tokens=text_completion_tokens,
total_tokens=total_tokens,
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=usage_metadata.get("totalTokenCount", 0) or (prompt_tokens + completion_tokens),
prompt_tokens_details=PromptTokensDetailsWrapper(
text_tokens=prompt_by_modality.get("TEXT"),
audio_tokens=prompt_by_modality.get("AUDIO"),
image_tokens=prompt_by_modality.get("IMAGE"),
video_tokens=prompt_by_modality.get("VIDEO"),
tool_use_tokens=usage_metadata.get("toolUsePromptTokenCount") or None,
web_search_requests=grounding_requests.web_search_requests,
google_maps_grounding_requests=grounding_requests.google_maps_grounding_requests,
),
completion_tokens_details=CompletionTokensDetailsWrapper(
text_tokens=candidates_by_modality.get("TEXT"),
audio_tokens=candidates_by_modality.get("AUDIO"),
image_tokens=candidates_by_modality.get("IMAGE"),
video_tokens=candidates_by_modality.get("VIDEO"),
),
)
def _session_usage(self, websocket_messages: Sequence[object], model: str) -> Usage | None:
usage_metadata: Final = self._extract_usage_metadata_from_websocket_messages(websocket_messages)
if usage_metadata is None:
return None
return self._create_usage_object_from_metadata(
usage_metadata=usage_metadata,
grounding_requests=_session_grounding_requests(websocket_messages),
model=model,
)
def _turn_cost(
self,
turn: Sequence[object],
model: str,
logging_obj: LiteLLMLoggingObj,
) -> tuple[float, CostBreakdown] | None:
usage: Final = self._session_usage(turn, model)
if usage is None:
return None
cost: Final = logging_obj._response_cost_calculator( # pyright: ignore[reportPrivateUsage] # the call's own calculator keeps custom pricing and the deployment's region in step with the spend row
result=ModelResponse(model=model, usage=usage),
litellm_model_name=model,
)
if cost is None:
return None
breakdown: Final = logging_obj.cost_breakdown
return None if breakdown is None else (cost, breakdown)
def _session_cost(
self,
websocket_messages: Sequence[object],
model: str,
logging_obj: LiteLLMLoggingObj,
) -> float | None:
"""Price each turn on its own tokens and grounding, so two grounded turns pay the query fee twice.
The fixed cost margin is a flat per-request fee, so the session's single spend row carries it once
rather than once per turn.
"""
turn_costs: Final = tuple(self._turn_cost(turn, model, logging_obj) for turn in _turns(websocket_messages))
priced: Final = tuple(turn_cost for turn_cost in turn_costs if turn_cost is not None)
if not priced or len(priced) != len(turn_costs):
return None
breakdowns: Final = tuple(breakdown for _, breakdown in priced)
first: Final = breakdowns[0]
fixed_margin: Final = first.get("margin_fixed_amount") or 0.0
duplicated_fixed_margin: Final = fixed_margin * (len(priced) - 1)
total_cost: Final = sum(cost for cost, _ in priced) - duplicated_fixed_margin
summed_margin_total: Final = _summed(breakdowns, "margin_total_amount")
margin_total_amount: Final = (
None if summed_margin_total is None else summed_margin_total - duplicated_fixed_margin
)
logging_obj.set_cost_breakdown(
input_cost=_summed(breakdowns, "input_cost") or 0.0,
output_cost=_summed(breakdowns, "output_cost") or 0.0,
total_cost=total_cost,
cost_for_built_in_tools_cost_usd_dollar=_summed(breakdowns, "tool_usage_cost") or 0.0,
original_cost=_summed(breakdowns, "original_cost"),
discount_percent=first.get("discount_percent"),
discount_amount=_summed(breakdowns, "discount_amount"),
margin_percent=first.get("margin_percent"),
margin_fixed_amount=first.get("margin_fixed_amount"),
margin_total_amount=margin_total_amount,
cache_read_cost=_summed(breakdowns, "cache_read_cost"),
cache_creation_cost=_summed(breakdowns, "cache_creation_cost"),
reasoning_cost=_summed(breakdowns, "reasoning_cost"),
service_tier=first.get("service_tier"),
data_residency=first.get("data_residency"),
vertex_location=first.get("vertex_location"),
)
return total_cost
def vertex_ai_live_passthrough_handler(
self,
websocket_messages: list[dict],
logging_obj,
websocket_messages: Sequence[object],
logging_obj: LiteLLMLoggingObj,
url_route: str,
start_time: datetime,
end_time: datetime,
@ -300,34 +387,25 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler):
"""
try:
# Extract model from request body or kwargs
model: Final = kwargs.get("model", "gemini-2.0-flash-live-preview-04-09")
requested_model: Final = kwargs.get("model")
model: Final = (
requested_model if isinstance(requested_model, str) else "gemini-2.0-flash-live-preview-04-09"
)
custom_llm_provider: Final = kwargs.get("custom_llm_provider", "vertex_ai")
verbose_proxy_logger.debug(
"Vertex AI Live API model: %s, custom_llm_provider: %s", model, custom_llm_provider
)
# Extract usage metadata from WebSocket messages
usage_metadata: Final = self._extract_usage_metadata_from_websocket_messages(websocket_messages)
usage: Final = self._session_usage(websocket_messages, model)
if not usage_metadata:
if usage is None:
verbose_proxy_logger.warning("No usage metadata found in Vertex AI Live API WebSocket messages")
return {
"result": None,
"kwargs": kwargs,
}
# Calculate cost using Live API specific pricing
response_cost: Final = self._calculate_live_api_cost(
model=model,
usage_metadata=usage_metadata,
custom_llm_provider=custom_llm_provider,
)
# Create Usage object for standard LiteLLM logging
usage: Final = self._create_usage_object_from_metadata(
usage_metadata=usage_metadata,
model=model,
)
response_cost: Final = self._session_cost(websocket_messages, model, logging_obj)
# Create a mock ModelResponse for standard logging
litellm_model_response: Final = ModelResponse(
@ -338,9 +416,9 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler):
usage=usage,
choices=[],
)
if response_cost is not None:
litellm_model_response._hidden_params["response_cost"] = response_cost # pyright: ignore[reportPrivateUsage] # the logger reads the cost off the response's hidden params; the constructor's hidden_params kwarg is reset by pydantic
# Update kwargs with cost information
kwargs["response_cost"] = response_cost
kwargs["model"] = model
kwargs["custom_llm_provider"] = custom_llm_provider
@ -348,12 +426,15 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler):
import re
allowed_pattern: Final = re.compile(r"^[A-Za-z0-9._\-:]+$")
safe_model: Final = model if isinstance(model, str) and allowed_pattern.match(model) else "[REDACTED]"
safe_model: Final = model if allowed_pattern.match(model) else "[REDACTED]"
verbose_proxy_logger.debug(
f"Vertex AI Live API passthrough cost tracking - "
f"Model: {safe_model}, Cost: ${response_cost:.6f}, "
f"Prompt tokens: {usage.prompt_tokens}, "
f"Completion tokens: {usage.completion_tokens}"
"Vertex AI Live API passthrough cost tracking - Model: %s, "
"Prompt tokens: %s %s, Completion tokens: %s %s",
safe_model,
usage.prompt_tokens,
usage.prompt_tokens_details,
usage.completion_tokens,
usage.completion_tokens_details,
)
return {

View file

@ -2090,6 +2090,22 @@ def _rewrite_vertex_live_setup_model(text_data: str, setup_model_rewriter: Calla
return json.dumps({**message, "setup": {**setup, "model": rewritten_model}}) # mutable-ok: one-shot json payload
def _resolved_vertex_live_setup(
setup_data: Mapping[str, object], setup_model_rewriter: Callable[[str], str] | None
) -> Mapping[str, object]:
"""
Give the model extractor the same fully qualified path the upstream will receive.
Clients may name a bare gateway alias, which the rewriter turns into a ``projects/...`` path before
it reaches Vertex. The extractor only reads a path containing ``/models/``, so running it on the raw
frame logs the session as ``unknown`` at no cost, which is precisely the supported client form
"""
setup_model: Final = setup_data.get("model")
if setup_model_rewriter is None or not isinstance(setup_model, str):
return setup_data
return {**setup_data, "model": setup_model_rewriter(setup_model)}
def _truncated_close_reason(reason: str) -> str:
"""
Fit a close reason inside the byte budget a WebSocket close frame allows, without splitting a character
@ -2314,7 +2330,9 @@ async def websocket_passthrough_request(
setup_data,
)
if isinstance(setup_data, dict) and "model" in setup_data:
extracted_model = _extract_model_from_vertex_ai_setup(setup_data)
extracted_model = _extract_model_from_vertex_ai_setup(
_resolved_vertex_live_setup(setup_data, setup_model_rewriter)
)
if extracted_model:
kwargs["model"] = extracted_model
kwargs["custom_llm_provider"] = "vertex_ai-language-models"

View file

@ -928,6 +928,7 @@ def cleanup_router_config_variables():
user_custom_auth_path, \
user_custom_key_generate, \
user_custom_key_update, \
user_custom_key_policy, \
user_custom_sso, \
user_custom_ui_sso_sign_in_handler, \
use_background_health_checks, \
@ -945,6 +946,7 @@ def cleanup_router_config_variables():
user_custom_auth_path = None
user_custom_key_generate = None
user_custom_key_update = None
user_custom_key_policy = None
TEAM_METADATA_VALIDATOR_REGISTRY.set(None)
TEAM_METADATA_SCHEMA_REGISTRY.set(())
user_custom_sso = None
@ -2369,6 +2371,7 @@ user_custom_key_generate = None
_pkce_no_redis_warning_emitted: bool = False
_cp_no_redis_warning_emitted: bool = False
user_custom_key_update = None
user_custom_key_policy = None
user_custom_sso = None
user_custom_ui_sso_sign_in_handler = None
use_background_health_checks = None
@ -4256,6 +4259,7 @@ _DB_OVERLAY_REMOTE_MODULE_STR_FIELDS: Final[dict[str, tuple[str, ...]]] = {
"custom_auth",
"custom_key_generate",
"custom_key_update",
"custom_key_policy",
"custom_team_metadata_validate",
"custom_sso",
"custom_ui_sso_sign_in_handler",
@ -5405,6 +5409,7 @@ class ProxyConfig:
user_custom_auth_path, \
user_custom_key_generate, \
user_custom_key_update, \
user_custom_key_policy, \
user_custom_sso, \
user_custom_ui_sso_sign_in_handler, \
use_background_health_checks, \
@ -5942,6 +5947,10 @@ class ProxyConfig:
if custom_key_update is not None:
user_custom_key_update = get_instance_fn(value=custom_key_update, config_file_path=config_file_path)
custom_key_policy: Final = general_settings.get("custom_key_policy", None)
if custom_key_policy is not None:
user_custom_key_policy = get_instance_fn(value=custom_key_policy, config_file_path=config_file_path)
custom_team_metadata_validate: Final = general_settings.get("custom_team_metadata_validate", None)
TEAM_METADATA_VALIDATOR_REGISTRY.set(
get_instance_fn(value=custom_team_metadata_validate, config_file_path=config_file_path)

View file

@ -30,6 +30,7 @@ from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import CallTypes, LlmProviders
from litellm.utils import ProviderConfigManager
from ..litellm_core_utils.credential_accessor import CredentialAccessor
from ..litellm_core_utils.get_litellm_params import get_litellm_params
from ..litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from ..llms.azure.common_utils import get_azure_ad_token
@ -54,6 +55,17 @@ xai_realtime: Final = XAIRealtime()
vertex_llm_base: Final = VertexBase()
base_llm_http_handler = BaseLLMHTTPHandler()
_EMPTY_MODEL_PARAMS: Final[Mapping[str, Any]] = MappingProxyType({})
_EMPTY_AUTH_HEADERS: Final[Mapping[str, str]] = MappingProxyType({})
def _model_params_with_stored_credentials(model_params: Mapping[str, Any]) -> Mapping[str, Any]:
credential_name: Final = model_params.get("litellm_credential_name")
credential_values: Final = (
CredentialAccessor.get_credential_values(credential_name)
if isinstance(credential_name, str)
else _EMPTY_MODEL_PARAMS
)
return MappingProxyType({**credential_values, **model_params})
def _with_resolved_session_model(session: dict[str, object], model_name: str) -> dict[str, object]:
@ -591,13 +603,15 @@ def _azure_realtime_health_protocol(
def _realtime_health_check_auth_headers(
custom_llm_provider: str, api_key: str | None, model_params: Mapping[str, Any]
) -> Mapping[str, str | None]:
if custom_llm_provider != "azure":
return MappingProxyType({"api-key": api_key})
return azure_realtime.get_auth_headers(
api_key=api_key,
azure_ad_token=(None if api_key else get_azure_ad_token(GenericLiteLLMParams(**model_params))),
)
) -> Mapping[str, str]:
if custom_llm_provider == "azure":
return azure_realtime.get_auth_headers(
api_key=api_key,
azure_ad_token=(None if api_key else get_azure_ad_token(GenericLiteLLMParams(**model_params))),
)
if api_key is None:
return _EMPTY_AUTH_HEADERS
return MappingProxyType({"Authorization": f"Bearer {api_key}"})
async def _realtime_health_check(
@ -629,34 +643,46 @@ async def _realtime_health_check(
"""
import websockets
resolved_params: Final = _model_params_with_stored_credentials(model_params or _EMPTY_MODEL_PARAMS)
resolved_api_key: Final = cast( # cast-ok: provider parameters expose optional string credentials
str | None, api_key or resolved_params.get("api_key")
)
resolved_api_base: Final = cast( # cast-ok: provider parameters expose optional string endpoints
str | None, api_base or resolved_params.get("api_base")
)
resolved_api_version: Final = cast( # cast-ok: provider parameters expose optional string versions
str | None, api_version or resolved_params.get("api_version")
)
url: str | None = None
auth_headers: Final = _realtime_health_check_auth_headers(
custom_llm_provider=custom_llm_provider,
api_key=api_key,
model_params=model_params or _EMPTY_MODEL_PARAMS,
api_key=resolved_api_key,
model_params=resolved_params,
)
if custom_llm_provider == "azure":
resolved_protocol, azure_query_params = _azure_realtime_health_protocol(
model=model,
realtime_protocol=realtime_protocol,
model_params=model_params or _EMPTY_MODEL_PARAMS,
model_params=resolved_params,
)
url = azure_realtime._construct_url(
api_base=api_base or "",
api_base=resolved_api_base or "",
model=model,
api_version=api_version or "2024-10-01-preview",
api_version=resolved_api_version or "2024-10-01-preview",
realtime_protocol=resolved_protocol,
query_params=azure_query_params,
)
elif custom_llm_provider == "openai":
url = openai_realtime._construct_url(
api_base=api_base or "https://api.openai.com/",
api_base=resolved_api_base or "https://api.openai.com/",
query_params={"model": model},
)
elif custom_llm_provider == "xai":
url = xai_realtime._construct_url(api_base=api_base or "https://api.x.ai/v1", query_params={"model": model})
url = xai_realtime._construct_url(
api_base=resolved_api_base or "https://api.x.ai/v1", query_params={"model": model}
)
elif custom_llm_provider == "vertex_ai":
vertex_model_params: Final = model_params or {}
vertex_model_params: Final = dict(resolved_params)
resolved_location: Final = vertex_llm_base.get_vertex_region(
vertex_region=VertexBase.safe_get_vertex_ai_location(vertex_model_params),
model=model,
@ -675,19 +701,19 @@ async def _realtime_health_check(
project=resolved_project,
location=resolved_location,
)
url = vertex_realtime_config.get_complete_url(api_base=api_base, model=model)
ssl_context = get_shared_realtime_ssl_context()
url = vertex_realtime_config.get_complete_url(api_base=resolved_api_base, model=model)
vertex_ssl_context: Final = get_shared_realtime_ssl_context()
headers: Final = vertex_realtime_config.validate_environment(headers={}, model=model, api_key=None)
async with websockets.connect(
url,
additional_headers=headers,
max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
ssl=ssl_context,
ssl=vertex_ssl_context,
):
return True
else:
raise ValueError(f"Unsupported model: {model}")
ssl_context = get_shared_realtime_ssl_context()
ssl_context: Final = get_shared_realtime_ssl_context()
async with websockets.connect(
url,
additional_headers=auth_headers,

View file

@ -2831,6 +2831,22 @@ class LiteLLMCompletionResponsesConfig:
if cache_write_tokens is not None
else MappingProxyType({})
)
# The cost path reads the grounding counters off the input details, and a realtime
# session's usage is rebuilt from its own response.done, so dropping them here bills
# no per-query grounding fee at all.
grounding_request_counts: Final[Mapping[str, int]] = MappingProxyType(
{
counter: count
for counter, count in (
("web_search_requests", getattr(prompt_details, "web_search_requests", None)),
(
"google_maps_grounding_requests",
getattr(prompt_details, "google_maps_grounding_requests", None),
),
)
if count is not None
}
)
response_usage.input_tokens_details = InputTokensDetails(
cached_tokens=prompt_details.cached_tokens if prompt_details.cached_tokens is not None else 0,
text_tokens=prompt_details.text_tokens,
@ -2839,6 +2855,7 @@ class LiteLLMCompletionResponsesConfig:
cached_tokens_details if isinstance(cached_tokens_details, CachedTokensDetails) else None
),
**cache_write_extra,
**grounding_request_counts,
)
# Translate completion_tokens_details to output_tokens_details

View file

@ -1183,6 +1183,10 @@ class ResponseAPILoggingUtils:
response_api_usage.input_tokens_details, "cached_tokens_details", None
),
cache_write_tokens=getattr(response_api_usage.input_tokens_details, "cache_write_tokens", None),
web_search_requests=getattr(response_api_usage.input_tokens_details, "web_search_requests", None),
google_maps_grounding_requests=getattr(
response_api_usage.input_tokens_details, "google_maps_grounding_requests", None
),
)
completion_tokens_details: CompletionTokensDetailsWrapper | None = None
output_tokens_details: Final[OutputTokensDetails | None] = getattr(

View file

@ -8377,7 +8377,8 @@ class Router:
def log_retry(self, kwargs: dict, e: Exception) -> dict:
"""
When a retry or fallback happens, record which model group, deployment and attempt just failed and why
When a retry or fallback happens, record which model group, deployment and attempt just failed and why,
and count it toward the request-wide num_retries_per_request cap
"""
from litellm.types.router import RetryAttemptRecord
@ -8401,7 +8402,10 @@ class Router:
else ()
)
breadcrumbs: Final = (*kept_breadcrumbs, attempt_record)
earlier: Final = request_metadata.get("request_retry_count")
request_retry_count: Final = (earlier if type(earlier) is int and 0 <= earlier else 0) + 1
kwargs[_metadata_var]["previous_models"] = breadcrumbs # rebind-ok: the logging object already holds this dict
kwargs[_metadata_var]["request_retry_count"] = request_retry_count # rebind-ok: same dict, read by the cap
return kwargs
def _update_usage(self, deployment_id: str, parent_otel_span: Span | None) -> int:

View file

@ -1,9 +1,12 @@
from datetime import datetime
from typing import Any, Final, Literal
from typing import Any, Final, Literal, TypeAlias
from pydantic import BaseModel, ConfigDict, model_validator
from typing_extensions import ReadOnly, TypedDict
from litellm.models.verification_token import LiteLLM_VerificationToken
from litellm.proxy._types import GenerateKeyRequest, RegenerateKeyRequest, UpdateKeyRequest
from litellm.types.llms.base import LiteLLMPydanticObjectBase
from litellm.types.proxy.management_endpoints.internal_user_endpoints import InsensitiveContains
@ -123,3 +126,24 @@ class BulkUpdateTeamKeysRequest(BaseModel):
if not has_key_ids and not self.all_keys_in_team:
raise ValueError("Must provide either `key_ids` (non-empty) or `all_keys_in_team=True`.")
return self
CustomKeyPolicyOperation: TypeAlias = Literal["generate", "update", "regenerate"]
class CustomKeyPolicyRequest(LiteLLMPydanticObjectBase):
"""What `general_settings.custom_key_policy` receives.
`effective_key` is the verification token row as it will be written: the existing row overlaid with the
requested changes, with `duration` resolved to `expires` and `budget_duration` to `budget_reset_at`. Values the
proxy fills in after the policy stay at their defaults: `token`, `key_name`, `created_by`, `updated_by` and the
soft-budget `budget_id` on generate, the rotated token on regenerate, and the `object_permission` relation on
every operation (`object_permission_id` is set; read `request.object_permission` for the requested change).
"""
model_config = ConfigDict(protected_namespaces=(), frozen=True)
operation: CustomKeyPolicyOperation
existing_key: LiteLLM_VerificationToken | None
effective_key: LiteLLM_VerificationToken
request: GenerateKeyRequest | UpdateKeyRequest | RegenerateKeyRequest

View file

@ -170,6 +170,7 @@ class StreamingResponse(BaseModel):
# the consumed body is elided, so this is the only place they surface.
stream_error: str | None = None
stream_done: bool = False
stream_done_positions: tuple[int, ...] = ()
@property
def ok(self) -> bool:
@ -647,6 +648,7 @@ def streaming_outcome(
stream_events=[payload for payload, _ in events],
stream_event_arrivals=[arrived for _, arrived in events],
stream_done=any(payload == _SSE_DONE for payload, _ in payloads),
stream_done_positions=tuple(index for index, (payload, _) in enumerate(payloads) if payload == _SSE_DONE),
stream_error=next(
(line.decode(errors="replace")[:300] for line, _ in stamped if _is_stream_error_line(line)),
None,

View file

@ -1,51 +1,90 @@
"""Vendor §12.3: chat completions streaming SSE contract (LIT-4778).
Asserts a streamed /chat/completions response is SSE, carries content chunks,
and terminates with the OpenAI [DONE] sentinel.
"""
from __future__ import annotations
from typing import Final
import pytest
from e2e_config import unique_marker
from e2e_config import provider_edge_base, unique_marker
from e2e_http import require_successful_call
from lifecycle import ResourceManager
from models import ChatBody, ChatMessage, LiteLLMParamsBody
from models import ChatBody, ChatMessage, ChatStreamOptions, LiteLLMParamsBody, Usage
from proxy_client import ProxyClient
from pydantic import BaseModel
pytestmark = pytest.mark.e2e
pytestmark = [pytest.mark.e2e, pytest.mark.replayable]
class _Delta(BaseModel):
content: str | None = None
class _Choice(BaseModel):
index: int
delta: _Delta
finish_reason: str | None = None
class _Chunk(BaseModel):
choices: tuple[_Choice, ...]
usage: Usage | None = None
class TestChatStreamContract:
@pytest.mark.covers("llm.chat_completions.openai.basic.stream.works")
def test_chat_stream_is_sse_and_ends_with_done(self, proxy: ProxyClient, resources: ResourceManager) -> None:
model = f"e2e-chat-stream-{unique_marker()}"
model_id = proxy.create_model(
model: Final = f"e2e-chat-stream-{unique_marker()}"
base: Final = provider_edge_base("openai")
model_id: Final = proxy.create_model(
model,
LiteLLMParamsBody(model="openai/gpt-4o-mini", api_key="os.environ/OPENAI_API_KEY"),
LiteLLMParamsBody(
model="openai/gpt-5.6",
api_key="os.environ/OPENAI_API_KEY",
api_base=f"{base}/v1" if base else None,
),
)
resources.defer(lambda: proxy.delete_model(model_id))
key = resources.key()
result = proxy.chat_stream(
key: Final = resources.key()
expected: Final = "The amber kite crosses the quiet lake."
result: Final = proxy.chat_stream(
key,
ChatBody(
model=model,
messages=[
ChatMessage(
role="user",
content=f"Reply with the single word ok. {unique_marker()}",
role="user", content=f"Repeat exactly this sentence, with no additional text: {expected}"
)
],
stream=True,
max_completion_tokens=32,
temperature=0.0,
stream_options=ChatStreamOptions(include_usage=True),
max_completion_tokens=256,
reasoning_effort="none",
),
)
require_successful_call(result)
assert result.is_streaming, f"expected SSE content-type, got {result.content_type!r}"
assert result.stream_events, "stream returned no data events"
assert result.stream_done, (
f"stream must terminate with [DONE]; "
f"chunks={result.chunks} done={result.stream_done} events={len(result.stream_events)}"
assert not result.stream_error, f"stream errored: {result.stream_error}"
assert result.stream_done, "stream must terminate with [DONE]"
assert result.stream_done_positions == (len(result.stream_events),), "[DONE] must occur once after all events"
chunks: Final = tuple(_Chunk.model_validate_json(event) for event in result.stream_events)
text_positions: Final = tuple(
i for i, chunk in enumerate(chunks) if any(c.delta.content for c in chunk.choices)
)
terminal_positions: Final = tuple(
i for i, chunk in enumerate(chunks) if any(c.finish_reason is not None for c in chunk.choices)
)
assert text_positions, "stream completed without meaningful text"
assert len(terminal_positions) == 1, "expected exactly one terminal choice"
assert text_positions[0] < terminal_positions[0], "meaningful text must arrive before termination"
assert text_positions[-1] <= terminal_positions[0], "text arrived after termination"
assert all(c.index == 0 for chunk in chunks for c in chunk.choices)
assert tuple(c.finish_reason for c in chunks[terminal_positions[0]].choices) == ("stop",)
text: Final = "".join(c.delta.content or "" for chunk in chunks for c in chunk.choices)
assert text.strip() == expected, f"streamed answer was altered or incomplete: {text!r}"
usage_positions: Final = tuple(i for i, chunk in enumerate(chunks) if chunk.usage is not None)
assert usage_positions == (len(chunks) - 1,), "expected one final usage chunk"
assert terminal_positions[0] < usage_positions[0], "usage must follow the terminal choice"
usage: Final = chunks[-1].usage
assert usage is not None
assert usage.prompt_tokens is not None and usage.prompt_tokens > 0
assert usage.completion_tokens is not None and usage.completion_tokens > 0
assert usage.total_tokens == usage.prompt_tokens + usage.completion_tokens

View file

@ -21,7 +21,12 @@ from e2e_http import assert_client_error, require_successful_call, unwrap
from endpoints_client import EndpointsClient, MessagesResult
from lifecycle import ResourceManager
from models import (
AnthropicAssistantTurn,
AnthropicContentBlock,
AnthropicCustomTool,
AnthropicToolChoice,
AnthropicToolResultBlock,
AnthropicToolResultTurn,
AnthropicMessagesBody,
ChatMessage,
JsonSchemaProperty,
@ -29,7 +34,7 @@ from models import (
SpendLogRow,
ToolInputSchema,
)
from pydantic import BaseModel
from pydantic import BaseModel, ConfigDict
pytestmark = [pytest.mark.e2e, pytest.mark.replayable]
@ -284,8 +289,139 @@ class TestAnthropicMessages:
result = endpoints_client.proxy.transport.send(
"/v1/messages",
headers=endpoints_client.proxy.transport.bearer(key),
json=_OptionalMessagesBody(
messages=[ChatMessage(role="user", content="hi")], max_tokens=50
),
json=_OptionalMessagesBody(messages=[ChatMessage(role="user", content="hi")], max_tokens=50),
)
assert_client_error(result, "messages missing model")
class _BridgeDelta(BaseModel):
type: str | None = None
partial_json: str | None = None
stop_reason: str | None = None
class _BridgeEvent(BaseModel):
type: str
index: int | None = None
content_block: AnthropicContentBlock | None = None
delta: _BridgeDelta | None = None
class _ParcelInput(BaseModel):
model_config = ConfigDict(extra="forbid", strict=True)
parcel: str
shelf: int
def _tool_from_stream(events: tuple[_BridgeEvent, ...]) -> AnthropicContentBlock:
starts: Final = tuple(
event
for event in events
if event.type == "content_block_start"
and event.content_block is not None
and event.content_block.type == "tool_use"
)
assert len(starts) == 1, "expected exactly one tool call"
start: Final = starts[0]
block: Final = start.content_block
assert block is not None and block.id and start.index is not None
fragments: Final = tuple(
event
for event in events
if event.type == "content_block_delta" and event.delta is not None and event.delta.type == "input_json_delta"
)
assert fragments, "tool stream contained no argument fragments"
assert all(event.index == start.index for event in fragments), "tool fragments changed index"
positions: Final = tuple(i for i, event in enumerate(events) if event in fragments)
stops: Final = tuple(
i for i, event in enumerate(events) if event.type == "content_block_stop" and event.index == start.index
)
assert len(stops) == 1 and events.index(start) < positions[0] <= positions[-1] < stops[0]
assert tuple(
event.delta.stop_reason for event in events if event.type == "message_delta" and event.delta is not None
) == ("tool_use",)
terminal_positions: Final = tuple(i for i, event in enumerate(events) if event.type == "message_delta")
assert len(terminal_positions) == 1 and stops[0] < terminal_positions[0] < len(events) - 1
assert tuple(i for i, event in enumerate(events) if event.type == "message_stop") == (len(events) - 1,), (
"tool stream did not terminate exactly once"
)
arguments: Final = _ParcelInput.model_validate_json(
"".join(event.delta.partial_json or "" for event in fragments if event.delta is not None)
)
return AnthropicContentBlock(type="tool_use", id=block.id, name=block.name, input=arguments.model_dump())
def _parcel_result(tool: AnthropicContentBlock, result: AnthropicToolResultBlock) -> AnthropicToolResultTurn:
assert tool.id and result.tool_use_id == tool.id, "tool result ID does not match the emitted call"
return AnthropicToolResultTurn(content=[result])
def _request_tool(
client: EndpointsClient, key: str, request: AnthropicMessagesBody, stream: bool
) -> AnthropicContentBlock:
if stream:
response: Final = client.proxy.messages_stream(key, request)
require_successful_call(response)
assert response.is_streaming and not response.stream_error
return _tool_from_stream(tuple(_BridgeEvent.model_validate_json(event) for event in response.stream_events))
response_body: Final = unwrap(client.proxy.messages(key, request))
blocks: Final = tuple(block for block in response_body.content or () if block.type == "tool_use")
assert len(blocks) == 1
return blocks[0]
class TestOpenAIMessagesToolContinuation:
@pytest.mark.parametrize("stream", [True, False], ids=["stream", "nonstream"])
def test_required_tool_arguments_and_correlated_result(
self, endpoints_client: EndpointsClient, resources: ResourceManager, stream: bool
) -> None:
model: Final = f"e2e-bridge-tool-{unique_marker()}"
base: Final = provider_edge_base("openai")
model_id: Final = endpoints_client.create_model(
model,
LiteLLMParamsBody(
model="openai/gpt-5.6", api_key="os.environ/OPENAI_API_KEY", api_base=f"{base}/v1" if base else None
),
)
resources.defer(lambda: endpoints_client.delete_model(model_id))
key: Final = resources.key(models=[model])
tool: Final = AnthropicCustomTool(
name="locate_parcel",
description="Look up the receipt for a parcel on a shelf. Return the receipt verbatim.",
input_schema=ToolInputSchema(
properties={"parcel": JsonSchemaProperty(type="string"), "shelf": JsonSchemaProperty(type="integer")},
required=["parcel", "shelf"],
),
)
question: Final = ChatMessage(
role="user",
content="Call locate_parcel with parcel exactly amber-kite and shelf exactly 7. After the tool result, reply with only the receipt returned by the tool.",
)
request: Final = AnthropicMessagesBody(
model=model,
max_tokens=2048,
messages=[question],
tools=[tool],
tool_choice=AnthropicToolChoice(type="tool", name=tool.name),
stream=stream,
)
emitted: Final = _request_tool(endpoints_client, key, request, stream)
assert emitted.id and emitted.name == "locate_parcel"
assert emitted.input == {"parcel": "amber-kite", "shelf": 7}, "required tool arguments were lost or changed"
receipt: Final = f"receipt-{unique_marker()}"
result_turn: Final = _parcel_result(emitted, AnthropicToolResultBlock(tool_use_id=emitted.id, content=receipt))
continuation: Final = unwrap(
endpoints_client.proxy.messages(
key,
AnthropicMessagesBody(
model=model,
max_tokens=2048,
tools=[tool],
tool_choice=AnthropicToolChoice(type="none"),
messages=[question, AnthropicAssistantTurn(content=[emitted]), result_turn],
),
)
)
answer: Final = "".join(block.text or "" for block in continuation.content or ())
assert answer.strip() == receipt, "continuation did not consume the correlated tool result"
assert all(block.type != "tool_use" for block in continuation.content or ())

View file

@ -283,10 +283,15 @@ class ChatToolResultTurn(BaseModel):
type ChatTurn = ChatMessage | ChatAssistantTurn | ChatToolResultTurn
class ChatStreamOptions(BaseModel):
include_usage: bool
class ChatBody(BaseModel):
model: str
messages: Sequence[ChatTurn]
stream: bool = False
stream_options: ChatStreamOptions | None = None
max_tokens: int | None = None
max_completion_tokens: int | None = None
temperature: float | None = None
@ -488,12 +493,18 @@ class AnthropicToolResultTurn(BaseModel):
type AnthropicMessage = ChatMessage | AnthropicAssistantTurn | AnthropicToolResultTurn
class AnthropicToolChoice(BaseModel):
type: Literal["auto", "any", "tool", "none"]
name: str | None = None
class AnthropicMessagesBody(BaseModel):
model: str
messages: list[AnthropicMessage]
max_tokens: int
stream: bool | None = None
tools: list[AnthropicTool] | None = None
tool_choice: AnthropicToolChoice | None = None
guardrails: list[str] | None = None
cache: dict[str, bool] | None = {"no-cache": True}

View file

@ -0,0 +1,113 @@
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass
from math import isclose
from typing import Final
from e2e_config import provider_edge_base, unique_marker
from e2e_http import unwrap
from lifecycle import ResourceManager
from models import ChatBody, ChatMessage, ChatResponse, KeyGenerateBody, LiteLLMParamsBody, TeamNewBody
from spend_e2e_client import SpendClient
INPUT_RATE: Final = 0.00004
OUTPUT_RATE: Final = 0.00008
@dataclass(frozen=True)
class TeamTraffic:
team_id: str
key: str
responses: tuple[ChatResponse, ...]
@property
def prompt_tokens(self) -> int:
return sum(response.usage.prompt_tokens or 0 for response in self.responses if response.usage)
@property
def completion_tokens(self) -> int:
return sum(response.usage.completion_tokens or 0 for response in self.responses if response.usage)
@property
def spend(self) -> float:
return self.prompt_tokens * INPUT_RATE + self.completion_tokens * OUTPUT_RATE
def create_traffic(client: SpendClient, resources: ResourceManager) -> tuple[TeamTraffic, ...]:
base: Final = provider_edge_base("openai")
model: Final = f"e2e-reconciliation-{unique_marker()}"
model_id: Final = client.proxy.create_model(
model,
LiteLLMParamsBody(
model="openai/gpt-5.6-luna",
api_key="os.environ/OPENAI_API_KEY",
api_base=None if base is None else f"{base}/v1",
input_cost_per_token=INPUT_RATE,
output_cost_per_token=OUTPUT_RATE,
),
)
resources.defer(lambda: client.proxy.delete_model(model_id))
def team_traffic() -> TeamTraffic:
team: Final = client.proxy.create_team(TeamNewBody(team_alias=f"e2e-spend-{unique_marker()}"))
resources.defer(lambda: client.proxy.delete_team(team))
key: Final = client.proxy.generate_key(KeyGenerateBody(team_id=team, models=[model]))
resources.defer(lambda: client.proxy.delete_key(key))
prompts: Final = tuple(f"Reply with one word. {index} {unique_marker()}" for index in range(7))
def call(index: int) -> ChatResponse:
response: Final = unwrap(
client.proxy.chat(
key,
ChatBody(
model=model,
messages=[ChatMessage(role="user", content=prompts[index])],
max_completion_tokens=128,
),
)
)
assert response.id, "successful response must have an ID"
assert response.usage is not None, "successful response must have usage"
assert response.usage.prompt_tokens is not None and response.usage.prompt_tokens > 0
assert response.usage.completion_tokens is not None and response.usage.completion_tokens > 0
assert response.usage.total_tokens == response.usage.prompt_tokens + response.usage.completion_tokens
assert not response.usage.cache_creation_input_tokens
assert not response.usage.cache_read_input_tokens
assert not response.usage.prompt_tokens_details or not response.usage.prompt_tokens_details.cached_tokens
return response
sequential: Final = call(0)
with ThreadPoolExecutor(max_workers=6) as pool:
concurrent: Final = tuple(pool.map(call, range(1, 7)))
return TeamTraffic(team, key, (sequential, *concurrent))
return tuple(team_traffic() for _ in range(2))
def assert_logs_match(client: SpendClient, traffic: TeamTraffic) -> None:
expected_ids: Final = frozenset(response.id for response in traffic.responses)
assert len(expected_ids) == len(traffic.responses), "responses must have distinct IDs"
rows: Final = client.poll_logs_for_key(
traffic.key,
min_rows=len(traffic.responses),
predicate=lambda values: frozenset(row.request_id for row in values) == expected_ids,
)
assert frozenset(row.request_id for row in rows) == expected_ids, "stored IDs must equal returned response IDs"
assert len(rows) == len(traffic.responses), "expected exactly one scoped spend row per response"
by_id: Final = {row.request_id: row for row in rows}
def assert_response(response: ChatResponse) -> None:
row: Final = by_id[response.id]
usage: Final = response.usage
assert usage is not None and usage.prompt_tokens is not None and usage.completion_tokens is not None
assert row.team_id == traffic.team_id
assert row.status == "success"
assert row.cache_hit != "True"
assert row.prompt_tokens == usage.prompt_tokens
assert row.completion_tokens == usage.completion_tokens
assert row.total_tokens == usage.total_tokens
expected_cost: Final = usage.prompt_tokens * INPUT_RATE + usage.completion_tokens * OUTPUT_RATE
assert row.spend is not None and isclose(row.spend, expected_cost, rel_tol=1e-6, abs_tol=1e-9)
for response in traffic.responses:
assert_response(response)

View file

@ -17,13 +17,13 @@ fails the test; a pricing or token-count drift does not.
import time
from collections.abc import Callable
from concurrent.futures import ThreadPoolExecutor
from math import isclose
from typing import Final
import pytest
from e2e_http import Result, Success
from e2e_http import Success
from lifecycle import ResourceManager
from models import ChatResponse, LiteLLMParamsBody, SpendLogs, SpendLogsParams
from models import LiteLLMParamsBody, SpendLogs, SpendLogsParams
from spend_e2e_client import SpendClient, SpendLogRow, is_ok, unique_marker, unwrap
pytestmark = pytest.mark.e2e
@ -280,51 +280,22 @@ def test_key_spend_equals_sum_of_logs(client: SpendClient, scoped_key: str) -> N
), f"key aggregate {key_spend} != sum of logs {logs_total}; rows: {_summarize(rows)}"
@pytest.mark.replayable
@pytest.mark.covers("quota_management.spend_tracking.concurrent_burst.loses_no_spend")
def test_burst_of_concurrent_calls_loses_no_spend(
client: SpendClient, scoped_key: str
client: SpendClient, resources: ResourceManager
) -> None:
"""Six concurrent calls on one key: every call lands its own spend row under a
distinct request_id and the key aggregate equals the sum of the rows.
Sequential accuracy is covered by test_key_spend_equals_sum_of_logs; this pins
the concurrent increment path (parallel writers racing on one key's counter),
where a lost update can never be reproduced by sequential calls."""
burst = 6
from spend_reconciliation import TeamTraffic, assert_logs_match, create_traffic
def call(idx: int) -> Result[ChatResponse]:
return client.chat(
scoped_key,
"gemini-2.5-flash",
f"burst call {idx} {unique_marker()}",
max_tokens=16,
)
traffic: Final = create_traffic(client, resources)
with ThreadPoolExecutor(max_workers=burst) as pool:
results = tuple(pool.map(call, range(burst)))
failed = [r for r in results if not is_ok(r)]
assert not failed, f"{len(failed)}/{burst} burst calls failed; first: {failed[0]}"
def assert_team(team: TeamTraffic) -> None:
assert_logs_match(client, team)
key_spend: Final = client.poll_key_spend(team.key, minimum=team.spend * 0.999999)
assert isclose(key_spend, team.spend, rel_tol=1e-6, abs_tol=1e-9)
rows = client.poll_logs_for_key(
scoped_key,
min_rows=burst,
predicate=lambda rs: len([r for r in rs if (r.spend or 0) > 0]) >= burst,
)
costed = [r for r in rows if (r.spend or 0) > 0]
assert len(costed) >= burst, (
f"only {len(costed)}/{burst} burst calls produced a costed row - "
f"rows lost under concurrency: {_summarize(rows)}"
)
request_ids = [r.request_id for r in costed]
assert len(set(request_ids)) == len(request_ids), (
f"concurrent rows collapsed onto shared request_ids: {_summarize(rows)}"
)
logs_total = sum((r.spend or 0) for r in rows)
key_spend = client.poll_key_spend(scoped_key, minimum=logs_total * 0.999)
assert _approx_equal(key_spend, logs_total), (
f"key aggregate {key_spend} != sum of {len(rows)} rows {logs_total} - "
f"spend increments lost under concurrency: {_summarize(rows)}"
)
for team in traffic:
assert_team(team)
@pytest.mark.covers("quota_management.spend_tracking.pagination.keeps_total")

View file

@ -7,13 +7,18 @@ missing start/end dates are rejected.
from __future__ import annotations
import time
from datetime import datetime, timedelta, timezone
from math import isclose
from typing import Final
import pytest
from e2e_http import ProbeResult
from models import DateRangeParams
from lifecycle import ResourceManager
from proxy_client import Converged, await_converged
from pydantic import BaseModel
from spend_e2e_client import SpendClient
from spend_reconciliation import TeamTraffic, assert_logs_match, create_traffic
pytestmark = pytest.mark.e2e
@ -24,22 +29,45 @@ class TeamDailyActivityParams(BaseModel):
start_date: str | None = None
end_date: str | None = None
page: int = 1
page_size: int = 1
team_ids: str | None = None
class TeamDailyActivityRow(BaseModel):
date: str
metrics: TeamDailyActivityMetrics
breakdown: TeamDailyActivityBreakdown
class TeamDailyActivityMetrics(BaseModel):
spend: float
total_tokens: int
prompt_tokens: int
completion_tokens: int
api_requests: int
successful_requests: int
failed_requests: int
class TeamDailyActivityEntity(BaseModel):
metrics: TeamDailyActivityMetrics
class TeamDailyActivityBreakdown(BaseModel):
entities: dict[str, TeamDailyActivityEntity]
class TeamDailyActivityMetadata(BaseModel):
page: int
total_pages: int
has_more: bool
total_spend: float
total_prompt_tokens: int
total_completion_tokens: int
total_tokens: int
total_api_requests: int
total_successful_requests: int
total_failed_requests: int
class TeamDailyActivityResponse(BaseModel):
@ -47,32 +75,128 @@ class TeamDailyActivityResponse(BaseModel):
metadata: TeamDailyActivityMetadata
def _range_days(days: int) -> DateRangeParams:
end = datetime.now(timezone.utc).date()
start = end - timedelta(days=days)
return DateRangeParams(start_date=start.isoformat(), end_date=end.isoformat())
def _probe(client: SpendClient, params: BaseModel) -> ProbeResult:
return client.proxy.transport.probe(ROUTE, params=params)
class TestTeamDailyActivity:
@pytest.mark.replayable
@pytest.mark.covers("mgmt.team.daily_activity.happy_path")
@pytest.mark.parametrize("days", [1, 7, 30])
def test_valid_date_range_returns_results_and_metadata(self, client: SpendClient, days: int) -> None:
result = _probe(client, _range_days(days))
assert result.status_code == 200, (
f"{ROUTE} range={days}d must be 200, got {result.status_code}: {result.body[:600]}"
def test_valid_date_range_returns_results_and_metadata(
self, client: SpendClient, resources: ResourceManager
) -> None:
started: Final = datetime.now(timezone.utc).date()
traffic: Final = create_traffic(client, resources)
for team in traffic:
assert_logs_match(client, team)
ended: Final = datetime.now(timezone.utc).date()
team_ids: Final = ",".join(team.team_id for team in traffic)
def fetch(
page: int, start: str = (started - timedelta(days=1)).isoformat(), end: str = ended.isoformat()
) -> TeamDailyActivityResponse:
result: Final = _probe(
client,
TeamDailyActivityParams(
start_date=start,
end_date=end,
page=page,
page_size=1,
team_ids=team_ids,
),
)
assert result.status_code == 200, f"daily activity failed: {result.status_code} {result.body[:300]}"
return TeamDailyActivityResponse.model_validate_json(result.body)
def pages() -> tuple[TeamDailyActivityResponse, ...]:
first: Final = fetch(1)
assert first.metadata.total_pages <= len(traffic) * 2, "unexpected extra scoped daily groups"
return (first, *(fetch(page) for page in range(2, first.metadata.total_pages + 1)))
outcome: Final = await_converged(
pages,
converged=lambda values: (
sum(page.metadata.total_api_requests for page in values) >= sum(len(team.responses) for team in traffic)
),
timeout=client.proxy.poll_timeout,
interval=client.proxy.poll_interval,
now=time.monotonic,
sleep=time.sleep,
)
parsed = TeamDailyActivityResponse.model_validate_json(result.body)
assert parsed.metadata.page == 1
assert parsed.metadata.total_pages >= 1
if parsed.results:
first = parsed.results[0]
assert first.date
assert first.metrics.spend >= 0
assert first.metrics.total_tokens >= 0
observed: Final = outcome.result if isinstance(outcome, Converged) else outcome.last_result
assert observed is not None, "daily aggregation must return a response before the deadline"
assert len(observed) >= 2, "two teams must exercise a page boundary"
def assert_page(index: int, page: TeamDailyActivityResponse) -> None:
assert page.metadata.page == index
assert page.metadata.total_pages == len(observed)
assert page.metadata.has_more == (index < len(observed))
assert len(page.results) == 1, "each fetched daily group must appear in results"
row: Final = page.results[0]
assert started <= datetime.fromisoformat(row.date).date() <= ended
assert len(row.breakdown.entities) == 1
assert row.metrics.total_tokens == page.metadata.total_tokens
assert row.metrics.prompt_tokens == page.metadata.total_prompt_tokens
assert row.metrics.completion_tokens == page.metadata.total_completion_tokens
assert row.metrics.api_requests == page.metadata.total_api_requests
assert row.metrics.successful_requests == page.metadata.total_successful_requests
assert row.metrics.failed_requests == page.metadata.total_failed_requests
assert isclose(row.metrics.spend, page.metadata.total_spend, rel_tol=1e-6, abs_tol=1e-9)
for index, page in enumerate(observed, 1):
assert_page(index, page)
entities: Final = tuple(
(team_id, entity.metrics)
for page in observed
for row in page.results
for team_id, entity in row.breakdown.entities.items()
)
assert frozenset(team_id for team_id, _ in entities) == frozenset(team.team_id for team in traffic)
def assert_team(team: TeamTraffic) -> None:
metrics: Final = tuple(metrics for team_id, metrics in entities if team_id == team.team_id)
assert sum(m.api_requests for m in metrics) == len(team.responses)
assert sum(m.successful_requests for m in metrics) == len(team.responses)
assert sum(m.failed_requests for m in metrics) == 0
assert sum(m.prompt_tokens for m in metrics) == team.prompt_tokens
assert sum(m.completion_tokens for m in metrics) == team.completion_tokens
assert sum(m.total_tokens for m in metrics) == team.prompt_tokens + team.completion_tokens
assert isclose(sum(m.spend for m in metrics), team.spend, rel_tol=1e-6, abs_tol=1e-9)
for team in traffic:
assert_team(team)
assert isclose(
sum(page.metadata.total_spend for page in observed),
sum(team.spend for team in traffic),
rel_tol=1e-6,
abs_tol=1e-9,
)
assert sum(page.metadata.total_tokens for page in observed) == sum(
team.prompt_tokens + team.completion_tokens for team in traffic
)
for days in (7, 30):
assert (
tuple(fetch(page, (started - timedelta(days=days)).isoformat()) for page in range(1, len(observed) + 1))
== observed
), f"{days}-day activity must preserve the same isolated groups and totals"
empty_date: Final = (started - timedelta(days=7)).isoformat()
empty: Final = fetch(1, empty_date, empty_date)
assert empty.results == []
assert empty.metadata.total_pages == 0
assert empty.metadata.page == 1
assert not empty.metadata.has_more
assert empty.metadata.total_spend == 0
assert empty.metadata.total_tokens == 0
assert empty.metadata.total_api_requests == 0
assert empty.metadata.total_prompt_tokens == 0
assert empty.metadata.total_completion_tokens == 0
assert empty.metadata.total_successful_requests == 0
assert empty.metadata.total_failed_requests == 0
@pytest.mark.covers("mgmt.team.daily_activity.missing_start_date_rejected")
def test_missing_start_date_is_rejected(self, client: SpendClient) -> None:

View file

@ -6,12 +6,15 @@ including the logging handler, cost tracking, and WebSocket message processing.
"""
import json
from collections.abc import Sequence
from datetime import datetime
from unittest.mock import AsyncMock, Mock, patch, MagicMock
from typing import Dict, List, Any, Optional
import pytest
import httpx
import litellm
from typing_extensions import NotRequired, ReadOnly, TypedDict
# Add the parent directory to the system path
@ -22,10 +25,16 @@ from litellm.proxy.pass_through_endpoints.success_handler import (
PassThroughEndpointLogging,
)
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.utils import LlmProviders
from litellm.types.utils import CostBreakdown, LlmProviders, Usage
from litellm.proxy._types import UserAPIKeyAuth
class _LiveTurn(TypedDict):
prompt: ReadOnly[tuple[int, int]]
candidates: ReadOnly[tuple[int, int]]
candidate_audio_token_count_missing: NotRequired[ReadOnly[bool]]
class TestVertexAILivePassthroughLoggingHandler:
"""Test the Vertex AI Live Passthrough Logging Handler"""
@ -39,6 +48,7 @@ class TestVertexAILivePassthroughLoggingHandler:
"""Create a mock logging object"""
mock = MagicMock(spec=LiteLLMLoggingObj)
mock.model_call_details = {}
mock._response_cost_calculator.return_value = None
return mock
@pytest.fixture
@ -201,88 +211,490 @@ class TestVertexAILivePassthroughLoggingHandler:
assert text_prompt["tokenCount"] == 10
assert audio_prompt["tokenCount"] == 10
@patch(
"litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_ai_live_passthrough_logging_handler.get_model_info"
)
def test_calculate_cost_basic(self, mock_get_model_info, handler):
"""Test basic cost calculation"""
mock_get_model_info.return_value = {
"input_cost_per_token": 0.000001,
"output_cost_per_token": 0.000002,
}
def test_usage_carries_every_modality(self, handler):
"""Regression: the Usage object reported only TEXT, so audio and image billed as nothing.
prompt_tokens must be the full count and the details must name each modality,
because the cost calculator prices audio and image from *_tokens_details.
"""
usage_metadata = {
"promptTokenCount": 100,
"candidatesTokenCount": 50,
"totalTokenCount": 150,
}
cost = handler._calculate_live_api_cost("gemini-1.5-pro", usage_metadata)
# The cost calculation may include additional factors, so we check it's reasonable
expected_min_cost = (100 * 0.000001) + (50 * 0.000002)
assert cost >= expected_min_cost
assert cost > 0
@patch(
"litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_ai_live_passthrough_logging_handler.get_model_info"
)
def test_calculate_cost_with_audio(self, mock_get_model_info, handler):
"""Test cost calculation with audio tokens"""
mock_get_model_info.return_value = {
"input_cost_per_token": 0.000001,
"output_cost_per_token": 0.000002,
"input_cost_per_audio_token": 0.0001,
"output_cost_per_audio_token": 0.0002,
}
usage_metadata = {
"promptTokenCount": 100,
"candidatesTokenCount": 50,
"totalTokenCount": 150,
"promptTokenCount": 1300,
"candidatesTokenCount": 124,
"totalTokenCount": 1424,
"promptTokensDetails": [
{"modality": "TEXT", "tokenCount": 80},
{"modality": "AUDIO", "tokenCount": 20},
{"modality": "TEXT", "tokenCount": 13},
{"modality": "AUDIO", "tokenCount": 127},
{"modality": "IMAGE", "tokenCount": 1160},
],
"candidatesTokensDetails": [
{"modality": "TEXT", "tokenCount": 30},
{"modality": "AUDIO", "tokenCount": 20},
{"modality": "TEXT", "tokenCount": 29},
{"modality": "AUDIO", "tokenCount": 95},
],
}
cost = handler._calculate_live_api_cost("gemini-1.5-pro", usage_metadata)
usage = handler._create_usage_object_from_metadata(
usage_metadata=usage_metadata, model="gemini-live-2.5-flash"
)
# Should include both text and audio costs
assert cost > 0
assert cost > (100 * 0.000001) + (
50 * 0.000002
) # Should be higher due to audio
assert usage.prompt_tokens == 1300, "the full prompt count must survive, not just its text share"
assert usage.completion_tokens == 124
assert usage.prompt_tokens_details.text_tokens == 13
assert usage.prompt_tokens_details.audio_tokens == 127
assert usage.prompt_tokens_details.image_tokens == 1160
assert usage.completion_tokens_details.text_tokens == 29
assert usage.completion_tokens_details.audio_tokens == 95
@patch(
"litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_ai_live_passthrough_logging_handler.get_model_info"
def test_usage_sums_repeated_modality_entries(self, handler):
"""A modality can appear more than once across aggregated turns; sum, don't overwrite."""
usage = handler._create_usage_object_from_metadata(
usage_metadata={
"promptTokenCount": 40,
"candidatesTokenCount": 0,
"promptTokensDetails": [
{"modality": "IMAGE", "tokenCount": 10},
{"modality": "IMAGE", "tokenCount": 25},
{"modality": "TEXT", "tokenCount": 5},
],
},
model="gemini-live-2.5-flash",
)
assert usage.prompt_tokens_details.image_tokens == 35
assert usage.prompt_tokens_details.text_tokens == 5
NATIVE_AUDIO_MODEL = "gemini-live-2.5-flash-preview-native-audio-09-2025"
# A four-turn native-audio session. Google charges per turn for the whole session context
# window, so the prompt side repeats the accumulated audio while the candidates side reports
# only that turn's own response. The last turn names AUDIO and omits its tokenCount, which is
# the shape Live really emits at the end of a spoken answer.
AUDIO_SESSION: tuple[_LiveTurn, ...] = (
{"prompt": (14, 122), "candidates": (8, 20)},
{"prompt": (21, 182), "candidates": (5, 50)},
{"prompt": (24, 203), "candidates": (13, 27)},
{"prompt": (24, 203), "candidates": (0, 3), "candidate_audio_token_count_missing": True},
)
def test_calculate_cost_with_web_search(self, mock_get_model_info, handler):
"""Test cost calculation with web search (tool use)"""
mock_get_model_info.return_value = {
"input_cost_per_token": 0.000001,
"output_cost_per_token": 0.000002,
"web_search_cost_per_request": 0.01,
}
usage_metadata = {
"promptTokenCount": 100,
"candidatesTokenCount": 50,
"totalTokenCount": 150,
"toolUsePromptTokenCount": 10,
}
@staticmethod
def _live_messages(turns: Sequence[_LiveTurn]) -> list[dict[str, object]]:
"""Wrap (text, audio) prompt/candidate pairs as the server messages a Live session emits."""
return [{"type": "session.created", "session": {"id": "s"}}] + [
{
"type": "response.done",
"usageMetadata": {
"promptTokenCount": sum(turn["prompt"]),
"candidatesTokenCount": sum(turn["candidates"]),
"totalTokenCount": sum(turn["prompt"]) + sum(turn["candidates"]),
"promptTokensDetails": [
{"modality": "TEXT", "tokenCount": turn["prompt"][0]},
{"modality": "AUDIO", "tokenCount": turn["prompt"][1]},
],
"candidatesTokensDetails": (
[{"modality": "AUDIO"}]
if turn.get("candidate_audio_token_count_missing")
else [
{"modality": "TEXT", "tokenCount": turn["candidates"][0]},
{"modality": "AUDIO", "tokenCount": turn["candidates"][1]},
]
),
},
}
for turn in turns
]
cost = handler._calculate_live_api_cost("gemini-1.5-pro", usage_metadata)
@staticmethod
def _session_usage(
handler: VertexAILivePassthroughLoggingHandler,
mock_logging_obj: MagicMock,
messages: list[dict[str, object]],
model: str,
) -> Usage:
result = handler.vertex_ai_live_passthrough_handler(
websocket_messages=messages,
logging_obj=mock_logging_obj,
url_route="/vertex_ai/live",
start_time=datetime.now(),
end_time=datetime.now(),
request_body={},
model=model,
)
assert result["result"] is not None, "the handler must produce a usage-bearing response to bill"
return result["result"].usage
# Should include web search cost
expected_base_cost = (100 * 0.000001) + (50 * 0.000002)
# The web search cost might be handled differently, so just check it's reasonable
assert cost >= expected_base_cost
assert cost > 0
@classmethod
def _session_cost(
cls,
handler: VertexAILivePassthroughLoggingHandler,
mock_logging_obj: MagicMock,
messages: list[dict[str, object]],
model: str,
) -> float:
from litellm.cost_calculator import completion_cost
from litellm.types.utils import ModelResponse
usage = cls._session_usage(handler, mock_logging_obj, messages, model)
return completion_cost(
completion_response=ModelResponse(
id="x", object="chat.completion", created=0, model=model, usage=usage, choices=[]
),
model=f"vertex_ai/{model}",
custom_llm_provider="vertex_ai",
call_type="acompletion",
)
@classmethod
def _expected_session_cost(cls, turns: Sequence[_LiveTurn]) -> float:
from litellm.utils import get_model_info
info = get_model_info(model=cls.NATIVE_AUDIO_MODEL, custom_llm_provider="vertex_ai")
return (
sum(turn["prompt"][0] for turn in turns) * info["input_cost_per_token"]
+ sum(turn["prompt"][1] for turn in turns) * info["input_cost_per_audio_token"]
+ sum(turn["candidates"][0] for turn in turns) * info["output_cost_per_token"]
+ sum(turn["candidates"][1] for turn in turns) * info["output_cost_per_audio_token"]
)
def test_every_turn_of_a_session_is_billed(self, handler, mock_logging_obj):
"""Google charges per turn for the whole context window, so every turn adds to the bill.
Billing one snapshot instead gives away all the other turns: on this session the
largest single turn is well under the session total, and its share of the audio is
priced 6x the text rate, so the gap is money rather than rounding.
"""
turns = self.AUDIO_SESSION[:3]
cost = self._session_cost(handler, mock_logging_obj, self._live_messages(turns), self.NATIVE_AUDIO_MODEL)
assert cost == pytest.approx(self._expected_session_cost(turns), rel=1e-9)
widest_single_turn = max(self._expected_session_cost([turn]) for turn in turns)
assert cost > widest_single_turn, "billing one snapshot drops every other turn of the session"
def test_audio_named_without_a_token_count_bills_at_the_audio_rate(self, handler, mock_logging_obj):
"""Live can name the modality carrying the rest of a turn and omit its tokenCount.
Reading the absent key as zero left those tokens inside candidatesTokenCount but outside
the breakdown, so the calculator charged real speech at the text output rate. At this
entry's rates the last turn's 3 audio tokens are $0.0000360 rather than $0.0000060.
"""
turns = self.AUDIO_SESSION
usage = self._session_usage(handler, mock_logging_obj, self._live_messages(turns), self.NATIVE_AUDIO_MODEL)
assert usage.completion_tokens_details.audio_tokens == 100, "the unpriced entry takes the turn's residual"
assert usage.completion_tokens_details.text_tokens == 26
assert usage.completion_tokens == 126
cost = self._session_cost(handler, mock_logging_obj, self._live_messages(turns), self.NATIVE_AUDIO_MODEL)
assert cost == pytest.approx(self._expected_session_cost(turns), rel=1e-9)
TOOL_USE_PER_TURN = (100, 250, 400)
def _grounded_messages(self):
"""The three-turn session again, with each turn's own toolUsePromptTokenCount attached."""
messages = self._live_messages(self.AUDIO_SESSION[:3])
head, turns = messages[0], messages[1:]
return [head] + [
{**message, "usageMetadata": {**message["usageMetadata"], "toolUsePromptTokenCount": tool_use}}
for message, tool_use in zip(turns, self.TOOL_USE_PER_TURN)
]
def test_server_side_tool_use_prompt_tokens_are_summed_over_the_session(self, handler, mock_logging_obj):
"""toolUsePromptTokenCount rode the unknown-key pass-through, so it took the first turn only.
Every other total beside it is summed across the session, and the first turn is the
smallest number in the series, so a grounded session logged far fewer tool-use tokens
than it used. This session's turns are deliberately distinct, so 750 can only come from
summing: first-turn selection gives 100, last-turn or max gives 400.
"""
grounded = self._grounded_messages()
usage = self._session_usage(handler, mock_logging_obj, grounded, self.NATIVE_AUDIO_MODEL)
assert usage.prompt_tokens_details.tool_use_tokens == sum(self.TOOL_USE_PER_TURN)
@staticmethod
def _grounding_frame(metadata: dict[str, object]) -> dict[str, object]:
"""One server frame carrying grounding metadata, the way Live reports it."""
return {"type": "response.done", "serverContent": {"groundingMetadata": metadata}}
def test_web_grounding_is_counted_so_it_can_be_billed(self, handler, mock_logging_obj):
"""Live reports grounding in the server frames and never in usageMetadata.
Nothing read those frames, so web_search_requests stayed unset and the cost path's only
trigger for the per-query grounding charge never fired. Google bills a grounded Live
prompt on top of its tokens, so the whole fee was missing from the bill.
"""
messages = [
self._grounding_frame(
{
"webSearchQueries": ["who won the 2026 world cup final"],
"groundingChunks": [{"web": {"uri": "https://example.com"}}],
}
),
*self._live_messages(self.AUDIO_SESSION[:1]),
]
usage = self._session_usage(handler, mock_logging_obj, messages, self.NATIVE_AUDIO_MODEL)
assert usage.prompt_tokens_details.web_search_requests == 1, "a grounded turn must report its query"
assert getattr(usage.prompt_tokens_details, "google_maps_grounding_requests", None) is None
def test_maps_grounding_is_counted_under_its_own_sku(self, handler, mock_logging_obj):
"""Maps grounding is a separate SKU from web search, so it needs its own counter.
A maps-only turn carries grounding chunks but no webSearchQueries, so counting queries
alone would report nothing and bill nothing.
"""
messages = [
self._grounding_frame({"groundingChunks": [{"maps": {"placeId": "abc123"}}]}),
*self._live_messages(self.AUDIO_SESSION[:1]),
]
usage = self._session_usage(handler, mock_logging_obj, messages, self.NATIVE_AUDIO_MODEL)
assert usage.prompt_tokens_details.google_maps_grounding_requests == 1
assert getattr(usage.prompt_tokens_details, "web_search_requests", None) is None
def test_an_ungrounded_session_reports_no_grounding(self, handler, mock_logging_obj):
"""The counters must stay absent when no tool ran, or every session pays a grounding fee."""
usage = self._session_usage(
handler, mock_logging_obj, self._live_messages(self.AUDIO_SESSION[:1]), self.NATIVE_AUDIO_MODEL
)
assert getattr(usage.prompt_tokens_details, "web_search_requests", None) is None
assert getattr(usage.prompt_tokens_details, "google_maps_grounding_requests", None) is None
def test_grounding_adds_its_query_fee_to_the_session_bill(self, handler, mock_logging_obj):
"""The counter only matters if it reaches the bill, so assert against the cost, not the field.
Same tokens either way: the difference between the two sessions is the grounding fee alone.
"""
turns = self.AUDIO_SESSION[:1]
plain = self._session_cost(handler, mock_logging_obj, self._live_messages(turns), self.NATIVE_AUDIO_MODEL)
grounded = self._session_cost(
handler,
mock_logging_obj,
[self._grounding_frame({"webSearchQueries": ["q"]}), *self._live_messages(turns)],
self.NATIVE_AUDIO_MODEL,
)
assert grounded > plain, "a grounded session must cost more than the same tokens ungrounded"
def _priced_logging_obj(self) -> LiteLLMLoggingObj:
"""A real logging object, since the session's price is handed to it turn by turn."""
logging_obj = LiteLLMLoggingObj(
model=self.NATIVE_AUDIO_MODEL,
messages=[],
stream=True,
call_type="pass_through_endpoint",
start_time=datetime.now(),
litellm_call_id="live-session",
function_id="live",
)
logging_obj.update_environment_variables(
model=self.NATIVE_AUDIO_MODEL,
user="u",
optional_params={},
litellm_params={},
call_type="pass_through_endpoint",
)
logging_obj.model_call_details["custom_llm_provider"] = "vertex_ai"
return logging_obj
def _billed_session(
self, handler: VertexAILivePassthroughLoggingHandler, messages: list[dict[str, object]]
) -> tuple[float, CostBreakdown]:
logging_obj = self._priced_logging_obj()
result = handler.vertex_ai_live_passthrough_handler(
websocket_messages=messages,
logging_obj=logging_obj,
url_route="/vertex_ai/live",
start_time=datetime.now(),
end_time=datetime.now(),
request_body={},
model=self.NATIVE_AUDIO_MODEL,
custom_llm_provider="vertex_ai",
)
assert result["result"] is not None, "the handler must produce a usage-bearing response to bill"
assert logging_obj.cost_breakdown is not None, "the session's price must reach the logging object"
return result["result"]._hidden_params["response_cost"], logging_obj.cost_breakdown
def test_each_grounded_turn_pays_its_own_query_fee(self, handler):
"""Google charges the grounding fee per grounded prompt, not per session.
Summing the session into one usage collapsed two grounded turns into one query, so the
second question was answered for free. The bill now grows by one fee per grounded turn.
"""
head, turn = self._live_messages(self.AUDIO_SESSION[:1])
grounding = self._grounding_frame({"webSearchQueries": ["q"]})
plain_cost, _ = self._billed_session(handler, [head, turn, turn])
one_cost, one_breakdown = self._billed_session(handler, [head, grounding, turn, turn])
two_cost, two_breakdown = self._billed_session(handler, [head, grounding, turn, grounding, turn])
fee = one_cost - plain_cost
assert fee > 0, "a grounded turn must cost more than the same tokens ungrounded"
assert two_cost - plain_cost == pytest.approx(2 * fee), "two grounded turns must pay the fee twice"
assert two_breakdown["total_cost"] == pytest.approx(two_cost)
assert two_breakdown["tool_usage_cost"] == pytest.approx(2 * one_breakdown["tool_usage_cost"])
def test_a_query_repeated_across_turns_is_reported_once_per_turn(self, handler):
"""The reported query count must agree with the bill, which charges every grounded turn.
The session usage collapsed duplicate query strings across turns while the price was
per turn, so two turns asking the same question paid two fees yet reported one query.
Duplicates within one turn still collapse, since that turn ran one search.
"""
head, turn = self._live_messages(self.AUDIO_SESSION[:1])
grounding = self._grounding_frame({"webSearchQueries": ["q"]})
logging_obj = self._priced_logging_obj()
result = handler.vertex_ai_live_passthrough_handler(
websocket_messages=[head, grounding, turn, grounding, turn],
logging_obj=logging_obj,
url_route="/vertex_ai/live",
start_time=datetime.now(),
end_time=datetime.now(),
request_body={},
model=self.NATIVE_AUDIO_MODEL,
custom_llm_provider="vertex_ai",
)
_, one_breakdown = self._billed_session(handler, [head, grounding, turn])
repeated_within_turn = handler._session_usage(
[head, self._grounding_frame({"webSearchQueries": ["q", "q"]}), turn], self.NATIVE_AUDIO_MODEL
)
assert result["result"].usage.prompt_tokens_details.web_search_requests == 2
assert logging_obj.cost_breakdown["tool_usage_cost"] == pytest.approx(2 * one_breakdown["tool_usage_cost"])
assert repeated_within_turn.prompt_tokens_details.web_search_requests == 1
def test_the_fixed_cost_margin_is_charged_once_per_session(self, handler):
"""A fixed cost margin is a flat per-request fee, and a Live session is one spend row.
Pricing each turn on its own applied the fixed margin per turn, so a two-turn session paid it
twice. The session now carries the fixed margin once no matter how many turns it billed.
"""
head, turn = self._live_messages(self.AUDIO_SESSION[:1])
grounding = self._grounding_frame({"webSearchQueries": ["q"]})
messages = [head, grounding, turn, grounding, turn]
plain_cost, _ = self._billed_session(handler, messages)
fixed_amount = 0.01
with patch.object(litellm, "cost_margin_config", {"vertex_ai": {"fixed_amount": fixed_amount}}):
margined_cost, breakdown = self._billed_session(handler, messages)
assert margined_cost - plain_cost == pytest.approx(
fixed_amount
), "a two-turn session must add the fixed margin once, not once per billed turn"
assert breakdown["margin_fixed_amount"] == pytest.approx(fixed_amount)
assert breakdown["margin_total_amount"] == pytest.approx(fixed_amount)
def test_reporting_tool_use_tokens_does_not_move_the_bill(self, handler, mock_logging_obj):
"""Deliberate boundary: these tokens are reported here, and priced nowhere.
generic_cost_per_token reads the input bill out of prompt_tokens_details, and falls
back to prompt_tokens only when the details carry no text or a cache hit overlaps them,
so adding tool-use tokens to prompt_tokens is worth nothing on an ordinary Live turn and
over-charges against the cache-overlap correction when it is not. Pricing them belongs
in the shared input-cost path, beside the modality terms that already read the details.
"""
turns = self.AUDIO_SESSION[:3]
plain_cost = self._session_cost(handler, mock_logging_obj, self._live_messages(turns), self.NATIVE_AUDIO_MODEL)
grounded_cost = self._session_cost(
handler, mock_logging_obj, self._grounded_messages(), self.NATIVE_AUDIO_MODEL
)
assert plain_cost == pytest.approx(self._expected_session_cost(turns), rel=1e-9)
assert grounded_cost == pytest.approx(plain_cost, rel=1e-9), "reporting tool use must not move the bill"
def test_a_malformed_details_entry_does_not_cost_the_whole_session(self, handler, mock_logging_obj):
"""A ``*TokensDetails`` value that is not a list of objects must not take the session down.
The handler's only error path returns no result at all, so one odd frame used to throw
while reading it and the whole session billed nothing. The good turns still bill.
"""
turns = self.AUDIO_SESSION[:3]
messages = self._live_messages(turns)
mangled = [dict(message) for message in messages]
mangled[1]["usageMetadata"] = {**mangled[1]["usageMetadata"], "promptTokensDetails": "TEXT"}
usage = self._session_usage(handler, mock_logging_obj, mangled, self.NATIVE_AUDIO_MODEL)
surviving = turns[1:]
assert usage.prompt_tokens_details.audio_tokens == sum(turn["prompt"][1] for turn in surviving)
assert usage.prompt_tokens_details.text_tokens == sum(turn["prompt"][0] for turn in surviving)
assert usage.prompt_tokens == sum(sum(turn["prompt"]) for turn in turns), "the totals still cover every turn"
direct = handler._create_usage_object_from_metadata(
usage_metadata={
"promptTokenCount": 40,
"candidatesTokenCount": 12,
"promptTokensDetails": [{"modality": "AUDIO", "tokenCount": 40}, "AUDIO"],
"candidatesTokensDetails": {"modality": "TEXT", "tokenCount": 12},
},
model=self.NATIVE_AUDIO_MODEL,
)
assert direct.prompt_tokens_details.audio_tokens == 40, "the well-formed entry beside a bad one still counts"
assert direct.completion_tokens == 12
@pytest.mark.parametrize(
"label,prompt_details,candidate_details",
[
("text only", [("TEXT", 6)], [("TEXT", 2)]),
("audio in", [("TEXT", 13), ("AUDIO", 127)], [("TEXT", 18)]),
("image in", [("TEXT", 10), ("IMAGE", 258)], [("TEXT", 24)]),
("frames in", [("TEXT", 11), ("IMAGE", 1032)], [("TEXT", 26)]),
("audio both ways", [("TEXT", 13), ("AUDIO", 127)], [("TEXT", 29), ("AUDIO", 95)]),
],
)
def test_live_session_bills_each_modality_at_its_own_rate(self, handler, label, prompt_details, candidate_details):
"""Every payload here is a real Vertex Live session's usageMetadata.
Before the fix these billed the text share only, from 1x (text) to 55x under.
The expected amount is derived from the entry's own rates rather than hardcoded,
so this stays correct as prices move, and it is asserted exactly, so dropping a
modality and double-charging one both fail.
"""
from litellm.cost_calculator import completion_cost
from litellm.types.utils import ModelResponse
from litellm.utils import get_model_info
model = self.NATIVE_AUDIO_MODEL
info = get_model_info(model=model, custom_llm_provider="vertex_ai")
text_in = info["input_cost_per_token"]
audio_in = info.get("input_cost_per_audio_token") or text_in
image_in = info.get("input_cost_per_image_token") or text_in
text_out = info["output_cost_per_token"]
audio_out = info.get("output_cost_per_audio_token") or text_out
rate_in = {"TEXT": text_in, "AUDIO": audio_in, "IMAGE": image_in}
rate_out = {"TEXT": text_out, "AUDIO": audio_out}
expected = sum(c * rate_in[m] for m, c in prompt_details) + sum(c * rate_out[m] for m, c in candidate_details)
usage = handler._create_usage_object_from_metadata(
usage_metadata={
"promptTokenCount": sum(c for _, c in prompt_details),
"candidatesTokenCount": sum(c for _, c in candidate_details),
"promptTokensDetails": [{"modality": m, "tokenCount": c} for m, c in prompt_details],
"candidatesTokensDetails": [{"modality": m, "tokenCount": c} for m, c in candidate_details],
},
model=model,
)
cost = completion_cost(
completion_response=ModelResponse(
id="x", object="chat.completion", created=0, model=model, usage=usage, choices=[]
),
model=f"vertex_ai/{model}",
custom_llm_provider="vertex_ai",
call_type="acompletion",
)
assert cost == pytest.approx(expected, rel=1e-9), label
text_only = sum(c for m, c in prompt_details if m == "TEXT") * text_in + sum(
c for m, c in candidate_details if m == "TEXT"
) * text_out
if any(m != "TEXT" for m, _ in prompt_details + candidate_details) and audio_in != text_in:
assert cost > text_only, f"{label}: non-text modalities must add cost"
def test_vertex_ai_live_passthrough_handler_integration(
self, handler, mock_logging_obj, sample_websocket_messages
@ -376,6 +788,7 @@ class TestVertexAILivePassthroughIntegration:
"""Create a mock logging object"""
mock = MagicMock(spec=LiteLLMLoggingObj)
mock.model_call_details = {}
mock._response_cost_calculator.return_value = None
return mock
@patch(
@ -509,6 +922,7 @@ class TestVertexAILivePassthroughErrorHandling:
"""Create a mock logging object"""
mock = MagicMock(spec=LiteLLMLoggingObj)
mock.model_call_details = {}
mock._response_cost_calculator.return_value = None
return mock
def test_invalid_websocket_messages_format(self):
@ -540,25 +954,24 @@ class TestVertexAILivePassthroughErrorHandling:
result = handler._extract_usage_metadata_from_websocket_messages(messages)
assert result is None
@patch(
"litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_ai_live_passthrough_logging_handler.get_model_info"
)
def test_cost_calculation_with_missing_model_info(self, mock_get_model_info):
"""Test cost calculation when model info is missing"""
def test_usage_without_modality_details(self):
"""Older payloads carry only the totals; fall back to them rather than reporting zero."""
handler = VertexAILivePassthroughLoggingHandler()
# Mock missing model info
mock_get_model_info.return_value = {}
usage = handler._create_usage_object_from_metadata(
usage_metadata={
"promptTokenCount": 100,
"candidatesTokenCount": 50,
"totalTokenCount": 150,
},
model="unknown-model",
)
usage_metadata = {
"promptTokenCount": 100,
"candidatesTokenCount": 50,
"totalTokenCount": 150,
}
# Should not raise an exception, should return 0 or handle gracefully
cost = handler._calculate_live_api_cost("unknown-model", usage_metadata)
assert cost == 0.0
assert usage.prompt_tokens == 100
assert usage.completion_tokens == 50
assert usage.total_tokens == 150
assert usage.prompt_tokens_details.audio_tokens is None
assert usage.prompt_tokens_details.image_tokens is None
def test_handler_with_none_websocket_messages(self, mock_logging_obj):
"""Test handler with None websocket messages"""

View file

@ -11,7 +11,7 @@ import litellm
from unittest.mock import patch, MagicMock, AsyncMock
from create_mock_standard_logging_payload import create_standard_logging_payload
from litellm.types.utils import StandardLoggingPayload
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo
from litellm.types.router import Deployment, DeploymentTypedDict, LiteLLM_Params, ModelInfo
from litellm.constants import DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS
@ -630,10 +630,12 @@ def test_deployment_callback_respects_cooldown_time(model_list):
@pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"])
def test_log_retry(model_list, metadata_key):
"""log_retry appends one flat record per failed attempt and copies neither the request kwargs nor
the request metadata into it"""
def test_log_retry(model_list: list[DeploymentTypedDict], metadata_key: str) -> None:
"""log_retry appends one flat record per failed attempt, copies neither the request kwargs nor the
request metadata into it, counts every failed attempt of the request independently of the
per-hop attempted_retries, and never trusts a negative count planted before the first failure"""
router = Router(model_list=model_list)
rate_limit_error = litellm.RateLimitError(message="slow down", llm_provider="openai", model="gpt-3.5-turbo")
new_kwargs = router.log_retry(
kwargs={
"model": "gpt-3.5-turbo",
@ -641,7 +643,7 @@ def test_log_retry(model_list, metadata_key):
"messages": [{"role": "user", "content": "hi"}],
metadata_key: {"model_info": {"id": "deployment-1"}, "attempted_retries": 2, "user_api_key": "sk-proxy"},
},
e=litellm.RateLimitError(message="slow down", llm_provider="openai", model="gpt-3.5-turbo"),
e=rate_limit_error,
)
assert json.loads(json.dumps(new_kwargs[metadata_key]["previous_models"])) == [
{
@ -652,6 +654,10 @@ def test_log_retry(model_list, metadata_key):
"attempted_retries": 2,
}
]
assert new_kwargs[metadata_key]["request_retry_count"] == 1
assert router.log_retry(kwargs=new_kwargs, e=rate_limit_error)[metadata_key]["request_retry_count"] == 2
planted_kwargs = {"model": "gpt-3.5-turbo", metadata_key: {"request_retry_count": -100}}
assert router.log_retry(kwargs=planted_kwargs, e=rate_limit_error)[metadata_key]["request_retry_count"] == 1
def test_update_usage(model_list):

View file

@ -13,6 +13,7 @@ import pytest
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.guardrail_translation.base_translation import StreamingScanKey
from litellm.llms.anthropic.chat.guardrail_translation.handler import (
AnthropicMessagesHandler,
@ -635,14 +636,19 @@ class TestAnthropicMessagesHandlerInputProcessing:
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
assert guardrail.inputs is not None
assert guardrail.inputs["texts"] == ["safe text", "prohibited correction"]
assert guardrail.inputs["texts"] == [
"trusted top-level system prompt",
"safe text",
"prohibited correction",
]
structured = guardrail.inputs["structured_messages"]
assert [m["role"] for m in structured] == ["system", "user", "system"]
assert structured[0]["content"] == "trusted top-level system prompt"
assert data["system"] == "trusted top-level system prompt"
assert data["messages"][1]["content"] == "[MASKED]"
@pytest.mark.asyncio
async def test_bedrock_masking_slice_is_unavailable_when_top_level_system_is_included(
async def test_bedrock_masking_slice_lines_up_when_top_level_system_is_included(
self,
):
from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
@ -668,25 +674,25 @@ class TestAnthropicMessagesHandlerInputProcessing:
structured = guardrail.inputs["structured_messages"]
bedrock = BedrockGuardrail(guardrailIdentifier="gi", guardrailVersion="1")
assert sum(bedrock._count_message_texts(m) for m in structured) == len(texts) + 1
assert sum(bedrock._count_message_texts(m) for m in structured) == len(texts)
latest_user_index = bedrock._find_latest_message_index(structured, target_role="user")
assert (
bedrock._locate_message_texts_slice(
structured_messages=structured,
target_index=latest_user_index,
texts=texts,
)
is None
)
assert (
bedrock._merge_masked_texts(
masked_texts=["{MASKED}"],
texts=texts,
scanned_slice=None,
scanned_role_subset=True,
)
== texts
scanned_slice = bedrock._locate_message_texts_slice(
structured_messages=structured,
target_index=latest_user_index,
texts=texts,
)
assert scanned_slice == (3, 1)
assert bedrock._merge_masked_texts(
masked_texts=["{MASKED}"],
texts=texts,
scanned_slice=scanned_slice,
scanned_role_subset=True,
) == [
"trusted top-level system prompt",
"safe text",
"prohibited correction",
"{MASKED}",
]
@pytest.mark.asyncio
@pytest.mark.parametrize("skip_system_message_in_guardrail", [True, None])
@ -1611,7 +1617,8 @@ class TestAnthropicMessagesIncrementalScan:
)
assert mock_api.call_count == 1
assert [m["content"] for m in mock_api.call_args.kwargs["messages"]] == [
"What is the capital of France?"
"You are a helpful geography assistant.",
"What is the capital of France?",
]
mock_api.reset_mock()
await handler.process_input_messages(
@ -2150,6 +2157,213 @@ class TestAnthropicMessagesScanOnlyToolResults:
assert guardrail.captured_inputs.get("images") == ["TOOL_IMG"]
class ToolCallArgumentsMaskingGuardrail(InputsRecordingGuardrail):
"""Masks the canary inside tool-call arguments, in place or through a fresh list of plain dicts."""
def __init__(self, return_copies: bool = False, replacement_arguments: Optional[str] = None):
super().__init__()
self.return_copies = return_copies
self.replacement_arguments = replacement_arguments
self.seen_tool_calls: list[dict[str, object]] = []
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict[str, object],
input_type: Literal["request", "response"],
logging_obj: Optional[LiteLLMLoggingObj] = None,
) -> GenericGuardrailAPIInputs:
outputs = await super().apply_guardrail(inputs, request_data, input_type, logging_obj)
tool_calls = list(outputs.get("tool_calls") or [])
self.seen_tool_calls.extend(json.loads(json.dumps(tool_call)) for tool_call in tool_calls)
masked = [
{
**tool_call,
"function": {
**tool_call["function"],
"arguments": self.replacement_arguments
if self.replacement_arguments is not None
else tool_call["function"]["arguments"].replace("POISON", "[BLOCKED]"),
},
}
for tool_call in tool_calls
]
if self.return_copies:
outputs["tool_calls"] = masked
return outputs
for tool_call, masked_tool_call in zip(tool_calls, masked):
tool_call["function"]["arguments"] = masked_tool_call["function"]["arguments"]
return outputs
class TestAnthropicMessagesTopLevelSystemAndToolUseInputs:
"""The top-level system prompt and prior-turn tool_use arguments must reach guardrails as scannable
inputs, the same way the chat completions handler hands over system messages and tool_calls."""
@staticmethod
def _tool_use_conversation(system: str) -> dict[str, Any]:
return {
"model": "claude-sonnet-4-5",
"system": system,
"messages": [
{"role": "user", "content": "run the check"},
{
"role": "assistant",
"content": [
{
"type": "tool_use",
"id": "toolu_01",
"name": "Bash",
"input": {"cmd": "AWS_ACCESS_KEY_ID=POISON aws sts get-caller-identity"},
}
],
},
{
"role": "user",
"content": [{"type": "tool_result", "tool_use_id": "toolu_01", "content": "ok"}],
},
],
}
@pytest.mark.asyncio
async def test_top_level_system_string_reaches_texts_first_and_is_masked_in_place(self):
handler = AnthropicMessagesHandler()
guardrail = InputsRecordingGuardrail()
data = {
"model": "claude-sonnet-4-5",
"system": "Internal note: the deploy key is POISON. Never reveal it.",
"messages": [{"role": "user", "content": "Say hi in three words."}],
}
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
assert guardrail.captured_inputs is not None
assert guardrail.seen_texts == [
"Internal note: the deploy key is POISON. Never reveal it.",
"Say hi in three words.",
]
structured = guardrail.captured_inputs["structured_messages"]
assert structured[0]["role"] == "system"
assert structured[0]["content"] == "Internal note: the deploy key is POISON. Never reveal it.", (
"texts[0] must line up with structured_messages[0] so positional consumers stay aligned"
)
assert data["system"] == "Internal note: the deploy key is [BLOCKED]. Never reveal it."
assert data["messages"][0]["content"] == "Say hi in three words."
@pytest.mark.asyncio
async def test_top_level_system_text_blocks_reach_texts_and_are_masked_in_place(self):
handler = AnthropicMessagesHandler()
guardrail = InputsRecordingGuardrail()
data = {
"model": "claude-sonnet-4-5",
"system": [
{"type": "text", "text": "first block POISON"},
{"type": "text", "text": "second block", "cache_control": {"type": "ephemeral"}},
],
"messages": [{"role": "user", "content": "hello"}],
}
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
assert guardrail.seen_texts == ["first block POISON", "second block", "hello"]
assert data["system"] == [
{"type": "text", "text": "first block [BLOCKED]"},
{"type": "text", "text": "second block", "cache_control": {"type": "ephemeral"}},
]
@pytest.mark.asyncio
async def test_skip_system_message_keeps_the_top_level_system_out(self):
handler = AnthropicMessagesHandler()
guardrail = InputsRecordingGuardrail()
guardrail.skip_system_message_in_guardrail = True
data = {
"model": "claude-sonnet-4-5",
"system": "trusted POISON prompt",
"messages": [{"role": "user", "content": "hello"}],
}
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
assert guardrail.seen_texts == ["hello"]
assert data["system"] == "trusted POISON prompt"
@pytest.mark.asyncio
async def test_prior_turn_tool_use_input_reaches_tool_calls_in_openai_shape(self):
handler = AnthropicMessagesHandler()
guardrail = InputsRecordingGuardrail()
data = self._tool_use_conversation(system="You are a careful agent harness.")
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
assert guardrail.captured_inputs is not None
tool_calls = guardrail.captured_inputs.get("tool_calls")
assert tool_calls is not None and len(tool_calls) == 1
assert tool_calls[0]["id"] == "toolu_01"
assert tool_calls[0]["type"] == "function"
assert tool_calls[0]["function"]["name"] == "Bash"
assert json.loads(tool_calls[0]["function"]["arguments"]) == {
"cmd": "AWS_ACCESS_KEY_ID=POISON aws sts get-caller-identity"
}
assert data["messages"][1]["content"][0]["input"] == {
"cmd": "AWS_ACCESS_KEY_ID=POISON aws sts get-caller-identity"
}, "a guardrail that leaves tool_calls alone must leave the tool_use input alone"
@pytest.mark.asyncio
@pytest.mark.parametrize("return_copies", [False, True])
async def test_masked_tool_call_arguments_write_back_into_the_tool_use_input(self, return_copies: bool):
handler = AnthropicMessagesHandler()
guardrail = ToolCallArgumentsMaskingGuardrail(return_copies=return_copies)
data = self._tool_use_conversation(system="You are a careful agent harness.")
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
assert [tool_call["function"]["name"] for tool_call in guardrail.seen_tool_calls] == ["Bash"]
tool_use = data["messages"][1]["content"][0]
assert tool_use == {
"type": "tool_use",
"id": "toolu_01",
"name": "Bash",
"input": {"cmd": "AWS_ACCESS_KEY_ID=[BLOCKED] aws sts get-caller-identity"},
}
assert data["messages"][2]["content"][0]["tool_use_id"] == "toolu_01"
@pytest.mark.asyncio
async def test_non_json_rewritten_arguments_are_rejected_by_name(self):
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
handler = AnthropicMessagesHandler()
guardrail = ToolCallArgumentsMaskingGuardrail(replacement_arguments="[REDACTED]")
data = self._tool_use_conversation(system="Internal note: the deploy key is POISON. Never reveal it.")
data["messages"][2]["content"][0]["content"] = "fetched POISON page"
original = json.loads(json.dumps(data))
with pytest.raises(UnappliableRequestRewrite) as excinfo:
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
assert excinfo.value.guardrail_name == "scan-only-capture"
assert data["system"] == original["system"], "a rejected rewrite must leave the request untouched"
assert data["messages"] == original["messages"], "a rejected rewrite must leave the request untouched"
@pytest.mark.asyncio
async def test_scan_only_tool_results_keeps_system_and_tool_use_out(self):
handler = AnthropicMessagesHandler()
guardrail = InputsRecordingGuardrail()
guardrail.scan_only_tool_results = True
data = self._tool_use_conversation(system="trusted POISON prompt")
data["messages"][2]["content"][0]["content"] = "fetched POISON page"
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
assert guardrail.seen_texts == ["fetched POISON page"]
assert guardrail.captured_inputs is not None
assert guardrail.captured_inputs.get("tool_calls") is None
assert data["system"] == "trusted POISON prompt"
assert data["messages"][1]["content"][0]["input"] == {
"cmd": "AWS_ACCESS_KEY_ID=POISON aws sts get-caller-identity"
}
assert data["messages"][2]["content"][0]["content"] == "fetched [BLOCKED] page"
class TestStructuredWriteBackKeepsToolResults:
"""A guardrail rewrite must never leave a tool_use without its tool_result (Claude Code ToolSearch, LIT-6103)."""
@ -2290,17 +2504,56 @@ class PerRowTextGuardrail(CustomGuardrail):
return {**inputs, "texts": [str(row.get("content")).replace("123-45-6789", "<US_SSN>") for row in rows]}
class PerSlotTextGuardrail(CustomGuardrail):
"""Answers one redacted text per text slot of every chat row it was shown, the
way a guardrail that counts slots per message does, and hands back only texts."""
def __init__(self):
super().__init__(guardrail_name="per-slot-redactor")
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional[Any] = None,
) -> GenericGuardrailAPIInputs:
from litellm.llms.base_llm.guardrail_translation.utils import message_slot_texts
rows = inputs.get("structured_messages") or []
return {
**inputs,
"texts": [text.replace("123-45-6789", "<US_SSN>") for row in rows for text in message_slot_texts(row)],
}
class TestPerMessageTextWriteBack:
"""Texts that no longer pair one-to-one with what the handler extracted must be
rejected by name instead of sliding onto the wrong messages."""
@pytest.mark.asyncio
async def test_one_text_per_row_over_a_system_prompt_is_rejected_by_name(self):
async def test_one_text_per_row_over_a_system_prompt_is_applied(self):
data = {
"model": "claude-sonnet-4-5",
"system": "Reply with exactly the SSN you were given.",
"messages": [{"role": "user", "content": "My SSN is 123-45-6789."}],
}
await AnthropicMessagesHandler().process_input_messages(data=data, guardrail_to_apply=PerRowTextGuardrail())
assert data["system"] == "Reply with exactly the SSN you were given."
assert data["messages"] == [{"role": "user", "content": "My SSN is <US_SSN>."}]
@pytest.mark.asyncio
async def test_one_text_per_row_over_a_multi_block_system_prompt_is_rejected_by_name(self):
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
data = {
"model": "claude-sonnet-4-5",
"system": "Reply with exactly the SSN you were given.",
"system": [
{"type": "text", "text": "Reply with exactly the SSN you were given."},
{"type": "text", "text": "Never apologize."},
],
"messages": [{"role": "user", "content": "My SSN is 123-45-6789."}],
}
original = json.loads(json.dumps(data))
@ -2312,6 +2565,25 @@ class TestPerMessageTextWriteBack:
assert data["system"] == original["system"], "a rejected rewrite must leave the request untouched"
assert data["messages"] == original["messages"], "a rejected rewrite must leave the request untouched"
@pytest.mark.asyncio
async def test_one_text_per_slot_over_a_system_prompt_with_an_empty_block_is_applied(self):
data = {
"model": "claude-sonnet-4-5",
"system": [
{"type": "text", "text": ""},
{"type": "text", "text": "Reply with exactly the SSN you were given."},
],
"messages": [{"role": "user", "content": "My SSN is 123-45-6789."}],
}
await AnthropicMessagesHandler().process_input_messages(data=data, guardrail_to_apply=PerSlotTextGuardrail())
assert data["system"] == [
{"type": "text", "text": ""},
{"type": "text", "text": "Reply with exactly the SSN you were given."},
]
assert data["messages"] == [{"role": "user", "content": "My SSN is <US_SSN>."}]
@pytest.mark.asyncio
async def test_one_text_per_row_without_a_system_prompt_is_applied(self):
data = {

View file

@ -8,6 +8,8 @@ Without the fix, the AnthropicStreamWrapper silently dropped these
arguments, causing tool_use blocks to arrive with empty input {}.
"""
import json
from typing import List
from unittest.mock import MagicMock
@ -139,9 +141,7 @@ async def test_async_stream_emits_input_json_delta_for_bundled_tool_args():
# Verify the delta carries the tool arguments
delta_event = events[input_json_delta_idx]
assert delta_event["delta"][
"partial_json"
], "input_json_delta should have non-empty partial_json"
assert json.loads(delta_event["delta"]["partial_json"]) == {"location": "Boston"}
@pytest.mark.asyncio
@ -300,7 +300,7 @@ def test_sync_stream_emits_input_json_delta_for_bundled_tool_args():
assert (
input_json_delta_idx == tool_start_idx + 1
), "input_json_delta should immediately follow the tool_use content_block_start"
assert events[input_json_delta_idx]["delta"]["partial_json"]
assert json.loads(events[input_json_delta_idx]["delta"]["partial_json"]) == {"location": "Boston"}
def test_sync_stream_no_extra_delta_when_tool_args_empty():

View file

@ -32,9 +32,14 @@ action.
import base64
import json
from datetime import datetime, timedelta, timezone
from types import MappingProxyType
from typing import Final
from unittest.mock import MagicMock, patch
import pytest
from pydantic import TypeAdapter
from litellm.llms.bedrock.base_aws_llm import WebIdentitySessionPolicy, _SessionPolicyStatement
# Actions the Claude Platform on AWS service is documented to call.
# Source: AWS IAM action reference + the #27678 surface area.
@ -49,9 +54,9 @@ _CLAUDE_PLATFORM_ACTIONS = {
}
def _captured_policy() -> dict:
"""Run _auth_with_web_identity_token under mocks + return the parsed
Policy dict that was actually sent to STS."""
def _captured_policy_document() -> str:
"""Run _auth_with_web_identity_token under mocks + return the Policy
JSON document that was actually sent to STS."""
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
base = BaseAWSLLM()
@ -84,11 +89,21 @@ def _captured_policy() -> dict:
mock_sts.assume_role_with_web_identity.assert_called_once()
kwargs = mock_sts.assume_role_with_web_identity.call_args.kwargs
policy_str = kwargs["Policy"]
return json.loads(policy_str)
return kwargs["Policy"]
def _statement_by_sid(policy: dict, sid: str) -> dict:
_SESSION_POLICY_ADAPTER: Final = TypeAdapter(WebIdentitySessionPolicy)
def _captured_policy() -> WebIdentitySessionPolicy:
return _SESSION_POLICY_ADAPTER.validate_python(json.loads(_captured_policy_document()))
def _granted_actions(policy: WebIdentitySessionPolicy) -> frozenset[str]:
return frozenset(action for stmt in policy["Statement"] for action in stmt["Action"])
def _statement_by_sid(policy: WebIdentitySessionPolicy, sid: str) -> _SessionPolicyStatement:
for stmt in policy["Statement"]:
if stmt.get("Sid") == sid:
return stmt
@ -102,7 +117,6 @@ class TestWebIdentitySessionPolicyShape:
def test_policy_parses_as_valid_iam_document(self):
policy = _captured_policy()
assert policy["Version"] == "2012-10-17"
assert isinstance(policy["Statement"], list)
assert len(policy["Statement"]) >= 2
def test_bedrock_statement_actions_preserved(self):
@ -137,16 +151,7 @@ class TestClaudePlatformActionsCovered:
@pytest.mark.parametrize("action", sorted(_CLAUDE_PLATFORM_ACTIONS))
def test_claude_platform_action_present(self, action: str):
policy = _captured_policy()
# Action may live in any Statement — search across all.
all_actions: set = set()
for stmt in policy["Statement"]:
stmt_actions = stmt.get("Action")
if isinstance(stmt_actions, str):
all_actions.add(stmt_actions)
elif isinstance(stmt_actions, list):
all_actions.update(stmt_actions)
assert action in all_actions, (
assert action in _granted_actions(_captured_policy()), (
f"{action} missing from session policy — "
f"bedrock/claude_platform/* requests will 403 on OIDC auth"
)
@ -179,15 +184,7 @@ class TestBedrockMantleActionsCovered:
action" even when the role's identity policy grants it."""
def test_bedrock_mantle_create_inference_present(self):
policy = _captured_policy()
all_actions: set = set()
for stmt in policy["Statement"]:
stmt_actions = stmt.get("Action")
if isinstance(stmt_actions, str):
all_actions.add(stmt_actions)
elif isinstance(stmt_actions, list):
all_actions.update(stmt_actions)
assert "bedrock-mantle:CreateInference" in all_actions, (
assert "bedrock-mantle:CreateInference" in _granted_actions(_captured_policy()), (
"bedrock-mantle:CreateInference missing from session policy — "
"bedrock_mantle/* requests will 403 on OIDC/WIF auth"
)
@ -233,7 +230,7 @@ class TestInvalidIdentityTokenSurfacesAudience:
operator can diagnose the mismatch without enabling LITELLM_LOG=DEBUG on a
prod instance."""
_AUD = "https://guidepoint.litellm-prod.ai"
_AUD = "https://gateway.example.com"
_ISS = "https://accounts.google.com"
_STS_MESSAGE = (
"An error occurred (InvalidIdentityToken) when calling the "
@ -308,3 +305,44 @@ class TestPolicyTransportConditions:
"ClaudePlatformLiteLLM must require aws:SecureTransport=true "
"to keep parity with the bedrock statement"
)
_STS_SESSION_POLICY_PLAINTEXT_LIMIT: Final = 2048
_BEDROCK_ROUTE_ACTIONS: Final = MappingProxyType(
{
"model/{model_id}/invoke": "bedrock:InvokeModel",
"model/{model_id}/invoke-with-response-stream": "bedrock:InvokeModelWithResponseStream",
"model/{model_id}/converse": "bedrock:InvokeModel",
"model/{model_id}/converse-stream": "bedrock:InvokeModelWithResponseStream",
"model/{model_id}/count-tokens": "bedrock:CountTokens",
"guardrail/{guardrail_id}/version/{version}/apply": "bedrock:ApplyGuardrail",
"rerank": "bedrock:Rerank",
"knowledgebases/{knowledge_base_id}/retrieve": "bedrock:Retrieve",
"knowledgebases": "bedrock:ListKnowledgeBases",
"agents/{agent_id}/agentAliases/{alias_id}/sessions/{session_id}/text": "bedrock:InvokeAgent",
"runtimes/{agent_runtime_arn}/invocations": "bedrock-agentcore:InvokeAgentRuntime",
"runtimes/{agent_runtime_arn}/invocations with X-Amzn-Bedrock-AgentCore-Runtime-User-Id": (
"bedrock-agentcore:InvokeAgentRuntimeForUser"
),
"mcp": "bedrock-agentcore:InvokeGateway",
}
)
class TestSessionPolicyGrantsEveryBedrockRoute:
"""LIT-7348: ``/rerank`` authorizes against ``bedrock:Rerank``, which the
ceiling never granted, so rerank 403d on web identity auth while static
credentials and IRSA worked. Each route the bedrock package signs with the
web identity session maps to the IAM action it authorizes against, and the
ceiling must grant every one of them."""
@pytest.mark.parametrize(("route", "action"), sorted(_BEDROCK_ROUTE_ACTIONS.items()))
def test_route_action_is_granted_by_the_ceiling(self, route: str, action: str):
assert action in _granted_actions(_captured_policy()), (
f"/{route} authorizes against {action}, which the session policy does not grant, "
"so it 403s on web identity auth"
)
def test_policy_document_fits_the_sts_plaintext_limit(self):
assert len(_captured_policy_document()) <= _STS_SESSION_POLICY_PLAINTEXT_LIMIT

View file

@ -1,4 +1,6 @@
import json
from collections.abc import Mapping
from typing import cast
from unittest.mock import MagicMock
import pytest
@ -6,6 +8,7 @@ import pytest
import litellm
from litellm.llms.gemini.realtime.transformation import GeminiRealtimeConfig
from litellm.types.llms.gemini import BidiGenerateContentServerMessage
def test_gemini_realtime_transformation_session_created():
@ -2178,3 +2181,71 @@ def test_unbilled_usage_on_session_close_flushes_trailing_audio(patch_gemini_tra
}
assert usage == expected
assert config.unbilled_usage_on_session_close("gemini-3.5-transcribe-live") is None
def _grounded_live_frame(grounding_metadata: Mapping[str, object] | None) -> Mapping[str, object]:
"""One Live server frame. Grounding metadata and usageMetadata arrive together, as Vertex sends them."""
from typing import Final
server_content: Final = {
"turnComplete": True,
**({} if grounding_metadata is None else {"groundingMetadata": grounding_metadata}),
}
return {
"serverContent": server_content,
"usageMetadata": {
"promptTokenCount": 19,
"candidatesTokenCount": 157,
"totalTokenCount": 176,
"promptTokensDetails": ({"modality": "TEXT", "tokenCount": 19},),
"candidatesTokensDetails": ({"modality": "AUDIO", "tokenCount": 157},),
},
}
def _response_done_input_details(message: Mapping[str, object]) -> Mapping[str, object]:
"""The ``input_tokens_details`` a ``response.done`` event carries, read off the emitted event."""
from typing import Final
config: Final = GeminiRealtimeConfig()
event: Final = config.transform_response_done_event(
message=cast( # cast-ok: a test fixture stands in for the server frame TypedDict
BidiGenerateContentServerMessage, message
),
current_response_id="resp_grounding",
current_conversation_id="conv_grounding",
output_items=None,
)
usage: Final = event["response"]["usage"]
assert usage, "response.done must carry a usage object"
return usage.get("input_tokens_details") or {}
def test_gemini_realtime_response_done_counts_web_grounding():
"""Regression: Live reports grounding in the server frames and never in usageMetadata.
Nothing read those frames on the realtime path, so web_search_requests stayed unset and the
cost path's only trigger for Google's per-query grounding charge never fired.
The counter is read off the emitted event, which is what the cost path is handed, so this covers
the grounding read and the usage bridge that carries it together
"""
input_details = _response_done_input_details(
_grounded_live_frame(
{
"webSearchQueries": ["who won the 2026 world cup final"],
"groundingChunks": [{"web": {"uri": "https://example.com"}}],
}
)
)
assert input_details.get("web_search_requests") == 1, "a grounded turn must report its query"
assert input_details.get("text_tokens") == 19, "the modality breakdown must survive alongside it"
def test_gemini_realtime_response_done_reports_no_grounding_when_none_ran():
"""The counter must stay unset on an ordinary turn, or every session pays a grounding fee."""
input_details = _response_done_input_details(_grounded_live_frame(None))
assert input_details.get("web_search_requests") is None
assert input_details.get("google_maps_grounding_requests") is None

View file

@ -252,11 +252,52 @@ class TestRender:
def test_savings_header_and_bars_against_the_routers_baseline(self, config_dir):
text = render("claude-sonnet-5", RECORDED, config_dir, use_color=False, bar_width=10)
assert text.splitlines() == [
"claude-auto · Routed to: claude-sonnet-5 -63% vs Claude Opus 5",
"LiteLLM ████░░░░░░ $0.14",
"Routed to: claude-sonnet-5 -63% vs Claude Opus 5",
"claude-auto ████░░░░░░ $0.14",
"Claude Opus 5 ██████████ $0.38",
]
def test_a_long_router_name_keeps_both_cost_bars_aligned(self, config_dir: Path) -> None:
session: Final = RECORDED._replace(router_name="engineering-smart-router")
text: Final = render("claude-sonnet-5", session, config_dir, use_color=False, bar_width=10)
assert text.splitlines()[1:] == [
"engineering-smart-router ████░░░░░░ $0.14",
"Claude Opus 5 ██████████ $0.38",
]
@pytest.mark.parametrize(
("router_name", "baseline_name", "router_padding", "baseline_padding"),
(
("路由-router", "Claude Opus 5", 3, 1),
("智能模型路由器", "Claude Opus 5", 1, 2),
("ABC-router", "Claude Opus 5", 1, 1),
("cafe\u0301-router", "Claude Opus 5", 3, 1),
("a\u20dd-router", "Claude Opus 5", 6, 1),
("カ\u3099-router", "Claude Opus 5", 5, 1),
("auto", "基準モデル", 7, 1),
("auto", "cafe\u0301", 1, 1),
),
)
@pytest.mark.parametrize("use_color", (False, True))
def test_unicode_labels_align_cost_bars_by_terminal_columns(
self,
config_dir: Path,
router_name: str,
baseline_name: str,
router_padding: int,
baseline_padding: int,
use_color: bool,
) -> None:
(config_dir / "cache" / "gateway-models.json").write_text(
json.dumps({"models": [{"id": "claude-opus-5", "display_name": baseline_name}]})
)
session: Final = RECORDED._replace(router_name=router_name)
text: Final = ANSI.sub("", render("claude-sonnet-5", session, config_dir, use_color, bar_width=10))
assert text.splitlines()[1:] == [
f"{router_name}{' ' * router_padding}████░░░░░░ $0.14",
f"{baseline_name}{' ' * baseline_padding}██████████ $0.38",
]
def test_control_characters_in_any_externally_sourced_label_never_reach_the_terminal(self, tmp_path, config_dir):
# The transcript, the proxy payload and Claude Code's model cache all feed labels straight into a
# terminal, and none is under this script's control. Only the control bytes are dropped (ESC, BEL,
@ -289,7 +330,7 @@ class TestRender:
assert "+25% vs Claude Opus 5" in render("m", dearer, config_dir, use_color=False)
def test_without_a_baseline_only_the_routed_line_shows(self, config_dir):
assert render("m", RECORDED._replace(baseline_model=None), config_dir, False) == "claude-auto · Routed to: m"
assert render("m", RECORDED._replace(baseline_model=None), config_dir, False) == "Routed to: m"
assert render("m", None, config_dir, False) == "Routed to: m"
def test_color_wraps_the_same_text(self, config_dir):
@ -311,7 +352,8 @@ class TestClaudeCodeMode:
return Fetched(RECORDED, definitive=True)
text: Final = _run(_payload(transcript), _env(tmp_path, config_dir), fetch)
assert text.startswith("claude-auto · Routed to: claude-sonnet-5 -63% vs Claude Opus 5\n")
assert text.startswith("Routed to: claude-sonnet-5 -63% vs Claude Opus 5\n")
assert text.splitlines()[1].startswith("claude-auto ")
def test_a_discovered_display_name_labels_the_sessions_model(
self, tmp_path: Path, transcript: Path, config_dir: Path
@ -322,7 +364,7 @@ class TestClaudeCodeMode:
return Fetched(session, definitive=True)
text: Final = _run(_payload(transcript), _env(tmp_path, config_dir), fetch)
assert text.startswith("claude-auto · Routed to: Claude Opus 5 -63% vs Claude Opus 5\n")
assert text.startswith("Routed to: Claude Opus 5 -63% vs Claude Opus 5\n")
def test_an_unrecorded_session_degrades_to_the_routed_line(self, tmp_path, transcript, config_dir):
assert _run(_payload(transcript), _env(tmp_path, config_dir), lambda c, s: Fetched(None, True)) == (
@ -378,7 +420,8 @@ class TestCodexMode:
out = _run({"hook_event_name": "Stop", "session_id": SESSION_ID, "transcript_path": "/nope"}, env, fetch)
message = json.loads(out)["systemMessage"]
assert message.splitlines()[1] == "claude-auto · Routed to: claude-sonnet-5 -63% vs Claude Opus 5"
assert message.splitlines()[1] == "Routed to: claude-sonnet-5 -63% vs Claude Opus 5"
assert message.splitlines()[2].startswith("claude-auto ")
assert message.startswith("\n")
assert seen == [Credentials("http://127.0.0.1:4000", "sk-codex")]

View file

@ -4620,46 +4620,27 @@ class TestPanwAirsLatestRoleMessageOnly:
@pytest.mark.asyncio
async def test_anthropic_system_plus_multiturn_no_fallback(self):
"""Anthropic with top-level system + multi-turn messages[]
— latest-user works, no scan-all fallback.
"""Anthropic with a top-level system prompt and multi-turn messages[]
scans only the latest user turn, with no scan-all fallback.
Key scenario: Anthropic top-level `system` field causes
structured_messages to have an injected system entry, but
request_data["messages"] does NOT include it.
The Anthropic handler hoists the top-level `system` field into both
`texts` and `structured_messages`, so the latest-user walk has to
count the same entries the framework flattened.
"""
handler = PanwPrismaAirsHandler(
guardrail_name="test_panw_airs",
api_key="test_api_key",
profile_name="test_profile",
default_on=True,
from litellm.llms.anthropic.chat.guardrail_translation.handler import (
AnthropicMessagesHandler,
)
# Original Anthropic messages (no system in messages array)
original_messages = [
{"role": "user", "content": "First user turn"},
{"role": "assistant", "content": "First assistant turn"},
{"role": "user", "content": "Latest user turn"},
]
# texts extracted from original_messages (3 text entries)
texts = ["First user turn", "First assistant turn", "Latest user turn"]
# structured_messages has an INJECTED system message from translation
structured_messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "First user turn"},
{"role": "assistant", "content": "First assistant turn"},
{"role": "user", "content": "Latest user turn"},
]
inputs: GenericGuardrailAPIInputs = {
"texts": texts,
"structured_messages": structured_messages,
}
handler = make_handler()
request_data = {
"litellm_call_id": "test-call-id",
"model": "anthropic/claude-sonnet-4-20250514",
"messages": original_messages,
"system": "You are a helpful assistant.",
"messages": [
{"role": "user", "content": "First user turn"},
{"role": "assistant", "content": "First assistant turn"},
{"role": "user", "content": "Latest user turn"},
],
"proxy_server_request": {
"url": "http://localhost:4000/v1/messages",
},
@ -4670,13 +4651,11 @@ class TestPanwAirsLatestRoleMessageOnly:
) as mock_api:
mock_api.return_value = {"action": "allow", "category": "benign"}
await handler.apply_guardrail(
inputs=inputs,
request_data=request_data,
input_type="request",
await AnthropicMessagesHandler().process_input_messages(
data=request_data,
guardrail_to_apply=handler,
)
# Should scan ONLY the latest user message, not fall back to scan-all
assert mock_api.call_count == 1
assert mock_api.call_args.kwargs["content"] == "Latest user turn"

View file

@ -5030,6 +5030,100 @@ async def test_websocket_passthrough_rewrites_gateway_alias_setup_model():
assert sent_setup["model"] == "projects/proj-db/locations/global/publishers/google/models/gemini-live-2.5-flash"
@pytest.mark.parametrize(
"setup_model",
["gemini-live-2.5-flash", "models/gemini-live-2.5-flash", "publishers/google/models/gemini-live-2.5-flash"],
)
def test_vertex_live_setup_model_resolves_before_extraction(setup_model):
"""A bare gateway alias left the session logged as ``unknown`` at zero cost.
The model was read off the raw client frame, and the extractor only yields a name when the string
already contains ``/models/``. The rewriter qualifies it a few lines later for the upstream, so a
client that addressed the gateway the documented way, by alias, logged no model and therefore
resolved no cost-map entry. Resolving first is what puts the real name on the logging object.
"""
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
_build_vertex_live_setup_model_rewriter,
)
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
_extract_model_from_vertex_ai_setup,
_resolved_vertex_live_setup,
)
rewriter = _build_vertex_live_setup_model_rewriter(
vertex_project="proj-db", vertex_location="global", llm_router=None
)
setup_data = {"model": setup_model}
resolved = _extract_model_from_vertex_ai_setup(_resolved_vertex_live_setup(setup_data, rewriter))
assert resolved == "gemini-live-2.5-flash", "an unresolved setup model logs the session as 'unknown'"
@pytest.mark.asyncio
async def test_websocket_passthrough_logs_a_bare_alias_setup_model():
"""End to end through the relay: a bare alias must reach the logging object as a real model name.
This is the call-site half of the fix. The helper tests above pass even if extraction moves back
before the rewrite, so this one drives the real websocket relay and asserts on what got logged,
which is the name the cost map is looked up by. An unbilled session logs ``unknown``.
"""
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
_build_vertex_live_setup_model_rewriter,
)
upstream_ws = RecordingUpstreamWebSocket()
setup_frame = json.dumps({"setup": {"model": "gemini-live-2.5-flash"}})
websocket = _client_websocket(
AsyncMock(
side_effect=[
{"type": "websocket.receive", "text": setup_frame},
{"type": "websocket.disconnect"},
]
)
)
built = []
real_logging = litellm.litellm_core_utils.litellm_logging.Logging
def _capture(*args, **kwargs):
obj = real_logging(*args, **kwargs)
built.append(obj)
return obj
with _patched_websocket_passthrough_environment(upstream_ws):
with patch("litellm.litellm_core_utils.litellm_logging.Logging", side_effect=_capture):
await websocket_passthrough_request(
websocket=websocket,
target="wss://aiplatform.googleapis.com/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent",
custom_headers={"Authorization": "Bearer token"},
user_api_key_dict=UserAPIKeyAuth(),
forward_headers=False,
endpoint="/vertex_ai/live",
accept_websocket=False,
setup_model_rewriter=_build_vertex_live_setup_model_rewriter(
vertex_project="proj-db", vertex_location="global", llm_router=None
),
)
assert built, "the relay should have built a logging object"
assert built[0].model == "gemini-live-2.5-flash", "a bare alias must not log as 'unknown'"
def test_vertex_live_setup_resolution_is_inert_without_a_rewriter():
"""Non-Live passthrough routes pass no rewriter, so the frame must be handed over untouched."""
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
_extract_model_from_vertex_ai_setup,
_resolved_vertex_live_setup,
)
setup_data = {"model": "projects/p/locations/global/publishers/google/models/gemini-live-2.5-flash"}
assert _resolved_vertex_live_setup(setup_data, None) is setup_data
assert _extract_model_from_vertex_ai_setup(_resolved_vertex_live_setup(setup_data, None)) == (
"gemini-live-2.5-flash"
)
@pytest.mark.asyncio
@pytest.mark.parametrize("rcvd_close", [None, "abnormal", "no_status"])
async def test_websocket_passthrough_does_not_relay_unsendable_upstream_close(rcvd_close):

View file

@ -7433,13 +7433,14 @@ def _reserved_stamp_key(key_metadata: dict | None = None) -> UserAPIKeyAuth:
_PLANTED_STAMPS = {
"attempted_fallbacks": 99,
"original_model_group": "spoofed-group",
"request_retry_count": -100,
"_client_output_ceiling": {"api_base": "https://attacker.example"},
"client_key": "client_value",
}
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_strips_router_reserved_stamps_from_both_buckets():
async def test_add_litellm_data_to_request_strips_router_reserved_stamps_from_both_buckets() -> None:
"""attempted_fallbacks and original_model_group are router-written facts the spend row
reads back; a client planting them in either bucket is dropped at the boundary so the
router never sees a reserved key it did not write."""
@ -7465,11 +7466,12 @@ async def test_add_litellm_data_to_request_strips_router_reserved_stamps_from_bo
assert "attempted_fallbacks" not in updated["metadata"]
assert "original_model_group" not in updated["metadata"]
assert "_client_output_ceiling" not in updated["metadata"]
assert "request_retry_count" not in updated["metadata"]
assert updated["metadata"]["client_key"] == "client_value"
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_strips_router_reserved_stamps_from_json_string_litellm_metadata():
async def test_add_litellm_data_to_request_strips_router_reserved_stamps_from_json_string_litellm_metadata() -> None:
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
data = {
@ -7490,11 +7492,12 @@ async def test_add_litellm_data_to_request_strips_router_reserved_stamps_from_js
assert "litellm_metadata" not in updated
assert "attempted_fallbacks" not in updated["metadata"]
assert "original_model_group" not in updated["metadata"]
assert "request_retry_count" not in updated["metadata"]
assert updated["metadata"]["client_key"] == "client_value"
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_strips_router_reserved_stamps_despite_pricing_override_opt_in():
async def test_add_litellm_data_to_request_strips_router_reserved_stamps_despite_pricing_override_opt_in() -> None:
"""The pricing strip is gated on allow_client_pricing_override; the reserved-stamp strip
is not, because no key or team setting makes a client-written fallback count valid."""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
@ -7518,6 +7521,7 @@ async def test_add_litellm_data_to_request_strips_router_reserved_stamps_despite
assert updated["metadata"]["model_info"] == {"input_cost_per_token": 0.0}
assert "attempted_fallbacks" not in updated["metadata"]
assert "original_model_group" not in updated["metadata"]
assert "request_retry_count" not in updated["metadata"]
@pytest.mark.asyncio

View file

@ -43,6 +43,7 @@ def test_litellm_settings_callback_list_strips_remote_urls(field):
"custom_auth",
"custom_key_generate",
"custom_key_update",
"custom_key_policy",
"custom_sso",
"custom_ui_sso_sign_in_handler",
],

View file

@ -7,6 +7,7 @@ from unittest.mock import MagicMock, patch
import pytest
import litellm
from litellm.models.credentials import CredentialItem
from litellm.realtime_api import main as realtime_main
from litellm.realtime_api.main import _with_resolved_session_model
@ -224,9 +225,11 @@ def test_client_secret_forwards_nested_transcription_model_untouched(monkeypatch
class _CapturingConnect:
def __init__(self) -> None:
self.url: str | None = None
self.kwargs: dict[str, object] = {}
def __call__(self, url: str, **kwargs: object) -> "_CapturingConnect":
self.url = url
self.kwargs = kwargs
return self
async def __aenter__(self) -> MagicMock:
@ -241,6 +244,72 @@ class _CapturingConnect:
return None
@pytest.mark.asyncio
async def test_azure_health_check_resolves_stored_credentials(monkeypatch):
monkeypatch.setattr(
litellm,
"credential_list",
[
CredentialItem(
credential_name="azure-rt",
credential_values={
"api_key": "sk-from-credential",
"api_base": "https://example.openai.azure.com",
"api_version": "2025-04-01-preview",
},
credential_info={},
)
],
)
connect = _CapturingConnect()
with patch("websockets.connect", connect):
assert await realtime_main._realtime_health_check(
model="gpt-realtime",
custom_llm_provider="azure",
api_key=None,
realtime_protocol="beta",
model_params={"model": "azure/gpt-realtime", "litellm_credential_name": "azure-rt"},
)
assert connect.kwargs["additional_headers"] == {"api-key": "sk-from-credential"}
assert connect.url is not None
assert connect.url.startswith("wss://example.openai.azure.com")
assert "api-version=2025-04-01-preview" in connect.url
@pytest.mark.asyncio
@pytest.mark.parametrize(
("custom_llm_provider", "model", "expected_url"),
[
("xai", "grok-voice-latest", "wss://api.x.ai/v1/realtime?model=grok-voice-latest"),
("openai", "gpt-realtime", "wss://api.openai.com/v1/realtime?model=gpt-realtime"),
],
)
async def test_bearer_health_check_sends_stored_credential_as_bearer_token(
monkeypatch, custom_llm_provider: str, model: str, expected_url: str
):
monkeypatch.setattr(
litellm,
"credential_list",
[
CredentialItem(
credential_name="voice-key",
credential_values={"api_key": "sk-from-credential"},
credential_info={},
)
],
)
connect = _CapturingConnect()
with patch("websockets.connect", connect):
assert await realtime_main._realtime_health_check(
model=model,
custom_llm_provider=custom_llm_provider,
api_key=None,
model_params={"model": f"{custom_llm_provider}/{model}", "litellm_credential_name": "voice-key"},
)
assert connect.kwargs["additional_headers"] == {"Authorization": "Bearer sk-from-credential"}
assert connect.url == expected_url
@pytest.mark.asyncio
async def test_azure_health_check_probes_ga_transcription_url_for_transcription_model(local_model_cost_map):
"""Regression for LIT-6240: transcription-only models (mode audio_transcription

View file

@ -8,7 +8,7 @@ from litellm.rust_bridge.lifecycle import check_limits
@pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"])
@pytest.mark.parametrize(
"cap, attempted_retries, refused",
"cap, request_retry_count, refused",
[(5, 5, True), (5, 4, False), (0, 0, False), (0, 1, True)],
ids=[
"cap-above-four-reached",
@ -17,12 +17,15 @@ from litellm.rust_bridge.lifecycle import check_limits
"cap-of-zero-refuses-first-retry",
],
)
def test_check_limits_reads_attempted_retries(
monkeypatch: pytest.MonkeyPatch, metadata_key: str, cap: int, attempted_retries: int, refused: bool
def test_check_limits_reads_request_retry_count(
monkeypatch: pytest.MonkeyPatch, metadata_key: str, cap: int, request_retry_count: int, refused: bool
) -> None:
monkeypatch.setattr(litellm, "num_retries_per_request", cap)
monkeypatch.setattr(litellm, "max_budget", None)
kwargs: Final = {"model": "mistral/mistral-ocr-latest", metadata_key: {"attempted_retries": attempted_retries}}
kwargs: Final = {
"model": "mistral/mistral-ocr-latest",
metadata_key: {"request_retry_count": request_retry_count},
}
if refused:
with pytest.raises(RuntimeError, match="Max retries per request hit!"):
check_limits(kwargs)

View file

@ -11066,6 +11066,66 @@ async def test_num_retries_per_request_stops_retries_at_caps_above_four(monkeypa
]
def _failing_group_with_healthy_fallback_router(num_retries: int) -> litellm.Router:
return litellm.Router(
model_list=[
{
"model_name": "broken-group",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "sk-fake",
"mock_response": "litellm.InternalServerError",
},
},
{
"model_name": "healthy-group",
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-fake", "mock_response": "ok"},
},
],
fallbacks=[{"broken-group": ["healthy-group"]}],
num_retries=num_retries,
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"cap, planted_count, hop_refused",
[(2, None, True), (4, None, False), (2, -100, True)],
ids=["cap-spent-before-the-hop", "cap-not-reached-by-the-hop", "planted-negative-count-does-not-lift-the-cap"],
)
async def test_num_retries_per_request_counts_retries_across_fallback_hops(
monkeypatch: pytest.MonkeyPatch, cap: int, planted_count: int | None, hop_refused: bool
) -> None:
"""num_retries_per_request caps the retries of one request, fallback hops included. Each hop starts a
fresh per-hop attempted_retries at zero, so a cap read from that counter let every hop retry from zero
and a request could spend far more retries than the cap allows. A caller who plants a negative count
in the request metadata must not push the cap further away either."""
monkeypatch.setattr(litellm, "num_retries_per_request", cap)
router = _failing_group_with_healthy_fallback_router(num_retries=1)
recorder = _FallbackAttemptRecorder()
litellm.callbacks.append(recorder)
try:
metadata = {} if planted_count is None else {"request_retry_count": planted_count}
request = router.acompletion(
model="broken-group", messages=[{"role": "user", "content": "hi"}], metadata=metadata
)
if not hop_refused:
assert (await request).choices[0].message.content == "ok"
return
with pytest.raises(litellm.InternalServerError):
await request
finally:
litellm.callbacks.remove(recorder)
assert recorder.failed_targets == ["healthy-group"]
hop_refusals = [
record["attempted_retries"]
for record in recorder.breadcrumbs_per_target[0]
if record["model_group"] == "healthy-group" and "Max retries per request hit!" in record["exception_string"]
]
assert hop_refusals == [0, 1]
@pytest.mark.asyncio
async def test_fallback_traceback_stays_available_at_debug_level():
"""Dropping the stack from the ERROR line is only safe because the fallback path still

View file

@ -4077,10 +4077,11 @@ class TestMetadataNoneHandling:
_RETRY_CAP_CASES: Final = (
pytest.param(5, {"attempted_retries": 5}, True, id="cap-above-four-reached"),
pytest.param(5, {"attempted_retries": 4}, False, id="cap-above-four-not-reached"),
pytest.param(0, {"attempted_retries": 0}, False, id="first-attempt-passes-cap-of-zero"),
pytest.param(0, {"attempted_retries": 1}, True, id="cap-of-zero-refuses-first-retry"),
pytest.param(5, {"request_retry_count": 5}, True, id="cap-above-four-reached"),
pytest.param(5, {"request_retry_count": 4}, False, id="cap-above-four-not-reached"),
pytest.param(0, {"request_retry_count": 0}, False, id="first-attempt-passes-cap-of-zero"),
pytest.param(0, {"request_retry_count": 1}, True, id="cap-of-zero-refuses-first-retry"),
pytest.param(0, {"attempted_retries": 1}, False, id="per-hop-attempted-retries-is-not-the-cap"),
pytest.param(5, {"previous_models": ("a", "b", "c", "d", "e")}, False, id="breadcrumb-count-is-not-the-cap"),
pytest.param(5, None, False, id="metadata-none"),
)
@ -4098,7 +4099,9 @@ def _capped_completion_kwargs(metadata_key: str, metadata: object) -> dict[str,
@pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"])
@pytest.mark.parametrize("cap, metadata, refused", _RETRY_CAP_CASES)
def test_num_retries_per_request_reads_attempted_retries_sync(monkeypatch, metadata_key, cap, metadata, refused):
def test_num_retries_per_request_reads_request_retry_count_sync(
monkeypatch: pytest.MonkeyPatch, metadata_key: str, cap: int, metadata: object, refused: bool
) -> None:
monkeypatch.setattr(litellm, "num_retries_per_request", cap)
kwargs: Final = _capped_completion_kwargs(metadata_key, metadata)
if refused:
@ -4111,7 +4114,9 @@ def test_num_retries_per_request_reads_attempted_retries_sync(monkeypatch, metad
@pytest.mark.asyncio
@pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"])
@pytest.mark.parametrize("cap, metadata, refused", _RETRY_CAP_CASES)
async def test_num_retries_per_request_reads_attempted_retries_async(monkeypatch, metadata_key, cap, metadata, refused):
async def test_num_retries_per_request_reads_request_retry_count_async(
monkeypatch: pytest.MonkeyPatch, metadata_key: str, cap: int, metadata: object, refused: bool
) -> None:
monkeypatch.setattr(litellm, "num_retries_per_request", cap)
kwargs: Final = _capped_completion_kwargs(metadata_key, metadata)
if refused:

View file

@ -806,7 +806,7 @@ async def test_shared_call_limits_still_reject_before_reading_ocr_file(
monkeypatch.setattr(litellm, "_current_cost", 2)
monkeypatch.setattr(litellm, "num_retries_per_request", 1 if limit == "retries" else None)
expected: Final = litellm.BudgetExceededError if limit == "budget" else RuntimeError
arguments: Final = {"document": {"type": "file", "file": File()}, "metadata": {"attempted_retries": 1}}
arguments: Final = {"document": {"type": "file", "file": File()}, "metadata": {"request_retry_count": 1}}
with pytest.raises(expected, match=r"Budget has been exceeded|Max retries per request hit"):
await call_aocr(ocr_server, **arguments) if asynchronous else call_ocr(ocr_server, **arguments)
assert reads == []