mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
merge: main into litellm_hotfix_7072_team_callbacks
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
commit
acba499e47
42 changed files with 3778 additions and 695 deletions
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
```
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 ())
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")]
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue