mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_fix_gateway_injection_scope
This commit is contained in:
commit
6d8c18d518
27 changed files with 1639 additions and 270 deletions
6
.github/workflows/image-scan.yml
vendored
6
.github/workflows/image-scan.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 \
|
||||
|
|
|
|||
|
|
@ -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 \
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
)
|
||||
469
tests/test_litellm/integrations/datadog/test_datadog_llm_obs.py
Normal file
469
tests/test_litellm/integrations/datadog/test_datadog_llm_obs.py
Normal 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"
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
56
tests/test_litellm/test_dockerfile_bedrock_realtime_extra.py
Normal file
56
tests/test_litellm/test_dockerfile_bedrock_realtime_extra.py
Normal 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'"
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue