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

This commit is contained in:
mateo-berri 2026-09-01 18:02:02 -07:00
commit 6d8c18d518
27 changed files with 1639 additions and 270 deletions

View file

@ -80,7 +80,7 @@ jobs:
LITELLM_IMAGE: litellm-image-scan:${{ github.sha }}
run: |
python -m pip install "pytest==9.0.3"
python -m pytest tests/proxy_migration_tests/test_offline_image_migration.py -v
python -m pytest tests/proxy_migration_tests/test_offline_image_migration.py tests/proxy_migration_tests/test_image_bedrock_realtime_extra.py -v
# Scans the whole shipped artifact: OS/apk plus every language package
# baked into the image, including ones no lockfile declares (e.g. prisma's
@ -124,7 +124,7 @@ jobs:
LITELLM_IMAGE: litellm-runtime-scan:${{ github.sha }}
run: |
python -m pip install "pytest==9.0.3"
python -m pytest tests/proxy_migration_tests/test_offline_image_migration.py -v
python -m pytest tests/proxy_migration_tests/test_offline_image_migration.py tests/proxy_migration_tests/test_image_bedrock_realtime_extra.py -v
migrations-image:
name: migrations-image
@ -185,7 +185,7 @@ jobs:
LITELLM_COMPONENT_PORT: "4000"
run: |
python -m pip install "pytest==9.0.3"
python -m pytest tests/proxy_migration_tests/test_component_image_serves_offline.py -v
python -m pytest tests/proxy_migration_tests/test_component_image_serves_offline.py tests/proxy_migration_tests/test_image_bedrock_realtime_extra.py -v
ui-image:
name: ui-image

View file

@ -66,6 +66,7 @@ RUN uv sync --frozen --no-install-project --no-install-workspace --no-default-gr
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--extra bedrock-realtime \
--python python3.13
# Copy full source tree
@ -87,6 +88,7 @@ RUN uv sync --frozen --no-default-groups --no-editable \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--extra bedrock-realtime \
--python python3.13
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \

View file

@ -64,6 +64,7 @@ RUN uv sync --frozen --no-install-project --no-install-workspace --no-default-gr
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--extra bedrock-realtime \
--python python3.13
# Copy full source tree
@ -85,6 +86,7 @@ RUN uv sync --frozen --no-default-groups --no-editable \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--extra bedrock-realtime \
--python python3.13
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \

View file

@ -70,6 +70,7 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--extra bedrock-realtime \
--python python3.13
# Copy full source tree
@ -97,6 +98,7 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--extra bedrock-realtime \
--python python3.13 \
--no-sources-package litellm-proxy-extras; \
else \
@ -106,6 +108,7 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--extra bedrock-realtime \
--python python3.13; \
fi

View file

@ -26,7 +26,7 @@ If `db.useStackgresOperator` is used (not yet implemented):
| `replicaCount` | The number of LiteLLM Proxy pods to be deployed | `1` |
| `masterkeySecretName` | The name of the Kubernetes Secret that contains the Master API Key for LiteLLM. If not specified, use the generated secret name. | N/A |
| `masterkeySecretKey` | The key within the Kubernetes Secret that contains the Master API Key for LiteLLM. If not specified, use `masterkey` as the key. | N/A |
| `masterkey` | The Master API Key for LiteLLM. If not specified, a random key in the `sk-...` format is generated. | N/A |
| `masterkey` | The Master API Key for LiteLLM. If not specified, a random key in the `sk-...` format is generated on first install and reused on upgrades. | N/A |
| `environmentSecrets` | An optional array of Secret object names. The keys and values in these secrets will be presented to the LiteLLM proxy pod as environment variables. See below for an example Secret object. | `[]` |
| `environmentConfigMaps` | An optional array of ConfigMap object names. The keys and values in these configmaps will be presented to the LiteLLM proxy pod as environment variables. See below for an example Secret object. | `[]` |
| `image.repository` | LiteLLM Proxy image repository | `ghcr.io/berriai/litellm` |
@ -212,6 +212,8 @@ service, the **Proxy Endpoint** should be set to `http://<RELEASE>-litellm:4000`
The **Proxy Key** is the value specified for `masterkey` or, if a `masterkey`
was not provided to the helm command line, the `masterkey` is a randomly
generated string in the `sk-...` format stored in the `<RELEASE>-litellm-masterkey` Kubernetes Secret.
The key is generated once on the first install; later `helm upgrade` runs reuse the
value already in that Secret, so upgrading never rotates the master key.
```bash
kubectl -n litellm get secret <RELEASE>-litellm-masterkey -o jsonpath="{.data.masterkey}"

View file

@ -1,9 +1,11 @@
{{- if not .Values.masterkeySecretName }}
{{ $masterkey := (.Values.masterkey | default (printf "sk-%s" (randAlphaNum 18))) }}
{{- $secretName := printf "%s-masterkey" (include "litellm.fullname" .) }}
{{- $existing := lookup "v1" "Secret" .Release.Namespace $secretName }}
{{- $masterkey := .Values.masterkey | default (dig "data" "masterkey" "" $existing | b64dec) | default (printf "sk-%s" (randAlphaNum 18)) }}
apiVersion: v1
kind: Secret
metadata:
name: {{ include "litellm.fullname" . }}-masterkey
name: {{ $secretName }}
data:
masterkey: {{ $masterkey | b64enc }}
type: Opaque

View file

@ -15,6 +15,53 @@ tests:
# Note: The masterkey is generated as "sk-<18-random-chars>" in plain text,
# but stored as base64 encoded in Kubernetes secret (requirement).
# "sk-" base64 encodes to "c2st", so we check for "^c2st" pattern.
- it: should reuse the master key already stored in the cluster instead of generating a new one on upgrade
template: secret-masterkey.yaml
set:
masterkeySecretName: ""
kubernetesProvider:
scheme:
"v1/Secret":
gvr:
version: "v1"
resource: "secrets"
namespaced: true
objects:
- kind: Secret
apiVersion: v1
metadata:
name: RELEASE-NAME-litellm-masterkey
namespace: NAMESPACE
data:
masterkey: c2stZXhpc3Rpbmcta2V5
asserts:
- equal:
path: data.masterkey
value: c2stZXhpc3Rpbmcta2V5
- it: should let an explicit masterkey value override the one already stored in the cluster
template: secret-masterkey.yaml
set:
masterkeySecretName: ""
masterkey: sk-explicit
kubernetesProvider:
scheme:
"v1/Secret":
gvr:
version: "v1"
resource: "secrets"
namespaced: true
objects:
- kind: Secret
apiVersion: v1
metadata:
name: RELEASE-NAME-litellm-masterkey
namespace: NAMESPACE
data:
masterkey: c2stZXhpc3Rpbmcta2V5
asserts:
- equal:
path: data.masterkey
value: c2stZXhwbGljaXQ=
- it: should not create a secret if masterkeySecretName is set
template: secret-masterkey.yaml
set:

View file

@ -11,6 +11,7 @@ import json
import os
from collections.abc import Mapping, Sequence
from datetime import datetime
from types import MappingProxyType
from typing import Any, Final, Literal
import httpx
@ -30,12 +31,16 @@ from litellm.integrations.datadog.datadog_mock_client import (
)
from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.litellm_core_utils.prompt_templates.common_utils import (
convert_content_list_to_str,
handle_any_messages_to_chat_completion_str_messages_conversion,
)
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.proxy.spend_tracking.savings import extract_cache_creation_tokens, extract_cache_read_tokens
from litellm.types.integrations.datadog_llm_obs import *
from litellm.types.utils import (
CallTypes,
@ -44,6 +49,189 @@ from litellm.types.utils import (
StandardLoggingPayloadErrorInformation,
)
_EMPTY_MAPPING: Final[Mapping[str, Any]] = MappingProxyType({})
_EMPTY_MESSAGE: Final[Message] = {"role": "", "content": ""}
_MAX_PARSED_TOOL_ARGUMENT_CHARS: Final = 256 * 1024
def _mapping_field(source: Mapping[str, Any], key: str) -> Mapping[str, Any]:
"""The value at `key` when it is a mapping, else an empty one."""
value: Final = source.get(key)
return value if isinstance(value, dict) else _EMPTY_MAPPING
def _content_blocks(message: Mapping[str, Any]) -> tuple[Mapping[str, Any], ...]:
content: Final = message.get("content")
if not isinstance(content, list):
return ()
return tuple(block for block in content if isinstance(block, dict))
def _to_dd_arguments(raw_arguments: object) -> dict[str, Any] | str:
"""
Arguments as the object LLM Obs types them as, or the raw string when they are not one.
Strings past the size bound ship unparsed: decoding multiplies memory on hostile compact
JSON, and the raw string is what the intake receives either way.
"""
if not isinstance(raw_arguments, str):
return raw_arguments if isinstance(raw_arguments, dict) else str(raw_arguments)
if len(raw_arguments) > _MAX_PARSED_TOOL_ARGUMENT_CHARS:
return raw_arguments
parsed: Final = safe_json_loads(raw_arguments)
return parsed if isinstance(parsed, dict) else raw_arguments
def _to_dd_tool_calls(message: Mapping[str, Any]) -> tuple[ToolCall, ...]:
"""
The tool calls a message carries, in LLM Obs' ToolCall schema, from either dialect.
OpenAI puts them in `tool_calls` with the callee nested under `function` and `arguments`
serialized; Anthropic puts them in `content` as `tool_use` blocks with `input` already an
object. LLM Obs reads `name` / `arguments` / `tool_id` either way.
"""
raw_tool_calls: Final = message.get("tool_calls")
openai_calls: Final = tuple(
ToolCall(
name=function.get("name", ""),
arguments=_to_dd_arguments(function.get("arguments", "")),
tool_id=tool_call.get("id", ""),
type=tool_call.get("type", "function"),
)
for tool_call in (raw_tool_calls if isinstance(raw_tool_calls, list) else ())
if isinstance(tool_call, dict)
for function in [_mapping_field(tool_call, "function")]
)
anthropic_calls: Final = tuple(
ToolCall(
name=block.get("name", ""),
arguments=_to_dd_arguments(block.get("input") or {}),
tool_id=block.get("id", ""),
type="tool_use",
)
for block in _content_blocks(message)
if block.get("type") == "tool_use"
)
return openai_calls + anthropic_calls
def _to_dd_tool_results(message: Mapping[str, Any], tool_call_names: Mapping[str, str]) -> tuple[ToolResult, ...]:
"""
The tool results a message carries, linked back to the call each answers.
OpenAI models a result as a whole `role: "tool"` message keyed by `tool_call_id`;
Anthropic nests `tool_result` blocks inside a user message, keyed by `tool_use_id`.
"""
def to_result(tool_id: str, result: object) -> ToolResult:
return ToolResult(
name=tool_call_names.get(tool_id, ""),
result=result if isinstance(result, str) else safe_dumps(result),
tool_id=tool_id,
type="function",
)
if message.get("role") == "tool":
return (to_result(str(message.get("tool_call_id", "")), message.get("content") or ""),)
return tuple(
to_result(str(block.get("tool_use_id", "")), block.get("content") or "")
for block in _content_blocks(message)
if block.get("type") == "tool_result"
)
def _tool_call_names_by_id(messages: Sequence[object]) -> Mapping[str, str]:
"""Ids to tool names for result linking; reads names structurally and parses nothing."""
openai_pairs: Final = tuple(
(tool_call.get("id"), function.get("name", ""))
for message in messages
if isinstance(message, dict) and isinstance(message.get("tool_calls"), list)
for tool_call in message["tool_calls"]
if isinstance(tool_call, dict)
for function in [_mapping_field(tool_call, "function")]
)
anthropic_pairs: Final = tuple(
(block.get("id"), block.get("name", ""))
for message in messages
if isinstance(message, dict)
for block in _content_blocks(message)
if block.get("type") == "tool_use"
)
return MappingProxyType({str(tool_id): str(name) for tool_id, name in openai_pairs + anthropic_pairs if tool_id})
def _to_dd_message(message: object, tool_call_names: Mapping[str, str]) -> Message:
"""
Map one chat message onto LLM Obs' Message schema, adding fields and never destroying content.
Content collapses to its text only when it has text; a content list with none (tool blocks,
images) rides along unchanged so nothing the caller logged is lost. Tool calls and results
move into the fields the LLM Obs Tools panel reads, from both the OpenAI and Anthropic shapes.
"""
if not isinstance(message, dict):
converted: Final = handle_any_messages_to_chat_completion_str_messages_conversion(message)
return converted[0] if converted else _EMPTY_MESSAGE
text: Final = convert_content_list_to_str(message) # pyright: ignore[reportArgumentType] # caller-supplied dict
original_content: Final = message.get("content")
content: Final = (
text if text or not isinstance(original_content, list) or not original_content else original_content
)
reasoning: Final = message.get("reasoning_content")
tool_calls: Final = _to_dd_tool_calls(message)
tool_results: Final = _to_dd_tool_results(message, tool_call_names)
dd_message: Final[Message] = {
"role": message.get("role", ""),
"content": content,
**({"reasoning_content": reasoning} if reasoning is not None else {}),
**({"tool_calls": tool_calls} if tool_calls else {}),
**({"tool_results": tool_results} if tool_results else {}),
}
return dd_message
def _to_dd_messages(messages: object) -> tuple[Message, ...]:
"""Map a whole conversation, resolving each tool result against the calls that precede it."""
if messages is None:
return ()
if not isinstance(messages, list):
return tuple(handle_any_messages_to_chat_completion_str_messages_conversion(messages))
tool_call_names: Final = _tool_call_names_by_id(messages)
return tuple(_to_dd_message(message, tool_call_names) for message in messages)
def _to_dd_tool_definition(entry: Mapping[str, Any]) -> ToolDefinition | None:
function: Final = entry.get("function")
declared: Final[Mapping[str, Any]] = function if isinstance(function, dict) else entry
name: Final = declared.get("name")
if not name:
return None
schema: Final = declared.get("parameters") or declared.get("input_schema")
description: Final = declared.get("description", "")
if not isinstance(schema, dict):
return ToolDefinition(name=name, description=description)
return ToolDefinition(name=name, description=description, schema=schema)
def _to_dd_tool_definitions(model_parameters: object) -> tuple[ToolDefinition, ...]:
"""
Map the request's declared tools onto LLM Obs' ToolDefinition schema.
Handles the wrapped chat-completions shape and the bare shape the Anthropic and
Responses surfaces use, since both reach this logger through `model_parameters`.
"""
if not isinstance(model_parameters, dict):
return ()
raw_tools: Final = model_parameters.get("tools") or model_parameters.get("functions")
if not isinstance(raw_tools, list):
return ()
return tuple(
definition
for entry in raw_tools
if isinstance(entry, dict)
if (definition := _to_dd_tool_definition(entry)) is not None
)
class DataDogLLMObsLogger(CustomBatchLogger):
def __init__(self, **kwargs):
@ -222,12 +410,9 @@ class DataDogLLMObsLogger(CustomBatchLogger):
if standard_logging_payload is None:
raise Exception("DataDogLLMObs: standard_logging_object is not set")
messages = standard_logging_payload["messages"]
messages = self._ensure_string_content(messages=messages)
metadata: Final = kwargs.get("litellm_params", {}).get("metadata", {})
input_meta: Final = InputMeta(messages=handle_any_messages_to_chat_completion_str_messages_conversion(messages))
input_meta: Final = InputMeta(messages=_to_dd_messages(standard_logging_payload["messages"]))
output_meta: Final = OutputMeta(
messages=self._get_response_messages(
standard_logging_payload=standard_logging_payload,
@ -241,22 +426,20 @@ class DataDogLLMObsLogger(CustomBatchLogger):
if isinstance(metadata, dict):
metadata_parent_id = metadata.get("parent_id")
meta: Final = Meta(
kind=self._get_datadog_span_kind(standard_logging_payload.get("call_type"), metadata_parent_id),
input=input_meta,
output=output_meta,
metadata=self._get_dd_llm_obs_payload_metadata(standard_logging_payload),
error=error_info,
)
tool_definitions: Final = _to_dd_tool_definitions(standard_logging_payload.get("model_parameters"))
span_kind: Final = self._get_datadog_span_kind(standard_logging_payload.get("call_type"), metadata_parent_id)
payload_metadata: Final = self._get_dd_llm_obs_payload_metadata(standard_logging_payload)
# Calculate metrics (you may need to adjust these based on available data)
metrics: Final = LLMMetrics(
input_tokens=float(standard_logging_payload.get("prompt_tokens", 0)),
output_tokens=float(standard_logging_payload.get("completion_tokens", 0)),
total_tokens=float(standard_logging_payload.get("total_tokens", 0)),
total_cost=float(standard_logging_payload.get("response_cost", 0)),
time_to_first_token=self._get_time_to_first_token_seconds(standard_logging_payload),
)
meta: Final[Meta] = {
"kind": span_kind,
"input": input_meta,
"output": output_meta,
"metadata": payload_metadata,
"error": error_info,
**({"tool_definitions": tool_definitions} if tool_definitions else {}),
}
metrics: Final = self._assemble_metrics(standard_logging_payload)
payload: Final[LLMObsPayload] = LLMObsPayload(
parent_id=metadata_parent_id if metadata_parent_id else "undefined",
@ -314,6 +497,45 @@ class DataDogLLMObsLogger(CustomBatchLogger):
)
return error_info
def _assemble_metrics(self, standard_logging_payload: StandardLoggingPayload) -> LLMMetrics:
"""
Build the span metrics, including the prompt-cache counts LLM Obs charts cache savings from.
Cache counts resolve through the same owners the savings dashboard uses, so every provider
spelling is covered, and `non_cached_input_tokens` subtracts BOTH cache categories because
litellm's normalized prompt count includes both (the invariant the cost calculator's custom
pricing helper documents). A zero residual on a fully cached request is real data and is
emitted; a zero read or write count is absence and is not.
"""
prompt_tokens: Final = float(standard_logging_payload.get("prompt_tokens", 0))
completion_tokens: Final = float(standard_logging_payload.get("completion_tokens", 0))
total_tokens: Final = float(standard_logging_payload.get("total_tokens", 0))
total_cost: Final = float(standard_logging_payload.get("response_cost", 0))
time_to_first_token: Final = self._get_time_to_first_token_seconds(standard_logging_payload)
raw_usage: Final = (standard_logging_payload.get("metadata") or {}).get("usage_object")
usage_object: Final = raw_usage if isinstance(raw_usage, dict) else None
cache_read: Final = float(extract_cache_read_tokens(usage_object))
cache_write: Final = float(extract_cache_creation_tokens(usage_object))
metrics: Final[LLMMetrics] = {
"input_tokens": prompt_tokens,
"output_tokens": completion_tokens,
"total_tokens": total_tokens,
"total_cost": total_cost,
"time_to_first_token": time_to_first_token,
**(
{
**({"cache_read_input_tokens": cache_read} if cache_read else {}),
**({"cache_write_input_tokens": cache_write} if cache_write else {}),
"non_cached_input_tokens": max(prompt_tokens - cache_read - cache_write, 0.0),
}
if cache_read or cache_write
else {}
),
}
return metrics
def _get_time_to_first_token_seconds(self, standard_logging_payload: StandardLoggingPayload) -> float:
"""
Get the time to first token in seconds
@ -335,7 +557,7 @@ class DataDogLLMObsLogger(CustomBatchLogger):
def _get_response_messages(
self, standard_logging_payload: StandardLoggingPayload, call_type: str | None
) -> list[object]:
) -> tuple[Message, ...]:
"""
Get the messages from the response object
@ -344,7 +566,7 @@ class DataDogLLMObsLogger(CustomBatchLogger):
response_obj = standard_logging_payload.get("response")
if response_obj is None:
return []
return ()
# edge case: handle response_obj is a string representation of a dict
if isinstance(response_obj, str):
@ -357,7 +579,7 @@ class DataDogLLMObsLogger(CustomBatchLogger):
# fallback to json parsing
response_obj = json.loads(str(response_obj))
except json.JSONDecodeError:
return []
return ()
if call_type in [
CallTypes.completion.value,
@ -375,12 +597,12 @@ class DataDogLLMObsLogger(CustomBatchLogger):
if isinstance(response_obj, dict) and "choices" in response_obj:
choices: Final = response_obj["choices"]
if choices and len(choices) > 0 and "message" in choices[0]:
return [choices[0]["message"]]
return []
return _to_dd_messages([choices[0]["message"]])
return ()
except (KeyError, IndexError, TypeError):
# In case of any error accessing the response structure, return empty list
return []
return []
return ()
return ()
def _get_datadog_span_kind(
self, call_type: str | None, parent_id: str | None = None
@ -485,17 +707,6 @@ class DataDogLLMObsLogger(CustomBatchLogger):
# Default fallback for unknown or passthrough operations
return "llm"
def _ensure_string_content(self, messages: str | Sequence[object] | Mapping[object, object] | None) -> list[object]:
if messages is None:
return []
if isinstance(messages, str):
return [messages]
elif isinstance(messages, list):
return [message for message in messages]
elif isinstance(messages, dict):
return [str(messages.get("content", ""))]
return []
def _get_dd_llm_obs_payload_metadata(self, standard_logging_payload: StandardLoggingPayload) -> dict[str, object]:
"""
Fields to track in DD LLM Observability metadata from litellm standard logging payload
@ -524,10 +735,6 @@ class DataDogLLMObsLogger(CustomBatchLogger):
spend_metrics: Final = self._get_spend_metrics(standard_logging_payload)
_metadata.update({"spend_metrics": dict(spend_metrics)})
## extract tool calls and add to metadata
tool_call_metadata: Final = self._extract_tool_call_metadata(standard_logging_payload)
_metadata.update(tool_call_metadata)
_standard_logging_metadata: Final[dict] = dict(standard_logging_payload.get("metadata", {})) or {}
_metadata.update(_standard_logging_metadata)
return _metadata
@ -647,107 +854,3 @@ class DataDogLLMObsLogger(CustomBatchLogger):
verbose_logger.debug("Original value: %s", user_api_key_budget_reset_at)
return spend_metrics
def _process_input_messages_preserving_tool_calls(self, messages: Sequence[object]) -> list[dict[str, object]]:
"""
Process input messages while preserving tool_calls and tool message types.
This bypasses the lossy string conversion when tool calls are present,
allowing complex nested tool_calls objects to be preserved for Datadog.
"""
processed: Final = []
for msg in messages:
if isinstance(msg, dict):
# Preserve messages with tool_calls or tool role as-is
if "tool_calls" in msg or msg.get("role") == "tool":
processed.append(msg)
else:
# For regular messages, still apply string conversion
converted = handle_any_messages_to_chat_completion_str_messages_conversion([msg])
processed.extend(converted)
else:
# For non-dict messages, apply string conversion
converted = handle_any_messages_to_chat_completion_str_messages_conversion([msg])
processed.extend(converted)
return processed
@staticmethod
def _tool_calls_kv_pair(tool_calls: list[dict[str, Any]]) -> dict[str, object]:
"""
Extract tool call information into key-value pairs for Datadog metadata.
Similar to OpenTelemetry's implementation but adapted for Datadog's format.
"""
kv_pairs: Final[dict[str, object]] = {}
for idx, tool_call in enumerate(tool_calls):
try:
# Extract tool call ID
tool_id = tool_call.get("id")
if tool_id:
kv_pairs[f"tool_calls.{idx}.id"] = tool_id
# Extract tool call type
tool_type = tool_call.get("type")
if tool_type:
kv_pairs[f"tool_calls.{idx}.type"] = tool_type
# Extract function information
function = tool_call.get("function")
if function:
function_name = function.get("name")
if function_name:
kv_pairs[f"tool_calls.{idx}.function.name"] = function_name
function_arguments = function.get("arguments")
if function_arguments:
# Store arguments as JSON string for Datadog
if isinstance(function_arguments, str):
kv_pairs[f"tool_calls.{idx}.function.arguments"] = function_arguments
else:
import json
kv_pairs[f"tool_calls.{idx}.function.arguments"] = json.dumps(function_arguments)
except (KeyError, TypeError, ValueError) as e:
verbose_logger.debug("DataDogLLMObs: Error processing tool call %s: %s", idx, e)
continue
return kv_pairs
def _extract_tool_call_metadata(self, standard_logging_payload: StandardLoggingPayload) -> dict[str, object]:
"""
Extract tool call information from both input messages and response for Datadog metadata.
"""
tool_call_metadata: Final[dict[str, object]] = {}
try:
# Extract tool calls from input messages
messages: Final = standard_logging_payload.get("messages", [])
if messages and isinstance(messages, list):
for message in messages:
if isinstance(message, dict) and "tool_calls" in message:
tool_calls = message.get("tool_calls")
if tool_calls:
input_tool_calls_kv = self._tool_calls_kv_pair(tool_calls)
# Prefix with "input_" to distinguish from response tool calls
for key, value in input_tool_calls_kv.items():
tool_call_metadata[f"input_{key}"] = value
# Extract tool calls from response
response_obj: Final = standard_logging_payload.get("response")
if response_obj and isinstance(response_obj, dict):
choices: Final = response_obj.get("choices", [])
for choice in choices:
if isinstance(choice, dict):
message = choice.get("message")
if message and isinstance(message, dict):
tool_calls = message.get("tool_calls")
if tool_calls:
response_tool_calls_kv = self._tool_calls_kv_pair(tool_calls)
# Prefix with "output_" to distinguish from input tool calls
for key, value in response_tool_calls_kv.items():
tool_call_metadata[f"output_{key}"] = value
except Exception as e:
verbose_logger.debug("DataDogLLMObs: Error extracting tool call metadata: %s", e)
return tool_call_metadata

View file

@ -12,6 +12,7 @@ import asyncio
import json
import os
import random
import time
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from datetime import datetime, timezone
@ -154,18 +155,6 @@ class GetModelCostMap:
return True
@staticmethod
def fetch_remote_model_cost_map(url: str, timeout: int = 5) -> dict:
"""
Fetch the model cost map from a remote URL.
Returns the parsed JSON dict. Raises on network/parse errors
(caller is expected to handle).
"""
response: Final = httpx.get(url, timeout=timeout)
response.raise_for_status()
return response.json()
RETRYABLE_FETCH_STATUS_CODES: Final = frozenset({429, 500, 502, 503, 504})
MODEL_COST_MAP_FETCH_MAX_ATTEMPTS: Final = 3
@ -212,6 +201,13 @@ class _AsyncGetClient(Protocol):
def get(self, url: str, *, timeout: float | None = None) -> Awaitable[httpx.Response]: ...
class _SyncGetClient(Protocol):
def get(self, url: str, *, timeout: float | None = None) -> httpx.Response: ...
_FetchAttemptOutcome = ModelCostMapReloaded | ModelCostMapReloadUnavailable | _FetchAttemptRetryable
def _default_reload_client() -> _AsyncGetClient:
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.custom_http import httpxSpecialProvider
@ -219,13 +215,30 @@ def _default_reload_client() -> _AsyncGetClient:
return get_async_httpx_client(llm_provider=httpxSpecialProvider.ModelCostMap)
async def _attempt_fetch(
client: _AsyncGetClient, url: str, timeout: int
) -> ModelCostMapReloaded | ModelCostMapReloadUnavailable | _FetchAttemptRetryable:
def _classify_fetch_error(error: httpx.HTTPError | httpx.InvalidURL, url: str) -> _FetchAttemptOutcome:
reason: Final = f"{type(error).__name__} fetching {url}: {error}"
if isinstance(error, (httpx.InvalidURL, httpx.UnsupportedProtocol)):
return ModelCostMapReloadUnavailable(reason=reason)
return _FetchAttemptRetryable(reason=reason, retry_after_seconds=None)
async def _attempt_fetch(client: _AsyncGetClient, url: str, timeout: int) -> _FetchAttemptOutcome:
try:
response: Final = await client.get(url, timeout=timeout)
except httpx.HTTPError as e:
return _FetchAttemptRetryable(reason=f"{type(e).__name__} fetching {url}: {e}", retry_after_seconds=None)
except (httpx.HTTPError, httpx.InvalidURL) as e:
return _classify_fetch_error(e, url)
return _classify_fetch_response(response, url)
def _attempt_fetch_sync(client: _SyncGetClient, url: str, timeout: int) -> _FetchAttemptOutcome:
try:
response: Final = client.get(url, timeout=timeout)
except (httpx.HTTPError, httpx.InvalidURL) as e:
return _classify_fetch_error(e, url)
return _classify_fetch_response(response, url)
def _classify_fetch_response(response: httpx.Response, url: str) -> _FetchAttemptOutcome:
if response.status_code in RETRYABLE_FETCH_STATUS_CODES:
return _FetchAttemptRetryable(
reason=f"HTTP {response.status_code} from {url}",
@ -242,6 +255,22 @@ async def _attempt_fetch(
return ModelCostMapReloaded(model_cost_map=parsed)
def _next_retry_wait(
outcome: _FetchAttemptRetryable, attempt: int, max_attempts: int, rng: random.Random
) -> float | ModelCostMapReloadUnavailable:
if attempt == max_attempts:
return ModelCostMapReloadUnavailable(reason=f"{outcome.reason} (after {max_attempts} attempts)")
wait_seconds: Final = _retry_wait_seconds(outcome=outcome, attempt=attempt, rng=rng)
verbose_logger.warning(
"LiteLLM: model cost map fetch attempt %d/%d failed (%s); retrying in %.1fs",
attempt,
max_attempts,
outcome.reason,
wait_seconds,
)
return wait_seconds
async def _fetch_remote_model_cost_map_with_retry(
url: str,
timeout: int,
@ -254,20 +283,32 @@ async def _fetch_remote_model_cost_map_with_retry(
outcome = await _attempt_fetch(client=client, url=url, timeout=timeout)
if not isinstance(outcome, _FetchAttemptRetryable):
return outcome
if attempt == max_attempts:
return ModelCostMapReloadUnavailable(reason=f"{outcome.reason} (after {max_attempts} attempts)")
wait_seconds = _retry_wait_seconds(outcome=outcome, attempt=attempt, rng=rng)
verbose_logger.warning(
"LiteLLM: model cost map fetch attempt %d/%d failed (%s); retrying in %.1fs",
attempt,
max_attempts,
outcome.reason,
wait_seconds,
)
wait_seconds = _next_retry_wait(outcome=outcome, attempt=attempt, max_attempts=max_attempts, rng=rng)
if isinstance(wait_seconds, ModelCostMapReloadUnavailable):
return wait_seconds
await sleep(wait_seconds)
return ModelCostMapReloadUnavailable(reason="model cost map fetch failed")
def _fetch_remote_model_cost_map_with_retry_sync(
url: str,
timeout: int,
max_attempts: int,
sleep: Callable[[float], None],
rng: random.Random,
client: _SyncGetClient,
) -> ModelCostMapReloadResult:
for attempt in range(1, max_attempts + 1):
outcome = _attempt_fetch_sync(client=client, url=url, timeout=timeout)
if not isinstance(outcome, _FetchAttemptRetryable):
return outcome
wait_seconds = _next_retry_wait(outcome=outcome, attempt=attempt, max_attempts=max_attempts, rng=rng)
if isinstance(wait_seconds, ModelCostMapReloadUnavailable):
return wait_seconds
sleep(wait_seconds)
return ModelCostMapReloadUnavailable(reason="model cost map fetch failed")
async def refetch_model_cost_map(
url: str,
timeout: int = 5,
@ -423,13 +464,21 @@ def _finalize_model_cost_map(model_cost: dict) -> dict:
return _expand_model_aliases(model_cost)
def get_model_cost_map(url: str) -> dict:
def get_model_cost_map(
url: str,
timeout: int = 5,
max_attempts: int = MODEL_COST_MAP_FETCH_MAX_ATTEMPTS,
sleep: Callable[[float], None] = time.sleep,
rng: random.Random | None = None,
client: "_SyncGetClient | None" = None,
) -> dict:
"""
Public entry point — returns the model cost map dict.
1. If ``LITELLM_LOCAL_MODEL_COST_MAP`` is set, uses the local backup only.
2. Otherwise fetches from ``url``, validates integrity, and falls back
to the local backup on any failure.
2. Otherwise fetches from ``url``, retrying transient HTTP errors
(429/5xx/transport) with Retry-After-aware backoff, validates
integrity, and falls back to the local backup on any failure.
Only the backup model count is cached (a single int) for validation.
The full backup dict is only parsed when it must be *returned* as a
@ -448,17 +497,24 @@ def get_model_cost_map(url: str) -> dict:
_cost_map_source_info.url = url
_cost_map_source_info.is_env_forced = False
try:
content: Final = GetModelCostMap.fetch_remote_model_cost_map(url)
except Exception as e:
result: Final = _fetch_remote_model_cost_map_with_retry_sync(
url=url,
timeout=timeout,
max_attempts=max_attempts,
sleep=sleep,
rng=rng if rng is not None else random.Random(),
client=client if client is not None else httpx,
)
if isinstance(result, ModelCostMapReloadUnavailable):
verbose_logger.warning(
"LiteLLM: Failed to fetch remote model cost map from %s: %s. Falling back to local backup.",
url,
str(e),
result.reason,
)
_cost_map_source_info.source = "local"
_cost_map_source_info.fallback_reason = f"Remote fetch failed: {e}"
_cost_map_source_info.fallback_reason = f"Remote fetch failed: {result.reason}"
return _finalize_model_cost_map(GetModelCostMap.load_local_model_cost_map())
content: Final = result.model_cost_map
# Validate using cached count (cheap int comparison, no file I/O)
if not GetModelCostMap.validate_model_cost_map(

View file

@ -4957,10 +4957,13 @@ def make_valid_bedrock_tool_name(input_tool_name: str) -> str:
def add_cache_point_tool_block(tool: dict, model: str | None = None) -> BedrockToolBlock | None:
from litellm.llms.bedrock.common_utils import is_claude_4_5_on_bedrock
from litellm.llms.bedrock.common_utils import (
bedrock_model_accepts_cache_points,
is_claude_4_5_on_bedrock,
)
cache_control: Final = tool.get("cache_control", None)
if cache_control is not None:
if cache_control is not None and bedrock_model_accepts_cache_points(model):
cache_point: Final = cache_control.get("type", "ephemeral")
if cache_point == "ephemeral":
cache_point_block: Final[CachePointBlock] = {"type": "default"}

View file

@ -87,6 +87,7 @@ from ..common_utils import (
BedrockError,
BedrockModelInfo,
bedrock_converse_supports_parallel_tool_use_config,
bedrock_model_accepts_cache_points,
get_anthropic_beta_from_headers,
get_bedrock_tool_name,
is_bedrock_application_inference_profile_arn,
@ -1149,7 +1150,7 @@ class AmazonConverseConfig(BaseConfig):
model: str | None = None,
) -> SystemContentBlock | ContentBlock | None:
cache_control: Final = message_block.get("cache_control", None)
if cache_control is None:
if cache_control is None or not bedrock_model_accepts_cache_points(model):
return None
cache_point: Final = self._build_cache_point_block(cache_control, model)
@ -1613,7 +1614,7 @@ class AmazonConverseConfig(BaseConfig):
# Append cachePoint to tools if cache_control_injection_points has tool_config
cache_injection_points: Final = additional_request_params.pop("cache_control_injection_points", None)
if cache_injection_points and len(bedrock_tools) > 0:
if cache_injection_points and len(bedrock_tools) > 0 and bedrock_model_accepts_cache_points(model):
for point in cache_injection_points:
if point.get("location") == "tool_config":
cache_point = self._build_cache_point_block(point.get("control"), model)

View file

@ -816,6 +816,30 @@ def bedrock_converse_supports_parallel_tool_use_config(model: str) -> bool:
)
def bedrock_model_accepts_cache_points(model: str | None) -> bool:
"""
Whether Converse ``cachePoint`` blocks may be sent to this model.
Bedrock rejects requests carrying cachePoint blocks for models without prompt
caching support ("You invoked an unsupported model or your request did not allow
prompt caching"), so a model whose cost-map entry does not declare
``supports_prompt_caching`` must not receive them. A model absent from the map
(an application inference profile ARN, a model newer than the map) keeps emitting
so existing caching setups never silently degrade. ``litellm.utils.supports_prompt_caching``
is not reusable here: it returns False for unmapped models, the opposite polarity.
"""
if model is None:
return True
entries: Final = tuple(
entry
for candidate in (model, get_bedrock_base_model(model))
if (entry := litellm.model_cost.get(candidate)) is not None
)
if not entries:
return True
return any(entry.get("supports_prompt_caching") is True for entry in entries)
def is_claude_4_5_on_bedrock(model: str) -> bool:
"""
Check if the model supports Bedrock prompt caching with an extended '1h' TTL

View file

@ -3,6 +3,7 @@ import concurrent.futures
import contextlib
import os
import ssl
import sys
import typing
import urllib.request
from collections.abc import Callable, Generator
@ -75,10 +76,22 @@ except ImportError:
pass
def _current_task_is_cancelling() -> bool:
task: Final = asyncio.current_task()
if task is None or sys.version_info < (3, 11):
return True
return task.cancelling() > 0
@contextlib.contextmanager
def map_aiohttp_exceptions() -> Generator[None, None, None]:
try:
yield
except asyncio.CancelledError as exc:
# a closing connector cancels its shielded DNS task; that surfaces here without the request task being cancelled
if _current_task_is_cancelling():
raise
raise httpx.ConnectError("aiohttp transport cancelled the request internally") from exc
except Exception as exc:
mapped_exc: type[Exception] | None = None

View file

@ -1,7 +1,8 @@
import asyncio
import importlib
from collections.abc import Awaitable, Callable, Mapping
from collections.abc import Awaitable, Callable, Mapping, Sequence
from datetime import datetime
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal
import anyio
@ -20,8 +21,11 @@ from litellm.proxy._experimental.mcp_server.exceptions import (
MCPUpstreamAuthError,
)
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
ServerListOk,
ServerOutcome,
classify_list_exception,
list_fault_http_status,
outcome_wire_value,
)
from litellm.proxy._experimental.mcp_server.ui_session_utils import (
acting_user_auth,
@ -99,6 +103,7 @@ if MCP_AVAILABLE:
ListMCPToolsRestAPIResponseObject,
MCPInfo,
MCPServer,
_aggregate_server_key, # pyright: ignore[reportPrivateUsage] # same per-server key as the tools/list _meta outcomes
_apply_toolset_scope,
_fire_mcp_tool_call_logging,
execute_mcp_tool,
@ -803,9 +808,6 @@ if MCP_AVAILABLE:
list(allowed_server_ids_set), _rest_client_ip
)
list_tools_result: Final = []
error_message = None
# If server_id is specified, only query that specific server
if server_id:
return await _list_tools_for_single_server(
@ -849,22 +851,19 @@ if MCP_AVAILABLE:
else {}
)
# Query all servers the user has access to
errors: Final = []
for allowed_server_id in allowed_server_ids:
server = global_mcp_server_manager.get_mcp_server_by_id(allowed_server_id)
if server is None:
continue
server_auth_header = _get_server_auth_header(server, mcp_server_auth_headers, mcp_auth_header)
user_oauth_extra_headers = await _get_user_oauth_extra_headers(
async def list_server(
server: MCPServer,
) -> tuple[Sequence[ListMCPToolsRestAPIResponseObject], ServerOutcome]:
server_auth_header: Final = _get_server_auth_header(
server, mcp_server_auth_headers, mcp_auth_header
)
user_oauth_extra_headers: Final = await _get_user_oauth_extra_headers(
server,
user_api_key_dict,
prefetched_creds=prefetched_oauth_creds,
)
try:
tools_result = await _get_tools_for_single_server(
tools_result: Final = await _get_tools_for_single_server(
server,
server_auth_header,
raw_headers_from_request,
@ -872,24 +871,36 @@ if MCP_AVAILABLE:
extra_headers=user_oauth_extra_headers,
apply_tool_filters=apply_tool_filters,
)
list_tools_result.extend(tools_result)
except Exception as e:
verbose_logger.exception("Error getting tools from %s: %s", server.name, e)
errors.append(
f"{get_server_prefix(server)}: {classify_list_exception(e).tag}"
if isinstance(e, (MCPServerListError, MCPUpstreamAuthError))
else f"{get_server_prefix(server)}: {e}"
)
continue
return (), classify_list_exception(e)
return tools_result, ServerListOk(tool_count=len(tools_result))
if errors and not list_tools_result:
error_message = "Failed to get tools from servers: " + "; ".join(errors)
return {
"tools": list_tools_result,
"error": "partial_failure" if error_message else None,
"message": (error_message if error_message else "Successfully retrieved tools"),
}
# Query all servers the user has access to
queried_servers: Final = tuple(
server
for server in map(global_mcp_server_manager.get_mcp_server_by_id, allowed_server_ids)
if server is not None
)
listings: Final = tuple([await list_server(server) for server in queried_servers])
list_tools_result: Final = [tool for tools, _ in listings for tool in tools]
server_outcomes: Final = MappingProxyType(
{_aggregate_server_key(server): outcome for server, (_, outcome) in zip(queried_servers, listings)}
)
errors: Final = tuple(
f"{key}: {outcome.tag}" for key, outcome in server_outcomes.items() if outcome.tag != "ok"
)
error_message: Final = (
"Failed to get tools from servers: " + "; ".join(errors)
if errors and not list_tools_result
else None
)
return {
"tools": list_tools_result,
"error": "partial_failure" if error_message else None,
"message": (error_message if error_message else "Successfully retrieved tools"),
"server_outcomes": {key: outcome_wire_value(outcome) for key, outcome in server_outcomes.items()},
}
except MCPUpstreamAuthError as e:
# Surface upstream pass-through 401/403 challenges to the client so

View file

@ -585,6 +585,37 @@ async def _users_named_by_member_value(
return tuple(dict.fromkeys(row.user_id for row in rows))
async def _accounts_named_by_member_value(value: str, prisma_client: PrismaClient) -> tuple[str, ...]:
"""Every user id this member value names, by user id, SSO identity or email.
Classification needs to know whether the value is one account's ``user_id`` and
whether it names any other account, so all three fields are read in one pass. The
id is compared exactly and unstripped, as a primary key lookup would; the
identities compare as ``_users_named_by_member_value`` describes. Two rows are
enough to tell one account from several, so the read stops there. Only a full
read that lacks the row keyed by the value leaves that row's existence open, and
only then is the id read on its own.
"""
subject: Final = value.strip()
email: Final[_CaseInsensitiveMatch] = {"equals": subject, "mode": "insensitive"}
users: Final = _table(UserRepository(prisma_client))
rows: Final = await users.find_many(
where={ # mutable-ok: Prisma filter
"OR": [ # mutable-ok: Prisma filter
{"user_id": value}, # mutable-ok: Prisma filter
{"sso_user_id": subject}, # mutable-ok: Prisma filter
{"user_email": email}, # mutable-ok: Prisma filter
],
},
take=2,
)
named: Final = tuple(dict.fromkeys(row.user_id for row in rows))
if len(named) < 2 or value in named:
return named
keyed: Final = await users.find_unique(where={"user_id": value})
return named if keyed is None else (value, *named)
async def _classify_group_member(member: SCIMMember, prisma_client: PrismaClient) -> _ClassifiedGroupMember:
"""
Decide what a single SCIM group member refers to.
@ -627,11 +658,9 @@ async def _classify_group_member(member: SCIMMember, prisma_client: PrismaClient
if member_type == "group":
return _SkippedGroupMember(value=value, reason="nested_group")
user: Final = await _table(UserRepository(prisma_client)).find_unique(where={"user_id": value})
if user is not None:
shared_with: Final = tuple(
other for other in await _users_named_by_member_value(value, prisma_client) if other != value
)
named: Final = await _accounts_named_by_member_value(value, prisma_client)
if value in named:
shared_with: Final = tuple(other for other in named if other != value)
if shared_with:
verbose_proxy_logger.warning(
"SCIM: group member '%s' is one account's user id and is also account '%s' by SSO identity or email, "
@ -651,7 +680,6 @@ async def _classify_group_member(member: SCIMMember, prisma_client: PrismaClient
if team is not None and _team_metadata_has_scim_provenance(team.metadata):
return _SkippedGroupMember(value=value, reason="existing_team")
named: Final = await _users_named_by_member_value(value, prisma_client)
if len(named) == 1:
verbose_proxy_logger.info(
"SCIM: group member '%s' matched user_id '%s' by SSO identity or email",

View file

@ -22,7 +22,7 @@ import traceback
import weakref
from collections import defaultdict
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Generator, Mapping, Sequence
from functools import lru_cache
from functools import lru_cache, partial
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypeAlias, TypeVar, Union, cast
@ -825,6 +825,7 @@ class Router:
self._zero_cost_cache: dict[str, bool] = {}
self._routing_group_rows: tuple[DeploymentTypedDict, ...] | None = None
self._init_routing_groups(None)
self._provider_unresolved_deployments: tuple[Callable[[], Deployment | None], ...] = ()
self.deployment_affinity_ttl_seconds = deployment_affinity_ttl_seconds
self.model_group_affinity_config = model_group_affinity_config
@ -8473,6 +8474,19 @@ class Router:
return deployment
except Exception as e:
if self.ignore_invalid_deployments:
if isinstance(e, litellm.BadRequestError):
self._provider_unresolved_deployments = (
*self._provider_unresolved_deployments,
partial(
self._create_deployment,
deployment_info=deployment_info,
_model_name=_model_name,
_litellm_params=_litellm_params,
_model_info=_model_info,
declared_id=declared_id,
duplicate_ids=duplicate_ids,
),
)
verbose_router_logger.exception(
"Error creating deployment: %s, ignoring and continuing with other deployments.", e
)
@ -8902,6 +8916,7 @@ class Router:
self.quality_routers = {}
self.complexity_routers = {}
self.auto_routers = {}
self._provider_unresolved_deployments = ()
self._invalidate_model_group_info_cache()
self._invalidate_access_groups_cache()
# we add api_base/api_key each model so load balancing between azure/gpt on api_base1 and api_base2 works
@ -9524,8 +9539,12 @@ class Router:
"""Re-assert this router's deployments onto a freshly fetched catalog.
Reads ``model_list`` at call time, so only deployments the router still
serves are restored.
serves are restored, plus any config deployment the fresh catalog now resolves.
"""
provider_unresolved: Final = self._provider_unresolved_deployments
self._provider_unresolved_deployments = ()
for create_deployment in provider_unresolved:
create_deployment()
for entry in tuple(self.model_list):
try:
deployment = entry if isinstance(entry, Deployment) else Deployment(**entry)

View file

@ -4,21 +4,58 @@ Payloads for Datadog LLM Observability Service (LLMObs)
API Reference: https://docs.datadoghq.com/llm_observability/setup/api/?tab=example#api-standards
"""
from collections.abc import Sequence
from typing import Any, Literal
from typing_extensions import TypedDict
from typing_extensions import ReadOnly, TypedDict
from litellm.types.integrations.custom_logger import StandardCustomLoggerInitParams
class ToolCall(TypedDict, total=False):
"""A tool call on a message, as LLM Obs names its fields."""
name: ReadOnly[str]
arguments: ReadOnly[dict[str, Any] | str] # parsed object, or the raw string when it will not parse to one
tool_id: ReadOnly[str]
type: ReadOnly[str]
class ToolResult(TypedDict, total=False):
"""The result of a tool call, as LLM Obs names its fields."""
name: ReadOnly[str]
result: ReadOnly[str]
tool_id: ReadOnly[str]
type: ReadOnly[str]
class ToolDefinition(TypedDict, total=False):
"""A tool the model was offered on the request."""
name: ReadOnly[str]
description: ReadOnly[str]
schema: ReadOnly[dict[str, Any]]
class Message(TypedDict, total=False):
"""A message on a span, as LLM Obs names its fields."""
content: ReadOnly[str]
role: ReadOnly[str]
reasoning_content: ReadOnly[str]
tool_calls: ReadOnly[Sequence[ToolCall]]
tool_results: ReadOnly[Sequence[ToolResult]]
class InputMeta(TypedDict):
messages: list[
dict[str, Any] # changed to fit with tool calls
messages: Sequence[
Message | dict[str, Any] # changed to fit with tool calls
] # Relevant Issue: https://github.com/BerriAI/litellm/issues/9494
class OutputMeta(TypedDict):
messages: list[Any]
messages: Sequence[Any]
class DDLLMObsError(TypedDict, total=False):
@ -36,6 +73,7 @@ class Meta(TypedDict, total=False):
output: OutputMeta # The span's output information.
metadata: dict[str, Any]
error: DDLLMObsError | None # Error information on the span
tool_definitions: ReadOnly[Sequence[ToolDefinition]] # The tools offered to the model on this request
class LLMMetrics(TypedDict, total=False):
@ -45,6 +83,9 @@ class LLMMetrics(TypedDict, total=False):
time_to_first_token: float
time_per_output_token: float
total_cost: float
cache_read_input_tokens: ReadOnly[float]
cache_write_input_tokens: ReadOnly[float]
non_cached_input_tokens: ReadOnly[float]
class LLMObsPayload(TypedDict, total=False):

View file

@ -0,0 +1,58 @@
"""Image-level check that the built proxy image can import the Bedrock realtime SDK.
Bedrock Nova Sonic (`/v1/realtime`) imports `aws_sdk_bedrock_runtime` lazily on the
first session, so an image whose `uv sync` stages skip the `bedrock-realtime` extra
boots, passes health checks, and then fails every Nova Sonic session with
"Missing aws_sdk_bedrock_runtime". Importing inside the built image is what catches
that class of regression (missing extra, lockfile drift, a stage that syncs a
different set of extras), which a static Dockerfile check cannot.
Gated on LITELLM_IMAGE like the other image checks in this directory; exercised
where an image has been built (the image-scan workflow). Requires a working docker CLI.
"""
import os
import shutil
import subprocess
from typing import Final
import pytest
IMAGE: Final = os.getenv("LITELLM_IMAGE")
NON_ROOT_UID: Final = "12345:0"
IMPORT_PROBE: Final = "import aws_sdk_bedrock_runtime, smithy_aws_core; print('bedrock-realtime ok')"
pytestmark = [
pytest.mark.skipif(IMAGE is None, reason="requires a built image (set LITELLM_IMAGE)"),
pytest.mark.skipif(shutil.which("docker") is None, reason="requires the docker CLI"),
]
def test_image_imports_bedrock_realtime_sdk():
assert IMAGE is not None
probe: Final = subprocess.run(
[
"docker",
"run",
"--rm",
"--network",
"none",
"--user",
NON_ROOT_UID,
"--entrypoint",
"python",
IMAGE,
"-c",
IMPORT_PROBE,
],
capture_output=True,
text=True,
check=False,
)
assert probe.returncode == 0 and "bedrock-realtime ok" in probe.stdout, (
f"{IMAGE} cannot import aws_sdk_bedrock_runtime as uid {NON_ROOT_UID}, so Bedrock Nova Sonic "
"/v1/realtime sessions fail with 'Missing aws_sdk_bedrock_runtime'. Is `--extra bedrock-realtime` "
f"passed to every `uv sync` in its Dockerfile?\nstdout:\n{probe.stdout}\nstderr:\n{probe.stderr}"
)

View file

@ -0,0 +1,469 @@
"""
Regression tests for the Datadog LLM Observability payload schema (issue #35786).
Datadog renders tool calls, tool results and prompt-cache savings only from the fields its
own schema names. These assert on the payload `create_llm_obs_payload` actually hands the
intake, so a regression that moves data back into `meta.metadata` fails here.
Fixtures mirror what a live proxy run recorded on the callback, including the provider
spelling of prompt-cache counts (`prompt_tokens_details.cached_tokens`).
"""
import json
import os
from datetime import datetime, timedelta
from typing import Any
from unittest.mock import patch
import pytest
from litellm.integrations.datadog.datadog_llm_obs import DataDogLLMObsLogger
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
TOOL_DEFINITION: dict[str, Any] = {
"type": "function",
"function": {
"name": "get_weather",
"description": "Get current weather for a city",
"parameters": {
"type": "object",
"properties": {"city": {"type": "string"}},
"required": ["city"],
},
},
}
ASSISTANT_TOOL_CALL: dict[str, Any] = {
"id": "call_abc123",
"type": "function",
"function": {"name": "get_weather", "arguments": '{"city":"Paris","unit":"c"}'},
}
@pytest.fixture
def logger() -> DataDogLLMObsLogger:
with patch.dict(os.environ, {"DD_API_KEY": "k", "DD_SITE": "us5.datadoghq.com"}, clear=True):
with patch("asyncio.create_task"):
return DataDogLLMObsLogger()
NOT_GIVEN: Any = object()
def build_payload(
messages: Any = NOT_GIVEN,
response_message: dict[str, Any] | None = None,
usage_object: dict[str, Any] | None = None,
model_parameters: dict[str, Any] | None = None,
prompt_tokens: int = 4447,
) -> dict[str, Any]:
return {
"standard_logging_object": {
"call_type": "acompletion",
"messages": [{"role": "user", "content": "hi"}] if messages is NOT_GIVEN else messages,
"response": {"choices": [{"message": response_message or {"role": "assistant", "content": "hello"}}]},
"model_parameters": model_parameters or {},
"metadata": {"usage_object": usage_object} if usage_object is not None else {},
"prompt_tokens": prompt_tokens,
"completion_tokens": 507,
"total_tokens": prompt_tokens + 507,
"response_cost": 0.02,
"status": "success",
},
"litellm_params": {"metadata": {}},
}
def build(logger: DataDogLLMObsLogger, **kwargs: Any) -> dict[str, Any]:
"""Build a span and read it back as the JSON the intake receives, not as Python objects."""
start = datetime(2026, 9, 1, 12, 0, 0)
payload = logger.create_llm_obs_payload(build_payload(**kwargs), start, start + timedelta(seconds=2))
return json.loads(safe_dumps(payload))
def test_output_tool_calls_use_the_datadog_tool_call_schema(logger: DataDogLLMObsLogger) -> None:
"""Datadog reads name/arguments/tool_id off the tool call; OpenAI nests them under `function`."""
payload = build(
logger,
response_message={"role": "assistant", "content": None, "tool_calls": [ASSISTANT_TOOL_CALL]},
)
message = payload["meta"]["output"]["messages"][0]
assert message["tool_calls"] == [
{
"name": "get_weather",
"arguments": {"city": "Paris", "unit": "c"},
"tool_id": "call_abc123",
"type": "function",
}
]
assert "function" not in message["tool_calls"][0]
def test_tool_calls_are_not_duplicated_into_metadata(logger: DataDogLLMObsLogger) -> None:
"""The flat `output_tool_calls.*` keys were a second copy of a fact that now has its own field."""
payload = build(
logger,
response_message={"role": "assistant", "content": None, "tool_calls": [ASSISTANT_TOOL_CALL]},
)
assert [key for key in payload["meta"]["metadata"] if "tool_calls." in key] == []
def test_tool_result_message_links_back_to_its_tool_call(logger: DataDogLLMObsLogger) -> None:
"""Datadog pairs a result with its call through tool_id, and names the tool from the call."""
payload = build(
logger,
messages=[
{"role": "user", "content": "Weather in Paris?"},
{"role": "assistant", "content": None, "tool_calls": [ASSISTANT_TOOL_CALL]},
{"role": "tool", "tool_call_id": "call_abc123", "content": '{"temp_c": 18}'},
],
)
tool_message = payload["meta"]["input"]["messages"][2]
assert tool_message["tool_results"] == [
{"name": "get_weather", "result": '{"temp_c": 18}', "tool_id": "call_abc123", "type": "function"}
]
def test_tool_result_without_a_matching_call_still_reports_its_id(logger: DataDogLLMObsLogger) -> None:
"""A truncated conversation loses the call, so the name is unknown but the link must survive."""
payload = build(
logger,
messages=[{"role": "tool", "tool_call_id": "call_orphan", "content": "42"}],
)
assert payload["meta"]["input"]["messages"][0]["tool_results"] == [
{"name": "", "result": "42", "tool_id": "call_orphan", "type": "function"}
]
def test_cache_tokens_are_reported_as_span_metrics(logger: DataDogLLMObsLogger) -> None:
"""
Datadog charts cache savings from span metrics; nested usage_object is not read for it.
litellm's normalized prompt count includes both cache categories, so the three cache
metrics must partition input_tokens: read + write + non_cached == input.
"""
payload = build(
logger,
usage_object={"prompt_tokens_details": {"cached_tokens": 4300, "cache_write_tokens": 95}},
)
metrics = payload["metrics"]
assert metrics["cache_read_input_tokens"] == 4300.0
assert metrics["cache_write_input_tokens"] == 95.0
assert metrics["non_cached_input_tokens"] == 4447.0 - 4300.0 - 95.0
assert (
metrics["cache_read_input_tokens"] + metrics["cache_write_input_tokens"] + metrics["non_cached_input_tokens"]
== metrics["input_tokens"]
)
def test_cache_write_tokens_are_not_counted_as_non_cached(logger: DataDogLLMObsLogger) -> None:
"""A cache-priming request must not report its primed prefix as full-price uncached input."""
payload = build(logger, usage_object={"prompt_tokens_details": {"cache_write_tokens": 4000}})
assert payload["metrics"]["cache_write_input_tokens"] == 4000.0
assert payload["metrics"]["non_cached_input_tokens"] == 4447.0 - 4000.0
assert "cache_read_input_tokens" not in payload["metrics"]
def test_a_fully_cached_request_reports_a_zero_non_cached_count(logger: DataDogLLMObsLogger) -> None:
"""Zero residual is real data: everything was served from cache. Inconsistent counts clamp to it."""
payload = build(
logger,
usage_object={"prompt_tokens_details": {"cached_tokens": 4352, "cache_write_tokens": 95}},
)
assert payload["metrics"]["non_cached_input_tokens"] == 0.0
def test_anthropic_top_level_cache_keys_are_read(logger: DataDogLLMObsLogger) -> None:
"""A raw Anthropic usage dict records the counts top level, not under prompt_tokens_details."""
payload = build(
logger,
usage_object={"cache_read_input_tokens": 4300, "cache_creation_input_tokens": 95},
)
metrics = payload["metrics"]
assert metrics["cache_read_input_tokens"] == 4300.0
assert metrics["cache_write_input_tokens"] == 95.0
assert metrics["non_cached_input_tokens"] == 4447.0 - 4300.0 - 95.0
def test_cache_metrics_come_from_the_normalized_field_not_the_anthropic_one(logger: DataDogLLMObsLogger) -> None:
"""
litellm normalizes every provider's cache counters into prompt_tokens_details.
A real cached request from a non-Anthropic provider carries only `cached_tokens`, so
reading the Anthropic-specific `cache_read_input_tokens` key reports nothing for it.
"""
payload = build(
logger,
usage_object={"prompt_tokens_details": {"audio_tokens": None, "cached_tokens": 4096}},
prompt_tokens=4335,
)
assert payload["metrics"]["cache_read_input_tokens"] == 4096.0
assert payload["metrics"]["non_cached_input_tokens"] == 4335.0 - 4096.0
@pytest.mark.parametrize(
"usage_object",
[
{"prompt_tokens_details": {"cache_write_tokens": 95}},
{"prompt_tokens_details": {"cache_creation_tokens": 95}},
{"cache_creation_input_tokens": 95},
],
)
def test_every_spelling_of_cache_write_tokens_is_read(
logger: DataDogLLMObsLogger, usage_object: dict[str, Any]
) -> None:
"""A raw usage dict that bypassed litellm's normalizer can carry any provider's spelling."""
payload = build(logger, usage_object=usage_object)
assert payload["metrics"]["cache_write_input_tokens"] == 95.0
def test_a_cache_read_does_not_emit_a_zero_cache_write(logger: DataDogLLMObsLogger) -> None:
"""A zero write on every cache-read span would drag Datadog's cache-write average to nothing."""
payload = build(logger, usage_object={"prompt_tokens_details": {"cached_tokens": 4096}})
assert payload["metrics"]["cache_read_input_tokens"] == 4096.0
assert "cache_write_input_tokens" not in payload["metrics"]
def test_no_cache_keys_when_the_provider_reports_no_caching(logger: DataDogLLMObsLogger) -> None:
"""An uncached request must not gain zero-valued cache metrics that dilute cache dashboards."""
payload = build(logger, usage_object={"prompt_tokens_details": None})
assert "cache_read_input_tokens" not in payload["metrics"]
assert "cache_write_input_tokens" not in payload["metrics"]
assert "non_cached_input_tokens" not in payload["metrics"]
def test_tool_definitions_are_sent_on_meta(logger: DataDogLLMObsLogger) -> None:
payload = build(logger, model_parameters={"tools": [TOOL_DEFINITION]})
assert payload["meta"]["tool_definitions"] == [
{
"name": "get_weather",
"description": "Get current weather for a city",
"schema": {"type": "object", "properties": {"city": {"type": "string"}}, "required": ["city"]},
}
]
def test_tool_definitions_accept_the_bare_anthropic_shape(logger: DataDogLLMObsLogger) -> None:
"""The Anthropic surface declares tools unwrapped, with input_schema instead of parameters."""
payload = build(
logger,
model_parameters={"tools": [{"name": "get_weather", "description": "d", "input_schema": {"type": "object"}}]},
)
assert payload["meta"]["tool_definitions"] == [
{"name": "get_weather", "description": "d", "schema": {"type": "object"}}
]
def test_meta_omits_tool_definitions_when_no_tools_were_offered(logger: DataDogLLMObsLogger) -> None:
assert "tool_definitions" not in build(logger)["meta"]
def test_unparseable_tool_arguments_are_preserved_rather_than_dropped(logger: DataDogLLMObsLogger) -> None:
"""A truncated argument string is still the only record of what the model tried to call."""
payload = build(
logger,
response_message={
"role": "assistant",
"content": None,
"tool_calls": [{"id": "c1", "type": "function", "function": {"name": "f", "arguments": '{"city":'}}],
},
)
assert payload["meta"]["output"]["messages"][0]["tool_calls"][0]["arguments"] == '{"city":'
def test_oversized_tool_arguments_ship_unparsed(logger: DataDogLLMObsLogger) -> None:
"""
Decoding attacker-sized compact JSON multiplies memory for a span that is only logging.
This payload is perfectly valid JSON, so the only reason it arrives as a string is the
size bound; a smaller copy of the same shape comes back as an object below.
"""
oversized = '{"a":"' + "x" * 300_000 + '"}'
payload = build(
logger,
response_message={
"role": "assistant",
"content": None,
"tool_calls": [{"id": "c1", "type": "function", "function": {"name": "f", "arguments": oversized}}],
},
)
assert payload["meta"]["output"]["messages"][0]["tool_calls"][0]["arguments"] == oversized
def test_valid_arguments_below_the_bound_still_parse(logger: DataDogLLMObsLogger) -> None:
"""The size bound must not swallow ordinary arguments; this is the oversized test's control."""
payload = build(
logger,
response_message={
"role": "assistant",
"content": None,
"tool_calls": [
{"id": "c1", "type": "function", "function": {"name": "f", "arguments": '{"a":"' + "x" * 64 + '"}'}}
],
},
)
assert payload["meta"]["output"]["messages"][0]["tool_calls"][0]["arguments"] == {"a": "x" * 64}
def test_a_result_is_named_even_when_its_call_had_unparseable_arguments(logger: DataDogLLMObsLogger) -> None:
"""Correlating a result to its call reads ids and names, so bad arguments cannot break linking."""
payload = build(
logger,
messages=[
{
"role": "assistant",
"content": None,
"tool_calls": [
{"id": "call_abc123", "type": "function", "function": {"name": "get_weather", "arguments": "{"}}
],
},
{"role": "tool", "tool_call_id": "call_abc123", "content": "18C"},
],
)
assert payload["meta"]["input"]["messages"][1]["tool_results"] == [
{"name": "get_weather", "result": "18C", "tool_id": "call_abc123", "type": "function"}
]
def test_deeply_nested_tool_arguments_do_not_drop_the_span(logger: DataDogLLMObsLogger) -> None:
"""json.loads raises RecursionError, not JSONDecodeError, on hostile nesting."""
payload = build(
logger,
response_message={
"role": "assistant",
"content": None,
"tool_calls": [{"id": "c1", "type": "function", "function": {"name": "f", "arguments": "[" * 50_000}}],
},
)
assert payload["meta"]["output"]["messages"][0]["tool_calls"][0]["arguments"] == "[" * 50_000
def test_tool_arguments_that_parse_to_a_non_object_stay_a_string(logger: DataDogLLMObsLogger) -> None:
"""Datadog types arguments as an object, so a bare JSON scalar must not land there as one."""
payload = build(
logger,
response_message={
"role": "assistant",
"content": None,
"tool_calls": [{"id": "c1", "type": "function", "function": {"name": "f", "arguments": "42"}}],
},
)
assert payload["meta"]["output"]["messages"][0]["tool_calls"][0]["arguments"] == "42"
def test_a_tool_without_a_name_is_not_offered_as_a_definition(logger: DataDogLLMObsLogger) -> None:
"""A nameless tool cannot be matched to a call, so it is dropped rather than sent blank."""
payload = build(logger, model_parameters={"tools": [{"function": {"description": "no name"}}, TOOL_DEFINITION]})
assert [tool["name"] for tool in payload["meta"]["tool_definitions"]] == ["get_weather"]
def test_a_tool_definition_without_a_schema_omits_the_field(logger: DataDogLLMObsLogger) -> None:
"""An empty schema object would read as a tool that takes no arguments, which is a different claim."""
payload = build(logger, model_parameters={"tools": [{"name": "ping", "description": "d"}]})
assert payload["meta"]["tool_definitions"] == [{"name": "ping", "description": "d"}]
def test_a_non_dict_message_still_reaches_datadog(logger: DataDogLLMObsLogger) -> None:
"""Callers can log arbitrary message payloads, and dropping the span over one loses the request."""
payload = build(logger, messages=["just a bare string"])
assert payload["meta"]["input"]["messages"] == [{"input": "just a bare string"}]
def test_messages_logged_as_a_bare_string_still_reach_datadog(logger: DataDogLLMObsLogger) -> None:
payload = build(logger, messages="the whole prompt as one string")
assert payload["meta"]["input"]["messages"] == [{"input": "the whole prompt as one string"}]
def test_non_chat_call_types_log_an_empty_input(logger: DataDogLLMObsLogger) -> None:
"""Embedding and image calls carry no messages; fabricating an "None" turn misreads in Datadog."""
payload = build(logger, messages=None)
assert payload["meta"]["input"]["messages"] == []
def test_anthropic_tool_blocks_map_to_tool_calls_and_results(logger: DataDogLLMObsLogger) -> None:
"""/v1/messages carries tool traffic as content blocks, not OpenAI fields."""
payload = build(
logger,
messages=[
{"role": "user", "content": [{"type": "text", "text": "Weather in Tokyo?"}]},
{
"role": "assistant",
"content": [{"type": "tool_use", "id": "toolu_1", "name": "get_weather", "input": {"city": "Tokyo"}}],
},
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_1", "content": "18C"}]},
],
)
assistant, result_turn = payload["meta"]["input"]["messages"][1:3]
assert assistant["tool_calls"] == [
{"name": "get_weather", "arguments": {"city": "Tokyo"}, "tool_id": "toolu_1", "type": "tool_use"}
]
assert result_turn["tool_results"] == [
{"name": "get_weather", "result": "18C", "tool_id": "toolu_1", "type": "function"}
]
def test_content_with_no_text_parts_is_preserved_not_blanked(logger: DataDogLLMObsLogger) -> None:
"""A content list the mapper does not understand must ride along, not be erased."""
blocks = [{"type": "image_url", "image_url": {"url": "https://example.com/x.png"}}]
payload = build(logger, messages=[{"role": "user", "content": blocks}])
assert payload["meta"]["input"]["messages"][0]["content"] == blocks
def test_multimodal_content_parts_are_flattened_to_text(logger: DataDogLLMObsLogger) -> None:
"""Datadog types Message.content as a string, so content lists collapse to their text."""
payload = build(
logger,
messages=[
{"role": "user", "content": [{"type": "text", "text": "describe "}, {"type": "text", "text": "this"}]}
],
)
assert payload["meta"]["input"]["messages"][0]["content"] == "describe this"
def test_mapping_input_messages_does_not_mutate_the_shared_payload(logger: DataDogLLMObsLogger) -> None:
"""Sibling callbacks read the same messages list, so flattening must not write through it."""
messages: list[dict[str, Any]] = [{"role": "user", "content": [{"type": "text", "text": "hi"}]}]
kwargs = build_payload(messages=messages)
start = datetime(2026, 9, 1, 12, 0, 0)
logger.create_llm_obs_payload(kwargs, start, start + timedelta(seconds=1))
assert messages[0]["content"] == [{"type": "text", "text": "hi"}]
def test_reasoning_content_survives_the_mapping(logger: DataDogLLMObsLogger) -> None:
payload = build(
logger,
response_message={"role": "assistant", "content": "answer", "reasoning_content": "thinking"},
)
assert payload["meta"]["output"]["messages"][0]["reasoning_content"] == "thinking"

View file

@ -2932,6 +2932,28 @@ def test_add_cache_point_tool_block_passes_ttl_for_claude_4_5(monkeypatch):
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", old_env)
def test_add_cache_point_tool_block_stands_down_for_model_without_prompt_caching(monkeypatch):
"""A tool carrying cache_control must not become a cachePoint for a Bedrock model
whose cost-map entry lacks prompt caching support, since Bedrock rejects the whole
request. An unmapped id keeps emitting so ARN deployments do not lose caching."""
from litellm.litellm_core_utils.prompt_templates.factory import (
add_cache_point_tool_block,
)
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
tool = {"cache_control": {"type": "ephemeral"}}
assert add_cache_point_tool_block(tool, model="nvidia.nemotron-super-3-120b") is None
assert add_cache_point_tool_block(tool, model="us.nvidia.nemotron-super-3-120b") is None
assert add_cache_point_tool_block(
tool, model="arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/abc123"
) == {"cachePoint": {"type": "default"}}
assert add_cache_point_tool_block(tool, model="us.anthropic.claude-sonnet-4-5-20250929-v1:0") == {
"cachePoint": {"type": "default"}
}
def test_bedrock_tools_pt_passes_ttl_for_claude_4_5(monkeypatch):
"""
End-to-end: _bedrock_tools_pt should produce cachePoint blocks with ttl

View file

@ -256,14 +256,12 @@ def test_get_model_cost_map_stamps_loaded_at(monkeypatch):
from litellm.litellm_core_utils import get_model_cost_map as module
monkeypatch.setattr(module._cost_map_source_info, "loaded_at", None)
monkeypatch.setattr(
module.GetModelCostMap,
"fetch_remote_model_cost_map",
staticmethod(lambda url, timeout=5: _load_root_cost_map()),
client, _calls = _mock_client(
[httpx.Response(200, content=_real_map_bytes())], client_cls=httpx.Client
)
before = datetime.now(timezone.utc)
module.get_model_cost_map(url="https://example.invalid/cost_map.json")
module.get_model_cost_map(url="https://example.invalid/cost_map.json", client=client)
loaded_at = module.get_model_cost_map_loaded_at()
assert loaded_at is not None
@ -308,7 +306,7 @@ def _unset_local_cost_map_env(monkeypatch):
monkeypatch.delenv("LITELLM_LOCAL_MODEL_COST_MAP", raising=False)
def _mock_client(outcomes):
def _mock_client(outcomes, client_cls=httpx.AsyncClient):
"""httpx client over a MockTransport serving one outcome per request; an exception instance is raised."""
calls = {"count": 0}
@ -320,7 +318,7 @@ def _mock_client(outcomes):
raise outcome
return outcome
return httpx.AsyncClient(transport=httpx.MockTransport(handler)), calls
return client_cls(transport=httpx.MockTransport(handler)), calls
@pytest.mark.asyncio
@ -450,3 +448,97 @@ async def test_refetch_respects_local_env_override(monkeypatch):
)
assert isinstance(result, ModelCostMapReloaded)
assert len(result.model_cost_map) > 100
# ---------------------------------------------------------------------------
# get_model_cost_map: the boot-time load retries transient failures like a reload does
# ---------------------------------------------------------------------------
from litellm.litellm_core_utils.get_model_cost_map import (
get_model_cost_map,
get_model_cost_map_source_info,
)
class _SyncSleepRecorder:
"""Injected in place of time.sleep so the boot path's waits are asserted without delay."""
def __init__(self):
self.waits = []
def __call__(self, seconds: float) -> None:
self.waits.append(seconds)
def test_boot_load_retries_transient_failures_instead_of_falling_back():
"""A refused connection then a 503 at pod boot used to pin the process to the bundled
backup for its lifetime; both are transient and must be retried before giving up."""
client, calls = _mock_client(
[
httpx.ConnectError("connection refused"),
httpx.Response(503),
httpx.Response(200, content=_real_map_bytes()),
],
client_cls=httpx.Client,
)
sleeper = _SyncSleepRecorder()
cost_map = get_model_cost_map(url=_URL, sleep=sleeper, rng=random.Random(0), client=client)
assert calls["count"] == 3
assert len(sleeper.waits) == 2
assert 2.0 <= sleeper.waits[0] < 3.0
assert 4.0 <= sleeper.waits[1] < 5.0
source = get_model_cost_map_source_info()
assert source["source"] == "remote"
assert source["fallback_reason"] is None
assert cost_map.keys() >= _load_root_cost_map().keys() - {"sample_spec", FALLBACK_GENERALIZATIONS_KEY}
def test_boot_load_honors_retry_after_then_falls_back_after_max_attempts():
"""An outage longer than the retry budget still ends on the bundled backup, and the
recorded fallback reason says how many attempts were spent so operators can tell."""
client, calls = _mock_client(
[httpx.Response(429, headers={"Retry-After": "7"})], client_cls=httpx.Client
)
sleeper = _SyncSleepRecorder()
cost_map = get_model_cost_map(url=_URL, sleep=sleeper, rng=random.Random(0), client=client)
assert calls["count"] == 3
assert sleeper.waits == [7.0, 7.0]
source = get_model_cost_map_source_info()
assert source["source"] == "local"
assert "after 3 attempts" in source["fallback_reason"]
assert len(cost_map) > 100
def test_boot_load_does_not_retry_permanent_failures():
"""A 404 or a malformed URL cannot heal by waiting: one attempt, no sleeps, backup."""
client, calls = _mock_client([httpx.Response(404)], client_cls=httpx.Client)
sleeper = _SyncSleepRecorder()
get_model_cost_map(url=_URL, sleep=sleeper, rng=random.Random(0), client=client)
assert calls["count"] == 1
assert sleeper.waits == []
assert get_model_cost_map_source_info()["source"] == "local"
get_model_cost_map(url="not a url", sleep=sleeper, rng=random.Random(0))
assert sleeper.waits == []
assert get_model_cost_map_source_info()["source"] == "local"
def test_boot_load_respects_local_env_override(monkeypatch):
"""LITELLM_LOCAL_MODEL_COST_MAP=True still short-circuits to the backup with zero HTTP."""
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
def _fail(request):
raise AssertionError("no HTTP request should be made when local map is forced")
cost_map = get_model_cost_map(
url=_URL,
sleep=_SyncSleepRecorder(),
client=httpx.Client(transport=httpx.MockTransport(_fail)),
)
assert len(cost_map) > 100
assert get_model_cost_map_source_info()["is_env_forced"] is True

View file

@ -5248,6 +5248,84 @@ def test_cache_control_injection_tool_config_drops_ttl_for_unsupported_model():
assert tools[-1] == {"cachePoint": {"type": "default"}}
@pytest.mark.parametrize(
("model", "expects_cache_points"),
[
pytest.param("nvidia.nemotron-super-3-120b", False, id="mapped-model-without-prompt-caching"),
pytest.param("us.nvidia.nemotron-super-3-120b", False, id="regional-prefix-resolves-through-base-model"),
pytest.param(
"us.anthropic.claude-3-5-sonnet-20240620-v1:0", False, id="claude-named-but-not-caching-on-bedrock"
),
pytest.param("us.anthropic.claude-sonnet-4-5-20250929-v1:0", True, id="mapped-model-with-prompt-caching"),
pytest.param(
"arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/abc123",
True,
id="unmapped-arn-keeps-emitting",
),
],
)
def test_cache_points_emitted_only_for_models_that_support_prompt_caching(model, expects_cache_points, monkeypatch):
"""Bedrock rejects cachePoint blocks for models without prompt caching support
("You invoked an unsupported model or your request did not allow prompt caching"),
and clients like Claude Code attach cache_control to every request, so a map-known
model without the capability must not receive them. Unmapped ids (application
inference profile ARNs, models newer than the map) keep emitting so existing
caching setups never silently degrade."""
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
body = AmazonConverseConfig().transform_request(
model=model,
messages=[
{"role": "system", "content": [{"type": "text", "text": "sys", "cache_control": {"type": "ephemeral"}}]},
{"role": "user", "content": [{"type": "text", "text": "hi", "cache_control": {"type": "ephemeral"}}]},
],
optional_params={},
litellm_params={},
headers={},
)
assert ("cachePoint" in json.dumps(body)) is expects_cache_points
assert body["system"][0]["text"] == "sys"
assert body["messages"][0]["content"][0]["text"] == "hi"
def test_tool_config_cachepoint_not_placed_or_credited_for_model_without_prompt_caching(monkeypatch):
"""The tool_config injection point must stand down with the rest of the cachePoint
emission when the model cannot cache, and spend attribution must not credit the
gateway for a breakpoint that was never placed."""
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
bucket: dict = {"user_api_key": "sk-test"}
data = AmazonConverseConfig()._transform_request_helper(
model="nvidia.nemotron-super-3-120b",
system_content_blocks=[],
optional_params={
"tools": [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get weather",
"parameters": {
"type": "object",
"properties": {"location": {"type": "string"}},
"required": ["location"],
},
},
}
],
"cache_control_injection_points": [{"location": "tool_config"}],
},
messages=[{"role": "user", "content": "hi"}],
litellm_params={"metadata": bucket, "litellm_metadata": None, "model_info": {"id": "dep-bedrock"}},
)
assert "cachePoint" not in json.dumps(data.get("toolConfig", {}))
assert "litellm_gateway_injected_cache" not in bucket
def test_translate_response_format_json_schema_still_injects_tool():
"""
response_format with an explicit json_schema should still use the
@ -6211,7 +6289,7 @@ def test_message_level_cache_control_drops_ttl_for_unsupported_model(ttl_target)
result = _bedrock_converse_messages_pt(
messages=_agentic_messages_with_ttl(ttl_target),
model="anthropic.claude-3-5-sonnet-20240620-v1:0",
model="anthropic.claude-3-5-sonnet-20241022-v2:0",
llm_provider="bedrock_converse",
)

View file

@ -1,7 +1,11 @@
import asyncio
import concurrent.futures
import socket
import sys
from typing import Final
import aiohttp
import aiohttp.abc
import aiohttp.client_exceptions
import aiohttp.http_exceptions
import httpx
@ -1140,3 +1144,55 @@ async def test_stopped_loop_session_disposed_synchronously_on_recycle():
finally:
await new_session.close()
result["loop"].close()
class _CancellingResolver(aiohttp.abc.AbstractResolver):
"""Cancels the given task (or, by default, aiohttp's shielded DNS child task) mid-lookup."""
def __init__(self, task_to_cancel: "asyncio.Task[object] | None" = None):
self._task_to_cancel: Final = task_to_cancel
async def resolve(
self, host: str, port: int = 0, family: socket.AddressFamily = socket.AF_INET
) -> list[aiohttp.abc.ResolveResult]:
target: Final = self._task_to_cancel or asyncio.current_task()
assert target is not None
target.cancel()
await asyncio.sleep(0)
raise OSError("resolver finished after the task was cancelled")
async def close(self) -> None:
return None
@pytest.mark.asyncio
@pytest.mark.skipif(
sys.version_info < (3, 11), reason="Task.cancelling() is needed to tell the two cancellations apart"
)
async def test_internal_dns_cancellation_maps_to_connect_error():
"""A CancelledError the request task never asked for must surface as a mapped httpx transport error."""
session = aiohttp.ClientSession(connector=aiohttp.TCPConnector(resolver=_CancellingResolver()))
transport = LiteLLMAiohttpTransport(client=session)
try:
with pytest.raises(httpx.ConnectError):
await transport.handle_async_request(httpx.Request("GET", "http://example.invalid/"))
current = asyncio.current_task()
assert current is not None and current.cancelling() == 0
finally:
await transport.aclose()
@pytest.mark.asyncio
async def test_genuine_request_cancellation_still_propagates():
"""Cancelling the request task itself (client disconnect, shutdown) must still propagate unmapped."""
current = asyncio.current_task()
assert current is not None
session = aiohttp.ClientSession(connector=aiohttp.TCPConnector(resolver=_CancellingResolver(current)))
transport = LiteLLMAiohttpTransport(client=session)
try:
with pytest.raises(asyncio.CancelledError):
await transport.handle_async_request(httpx.Request("GET", "http://example.invalid/"))
finally:
if sys.version_info >= (3, 11):
current.uncancel()
await transport.aclose()

View file

@ -1431,7 +1431,11 @@ class TestListToolsRestAPI:
async def test_aggregate_list_absorbs_one_server_auth_failure(self, monkeypatch):
"""The multi-server aggregate listing degrades a server whose upstream
rejects auth to an empty contribution and still returns the healthy
server's tools with a 200, rather than surfacing a 401."""
server's tools with a 200, rather than surfacing a 401. The absorbed
server must still show up as a classified per-server outcome so a REST
caller can tell "needs upstream auth" apart from "has no tools"."""
from pydantic import TypeAdapter
from litellm.proxy._experimental.mcp_server.exceptions import (
MCPUpstreamAuthError,
)
@ -1497,6 +1501,11 @@ class TestListToolsRestAPI:
assert result["tools"] == ["good-tool"]
assert result["error"] is None
wire_body = json.loads(TypeAdapter(dict).dump_json(result))
assert wire_body["server_outcomes"] == {
"good": {"status": "ok", "tool_count": 1},
"bad": {"status": "auth_required", "http_status": 401},
}
async def test_name_resolution_finds_server_by_uuid(self, monkeypatch):
"""When server_id is a name string, it should be resolved to its UUID

View file

@ -1,6 +1,6 @@
import logging
import time
from collections.abc import Mapping
from collections.abc import Callable, Mapping
from itertools import chain
from typing import Final
from unittest.mock import AsyncMock, MagicMock, call
@ -1645,6 +1645,25 @@ async def test_update_group_e2e(mocker):
ScimTransformations.transform_litellm_team_to_scim_group.assert_called_once_with(updated_team)
def _rows_by_exact_id(
user_row: Callable[[Mapping[str, str]], LiteLLM_UserTable | MagicMock | None],
) -> Callable[..., tuple[LiteLLM_UserTable | MagicMock, ...]]:
"""``find_many`` stand-in for the classifier's cross-field read on a table where a
member value only ever matches as an exact ``user_id``."""
def rows(where: Mapping[str, object], take: int | None = None) -> tuple[LiteLLM_UserTable | MagicMock, ...]:
clauses: Final = where["OR"]
assert isinstance(clauses, list)
found: Final = tuple(user_row(clause) for clause in clauses if "user_id" in clause)
return tuple(row for row in found if row is not None)
return rows
def _user_row_for(where: Mapping[str, str]) -> LiteLLM_UserTable:
return LiteLLM_UserTable(user_id=where["user_id"])
@pytest.mark.asyncio
async def test_create_group_with_nonexistent_users_rejects(mocker, monkeypatch):
"""
@ -1696,9 +1715,8 @@ async def test_create_group_with_nonexistent_users_rejects(mocker, monkeypatch):
return mock_user
return None # new-user-1 and new-user-2 don't exist
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=mock_user_lookup)
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None)
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=[])
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(side_effect=_rows_by_exact_id(mock_user_lookup))
# Mock dependencies
mocker.patch(
@ -1782,9 +1800,8 @@ async def test_update_group_with_nonexistent_users_rejects(mocker, monkeypatch):
return mock_user
return None # new-user-3 and new-user-4 don't exist
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=mock_user_lookup)
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None)
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=[])
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(side_effect=_rows_by_exact_id(mock_user_lookup))
# Mock dependencies
mocker.patch(
@ -1853,9 +1870,8 @@ async def test_create_group_with_nonexistent_users_creates_when_flag_true(mocker
return mock_user
return None # new-user-1 and new-user-2 don't exist
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=mock_user_lookup)
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None)
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=[])
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(side_effect=_rows_by_exact_id(mock_user_lookup))
# Mock user creation
created_user_1 = NewUserResponse(user_id="new-user-1", key="test-key-1")
@ -1943,9 +1959,8 @@ async def test_extract_group_member_ids_with_flag_true_creates_users(mocker, mon
return mock_user
return None # new-user-1 doesn't exist
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=mock_user_lookup)
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None)
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=[])
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(side_effect=_rows_by_exact_id(mock_user_lookup))
# Mock user creation
created_user = NewUserResponse(user_id="new-user-1", key="test-key-1")
@ -2013,9 +2028,8 @@ async def test_extract_group_member_ids_with_flag_false_rejects(mocker, monkeypa
return mock_user
return None # new-user-1 doesn't exist
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=mock_user_lookup)
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None)
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=[])
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(side_effect=_rows_by_exact_id(mock_user_lookup))
# Mock dependencies
mocker.patch(
@ -3121,8 +3135,7 @@ async def test_process_group_patch_operations_add_retains_existing_members(mocke
mock_prisma_client.db = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable = mocker.MagicMock()
# new-user already exists in the DB
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=mocker.MagicMock(user_id="new-user"))
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=())
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=(mocker.MagicMock(user_id="new-user"),))
_, final_members, _ = await _process_group_patch_operations(
patch_ops=patch_ops,
@ -3415,8 +3428,7 @@ async def test_patch_group_add_applies_delta_and_keeps_concurrent_add(mocker):
)
mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=final_team)
mock_prisma_client.db.litellm_usertable = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=mocker.MagicMock())
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=())
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(side_effect=_rows_by_exact_id(_user_row_for))
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
@ -3509,8 +3521,7 @@ async def test_patch_group_replace_stays_absolute_against_concurrent_roster(mock
)
mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=final_team)
mock_prisma_client.db.litellm_usertable = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=mocker.MagicMock())
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=())
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(side_effect=_rows_by_exact_id(_user_row_for))
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
@ -3640,8 +3651,7 @@ async def test_process_group_patch_add_filtered_path_without_value(mocker):
prisma_client = mocker.MagicMock()
prisma_client.db = mocker.MagicMock()
prisma_client.db.litellm_usertable = mocker.MagicMock()
prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=LiteLLM_UserTable(user_id="user-3"))
prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=())
prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=(LiteLLM_UserTable(user_id="user-3"),))
_, final_members, _ = await _process_group_patch_operations(
patch_ops=patch_ops,
@ -3733,12 +3743,14 @@ def _member_resolution_prisma(
starts folding it, fails here instead of passing.
A caller that must know which accounts match rather than merely how many
passes take=None, so an unbounded read returns every match.
passes take=None, so an unbounded read returns every match. The row keyed by
the value comes last, the order a bounded read is least prepared for, since
the database promises no order at all.
"""
clauses: Final = where["OR"]
assert isinstance(clauses, list)
fields: Final = tuple(next(iter(clause)) for clause in clauses)
assert fields == ("sso_user_id", "user_email"), fields
assert fields in (("user_id", "sso_user_id", "user_email"), ("sso_user_id", "user_email")), fields
def comparison(clause: Mapping[str, object]) -> tuple[str, bool]:
"""The needle and whether production asked for a case-insensitive compare,
@ -3749,8 +3761,9 @@ def _member_resolution_prisma(
assert isinstance(criterion, dict), criterion
return criterion["equals"], criterion.get("mode") == "insensitive"
sso_needle, sso_insensitive = comparison(clauses[0])
email_needle, email_insensitive = comparison(clauses[1])
by_field: Final = dict(zip(fields, (comparison(clause) for clause in clauses)))
sso_needle, sso_insensitive = by_field["sso_user_id"]
email_needle, email_insensitive = by_field["user_email"]
def same(stored: str, needle: str, insensitive: bool) -> bool:
return stored.casefold() == needle.casefold() if insensitive else stored == needle
@ -3768,6 +3781,11 @@ def _member_resolution_prisma(
if same(email, email_needle, email_insensitive)
for user_id in user_ids
),
(
user_id
for user_id in users
if "user_id" in by_field and same(user_id, by_field["user_id"][0], by_field["user_id"][1])
),
)
)
found: Final = tuple(dict.fromkeys(matched))
@ -4611,9 +4629,15 @@ async def test_resolve_group_member_ids_dedupes_repeated_member(mocker, scim_ups
def _identity_lookup(value: str) -> object:
"""The single cross-field lookup the classifier is expected to issue."""
"""The single cross-field lookup the classifier is expected to issue per member."""
return call(
where={"OR": [{"sso_user_id": value}, {"user_email": {"equals": value, "mode": "insensitive"}}]},
where={
"OR": [
{"user_id": value},
{"sso_user_id": value},
{"user_email": {"equals": value, "mode": "insensitive"}},
]
},
take=2,
)
@ -5152,6 +5176,77 @@ async def test_resolve_group_member_ids_refuses_a_user_id_that_names_another_acc
)
@pytest.mark.asyncio
async def test_resolve_group_member_ids_reads_the_exact_id_when_two_other_accounts_fill_the_lookup(
mocker, scim_upsert_user_enabled
):
"""A value that is one account's id and two other accounts' identities fills the
bounded lookup with the other two. The account keyed by the value must still be
found, or the id would lose its precedence and a non-canonical type would skip
a member that names a real user."""
prisma_client = _member_resolution_prisma(
mocker,
users={"shared"},
teams=set(),
sso_user_id_to_user_id={"shared": "by-sso"},
email_to_user_id={"shared": "by-email"},
)
create_user_mock = mocker.patch( # test-quality-ok: user creation is module-level, not injectable into the resolver
"litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists",
AsyncMock(return_value=None),
)
with pytest.raises(HTTPException) as exc_info:
await _resolve_group_member_ids(
members=[SCIMMember(value="shared", type="direct")],
created_via="scim_group_membership",
prisma_client=prisma_client,
)
assert exc_info.value.status_code == 400
assert "shared" in str(exc_info.value.detail)
create_user_mock.assert_not_called()
assert prisma_client.db.litellm_usertable.find_many.await_args_list == [_identity_lookup("shared")]
prisma_client.db.litellm_usertable.find_unique.assert_awaited_once_with(where={"user_id": "shared"})
@pytest.mark.asyncio
async def test_resolve_group_member_ids_reads_the_user_table_once_per_member(mocker, scim_upsert_user_enabled):
"""Every member costs one read of the user table, however it resolves: by its exact
id (which still outranks a non-canonical type), by identity, as a SCIM team, or not
at all. Looking the exact id up on its own before the identity read doubled the
reads of a push, and the identity read is a scan."""
prisma_client = _member_resolution_prisma(
mocker,
users={"by-id"},
teams={"by-team"},
email_to_user_id={"by-email@example.com": "email-user"},
)
mocker.patch( # test-quality-ok: user creation is module-level, not injectable into the resolver
"litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists",
AsyncMock(return_value=NewUserResponse(user_id="nobody", key="key")),
)
result = await _resolve_group_member_ids(
members=[
SCIMMember(value="by-id", type="direct"),
SCIMMember(value="by-email@example.com"),
SCIMMember(value="by-team"),
SCIMMember(value="nobody"),
],
created_via="scim_group_membership",
prisma_client=prisma_client,
)
assert result.all_member_ids == ["by-id", "email-user", "nobody"]
prisma_client.db.litellm_usertable.find_unique.assert_not_awaited()
assert prisma_client.db.litellm_usertable.find_many.await_args_list == [
_identity_lookup("by-id"),
_identity_lookup("by-email@example.com"),
_identity_lookup("by-team"),
_identity_lookup("nobody"),
]
@pytest.mark.asyncio
async def test_resolve_group_member_ids_warns_before_creating_unmatched_placeholder(
@ -5536,10 +5631,7 @@ async def test_resolve_group_member_ids_admits_member_created_concurrently(mocke
the member is still admitted: the id resolves to a real user row, so failing
or dropping it would be wrong either way."""
prisma_client = _member_resolution_prisma(mocker, users=set(), teams=set())
prisma_client.db.litellm_usertable.find_unique = AsyncMock(
side_effect=[None, LiteLLM_UserTable(user_id="raced-user")]
)
prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=())
prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=LiteLLM_UserTable(user_id="raced-user"))
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists",
AsyncMock(return_value=None),

View file

@ -0,0 +1,56 @@
"""
Static checks that every proxy Docker image installs the `bedrock-realtime` extra.
Bedrock Nova Sonic speech-to-speech (`/v1/realtime`) needs `aws-sdk-bedrock-runtime`,
which only ships in the `bedrock-realtime` extra. An image whose `uv sync` stages
omit the extra fails every Nova Sonic realtime session with
"Missing aws_sdk_bedrock_runtime. Install with: pip install aws-sdk-bedrock-runtime".
"""
import os
import re
from typing import Final
import pytest
REPO_ROOT: Final = os.path.join(os.path.dirname(__file__), "..", "..")
PROXY_DOCKERFILES: Final = (
"Dockerfile",
os.path.join("docker", "Dockerfile.non_root"),
os.path.join("docker", "Dockerfile.database"),
os.path.join("gateway", "Dockerfile"),
)
CONTINUED_LINE_RE: Final = re.compile(r"(?:\\\n|[^\n])+")
UV_SYNC_BOUNDARY_RE: Final = re.compile(r"(?=uv sync)")
def _uv_sync_invocations(dockerfile_text: str) -> tuple[str, ...]:
"""Return each `uv sync ...` command, split apart when one RUN holds several (if/else branches)."""
return tuple(
part
for line in CONTINUED_LINE_RE.finditer(dockerfile_text)
for part in UV_SYNC_BOUNDARY_RE.split(line.group(0))
if part.startswith("uv sync")
)
@pytest.mark.parametrize("relative_path", PROXY_DOCKERFILES)
def test_every_uv_sync_installs_bedrock_realtime_extra(relative_path: str):
dockerfile_path: Final = os.path.join(REPO_ROOT, relative_path)
if not os.path.exists(dockerfile_path):
pytest.skip(f"{relative_path} not present in this checkout")
with open(dockerfile_path, "r", encoding="utf-8") as f:
contents: Final = f.read()
invocations: Final = _uv_sync_invocations(contents)
assert invocations, f"{relative_path} has no `uv sync` invocation"
missing: Final = tuple(invocation for invocation in invocations if "--extra bedrock-realtime" not in invocation)
assert not missing, (
f"{relative_path}: {len(missing)} of {len(invocations)} `uv sync` invocations omit "
"`--extra bedrock-realtime`, so aws-sdk-bedrock-runtime is absent and Bedrock Nova Sonic "
"/v1/realtime sessions fail with 'Missing aws_sdk_bedrock_runtime'"
)

View file

@ -2281,3 +2281,83 @@ def test_every_declaring_deployment_is_named(caplog):
assert "azure-ptu-east" in warnings[0]
assert "azure-ptu-west" in warnings[0]
assert "plain-gpt-4o" not in warnings[0]
def _simulate_price_data_reload_with_provider_sets(monkeypatch, fetched_catalog):
"""Like `_simulate_price_data_reload`, plus the provider model-set refresh the proxy's
`_swap_in_model_cost_map` does before replaying, so bare names in the new catalog resolve."""
monkeypatch.setattr(litellm, "model_cost", fetched_catalog)
_invalidate_model_cost_lowercase_map()
litellm.add_known_models(model_cost_map=fetched_catalog)
reapply_runtime_model_cost_registrations()
def test_a_config_deployment_dropped_by_a_stale_cost_map_comes_back_on_reload(monkeypatch):
"""
Booting on the bundled backup, a bare model that only the remote catalog knows
cannot be provider-resolved, so the proxy router (ignore_invalid_deployments) drops
it. Once a reload brings in a catalog that knows the model, the deployment must be
served again with its access groups, and exactly once however many reloads follow.
"""
backend = "lit-5766-only-in-remote-catalog"
try:
router = Router(
model_list=[
{
"model_name": "new-model",
"litellm_params": {"model": backend, "api_key": "k"},
"model_info": {"id": "new-id", "access_groups": ["team-models"]},
},
{
"model_name": "control-model",
"litellm_params": {"model": "hosted_vllm/control-backend", "api_key": "k"},
"model_info": {"id": "control-id", "access_groups": ["team-models"]},
},
],
ignore_invalid_deployments=True,
)
assert router.get_model_names() == ["control-model"]
assert router.get_model_access_groups(model_name="new-model") == {}
fresh_catalog = {**litellm.model_cost, backend: {"litellm_provider": "openai", "mode": "chat"}}
_simulate_price_data_reload_with_provider_sets(monkeypatch, fresh_catalog)
_simulate_price_data_reload_with_provider_sets(monkeypatch, fresh_catalog)
assert sorted(router.get_model_names()) == ["control-model", "new-model"]
assert router.get_model_access_groups(model_name="new-model") == {"team-models": ["new-model"]}
assert [d["model_info"]["id"] for d in router.model_list] == ["control-id", "new-id"]
assert "new-id" in litellm.model_cost
finally:
litellm.open_ai_chat_completion_models.discard(backend)
litellm.models_by_provider["openai"].discard(backend)
def test_a_config_deployment_dropped_for_a_permanent_reason_is_not_retried_on_reload(monkeypatch):
"""
Only provider-resolution drops can be healed by a fresh catalog. A deployment that
fails after its provider resolved (here a pass-through vertex entry with no project)
has already touched router state, so replaying it on every reload would leak into
`deployment_names` each time.
"""
router = Router(
model_list=[
{
"model_name": "vertex-passthrough",
"litellm_params": {"model": "vertex_ai/gemini-2.5-flash", "use_in_pass_through": True},
"model_info": {"id": "vertex-id"},
},
{
"model_name": "control-model",
"litellm_params": {"model": "hosted_vllm/control-backend", "api_key": "k"},
"model_info": {"id": "control-id"},
},
],
ignore_invalid_deployments=True,
)
assert router.get_model_names() == ["control-model"]
names_after_boot = list(router.deployment_names)
_simulate_price_data_reload_with_provider_sets(monkeypatch, dict(litellm.model_cost))
assert router.get_model_names() == ["control-model"]
assert router.deployment_names == names_after_boot