Merge remote-tracking branch 'origin/litellm_internal_staging' into fix/mcp-failure-call-type-logging

This commit is contained in:
CrypticCortex 2026-06-24 00:54:53 +05:30
commit 3cbd219fe4
170 changed files with 17406 additions and 3780 deletions

View file

@ -14,7 +14,7 @@ permissions:
jobs:
lint:
runs-on: ubuntu-latest
timeout-minutes: 10
timeout-minutes: 15
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
@ -87,9 +87,11 @@ jobs:
run: |
uv run --no-sync python -c "import openai; print(f'OpenAI version: {openai.__version__}')"
- name: Run basedpyright type checking
- name: Check basedpyright budget (delta vs base)
env:
BASE_SHA: ${{ github.event.pull_request.base.sha }}
run: |
(uv run --no-sync basedpyright --outputjson || true) | uv run --no-sync python scripts/type_check_gate.py
(uv run --no-sync basedpyright --outputjson || true) | uv run --no-sync python scripts/type_check_gate.py --base "$BASE_SHA"
- name: Check for circular imports
run: |

View file

@ -33,6 +33,7 @@ jobs:
tests/test_litellm/images
tests/test_litellm/interactions
tests/test_litellm/passthrough
tests/test_litellm/sandbox
tests/test_litellm/vector_stores
tests/test_litellm/test_*.py
workers: 2

View file

@ -125,7 +125,8 @@ lint-ruff-FULL-dev: install-dev
else echo "No changed .py files to check."; fi
lint-basedpyright: install-dev
($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py
git fetch origin litellm_internal_staging
($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --base origin/litellm_internal_staging
lint-basedpyright-budget-update: install-dev
($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --update

View file

@ -121,7 +121,7 @@
},
"reportReturnType": {
"baseline": 126,
"slack": 13
"slack": 100
},
"reportTypedDictNotRequiredAccess": {
"baseline": 20,
@ -157,7 +157,7 @@
},
"reportUnnecessaryComparison": {
"baseline": 683,
"slack": 10
"slack": 100
},
"reportUnnecessaryContains": {
"baseline": 4,

View file

@ -673,6 +673,7 @@ elevenlabs_models: Set = set()
dashscope_models: Set = set()
moonshot_models: Set = set()
publicai_models: Set = set()
darkbloom_models: Set = set()
v0_models: Set = set()
morph_models: Set = set()
lambda_ai_models: Set = set()
@ -927,6 +928,8 @@ def add_known_models(model_cost_map: Optional[Dict] = None):
moonshot_models.add(key)
elif value.get("litellm_provider") == "publicai":
publicai_models.add(key)
elif value.get("litellm_provider") == "darkbloom":
darkbloom_models.add(key)
elif value.get("litellm_provider") == "v0":
v0_models.add(key)
elif value.get("litellm_provider") == "morph":
@ -1075,6 +1078,7 @@ model_list = list(
| dashscope_models
| moonshot_models
| publicai_models
| darkbloom_models
| v0_models
| morph_models
| lambda_ai_models
@ -1179,6 +1183,7 @@ models_by_provider: dict = {
"modelscope": modelscope_models,
"moonshot": moonshot_models,
"publicai": publicai_models,
"darkbloom": darkbloom_models,
"v0": v0_models,
"morph": morph_models,
"lambda_ai": lambda_ai_models,
@ -1922,9 +1927,6 @@ if TYPE_CHECKING:
from .llms.fireworks_ai.completion.transformation import (
FireworksAITextCompletionConfig as FireworksAITextCompletionConfig,
)
from .llms.fireworks_ai.audio_transcription.transformation import (
FireworksAIAudioTranscriptionConfig as FireworksAIAudioTranscriptionConfig,
)
from .llms.fireworks_ai.embed.fireworks_ai_transformation import (
FireworksAIEmbeddingConfig as FireworksAIEmbeddingConfig,
)

View file

@ -260,7 +260,6 @@ LLM_CONFIG_NAMES = (
"SambaNovaEmbeddingConfig",
"FireworksAIConfig",
"FireworksAITextCompletionConfig",
"FireworksAIAudioTranscriptionConfig",
"FireworksAIEmbeddingConfig",
"FriendliaiChatConfig",
"JinaAIEmbeddingConfig",
@ -1027,10 +1026,6 @@ _LLM_CONFIGS_IMPORT_MAP = {
".llms.fireworks_ai.completion.transformation",
"FireworksAITextCompletionConfig",
),
"FireworksAIAudioTranscriptionConfig": (
".llms.fireworks_ai.audio_transcription.transformation",
"FireworksAIAudioTranscriptionConfig",
),
"FireworksAIEmbeddingConfig": (
".llms.fireworks_ai.embed.fireworks_ai_transformation",
"FireworksAIEmbeddingConfig",

View file

@ -201,6 +201,18 @@ DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET = int(
# Provider-specific API base URLs
XAI_API_BASE = "https://api.x.ai/v1"
OPEN_SANDBOX_API_BASE_ENV_VAR = "OPEN_SANDBOX_API_BASE"
OPEN_SANDBOX_API_KEY_ENV_VAR = "OPEN_SANDBOX_API_KEY"
OPEN_SANDBOX_DEFAULT_TEMPLATE = "opensandbox/code-interpreter:v1.1.0"
_OPEN_SANDBOX_FALLBACK_ENTRYPOINT = "/opt/code-interpreter/code-interpreter.sh"
OPEN_SANDBOX_DEFAULT_ENTRYPOINT = (_OPEN_SANDBOX_FALLBACK_ENTRYPOINT,)
OPEN_SANDBOX_DEFAULT_LANGUAGE = "python"
OPEN_SANDBOX_DEFAULT_CPU_LIMIT = "1"
OPEN_SANDBOX_DEFAULT_MEMORY_LIMIT = "2Gi"
OPEN_SANDBOX_EXECD_PORT = 44772
OPEN_SANDBOX_DEFAULT_TIMEOUT = 300
OPEN_SANDBOX_READY_TIMEOUT = 30.0
OPEN_SANDBOX_POLL_INTERVAL = 0.2
DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET = int(
os.getenv("DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET", 1024)
@ -867,6 +879,7 @@ openai_compatible_providers: List = [
"docker_model_runner",
"ragflow",
"pinstripes", # Pinstripes - JSON-configured provider
"darkbloom",
]
openai_text_completion_compatible_providers: List = (
[ # providers that support `/v1/completions`

View file

@ -9,9 +9,11 @@ captured stdout back through the typed agentic loop plan.
import json
import time
import uuid
from typing import Any, cast
from typing import Any, Literal, TypedDict, cast
import litellm
from pydantic import ValidationError
from litellm._logging import verbose_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.integrations.code_interpreter_interception import (
@ -20,15 +22,93 @@ from litellm.types.integrations.code_interpreter_interception import (
from litellm.types.integrations.custom_logger import (
AgenticLoopPlan,
AgenticLoopRequestPatch,
CHAT_COMPLETION_AGENTIC_SURFACE,
NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES,
is_interception_internal_key,
)
from litellm.types.llms.openai import (
ChatCompletionAssistantMessage,
ChatCompletionAssistantToolCall,
ChatCompletionToolMessage,
)
from litellm.types.utils import (
CallTypes,
ChatCompletionMessageToolCall,
ModelResponse,
)
from litellm.types.utils import CallTypes
LITELLM_CODE_EXECUTION_TOOL_NAME = "litellm_code_execution"
_INTERCEPTION_ACTIVE_KEY = "_code_interpreter_interception_active"
_SANDBOX_KEY = "_code_interpreter_interception_sandbox_key"
_CONVERTED_STREAM_KEY = "_code_interpreter_interception_converted_stream"
_LITELLM_METADATA_KEY = "litellm_metadata"
_CACHE_TTL_SECONDS = 15 * 60
class CodeExecutionToolCall(TypedDict, total=False):
id: str | None
call_id: str | None
type: Literal["function"]
name: str
arguments: str
class CodeInterpreterLogOutput(TypedDict):
type: Literal["logs"]
logs: str
class CodeInterpreterCall(TypedDict):
id: str
type: Literal["code_interpreter_call"]
status: Literal["completed"]
code: str
container_id: str | None
outputs: list[CodeInterpreterLogOutput]
class CodeExecutionFunctionParameters(TypedDict):
type: Literal["object"]
properties: dict[str, dict[str, str]]
required: list[str]
class ResponsesFunctionTool(TypedDict):
type: Literal["function"]
name: str
description: str
parameters: CodeExecutionFunctionParameters
class ChatCompletionFunctionDefinition(TypedDict):
name: str
description: str
parameters: CodeExecutionFunctionParameters
class ChatCompletionFunctionTool(TypedDict):
type: Literal["function"]
function: ChatCompletionFunctionDefinition
CodeExecutionFunctionTool = ResponsesFunctionTool | ChatCompletionFunctionTool
class ResponsesFunctionToolChoice(TypedDict):
type: Literal["function"]
name: str
class ChatCompletionFunctionToolChoice(TypedDict):
type: Literal["function"]
function: dict[str, str]
CodeExecutionFunctionToolChoice = (
ResponsesFunctionToolChoice | ChatCompletionFunctionToolChoice
)
def _resolve_sandbox_tool(sandbox_tool_name: str | None) -> dict[str, Any] | None:
try:
from litellm.sandbox.sandbox_tools import resolve_sandbox_tool
@ -97,9 +177,15 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
if not kwargs.get("_agentic_loop_depth"):
kwargs.pop(_INTERCEPTION_ACTIVE_KEY, None)
kwargs.pop(_SANDBOX_KEY, None)
self._strip_interception_metadata(kwargs)
if not self.enabled:
return None
if call_type not in (CallTypes.responses, CallTypes.aresponses):
if call_type not in (
CallTypes.responses,
CallTypes.aresponses,
CallTypes.completion,
CallTypes.acompletion,
):
return None
if (
self.enabled_providers is not None
@ -120,18 +206,10 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
kwargs[_SANDBOX_KEY] = uuid.uuid4().hex
if kwargs.get("stream"):
kwargs["stream"] = False
kwargs["_code_interpreter_interception_converted_stream"] = True
kwargs[_CONVERTED_STREAM_KEY] = True
self._write_interception_metadata(kwargs)
function_tool = {
"type": "function",
"name": LITELLM_CODE_EXECUTION_TOOL_NAME,
"description": "Execute python code in a sandbox and return stdout.",
"parameters": {
"type": "object",
"properties": {"code": {"type": "string"}},
"required": ["code"],
},
}
function_tool = self._get_function_tool(call_type=call_type)
kwargs["tools"] = [
(
function_tool
@ -141,19 +219,90 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
for tool in tools
]
if self._tool_choice_targets_code_interpreter(kwargs.get("tool_choice")):
kwargs["tool_choice"] = {
"type": "function",
"name": LITELLM_CODE_EXECUTION_TOOL_NAME,
}
kwargs["tool_choice"] = self._get_function_tool_choice(call_type=call_type)
return kwargs
@staticmethod
def _strip_interception_metadata(kwargs: dict[str, Any]) -> None:
metadata = kwargs.get(_LITELLM_METADATA_KEY)
if not isinstance(metadata, dict):
return
filtered_metadata = {
key: value
for key, value in metadata.items()
if not is_interception_internal_key(key)
and not key.startswith("_agentic_loop")
and key != "max_agentic_loops"
}
if filtered_metadata:
kwargs[_LITELLM_METADATA_KEY] = filtered_metadata
else:
kwargs.pop(_LITELLM_METADATA_KEY, None)
@staticmethod
def _write_interception_metadata(kwargs: dict[str, Any]) -> None:
metadata = kwargs.get(_LITELLM_METADATA_KEY)
metadata = dict(metadata) if isinstance(metadata, dict) else {}
for key in (_INTERCEPTION_ACTIVE_KEY, _SANDBOX_KEY, _CONVERTED_STREAM_KEY):
if key in kwargs:
metadata[key] = kwargs[key]
kwargs[_LITELLM_METADATA_KEY] = metadata
@staticmethod
def _get_function_parameters() -> CodeExecutionFunctionParameters:
return {
"type": "object",
"properties": {"code": {"type": "string"}},
"required": ["code"],
}
def _get_function_tool(
self, call_type: CallTypes | None
) -> CodeExecutionFunctionTool:
description = "Execute python code in a sandbox and return stdout."
if call_type in (CallTypes.completion, CallTypes.acompletion):
return {
"type": "function",
"function": {
"name": LITELLM_CODE_EXECUTION_TOOL_NAME,
"description": description,
"parameters": self._get_function_parameters(),
},
}
return {
"type": "function",
"name": LITELLM_CODE_EXECUTION_TOOL_NAME,
"description": description,
"parameters": self._get_function_parameters(),
}
@staticmethod
def _get_function_tool_choice(
call_type: CallTypes | None,
) -> CodeExecutionFunctionToolChoice:
if call_type in (CallTypes.completion, CallTypes.acompletion):
return {
"type": "function",
"function": {"name": LITELLM_CODE_EXECUTION_TOOL_NAME},
}
return {
"type": "function",
"name": LITELLM_CODE_EXECUTION_TOOL_NAME,
}
@staticmethod
def _tool_choice_targets_code_interpreter(tool_choice: Any) -> bool:
if not isinstance(tool_choice, dict):
return False
function = tool_choice.get("function")
return (
tool_choice.get("type") == "code_interpreter"
or tool_choice.get("name") == "code_interpreter"
or tool_choice.get("name") == LITELLM_CODE_EXECUTION_TOOL_NAME
or (
isinstance(function, dict)
and function.get("name") == LITELLM_CODE_EXECUTION_TOOL_NAME
)
)
def _resolve_provider(self, kwargs: dict[str, Any]) -> str | None:
@ -188,7 +337,12 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
):
return False, {}
tool_calls = self._extract_code_execution_tool_calls(response=response)
tool_calls = (
self._extract_chat_completion_code_execution_tool_calls(response=response)
if kwargs.get("_agentic_loop_api_surface")
== CHAT_COMPLETION_AGENTIC_SURFACE
else self._extract_code_execution_tool_calls(response=response)
)
if not tool_calls:
return False, {}
@ -206,15 +360,24 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
stream: bool,
kwargs: dict,
) -> AgenticLoopPlan:
if kwargs.get("_agentic_loop_api_surface") == CHAT_COMPLETION_AGENTIC_SURFACE:
return await self._build_chat_completion_agentic_loop_plan(
tools=tools,
model=model,
messages=messages,
optional_params=anthropic_messages_optional_request_params,
kwargs=kwargs,
)
await self._prune_expired_cache()
tool_calls = cast(list[dict[str, Any]], tools.get("tool_calls", []))
tool_calls = cast(list[CodeExecutionToolCall], tools.get("tool_calls", []))
sandbox_key = kwargs.get(_SANDBOX_KEY)
container, params = await self._get_or_create_container(cache_key=sandbox_key)
try:
container_id = getattr(container, "id", None)
container_id = cast(str | None, getattr(container, "id", None))
input_list = self._normalize_messages(messages)
code_interpreter_calls = []
code_interpreter_calls: list[CodeInterpreterCall] = []
for tool_call in tool_calls:
arguments = tool_call.get("arguments", "")
code = self._parse_code(arguments)
@ -256,9 +419,12 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
request_patch = AgenticLoopRequestPatch(
model=model,
messages=input_list,
tools=optional_params.get("tools"),
optional_params={k: v for k, v in optional_params.items() if k != "tools"},
kwargs={k: v for k, v in kwargs.items() if k != "litellm_logging_obj"},
tools=self._get_followup_tools(
tools=optional_params.get("tools"),
call_type=CallTypes.responses,
),
optional_params=self._get_followup_optional_params(optional_params),
kwargs=self._filter_agentic_loop_kwargs(kwargs),
)
return AgenticLoopPlan(
@ -271,12 +437,134 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
},
)
async def _build_chat_completion_agentic_loop_plan(
self,
tools: dict[str, object],
model: str,
messages: list[dict],
optional_params: dict[str, object],
kwargs: dict[str, object],
) -> AgenticLoopPlan:
await self._prune_expired_cache()
tool_calls = cast(list[CodeExecutionToolCall], tools.get("tool_calls", []))
sandbox_key = cast(str | None, kwargs.get(_SANDBOX_KEY))
container, params = await self._get_or_create_container(cache_key=sandbox_key)
try:
container_id = cast(str | None, getattr(container, "id", None))
tool_results = [
await self._build_chat_completion_tool_result(
container=container,
params=params,
tool_call=tool_call,
container_id=container_id,
)
for tool_call in tool_calls
]
except Exception:
await self._delete_container_for_cache_key(sandbox_key)
raise
tool_messages = [result[0] for result in tool_results]
code_interpreter_calls = [result[1] for result in tool_results]
request_patch = AgenticLoopRequestPatch(
model=model,
messages=list(messages)
+ [self._build_chat_completion_assistant_message(tool_calls)]
+ tool_messages,
tools=self._get_followup_tools(
tools=optional_params.get("tools"),
call_type=CallTypes.completion,
),
optional_params=self._get_followup_optional_params(optional_params),
kwargs=self._filter_agentic_loop_kwargs(kwargs),
)
return AgenticLoopPlan(
run_agentic_loop=True,
request_patch=request_patch,
metadata={
"tool_type": "code_interpreter",
"sandbox_key": sandbox_key or "",
"code_interpreter_calls": code_interpreter_calls,
"response_format": "openai",
},
)
async def _build_chat_completion_tool_result(
self,
container: object,
params: dict[str, Any] | None,
tool_call: CodeExecutionToolCall,
container_id: str | None,
) -> tuple[ChatCompletionToolMessage, CodeInterpreterCall]:
arguments = tool_call.get("arguments", "")
code = self._parse_code(arguments)
stdout = await self._run_tool_call(
container=container, params=params, arguments=arguments
)
tool_call_id = (
tool_call.get("id") or tool_call.get("call_id") or uuid.uuid4().hex
)
return (
{
"role": "tool",
"tool_call_id": tool_call_id,
"content": stdout,
},
{
"id": f"ci_{uuid.uuid4().hex}",
"type": "code_interpreter_call",
"status": "completed",
"code": code,
"container_id": container_id,
"outputs": [{"type": "logs", "logs": stdout}] if stdout else [],
},
)
async def async_agentic_loop_cleanup_hook(
self, plan: AgenticLoopPlan, kwargs: dict
) -> None:
metadata = plan.metadata or {} if plan else {}
await self._delete_container_for_cache_key(metadata.get("sandbox_key"))
@staticmethod
def _filter_agentic_loop_kwargs(kwargs: dict[str, object]) -> dict[str, object]:
return {
k: v
for k, v in kwargs.items()
if k not in {"litellm_logging_obj", "acompletion"}
and not is_interception_internal_key(
k, prefixes=NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES
)
}
def _get_followup_tools(
self, tools: object, call_type: CallTypes | None
) -> list[dict[str, Any]] | None:
if not isinstance(tools, list):
return None
return [
(
self._get_function_tool(call_type=call_type)
if isinstance(tool, dict) and tool.get("type") == "code_interpreter"
else tool
)
for tool in tools
]
def _get_followup_optional_params(
self, optional_params: dict[str, object]
) -> dict[str, object]:
drop_tool_choice = self._tool_choice_targets_code_interpreter(
optional_params.get("tool_choice")
)
return {
k: v
for k, v in optional_params.items()
if k != "tools" and not (k == "tool_choice" and drop_tool_choice)
}
async def async_post_agentic_loop_response_hook(
self, response: Any, plan: AgenticLoopPlan, kwargs: dict
) -> Any:
@ -420,7 +708,9 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
return list(messages)
return []
def _extract_code_execution_tool_calls(self, response: Any) -> list[dict[str, Any]]:
def _extract_code_execution_tool_calls(
self, response: object
) -> list[CodeExecutionToolCall]:
if isinstance(response, dict):
output = response.get("output", [])
else:
@ -446,6 +736,82 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
if self._is_code_execution_call(item)
]
def _extract_chat_completion_code_execution_tool_calls(
self, response: ModelResponse | dict[str, Any]
) -> list[CodeExecutionToolCall]:
model_response = self._to_model_response(response)
if model_response is None:
return []
choices = model_response.choices or []
if not choices:
return []
message = choices[0].message
tool_calls = message.tool_calls or []
return [
normalized
for tool_call in tool_calls
if (normalized := self._normalize_chat_completion_tool_call(tool_call))
is not None
]
@staticmethod
def _normalize_chat_completion_tool_call(
tool_call: ChatCompletionMessageToolCall,
) -> CodeExecutionToolCall | None:
if (
tool_call.type != "function"
or tool_call.function.name != LITELLM_CODE_EXECUTION_TOOL_NAME
):
return None
arguments = tool_call.function.arguments
if isinstance(arguments, dict):
arguments = json.dumps(arguments)
elif not isinstance(arguments, str):
arguments = "" if arguments is None else str(arguments)
return {
"id": tool_call.id,
"call_id": tool_call.id,
"type": "function",
"name": LITELLM_CODE_EXECUTION_TOOL_NAME,
"arguments": arguments,
}
@staticmethod
def _build_chat_completion_assistant_message(
tool_calls: list[CodeExecutionToolCall],
) -> ChatCompletionAssistantMessage:
return {
"role": "assistant",
"tool_calls": [
cast(
ChatCompletionAssistantToolCall,
{
"id": tool_call.get("id"),
"type": "function",
"function": {
"name": LITELLM_CODE_EXECUTION_TOOL_NAME,
"arguments": tool_call.get("arguments", ""),
},
},
)
for tool_call in tool_calls
],
}
@staticmethod
def _to_model_response(
response: ModelResponse | dict[str, Any],
) -> ModelResponse | None:
if isinstance(response, ModelResponse):
return response
try:
return ModelResponse(**response)
except (TypeError, ValidationError):
return None
def _is_code_execution_call(self, item: Any) -> bool:
if isinstance(item, dict):
return (

View file

@ -1,6 +1,7 @@
"""Typed configuration for the OpenTelemetry instrumentation."""
from enum import Enum
from functools import lru_cache
from typing import Any, List
from pydantic import AliasChoices, BaseModel, Field, field_validator, model_validator
@ -47,7 +48,12 @@ class _OTelV2Flag(BaseSettings):
enabled: bool = Field(default=False, validation_alias=AliasChoices(OTEL_V2_ENV))
@lru_cache(maxsize=1)
def is_otel_v2_enabled() -> bool:
# Resolved once at startup and cached: constructing the pydantic-settings
# model re-scans the environment and cost ~28us, which on the proxy hot path
# (auth, logging-callback setup) compounded into a measurable throughput
# regression. Tests that toggle the env must call ``is_otel_v2_enabled.cache_clear()``.
return _OTelV2Flag().enabled

View file

@ -300,9 +300,6 @@ class LiteLLMResponsesInteractionsConfig:
"total_output_tokens": getattr(usage, "output_tokens", 0),
}
# Add role
interactions_response_dict["role"] = "model"
# Add updated (same as created for now)
interactions_response_dict["updated"] = created

View file

@ -0,0 +1,332 @@
# this is a patch to allow for agentic loops covering llm_http_handler.py and openai sdk based calling flows for the .completion() api
import json
from typing import cast
from litellm._logging import verbose_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.integrations.custom_logger import (
CHAT_COMPLETION_AGENTIC_SURFACE,
NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES,
AgenticLoopPlan,
AgenticLoopRequestPatch,
is_interception_internal_key,
)
from litellm.types.utils import ModelResponse
from litellm.utils import CustomStreamWrapper
_FOLLOWUP_INTERNAL_PARAMS = frozenset(
(
"acompletion",
"litellm_logging_obj",
"custom_llm_provider",
"model_alias_map",
"stream_response",
"custom_prompt_dict",
"_agentic_loop_api_surface",
)
)
def _gate_overridden(callback: CustomLogger) -> bool:
base = CustomLogger.async_should_run_agentic_loop
func = type(callback).async_should_run_agentic_loop
return getattr(func, "__func__", func) is not getattr(base, "__func__", base)
def _build_plan_overridden(callback: CustomLogger) -> bool:
base = CustomLogger.async_build_agentic_loop_plan
func = type(callback).async_build_agentic_loop_plan
return getattr(func, "__func__", func) is not getattr(base, "__func__", base)
def _post_hook_overridden(callback: CustomLogger) -> bool:
base = CustomLogger.async_post_agentic_loop_response_hook
func = type(callback).async_post_agentic_loop_response_hook
return getattr(func, "__func__", func) is not getattr(base, "__func__", base)
def _coerce_int(value: object, default: int) -> int:
return int(value) if isinstance(value, (int, str)) else default
def _agentic_loop_settings(kwargs: dict[str, object]) -> tuple[int, int, list[str]]:
depth = _coerce_int(kwargs.get("_agentic_loop_depth"), 0)
max_loops = max(_coerce_int(kwargs.get("max_agentic_loops"), 3), 1)
raw_fingerprints = kwargs.get("_agentic_loop_fingerprints")
fingerprints = (
[str(fp) for fp in raw_fingerprints]
if isinstance(raw_fingerprints, list)
else []
)
return depth, max_loops, fingerprints
def _fingerprint_tools(tool_calls: object) -> str:
try:
return json.dumps(tool_calls, sort_keys=True, default=str)
except Exception:
return str(tool_calls)
def _check_agentic_loop_safety(
tool_calls: object,
fingerprints: list[str],
depth: int,
max_loops: int,
model: str,
) -> str:
fingerprint = _fingerprint_tools(tool_calls)
if fingerprint in fingerprints:
raise ValueError(
"Agentic loop detected repeated tool-call fingerprint; aborting rerun"
)
if depth >= max_loops:
raise ValueError(f"Exceeded max_agentic_loops={max_loops} for model={model}")
return fingerprint
def _wrap_response_as_fake_stream(response: object) -> object:
if getattr(response, "object", None) == "chat.completion.chunk":
return response
if not hasattr(response, "choices"):
return response
from litellm.llms.base_llm.base_model_iterator import (
convert_model_response_to_streaming,
)
return convert_model_response_to_streaming(cast(ModelResponse, response))
def _add_agentic_loop_metadata(kwargs_for_followup: dict[str, object]) -> None:
metadata = kwargs_for_followup.get("litellm_metadata")
metadata = dict(metadata) if isinstance(metadata, dict) else {}
for key, value in kwargs_for_followup.items():
if (
key.startswith("_agentic_loop")
or key == "max_agentic_loops"
or is_interception_internal_key(key)
):
metadata[key] = value
kwargs_for_followup["litellm_metadata"] = metadata
def _filter_followup_kwargs(source: dict[str, object]) -> dict[str, object]:
return {
k: v
for k, v in source.items()
if not is_interception_internal_key(
k, prefixes=NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES
)
and k not in _FOLLOWUP_INTERNAL_PARAMS
}
async def _execute_chat_completion_agentic_plan(
*,
plan: AgenticLoopPlan,
callback: CustomLogger,
model: str,
optional_params: dict[str, object],
kwargs: dict[str, object],
logging_obj: object,
custom_llm_provider: str,
depth: int,
max_loops: int,
fingerprints: list[str],
fingerprint: str,
) -> object:
import litellm
patch = plan.request_patch or AgenticLoopRequestPatch()
if patch.messages is None:
raise ValueError("Agentic loop plan missing patched messages")
full_model_name = patch.model or model
if "/" not in full_model_name:
full_model_name = f"{custom_llm_provider}/{full_model_name}"
optional_params_for_followup = {**optional_params, **patch.optional_params}
if patch.tools is not None:
optional_params_for_followup["tools"] = patch.tools
if "tool_choice" not in patch.optional_params:
optional_params_for_followup.pop("tool_choice", None)
kwargs_for_followup = _filter_followup_kwargs(kwargs)
kwargs_for_followup.update(
{
k: v
for k, v in _filter_followup_kwargs(patch.kwargs).items()
if k not in optional_params_for_followup
}
)
kwargs_for_followup["_agentic_loop_depth"] = depth + 1
kwargs_for_followup["max_agentic_loops"] = max_loops
kwargs_for_followup["_agentic_loop_fingerprints"] = fingerprints + [fingerprint]
_add_agentic_loop_metadata(kwargs_for_followup)
try:
response_followup = await litellm.acompletion(
model=full_model_name,
messages=patch.messages,
**optional_params_for_followup,
**kwargs_for_followup,
)
if _post_hook_overridden(callback):
try:
response_followup = (
await callback.async_post_agentic_loop_response_hook(
response=response_followup, plan=plan, kwargs=kwargs
)
)
except Exception as e:
_call_id = getattr(logging_obj, "litellm_call_id", "unknown")
verbose_logger.exception(
"LiteLLM.AgenticHookError: Exception in "
"async_post_agentic_loop_response_hook [call_id=%s model=%s]: %s",
_call_id,
model,
str(e),
)
if kwargs.get("_code_interpreter_interception_converted_stream") and not depth:
return _wrap_response_as_fake_stream(response_followup)
return response_followup
finally:
try:
await callback.async_agentic_loop_cleanup_hook(plan=plan, kwargs=kwargs)
except Exception as e:
_call_id = getattr(logging_obj, "litellm_call_id", "unknown")
verbose_logger.exception(
"LiteLLM.AgenticHookError: Exception in "
"async_agentic_loop_cleanup_hook [call_id=%s model=%s]: %s",
_call_id,
model,
str(e),
)
async def maybe_run_chat_completion_agentic_loop(
*,
response: ModelResponse,
model: str,
messages: list,
optional_params: dict,
kwargs: dict,
logging_obj: object,
custom_llm_provider: str,
stream: bool,
) -> ModelResponse | CustomStreamWrapper | None:
import litellm
callbacks = litellm.callbacks + (
getattr(logging_obj, "dynamic_success_callbacks", None) or []
)
depth, max_loops, fingerprints = _agentic_loop_settings(kwargs)
tools = optional_params.get("tools", [])
for callback in callbacks:
if not isinstance(callback, CustomLogger):
continue
if not _gate_overridden(callback):
continue
gate_kwargs = {
**kwargs,
"_agentic_loop_api_surface": CHAT_COMPLETION_AGENTIC_SURFACE,
"custom_llm_provider": custom_llm_provider,
}
try:
should_run, tool_calls = await callback.async_should_run_agentic_loop(
response=response,
model=model,
messages=messages,
tools=tools,
stream=stream,
custom_llm_provider=custom_llm_provider,
kwargs=gate_kwargs,
)
except Exception as e:
verbose_logger.exception(
"LiteLLM.AgenticHookError: Exception in chat completion agentic gate: %s",
str(e),
)
continue
if not should_run:
continue
fingerprint = _check_agentic_loop_safety(
tool_calls=tool_calls,
fingerprints=fingerprints,
depth=depth,
max_loops=max_loops,
model=model,
)
try:
plan_kwargs = {
**kwargs,
"_agentic_loop_api_surface": CHAT_COMPLETION_AGENTIC_SURFACE,
"custom_llm_provider": custom_llm_provider,
}
if not _build_plan_overridden(callback):
return await callback.async_run_agentic_loop(
tools=tool_calls,
model=model,
messages=messages,
response=response,
anthropic_messages_provider_config=None,
anthropic_messages_optional_request_params=optional_params,
logging_obj=logging_obj,
stream=stream,
kwargs=plan_kwargs,
)
plan = await callback.async_build_agentic_loop_plan(
tools=tool_calls,
model=model,
messages=messages,
response=response,
anthropic_messages_provider_config=None,
anthropic_messages_optional_request_params=optional_params,
logging_obj=logging_obj,
stream=stream,
kwargs=plan_kwargs,
)
if plan.response_override is not None:
return plan.response_override
if plan.terminate:
return response
if not plan.run_agentic_loop:
continue
return await _execute_chat_completion_agentic_plan(
plan=plan,
callback=callback,
model=model,
optional_params=optional_params,
kwargs=kwargs,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
depth=depth,
max_loops=max_loops,
fingerprints=fingerprints,
fingerprint=fingerprint,
)
except Exception as e:
verbose_logger.exception(
"LiteLLM.AgenticHookError: Exception in chat completion agentic hooks: %s",
str(e),
)
if (
kwargs.get("_code_interpreter_interception_converted_stream")
and not depth
and hasattr(response, "choices")
):
return cast(
"ModelResponse | CustomStreamWrapper",
_wrap_response_as_fake_stream(response),
)
return None

View file

@ -86,9 +86,7 @@ def get_supported_openai_params(
model=model
)
elif request_type == "transcription":
return litellm.FireworksAIAudioTranscriptionConfig().get_supported_openai_params(
model=model
)
return None
else:
return litellm.FireworksAIConfig().get_supported_openai_params(model=model)
elif custom_llm_provider == "nvidia_nim":
@ -191,7 +189,9 @@ def get_supported_openai_params(
)
elif custom_llm_provider == "sambanova":
if request_type == "embeddings":
litellm.SambaNovaEmbeddingConfig().get_supported_openai_params(model=model)
return litellm.SambaNovaEmbeddingConfig().get_supported_openai_params(
model=model
)
else:
return litellm.SambanovaConfig().get_supported_openai_params(model=model)
elif custom_llm_provider == "nebius":

View file

@ -12,6 +12,7 @@ class SensitiveDataMasker:
visible_prefix: int = 4,
visible_suffix: int = 4,
mask_char: str = "*",
mask_short_values: bool = True,
):
self.sensitive_patterns = sensitive_patterns or {
"password",
@ -38,12 +39,17 @@ class SensitiveDataMasker:
self.visible_prefix = visible_prefix
self.visible_suffix = visible_suffix
self.mask_char = mask_char
self.mask_short_values = mask_short_values
def _mask_value(self, value: str) -> str:
if not value or len(str(value)) < (self.visible_prefix + self.visible_suffix):
return value
value_str = str(value)
if not value_str:
return value
if len(value_str) <= (self.visible_prefix + self.visible_suffix):
return (
self.mask_char * len(value_str) if self.mask_short_values else value_str
)
masked_length = len(value_str) - (self.visible_prefix + self.visible_suffix)
# Handle the case where visible_suffix is 0 to avoid showing the entire string

View file

@ -2005,11 +2005,29 @@ class CustomStreamWrapper:
except StopIteration:
if self.sent_last_chunk is True:
complete_streaming_response = litellm.stream_chunk_builder(
chunks=self.chunks,
messages=self.messages,
logging_obj=self.logging_obj,
)
try:
complete_streaming_response = litellm.stream_chunk_builder(
chunks=self.chunks,
messages=self.messages,
logging_obj=self.logging_obj,
)
except Exception as e:
# stream_chunk_builder can re-raise (as APIError) on large agentic
# streams. The raise originates inside this except-StopIteration block,
# so the sibling `except Exception` below does not catch it; it would
# escape __next__ and drop the request from SpendLogs. Recover
# best-effort usage from the raw chunks so cost is still tracked
verbose_logger.warning(
"stream_chunk_builder raised at end-of-stream (%s); logging "
"best-effort usage from chunks.",
str(e),
)
try:
complete_streaming_response = self.model_response_creator(
chunk={"usage": calculate_total_usage(chunks=self.chunks)}
)
except Exception:
complete_streaming_response = None
response = self.model_response_creator()
if complete_streaming_response is not None:
@ -2234,11 +2252,27 @@ class CustomStreamWrapper:
except (StopAsyncIteration, StopIteration):
if self.sent_last_chunk is True:
# log the final chunk with accurate streaming values
complete_streaming_response = litellm.stream_chunk_builder(
chunks=self.chunks,
messages=self.messages,
logging_obj=self.logging_obj,
)
try:
complete_streaming_response = litellm.stream_chunk_builder(
chunks=self.chunks,
messages=self.messages,
logging_obj=self.logging_obj,
)
except Exception as e:
# see sync __next__: a raise from stream_chunk_builder inside this
# except handler escapes __anext__ and drops the request from SpendLogs.
# Recover best-effort usage from the raw chunks so cost is still tracked
verbose_logger.warning(
"stream_chunk_builder raised at end-of-stream (%s); logging "
"best-effort usage from chunks.",
str(e),
)
try:
complete_streaming_response = self.model_response_creator(
chunk={"usage": calculate_total_usage(chunks=self.chunks)}
)
except Exception:
complete_streaming_response = None
response = self.model_response_creator()
if complete_streaming_response is not None:

View file

@ -84,8 +84,14 @@ class AdvisorOrchestrationHandler(MessagesInterceptor):
)
# Optional routing overrides for the advisor sub-call (e.g. proxy routing).
# If not set in the tool definition, litellm resolves from env vars.
advisor_api_key: Optional[str] = advisor_tool.get("api_key")
advisor_api_base: Optional[str] = advisor_tool.get("api_base")
# The advisor tool is caller-controlled; only honor a client-supplied
# api_base/api_key when the proxy has enabled clientside credentials,
# otherwise let litellm resolve from server config.
advisor_api_key: Optional[str] = None
advisor_api_base: Optional[str] = None
if _allow_client_side_advisor_credentials():
advisor_api_key = advisor_tool.get("api_key")
advisor_api_base = advisor_tool.get("api_base")
# Build the synthetic tool definition the provider will receive.
synthetic_advisor_tool = _make_synthetic_advisor_tool()
@ -181,6 +187,20 @@ class AdvisorOrchestrationHandler(MessagesInterceptor):
# ---------------------------------------------------------------------------
def _allow_client_side_advisor_credentials() -> bool:
"""Whether a caller-supplied advisor api_base/api_key may be honored.
Gated on the proxy's ``allow_client_side_credentials`` opt-in. When the
interceptor runs outside the proxy (SDK use), there is no admin boundary
to protect, so client-supplied routing is allowed.
"""
try:
from litellm.proxy.proxy_server import general_settings
except (ImportError, ModuleNotFoundError):
return True
return general_settings.get("allow_client_side_credentials") is True
def _make_synthetic_advisor_tool() -> Dict:
"""Build a regular tool definition the executor provider can understand."""
return {

View file

@ -8,10 +8,14 @@ run code -> delete container; `code_interpreter_tool` combines all three.
from typing import Any, Union
import httpx
from pydantic import Field, PrivateAttr
from litellm.types.llms.base import LiteLLMPydanticObjectBase
SANDBOX_MAX_OUTPUT_BYTES = 10 * 1024 * 1024
class ContainerHandle(LiteLLMPydanticObjectBase):
"""A live sandbox container. Carries everything needed to reach it again."""
@ -53,7 +57,7 @@ class BaseSandboxConfig:
*,
template: str | None = None,
timeout: int | None = None,
allow_internet_access: bool = True,
allow_internet_access: bool | None = None,
api_key: str | None = None,
**kwargs,
) -> ContainerHandle:
@ -77,3 +81,16 @@ class BaseSandboxConfig:
**kwargs,
) -> bool:
raise NotImplementedError("adelete_sandbox must be implemented by provider")
async def _read_capped_lines(self, response: httpx.Response) -> list[str]:
lines: list[str] = []
total = 0
async for line in response.aiter_lines():
total += len(line.encode("utf-8"))
if total > SANDBOX_MAX_OUTPUT_BYTES:
raise ValueError(
f"Sandbox output exceeded {SANDBOX_MAX_OUTPUT_BYTES} bytes; aborting "
"to avoid unbounded memory use."
)
lines.append(line)
return lines

View file

@ -10,7 +10,6 @@ from typing import (
Callable,
ClassVar,
Dict,
List,
Literal,
Optional,
Tuple,
@ -210,32 +209,11 @@ class BaseAWSLLM:
"""
Return a boto3.Credentials object
"""
## CHECK IS 'os.environ/' passed in
params_to_check: List[Optional[str]] = [
aws_access_key_id,
aws_secret_access_key,
aws_session_token,
aws_region_name,
aws_session_name,
aws_profile_name,
aws_role_name,
aws_web_identity_token,
aws_sts_endpoint,
aws_external_id,
]
# Iterate over parameters and update if needed
for i, param in enumerate(params_to_check):
if param and param.startswith("os.environ/"):
_v = get_secret(param)
if _v is not None and isinstance(_v, str):
params_to_check[i] = _v
elif param is None: # check if uppercase value in env
key = self.aws_authentication_params[i]
if key.upper() in os.environ:
params_to_check[i] = os.getenv(key.upper())
# Assign updated values back to parameters
# Only config-sourced credentials are expanded against the environment.
# os.environ/<VAR> references in the model config are resolved at load time,
# so any reference still present at this point is caller-supplied input and is
# left as-is rather than expanded into a process environment variable. Each
# unset param falls back to its matching fixed AWS_* ambient env var.
(
aws_access_key_id,
aws_secret_access_key,
@ -247,7 +225,21 @@ class BaseAWSLLM:
aws_web_identity_token,
aws_sts_endpoint,
aws_external_id,
) = params_to_check
) = tuple(
value if value is not None else os.getenv(env_var)
for value, env_var in (
(aws_access_key_id, "AWS_ACCESS_KEY_ID"),
(aws_secret_access_key, "AWS_SECRET_ACCESS_KEY"),
(aws_session_token, "AWS_SESSION_TOKEN"),
(aws_region_name, "AWS_REGION_NAME"),
(aws_session_name, "AWS_SESSION_NAME"),
(aws_profile_name, "AWS_PROFILE_NAME"),
(aws_role_name, "AWS_ROLE_NAME"),
(aws_web_identity_token, "AWS_WEB_IDENTITY_TOKEN"),
(aws_sts_endpoint, "AWS_STS_ENDPOINT"),
(aws_external_id, "AWS_EXTERNAL_ID"),
)
)
verbose_logger.debug(
"in get credentials\n"
@ -845,6 +837,20 @@ class BaseAWSLLM:
f"IN Web Identity Token: {aws_web_identity_token} | Role Name: {aws_role_name} | Session Name: {aws_session_name}"
)
# get_secret() expands environment-variable references (an os.environ/<VAR>
# prefix, or a bare name matching an environment variable). Config-sourced
# references are expanded at load time, so such a reference reaching here is
# caller-supplied input; reject it rather than expanding a process-environment
# value for use as the token.
if (
aws_web_identity_token.startswith("os.environ/")
or aws_web_identity_token in os.environ
):
raise AwsAuthError(
message="Invalid web identity token reference.",
status_code=400,
)
oidc_token = get_secret(aws_web_identity_token)
if oidc_token is None:

View file

@ -70,6 +70,7 @@ from ..base_aws_llm import BaseAWSLLM
from ..common_utils import (
BedrockError,
ModelResponseIterator,
build_bedrock_stream_error,
get_bedrock_response_stream_shape,
get_bedrock_tool_name,
)
@ -1841,23 +1842,7 @@ class AWSEventStreamDecoder:
parsed_response = self.parser.parse(response_dict, response_stream_shape)
if response_dict["status_code"] != 200:
decoded_body = response_dict["body"].decode()
if isinstance(decoded_body, dict):
error_message = decoded_body.get("message")
elif isinstance(decoded_body, str):
error_message = decoded_body
else:
error_message = ""
exception_status = response_dict["headers"].get(":exception-type")
error_message = exception_status + " " + error_message
raise BedrockError(
status_code=response_dict["status_code"],
message=(
json.dumps(error_message)
if isinstance(error_message, dict)
else error_message
),
)
raise build_bedrock_stream_error(response_dict, response_stream_shape)
if "chunk" in parsed_response:
chunk = parsed_response.get("chunk")
if not chunk:

View file

@ -7,9 +7,21 @@ Common utilities used across bedrock chat/embedding/image generation
import functools
import json
import os
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union
from typing import (
TYPE_CHECKING,
Any,
Dict,
List,
Literal,
Mapping,
Optional,
TypedDict,
Union,
)
if TYPE_CHECKING:
from botocore.model import Shape
from litellm.types.llms.bedrock import BedrockCreateBatchRequest
import httpx
@ -1132,6 +1144,39 @@ def get_bedrock_response_stream_shape():
return _load_bedrock_response_stream_shape()
class BedrockEventStreamResponseDict(TypedDict):
status_code: int
headers: Mapping[str, str]
body: bytes
def build_bedrock_stream_error(
response_dict: BedrockEventStreamResponseDict,
response_stream_shape: Shape | None,
) -> BedrockError:
"""Build a BedrockError for a non-200 event-stream error event.
botocore hard-codes HTTP 400 on every mid-stream error event, so the modeled
ResponseStream member's httpStatusCode is the real status. Resolve it from the
shape and fall back to the raw status when the type is not modeled.
"""
exception_type = response_dict["headers"].get(":exception-type")
decoded_body = response_dict["body"].decode()
message = f"{exception_type} {decoded_body}" if exception_type else decoded_body
status_code = response_dict["status_code"]
if exception_type is not None and response_stream_shape is not None:
member = response_stream_shape.members.get(exception_type)
if member is not None:
modeled_status = (
(member.metadata or {}).get("error", {}).get("httpStatusCode")
)
if modeled_status is not None:
status_code = int(modeled_status)
return BedrockError(status_code=status_code, message=message)
class BedrockEventStreamDecoderBase:
"""
Base class for event stream decoding for Bedrock
@ -1156,23 +1201,7 @@ class BedrockEventStreamDecoderBase:
parsed_response = self.parser.parse(response_dict, response_stream_shape)
if response_dict["status_code"] != 200:
decoded_body = response_dict["body"].decode()
if isinstance(decoded_body, dict):
error_message = decoded_body.get("message")
elif isinstance(decoded_body, str):
error_message = decoded_body
else:
error_message = ""
exception_status = response_dict["headers"].get(":exception-type")
error_message = exception_status + " " + error_message
raise BedrockError(
status_code=response_dict["status_code"],
message=(
json.dumps(error_message)
if isinstance(error_message, dict)
else error_message
),
)
raise build_bedrock_stream_error(response_dict, response_stream_shape)
if "chunk" in parsed_response:
chunk = parsed_response.get("chunk")
if not chunk:

View file

@ -1,26 +1,15 @@
import json
import time
from typing import AsyncIterator, Iterator, List, Optional, Union
from typing import List, Optional, Union
import httpx
import litellm
from litellm.litellm_core_utils.url_utils import encode_url_path_segments
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
from litellm.llms.base_llm.chat.transformation import (
BaseConfig,
BaseLLMException,
LiteLLMLoggingObj,
from litellm._logging import verbose_logger
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
from litellm.secret_managers.main import (
get_secret_str,
normalize_nonempty_secret_str,
)
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import (
ChatCompletionToolCallChunk,
ChatCompletionUsageBlock,
GenericStreamingChunk,
ModelResponse,
Usage,
)
class CloudflareError(BaseLLMException):
@ -34,26 +23,46 @@ class CloudflareError(BaseLLMException):
message=message,
request=self.request,
response=self.response,
) # Call the base class constructor with the parameters it needs
)
class CloudflareChatConfig(BaseConfig):
max_tokens: Optional[int] = None
stream: Optional[bool] = None
def __init__(
class CloudflareChatConfig(OpenAIGPTConfig):
def get_complete_url(
self,
max_tokens: Optional[int] = None,
api_base: Optional[str],
api_key: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> None:
locals_ = locals().copy()
for key, value in locals_.items():
if key != "self" and value is not None:
setattr(self.__class__, key, value)
) -> str:
return super().get_complete_url(
api_base=self._resolve_api_base(api_base),
api_key=api_key,
model=model,
optional_params=optional_params,
litellm_params=litellm_params,
stream=stream,
)
@classmethod
def get_config(cls):
return super().get_config()
@staticmethod
def _resolve_api_base(api_base: Optional[str]) -> str:
if not api_base:
account_id = normalize_nonempty_secret_str(
get_secret_str("CLOUDFLARE_ACCOUNT_ID")
)
if account_id is None:
raise ValueError(
"Missing CLOUDFLARE_ACCOUNT_ID - set CLOUDFLARE_ACCOUNT_ID in the environment or pass api_base explicitly"
)
return f"https://api.cloudflare.com/client/v4/accounts/{account_id}/ai/v1"
trimmed = api_base.rstrip("/")
if trimmed.endswith("/ai/run"):
verbose_logger.warning(
"Cloudflare api_base ending in '/ai/run' is the legacy Workers AI path and no longer serves OpenAI-compatible requests; rewriting to the '/ai/v1' endpoint"
)
return f"{trimmed[: -len('/ai/run')]}/ai/v1"
return api_base
def validate_environment(
self,
@ -67,107 +76,18 @@ class CloudflareChatConfig(BaseConfig):
) -> dict:
if api_key is None:
raise ValueError(
"Missing CloudflareError API Key - A call is being made to cloudflare but no key is set either in the environment variables or via params"
"Missing Cloudflare API Key - A call is being made to cloudflare but no key is set either in the environment variables or via params"
)
headers = {
"accept": "application/json",
"content-type": "apbplication/json",
"Authorization": "Bearer " + api_key,
}
return headers
def get_complete_url(
self,
api_base: Optional[str],
api_key: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
if api_base is None:
account_id = get_secret_str("CLOUDFLARE_ACCOUNT_ID")
api_base = (
f"https://api.cloudflare.com/client/v4/accounts/{account_id}/ai/run/"
)
encoded_model = encode_url_path_segments(model, field_name="model")
return api_base + encoded_model
def get_supported_openai_params(self, model: str) -> List[str]:
return [
"stream",
"max_tokens",
]
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
supported_openai_params = self.get_supported_openai_params(model=model)
for param, value in non_default_params.items():
if param == "max_completion_tokens":
optional_params["max_tokens"] = value
elif param in supported_openai_params:
optional_params[param] = value
return optional_params
def transform_request(
self,
model: str,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
headers: dict,
) -> dict:
config = litellm.CloudflareChatConfig.get_config()
for k, v in config.items():
if k not in optional_params:
optional_params[k] = v
data = {
"messages": messages,
**optional_params,
}
return data
def transform_response(
self,
model: str,
raw_response: httpx.Response,
model_response: ModelResponse,
logging_obj: LiteLLMLoggingObj,
request_data: dict,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
encoding: str,
api_key: Optional[str] = None,
json_mode: Optional[bool] = None,
) -> ModelResponse:
completion_response = raw_response.json()
# Support both "response" and "response_text" keys (newer models like Nemotron use "response_text")
result = completion_response["result"]
model_response.choices[0].message.content = result.get("response") if result.get("response") is not None else result.get("response_text", "") # type: ignore
prompt_tokens = litellm.utils.get_token_count(messages=messages, model=model)
completion_tokens = len(
encoding.encode(model_response["choices"][0]["message"].get("content", ""))
return super().validate_environment(
headers=headers,
model=model,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
api_key=api_key,
api_base=api_base,
)
model_response.created = int(time.time())
model_response.model = "cloudflare/" + model
usage = Usage(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens,
)
setattr(model_response, "usage", usage)
return model_response
def get_error_class(
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
) -> BaseLLMException:
@ -175,48 +95,3 @@ class CloudflareChatConfig(BaseConfig):
status_code=status_code,
message=error_message,
)
def get_model_response_iterator(
self,
streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse],
sync_stream: bool,
json_mode: Optional[bool] = False,
):
return CloudflareChatResponseIterator(
streaming_response=streaming_response,
sync_stream=sync_stream,
json_mode=json_mode,
)
class CloudflareChatResponseIterator(BaseModelResponseIterator):
def chunk_parser(self, chunk: dict) -> GenericStreamingChunk:
try:
text = ""
tool_use: Optional[ChatCompletionToolCallChunk] = None
is_finished = False
finish_reason = ""
usage: Optional[ChatCompletionUsageBlock] = None
provider_specific_fields = None
index = int(chunk.get("index", 0))
if "response" in chunk and chunk["response"] is not None:
text = chunk["response"]
elif "response_text" in chunk and chunk["response_text"] is not None:
text = chunk["response_text"]
returned_chunk = GenericStreamingChunk(
text=text,
tool_use=tool_use,
is_finished=is_finished,
finish_reason=finish_reason,
usage=usage,
index=index,
provider_specific_fields=provider_specific_fields,
)
return returned_chunk
except json.JSONDecodeError:
raise ValueError(f"Failed to decode JSON from chunk: {chunk}")

View file

@ -1,5 +1,6 @@
import json
import ssl
from functools import lru_cache
from urllib.parse import parse_qs, urlencode, urlparse, urlunparse
from typing import (
TYPE_CHECKING,
@ -13,6 +14,7 @@ from typing import (
Tuple,
Union,
cast,
get_type_hints,
)
import httpx # type: ignore
@ -26,6 +28,7 @@ from litellm._logging import _redact_string, verbose_logger
from litellm.anthropic_beta_headers_manager import update_headers_with_filtered_beta
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming
from litellm.litellm_core_utils.asyncify import run_async_function
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
from litellm.llms.base_llm.anthropic_messages.transformation import (
BaseAnthropicMessagesConfig,
@ -101,6 +104,7 @@ from litellm.types.llms.openai import (
HttpxBinaryResponseContent,
OpenAIFileObject,
ResponseInputParam,
ResponsesAPIOptionalRequestParams,
ResponsesAPIResponse,
)
from litellm.types.rerank import RerankResponse
@ -135,6 +139,7 @@ from litellm.utils import (
ImageResponse,
ModelResponse,
ProviderConfigManager,
async_pre_call_deployment_hook,
)
from .http_handler import get_shared_realtime_ssl_context
@ -184,6 +189,47 @@ def _google_genai_streaming_hidden_params(
}
@lru_cache(maxsize=None)
def _responses_api_optional_request_param_names() -> frozenset[str]:
return frozenset(get_type_hints(ResponsesAPIOptionalRequestParams).keys())
def _custom_logger_callbacks(logging_obj: Any) -> list[Any]:
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import (
get_custom_logger_compatible_class,
)
dynamic_success_callbacks = getattr(logging_obj, "dynamic_success_callbacks", None)
callbacks = list(litellm.callbacks)
if isinstance(dynamic_success_callbacks, (list, tuple)):
callbacks.extend(dynamic_success_callbacks)
custom_loggers: list[Any] = []
for cb in callbacks:
if isinstance(cb, str):
resolved = get_custom_logger_compatible_class(cb) # type: ignore[arg-type]
if resolved is None:
continue
cb = resolved
if isinstance(cb, CustomLogger):
custom_loggers.append(cb)
return custom_loggers
def _has_pre_call_deployment_hook(logging_obj: Any) -> bool:
from litellm.integrations.custom_logger import CustomLogger
base_func = CustomLogger.async_pre_call_deployment_hook
for cb in _custom_logger_callbacks(logging_obj):
cb_func = getattr(type(cb), "async_pre_call_deployment_hook", base_func)
if getattr(cb_func, "__func__", cb_func) is not getattr(
base_func, "__func__", base_func
):
return True
return False
class BaseLLMHTTPHandler:
async def _make_common_async_call(
self,
@ -2224,12 +2270,92 @@ class BaseLLMHTTPHandler:
)
raise ValueError("anthropic_messages_handler is not implemented for sync calls")
def _run_sync_responses_pre_call_deployment_hook(
self,
*,
model: str,
input: Union[str, ResponseInputParam],
custom_llm_provider: str,
response_api_optional_request_params: dict[str, Any],
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
) -> tuple[
str,
Union[str, ResponseInputParam],
str,
dict[str, Any],
GenericLiteLLMParams,
]:
if not _has_pre_call_deployment_hook(logging_obj):
return (
model,
input,
custom_llm_provider,
response_api_optional_request_params,
litellm_params,
)
modified_kwargs = run_async_function(
async_pre_call_deployment_hook,
{
**dict(litellm_params),
**response_api_optional_request_params,
"model": model,
"input": input,
"custom_llm_provider": custom_llm_provider,
},
CallTypes.responses.value,
)
if modified_kwargs is None:
return (
model,
input,
custom_llm_provider,
response_api_optional_request_params,
litellm_params,
)
optional_param_names = _responses_api_optional_request_param_names()
updated_response_params = {
**response_api_optional_request_params,
**{
key: value
for key, value in modified_kwargs.items()
if key in optional_param_names
},
}
updated_litellm_params = GenericLiteLLMParams(
**{
**dict(litellm_params),
**{
key: value
for key, value in modified_kwargs.items()
if key not in optional_param_names
and key not in {"model", "input", "custom_llm_provider"}
},
}
)
return (
str(modified_kwargs["model"]) if "model" in modified_kwargs else model,
cast(
Union[str, ResponseInputParam],
modified_kwargs["input"] if "input" in modified_kwargs else input,
),
(
str(modified_kwargs["custom_llm_provider"])
if "custom_llm_provider" in modified_kwargs
else custom_llm_provider
),
updated_response_params,
updated_litellm_params,
)
def response_api_handler(
self,
model: str,
input: Union[str, ResponseInputParam],
responses_api_provider_config: BaseResponsesAPIConfig,
response_api_optional_request_params: Dict,
response_api_optional_request_params: dict[str, Any],
custom_llm_provider: str,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
@ -2276,6 +2402,21 @@ class BaseLLMHTTPHandler:
shared_session=shared_session,
)
(
model,
input,
custom_llm_provider,
response_api_optional_request_params,
litellm_params,
) = self._run_sync_responses_pre_call_deployment_hook(
model=model,
input=input,
custom_llm_provider=custom_llm_provider,
response_api_optional_request_params=response_api_optional_request_params,
litellm_params=litellm_params,
logging_obj=logging_obj,
)
if client is None or not isinstance(client, HTTPHandler):
sync_httpx_client = _get_httpx_client(
params={"ssl_verify": litellm_params.get("ssl_verify", None)}
@ -2414,9 +2555,27 @@ class BaseLLMHTTPHandler:
logging_obj=logging_obj,
)
)
# Responses agentic interception (e.g. code interpreter) runs the follow-up
# loop via the async hook, so it is async-only for now; the sync path returns
# the initial response unchanged.
if self._has_agentic_completion_hook(logging_obj):
final_response = run_async_function(
self._call_agentic_completion_hooks,
response=initial_response,
model=model,
messages=(
input
if isinstance(input, list)
else [{"role": "user", "content": input}]
),
anthropic_messages_provider_config=responses_api_provider_config,
anthropic_messages_optional_request_params=response_api_optional_request_params,
logging_obj=logging_obj,
stream=False,
custom_llm_provider=custom_llm_provider,
kwargs=dict(litellm_params),
api_surface="responses",
)
return final_response if final_response is not None else initial_response
return initial_response
async def async_response_api_handler(
@ -4772,22 +4931,9 @@ class BaseLLMHTTPHandler:
agentic callback is detected too.
"""
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import (
get_custom_logger_compatible_class,
)
base_func = CustomLogger.async_should_run_agentic_loop
callbacks = litellm.callbacks + (
getattr(logging_obj, "dynamic_success_callbacks", None) or []
)
for cb in callbacks:
if isinstance(cb, str):
resolved = get_custom_logger_compatible_class(cb) # type: ignore[arg-type]
if resolved is None:
continue
cb = resolved
if not isinstance(cb, CustomLogger):
continue
for cb in _custom_logger_callbacks(logging_obj):
cb_func = getattr(type(cb), "async_should_run_agentic_loop", base_func)
if getattr(cb_func, "__func__", cb_func) is not getattr(
base_func, "__func__", base_func
@ -5537,9 +5683,7 @@ class BaseLLMHTTPHandler:
import websockets
from websockets.asyncio.client import ClientConnection
url = self._append_query_params(
provider_config.get_complete_url(api_base, model, api_key), query_params
)
url = provider_config.get_complete_url(api_base, model, api_key)
headers = provider_config.validate_environment(
headers=headers,
model=model,

View file

@ -16,6 +16,7 @@ from litellm.llms.base_llm.sandbox.transformation import (
BaseSandboxConfig,
CodeExecutionResult,
ContainerHandle,
SANDBOX_MAX_OUTPUT_BYTES,
)
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
@ -29,7 +30,7 @@ E2B_DEFAULT_TEMPLATE = "code-interpreter-v1"
E2B_DEFAULT_DOMAIN = "e2b.app"
JUPYTER_PORT = 49999
DEFAULT_SANDBOX_TIMEOUT = 300
MAX_OUTPUT_BYTES = 10 * 1024 * 1024
MAX_OUTPUT_BYTES = SANDBOX_MAX_OUTPUT_BYTES
class E2BSandboxConfig(BaseSandboxConfig):
@ -49,7 +50,7 @@ class E2BSandboxConfig(BaseSandboxConfig):
*,
template: str | None = None,
timeout: int | None = None,
allow_internet_access: bool = True,
allow_internet_access: bool | None = None,
api_key: str | None = None,
api_base: str | None = None,
metadata: dict | None = None,
@ -62,7 +63,9 @@ class E2BSandboxConfig(BaseSandboxConfig):
"templateID": template or E2B_DEFAULT_TEMPLATE,
"timeout": timeout if timeout is not None else DEFAULT_SANDBOX_TIMEOUT,
"secure": True,
"allow_internet_access": allow_internet_access,
"allow_internet_access": (
True if allow_internet_access is None else allow_internet_access
),
}
if metadata:
body["metadata"] = metadata
@ -168,20 +171,6 @@ class E2BSandboxConfig(BaseSandboxConfig):
handle._hidden_params = {}
return handle
@staticmethod
async def _read_capped_lines(response: httpx.Response) -> list[str]:
lines: list[str] = []
total = 0
async for line in response.aiter_lines():
total += len(line.encode("utf-8"))
if total > MAX_OUTPUT_BYTES:
raise ValueError(
f"Sandbox output exceeded {MAX_OUTPUT_BYTES} bytes; aborting to "
"avoid unbounded memory use."
)
lines.append(line)
return lines
@staticmethod
def _parse_lines(lines: list[str]) -> CodeExecutionResult:
def _try_parse(stripped: str):
@ -192,10 +181,9 @@ class E2BSandboxConfig(BaseSandboxConfig):
messages = tuple(
parsed
for stripped in (line.strip() for line in lines)
if stripped
for parsed in (_try_parse(stripped),)
if parsed is not None
for line in lines
if (stripped := line.strip())
if (parsed := _try_parse(stripped)) is not None
)
def of_type(message_type: str):

View file

@ -1,17 +0,0 @@
from typing import List
from litellm.types.llms.openai import OpenAIAudioTranscriptionOptionalParams
from ...openai.transcriptions.whisper_transformation import (
OpenAIWhisperAudioTranscriptionConfig,
)
from ..common_utils import FireworksAIMixin
class FireworksAIAudioTranscriptionConfig(
FireworksAIMixin, OpenAIWhisperAudioTranscriptionConfig
):
def get_supported_openai_params(
self, model: str
) -> List[OpenAIAudioTranscriptionOptionalParams]:
return ["language", "prompt", "response_format", "timestamp_granularities"]

View file

@ -103,6 +103,10 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
# bypassing spend and budget accounting.
self._pending_usage_metadata: Optional[dict] = None
def _include_function_response_id(self) -> bool:
"""Google AI Studio Gemini 3.5+ accepts ``id`` on functionResponses; Vertex AI rejects it."""
return True
@staticmethod
def _usage_detail_alias(details: Any, defaults: Dict[str, int]) -> Dict[str, Any]:
if not isinstance(details, dict):
@ -604,10 +608,9 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
)
# Build Gemini toolResponse format
function_response = {
"id": call_id,
"response": output_dict,
}
function_response: dict[str, Any] = {"response": output_dict}
if self._include_function_response_id() and call_id:
function_response["id"] = call_id
if function_name:
function_response["name"] = function_name

View file

@ -247,6 +247,8 @@ class MistralConfig(OpenAIGPTConfig):
The above statement is not valid now. Need to plan to remove all the #1,2,3
Mistral API supports content as a list.
"""
messages = [self._strip_output_only_fields(m) for m in messages]
## 1. If 'image_url' or 'file' in content, then transform with base class and mistral-specific handling
for m in messages:
_content_block = m.get("content")
@ -409,6 +411,25 @@ class MistralConfig(OpenAIGPTConfig):
return cleaned_tools
@classmethod
def _strip_output_only_fields(cls, message: AllMessageValues) -> AllMessageValues:
"""
``reasoning_content`` and ``thinking_blocks`` are output-only fields that
LiteLLM attaches to assistant responses. Mistral's input schema forbids
unknown fields, so replaying them verbatim in a follow-up turn triggers a
422 ``extra_forbidden``. Drop them before the request is sent.
"""
if message["role"] != "assistant":
return message
return cast(
AllMessageValues,
{
k: v
for k, v in message.items()
if k not in ("reasoning_content", "thinking_blocks")
},
)
@classmethod
def _handle_name_in_message(cls, message: AllMessageValues) -> AllMessageValues:
"""

View file

@ -115,6 +115,14 @@
"max_completion_tokens": "max_tokens"
}
},
"darkbloom": {
"base_url": "https://api.darkbloom.dev/v1",
"api_key_env": "DARKBLOOM_API_KEY",
"api_base_env": "DARKBLOOM_API_BASE",
"param_mappings": {
"max_completion_tokens": "max_tokens"
}
},
"neosantara": {
"base_url": "https://api.neosantara.xyz/v1",
"api_key_env": "NEOSANTARA_API_KEY",

View file

@ -0,0 +1 @@

View file

@ -0,0 +1 @@

View file

@ -0,0 +1,598 @@
import asyncio
import json
import time
from typing import Union, cast
import httpx
from litellm.constants import (
OPEN_SANDBOX_API_BASE_ENV_VAR,
OPEN_SANDBOX_API_KEY_ENV_VAR,
OPEN_SANDBOX_DEFAULT_CPU_LIMIT,
OPEN_SANDBOX_DEFAULT_ENTRYPOINT,
OPEN_SANDBOX_DEFAULT_LANGUAGE,
OPEN_SANDBOX_DEFAULT_MEMORY_LIMIT,
OPEN_SANDBOX_DEFAULT_TEMPLATE,
OPEN_SANDBOX_DEFAULT_TIMEOUT,
OPEN_SANDBOX_EXECD_PORT,
OPEN_SANDBOX_POLL_INTERVAL,
OPEN_SANDBOX_READY_TIMEOUT,
)
from litellm.llms.base_llm.sandbox.transformation import (
BaseSandboxConfig,
CodeExecutionResult,
ContainerHandle,
SANDBOX_MAX_OUTPUT_BYTES,
)
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
get_async_httpx_client,
)
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.custom_http import httpxSpecialProvider
DEFAULT_SANDBOX_TIMEOUT = OPEN_SANDBOX_DEFAULT_TIMEOUT
DEFAULT_READY_TIMEOUT = OPEN_SANDBOX_READY_TIMEOUT
DEFAULT_POLL_INTERVAL = OPEN_SANDBOX_POLL_INTERVAL
MAX_OUTPUT_BYTES = SANDBOX_MAX_OUTPUT_BYTES
class OpenSandboxSandboxConfig(BaseSandboxConfig):
def _http(self, client: AsyncHTTPHandler | None) -> AsyncHTTPHandler:
if client is not None:
return client
return get_async_httpx_client(llm_provider=httpxSpecialProvider.Sandbox)
def validate_environment(self, api_key: str | None = None, **kwargs) -> str:
if api_key is not None:
return api_key
return get_secret_str(OPEN_SANDBOX_API_KEY_ENV_VAR) or ""
async def acreate_sandbox(
self,
*,
template: str | None = None,
timeout: int | None = None,
allow_internet_access: bool | None = None,
api_key: str | None = None,
api_base: str | None = None,
metadata: dict[str, str] | None = None,
env_vars: dict[str, str] | None = None,
resource_limits: dict[str, str] | None = None,
resource_requests: dict[str, str] | None = None,
entrypoint: list[str] | tuple[str, ...] | None = None,
network_policy: dict[str, object] | None = None,
secure_access: bool = False,
use_server_proxy: bool = False,
ready_timeout: float | None = None,
poll_interval: float | None = None,
client: AsyncHTTPHandler | None = None,
**kwargs,
) -> ContainerHandle:
key = self.validate_environment(api_key=api_key)
base = self._api_base(api_base)
ready_timeout_seconds = (
float(ready_timeout) if ready_timeout is not None else DEFAULT_READY_TIMEOUT
)
poll_interval_seconds = (
float(poll_interval) if poll_interval is not None else DEFAULT_POLL_INTERVAL
)
body = self._create_body(
template=template,
timeout=timeout,
allow_internet_access=allow_internet_access,
metadata=metadata,
env_vars=env_vars,
resource_limits=resource_limits,
resource_requests=resource_requests,
entrypoint=entrypoint,
network_policy=network_policy,
secure_access=secure_access,
)
response = cast(
httpx.Response,
await self._http(client).post(
url=f"{base}/sandboxes",
headers=self._lifecycle_headers(key),
json=body,
),
)
data = response.json()
sandbox_id = str(data["id"])
if self._sandbox_state(data) != "Running":
await self._wait_until_running(
sandbox_id=sandbox_id,
api_base=base,
headers=self._lifecycle_headers(key),
client=client,
ready_timeout=ready_timeout_seconds,
poll_interval=poll_interval_seconds,
)
endpoint, endpoint_headers = await self._wait_for_execd_endpoint(
sandbox_id=sandbox_id,
api_base=base,
headers=self._lifecycle_headers(key),
use_server_proxy=use_server_proxy,
client=client,
ready_timeout=ready_timeout_seconds,
poll_interval=poll_interval_seconds,
)
handle = ContainerHandle(id=sandbox_id, provider="opensandbox", domain=base)
handle._hidden_params = {
"api_base": base,
"api_key": key,
"execd_endpoint": endpoint,
"execd_headers": endpoint_headers,
"use_server_proxy": use_server_proxy,
}
return handle
async def arun_code(
self,
*,
container: Union[ContainerHandle, str],
code: str,
api_key: str | None = None,
api_base: str | None = None,
language: str = OPEN_SANDBOX_DEFAULT_LANGUAGE,
use_server_proxy: bool = False,
ready_timeout: float | None = None,
poll_interval: float | None = None,
client: AsyncHTTPHandler | None = None,
**kwargs,
) -> CodeExecutionResult:
handle = await self._ensure_handle(
container=container,
api_key=api_key,
api_base=api_base,
use_server_proxy=use_server_proxy,
ready_timeout=(
float(ready_timeout)
if ready_timeout is not None
else DEFAULT_READY_TIMEOUT
),
poll_interval=(
float(poll_interval)
if poll_interval is not None
else DEFAULT_POLL_INTERVAL
),
client=client,
)
endpoint = str(handle._hidden_params["execd_endpoint"])
endpoint_headers = self._as_str_dict(handle._hidden_params.get("execd_headers"))
base = str(
handle._hidden_params.get("api_base")
or handle.domain
or self._api_base(api_base)
)
lines = await self._post_code(
url=f"{self._endpoint_base_url(endpoint, base)}/code",
headers={
"Content-Type": "application/json",
"Accept": "text/event-stream",
"Cache-Control": "no-cache",
**endpoint_headers,
},
body={
"code": code,
"context": {"language": language},
},
client=client,
)
return self._parse_lines(lines)
async def adelete_sandbox(
self,
*,
container: Union[ContainerHandle, str],
api_key: str | None = None,
api_base: str | None = None,
client: AsyncHTTPHandler | None = None,
**kwargs,
) -> bool:
handle = self._as_handle(container, api_base=api_base)
base = str(handle._hidden_params.get("api_base") or self._api_base(api_base))
key = self._api_key(api_key=api_key, handle=handle)
try:
response = cast(
httpx.Response,
await self._http(client).delete(
url=f"{base}/sandboxes/{handle.id}",
headers=self._lifecycle_headers(key),
),
)
except httpx.HTTPStatusError as e:
if e.response.status_code == 404:
return False
raise
return 200 <= response.status_code < 300
async def _ensure_handle(
self,
*,
container: Union[ContainerHandle, str],
api_key: str | None,
api_base: str | None,
use_server_proxy: bool,
ready_timeout: float,
poll_interval: float,
client: AsyncHTTPHandler | None,
) -> ContainerHandle:
handle = self._as_handle(container, api_base=api_base)
if handle._hidden_params.get("execd_endpoint"):
return handle
base = str(handle._hidden_params.get("api_base") or self._api_base(api_base))
key = self._api_key(api_key=api_key, handle=handle)
resolved_use_server_proxy = bool(
handle._hidden_params.get("use_server_proxy", use_server_proxy)
)
endpoint, endpoint_headers = await self._wait_for_execd_endpoint(
sandbox_id=handle.id,
api_base=base,
headers=self._lifecycle_headers(key),
use_server_proxy=resolved_use_server_proxy,
client=client,
ready_timeout=ready_timeout,
poll_interval=poll_interval,
)
handle.domain = base
handle._hidden_params = {
**handle._hidden_params,
"api_base": base,
"api_key": key,
"execd_endpoint": endpoint,
"execd_headers": endpoint_headers,
"use_server_proxy": resolved_use_server_proxy,
}
return handle
async def _wait_until_running(
self,
*,
sandbox_id: str,
api_base: str,
headers: dict[str, str],
client: AsyncHTTPHandler | None,
ready_timeout: float,
poll_interval: float,
) -> None:
deadline = time.monotonic() + ready_timeout
while True:
response = cast(
httpx.Response,
await self._http(client).get(
url=f"{api_base}/sandboxes/{sandbox_id}",
headers=headers,
),
)
data = response.json()
state = self._sandbox_state(data)
if state == "Running":
return
if state in {"Failed", "Stopping", "Terminated"}:
raise ValueError(f"OpenSandbox sandbox {sandbox_id} entered {state}")
if time.monotonic() >= deadline:
raise TimeoutError(
f"OpenSandbox sandbox {sandbox_id} was not Running within "
f"{ready_timeout} seconds"
)
await asyncio.sleep(poll_interval)
async def _wait_for_execd_endpoint(
self,
*,
sandbox_id: str,
api_base: str,
headers: dict[str, str],
use_server_proxy: bool,
client: AsyncHTTPHandler | None,
ready_timeout: float,
poll_interval: float,
) -> tuple[str, dict[str, str]]:
deadline = time.monotonic() + ready_timeout
last_error: Exception | None = None
while True:
try:
return await self._get_execd_endpoint(
sandbox_id=sandbox_id,
api_base=api_base,
headers=headers,
use_server_proxy=use_server_proxy,
client=client,
)
except httpx.HTTPStatusError as e:
if e.response.status_code != 404:
raise
last_error = e
except ValueError as e:
last_error = e
if time.monotonic() >= deadline:
raise TimeoutError(
f"OpenSandbox execd endpoint for {sandbox_id} was not ready within "
f"{ready_timeout} seconds"
) from last_error
await asyncio.sleep(poll_interval)
async def _get_execd_endpoint(
self,
*,
sandbox_id: str,
api_base: str,
headers: dict[str, str],
use_server_proxy: bool,
client: AsyncHTTPHandler | None,
) -> tuple[str, dict[str, str]]:
response = cast(
httpx.Response,
await self._http(client).get(
url=f"{api_base}/sandboxes/{sandbox_id}/endpoints/{OPEN_SANDBOX_EXECD_PORT}",
headers=headers,
params={"use_server_proxy": use_server_proxy},
),
)
data = response.json()
endpoint = data.get("endpoint")
if not endpoint:
raise ValueError(
f"OpenSandbox did not return an execd endpoint for {sandbox_id}"
)
return str(endpoint), self._as_str_dict(data.get("headers"))
async def _post_code(
self,
*,
url: str,
headers: dict[str, str],
body: dict[str, object],
client: AsyncHTTPHandler | None,
) -> list[str]:
timeout = httpx.Timeout(connect=30.0, read=None, write=30.0, pool=None)
response = cast(
httpx.Response,
await self._http(client).post(
url=url,
headers=headers,
timeout=timeout,
json=body,
stream=True,
),
)
return await self._read_capped_lines(response)
def _api_key(self, *, api_key: str | None, handle: ContainerHandle) -> str:
if api_key is not None:
return api_key
if "api_key" in handle._hidden_params:
return str(handle._hidden_params["api_key"])
return self.validate_environment()
@staticmethod
def _create_body(
*,
template: str | None,
timeout: int | None,
allow_internet_access: bool | None,
metadata: dict[str, str] | None,
env_vars: dict[str, str] | None,
resource_limits: dict[str, str] | None,
resource_requests: dict[str, str] | None,
entrypoint: list[str] | tuple[str, ...] | None,
network_policy: dict[str, object] | None,
secure_access: bool,
) -> dict[str, object]:
body: dict[str, object] = {
"image": {"uri": template or OPEN_SANDBOX_DEFAULT_TEMPLATE},
"entrypoint": list(entrypoint or OPEN_SANDBOX_DEFAULT_ENTRYPOINT),
"timeout": timeout if timeout is not None else DEFAULT_SANDBOX_TIMEOUT,
"resourceLimits": resource_limits
or OpenSandboxSandboxConfig._default_resource_limits(),
}
if metadata:
body["metadata"] = metadata
if env_vars:
body["env"] = env_vars
if resource_requests:
body["resourceRequests"] = resource_requests
if network_policy is not None:
body["networkPolicy"] = network_policy
elif allow_internet_access is not True:
body["networkPolicy"] = {"defaultAction": "deny", "egress": []}
if secure_access:
body["secureAccess"] = True
return body
@staticmethod
def _default_resource_limits() -> dict[str, str]:
return {
"cpu": OPEN_SANDBOX_DEFAULT_CPU_LIMIT,
"memory": OPEN_SANDBOX_DEFAULT_MEMORY_LIMIT,
}
@staticmethod
def _sandbox_state(data: object) -> str | None:
if not isinstance(data, dict):
return None
status = data.get("status")
if not isinstance(status, dict):
return None
state = status.get("state")
return str(state) if state is not None else None
@staticmethod
def _as_str_dict(value: object) -> dict[str, str]:
if not isinstance(value, dict):
return {}
return {str(k): str(v) for k, v in value.items()}
@staticmethod
def _api_base(api_base: str | None) -> str:
base = api_base or get_secret_str(OPEN_SANDBOX_API_BASE_ENV_VAR)
if not base:
raise ValueError(
"OpenSandbox api_base is required. Pass api_base or set "
f"{OPEN_SANDBOX_API_BASE_ENV_VAR}."
)
return str(base).rstrip("/")
@staticmethod
def _lifecycle_headers(api_key: str) -> dict[str, str]:
headers = {"Content-Type": "application/json"}
if api_key:
headers["OPEN-SANDBOX-API-KEY"] = api_key
return headers
@staticmethod
def _endpoint_base_url(endpoint: str, api_base: str) -> str:
normalized_endpoint = endpoint.rstrip("/")
if normalized_endpoint.startswith(("http://", "https://")):
return normalized_endpoint
protocol = api_base.split("://", 1)[0] if "://" in api_base else "http"
return f"{protocol}://{normalized_endpoint}"
@staticmethod
def _as_handle(
container: Union[ContainerHandle, str], *, api_base: str | None
) -> ContainerHandle:
if isinstance(container, ContainerHandle):
return container
handle = ContainerHandle(
id=str(container),
provider="opensandbox",
domain=OpenSandboxSandboxConfig._api_base(api_base),
)
handle._hidden_params = {}
return handle
@staticmethod
def _parse_lines(lines: list[str]) -> CodeExecutionResult:
messages = tuple(
event
for line in lines
if (event := OpenSandboxSandboxConfig._parse_sse_line(line)) is not None
)
def of_type(message_type: str):
return (m for m in messages if m.get("type") == message_type)
error = next(
(OpenSandboxSandboxConfig._normalize_error(m) for m in of_type("error")),
None,
)
execution_count = next(
(
OpenSandboxSandboxConfig._as_int(m.get("execution_count"))
for m in of_type("execution_count")
if OpenSandboxSandboxConfig._as_int(m.get("execution_count"))
is not None
),
None,
)
return CodeExecutionResult(
stdout="".join(str(m.get("text", "")) for m in of_type("stdout")),
stderr="".join(str(m.get("text", "")) for m in of_type("stderr")),
results=[
OpenSandboxSandboxConfig._normalize_result(m) for m in of_type("result")
],
error=error,
execution_count=execution_count,
)
@staticmethod
def _parse_sse_line(line: str) -> dict[str, object] | None:
stripped = line.strip()
if not stripped or stripped.startswith(
(
":",
"event:",
"id:",
"retry:",
)
):
return None
data = stripped[5:].strip() if stripped.startswith("data:") else stripped
if not data:
return None
try:
parsed = json.loads(data)
except json.JSONDecodeError:
return None
if not isinstance(parsed, dict):
return None
if "type" not in parsed and "code" in parsed and "message" in parsed:
return {
"type": "error",
"error": {
"ename": str(parsed["code"]),
"evalue": str(parsed["message"]),
"traceback": [],
},
}
return parsed
@staticmethod
def _normalize_result(message: dict[str, object]) -> dict[str, object]:
results = message.get("results")
if isinstance(results, dict):
return {str(k): v for k, v in results.items()}
return {
str(k): v
for k, v in message.items()
if k not in {"type", "timestamp", "execution_count"}
}
@staticmethod
def _normalize_error(message: dict[str, object]) -> dict[str, object]:
raw_error = message.get("error")
if isinstance(raw_error, dict):
name = OpenSandboxSandboxConfig._first_non_none_value(
raw_error, "ename", "name", default=""
)
value = OpenSandboxSandboxConfig._first_non_none_value(
raw_error, "evalue", "value", default=""
)
traceback = OpenSandboxSandboxConfig._first_non_none_value(
raw_error, "traceback", default=[]
)
return {
"name": name,
"value": value,
"traceback": traceback,
}
return {
"name": OpenSandboxSandboxConfig._first_non_none_value(
message, "name", default=""
),
"value": OpenSandboxSandboxConfig._first_non_none_value(
message, "value", "text", default=""
),
"traceback": OpenSandboxSandboxConfig._first_non_none_value(
message, "traceback", default=[]
),
}
@staticmethod
def _as_int(value: object) -> int | None:
if isinstance(value, int):
return value
if isinstance(value, str):
try:
return int(value)
except ValueError:
return None
return None
@staticmethod
def _first_non_none_value(
values: dict[str, object], *keys: str, default: object
) -> object:
return next(
(values[key] for key in keys if key in values and values[key] is not None),
default,
)

View file

@ -98,10 +98,11 @@ def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]:
if num_search_queries > 0 and search_cost_value is not None:
# Handle both dict and float formats
if isinstance(search_cost_value, dict):
# Use the "low" size as default - tests expect 0.005 / 1000
search_cost_per_query = (
_safe_float_cast(search_cost_value.get("search_context_size_low", 0))
/ 1000
# search_context_cost_per_query stores the per-request price in USD
# (e.g. sonar low = $0.005/request). Use it directly, matching the
# gemini cost calculator which reads the same field per request.
search_cost_per_query = _safe_float_cast(
search_cost_value.get("search_context_size_low", 0)
)
else:
search_cost_per_query = _safe_float_cast(search_cost_value)

View file

@ -32,6 +32,9 @@ class VertexAIRealtimeConfig(GeminiRealtimeConfig):
self._project = project
self._location = location
def _include_function_response_id(self) -> bool:
return False
# ------------------------------------------------------------------
# URL
# ------------------------------------------------------------------

File diff suppressed because it is too large Load diff

View file

@ -571,7 +571,7 @@
"output_vector_size": 1536
},
"amazon.titan-embed-text-v2:0": {
"input_cost_per_token": 2e-07,
"input_cost_per_token": 2e-08,
"litellm_provider": "bedrock",
"max_input_tokens": 8192,
"max_tokens": 8192,
@ -10684,6 +10684,268 @@
"mode": "chat",
"output_cost_per_token": 1.923e-06
},
"cloudflare/@cf/openai/gpt-oss-120b": {
"input_cost_per_token": 3.5e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 7.5e-07,
"supports_function_calling": true,
"supports_reasoning": true
},
"cloudflare/@cf/google/gemma-2b-it-lora": {
"input_cost_per_token": 0.0,
"litellm_provider": "cloudflare",
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 0.0
},
"cloudflare/@cf/meta/llama-3.2-3b-instruct": {
"input_cost_per_token": 5.09e-08,
"litellm_provider": "cloudflare",
"max_input_tokens": 80000,
"max_output_tokens": 80000,
"max_tokens": 80000,
"mode": "chat",
"output_cost_per_token": 3.35e-07
},
"cloudflare/@cf/meta/llama-guard-3-8b": {
"input_cost_per_token": 4.84e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 3e-08
},
"cloudflare/@cf/mistral/mistral-7b-instruct-v0.2-lora": {
"input_cost_per_token": 0.0,
"litellm_provider": "cloudflare",
"max_input_tokens": 15000,
"max_output_tokens": 15000,
"max_tokens": 15000,
"mode": "chat",
"output_cost_per_token": 0.0
},
"cloudflare/@cf/moonshotai/kimi-k2.7-code": {
"cache_read_input_token_cost": 1.9e-07,
"input_cost_per_token": 9.5e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 4e-06,
"supports_function_calling": true,
"supports_reasoning": true
},
"cloudflare/@cf/deepseek-ai/deepseek-r1-distill-qwen-32b": {
"input_cost_per_token": 4.97e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 80000,
"max_output_tokens": 80000,
"max_tokens": 80000,
"mode": "chat",
"output_cost_per_token": 4.881e-06,
"supports_reasoning": true
},
"cloudflare/@cf/meta/llama-3.1-8b-instruct-fp8": {
"input_cost_per_token": 1.52e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 32000,
"max_output_tokens": 32000,
"max_tokens": 32000,
"mode": "chat",
"output_cost_per_token": 2.87e-07
},
"cloudflare/@cf/meta/llama-3.2-1b-instruct": {
"input_cost_per_token": 2.7e-08,
"litellm_provider": "cloudflare",
"max_input_tokens": 60000,
"max_output_tokens": 60000,
"max_tokens": 60000,
"mode": "chat",
"output_cost_per_token": 2.01e-07
},
"cloudflare/@cf/moonshotai/kimi-k2.6": {
"cache_read_input_token_cost": 1.6e-07,
"input_cost_per_token": 9.5e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 4e-06,
"supports_function_calling": true,
"supports_reasoning": true
},
"cloudflare/@cf/zai-org/glm-4.7-flash": {
"input_cost_per_token": 6.05e-08,
"litellm_provider": "cloudflare",
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4e-07,
"supports_function_calling": true,
"supports_reasoning": true
},
"cloudflare/@cf/meta-llama/llama-2-7b-chat-hf-lora": {
"input_cost_per_token": 0.0,
"litellm_provider": "cloudflare",
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 0.0
},
"cloudflare/@cf/meta/llama-3.3-70b-instruct-fp8-fast": {
"input_cost_per_token": 2.93e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 24000,
"max_output_tokens": 24000,
"max_tokens": 24000,
"mode": "chat",
"output_cost_per_token": 2.253e-06,
"supports_function_calling": true
},
"cloudflare/@cf/ibm-granite/granite-4.0-h-micro": {
"input_cost_per_token": 1.7e-08,
"litellm_provider": "cloudflare",
"max_input_tokens": 131000,
"max_output_tokens": 131000,
"max_tokens": 131000,
"mode": "chat",
"output_cost_per_token": 1.12e-07,
"supports_function_calling": true
},
"cloudflare/@cf/qwen/qwen2.5-coder-32b-instruct": {
"input_cost_per_token": 6.6e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 32768,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 1e-06
},
"cloudflare/@cf/zai-org/glm-5.2": {
"cache_read_input_token_cost": 2.6e-07,
"input_cost_per_token": 1.4e-06,
"litellm_provider": "cloudflare",
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 4.4e-06,
"supports_function_calling": true,
"supports_reasoning": true
},
"cloudflare/@cf/nvidia/nemotron-3-120b-a12b": {
"input_cost_per_token": 5e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 256000,
"max_output_tokens": 256000,
"max_tokens": 256000,
"mode": "chat",
"output_cost_per_token": 1.5e-06,
"supports_function_calling": true,
"supports_reasoning": true
},
"cloudflare/@cf/aisingapore/gemma-sea-lion-v4-27b-it": {
"input_cost_per_token": 3.51e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 5.55e-07
},
"cloudflare/@cf/qwen/qwen3-30b-a3b-fp8": {
"input_cost_per_token": 5.09e-08,
"litellm_provider": "cloudflare",
"max_input_tokens": 32768,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 3.35e-07,
"supports_function_calling": true,
"supports_reasoning": true
},
"cloudflare/@cf/google/gemma-7b-it-lora": {
"input_cost_per_token": 0.0,
"litellm_provider": "cloudflare",
"max_input_tokens": 3500,
"max_output_tokens": 3500,
"max_tokens": 3500,
"mode": "chat",
"output_cost_per_token": 0.0
},
"cloudflare/@cf/google/gemma-4-26b-a4b-it": {
"input_cost_per_token": 1e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 256000,
"max_output_tokens": 256000,
"max_tokens": 256000,
"mode": "chat",
"output_cost_per_token": 3e-07,
"supports_function_calling": true,
"supports_reasoning": true
},
"cloudflare/@cf/mistralai/mistral-small-3.1-24b-instruct": {
"input_cost_per_token": 3.51e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 5.55e-07,
"supports_function_calling": true
},
"cloudflare/@cf/meta/llama-3.2-11b-vision-instruct": {
"input_cost_per_token": 4.85e-08,
"litellm_provider": "cloudflare",
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 6.76e-07,
"supports_vision": true
},
"cloudflare/@cf/openai/gpt-oss-20b": {
"input_cost_per_token": 2e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 3e-07,
"supports_function_calling": true,
"supports_reasoning": true
},
"cloudflare/@cf/meta/llama-4-scout-17b-16e-instruct": {
"input_cost_per_token": 2.7e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 131000,
"max_output_tokens": 131000,
"max_tokens": 131000,
"mode": "chat",
"output_cost_per_token": 8.5e-07,
"supports_function_calling": true
},
"cloudflare/@cf/qwen/qwq-32b": {
"input_cost_per_token": 6.6e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 24000,
"max_output_tokens": 24000,
"max_tokens": 24000,
"mode": "chat",
"output_cost_per_token": 1e-06,
"supports_reasoning": true
},
"codestral/codestral-2405": {
"input_cost_per_token": 0.0,
"litellm_provider": "codestral",
@ -39908,24 +40170,6 @@
"litellm_provider": "fireworks_ai",
"mode": "chat"
},
"fireworks_ai/accounts/fireworks/models/whisper-v3": {
"max_tokens": 4096,
"max_input_tokens": 4096,
"max_output_tokens": 4096,
"input_cost_per_token": 0.0,
"output_cost_per_token": 0.0,
"litellm_provider": "fireworks_ai",
"mode": "audio_transcription"
},
"fireworks_ai/accounts/fireworks/models/whisper-v3-turbo": {
"max_tokens": 4096,
"max_input_tokens": 4096,
"max_output_tokens": 4096,
"input_cost_per_token": 0.0,
"output_cost_per_token": 0.0,
"litellm_provider": "fireworks_ai",
"mode": "audio_transcription"
},
"fireworks_ai/accounts/fireworks/models/yi-34b": {
"max_tokens": 4096,
"max_input_tokens": 4096,
@ -43061,6 +43305,40 @@
"supports_tool_choice": true,
"supports_vision": false
},
"darkbloom/gemma-4-26b": {
"input_cost_per_token": 3e-08,
"litellm_provider": "darkbloom",
"max_input_tokens": 131072,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 1.65e-07,
"source": "https://www.darkbloom.dev/",
"supported_endpoints": [
"/v1/chat/completions"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_system_messages": true,
"supports_tool_choice": true
},
"darkbloom/gpt-oss-20b": {
"input_cost_per_token": 1.45e-08,
"litellm_provider": "darkbloom",
"max_input_tokens": 131072,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 7e-08,
"source": "https://www.darkbloom.dev/",
"supported_endpoints": [
"/v1/chat/completions"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_system_messages": true,
"supports_tool_choice": true
},
"deepseek/deepseek-v4-pro": {
"cache_creation_input_token_cost": 0.0,
"cache_read_input_token_cost": 3.625e-09,

View file

@ -1835,6 +1835,23 @@
"interactions": true
}
},
"darkbloom": {
"display_name": "Darkbloom (`darkbloom`)",
"url": "https://docs.litellm.ai/docs/providers/darkbloom",
"endpoints": {
"chat_completions": true,
"messages": false,
"responses": false,
"embeddings": false,
"image_generations": false,
"audio_transcriptions": false,
"audio_speech": false,
"moderations": false,
"batches": false,
"rerank": false,
"a2a": false
}
},
"predibase": {
"display_name": "Predibase (`predibase`)",
"url": "https://docs.litellm.ai/docs/providers/predibase",

View file

@ -0,0 +1,95 @@
# Experimental MCP Server Change Guidelines
Read @../../../../CLAUDE.md and @CLAUDE.md before changing this package.
This directory owns the proxy-hosted MCP server implementation. Keep changes
inside the module that owns the behavior, and only reach outside this package
when the public type contract, database schema, dashboard, or cross-proxy route
wiring must change with it.
## File Structure
Respect the current package boundaries:
```text
litellm/proxy/_experimental/mcp_server/
AGENTS.md
CLAUDE.md
server.py # ASGI/MCP route handling, sessions, tool calls [PR7: 7-arm only — move BYOK/OAuth pre-fetch into resolver]
mcp_server_manager.py # upstream server registry, clients, tool routing [PR7: _create_mcp_client swaps resolve_mcp_auth -> resolve_credentials]
auth/
user_api_key_auth_mcp.py # LiteLLM admission auth and MCP request headers
token_exchange.py # OAuth token exchange handling [unchanged; V1TokenExchangeAdapter delegates here]
litellm_auth_handler.py # authenticated-user adapter for MCP sessions
outbound_credentials/ # NEW — typed upstream-credential resolution (resolve_credentials + arms)
__init__.py # public surface: resolve_credentials, the configs, CredError
result.py # Ok | Error union (pure stdlib)
types.py # AuthConfig union, CredError, Subject, ServerSpec
httpx_auth.py # NoOpAuth, StaticHeaderAuth (every mode -> one httpx.Auth)
resolver.py # resolve_credentials(): exhaustive per-mode match + assert_never
seams.py # injected Protocols (one per cache-touching mode)
v1_adapters.py # v1-backed seam bodies; delegate to auth/oauth2/db owners
adapter.py # to_subject / to_server_spec / raise_public (v1 <-> v2 boundary)
discoverable_endpoints.py # MCP OAuth metadata, authorize, token, callback
byok_oauth_endpoints.py # BYOK OAuth UI/API flow
oauth_utils.py # redirect URI and proxy base URL validation
oauth2_token_cache.py # OAuth2 and per-user token resolution/cache [PR7: resolve_mcp_auth removed; cache class stays, V1OAuth2CacheAdapter delegates to async_get_token]
db.py # MCP server, credential, env var, submission DB access [unchanged; V1ByokStore delegates to _get_byok_credential / get_user_credential]
toolset_db.py # MCP toolset DB access
rest_endpoints.py # proxy REST facade for listing/calling MCP tools [PR7: 7-arm only — pass identity + inbound token down instead of mcp_auth_header]
openapi_to_mcp_generator.py# OpenAPI spec to MCP tool generation
sampling_handler.py # MCP sampling to LiteLLM completion flow
elicitation_handler.py # MCP elicitation relay flow
semantic_tool_filter.py # semantic filtering of available MCP tools
guardrail_translation/
handler.py # MCP guardrail result translation
sse_transport.py # SSE transport implementation
mcp_context.py # contextvars for MCP request/session metadata
mcp_debug.py # debug helpers
tool_registry.py # in-memory MCP tool registry helpers
cost_calculator.py # MCP tool cost calculation
ui_session_utils.py # dashboard session auth context helpers
utils.py # shared primitives used by several modules
```
Do not add broad catch-all modules. Prefer the existing owner above, and add a
new file only for a distinct capability that would otherwise make an existing
module materially harder to understand.
## Implementation Rules
- Preserve the boundary between LiteLLM admission auth and upstream MCP auth.
Admission belongs in `auth/user_api_key_auth_mcp.py`; upstream token exchange,
delegated auth, per-user OAuth, BYOK, and raw header forwarding belong in the
dedicated OAuth/header modules.
- Treat `none`, bearer/API key, OAuth, OAuth token exchange, delegated upstream
auth, SSE, streamable HTTP, and stdio as separate flows. Do not collapse them
behind a single generic branch unless tests prove every mode still behaves
correctly.
- Be especially careful with `available_on_public_internet: false` combined with
`delegate_auth_to_upstream: true`. The local `CLAUDE.md` explains the anonymous
upstream PKCE path that must remain intentional.
- Keep database-backed fields in sync across migrations, typed models under
`litellm/types/mcp.py` or `litellm/types/mcp_server/`, config loading, this
package, and dashboard state when the field is user-visible.
- Use the official MCP SDK types and established LiteLLM Pydantic models where
they exist. Avoid untyped protocol dictionaries at package boundaries.
- Keep security-sensitive logic easy to audit. Header forwarding, IP filtering,
public internet checks, token storage, env var interpolation, and credential
encryption need focused tests for both allowed and rejected paths.
- Avoid adding comments to new code unless they explain non-obvious security or
protocol behavior. Prefer clear names and small functions.
## Tests
Mirror this package under `tests/test_litellm/proxy/_experimental/mcp_server/`.
For regressions, extend the existing mapped test file instead of creating a new
one. Use subdirectories that match the implementation path, such as
`auth/test_token_exchange.py` for `auth/token_exchange.py` and
`guardrail_translation/test_mcp_guardrail_handler.py` for
`guardrail_translation/handler.py`.
Use `tests/mcp_tests/` only when extending an existing broader MCP integration
scenario that already lives there. Route, auth, tool listing, tool execution,
OAuth, sampling, elicitation, DB, and dashboard-session changes should have
focused coverage in the mirrored `tests/test_litellm/...` path first.

View file

@ -12,6 +12,7 @@ from litellm.proxy._types import (
LiteLLM_TeamTable,
ProxyException,
SpecialHeaders,
SpecialMCPServerNames,
UserAPIKeyAuth,
)
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
@ -642,6 +643,15 @@ class MCPRequestHandler:
user_api_key_auth
)
)
# The key explicitly opted out of every MCP server. This overrides
# team inheritance and additive grants (mirrors no-default-models).
if (
SpecialMCPServerNames.no_mcp_servers.value
in allowed_mcp_servers_for_key
):
return []
allowed_mcp_servers_for_team = (
await MCPRequestHandler._get_allowed_mcp_servers_for_team(
user_api_key_auth
@ -1058,6 +1068,13 @@ class MCPRequestHandler:
if key_object_permission is None:
return []
# Sentinel opt-out: surface it unexpanded so the caller can short-circuit
# to zero servers instead of inheriting the team.
if SpecialMCPServerNames.no_mcp_servers.value in (
key_object_permission.mcp_servers or []
):
return [SpecialMCPServerNames.no_mcp_servers.value]
# Permission entries may be server_ids OR names/aliases — expand to ids.
direct_mcp_servers = global_mcp_server_manager.expand_permission_list(
key_object_permission.mcp_servers or []

View file

@ -80,6 +80,7 @@ from litellm.proxy._types import (
MCPEnvVar,
MCPTransport,
MCPTransportType,
SpecialMCPServerNames,
UserAPIKeyAuth,
)
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
@ -1349,6 +1350,17 @@ class MCPServerManager:
allow_all_server_ids = self.get_allow_all_keys_server_ids()
try:
# The key explicitly opted out of every MCP server. Return zero before
# layering on allow_all_keys servers so the opt-out is absolute.
key_object_permission = (
user_api_key_auth.object_permission if user_api_key_auth else None
)
if key_object_permission is not None and (
SpecialMCPServerNames.no_mcp_servers.value
in (key_object_permission.mcp_servers or [])
):
return []
# Check if object_permission.mcp_servers is explicitly set
has_explicit_object_permission = False
if user_api_key_auth and user_api_key_auth.object_permission:

View file

@ -0,0 +1,73 @@
"""Typed upstream-credential resolution for MCP servers.
This subpackage houses the typed credential vocabulary and the ``resolve_credentials``
dispatch. A server declares one per-mode config from the ``AuthConfig`` discriminated union;
``UpstreamCredentialProvider.resolve_credentials`` selects one arm and returns an ``httpx.Auth``
or a typed ``CredError``. Failures are modeled as values via :mod:`.result` (``Result[T,
CredError]``) rather than raised, so every seam is total. Nothing here is wired onto a live
request path yet.
"""
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
NoOpAuth,
StaticHeaderAuth,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.resolver import (
UpstreamCredentialProvider,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import (
Error,
Ok,
Result,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
Ambient,
ApiKeyConfig,
ApiKeySource,
AssumeRole,
AuthConfig,
AuthorizationCodeConfig,
AuthSpecKind,
AwsCredentialSource,
AwsSigV4Config,
Byok,
ClientCredentialsConfig,
CredError,
NoneConfig,
PassthroughConfig,
ServerSpec,
SharedKey,
StaticKeys,
Subject,
TokenExchangeConfig,
parse_auth_spec_kind,
)
__all__ = [
"Ok",
"Error",
"Result",
"NoOpAuth",
"StaticHeaderAuth",
"UpstreamCredentialProvider",
"AuthSpecKind",
"CredError",
"Subject",
"ServerSpec",
"AuthConfig",
"parse_auth_spec_kind",
"AuthorizationCodeConfig",
"ClientCredentialsConfig",
"TokenExchangeConfig",
"ApiKeyConfig",
"ApiKeySource",
"SharedKey",
"Byok",
"PassthroughConfig",
"NoneConfig",
"AwsSigV4Config",
"AwsCredentialSource",
"StaticKeys",
"AssumeRole",
"Ambient",
]

View file

@ -0,0 +1,45 @@
"""Concrete `httpx.Auth` objects the resolver returns for the self-contained modes.
These are the egress credential as the SDK consumes it: an `httpx.Auth` attached to the
upstream `AsyncClient`. The OAuth-flow modes (`authorization_code`, `client_credentials`,
`token_exchange`) return SDK-provided auth objects instead and land later.
`auth_flow` mutating the outbound request is the `httpx.Auth` contract, not a house-style
violation: the request is httpx's object, and these carry no state of their own.
"""
from __future__ import annotations
from collections.abc import Generator
import httpx
from pydantic import SecretStr
class NoOpAuth(httpx.Auth):
"""Attaches nothing — the `none` mode (and the seam-level default)."""
def auth_flow(
self, request: httpx.Request
) -> Generator[httpx.Request, httpx.Response, None]:
yield request
class StaticHeaderAuth(httpx.Auth):
"""Sets one fixed header on every request — the `api_key` family and `passthrough`.
The header value is a live credential (a bearer token, an API key, a forwarded user
token), so it is held as a `SecretStr` and unwrapped only when written onto the request.
That keeps it masked in reprs, `vars()`, tracebacks, and structured logs, matching the
`SecretStr` discipline the config models use.
"""
def __init__(self, header_value: str, header_name: str = "Authorization") -> None:
self.header_name = header_name
self._header_value = SecretStr(header_value)
def auth_flow(
self, request: httpx.Request
) -> Generator[httpx.Request, httpx.Response, None]:
request.headers[self.header_name] = self._header_value.get_secret_value()
yield request

View file

@ -0,0 +1,70 @@
"""The one credential resolver: dispatch on the declared mode, fail closed.
`resolve_credentials` selects exactly one arm off the server's typed `config` and either
produces an `httpx.Auth` or returns a typed `CredError`. The `match` is over the `AuthConfig`
variant, so each arm receives its own fully-typed config with no field-presence inference and
no precedence cascade. It is wildcard-free with an `assert_never` tail, so adding a mode without
an arm fails the type gate (basedpyright `reportMatchNotExhaustive`); a bypassed gate fails loudly
at runtime instead of returning `None`.
This skeleton ships every arm as a `not_implemented` stub. Each mode's real body, with its
injected seam, lands in its own follow-up PR; until then the arm returns a typed error rather
than silently producing no credential. Pure v2: no imports from v1.
"""
from __future__ import annotations
import httpx
from typing_extensions import assert_never
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import (
Error,
Result,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
ApiKeyConfig,
AuthorizationCodeConfig,
AuthSpecKind,
AwsSigV4Config,
ClientCredentialsConfig,
CredError,
NoneConfig,
PassthroughConfig,
ServerSpec,
Subject,
TokenExchangeConfig,
)
class UpstreamCredentialProvider:
"""Produces the one `httpx.Auth` for a `(subject, upstream)` pair, per declared mode.
Collaborators (the per-mode credential stores and token fetchers) are injected as each arm
is built; the skeleton needs none, since every arm is a stub.
"""
async def resolve_credentials(
self, subject: Subject, server: ServerSpec
) -> Result[httpx.Auth, CredError]:
match server.config:
case NoneConfig():
return _not_implemented(AuthSpecKind.none)
case ApiKeyConfig():
return _not_implemented(AuthSpecKind.api_key)
case PassthroughConfig():
return _not_implemented(AuthSpecKind.passthrough)
case ClientCredentialsConfig():
return _not_implemented(AuthSpecKind.client_credentials)
case TokenExchangeConfig():
return _not_implemented(AuthSpecKind.token_exchange)
case AuthorizationCodeConfig():
return _not_implemented(AuthSpecKind.authorization_code)
case AwsSigV4Config():
return _not_implemented(AuthSpecKind.aws_sigv4)
assert_never(server.config)
def _not_implemented(kind: AuthSpecKind) -> Result[httpx.Auth, CredError]:
return Error(
CredError.of_not_implemented(f"{kind.value}: resolver arm not implemented yet")
)

View file

@ -0,0 +1,54 @@
"""A tagged-union ``Result`` the type checker can actually narrow.
``Ok`` and ``Error`` are separate frozen classes joined by a ``Union`` alias, so
reaching for ``result.ok`` before eliminating the ``Error`` arm (via ``isinstance``
or a ``match`` pattern) is a type error rather than a runtime ``AttributeError``. A
single class carrying both payload fields would make that unguarded access invisible
to the type checker.
Both variants are covariant and frozen; the absent side defaults to ``Never`` so a
bare ``Ok(value)`` or ``Error(err)`` infers fully and is assignable to any ``Result``
whose matching side fits.
``is_ok`` / ``is_error`` are runtime predicates that also narrow via their ``Literal``
returns; inside strictly typed code, discriminate with ``match`` or ``isinstance``.
This is the shared ``Result`` shape for the ``outbound_credentials`` resolver: every
seam returns ``Result[T, CredError]`` instead of raising, so each failure is a value
the caller must handle rather than an exception that can slip past the type checker.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Generic, Literal, TypeAlias
from typing_extensions import Never, TypeVar
_TOk_co = TypeVar("_TOk_co", covariant=True, default=Never)
_TError_co = TypeVar("_TError_co", covariant=True, default=Never)
@dataclass(frozen=True)
class Ok(Generic[_TOk_co, _TError_co]):
ok: _TOk_co
def is_ok(self) -> Literal[True]:
return True
def is_error(self) -> Literal[False]:
return False
@dataclass(frozen=True)
class Error(Generic[_TOk_co, _TError_co]):
error: _TError_co
def is_ok(self) -> Literal[False]:
return False
def is_error(self) -> Literal[True]:
return True
Result: TypeAlias = Ok[_TOk_co, _TError_co] | Error[_TOk_co, _TError_co]

View file

@ -0,0 +1,334 @@
"""The upstream-credential vocabulary — the typed seam the resolver dispatches on.
This module ships the data types only; the resolver lands in a later PR. It is the contract
the credential build implements and the spec tests assert against.
Design invariants encoded here:
- **Mode is the single source of truth.** A server declares exactly one per-mode `config`
(the `AuthConfig` discriminated union); `auth_spec_kind` is *derived* from it, never a
second field that can drift. The resolver dispatches on the config variant, one arm per
mode. No field-presence inference, no precedence cascade.
- **Illegal states unrepresentable.** Each mode's config is its own frozen model holding
only that mode's fields — an `aws_sigv4` server cannot hold OAuth fields, and a config
missing a required field is rejected at construction, not at call time.
- **Fail-closed at the boundary.** A raw mode string can only enter through
`parse_auth_spec_kind()`, which returns a typed `CredError`.
- **Errors as values.** Every seam returns `Result[_, CredError]`; only edge adapters raise.
- **No v1 imports.** This vocabulary stays free of `MCPServer` and the rest of v1; the
v1 -> v2 adapter maps onto these types in a later PR.
Sum types are Expression `@tagged_union`s discriminated on a `Literal` `tag`, matched via
`self.tag` with an `assert_never` tail; `Result` is this package's vendored `Ok | Error`
union (see `result.py`), not `expression.Result`.
"""
from __future__ import annotations
from enum import Enum
from typing import Annotated, Literal
from expression import case, tag, tagged_union
from pydantic import BaseModel, ConfigDict, Field, SecretStr
from typing_extensions import assert_never
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import (
Error,
Ok,
Result,
)
class AuthSpecKind(str, Enum):
"""The server's statically-declared upstream-auth mode — derived from its `config`.
Covers v1's full `MCPAuth` surface, not only OAuth grants: the three grant modes, the
collapsed static-header family, client passthrough, no-auth, and AWS request signing.
BYOK is *not* a member: it is the `api_key` mode seeded per-user, a source selector
inside that arm. The static-header schemes v1 splits into separate `MCPAuth` values
(`bearer_token`/`api_key`/`basic`/`token`/`authorization`) collapse into `api_key`; the
scheme is a parameter the arm carries, not its own mode.
"""
authorization_code = "authorization_code" # per-user 3LO; gateway-stored token
client_credentials = "client_credentials" # gateway service account (M2M)
token_exchange = "token_exchange" # RFC 8693: token endpoint + subject_token (OBO)
api_key = "api_key" # static header, any scheme (BYOK = per-user-seeded source)
passthrough = "passthrough" # client forwards an upstream-audience token
none = "none" # no upstream credential; resolve yields a no-op auth, never an error
aws_sigv4 = "aws_sigv4" # AWS SigV4 per-request signing (e.g. Bedrock AgentCore)
@tagged_union(frozen=True)
class CredError:
"""Why a credential could not be produced. Fail-closed: an arm yields this or an `httpx.Auth`.
Discriminated on the `Literal` `tag`; consumers `match self.tag` (see `summary`) so the
type checker can prove exhaustiveness. Construct via the `of_*` factories.
"""
tag: Literal[
"unauthorized",
"misconfigured",
"upstream_unavailable",
"unsupported_mode",
"precondition_required",
"not_implemented",
] = tag()
unauthorized: str = (
case()
) # no usable credential for this (subject, server) -> 401 challenge
misconfigured: str = (
case()
) # the declared mode is missing required config -> 5xx (operator)
upstream_unavailable: str = (
case()
) # the IdP / token endpoint could not be reached -> 503
unsupported_mode: str = (
case()
) # a raw mode string did not parse into AuthSpecKind (boundary)
precondition_required: str = (
case()
) # a required per-user value (e.g. an env var) has not been provided -> 412
not_implemented: str = (
case()
) # the declared mode's resolver arm is not built yet -> 501 (not operator error)
@staticmethod
def of_unauthorized(detail: str) -> CredError:
return CredError(unauthorized=detail)
@staticmethod
def of_misconfigured(detail: str) -> CredError:
return CredError(misconfigured=detail)
@staticmethod
def of_upstream_unavailable(detail: str) -> CredError:
return CredError(upstream_unavailable=detail)
@staticmethod
def of_unsupported_mode(detail: str) -> CredError:
return CredError(unsupported_mode=detail)
@staticmethod
def of_precondition_required(detail: str) -> CredError:
return CredError(precondition_required=detail)
@staticmethod
def of_not_implemented(detail: str) -> CredError:
return CredError(not_implemented=detail)
@property
def summary(self) -> str:
# Exhaustiveness: every Literal tag has an arm; the trailing assert_never typechecks
# only while that stays true (a `case _` would defeat reportMatchNotExhaustive).
match self.tag:
case "unauthorized":
return f"unauthorized: {self.unauthorized}"
case "misconfigured":
return f"misconfigured: {self.misconfigured}"
case "upstream_unavailable":
return f"upstream unavailable: {self.upstream_unavailable}"
case "unsupported_mode":
return self.unsupported_mode
case "precondition_required":
return f"precondition required: {self.precondition_required}"
case "not_implemented":
return f"not implemented: {self.not_implemented}"
assert_never(self.tag)
class AuthorizationCodeConfig(BaseModel):
"""Per-user 3LO; the gateway is the OAuth client and stores the user's token.
Endpoints are discovered (RFC 9728 -> RFC 8414) and the client is registered via DCR
(RFC 7591), so the common case carries none of the fields below; they are optional manual
overrides for IdPs without discovery / DCR. The per-user token is read from the token store
at resolve time, not held here.
"""
model_config = ConfigDict(frozen=True)
kind: Literal[AuthSpecKind.authorization_code] = AuthSpecKind.authorization_code
scopes: tuple[str, ...] = ()
client_id: str | None = None
client_secret: SecretStr | None = None
authorization_url: str | None = None
token_url: str | None = None
class ClientCredentialsConfig(BaseModel):
"""M2M service account; one upstream identity for every user.
Fields are optional so the config can be built incomplete: a value may be supplied at
runtime (`token_url` via RFC 8414 discovery, `client_id`/`secret` via DCR), and the
resolver arm raises `CredError.misconfigured` when a needed field is still absent.
"""
model_config = ConfigDict(frozen=True)
kind: Literal[AuthSpecKind.client_credentials] = AuthSpecKind.client_credentials
client_id: str | None = None
client_secret: SecretStr | None = None
token_url: str | None = None
scopes: tuple[str, ...] = ()
class TokenExchangeConfig(BaseModel):
"""RFC 8693 OBO; swap the caller's live subject_token for a token bound to the upstream's
audience (`server.resource`, RFC 8707). The gateway authenticates to the exchange endpoint
as an OAuth client (`client_id`/`client_secret`); the inbound token is sent only to that
endpoint, never to the upstream.
"""
model_config = ConfigDict(frozen=True)
kind: Literal[AuthSpecKind.token_exchange] = AuthSpecKind.token_exchange
subject_token_type: str = "urn:ietf:params:oauth:token-type:access_token"
token_exchange_endpoint: str | None = None
client_id: str | None = None
client_secret: SecretStr | None = None
scopes: tuple[str, ...] = ()
class SharedKey(BaseModel):
"""A fixed key configured on the server, identical for every caller."""
model_config = ConfigDict(frozen=True)
source: Literal["shared"] = "shared"
value: SecretStr
class Byok(BaseModel):
"""A key the user brings via the entry flow, stored per-user and pulled from the credential
store at resolve time. Missing means the user must provide it, a 401 + WWW-Authenticate
challenge."""
model_config = ConfigDict(frozen=True)
source: Literal["byok"] = "byok"
ApiKeySource = Annotated[SharedKey | Byok, Field(discriminator="source")]
class ApiKeyConfig(BaseModel):
"""A fixed credential injected as a header. The value is shared (in config) or seeded
per-user (pulled from the store); `header_name` and `value_prefix` say where and how it is
written, modeled like OpenAPI's apiKey scheme so any upstream convention is expressible
(Authorization + Bearer, a raw value on X-API-Key, Ocp-Apim-Subscription-Key, etc.).
"""
model_config = ConfigDict(frozen=True)
kind: Literal[AuthSpecKind.api_key] = AuthSpecKind.api_key
header_name: str = "Authorization"
value_prefix: str = "Bearer"
key_source: ApiKeySource
def header(self, value: str) -> tuple[str, str]:
formatted = f"{self.value_prefix} {value}" if self.value_prefix else value
return self.header_name, formatted
class PassthroughConfig(BaseModel):
"""Client-driven upstream OAuth; the gateway forwards the client's upstream token."""
model_config = ConfigDict(frozen=True)
kind: Literal[AuthSpecKind.passthrough] = AuthSpecKind.passthrough
class NoneConfig(BaseModel):
"""No upstream credential; the request is sent unauthenticated."""
model_config = ConfigDict(frozen=True)
kind: Literal[AuthSpecKind.none] = AuthSpecKind.none
class StaticKeys(BaseModel):
"""Long-lived AWS access keys configured on the server."""
model_config = ConfigDict(frozen=True)
source: Literal["static_keys"] = "static_keys"
access_key_id: str
secret_access_key: SecretStr
session_token: SecretStr | None = None
class AssumeRole(BaseModel):
"""An IAM role the gateway assumes via STS for short-lived, auto-refreshed credentials."""
model_config = ConfigDict(frozen=True)
source: Literal["assume_role"] = "assume_role"
role_arn: str
session_name: str | None = None
external_id: str | None = None
class Ambient(BaseModel):
"""The environment's default AWS credential chain (instance profile, IRSA, env vars)."""
model_config = ConfigDict(frozen=True)
source: Literal["ambient"] = "ambient"
AwsCredentialSource = Annotated[
StaticKeys | AssumeRole | Ambient, Field(discriminator="source")
]
class AwsSigV4Config(BaseModel):
"""AWS SigV4 per-request signing for an AWS-hosted upstream (e.g. Bedrock AgentCore). The
gateway signs with its own AWS identity, never the caller's; `credentials` selects how that
identity is obtained, defaulting to the ambient credential chain."""
model_config = ConfigDict(frozen=True)
kind: Literal[AuthSpecKind.aws_sigv4] = AuthSpecKind.aws_sigv4
region: str
service: str = "bedrock-agentcore"
credentials: AwsCredentialSource = Ambient()
AuthConfig = Annotated[
AuthorizationCodeConfig
| ClientCredentialsConfig
| TokenExchangeConfig
| ApiKeyConfig
| PassthroughConfig
| NoneConfig
| AwsSigV4Config,
Field(discriminator="kind"),
]
class Subject(BaseModel):
"""The validated inbound principal. NOT the v1 request object and NOT the LiteLLM key."""
model_config = ConfigDict(frozen=True)
tenant_id: str
subject_id: str
# Opaque, already-validated inbound identity. Only `token_exchange` / `passthrough` read it.
inbound_token: SecretStr | None = None
class ServerSpec(BaseModel):
"""The declared upstream. A v2-native type; the v1 -> v2 adapter maps onto this."""
model_config = ConfigDict(frozen=True)
server_id: str
resource: str # RFC 8707 audience URI this upstream's tokens are bound to
config: AuthConfig
@property
def auth_spec_kind(self) -> AuthSpecKind:
return self.config.kind
def parse_auth_spec_kind(raw: str) -> Result[AuthSpecKind, CredError]:
"""Boundary parser — the *only* place an unknown mode is handled, and it fails closed.
Inside the core the mode is always a valid `AuthSpecKind`, so the resolver never needs a
wildcard arm and basedpyright can prove its `match` exhaustive.
"""
try:
return Ok(AuthSpecKind(raw))
except ValueError:
return Error(CredError.of_unsupported_mode(f"unknown auth_spec_kind: {raw!r}"))

View file

@ -63,7 +63,11 @@ from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy._types import (
ProxyException,
SpecialMCPServerNames,
UserAPIKeyAuth,
)
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
from litellm.proxy.litellm_pre_call_utils import (
LiteLLMProxyRequestSetup,
@ -229,6 +233,28 @@ def _jsonrpc_text_has_top_level_method(text: str) -> bool:
return False
def _proxy_exception_to_http_exception(exc: ProxyException) -> HTTPException:
"""Map a ``ProxyException`` to an ``HTTPException`` that preserves its real
status code and headers.
``user_api_key_auth`` raises ``ProxyException`` (not ``HTTPException``) on
auth failures. The MCP ASGI handlers re-raise ``HTTPException`` to keep the
status and any ``WWW-Authenticate`` challenge, but a ``ProxyException`` would
otherwise fall through to their generic handler and be flattened to a 500 —
dropping the 401 + challenge an OAuth client needs to re-authenticate, so the
tool call surfaces as a cancelled/terminated session instead.
"""
try:
status_code = int(exc.code)
except (TypeError, ValueError):
status_code = 500
return HTTPException(
status_code=status_code,
detail=exc.message,
headers=exc.headers or None,
)
if MCP_AVAILABLE:
from mcp.server import Server
from mcp.server.lowlevel.server import NotificationOptions
@ -3380,6 +3406,19 @@ if MCP_AVAILABLE:
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
# A key scoped to no MCP servers opts out of every MCP path. Enforce it
# here too, since toolset scoping replaces mcp_servers and would otherwise
# drop the sentinel. Checked before the admin branch, mirroring
# get_allowed_mcp_servers.
original_op = user_api_key_auth.object_permission
if original_op is not None and SpecialMCPServerNames.no_mcp_servers.value in (
original_op.mcp_servers or []
):
raise HTTPException(
status_code=403,
detail="API key is scoped to no MCP servers; toolset access is denied.",
)
# Access control: non-admin keys must have this toolset in their grant list.
# Use _user_has_admin_view so that PROXY_ADMIN_VIEW_ONLY is also treated as admin.
is_admin = _user_has_admin_view(user_api_key_auth)
@ -4034,6 +4073,12 @@ if MCP_AVAILABLE:
except HTTPException:
# Re-raise HTTP exceptions to preserve status codes and details
raise
except ProxyException as e:
# Auth failures from user_api_key_auth arrive as ProxyException, not
# HTTPException. Preserve the real status (e.g. 401 + WWW-Authenticate)
# so OAuth clients can re-authenticate instead of receiving a generic
# 500 that surfaces as a cancelled tool call.
raise _proxy_exception_to_http_exception(e)
except Exception as e:
verbose_logger.exception(f"Error handling MCP request: {e}")
# Try to send a graceful error response for non-HTTP exceptions
@ -4151,6 +4196,12 @@ if MCP_AVAILABLE:
# Re-raise HTTP exceptions to preserve status codes and details
# (e.g. 401 + WWW-Authenticate challenges from OAuth pass-through).
raise
except ProxyException as e:
# Auth failures from user_api_key_auth arrive as ProxyException, not
# HTTPException. Preserve the real status (e.g. 401 + WWW-Authenticate)
# so OAuth clients can re-authenticate instead of receiving a generic
# 500 that surfaces as a cancelled tool call.
raise _proxy_exception_to_http_exception(e)
except Exception as e:
verbose_logger.exception(f"Error handling MCP request: {e}")
# Try to send a graceful error response for non-HTTP exceptions

View file

@ -461,6 +461,7 @@ class LiteLLMRoutes(enum.Enum):
"/mcp/tools/call",
"/mcp-rest/tools/list",
"/mcp-rest/tools/call",
"/v1/mcp/tools",
]
# MCP server CRUD routes — control-plane. Gated by DISABLE_ADMIN_ENDPOINTS.
@ -2977,6 +2978,10 @@ class SpecialModelNames(enum.Enum):
no_default_models = "no-default-models"
class SpecialMCPServerNames(enum.Enum):
no_mcp_servers = "no-mcp-servers"
class SpecialProxyStrings(enum.Enum):
default_user_id = "default_user_id" # global proxy admin
@ -3353,7 +3358,9 @@ class ProxyException(Exception):
class CommonProxyErrors(str, enum.Enum):
db_not_connected_error = (
"DB not connected. See https://docs.litellm.ai/docs/proxy/virtual_keys"
"DB not connected. This endpoint needs a database; set DATABASE_URL to a "
"PostgreSQL connection string (postgresql://...) to enable it. "
"See https://docs.litellm.ai/docs/proxy/virtual_keys"
)
no_llm_router = "No models configured on proxy"
not_allowed_access = "Admin-only endpoint. Not allowed to access this."

View file

@ -699,6 +699,11 @@ async def common_checks(
if valid_token is not None:
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
LiteLLMProxyRequestSetup.pre_seed_litellm_metadata_for_route(
request_data=request_body,
route=route,
)
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
request_data=request_body,
user_api_key_dict=valid_token,
@ -2949,6 +2954,26 @@ async def _get_agent_ids_from_access_groups(
)
def _resolve_all_team_model_sentinel_for_auth_check(
models: List[str],
llm_router: Optional[Router],
team_id: Optional[str],
) -> List[str]:
if (
SpecialModelNames.all_team_models.value not in models
or team_id is None
or llm_router is None
):
return models
proxy_models = llm_router.get_model_names()
non_sentinel_models = [
model for model in models if model != SpecialModelNames.all_team_models.value
]
if not proxy_models:
return non_sentinel_models or models
return list(dict.fromkeys(non_sentinel_models + proxy_models))
def _check_model_access_helper(
model: str,
llm_router: Optional[Router],
@ -2966,6 +2991,12 @@ def _check_model_access_helper(
model_name=model, team_id=team_id
)
models = _resolve_all_team_model_sentinel_for_auth_check(
models=models,
llm_router=llm_router,
team_id=team_id,
)
if (
len(access_groups) > 0 and llm_router is not None
): # check if token contains any model access groups
@ -3658,9 +3689,18 @@ async def _virtual_key_max_budget_check(
# so a NaN max_budget would silently disable enforcement. Treat a
# non-finite max_budget as "no configured limit" rather than as a bypass.
if math.isfinite(valid_token.max_budget) and spend >= valid_token.max_budget:
# name the key in the error so operators don't have to reverse-map
# spend back to a key; key_name is the masked form (last 4 chars)
key_label = valid_token.key_alias or "key"
key_descriptor = (
f"{key_label} ({valid_token.key_name})"
if valid_token.key_name
else key_label
)
raise litellm.BudgetExceededError(
current_cost=spend,
max_budget=valid_token.max_budget,
message=f"Budget has been exceeded! Key={key_descriptor} Current cost: {spend}, Max budget: {valid_token.max_budget}",
)

View file

@ -285,6 +285,8 @@ _BANNED_REQUEST_BODY_PARAMS: Tuple[str, ...] = (
"s3_endpoint_url",
"sagemaker_base_url",
"deployment_url",
# SDK-only field; also rejected outright in is_request_body_safe.
"model_list",
# Observability credentials, hosts, and project identifiers: derived
# from the canonical ``_supported_callback_params`` allowlist so new
# integrations are covered automatically. Sorted for stable iteration
@ -365,6 +367,10 @@ def is_request_body_safe(
``litellm_embedding_config.api_base`` (VERIA-6) without exposing a
recursion-depth DoS surface.
"""
if "model_list" in request_body:
raise ValueError(
"Rejected Request: model_list is not allowed in the request body."
)
_check_banned_params(request_body, general_settings, llm_router, model)
for nested_key in _NESTED_CONFIG_KEYS:
nested = _coerce_metadata_to_dict(request_body.get(nested_key))

View file

@ -122,9 +122,16 @@ def get_key_models(
SpecialModelNames.all_team_models.value in all_models
and user_api_key_dict.team_id is not None
):
all_models = list(
user_api_key_dict.team_models
) # copy to avoid mutating cached objects
all_models = list(user_api_key_dict.team_models)
if SpecialModelNames.all_team_models.value in all_models:
all_models = [
model
for model in all_models
if model != SpecialModelNames.all_team_models.value
]
all_models.extend(proxy_model_list)
if include_model_access_groups:
all_models.extend(model_access_groups.keys())
if SpecialModelNames.all_proxy_models.value in all_models:
all_models = list(proxy_model_list) # copy to avoid mutating caller's list
if include_model_access_groups:
@ -160,6 +167,12 @@ def get_team_models(
all_models_set.update(team_models)
if SpecialModelNames.all_team_models.value in all_models_set:
all_models_set.update(team_models)
# GH#30619: expand all-team-models sentinel
# to the actual proxy model list
all_models_set.discard(SpecialModelNames.all_team_models.value)
all_models_set.update(proxy_model_list)
if include_model_access_groups:
all_models_set.update(model_access_groups.keys())
if SpecialModelNames.all_proxy_models.value in all_models_set:
all_models_set.update(proxy_model_list)
if include_model_access_groups:

View file

@ -2396,6 +2396,17 @@ async def _run_centralized_common_checks(
llm_router=llm_router,
)
# Pin the metadata variable name (litellm_metadata vs metadata) before
# any tag merge runs. Without this, header tags from
# apply_client_tag_policy_pre_auth would land in `metadata` while the
# later seed in common_checks pushes key tags and the
# _tag_max_budget_check read into `litellm_metadata`, hiding header
# tags from per-tag budget enforcement on LITELLM_METADATA_ROUTES.
LiteLLMProxyRequestSetup.pre_seed_litellm_metadata_for_route(
request_data=request_data,
route=route,
)
# Merge x-litellm-tags into request_data BEFORE common_checks runs.
# _tag_max_budget_check inside common_checks only inspects request_data;
# without this pre-merge, header-supplied tags bypass tag-budget

View file

@ -1037,6 +1037,8 @@ class ProxyBaseLLMRequestProcessing:
version=version,
proxy_config=proxy_config,
)
if not general_settings.get("expose_fallback_errors_to_caller"):
self.data.pop("include_fallback_errors", None)
if route_type in {"aresponses", "_aresponses_websocket"}:
await _authorize_response_file_search_vector_stores(
data=self.data,

View file

@ -32,7 +32,7 @@ password when their ``*_READ_REPLICA`` counterpart is unset.
import os
import urllib.parse
from typing import Optional, cast
from typing import Final, cast
from pydantic import AliasChoices, Field
from pydantic_settings import BaseSettings, SettingsConfigDict
@ -44,6 +44,41 @@ from litellm.proxy.auth import rds_iam_token
_IAM_ENV_KEY = "IAM_TOKEN_DB_AUTH"
_DEFAULT_PG_PORT = "5432"
# schema.prisma pins `provider = "postgresql"`, so these are the only schemes
# Prisma can actually connect with.
SUPPORTED_DB_SCHEMES: Final[frozenset[str]] = frozenset({"postgresql", "postgres"})
_MISSING_SCHEME = "<missing scheme>"
def unsupported_db_scheme(database_url: str) -> str | None:
"""Return the connection URL scheme when it is not PostgreSQL, else None.
A `sqlite://` / `mysql://` URL can never connect against the
postgresql-only datasource, but the resulting Prisma failure is opaque and
version-dependent (a confusing migration error, or a startup that never
binds). Callers use this to reject the URL up front with an actionable
error instead.
A schemeless value (e.g. a malformed DSN like ``user:pass@host/db``) yields
the ``_MISSING_SCHEME`` placeholder rather than the raw URL, so callers that
log the return value never echo embedded credentials.
"""
scheme = urllib.parse.urlsplit(database_url).scheme.lower()
if scheme in SUPPORTED_DB_SCHEMES:
return None
return scheme or _MISSING_SCHEME
def unsupported_db_scheme_message(env_var: str, scheme: str) -> str:
"""Operator-facing message naming the offending env var and scheme."""
return (
f"{env_var} uses unsupported scheme '{scheme}'. LiteLLM's database "
"features (virtual keys, store_model_in_db, spend tracking) require "
"PostgreSQL; use a 'postgresql://' connection string. SQLite and other "
"engines are not supported. "
"See https://docs.litellm.ai/docs/proxy/virtual_keys"
)
class DatabaseURLSettings(BaseSettings):
"""Discrete ``DATABASE_*`` env vars, loaded once at process start.
@ -58,46 +93,47 @@ class DatabaseURLSettings(BaseSettings):
iam_token_db_auth: bool = Field(default=False, validation_alias=_IAM_ENV_KEY)
# Writer
database_url: Optional[str] = Field(default=None, validation_alias="DATABASE_URL")
database_host: Optional[str] = Field(default=None, validation_alias="DATABASE_HOST")
database_url: str | None = Field(default=None, validation_alias="DATABASE_URL")
direct_url: str | None = Field(default=None, validation_alias="DIRECT_URL")
database_host: str | None = Field(default=None, validation_alias="DATABASE_HOST")
database_port: str = Field(
default=_DEFAULT_PG_PORT, validation_alias="DATABASE_PORT"
)
database_user: Optional[str] = Field(
database_user: str | None = Field(
default=None,
validation_alias=AliasChoices("DATABASE_USER", "DATABASE_USERNAME"),
)
database_name: Optional[str] = Field(default=None, validation_alias="DATABASE_NAME")
database_schema: Optional[str] = Field(
database_name: str | None = Field(default=None, validation_alias="DATABASE_NAME")
database_schema: str | None = Field(
default=None, validation_alias="DATABASE_SCHEMA"
)
database_password: Optional[str] = Field(
database_password: str | None = Field(
default=None, validation_alias="DATABASE_PASSWORD"
)
# Read replica
database_url_read_replica: Optional[str] = Field(
database_url_read_replica: str | None = Field(
default=None, validation_alias="DATABASE_URL_READ_REPLICA"
)
database_host_read_replica: Optional[str] = Field(
database_host_read_replica: str | None = Field(
default=None, validation_alias="DATABASE_HOST_READ_REPLICA"
)
database_port_read_replica: Optional[str] = Field(
database_port_read_replica: str | None = Field(
default=None, validation_alias="DATABASE_PORT_READ_REPLICA"
)
database_user_read_replica: Optional[str] = Field(
database_user_read_replica: str | None = Field(
default=None,
validation_alias=AliasChoices(
"DATABASE_USER_READ_REPLICA", "DATABASE_USERNAME_READ_REPLICA"
),
)
database_name_read_replica: Optional[str] = Field(
database_name_read_replica: str | None = Field(
default=None, validation_alias="DATABASE_NAME_READ_REPLICA"
)
database_schema_read_replica: Optional[str] = Field(
database_schema_read_replica: str | None = Field(
default=None, validation_alias="DATABASE_SCHEMA_READ_REPLICA"
)
database_password_read_replica: Optional[str] = Field(
database_password_read_replica: str | None = Field(
default=None, validation_alias="DATABASE_PASSWORD_READ_REPLICA"
)
@ -106,7 +142,7 @@ class DatabaseURLSettings(BaseSettings):
"""Load the settings from ``os.environ`` (read at call time)."""
return cls()
def build_writer_url(self) -> Optional[str]:
def build_writer_url(self) -> str | None:
"""Return the writer URL to set, or ``None`` to leave it as-is.
Raises ``RuntimeError`` (naming the offending vars) when IAM auth is
@ -156,7 +192,7 @@ class DatabaseURLSettings(BaseSettings):
)
return None
def build_reader_url(self) -> Optional[str]:
def build_reader_url(self) -> str | None:
"""Return the read-replica URL to set, or ``None`` to leave it as-is.
Opt-in via ``DATABASE_HOST_READ_REPLICA``; never clobbers a
@ -217,11 +253,11 @@ class DatabaseURLSettings(BaseSettings):
def _password_url(
*,
user: str,
password: Optional[str],
password: str | None,
host: str,
port: str,
name: str,
schema: Optional[str],
schema: str | None,
) -> str:
"""Percent-encode credentials into a ``postgresql://`` URL.
@ -239,6 +275,26 @@ class DatabaseURLSettings(BaseSettings):
url += f"?schema={schema}"
return url
def _raise_for_unsupported_scheme(self) -> None:
"""Reject an operator-pinned non-PostgreSQL writer / direct / reader URL.
The componentized entrypoints (gateway / backend / migrations) call
``apply_to_env`` and then hand the URL straight to Prisma, bypassing
the CLI's own guard. A pinned URL flows through untouched, so validate
the same three vars the CLI guard checks (DATABASE_URL, DIRECT_URL, and
the read replica) rather than letting Prisma stall on an unusable scheme.
"""
for env_var, url in (
("DATABASE_URL", self.database_url),
("DIRECT_URL", self.direct_url),
("DATABASE_URL_READ_REPLICA", self.database_url_read_replica),
):
if not url:
continue
bad_scheme = unsupported_db_scheme(url)
if bad_scheme is not None:
raise RuntimeError(unsupported_db_scheme_message(env_var, bad_scheme))
def apply_to_env(self) -> bool:
"""Write the assembled URL(s) into ``os.environ``.
@ -246,6 +302,7 @@ class DatabaseURLSettings(BaseSettings):
password auth that assembled a fresh URL). False means there was
nothing to do — an operator-pinned URL, or no discrete fields.
"""
self._raise_for_unsupported_scheme()
wrote_writer = False
writer_url = self.build_writer_url()
if writer_url is not None:

View file

@ -123,11 +123,70 @@ class SemanticToolFilterHook(CustomLogger):
return openai_tools_as_dicts
def _is_mcp_tool(self, tool: object) -> bool:
"""
Check whether *tool* is registered in the MCP semantic router.
Classification strategy (shape-first, lookup-second):
1. Chat Completions format dicts are always native.
2. Responses API function tools are always native.
3. Everything else is looked up by name in the MCP registry.
"""
if (
isinstance(tool, dict)
and tool.get("type") == "function"
and isinstance(tool.get("function"), dict)
):
return False
if (
isinstance(tool, dict)
and tool.get("type") == "function"
and isinstance(tool.get("name"), str)
):
return False
name, _ = self.filter._extract_tool_info(tool)
return bool(name) and name in self.filter._tool_map
def _get_metadata_variable_name(self, data: dict) -> str:
if "litellm_metadata" in data:
return "litellm_metadata"
return "metadata"
def _emit_filter_metadata(
self,
data: dict,
mcp_tools: list[object],
filtered_mcp_tools: list[object],
native_tools: list[object],
filtered_tools: list[object],
) -> None:
"""
Emit response-header metadata when MCP tools were filtered.
Stats report MCP-only counts so downstream consumers see accurate
semantic filter metrics. Skips metadata entirely for purely-native
requests to avoid spurious headers.
"""
if mcp_tools:
filter_stats = f"{len(mcp_tools)}->{len(filtered_mcp_tools)}"
tool_names_csv = self._get_tool_names_csv(filtered_mcp_tools)
_metadata_variable_name = self._get_metadata_variable_name(data)
metadata = data.setdefault(_metadata_variable_name, {})
metadata["litellm_semantic_filter_stats"] = filter_stats
metadata["litellm_semantic_filter_tools"] = tool_names_csv
verbose_proxy_logger.info(
f"Semantic tool filter: {filter_stats} MCP tools "
f"({len(native_tools)} native preserved, "
f"{len(filtered_tools)} total)"
)
else:
verbose_proxy_logger.info(
f"Semantic tool filter: all {len(native_tools)} tools "
f"are native, no MCP filtering applied"
)
async def async_pre_call_hook(
self,
user_api_key_dict: "UserAPIKeyAuth",
@ -140,53 +199,55 @@ class SemanticToolFilterHook(CustomLogger):
This hook is called before the LLM request is made. It filters the
tools list to only include semantically relevant tools.
Args:
user_api_key_dict: User authentication
cache: Cache instance
data: Request data containing messages and tools
call_type: Type of call (completion, acompletion, etc.)
Returns:
Modified data dict with filtered tools, or None if no changes
"""
# Only filter endpoints that support tools
if call_type not in ("completion", "acompletion", "aresponses"):
verbose_proxy_logger.debug(
f"Skipping semantic filter for call_type={call_type}"
)
return None
# Check if tools are present
tools = data.get("tools")
if not tools:
verbose_proxy_logger.debug("No tools in request, skipping semantic filter")
return None
original_tool_count = len(tools)
# Check for MCP references (server_url="litellm_proxy") and expand them
# Expanded MCP tools are in OpenAI nested format which
# filter_tools/_extract_tool_info cannot name-match, so we skip
# semantic filtering and return early.
if self._should_expand_mcp_tools(tools):
verbose_proxy_logger.debug(
"Detected litellm_proxy MCP references, expanding before semantic filtering"
)
try:
native_tools_before_expand = [
t
for t in tools
if not (isinstance(t, dict) and t.get("type") == "mcp")
]
expanded_tools = await self._expand_mcp_tools(tools, user_api_key_dict)
if not expanded_tools:
if native_tools_before_expand:
data["tools"] = native_tools_before_expand
verbose_proxy_logger.warning(
"No MCP tools expanded, preserving "
f"{len(native_tools_before_expand)} native tools"
)
return data
verbose_proxy_logger.warning(
"No tools expanded from MCP references"
)
return None
data["tools"] = native_tools_before_expand + expanded_tools
verbose_proxy_logger.info(
f"Expanded {len(tools)} MCP reference(s) to {len(expanded_tools)} tools"
f"Expanded MCP references to {len(expanded_tools)} tools "
f"({len(native_tools_before_expand)} native preserved), "
f"skipping semantic filter (OpenAI nested format)"
)
# Update tools for filtering
tools = expanded_tools
original_tool_count = len(tools)
return data
except Exception as e:
verbose_proxy_logger.error(
@ -194,7 +255,6 @@ class SemanticToolFilterHook(CustomLogger):
)
return None
# Check if messages are present (try both "messages" and "input" for responses API)
messages = data.get("messages", [])
if not messages:
messages = data.get("input", [])
@ -204,13 +264,11 @@ class SemanticToolFilterHook(CustomLogger):
)
return None
# Check if filter is enabled
if not self.filter.enabled:
verbose_proxy_logger.debug("Semantic filter disabled, skipping")
return None
try:
# Extract user query from messages
user_query = self.filter.extract_user_query(messages)
if not user_query:
verbose_proxy_logger.debug(
@ -218,33 +276,60 @@ class SemanticToolFilterHook(CustomLogger):
)
return None
native_tools: list[object] = []
mcp_tools: list[object] = []
mcp_indices: set[int] = set()
for i, t in enumerate(tools):
if self._is_mcp_tool(t):
mcp_tools.append(t)
mcp_indices.add(i)
else:
native_tools.append(t)
verbose_proxy_logger.debug(
f"Applying semantic filter to {len(tools)} tools "
f"with query: '{user_query[:50]}...'"
f"Applying semantic filter: {len(mcp_tools)} MCP tools, "
f"{len(native_tools)} native tools, "
f"query: '{user_query[:50]}...'"
)
# Filter tools semantically
filtered_tools = await self.filter.filter_tools(
query=user_query,
available_tools=tools, # type: ignore
)
if mcp_tools:
filtered_mcp_tools = await self.filter.filter_tools(
query=user_query,
available_tools=mcp_tools, # type: ignore
)
else:
filtered_mcp_tools = []
filtered_mcp_names: set[str] = set()
for t in filtered_mcp_tools:
name, _ = self.filter._extract_tool_info(t)
if name:
filtered_mcp_names.add(name)
filtered_tools: list[object] = []
for i, t in enumerate(tools):
if i in mcp_indices:
name, _ = self.filter._extract_tool_info(t)
if name in filtered_mcp_names:
filtered_tools.append(t)
else:
filtered_tools.append(t)
# Always update tools and emit header (even if count unchanged)
data["tools"] = filtered_tools
# Store filter stats and tool names for response header
filter_stats = f"{original_tool_count}->{len(filtered_tools)}"
tool_names_csv = self._get_tool_names_csv(filtered_tools)
_metadata_variable_name = self._get_metadata_variable_name(data)
data[_metadata_variable_name][
"litellm_semantic_filter_stats"
] = filter_stats
data[_metadata_variable_name][
"litellm_semantic_filter_tools"
] = tool_names_csv
verbose_proxy_logger.info(f"Semantic tool filter: {filter_stats} tools")
try:
self._emit_filter_metadata(
data=data,
mcp_tools=mcp_tools,
filtered_mcp_tools=filtered_mcp_tools,
native_tools=native_tools,
filtered_tools=filtered_tools,
)
except Exception as e:
verbose_proxy_logger.warning(
f"Failed to emit semantic filter metadata: {e}",
exc_info=True,
)
return data
@ -266,7 +351,7 @@ class SemanticToolFilterHook(CustomLogger):
from litellm.constants import MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH
_metadata_variable_name = self._get_metadata_variable_name(data)
metadata = data[_metadata_variable_name]
metadata = data.get(_metadata_variable_name, {})
filter_stats = metadata.get("litellm_semantic_filter_stats")
if not filter_stats:

View file

@ -108,6 +108,7 @@ def parse_cache_control(cache_control):
LITELLM_METADATA_ROUTES = (
"batches",
"bedrock",
"/v1/messages",
"responses",
"files",
@ -1237,6 +1238,27 @@ class LiteLLMProxyRequestSetup:
return tags
@staticmethod
def pre_seed_litellm_metadata_for_route(
request_data: dict,
route: str,
) -> None:
"""Pre-seed ``litellm_metadata`` for routes that track tags there.
Routes in ``LITELLM_METADATA_ROUTES`` (e.g. Bedrock, ``/v1/messages``,
responses, batches, files) store request-scoped tag metadata in
``litellm_metadata`` rather than the provider-facing ``metadata``
field. ``get_metadata_variable_name_from_kwargs`` picks the target
based on whether ``litellm_metadata`` is present, so it must be
seeded BEFORE any tag merge runs; otherwise header tags from
``apply_client_tag_policy_pre_auth`` land in ``metadata`` while
key tags from ``apply_key_tags_pre_auth`` and the read in
``_tag_max_budget_check`` resolve to ``litellm_metadata``, leaving
header tags invisible to per-tag budget enforcement.
"""
if any(metadata_route in route for metadata_route in LITELLM_METADATA_ROUTES):
request_data.setdefault("litellm_metadata", {})
@staticmethod
def apply_key_tags_pre_auth(
request_data: dict,
@ -1468,8 +1490,7 @@ async def add_litellm_data_to_request(
_metadata_variable_name=_metadata_variable_name,
)
# Add headers to metadata for guardrails to access (fixes #17477)
# Guardrails use metadata["headers"] to access request headers (e.g., User-Agent)
# Expose request headers under the metadata field for guardrails (fixes #17477)
if _metadata_variable_name in data and isinstance(
data[_metadata_variable_name], dict
):

View file

@ -1,7 +1,7 @@
import asyncio
from datetime import datetime
from types import SimpleNamespace
from typing import Any, Callable, Dict, List, Optional, Set, Tuple, Union
from typing import Any, Awaitable, Callable, Dict, List, Optional, Set, Tuple, Union
from fastapi import HTTPException, status
@ -887,8 +887,17 @@ async def get_daily_activity(
exclude_entity_ids: Optional[List[str]] = None,
metadata_metrics_func: Optional[Callable[[List[Any]], SpendMetrics]] = None,
timezone_offset_minutes: Optional[int] = None,
resolve_entity_metadata: Optional[
Callable[[list[Any]], Awaitable[dict[str, dict]]]
] = None,
) -> SpendAnalyticsPaginatedResponse:
"""Common function to get daily activity for any entity type."""
"""Common function to get daily activity for any entity type.
``resolve_entity_metadata`` lets a caller resolve entity metadata from the
rows actually on the page (e.g. user_id -> user_email) instead of fetching
the whole entity table upfront, which matters when the entity set is
unbounded.
"""
if prisma_client is None:
raise HTTPException(
@ -939,11 +948,18 @@ async def get_daily_activity(
take=page_size,
)
resolved_entity_metadata = entity_metadata_field
if resolve_entity_metadata is not None:
resolved_entity_metadata = {
**(entity_metadata_field or {}),
**(await resolve_entity_metadata(daily_spend_data)),
}
aggregated = await _aggregate_spend_records(
prisma_client=prisma_client,
records=daily_spend_data,
entity_id_field=entity_id_field,
entity_metadata_field=entity_metadata_field,
entity_metadata_field=resolved_entity_metadata,
)
metadata_metrics = aggregated["totals"]

View file

@ -57,6 +57,9 @@ from litellm.repositories.verification_token_repository import (
from litellm.types.proxy.management_endpoints.common_daily_activity import (
SpendAnalyticsPaginatedResponse,
)
from litellm.types.proxy.management_endpoints.scim_v2 import (
SCIM_ENTERPRISE_METADATA_KEY,
)
from litellm.types.proxy.management_endpoints.internal_user_endpoints import (
BulkUpdateUserRequest,
BulkUpdateUserResponse,
@ -719,6 +722,17 @@ async def _get_user_info_teams(
return team_list, teams_1
def _redact_scim_enterprise_metadata(
metadata: Optional[Dict[str, Any]],
) -> Optional[Dict[str, Any]]:
"""SCIM enterprise attributes are persisted in user metadata so reporting can
group on them, but they are directory-only fields that generic user-info
endpoints must not surface; SCIM clients read them through the SCIM endpoints."""
if not isinstance(metadata, dict) or SCIM_ENTERPRISE_METADATA_KEY not in metadata:
return metadata
return {k: v for k, v in metadata.items() if k != SCIM_ENTERPRISE_METADATA_KEY}
def _build_user_info_response(
user_id: Optional[str],
user_info: Optional[Any],
@ -739,6 +753,9 @@ def _build_user_info_response(
)
if isinstance(_user_info, dict):
_user_info.pop("password", None)
_user_info["metadata"] = _redact_scim_enterprise_metadata(
_user_info.get("metadata")
)
return UserInfoResponse(
user_id=user_id,
@ -983,7 +1000,7 @@ async def user_info_v2(
models=user_data.get("models") or [],
budget_duration=user_data.get("budget_duration"),
budget_reset_at=user_data.get("budget_reset_at"),
metadata=user_data.get("metadata"),
metadata=_redact_scim_enterprise_metadata(user_data.get("metadata")),
created_at=user_data.get("created_at"),
updated_at=user_data.get("updated_at"),
sso_user_id=user_data.get("sso_user_id"),
@ -2098,9 +2115,13 @@ async def get_users(
user_list: List[LiteLLM_UserTableWithKeyCount] = []
if users is not None:
for user in users:
user_dump = user.model_dump()
user_dump["metadata"] = _redact_scim_enterprise_metadata(
user_dump.get("metadata")
)
user_list.append(
LiteLLM_UserTableWithKeyCount(
**user.model_dump(), key_count=user_key_counts.get(user.user_id, 0)
**user_dump, key_count=user_key_counts.get(user.user_id, 0)
)
)
else:
@ -2596,6 +2617,25 @@ async def ui_view_users(
# Using shared metric helper implementations from common_daily_activity
async def _resolve_user_email_metadata(
prisma_client: "PrismaClient", records: list[Any]
) -> dict[str, dict]:
"""Map each user_id on the page to its email/alias so the Usage dashboard can
label the 'Spend Per User' chart with the email instead of the raw UUID."""
user_ids = {
record.user_id for record in records if getattr(record, "user_id", None)
}
if not user_ids:
return {}
users = await UserRepository(prisma_client).table.find_many(
where={"user_id": {"in": list(user_ids)}}
)
return {
user.user_id: {"user_email": user.user_email, "user_alias": user.user_alias}
for user in users
}
@router.get(
"/user/daily/activity",
tags=["Budget & Spend Tracking", "Internal User management"],
@ -2698,6 +2738,9 @@ async def get_user_daily_activity(
page=page,
page_size=page_size,
timezone_offset_minutes=timezone,
resolve_entity_metadata=lambda records: _resolve_user_email_metadata(
prisma_client, records
),
)
except HTTPException:

View file

@ -50,8 +50,16 @@ class ScimTransformations:
scim_active = metadata.get("scim_active")
active = True if scim_active is None else bool(scim_active)
schemas = ["urn:ietf:params:scim:schemas:core:2.0:User"]
enterprise_user = None
if metadata.get(SCIM_ENTERPRISE_METADATA_KEY):
enterprise_user = SCIMEnterpriseUser.model_validate(
metadata[SCIM_ENTERPRISE_METADATA_KEY]
)
schemas.append(SCIM_ENTERPRISE_USER_SCHEMA)
return SCIMUser(
schemas=["urn:ietf:params:scim:schemas:core:2.0:User"],
schemas=schemas,
id=user.user_id,
userName=ScimTransformations._get_scim_user_name(user),
displayName=ScimTransformations._get_scim_user_name(user),
@ -62,6 +70,7 @@ class ScimTransformations:
emails=emails,
groups=groups,
active=active,
enterprise_user=enterprise_user,
meta={
"resourceType": "User",
"created": user_created_at,

View file

@ -5,7 +5,7 @@ This is an enterprise feature and requires a premium license.
"""
import re
from typing import Any, Dict, List, Optional, Set, Tuple
from typing import Any, Dict, Iterable, List, Optional, Set, Tuple
from fastapi import (
APIRouter,
@ -69,14 +69,21 @@ class UserProvisionerHelpers:
@staticmethod
async def handle_existing_user_by_email(
prisma_client, new_user_request: NewUserRequest
prisma_client,
new_user_request: NewUserRequest,
admin_group: Optional[str] = None,
) -> Optional[SCIMUser]:
"""
Check if a user with the given email already exists and update them if found.
When admin_group is configured the resolved global role on new_user_request
is persisted too, so re-upserting an existing email demotes a user who is no
longer in the admin group instead of leaving the stale role.
Args:
prisma_client: Database client
new_user_request: New user request data
admin_group: Configured SCIM admin group, or None to leave role untouched
Returns:
SCIMUser if user was updated, None if no existing user found
@ -100,6 +107,11 @@ class UserProvisionerHelpers:
"user_alias": new_user_request.user_alias,
"teams": new_user_request.teams,
"metadata": safe_dumps(new_user_request.metadata),
**(
{"user_role": new_user_request.user_role}
if admin_group is not None
else {}
),
},
)
@ -118,6 +130,7 @@ class ScimUserData(TypedDict):
given_name: Optional[str]
family_name: Optional[str]
active: Optional[bool]
enterprise: Optional[SCIMEnterpriseUser]
class GroupMemberExtractionResult(BaseModel):
@ -199,11 +212,15 @@ def _extract_scim_user_data(user: SCIMUser) -> ScimUserData:
"given_name": user.name.givenName if user.name else None,
"family_name": user.name.familyName if user.name else None,
"active": user.active,
"enterprise": user.enterprise_user,
}
def _build_scim_metadata(
given_name: Optional[str], family_name: Optional[str], active: Optional[bool] = None
given_name: Optional[str],
family_name: Optional[str],
active: Optional[bool] = None,
enterprise: Optional[SCIMEnterpriseUser] = None,
) -> Dict[str, Any]:
"""Build metadata dictionary with SCIM data."""
metadata: Dict[str, Any] = {
@ -216,6 +233,11 @@ def _build_scim_metadata(
if active is not None:
metadata["scim_active"] = active
if enterprise is not None:
metadata[SCIM_ENTERPRISE_METADATA_KEY] = enterprise.model_dump(
by_alias=True, exclude_none=True
)
return metadata
@ -244,6 +266,117 @@ async def _get_scim_upsert_user_setting() -> bool:
return True
ScimUserRole = Literal[
LitellmUserRoles.PROXY_ADMIN,
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
LitellmUserRoles.INTERNAL_USER,
LitellmUserRoles.INTERNAL_USER_VIEW_ONLY,
]
def _default_scim_user_role() -> ScimUserRole:
"""Non-admin default role for SCIM-provisioned users."""
if litellm.default_internal_user_params:
configured_role = litellm.default_internal_user_params.get("user_role")
if configured_role is not None:
return configured_role
return LitellmUserRoles.INTERNAL_USER_VIEW_ONLY
async def _get_scim_admin_group() -> Optional[str]:
"""
Get the scim_admin_group setting from litellm_settings.
Returns the configured admin group identifier, or None when unset so callers
leave a user's global role untouched (default-safe).
"""
try:
from litellm.proxy.proxy_server import proxy_config
config = await proxy_config.get_config()
litellm_settings = config.get("litellm_settings", {}) or {}
return litellm_settings.get("scim_admin_group") or None
except Exception as e:
verbose_proxy_logger.warning(
f"Error reading scim_admin_group setting, defaulting to None: {e}"
)
return None
def _resolve_scim_user_role(
groups: list[SCIMUserGroup],
admin_group: Optional[str],
default_role: ScimUserRole,
) -> Optional[LitellmUserRoles]:
"""
Resolve a user's global proxy role from their SCIM groups.
Returns None when no admin group is configured, signalling callers to leave
the role unchanged. Otherwise grants PROXY_ADMIN when any group matches the
admin group by value or display, and falls back to the non-admin default.
"""
if admin_group is None:
return None
for group in groups:
if group.value == admin_group or group.display == admin_group:
return LitellmUserRoles.PROXY_ADMIN
return default_role
async def _scim_groups_from_team_ids(
prisma_client: Any, team_ids: list[str]
) -> list[SCIMUserGroup]:
"""
Build SCIMUserGroup objects from team ids, populating display from each
team's alias so admin-group matching by display name works the same way it
does on PUT (where SCIM groups carry display names natively).
"""
teams = [
await TeamRepository(prisma_client).table.find_unique(
where={"team_id": team_id}
)
for team_id in team_ids
]
return [
SCIMUserGroup(
value=team_id,
display=team.team_alias if team is not None else None,
)
for team_id, team in zip(team_ids, teams)
]
async def _recompute_scim_member_roles(
prisma_client: Any, user_ids: Iterable[str]
) -> None:
"""
Recompute and persist each user's global proxy role from their resulting team
membership. No-op unless scim_admin_group is configured, so a SCIM group write
that drops a member from the admin group demotes them just like the user
endpoints do, and the role is left untouched when the feature is off.
"""
admin_group = await _get_scim_admin_group()
if admin_group is None:
return
default_role = _default_scim_user_role()
for user_id in user_ids:
user = await UserRepository(prisma_client).table.find_unique(
where={"user_id": user_id}
)
if user is None:
continue
resolved_role = _resolve_scim_user_role(
await _scim_groups_from_team_ids(prisma_client, user.teams or []),
admin_group,
default_role,
)
await UserRepository(prisma_client).table.update(
where={"user_id": user_id},
data={"user_role": resolved_role},
)
async def _extract_group_member_ids(group: SCIMGroup) -> GroupMemberExtractionResult:
"""
Extract member IDs from SCIMGroup, validating that all users exist.
@ -999,19 +1132,16 @@ async def create_user(
# Create user in database
user_id = user.userName or str(uuid.uuid4())
metadata = _build_scim_metadata(
user_data["given_name"], user_data["family_name"]
user_data["given_name"],
user_data["family_name"],
enterprise=user_data["enterprise"],
)
default_role: Optional[
Literal[
LitellmUserRoles.PROXY_ADMIN,
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
LitellmUserRoles.INTERNAL_USER,
LitellmUserRoles.INTERNAL_USER_VIEW_ONLY,
]
] = LitellmUserRoles.INTERNAL_USER_VIEW_ONLY
if litellm.default_internal_user_params:
default_role = litellm.default_internal_user_params.get("user_role")
default_role = _default_scim_user_role()
admin_group = await _get_scim_admin_group()
resolved_role = _resolve_scim_user_role(
user.groups or [], admin_group, default_role
)
new_user_request = NewUserRequest(
user_id=user_id,
@ -1020,12 +1150,14 @@ async def create_user(
teams=user_data["teams"],
metadata=metadata,
auto_create_key=False,
user_role=default_role,
user_role=resolved_role if admin_group is not None else default_role,
)
# Check if user with email already exists and update if found
existing_user_scim = await UserProvisionerHelpers.handle_existing_user_by_email(
prisma_client=prisma_client, new_user_request=new_user_request
prisma_client=prisma_client,
new_user_request=new_user_request,
admin_group=admin_group,
)
if existing_user_scim:
@ -1088,6 +1220,7 @@ async def update_user(
user_data["given_name"],
user_data["family_name"],
scim_active_for_metadata,
enterprise=user_data["enterprise"],
)
await _handle_team_membership_changes(
@ -1104,6 +1237,12 @@ async def update_user(
"metadata": safe_dumps(metadata),
}
admin_group = await _get_scim_admin_group()
if admin_group is not None:
update_data["user_role"] = _resolve_scim_user_role(
user.groups or [], admin_group, _default_scim_user_role()
)
updated_user = await UserRepository(prisma_client).table.update(
where={"user_id": user_id},
data=update_data,
@ -1417,6 +1556,14 @@ async def patch_user(
update_data["teams"] = list(final_team_set)
admin_group = await _get_scim_admin_group()
if admin_group is not None:
update_data["user_role"] = _resolve_scim_user_role(
await _scim_groups_from_team_ids(prisma_client, list(final_team_set)),
admin_group,
_default_scim_user_role(),
)
# Serialize metadata to JSON string for Prisma to avoid GraphQL parsing issues
if "metadata" in update_data and isinstance(update_data["metadata"], dict):
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
@ -1599,6 +1746,8 @@ async def create_group(
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
)
await _recompute_scim_member_roles(prisma_client, member_result.all_member_ids)
scim_group = await ScimTransformations.transform_litellm_team_to_scim_group(
created_team
)
@ -1665,6 +1814,19 @@ async def update_group(
final_members=final_members,
)
# A rename can flip whether this group matches scim_admin_group by display
# name, so retained members must be re-resolved too, not just the ones whose
# membership changed.
alias_changed = existing_team.team_alias != group.displayName
await _recompute_scim_member_roles(
prisma_client,
(
current_members | final_members
if alias_changed
else current_members ^ final_members
),
)
# Convert to SCIM format and return
scim_group = await ScimTransformations.transform_litellm_team_to_scim_group(
updated_team
@ -1691,8 +1853,10 @@ async def delete_group(
prisma_client = await _get_prisma_client_or_raise_exception()
existing_team = await _check_team_exists(group_id)
member_ids = await _get_team_member_user_ids_from_team(existing_team)
# For each member, remove this team from their teams list
for member_id in existing_team.members or []:
for member_id in member_ids:
user = await UserRepository(prisma_client).table.find_unique(
where={"user_id": member_id}
)
@ -1704,6 +1868,8 @@ async def delete_group(
where={"user_id": member_id}, data={"teams": new_teams}
)
await _recompute_scim_member_roles(prisma_client, member_ids)
# Delete team
await TeamRepository(prisma_client).table.delete(where={"team_id": group_id})
@ -1903,6 +2069,20 @@ async def patch_group(
# Handle user-team relationship changes
await _handle_group_membership_changes(group_id, current_members, final_members)
# A rename can flip whether this group matches scim_admin_group by display
# name, so retained members must be re-resolved too, not just the ones whose
# membership changed.
new_alias = update_data.get("team_alias", existing_team.team_alias)
alias_changed = new_alias != existing_team.team_alias
await _recompute_scim_member_roles(
prisma_client,
(
current_members | final_members
if alias_changed
else current_members ^ final_members
),
)
# Refresh team one more time to get final state after membership changes
final_team = await TeamRepository(prisma_client).table.find_unique(
where={"team_id": group_id}

View file

@ -11,6 +11,7 @@ from fastapi import HTTPException, status
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.proxy._types import SpecialMCPServerNames
from litellm.proxy.utils import PrismaClient
from litellm.repositories.object_permission_repository import ObjectPermissionRepository
from litellm.repositories.table_repositories import MCPServerRepository
@ -287,6 +288,9 @@ def _rewrite_object_permission_mcp_servers(
normalized_servers: List[str] = []
for identifier in mcp_servers:
if identifier == SpecialMCPServerNames.no_mcp_servers.value:
normalized_servers.append(SpecialMCPServerNames.no_mcp_servers.value)
continue
normalized_servers.extend(sorted(identifier_to_server_ids.get(identifier, [])))
object_permission["mcp_servers"] = _dedupe_preserving_order(normalized_servers)
@ -426,6 +430,7 @@ def _extract_requested_mcp_server_ids(
mcp_servers = object_permission.get("mcp_servers")
if isinstance(mcp_servers, list):
server_ids.update(mcp_servers)
server_ids.discard(SpecialMCPServerNames.no_mcp_servers.value)
mcp_tool_permissions = object_permission.get("mcp_tool_permissions")
if isinstance(mcp_tool_permissions, dict):

View file

@ -6,6 +6,7 @@ import httpx
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.core_helpers import map_finish_reason
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.litellm_logging import use_custom_pricing_for_model
from litellm.litellm_core_utils.prompt_templates.common_utils import (
@ -15,12 +16,19 @@ from litellm.llms.anthropic import get_anthropic_config
from litellm.llms.anthropic.chat.handler import (
ModelResponseIterator as AnthropicModelResponseIterator,
)
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
from litellm.proxy._types import PassThroughEndpointLoggingTypedDict
from litellm.proxy.auth.auth_utils import get_end_user_id_from_request_body
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
PassthroughStandardLoggingPayload,
)
from litellm.types.utils import LiteLLMBatch, ModelResponse, TextCompletionResponse
from litellm.types.utils import (
Choices,
LiteLLMBatch,
Message,
ModelResponse,
TextCompletionResponse,
)
if TYPE_CHECKING:
from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType
@ -272,6 +280,9 @@ class AnthropicPassthroughLoggingHandler:
kwargs["response_cost"] = response_cost
kwargs["model"] = model
# the pass-through success path reads spend from
# model_call_details["response_cost"], not from kwargs
logging_obj.model_call_details["response_cost"] = response_cost
passthrough_logging_payload: Optional[PassthroughStandardLoggingPayload] = ( # type: ignore
kwargs.get("passthrough_logging_payload")
)
@ -343,13 +354,42 @@ class AnthropicPassthroughLoggingHandler:
if chunk_model:
model = chunk_model
complete_streaming_response = (
AnthropicPassthroughLoggingHandler._build_complete_streaming_response(
all_chunks=all_chunks,
litellm_logging_obj=litellm_logging_obj,
model=model,
try:
complete_streaming_response = (
AnthropicPassthroughLoggingHandler._build_complete_streaming_response(
all_chunks=all_chunks,
litellm_logging_obj=litellm_logging_obj,
model=model,
)
)
)
except Exception as e:
# stream_chunk_builder re-raises assembly failures (as litellm.APIError)
# on large agentic tool-use / thinking streams; treat that the same as a
# None result so the usage-only fallback below still recovers cost
verbose_proxy_logger.warning(
"Anthropic passthrough: stream assembly raised (model=%s): %s; falling "
"back to usage-only cost from raw SSE events.",
model,
e,
)
complete_streaming_response = None
if complete_streaming_response is None:
# stream_chunk_builder cannot always reassemble large agentic streams, but
# Anthropic still emits token usage in the message_start / message_delta SSE
# events regardless of content shape; recover usage-only so cost is tracked.
# Guard it too: a raise here would defeat the point and drop the request
try:
complete_streaming_response = AnthropicPassthroughLoggingHandler._build_usage_only_response_from_chunks(
all_chunks=all_chunks,
model=model,
)
except Exception as e:
verbose_proxy_logger.warning(
"Anthropic passthrough: usage-only fallback failed (model=%s): %s",
model,
e,
)
complete_streaming_response = None
if complete_streaming_response is None:
verbose_proxy_logger.error(
"Unable to build complete streaming response for Anthropic passthrough endpoint, not logging..."
@ -636,6 +676,141 @@ class AnthropicPassthroughLoggingHandler:
)
return complete_streaming_response
@staticmethod
def _extract_sse_data(event_str: str) -> Optional[dict]:
"""Parse the JSON object from the ``data:`` line of an Anthropic SSE event."""
for line in event_str.splitlines():
stripped = line.strip()
if stripped.startswith("data:"):
payload = stripped[len("data:") :].strip()
if not payload or payload == "[DONE]":
return None
try:
return cast(dict, json.loads(payload))
except (ValueError, TypeError):
return None
return None
@staticmethod
def _build_usage_only_response_from_chunks(
all_chunks: Sequence[Union[str, bytes]],
model: str,
) -> Optional[ModelResponse]:
"""
Build a usage-bearing ModelResponse from Anthropic SSE token-usage events, for
cost tracking when stream_chunk_builder cannot reassemble the stream.
Anthropic emits usage in ``message_start`` (uncached input + cache tokens, and an
initial output_tokens) and the final ``message_delta`` (cumulative output_tokens)
regardless of the content/tool shape, so cost is recoverable even when full
content assembly fails. Returns ``None`` if no usage event is found.
"""
input_tokens = 0
cache_read = 0
cache_creation = 0
cache_creation_5m: Optional[int] = None
cache_creation_1h: Optional[int] = None
output_tokens = 0
web_search_requests: Optional[int] = None
tool_search_requests: Optional[int] = None
inference_geo: Optional[str] = None
stop_reason: Optional[str] = None
found_usage = False
resolved_model = model
for _chunk_str in all_chunks:
for (
event_str
) in AnthropicPassthroughLoggingHandler._split_sse_chunk_into_events(
_chunk_str
):
data = AnthropicPassthroughLoggingHandler._extract_sse_data(event_str)
if not data:
continue
event_type = data.get("type")
if event_type == "message_start":
message = data.get("message") or {}
if not resolved_model or resolved_model == "unknown":
resolved_model = message.get("model") or resolved_model
usage = message.get("usage") or {}
input_tokens = usage.get("input_tokens") or input_tokens
cache_read = usage.get("cache_read_input_tokens") or cache_read
cache_creation = (
usage.get("cache_creation_input_tokens") or cache_creation
)
_cc = usage.get("cache_creation")
if isinstance(_cc, dict):
cache_creation_5m = _cc.get("ephemeral_5m_input_tokens")
cache_creation_1h = _cc.get("ephemeral_1h_input_tokens")
if usage.get("inference_geo") is not None:
inference_geo = usage.get("inference_geo")
if usage.get("output_tokens") is not None:
output_tokens = usage.get("output_tokens")
found_usage = True
elif event_type == "message_delta":
_delta_stop = (data.get("delta") or {}).get("stop_reason")
if _delta_stop:
stop_reason = _delta_stop
usage = data.get("usage") or {}
if usage.get("output_tokens") is not None:
output_tokens = usage.get("output_tokens")
_stu = usage.get("server_tool_use")
if isinstance(_stu, dict):
if _stu.get("web_search_requests") is not None:
web_search_requests = _stu.get("web_search_requests")
if _stu.get("tool_search_requests") is not None:
tool_search_requests = _stu.get("tool_search_requests")
if usage.get("cache_read_input_tokens") is not None:
cache_read = usage.get("cache_read_input_tokens")
if usage.get("inference_geo") is not None:
inference_geo = usage.get("inference_geo")
found_usage = True
if not found_usage:
return None
# If only the 5m/1h split was provided, derive the cache_creation total from it.
if not cache_creation and (cache_creation_5m or cache_creation_1h):
cache_creation = (cache_creation_5m or 0) + (cache_creation_1h or 0)
# build usage via the same AnthropicConfig.calculate_usage path the success
# cases use, so prompt_tokens are cache-inclusive and cache / server_tool_use /
# inference_geo tokens are priced instead of left at $0
usage_object: dict = {
"input_tokens": input_tokens,
"output_tokens": output_tokens,
}
if cache_read:
usage_object["cache_read_input_tokens"] = cache_read
if cache_creation:
usage_object["cache_creation_input_tokens"] = cache_creation
if cache_creation_5m is not None or cache_creation_1h is not None:
usage_object["cache_creation"] = {
"ephemeral_5m_input_tokens": cache_creation_5m or 0,
"ephemeral_1h_input_tokens": cache_creation_1h or 0,
}
if web_search_requests is not None or tool_search_requests is not None:
_server_tool_use: dict = {}
if web_search_requests is not None:
_server_tool_use["web_search_requests"] = web_search_requests
if tool_search_requests is not None:
_server_tool_use["tool_search_requests"] = tool_search_requests
usage_object["server_tool_use"] = _server_tool_use
if inference_geo is not None:
usage_object["inference_geo"] = inference_geo
usage_obj = AnthropicConfig().calculate_usage(
usage_object=usage_object, reasoning_content=None
)
return ModelResponse(
model=resolved_model,
choices=[
Choices(
finish_reason=(
map_finish_reason(stop_reason) if stop_reason else "stop"
),
index=0,
message=Message(role="assistant", content=""),
)
],
usage=usage_obj,
)
@staticmethod
def batch_creation_handler(
httpx_response: httpx.Response,

View file

@ -116,6 +116,9 @@ class BasePassthroughLoggingHandler(ABC):
kwargs["response_cost"] = response_cost
kwargs["model"] = model
# the pass-through success path reads spend from
# model_call_details["response_cost"], not from kwargs
logging_obj.model_call_details["response_cost"] = response_cost
passthrough_logging_payload: Optional[PassthroughStandardLoggingPayload] = ( # type: ignore
kwargs.get("passthrough_logging_payload")
)

View file

@ -285,8 +285,10 @@ class PassThroughStreamingHandler:
Returns:
List of string lines, with each line being a complete data: {} chunk
"""
# Combine all bytes and decode to string
combined_str = b"".join(raw_bytes).decode("utf-8")
# errors="replace" so a stream cut mid-multibyte-sequence (client disconnect)
# still decodes and logs the usage events already received, instead of raising
# and dropping the whole request from SpendLogs
combined_str = b"".join(raw_bytes).decode("utf-8", errors="replace")
# Split by newlines and filter out empty lines
lines = [line.strip() for line in combined_str.split("\n") if line.strip()]

View file

@ -1195,6 +1195,25 @@ def run_server(
os.getenv("DATABASE_URL", None) is not None
or os.getenv("DIRECT_URL", None) is not None
):
from litellm.proxy.db.db_url_settings import (
unsupported_db_scheme,
unsupported_db_scheme_message,
)
for _db_env in ("DATABASE_URL", "DIRECT_URL"):
_candidate_url = os.getenv(_db_env)
if _candidate_url is None:
continue
_bad_scheme = unsupported_db_scheme(_candidate_url)
if _bad_scheme is not None:
print(
f"\033[1;31mLiteLLM Proxy: "
f"{unsupported_db_scheme_message(_db_env, _bad_scheme)}"
"\033[0m",
file=sys.stderr,
flush=True,
)
sys.exit(1)
try:
from litellm.secret_managers.main import get_secret

View file

@ -106,6 +106,10 @@ from litellm.proxy.common_utils.callback_utils import (
process_callback,
)
from litellm.proxy.common_utils.realtime_utils import _realtime_request_body
from litellm.router_utils.add_retry_fallback_headers import (
get_fallback_errors_from_headers,
get_hidden_params_dict,
)
from litellm.types.utils import (
ModelResponse,
ModelResponseStream,
@ -7085,57 +7089,122 @@ def _get_client_requested_model_for_streaming(request_data: dict) -> str:
return requested_model if isinstance(requested_model, str) else ""
def _is_positive_int_like(value: Any) -> bool:
try:
return int(value) > 0
except (TypeError, ValueError):
return False
def _should_include_fallback_errors(request_data: dict[str, object]) -> bool:
if not general_settings.get("expose_fallback_errors_to_caller"):
return False
return request_data.get("include_fallback_errors") is True
def _get_streaming_fallback_metadata(
response_obj: object,
) -> tuple[bool, str | None, list[dict[str, object]]]:
additional_headers = get_hidden_params_dict(response_obj).get("additional_headers")
if not isinstance(additional_headers, dict):
return False, None, []
if not _is_positive_int_like(
additional_headers.get("x-litellm-attempted-fallbacks")
):
return False, None, []
fallback_model = additional_headers.get("x-litellm-model-group")
fallback_errors = get_fallback_errors_from_headers(additional_headers)
if isinstance(fallback_model, str) and fallback_model:
return True, fallback_model, fallback_errors
return True, None, fallback_errors
def _format_fallback_metadata_sse_event(
*,
fallback_model: str | None,
fallback_errors: list[dict[str, object]],
) -> str:
import time
payload = {
"id": "litellm-fallback-metadata",
"object": "chat.completion.chunk",
"created": int(time.time()),
"model": fallback_model or "",
"choices": [],
"litellm_fallback": {
"fallback_model": fallback_model,
"errors": fallback_errors,
},
}
return f"data: {json.dumps(payload)}\n\n"
def _restamp_streaming_chunk_model(
*,
chunk: Any,
requested_model_from_client: str,
request_data: dict,
model_mismatch_logged: bool,
) -> Tuple[Any, bool]:
fallback_was_attempted: bool = False,
fallback_model_from_metadata: str | None = None,
) -> tuple[Any, bool]:
target_model = (
fallback_model_from_metadata
if fallback_was_attempted
else requested_model_from_client
)
# Always return the client-requested model name (not provider-prefixed internal identifiers)
# on streaming chunks.
# On fallback, use the public OpenAI-compatible model name. This keeps
# provider-prefixed internal identifiers from leaking into the public API.
#
# Note: This warning is intentionally verbose. A mismatch is a useful signal that an
# internal provider/deployment identifier is leaking into the public API, and helps
# maintainers/operators catch regressions while preserving OpenAI-compatible output.
if not requested_model_from_client or not isinstance(chunk, (BaseModel, dict)):
if not target_model or not isinstance(chunk, (BaseModel, dict)):
return chunk, model_mismatch_logged
# For Azure Model Router, preserve the actual model used in each chunk
if _is_azure_model_router_request(requested_model_from_client):
if not fallback_was_attempted and _is_azure_model_router_request(
requested_model_from_client
):
return chunk, model_mismatch_logged
# For fastest_response batch completions, preserve the winning model's name
# instead of stamping the comma-separated list the client sent.
if request_data.get("fastest_response", False):
if not fallback_was_attempted and request_data.get("fastest_response", False):
return chunk, model_mismatch_logged
downstream_model = (
chunk.get("model") if isinstance(chunk, dict) else getattr(chunk, "model", None)
)
if downstream_model == requested_model_from_client:
if downstream_model == target_model:
return chunk, model_mismatch_logged
if not model_mismatch_logged and downstream_model != requested_model_from_client:
if not model_mismatch_logged and downstream_model != target_model:
verbose_proxy_logger.debug(
"litellm_call_id=%s: streaming chunk model mismatch - requested=%r downstream=%r. Overriding model to requested.",
"litellm_call_id=%s: streaming chunk model mismatch - target=%r downstream=%r fallback_was_attempted=%s. Overriding chunk model to target.",
request_data.get("litellm_call_id"),
requested_model_from_client,
target_model,
downstream_model,
fallback_was_attempted,
)
model_mismatch_logged = True
if isinstance(chunk, dict):
chunk["model"] = requested_model_from_client
chunk["model"] = target_model
return chunk, model_mismatch_logged
try:
setattr(chunk, "model", requested_model_from_client)
chunk.model = target_model
except Exception as e:
verbose_proxy_logger.error(
"litellm_call_id=%s: failed to override chunk.model=%r on chunk_type=%s. error=%s",
request_data.get("litellm_call_id"),
requested_model_from_client,
target_model,
type(chunk),
str(e),
exc_info=True,
@ -7294,7 +7363,14 @@ async def async_data_generator(
requested_model_from_client = _get_client_requested_model_for_streaming(
request_data=request_data
)
(
fallback_was_attempted,
fallback_model_from_metadata,
fallback_errors,
) = _get_streaming_fallback_metadata(response)
model_mismatch_logged = False
fallback_metadata_event_sent = False
include_fallback_errors = _should_include_fallback_errors(request_data)
# Use a running string instead of list + join to avoid O(n^2) overhead.
# Previously "".join(str_so_far_parts) was called every chunk, re-joining
# the entire accumulated response. String += is O(n) amortized total.
@ -7332,13 +7408,37 @@ async def async_data_generator(
str_so_far=_str_so_far,
)
# Mid-stream fallbacks surface metadata on individual chunks rather than
# the response wrapper. Keep scanning chunks until a fallback model is
# resolved, then latch it for the rest of the stream.
if fallback_model_from_metadata is None:
(
chunk_fallback_was_attempted,
chunk_fallback_model,
chunk_fallback_errors,
) = _get_streaming_fallback_metadata(chunk)
if chunk_fallback_was_attempted:
fallback_was_attempted = True
fallback_model_from_metadata = chunk_fallback_model
fallback_errors = fallback_errors or chunk_fallback_errors
pending_fallback_event = (
include_fallback_errors
and fallback_was_attempted
and fallback_errors
and not fallback_metadata_event_sent
)
chunk, model_mismatch_logged = _restamp_streaming_chunk_model(
chunk=chunk,
requested_model_from_client=requested_model_from_client,
request_data=request_data,
model_mismatch_logged=model_mismatch_logged,
fallback_was_attempted=fallback_was_attempted,
fallback_model_from_metadata=fallback_model_from_metadata,
)
raw_passthrough = False
if isinstance(chunk, BaseModel):
chunk = _serialize_streaming_chunk(chunk)
elif isinstance(chunk, bytes):
@ -7354,14 +7454,14 @@ async def async_data_generator(
raise ValueError(
"Raw SSE stream exceeded maximum buffered size without a frame delimiter"
)
continue
if chunk.startswith(("data:", "event:", ":")):
raw_passthrough = True
elif chunk.startswith(("data:", "event:", ":")):
yield (
chunk
if chunk.endswith(_SSE_FRAME_DELIMITERS)
else chunk + "\n\n"
)
continue
raw_passthrough = True
elif isinstance(chunk, str) and is_raw_sse_stream:
raw_sse_buffer += chunk
while True:
@ -7373,15 +7473,23 @@ async def async_data_generator(
raise ValueError(
"Raw SSE stream exceeded maximum buffered size without a frame delimiter"
)
continue
raw_passthrough = True
elif isinstance(chunk, str) and chunk.startswith("data: "):
error_message = chunk
break
try:
yield _format_streaming_sse_chunk(chunk=chunk)
except Exception as e:
yield f"data: {str(e)}\n\n"
if not raw_passthrough:
try:
yield _format_streaming_sse_chunk(chunk=chunk)
except Exception as e:
yield f"data: {str(e)}\n\n"
if pending_fallback_event:
yield _format_fallback_metadata_sse_event(
fallback_model=fallback_model_from_metadata,
fallback_errors=fallback_errors,
)
fallback_metadata_event_sent = True
stream_completed = True
if not needs_iterator_wrap:
@ -11581,9 +11689,12 @@ async def _get_caller_byok_team_scope(
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
):
return None
key_team_scope: set[str] = (
{user_api_key_dict.team_id} if user_api_key_dict.team_id else set()
)
user_id = user_api_key_dict.user_id
if user_id is None:
return set()
return key_team_scope
try:
user_row = await UserRepository(prisma_client).table.find_unique(
where={"user_id": user_id}
@ -11591,12 +11702,12 @@ async def _get_caller_byok_team_scope(
except Exception:
verbose_proxy_logger.exception(
"Failed to look up caller teams while scoping BYOK search; "
"defaulting to no team access."
"defaulting to key team scope only."
)
return set()
return key_team_scope
if user_row is None:
return set()
return set(user_row.teams or [])
return key_team_scope
return key_team_scope | set(user_row.teams or [])
def _byok_row_outside_caller_teams(

View file

@ -209,6 +209,114 @@
],
"default_model_placeholder": "claude-3-opus"
},
{
"provider": "BedrockMantle",
"provider_display_name": "Amazon Bedrock Mantle",
"litellm_provider": "bedrock_mantle",
"credential_fields": [
{
"key": "api_key",
"label": "Bedrock Mantle API Key",
"placeholder": null,
"tooltip": "Bearer token for the Bedrock Mantle OpenAI-compatible endpoint. You can provide the raw token or the environment variable (e.g. `os.environ/BEDROCK_MANTLE_API_KEY`). Leave blank to authenticate with AWS SigV4 credentials instead.",
"required": false,
"field_type": "password",
"options": null,
"default_value": null
},
{
"key": "aws_access_key_id",
"label": "AWS Access Key ID",
"placeholder": null,
"tooltip": "Used for AWS SigV4 auth when no API key is set. You can provide the raw key or the environment variable (e.g. `os.environ/MY_SECRET_KEY`).",
"required": false,
"field_type": "password",
"options": null,
"default_value": null
},
{
"key": "aws_secret_access_key",
"label": "AWS Secret Access Key",
"placeholder": null,
"tooltip": "Used for AWS SigV4 auth when no API key is set. You can provide the raw key or the environment variable (e.g. `os.environ/MY_SECRET_KEY`).",
"required": false,
"field_type": "password",
"options": null,
"default_value": null
},
{
"key": "aws_session_token",
"label": "AWS Session Token",
"placeholder": null,
"tooltip": "Temporary credentials session token. You can provide the raw token or the environment variable (e.g. `os.environ/MY_SESSION_TOKEN`).",
"required": false,
"field_type": "password",
"options": null,
"default_value": null
},
{
"key": "aws_region_name",
"label": "AWS Region Name",
"placeholder": "us-east-1",
"tooltip": "Region of the Bedrock Mantle endpoint. Defaults to us-east-1. You can provide the raw value or the environment variable (e.g. `os.environ/AWS_REGION_NAME`).",
"required": false,
"field_type": "text",
"options": null,
"default_value": null
},
{
"key": "aws_session_name",
"label": "AWS Session Name",
"placeholder": "my-session",
"tooltip": "Name for the AWS session. You can provide the raw value or the environment variable (e.g. `os.environ/MY_SESSION_NAME`).",
"required": false,
"field_type": "text",
"options": null,
"default_value": null
},
{
"key": "aws_profile_name",
"label": "AWS Profile Name",
"placeholder": "default",
"tooltip": "AWS profile name to use for authentication. You can provide the raw value or the environment variable (e.g. `os.environ/MY_PROFILE_NAME`).",
"required": false,
"field_type": "text",
"options": null,
"default_value": null
},
{
"key": "aws_role_name",
"label": "AWS Role Name",
"placeholder": "MyRole",
"tooltip": "AWS IAM role name to assume. You can provide the raw value or the environment variable (e.g. `os.environ/MY_ROLE_NAME`).",
"required": false,
"field_type": "text",
"options": null,
"default_value": null
},
{
"key": "aws_web_identity_token",
"label": "AWS Web Identity Token",
"placeholder": null,
"tooltip": "Web identity token for OIDC authentication. You can provide the raw token or the environment variable (e.g. `os.environ/MY_WEB_IDENTITY_TOKEN`).",
"required": false,
"field_type": "password",
"options": null,
"default_value": null
},
{
"key": "api_base",
"label": "API Base",
"placeholder": "https://bedrock-mantle.us-east-1.api.aws",
"tooltip": "Optional. Custom Bedrock Mantle endpoint. Defaults to https://bedrock-mantle.<region>.api.aws. You can provide the raw value or the environment variable (e.g. `os.environ/BEDROCK_MANTLE_API_BASE`).",
"required": false,
"field_type": "text",
"options": null,
"default_value": null
}
],
"default_model_placeholder": "bedrock_mantle/openai.gpt-oss-120b"
},
{
"provider": "Anthropic",
"provider_display_name": "Anthropic",

View file

@ -3792,6 +3792,10 @@ class PrismaClient:
db_data["members_with_roles"], list
):
db_data["members_with_roles"] = json.dumps(db_data["members_with_roles"])
if db_data.get("budget_limits", None) is not None and isinstance(
db_data["budget_limits"], list
):
db_data["budget_limits"] = json.dumps(db_data["budget_limits"])
return db_data
# Define a retrying strategy with exponential backoff

View file

@ -58,7 +58,10 @@ from litellm.llms.openai.data_residency import infer_openai_data_residency
from litellm.secret_managers.main import get_secret_str
from litellm.types.responses.main import *
from litellm.types.router import GenericLiteLLMParams
from litellm.utils import ProviderConfigManager, client
from litellm.utils import (
ProviderConfigManager,
client,
)
if TYPE_CHECKING:
from mcp.types import Tool as MCPTool

View file

@ -40,7 +40,6 @@ import anyio
import httpx
import openai
from openai import AsyncOpenAI
from pydantic import BaseModel
from typing_extensions import overload
import litellm
@ -81,8 +80,10 @@ from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2
from litellm.router_strategy.simple_shuffle import simple_shuffle
from litellm.router_strategy.tag_based_routing import get_deployments_for_tag
from litellm.router_utils.add_retry_fallback_headers import (
_HiddenParamsHost,
add_fallback_headers_to_response,
add_retry_headers_to_response,
get_hidden_params_dict,
)
from litellm.router_utils.batch_utils import (
_get_router_metadata_variable_name,
@ -2165,6 +2166,36 @@ class Router:
)
setattr(fallback_item, "usage", combined_usage)
@staticmethod
def _prepare_fallback_hidden_params(
fallback_response: object,
) -> tuple[dict[str, object], dict[str, object]]:
fallback_hidden_params = get_hidden_params_dict(fallback_response)
fallback_headers = fallback_hidden_params.get("additional_headers")
if not isinstance(fallback_headers, dict):
return fallback_hidden_params, {}
return fallback_hidden_params, cast("dict[str, object]", fallback_headers)
@staticmethod
def _apply_fallback_hidden_params_to_item(
fallback_item: object,
prepared_fallback_hidden_params: tuple[dict[str, object], dict[str, object]],
) -> None:
if fallback_item is None or not hasattr(fallback_item, "_hidden_params"):
return
fallback_hidden_params, fallback_headers = prepared_fallback_hidden_params
item_hidden_params = get_hidden_params_dict(fallback_item)
item_headers = item_hidden_params.get("additional_headers")
if not isinstance(item_headers, dict):
item_headers = {}
cast(_HiddenParamsHost, fallback_item)._hidden_params = {
**item_hidden_params,
**fallback_hidden_params,
"additional_headers": {**item_headers, **fallback_headers},
}
async def _acompletion_streaming_iterator(
self,
model_response: CustomStreamWrapper,
@ -2257,12 +2288,22 @@ class Router:
model_group=model_group,
args=(),
kwargs=initial_kwargs,
include_fallback_errors=initial_kwargs.get(
"include_fallback_errors", False
)
is True,
)
)
# If fallback returns a streaming response, iterate over it
if hasattr(fallback_response, "__aiter__"):
prepared_fallback_hidden_params = (
Router._prepare_fallback_hidden_params(fallback_response)
)
async for fallback_item in fallback_response: # type: ignore
Router._apply_fallback_hidden_params_to_item(
fallback_item, prepared_fallback_hidden_params
)
if (
fallback_item
and isinstance(fallback_item, ModelResponseStream)
@ -2686,11 +2727,21 @@ class Router:
model_group=model_group,
args=(),
kwargs=initial_kwargs,
include_fallback_errors=initial_kwargs.get(
"include_fallback_errors", False
)
is True,
)
)
if hasattr(fallback_response, "__aiter__"):
prepared_fallback_hidden_params = (
Router._prepare_fallback_hidden_params(fallback_response)
)
async for fallback_item in fallback_response: # type: ignore
Router._apply_fallback_hidden_params_to_item(
fallback_item, prepared_fallback_hidden_params
)
if partial_usage is not None:
Router._combine_responses_fallback_usage(
fallback_item, partial_usage
@ -2815,7 +2866,13 @@ class Router:
)
if hasattr(fallback_response, "__iter__"):
prepared_fallback_hidden_params = (
Router._prepare_fallback_hidden_params(fallback_response)
)
for fallback_item in fallback_response:
Router._apply_fallback_hidden_params_to_item(
fallback_item, prepared_fallback_hidden_params
)
if (
fallback_item
and isinstance(fallback_item, ModelResponseStream)
@ -2972,6 +3029,7 @@ class Router:
**kwargs,
}
input_kwargs.pop("silent_model", None)
input_kwargs.pop("include_fallback_errors", None)
_response = litellm.acompletion(**input_kwargs)
@ -3076,7 +3134,18 @@ class Router:
- litellm_trace_id
- metadata
"""
kwargs["num_retries"] = kwargs.get("num_retries", self.num_retries)
# Normalise an explicit num_retries=None to the router default here (dict.get()
# only falls back when the key is absent, not when its value is None), then to 0
# if the router default is itself None - mirroring the guard in
# async_function_with_retries, which remains the safety net for paths that bypass
# this setter.
_req_num_retries = kwargs.get("num_retries")
if _req_num_retries is not None:
kwargs["num_retries"] = _req_num_retries
else:
kwargs["num_retries"] = (
self.num_retries if self.num_retries is not None else 0
)
kwargs.setdefault("litellm_trace_id", str(uuid.uuid4()))
model_group_alias: Optional[str] = None
if self._get_model_from_alias(model=model):
@ -6478,6 +6547,7 @@ class Router:
model_group: Optional[str],
args: tuple,
kwargs: dict,
include_fallback_errors: bool = False,
):
"""
Common utilities for async_function_with_fallbacks
@ -6501,6 +6571,8 @@ class Router:
input_kwargs["max_fallbacks"] = self.max_fallbacks
if "fallback_depth" not in input_kwargs:
input_kwargs["fallback_depth"] = 0
if include_fallback_errors:
input_kwargs["include_fallback_errors"] = True
# ORDER-BASED FALLBACKS: prepend higher order levels to the fallback list
# Skip for error types that have their own dedicated fallback handlers
@ -6759,6 +6831,7 @@ class Router:
If it fails after num_retries, fall back to another model group
"""
model_group: Optional[str] = kwargs.get("model")
include_fallback_errors = kwargs.get("include_fallback_errors", False) is True
disable_fallbacks: Optional[bool] = kwargs.pop("disable_fallbacks", False)
fallbacks: Optional[List] = kwargs.get("fallbacks", self.fallbacks)
context_window_fallbacks: Optional[List] = kwargs.get(
@ -6802,6 +6875,7 @@ class Router:
model_group,
args,
kwargs,
include_fallback_errors=include_fallback_errors,
)
def _handle_mock_testing_fallbacks(
@ -6868,7 +6942,11 @@ class Router:
"model_group_retry_policy", self.model_group_retry_policy
)
model_group: Optional[str] = kwargs.get("model")
num_retries = kwargs.pop("num_retries")
num_retries = kwargs.pop("num_retries", None)
if num_retries is None:
# Fall back to the router setting (then 0) so the comparisons below never
# hit `None > int`, which would mask the real upstream error with a TypeError.
num_retries = self.num_retries if self.num_retries is not None else 0
## ADD MODEL GROUP SIZE TO METADATA - used for model_group_rate_limit_error tracking
_metadata: dict = kwargs.get("litellm_metadata", kwargs.get("metadata")) or {}
@ -9725,17 +9803,19 @@ class Router:
# - if healthy_deployments > 1, return model group rate limit headers
# - else return the model's rate limit headers
"""
if (
isinstance(response, BaseModel)
and hasattr(response, "_hidden_params")
and isinstance(response._hidden_params, dict) # type: ignore
):
response._hidden_params.setdefault("additional_headers", {}) # type: ignore
response._hidden_params["additional_headers"][ # type: ignore
"x-litellm-model-group"
] = model_group
if response is not None and hasattr(response, "_hidden_params"):
hidden_params = getattr(response, "_hidden_params", {}) or {}
if hasattr(hidden_params, "model_dump"):
hidden_params = hidden_params.model_dump()
if not isinstance(hidden_params, dict):
return response
response._hidden_params = hidden_params
additional_headers = response._hidden_params["additional_headers"] # type: ignore
additional_headers = hidden_params.get("additional_headers")
if not isinstance(additional_headers, dict):
additional_headers = {}
hidden_params["additional_headers"] = additional_headers
additional_headers["x-litellm-model-group"] = model_group
# Lift QualityRouter routing decision into response headers for
# transparency. The decision is stashed in request_kwargs.metadata

View file

@ -1,44 +1,99 @@
from typing import Any, Optional, Union
import json
from typing import Protocol, TypedDict, cast
from pydantic import BaseModel
from litellm.types.utils import HiddenParams
class FallbackErrorInfo(TypedDict):
message: str
type: str
param: str | None
code: str | None
def _add_headers_to_response(response: Any, headers: dict) -> Any:
class _HiddenParamsHost(Protocol):
_hidden_params: dict[str, object]
def get_hidden_params_dict(response: object) -> dict[str, object]:
hidden_params: object = cast(object, getattr(response, "_hidden_params", None))
if isinstance(hidden_params, BaseModel):
return cast("dict[str, object]", hidden_params.model_dump())
if isinstance(hidden_params, dict):
return cast("dict[str, object]", hidden_params)
return {}
def _ensure_additional_headers_dict(
hidden_params: dict[str, object],
) -> dict[str, object]:
additional_headers = hidden_params.get("additional_headers")
if isinstance(additional_headers, dict):
return cast("dict[str, object]", additional_headers)
return {}
def get_fallback_error_info(error: Exception) -> FallbackErrorInfo:
message = cast(object, getattr(error, "message", str(error)))
error_type = cast(object, getattr(error, "type", error.__class__.__name__))
param = cast(object, getattr(error, "param", None))
code = cast(object, getattr(error, "status_code", getattr(error, "code", None)))
return FallbackErrorInfo(
message=str(message),
type=str(error_type),
param=str(param) if param is not None else None,
code=str(code) if code is not None else None,
)
def _coerce_error_dicts(items: list[object]) -> list[dict[str, object]]:
return [cast("dict[str, object]", item) for item in items if isinstance(item, dict)]
def get_fallback_errors_from_headers(
additional_headers: dict[str, object],
) -> list[dict[str, object]]:
existing_errors = additional_headers.get("x-litellm-fallback-errors")
if isinstance(existing_errors, list):
return _coerce_error_dicts(cast("list[object]", existing_errors))
if isinstance(existing_errors, str):
try:
parsed_errors: object = cast(object, json.loads(existing_errors))
except json.JSONDecodeError:
return []
if isinstance(parsed_errors, list):
return _coerce_error_dicts(cast("list[object]", parsed_errors))
return []
def _add_headers_to_response(response: object, headers: dict[str, object]) -> object:
"""
Helper function to add headers to a response's hidden params
"""
if response is None or not isinstance(response, BaseModel):
if response is None:
return response
hidden_params: Optional[Union[dict, HiddenParams]] = getattr(
response, "_hidden_params", {}
)
if not isinstance(response, BaseModel) and not hasattr(response, "_hidden_params"):
return response
if hidden_params is None:
hidden_params_dict = {}
elif isinstance(hidden_params, HiddenParams):
hidden_params_dict = hidden_params.model_dump()
else:
hidden_params_dict = hidden_params
hidden_params = get_hidden_params_dict(response)
additional_headers = _ensure_additional_headers_dict(hidden_params)
additional_headers.update(headers)
hidden_params["additional_headers"] = additional_headers
hidden_params_dict.setdefault("additional_headers", {})
hidden_params_dict["additional_headers"].update(headers)
setattr(response, "_hidden_params", hidden_params_dict)
cast(_HiddenParamsHost, response)._hidden_params = hidden_params
return response
def add_retry_headers_to_response(
response: Any,
response: object,
attempted_retries: int,
max_retries: Optional[int] = None,
) -> Any:
max_retries: int | None = None,
) -> object:
"""
Add retry headers to the request
"""
retry_headers = {
retry_headers: dict[str, object] = {
"x-litellm-attempted-retries": attempted_retries,
}
if max_retries is not None:
@ -48,9 +103,10 @@ def add_retry_headers_to_response(
def add_fallback_headers_to_response(
response: Any,
response: object,
attempted_fallbacks: int,
) -> Any:
fallback_errors: list[FallbackErrorInfo] | None = None,
) -> object:
"""
Add fallback headers to the response
@ -64,7 +120,19 @@ def add_fallback_headers_to_response(
Note: It's intentional that we don't add max_fallbacks in response headers
Want to avoid bloat in the response headers for performance.
"""
fallback_headers = {
fallback_headers: dict[str, object] = {
"x-litellm-attempted-fallbacks": attempted_fallbacks,
}
return _add_headers_to_response(response, fallback_headers)
response = _add_headers_to_response(response, fallback_headers)
if fallback_errors is None or response is None:
return response
hidden_params = get_hidden_params_dict(response)
additional_headers = _ensure_additional_headers_dict(hidden_params)
merged_errors = get_fallback_errors_from_headers(additional_headers) + [
cast("dict[str, object]", error) for error in fallback_errors
]
additional_headers["x-litellm-fallback-errors"] = json.dumps(merged_errors)
hidden_params["additional_headers"] = additional_headers
cast(_HiddenParamsHost, response)._hidden_params = hidden_params
return response

View file

@ -38,6 +38,7 @@ class CooldownCache:
visible_prefix=50, # Show first 50 characters
visible_suffix=0, # Show last 0 characters
mask_char="*", # Use * for masking
mask_short_values=False, # Truncate long messages only; keep short ones readable
)
def _common_add_cooldown_logic(

View file

@ -6,6 +6,7 @@ from litellm._logging import verbose_router_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.router_utils.add_retry_fallback_headers import (
add_fallback_headers_to_response,
get_fallback_error_info,
)
from litellm.types.router import LiteLLMParamsTypedDict
@ -90,6 +91,7 @@ async def run_async_fallback(
original_exception: Exception,
max_fallbacks: int,
fallback_depth: int,
include_fallback_errors: bool = False,
**kwargs,
) -> Any:
"""
@ -118,6 +120,7 @@ async def run_async_fallback(
raise original_exception
error_from_fallbacks = original_exception
fallback_errors = (get_fallback_error_info(original_exception),)
for mg in fallback_model_group:
if mg == original_model_group:
@ -136,6 +139,8 @@ async def run_async_fallback(
fallback_depth = fallback_depth + 1
kwargs["fallback_depth"] = fallback_depth
kwargs["max_fallbacks"] = max_fallbacks
if include_fallback_errors:
kwargs["include_fallback_errors"] = include_fallback_errors
response = await litellm_router.async_function_with_fallbacks(
*args, **kwargs
)
@ -143,6 +148,9 @@ async def run_async_fallback(
response = add_fallback_headers_to_response(
response=response,
attempted_fallbacks=fallback_depth,
fallback_errors=(
list(fallback_errors) if include_fallback_errors else None
),
)
# callback for successfull_fallback_event():
await log_success_fallback_event(
@ -153,6 +161,7 @@ async def run_async_fallback(
return response
except Exception as e:
error_from_fallbacks = e
fallback_errors = fallback_errors + (get_fallback_error_info(e),)
await log_failure_fallback_event(
original_model_group=original_model_group,
kwargs=kwargs,

View file

@ -68,7 +68,7 @@ async def acreate_sandbox(
provider: str,
template: str | None = None,
timeout: int | None = None,
allow_internet_access: bool = True,
allow_internet_access: bool | None = None,
api_key: str | None = None,
api_base: str | None = None,
**kwargs,

View file

@ -1,8 +1,28 @@
from typing import Iterable, List, Optional, Union
from __future__ import annotations
from dataclasses import dataclass
from typing import (
TYPE_CHECKING,
Any,
Callable,
Coroutine,
Iterable,
List,
Optional,
Union,
)
from pydantic import BaseModel, ConfigDict
from typing_extensions import Literal, Required, TypedDict
if TYPE_CHECKING:
import httpx
from aiohttp import ClientSession
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm import BaseConfig
from litellm.utils import CustomStreamWrapper, ModelResponse
class ChatCompletionSystemMessageParam(TypedDict, total=False):
content: Required[str]
@ -191,3 +211,44 @@ class CompletionRequest(BaseModel):
model_list: Optional[List[str]] = None
model_config = ConfigDict(protected_namespaces=(), extra="allow")
@dataclass(frozen=True, slots=True)
class _CompletionDispatchContext:
_azure_detection_model: str
acompletion: bool
api_base: Optional[str]
api_key: Optional[str]
api_version: Optional[str]
client: Any
custom_llm_provider: str
custom_prompt_dict: dict
extra_headers: Optional[dict]
headers: dict
hf_model_name: Optional[str]
kwargs: dict
litellm_params: dict
logger_fn: Optional[Callable]
logging: LiteLLMLoggingObj
max_retries: Optional[int]
max_tokens: Optional[int]
messages: list
metadata: Optional[dict]
model: str
model_response: ModelResponse
optional_params: dict
organization: Optional[str]
provider_config: Optional[BaseConfig]
shared_session: Optional[ClientSession]
stream: Optional[bool]
temperature: Optional[float]
text_completion: bool
timeout: Optional[Union[float, str, httpx.Timeout]]
top_p: Optional[float]
_CompletionDispatchResult = Union[
Coroutine[Any, Any, Union["ModelResponse", "CustomStreamWrapper"]],
"ModelResponse",
"CustomStreamWrapper",
]

View file

@ -2,6 +2,25 @@ from typing import Any, Dict, List, Optional
from pydantic import BaseModel, Field
CHAT_COMPLETION_AGENTIC_SURFACE = "chat_completions"
CODE_INTERPRETER_INTERCEPTION_PREFIX = "_code_interpreter_interception"
NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES = frozenset(
("_websearch_interception", "_compression_interception")
)
INTERCEPTION_INTERNAL_PREFIXES = frozenset(
(
*NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES,
CODE_INTERPRETER_INTERCEPTION_PREFIX,
)
)
def is_interception_internal_key(
key: str,
prefixes: frozenset[str] = INTERCEPTION_INTERNAL_PREFIXES,
) -> bool:
return any(key.startswith(prefix) for prefix in prefixes)
class StandardCustomLoggerInitParams(BaseModel):
"""

View file

@ -954,9 +954,6 @@ class Interaction(BaseModel):
None,
description="Output only. The time at which the response was last updated in ISO 8601 format\n(YYYY-MM-DDThh:mm:ssZ).",
)
role: Optional[str] = Field(
None, description="Output only. The role of the interaction."
)
outputs: Optional[List[Content]] = Field(
None, description="Output only. Responses from the model."
)
@ -1031,9 +1028,6 @@ class CreateModelInteractionParams(BaseModel):
None,
description="Output only. The time at which the response was last updated in ISO 8601 format\n(YYYY-MM-DDThh:mm:ssZ).",
)
role: Optional[str] = Field(
None, description="Output only. The role of the interaction."
)
outputs: Optional[List[Content]] = Field(
None, description="Output only. Responses from the model."
)
@ -1101,9 +1095,6 @@ class CreateAgentInteractionParams(BaseModel):
None,
description="Output only. The time at which the response was last updated in ISO 8601 format\n(YYYY-MM-DDThh:mm:ssZ).",
)
role: Optional[str] = Field(
None, description="Output only. The role of the interaction."
)
outputs: Optional[List[Content]] = Field(
None, description="Output only. Responses from the model."
)
@ -1323,7 +1314,6 @@ class InteractionsAPIResponse(BaseLiteLLMOpenAIResponseObject):
status: Optional[str] = None
created: Optional[str] = None
updated: Optional[str] = None
role: Optional[str] = None
# Legacy schema field (Api-Revision: 2026-05-07). Remove after June 8, 2026.
outputs: Optional[List[Dict[str, Any]]] = None
# New schema field (Api-Revision: 2026-05-20).
@ -1356,7 +1346,6 @@ class InteractionsAPIStreamingResponse(BaseLiteLLMOpenAIResponseObject):
status: Optional[str] = None
created: Optional[str] = None
updated: Optional[str] = None
role: Optional[str] = None
# Legacy schema field (Api-Revision: 2026-05-07). Remove after June 8, 2026.
outputs: Optional[List[Dict[str, Any]]] = None
# New schema field (Api-Revision: 2026-05-20).

View file

@ -69,7 +69,6 @@ class MCPPublicServer(BaseModel):
name: str
alias: Optional[str] = None
server_name: Optional[str] = None
url: Optional[str] = None
transport: MCPTransportType
spec_path: Optional[str] = None
auth_type: Optional[MCPAuthType] = None

View file

@ -1,7 +1,20 @@
from typing import Any, Dict, List, Literal, Optional, Union
from fastapi import HTTPException
from pydantic import BaseModel, ConfigDict, EmailStr, field_validator
from pydantic import (
BaseModel,
ConfigDict,
EmailStr,
Field,
field_validator,
model_serializer,
)
from pydantic_core.core_schema import SerializerFunctionWrapHandler
SCIM_ENTERPRISE_USER_SCHEMA = (
"urn:ietf:params:scim:schemas:extension:enterprise:2.0:User"
)
SCIM_ENTERPRISE_METADATA_KEY = "scim_enterprise"
class LiteLLM_UserScimMetadata(BaseModel):
@ -42,13 +55,49 @@ class SCIMUserGroup(BaseModel):
type: Optional[str] = "direct" # direct or indirect
class SCIMUserManager(BaseModel):
model_config = ConfigDict(populate_by_name=True)
value: Optional[str] = None
displayName: Optional[str] = None
ref: Optional[str] = Field(default=None, alias="$ref")
class SCIMEnterpriseUser(BaseModel):
model_config = ConfigDict(populate_by_name=True)
employeeNumber: Optional[str] = None
costCenter: Optional[str] = None
organization: Optional[str] = None
division: Optional[str] = None
department: Optional[str] = None
manager: Optional[SCIMUserManager] = None
class SCIMUser(SCIMResource):
model_config = ConfigDict(populate_by_name=True)
userName: Optional[str] = None
name: Optional[SCIMUserName] = None
displayName: Optional[str] = None
active: bool = True
emails: Optional[List[SCIMUserEmail]] = None
groups: Optional[List[SCIMUserGroup]] = None
enterprise_user: Optional[SCIMEnterpriseUser] = Field(
default=None,
alias=SCIM_ENTERPRISE_USER_SCHEMA,
serialization_alias=SCIM_ENTERPRISE_USER_SCHEMA,
)
@model_serializer(mode="wrap")
def _omit_absent_enterprise(
self, handler: SerializerFunctionWrapHandler
) -> Dict[str, Any]:
dumped = handler(self)
if self.enterprise_user is None:
dumped.pop(SCIM_ENTERPRISE_USER_SCHEMA, None)
dumped.pop("enterprise_user", None)
return dumped
class SCIMMember(BaseModel):

View file

@ -37,6 +37,8 @@ from pydantic import (
ConfigDict,
Field,
PrivateAttr,
SkipValidation,
field_serializer,
field_validator,
)
from typing_extensions import Required, TypedDict
@ -3146,10 +3148,41 @@ class CustomPricingLiteLLMParams(BaseModel):
search_context_cost_per_query: Optional[Dict[str, Any]] = None
citation_cost_per_token: Optional[float] = None
tiered_pricing: Optional[List[Dict[str, Any]]] = None
cache_read_input_token_cost_above_272k_tokens: Optional[float] = None
cache_read_input_token_cost_above_512k_tokens: Optional[float] = None
input_cost_per_image_token: Optional[float] = None
input_cost_per_token_above_272k_tokens: Optional[float] = None
input_cost_per_token_above_512k_tokens: Optional[float] = None
output_cost_per_token_above_272k_tokens: Optional[float] = None
output_cost_per_token_above_512k_tokens: Optional[float] = None
output_vector_size: Optional[int] = None
ocr_cost_per_page: Optional[float] = None
ocr_cost_per_credit: Optional[float] = None
annotation_cost_per_page: Optional[float] = None
regional_processing_uplift_multiplier_eu: Optional[float] = None
regional_processing_uplift_multiplier_us: Optional[float] = None
# Server-controlled fields that bound or drive an interceptor's agentic loop
# (depth, cycle fingerprints, ceiling, code-interpreter sandbox state). Listed
# in all_litellm_params so they are treated as LiteLLM-level and excluded from
# get_non_default_completion_params; otherwise the OpenAI param builder sweeps
# any unrecognized top-level key into extra_body and leaks them to the provider.
# This is what lets the loop carry state across rerun calls without a provider
# scrubber.
agentic_loop_internal_litellm_params = [
"_agentic_loop_depth",
"_agentic_loop_fingerprints",
"_agentic_loop_api_surface",
"max_agentic_loops",
"_code_interpreter_interception_active",
"_code_interpreter_interception_sandbox_key",
"_code_interpreter_interception_converted_stream",
]
all_litellm_params = (
[
agentic_loop_internal_litellm_params
+ [
"metadata",
"litellm_metadata",
"litellm_trace_id",
@ -3448,6 +3481,7 @@ class LlmProviders(str, Enum):
TENSORMESH = "tensormesh"
LIBERTAI = "libertai"
PINSTRIPES = "pinstripes"
DARKBLOOM = "darkbloom"
LITELLM_AGENT = "litellm_agent"
CURSOR = "cursor"
BEDROCK_MANTLE = "bedrock_mantle"
@ -3505,6 +3539,7 @@ class SandboxProviders(str, Enum):
"""
E2B = "e2b"
OPENSANDBOX = "opensandbox"
class LiteLLMLoggingBaseClass:
@ -3610,10 +3645,20 @@ class LiteLLMBatch(Batch):
class LiteLLMRealtimeStreamLoggingObject(LiteLLMPydanticObjectBase):
results: OpenAIRealtimeStreamList
# Events are already well-formed provider dicts. Validating them against the
# OpenAIRealtimeEvents union makes Pydantic try every member per event, which
# floods thousands of ValidationErrors for events outside the union (e.g.
# rate_limits.updated), blocks the event loop, and discards the session usage.
results: SkipValidation[OpenAIRealtimeStreamList]
usage: Usage
_hidden_params: dict = {}
@field_serializer("results")
def _serialize_results(
self, results: OpenAIRealtimeStreamList
) -> List[Dict[str, Any]]:
return [dict(event) for event in results]
def __contains__(self, key):
# Define custom behavior for the 'in' operator
return hasattr(self, key)

View file

@ -3191,7 +3191,7 @@ def get_optional_params_transcription(
model=model,
drop_params=drop_params if drop_params is not None else False,
)
elif provider_config is not None: # handles fireworks ai, and any future providers
elif provider_config is not None: # custom audio transcription config
supported_params = provider_config.get_supported_openai_params(model=model)
_check_valid_arg(supported_params=supported_params)
optional_params = provider_config.map_openai_params(
@ -8915,8 +8915,6 @@ class ProviderConfigManager:
)
return AzureSpeechAudioTranscriptionConfig()
if litellm.LlmProviders.FIREWORKS_AI == provider:
return litellm.FireworksAIAudioTranscriptionConfig()
elif litellm.LlmProviders.DEEPGRAM == provider:
return litellm.DeepgramAudioTranscriptionConfig()
elif litellm.LlmProviders.ELEVENLABS == provider:
@ -9733,9 +9731,14 @@ class ProviderConfigManager:
Get sandbox (code execution) configuration for a given provider.
"""
from litellm.llms.e2b.sandbox.transformation import E2BSandboxConfig
from litellm.llms.opensandbox.sandbox.transformation import (
OpenSandboxSandboxConfig,
)
if provider == SandboxProviders.E2B:
return E2BSandboxConfig()
if provider == SandboxProviders.OPENSANDBOX:
return OpenSandboxSandboxConfig()
return None
@staticmethod

View file

@ -571,7 +571,7 @@
"output_vector_size": 1536
},
"amazon.titan-embed-text-v2:0": {
"input_cost_per_token": 2e-07,
"input_cost_per_token": 2e-08,
"litellm_provider": "bedrock",
"max_input_tokens": 8192,
"max_tokens": 8192,
@ -10684,6 +10684,268 @@
"mode": "chat",
"output_cost_per_token": 1.923e-06
},
"cloudflare/@cf/openai/gpt-oss-120b": {
"input_cost_per_token": 3.5e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 7.5e-07,
"supports_function_calling": true,
"supports_reasoning": true
},
"cloudflare/@cf/google/gemma-2b-it-lora": {
"input_cost_per_token": 0.0,
"litellm_provider": "cloudflare",
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 0.0
},
"cloudflare/@cf/meta/llama-3.2-3b-instruct": {
"input_cost_per_token": 5.09e-08,
"litellm_provider": "cloudflare",
"max_input_tokens": 80000,
"max_output_tokens": 80000,
"max_tokens": 80000,
"mode": "chat",
"output_cost_per_token": 3.35e-07
},
"cloudflare/@cf/meta/llama-guard-3-8b": {
"input_cost_per_token": 4.84e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 3e-08
},
"cloudflare/@cf/mistral/mistral-7b-instruct-v0.2-lora": {
"input_cost_per_token": 0.0,
"litellm_provider": "cloudflare",
"max_input_tokens": 15000,
"max_output_tokens": 15000,
"max_tokens": 15000,
"mode": "chat",
"output_cost_per_token": 0.0
},
"cloudflare/@cf/moonshotai/kimi-k2.7-code": {
"cache_read_input_token_cost": 1.9e-07,
"input_cost_per_token": 9.5e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 4e-06,
"supports_function_calling": true,
"supports_reasoning": true
},
"cloudflare/@cf/deepseek-ai/deepseek-r1-distill-qwen-32b": {
"input_cost_per_token": 4.97e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 80000,
"max_output_tokens": 80000,
"max_tokens": 80000,
"mode": "chat",
"output_cost_per_token": 4.881e-06,
"supports_reasoning": true
},
"cloudflare/@cf/meta/llama-3.1-8b-instruct-fp8": {
"input_cost_per_token": 1.52e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 32000,
"max_output_tokens": 32000,
"max_tokens": 32000,
"mode": "chat",
"output_cost_per_token": 2.87e-07
},
"cloudflare/@cf/meta/llama-3.2-1b-instruct": {
"input_cost_per_token": 2.7e-08,
"litellm_provider": "cloudflare",
"max_input_tokens": 60000,
"max_output_tokens": 60000,
"max_tokens": 60000,
"mode": "chat",
"output_cost_per_token": 2.01e-07
},
"cloudflare/@cf/moonshotai/kimi-k2.6": {
"cache_read_input_token_cost": 1.6e-07,
"input_cost_per_token": 9.5e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 4e-06,
"supports_function_calling": true,
"supports_reasoning": true
},
"cloudflare/@cf/zai-org/glm-4.7-flash": {
"input_cost_per_token": 6.05e-08,
"litellm_provider": "cloudflare",
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4e-07,
"supports_function_calling": true,
"supports_reasoning": true
},
"cloudflare/@cf/meta-llama/llama-2-7b-chat-hf-lora": {
"input_cost_per_token": 0.0,
"litellm_provider": "cloudflare",
"max_input_tokens": 8192,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 0.0
},
"cloudflare/@cf/meta/llama-3.3-70b-instruct-fp8-fast": {
"input_cost_per_token": 2.93e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 24000,
"max_output_tokens": 24000,
"max_tokens": 24000,
"mode": "chat",
"output_cost_per_token": 2.253e-06,
"supports_function_calling": true
},
"cloudflare/@cf/ibm-granite/granite-4.0-h-micro": {
"input_cost_per_token": 1.7e-08,
"litellm_provider": "cloudflare",
"max_input_tokens": 131000,
"max_output_tokens": 131000,
"max_tokens": 131000,
"mode": "chat",
"output_cost_per_token": 1.12e-07,
"supports_function_calling": true
},
"cloudflare/@cf/qwen/qwen2.5-coder-32b-instruct": {
"input_cost_per_token": 6.6e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 32768,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 1e-06
},
"cloudflare/@cf/zai-org/glm-5.2": {
"cache_read_input_token_cost": 2.6e-07,
"input_cost_per_token": 1.4e-06,
"litellm_provider": "cloudflare",
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 4.4e-06,
"supports_function_calling": true,
"supports_reasoning": true
},
"cloudflare/@cf/nvidia/nemotron-3-120b-a12b": {
"input_cost_per_token": 5e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 256000,
"max_output_tokens": 256000,
"max_tokens": 256000,
"mode": "chat",
"output_cost_per_token": 1.5e-06,
"supports_function_calling": true,
"supports_reasoning": true
},
"cloudflare/@cf/aisingapore/gemma-sea-lion-v4-27b-it": {
"input_cost_per_token": 3.51e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 5.55e-07
},
"cloudflare/@cf/qwen/qwen3-30b-a3b-fp8": {
"input_cost_per_token": 5.09e-08,
"litellm_provider": "cloudflare",
"max_input_tokens": 32768,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 3.35e-07,
"supports_function_calling": true,
"supports_reasoning": true
},
"cloudflare/@cf/google/gemma-7b-it-lora": {
"input_cost_per_token": 0.0,
"litellm_provider": "cloudflare",
"max_input_tokens": 3500,
"max_output_tokens": 3500,
"max_tokens": 3500,
"mode": "chat",
"output_cost_per_token": 0.0
},
"cloudflare/@cf/google/gemma-4-26b-a4b-it": {
"input_cost_per_token": 1e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 256000,
"max_output_tokens": 256000,
"max_tokens": 256000,
"mode": "chat",
"output_cost_per_token": 3e-07,
"supports_function_calling": true,
"supports_reasoning": true
},
"cloudflare/@cf/mistralai/mistral-small-3.1-24b-instruct": {
"input_cost_per_token": 3.51e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 5.55e-07,
"supports_function_calling": true
},
"cloudflare/@cf/meta/llama-3.2-11b-vision-instruct": {
"input_cost_per_token": 4.85e-08,
"litellm_provider": "cloudflare",
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 6.76e-07,
"supports_vision": true
},
"cloudflare/@cf/openai/gpt-oss-20b": {
"input_cost_per_token": 2e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 3e-07,
"supports_function_calling": true,
"supports_reasoning": true
},
"cloudflare/@cf/meta/llama-4-scout-17b-16e-instruct": {
"input_cost_per_token": 2.7e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 131000,
"max_output_tokens": 131000,
"max_tokens": 131000,
"mode": "chat",
"output_cost_per_token": 8.5e-07,
"supports_function_calling": true
},
"cloudflare/@cf/qwen/qwq-32b": {
"input_cost_per_token": 6.6e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 24000,
"max_output_tokens": 24000,
"max_tokens": 24000,
"mode": "chat",
"output_cost_per_token": 1e-06,
"supports_reasoning": true
},
"codestral/codestral-2405": {
"input_cost_per_token": 0.0,
"litellm_provider": "codestral",
@ -39946,24 +40208,6 @@
"litellm_provider": "fireworks_ai",
"mode": "chat"
},
"fireworks_ai/accounts/fireworks/models/whisper-v3": {
"max_tokens": 4096,
"max_input_tokens": 4096,
"max_output_tokens": 4096,
"input_cost_per_token": 0.0,
"output_cost_per_token": 0.0,
"litellm_provider": "fireworks_ai",
"mode": "audio_transcription"
},
"fireworks_ai/accounts/fireworks/models/whisper-v3-turbo": {
"max_tokens": 4096,
"max_input_tokens": 4096,
"max_output_tokens": 4096,
"input_cost_per_token": 0.0,
"output_cost_per_token": 0.0,
"litellm_provider": "fireworks_ai",
"mode": "audio_transcription"
},
"fireworks_ai/accounts/fireworks/models/yi-34b": {
"max_tokens": 4096,
"max_input_tokens": 4096,
@ -43543,5 +43787,39 @@
"supports_assistant_prefill": true,
"supports_reasoning": false,
"source": "https://pinstripes.io/pricing"
},
"darkbloom/gemma-4-26b": {
"input_cost_per_token": 3e-08,
"litellm_provider": "darkbloom",
"max_input_tokens": 131072,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 1.65e-07,
"source": "https://www.darkbloom.dev/",
"supported_endpoints": [
"/v1/chat/completions"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_system_messages": true,
"supports_tool_choice": true
},
"darkbloom/gpt-oss-20b": {
"input_cost_per_token": 1.45e-08,
"litellm_provider": "darkbloom",
"max_input_tokens": 131072,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 7e-08,
"source": "https://www.darkbloom.dev/",
"supported_endpoints": [
"/v1/chat/completions"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_system_messages": true,
"supports_tool_choice": true
}
}

View file

@ -1833,6 +1833,23 @@
"text_completion": true
}
},
"opensandbox": {
"display_name": "OpenSandbox (`opensandbox`)",
"url": "https://open-sandbox.ai/api/",
"endpoints": {
"chat_completions": false,
"messages": false,
"responses": false,
"embeddings": false,
"image_generations": false,
"audio_transcriptions": false,
"audio_speech": false,
"moderations": false,
"batches": false,
"rerank": false,
"sandbox": true
}
},
"openai_like": {
"display_name": "OpenAI-like (`openai_like`)",
"url": "https://docs.litellm.ai/docs/providers/openai_compatible",
@ -2008,6 +2025,23 @@
"interactions": true
}
},
"darkbloom": {
"display_name": "Darkbloom (`darkbloom`)",
"url": "https://docs.litellm.ai/docs/providers/darkbloom",
"endpoints": {
"chat_completions": true,
"messages": false,
"responses": false,
"embeddings": false,
"image_generations": false,
"audio_transcriptions": false,
"audio_speech": false,
"moderations": false,
"batches": false,
"rerank": false,
"a2a": false
}
},
"predibase": {
"display_name": "Predibase (`predibase`)",
"url": "https://docs.litellm.ai/docs/providers/predibase",

View file

@ -70,6 +70,7 @@ proxy = [
"soundfile>=0.12.1,<1.0",
"pyroscope-io>=0.8.16,<1.0; sys_platform != 'win32'",
"pydantic-settings>=2.14.1,<3.0",
"expression>=5.6.0,<6.0",
]
# Thin client install for the `lite` CLI on developer laptops. The CLI's heavy
# imports (fastapi, cryptography, ...) are all guarded, so it runs on the base

View file

@ -300,7 +300,7 @@
"slack": 3
},
"RET504": {
"baseline": 709,
"baseline": 702,
"slack": 20
},
"RUF010": {

View file

@ -1,21 +1,22 @@
#!/usr/bin/env python3
"""Per-rule count gate for basedpyright.
"""Delta-vs-base per-rule gate for basedpyright.
basedpyright's ``--outputjson`` is reduced to a count of errors per *rule*
(``reportAny``, ``reportArgumentType``, ...) and checked against a committed
budget of the form ``{rule: {baseline, slack}}``, the same shape as
``ruff-strict-budget.json``. A rule fails when its codebase-wide total exceeds
``baseline + slack``. Counts ignore file, line, and column, so a violation
moving anywhere in the tree is invisible; only the per-rule total moves the
needle.
``ruff-strict-budget.json``. A rule fails only when its codebase-wide total is
both over its ceiling (``baseline + slack``) *and* higher than the count on the
base it merges into, so a change is blamed for the errors it adds, never for
drift that already sits in the base. That ``> base`` guard is what stops an
unrelated PR from inheriting a red once two PRs each land near the ceiling and
their sum crosses it: the bystander's count equals its base, so it is spared,
while any PR that actually grows the rule past the cap still fails.
Unlike ``ruff_strict_gate.py`` this does *not* re-run the tool on the merge base
to compute a delta: a second basedpyright pass is minutes and gigabytes, whereas
ruff is milliseconds. The committed budget is the baseline instead -- exactly
how the previous per-file gate worked -- so keep it fresh with ``--update``
(ratchet), which re-captures every rule's count from the current tree while
preserving each rule's slack. Tool output is read from stdin, so the caller
decides how to invoke basedpyright (and from which cwd).
Head counts are read from stdin (the caller runs basedpyright once and pipes
``--outputjson`` in); the base count is a second basedpyright pass over a
detached worktree at the merge-base, run under the same environment so import
resolution matches. ``--update`` re-captures the absolute per-rule baselines for
the ratchet, preserving each rule's slack.
``--outputjson`` is used rather than text diagnostics because the latter wrap
across lines, leaving the ``(reportRule)`` on a continuation line away from the
@ -24,13 +25,21 @@ carries an unambiguous ``rule`` field.
"""
import argparse
import contextlib
import json
import shutil
import subprocess
import sys
import tempfile
from collections import Counter
from collections.abc import Iterator, Mapping
from pathlib import Path
from typing import Mapping, NamedTuple
from typing import NamedTuple
REPO_ROOT = Path(__file__).resolve().parent.parent
BUDGET_PATH = REPO_ROOT / "basedpyright-code-budget.json"
PYRIGHT_CONFIG = REPO_ROOT / "pyrightconfig.json"
DEFAULT_BASE = "origin/litellm_internal_staging"
# Bucket for a basedpyright diagnostic with no `rule`. Counted so it's gated.
UNCODED = "<uncoded>"
@ -45,6 +54,7 @@ class Breach(NamedTuple):
code: str
total: int
cap: int
added: int
def _seed_slack(baseline: int) -> int:
@ -54,18 +64,19 @@ def _seed_slack(baseline: int) -> int:
return 10 if baseline >= 50 else 3
def _to_repo_relative(raw: str) -> str | None:
def _to_relative(raw: str, root: Path) -> str | None:
path = Path(raw)
absolute = path if path.is_absolute() else Path.cwd() / path
absolute = path if path.is_absolute() else root / path
try:
return absolute.resolve().relative_to(REPO_ROOT).as_posix()
return absolute.resolve().relative_to(root).as_posix()
except ValueError:
return None
def count_basedpyright(payload: str) -> dict[str, int]:
"""Count in-repo basedpyright errors per rule from `--outputjson`. Warnings
and information are ignored; only `severity == "error"` is gated."""
def count_basedpyright(payload: str, root: Path = REPO_ROOT) -> dict[str, int]:
"""Count in-tree basedpyright errors per rule from `--outputjson`. Warnings
and information are ignored; only `severity == "error"` is gated. Files
outside `root` (the venv's site-packages, say) are dropped."""
try:
data = json.loads(payload or "{}")
except json.JSONDecodeError as exc:
@ -79,21 +90,62 @@ def count_basedpyright(payload: str) -> dict[str, int]:
for diag in data.get("generalDiagnostics", []):
if diag.get("severity") != "error":
continue
if _to_repo_relative(diag.get("file", "")) is None:
if _to_relative(diag.get("file", ""), root) is None:
continue
counts[diag.get("rule") or UNCODED] += 1
return dict(counts)
def _run(cmd: list[str], cwd: Path = REPO_ROOT) -> str:
proc = subprocess.run(cmd, cwd=cwd, capture_output=True, text=True)
if proc.returncode not in (0, 1):
sys.stderr.write(proc.stderr)
raise SystemExit(f"{cmd[0]} exited {proc.returncode}")
return proc.stdout
@contextlib.contextmanager
def _temp_worktree(ref: str) -> Iterator[Path]:
parent = Path(tempfile.mkdtemp(prefix="bpr_base_"))
worktree = parent / "wt"
try:
_run(["git", "worktree", "add", "--detach", str(worktree), ref])
yield worktree
finally:
subprocess.run(
["git", "worktree", "remove", "--force", str(worktree)],
cwd=REPO_ROOT,
capture_output=True,
text=True,
)
shutil.rmtree(parent, ignore_errors=True)
def base_counts(ref: str) -> dict[str, int]:
"""basedpyright error counts per rule for the merge-base tree. The head
config is copied in so the base is judged by today's rules, and the run uses
the head environment's basedpyright (on PATH) so imports resolve the same."""
exe = shutil.which("basedpyright") or "basedpyright"
with _temp_worktree(ref) as worktree:
shutil.copy(PYRIGHT_CONFIG, worktree / "pyrightconfig.json")
proc = subprocess.run(
[exe, "--outputjson"], cwd=worktree, capture_output=True, text=True
)
return count_basedpyright(proc.stdout, root=worktree)
def evaluate(
counts: Mapping[str, int], budget: Mapping[str, Mapping[str, int]]
head: Mapping[str, int],
base: Mapping[str, int],
budget: Mapping[str, Mapping[str, int]],
) -> list[Breach]:
breaches = []
for code, total in counts.items():
for code, total in head.items():
spec = budget.get(code)
cap = spec["baseline"] + spec["slack"] if spec else DEFAULT_SLACK
if total > cap:
breaches.append(Breach(code, total, cap))
prior = base.get(code, 0)
if total > cap and total > prior:
breaches.append(Breach(code, total, cap, total - prior))
return sorted(breaches)
@ -107,9 +159,6 @@ def is_vacuous_run(
return not counts and any(spec["baseline"] for spec in budget.values())
BUDGET_PATH = REPO_ROOT / "basedpyright-code-budget.json"
def cmd_update(counts: Mapping[str, int]) -> None:
existing = json.loads(BUDGET_PATH.read_text()) if BUDGET_PATH.exists() else {}
budget = {
@ -127,9 +176,10 @@ def cmd_update(counts: Mapping[str, int]) -> None:
)
def cmd_check(counts: Mapping[str, int]) -> None:
def cmd_check(base_ref: str) -> None:
budget = json.loads(BUDGET_PATH.read_text())
if is_vacuous_run(counts, budget):
head = count_basedpyright(sys.stdin.read())
if is_vacuous_run(head, budget):
expected = sum(spec["baseline"] for spec in budget.values())
print(
f"FAIL: basedpyright produced no errors, but {BUDGET_PATH.name} expects "
@ -137,27 +187,44 @@ def cmd_check(counts: Mapping[str, int]) -> None:
f"nothing; refusing to certify a vacuous run."
)
raise SystemExit(1)
breaches = evaluate(counts, budget)
base_point = _run(["git", "merge-base", base_ref, "HEAD"]).strip() or base_ref
base = base_counts(base_point)
if is_vacuous_run(base, budget):
print(
f"FAIL: basedpyright produced no errors for the base tree at "
f"{base_point[:12]}, so every rule would look freshly added. The base "
f"pass almost certainly crashed; refusing to blame this change for it."
)
raise SystemExit(1)
breaches = evaluate(head, base, budget)
if not breaches:
print(
f"OK: every rule is within its basedpyright ceiling ({sum(counts.values())} errors total)"
f"OK: every rule is within its basedpyright ceiling or no higher than base ({sum(head.values())} errors total)"
)
return
print("FAIL: basedpyright errors exceed the per-rule ceiling:")
for breach in breaches:
print(f" {breach.code}: {breach.total} errors over cap {breach.cap}")
print(
f" {breach.code}: total {breach.total} over cap {breach.cap} (this change added {breach.added})"
)
print(
"Resolve the new errors, or run 'make lint-basedpyright-budget-update' if the ceiling should move."
"Reduce the new errors or remove an equal number elsewhere; the ceiling is "
"baseline + slack in basedpyright-code-budget.json."
)
summary = "; ".join(f"{b.code} {b.total}/{b.cap} (+{b.added})" for b in breaches)
print(f"BREACHED RULES: {summary}")
raise SystemExit(1)
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--base", default=DEFAULT_BASE)
parser.add_argument("--update", action="store_true")
args = parser.parse_args()
counts = count_basedpyright(sys.stdin.read())
cmd_update(counts) if args.update else cmd_check(counts)
if args.update:
cmd_update(count_basedpyright(sys.stdin.read()))
else:
cmd_check(args.base)
if __name__ == "__main__":

View file

@ -2502,19 +2502,34 @@ async def test_bedrock_image_url_sync_client():
mock_post.assert_called_once()
def test_bedrock_error_handling_streaming():
@pytest.mark.parametrize(
"exception_type, expected_status_code",
[
("internalServerException", 500),
("serviceUnavailableException", 503),
("modelTimeoutException", 408),
("modelStreamErrorException", 424),
("validationException", 400),
],
)
def test_bedrock_error_handling_streaming(exception_type, expected_status_code):
"""Bedrock event-stream error events arrive with botocore's hard-coded
status_code=400; the decoder must surface the modeled HTTP status instead
(e.g. internalServerException -> 500). For 5xx this is what makes the error
retryable downstream; for all types it replaces the misleading 400 with the
true code. Regression for #24608."""
from litellm.llms.bedrock.chat.invoke_handler import (
AWSEventStreamDecoder,
BedrockError,
)
from unittest.mock import patch, Mock
from unittest.mock import Mock
event = Mock()
event.to_response_dict = Mock(
return_value={
"status_code": 400,
"headers": {
":exception-type": "serviceUnavailableException",
":exception-type": exception_type,
":content-type": "application/json",
":message-type": "exception",
},
@ -2525,11 +2540,10 @@ def test_bedrock_error_handling_streaming():
decoder = AWSEventStreamDecoder(
model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0"
)
with pytest.raises(Exception) as e:
with pytest.raises(BedrockError) as e:
decoder._parse_message_from_event(event)
assert isinstance(e.value, BedrockError)
assert "Bedrock is unable to process your request." in e.value.message
assert e.value.status_code == 400
assert e.value.status_code == expected_status_code
@pytest.mark.parametrize(

View file

@ -0,0 +1,34 @@
"""
Tests for AWS Bedrock embedding model pricing in the model cost map.
Regression test for the Amazon Titan Text Embeddings V2 commercial price,
which was previously set 10x too high (2e-07 instead of 2e-08).
AWS lists Titan Text Embeddings V2 at $0.02 per 1M input tokens
(= $0.00002 per 1K tokens = 2e-08 per token).
"""
import importlib
class TestBedrockEmbeddingPricing:
"""Test suite for Bedrock embedding model pricing in the cost map."""
def test_titan_embed_v2_commercial_input_cost(self, monkeypatch):
"""Titan Text Embeddings V2 should be priced at $0.02 / 1M tokens (2e-08)."""
# Scope the local-cost-map flag to this test only, so it does not leak
# into sibling tests. monkeypatch restores the environment on teardown.
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
import litellm.litellm_core_utils.get_model_cost_map
import litellm
# Reload so the cost map is re-read from the local file with the flag set.
importlib.reload(litellm.litellm_core_utils.get_model_cost_map)
importlib.reload(litellm)
model = litellm.model_cost["amazon.titan-embed-text-v2:0"]
assert model["input_cost_per_token"] == 2e-08
assert model["output_cost_per_token"] == 0.0
assert model["litellm_provider"] == "bedrock"
assert model["mode"] == "embedding"

View file

@ -9,9 +9,7 @@ import pytest
from litellm import acompletion, completion
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
FAKE_API_BASE = (
"https://fake-cloudflare.example.com/client/v4/accounts/fake-acct/ai/run/"
)
FAKE_API_BASE = "https://fake-cloudflare.example.com/client/v4/accounts/fake-acct/ai/v1"
FAKE_API_KEY = "fake-cf-api-key"
@ -26,28 +24,78 @@ def _make_mock_response(json_data: Dict[str, Any]) -> MagicMock:
def _chat_response() -> Dict[str, Any]:
return {
"result": {
"response": "I am a large language model created to assist you.",
},
"success": True,
"errors": [],
"messages": [],
"id": "chatcmpl-cf",
"object": "chat.completion",
"created": 1234567890,
"model": "@cf/meta/llama-2-7b-chat-int8",
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": "I am a large language model created to assist you.",
},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 8, "completion_tokens": 11, "total_tokens": 19},
}
def _tool_call_response() -> Dict[str, Any]:
return {
"id": "chatcmpl-cf-tools",
"object": "chat.completion",
"created": 1234567890,
"model": "@cf/meta/llama-2-7b-chat-int8",
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {
"name": "get_weather",
"arguments": '{"city": "New York"}',
},
}
],
},
"finish_reason": "tool_calls",
}
],
"usage": {"prompt_tokens": 20, "completion_tokens": 9, "total_tokens": 29},
}
def _streaming_chunks() -> list[str]:
base = {
"id": "chatcmpl-cf",
"object": "chat.completion.chunk",
"created": 1234567890,
"model": "@cf/meta/llama-2-7b-chat-int8",
}
return [
json.dumps({"response": "I am"}),
json.dumps({"response": " a language"}),
json.dumps({"response": " model."}),
]
def _streaming_chunks_response_text() -> list[str]:
return [
json.dumps({"response_text": "I am"}),
json.dumps({"response_text": " a language"}),
json.dumps({"response_text": " model."}),
json.dumps({**base, "choices": [{"index": 0, "delta": {"content": "I am"}}]}),
json.dumps(
{**base, "choices": [{"index": 0, "delta": {"content": " a language"}}]}
),
json.dumps(
{
**base,
"choices": [
{
"index": 0,
"delta": {"content": " model."},
"finish_reason": "stop",
}
],
}
),
]
@ -85,6 +133,48 @@ def test_completion_cloudflare(sync_mode):
assert response.choices[0].message.content is not None
assert "language model" in response.choices[0].message.content.lower()
called_url = mock_post.call_args.kwargs.get("url") or mock_post.call_args.args[0]
assert called_url.endswith("/ai/v1/chat/completions")
assert "/ai/run/" not in called_url
def test_completion_cloudflare_tool_calls_sent_to_openai_endpoint():
messages = [{"role": "user", "content": "weather in New York?"}]
tools = [
{
"type": "function",
"function": {
"name": "get_weather",
"parameters": {
"type": "object",
"properties": {"city": {"type": "string"}},
"required": ["city"],
},
},
}
]
mock_resp = _make_mock_response(_tool_call_response())
with patch.object(HTTPHandler, "post", return_value=mock_resp) as mock_post:
response = completion(
model="cloudflare/@cf/meta/llama-2-7b-chat-int8",
messages=messages,
tools=tools,
tool_choice="auto",
api_base=FAKE_API_BASE,
api_key=FAKE_API_KEY,
)
mock_post.assert_called_once()
sent_body = json.loads(mock_post.call_args.kwargs["data"])
assert sent_body["tools"] == tools
assert sent_body["tool_choice"] == "auto"
assert response.choices[0].finish_reason == "tool_calls"
tool_calls = response.choices[0].message.tool_calls
assert tool_calls is not None and len(tool_calls) == 1
assert tool_calls[0].function.name == "get_weather"
@pytest.mark.parametrize("sync_mode", [True, False])
def test_completion_cloudflare_stream(sync_mode):
@ -153,76 +243,3 @@ def test_completion_cloudflare_stream(sync_mode):
if c.choices[0].delta.content
)
assert "language" in content.lower()
@pytest.mark.parametrize("sync_mode", [True, False])
def test_completion_cloudflare_stream_response_text(sync_mode):
"""Newer Cloudflare Workers AI models (e.g. Nemotron) emit `response_text`
instead of `response` in streamed chunks. The iterator must surface that
text so streaming output is not silently empty.
"""
messages = [{"role": "user", "content": "what llm are you"}]
raw_chunks = _streaming_chunks_response_text()
if sync_mode:
def _iter_lines():
for chunk in raw_chunks:
yield f"data: {chunk}"
yield "data: [DONE]"
mock_resp = MagicMock()
mock_resp.iter_lines.return_value = _iter_lines()
mock_resp.status_code = 200
mock_resp.headers = {"content-type": "text/event-stream"}
with patch.object(HTTPHandler, "post", return_value=mock_resp) as mock_post:
response = completion(
model="cloudflare/@cf/nvidia/nemotron-mini-4b-instruct",
messages=messages,
max_tokens=15,
stream=True,
api_base=FAKE_API_BASE,
api_key=FAKE_API_KEY,
)
chunks_received = list(response)
mock_post.assert_called_once()
else:
async def _aiter_lines():
for chunk in raw_chunks:
yield f"data: {chunk}"
yield "data: [DONE]"
mock_resp = MagicMock()
mock_resp.aiter_lines.return_value = _aiter_lines()
mock_resp.status_code = 200
mock_resp.headers = {"content-type": "text/event-stream"}
async def _run():
with patch.object(
AsyncHTTPHandler, "post", new_callable=AsyncMock, return_value=mock_resp
) as mock_post:
resp = await acompletion(
model="cloudflare/@cf/nvidia/nemotron-mini-4b-instruct",
messages=messages,
max_tokens=15,
stream=True,
api_base=FAKE_API_BASE,
api_key=FAKE_API_KEY,
)
received = []
async for chunk in resp:
received.append(chunk)
mock_post.assert_called_once()
return received
chunks_received = asyncio.run(_run())
assert len(chunks_received) > 0
content = "".join(
c.choices[0].delta.content
for c in chunks_received
if c.choices[0].delta.content
)
assert "language" in content.lower()

View file

@ -7,9 +7,10 @@ sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
import litellm
from litellm import transcription
from litellm.litellm_core_utils.get_supported_openai_params import (
get_supported_openai_params,
)
from litellm.llms.fireworks_ai.chat.transformation import FireworksAIConfig
from base_audio_transcription_unit_tests import BaseLLMAudioTranscriptionTest
fireworks = FireworksAIConfig()
@ -69,74 +70,16 @@ def test_map_response_format():
assert result == {"response_format": response_format}
_AUDIO_FILE_PATH = os.path.join(
os.path.dirname(os.path.realpath(__file__)), "gettysburg.wav"
)
class TestFireworksAIAudioTranscription(BaseLLMAudioTranscriptionTest):
def get_base_audio_transcription_call_args(self) -> dict:
return {
"model": "fireworks_ai/whisper-v3",
"api_base": "https://audio-prod.api.fireworks.ai/v1",
}
def get_custom_llm_provider(self) -> litellm.LlmProviders:
return litellm.LlmProviders.FIREWORKS_AI
def test_audio_transcription(self):
from unittest.mock import MagicMock
from openai.types.audio import Transcription
audio_file = open(_AUDIO_FILE_PATH, "rb")
mock_client = MagicMock()
mock_client.audio.transcriptions.create.return_value = Transcription(
text="four score and seven years ago"
)
transcript = transcription(
**self.get_base_audio_transcription_call_args(),
file=audio_file,
api_key="fw-test-key",
client=mock_client,
)
assert transcript.text == "four score and seven years ago"
sent = mock_client.audio.transcriptions.create.call_args.kwargs
assert sent["model"] == "whisper-v3"
assert sent["file"] is audio_file
@pytest.mark.asyncio
async def test_audio_transcription_async(self):
from unittest.mock import AsyncMock, MagicMock
from openai.types.audio import Transcription
audio_file = open(_AUDIO_FILE_PATH, "rb")
raw_response = MagicMock()
raw_response.headers = {}
raw_response.parse.return_value = Transcription(
text="four score and seven years ago"
)
mock_client = MagicMock()
mock_client.audio.transcriptions.with_raw_response.create = AsyncMock(
return_value=raw_response
)
transcript = await litellm.atranscription(
**self.get_base_audio_transcription_call_args(),
file=audio_file,
api_key="fw-test-key",
client=mock_client,
)
assert transcript.text == "four score and seven years ago"
sent = (
mock_client.audio.transcriptions.with_raw_response.create.call_args.kwargs
)
assert sent["model"] == "whisper-v3"
assert sent["file"] is audio_file
def test_get_supported_openai_params_transcription_returns_none():
# Fireworks AI deprecated audio transcription on 2026-06-10; the endpoint
# is decommissioned. Returning None (not chat-completion params) signals
# to callers that transcription is unsupported for this provider.
result = get_supported_openai_params(
model="fireworks_ai/accounts/fireworks/models/whisper-v3",
custom_llm_provider="fireworks_ai",
request_type="transcription",
)
assert result is None
@pytest.mark.parametrize(

View file

@ -605,8 +605,33 @@ def test_no_messages_yields_user_text():
assert contents == expected_output
def test_convert_url():
convert_url_to_base64("https://picsum.photos/id/237/200/300")
def test_convert_url(monkeypatch):
import base64
from unittest.mock import MagicMock
import httpx
from litellm.litellm_core_utils.prompt_templates.image_handling import (
in_memory_cache,
)
url = "https://picsum.photos/id/237/200/300"
image_bytes = b"\x89PNG\r\n\x1a\nfake-png-bytes"
mock_client = MagicMock()
mock_client.get.return_value = httpx.Response(
200, content=image_bytes, headers={"Content-Type": "image/png"}
)
monkeypatch.setattr(litellm, "user_url_validation", False, raising=False)
monkeypatch.setattr(litellm, "module_level_client", mock_client, raising=False)
in_memory_cache.flush_cache()
result = convert_url_to_base64(url)
expected = "data:image/png;base64," + base64.b64encode(image_bytes).decode("utf-8")
assert result == expected
mock_client.get.assert_called_once()
def test_azure_tool_call_invoke_helper():

View file

@ -1,8 +1,8 @@
"""
Unit tests for CodeInterpreterInterceptionLogger.
All sandbox dependencies are injected (dependency injection, no monkeypatch):
a FakeSandbox stands in for the real e2b config and records how it is called.
All sandbox dependencies are injected: a FakeSandbox stands in for the real e2b
config and records how it is called.
"""
import time
@ -12,13 +12,17 @@ import pytest
from litellm.integrations.code_interpreter_interception.handler import (
CodeInterpreterInterceptionLogger,
LITELLM_CODE_EXECUTION_TOOL_NAME,
_INTERCEPTION_ACTIVE_KEY as _ACTIVE_KEY,
_SANDBOX_KEY,
)
from litellm.types.integrations.custom_logger import (
CHAT_COMPLETION_AGENTIC_SURFACE,
NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES,
is_interception_internal_key,
)
from litellm.llms.base_llm.sandbox.transformation import CodeExecutionResult
from litellm.types.utils import CallTypes
_ACTIVE_KEY = "_code_interpreter_interception_active"
_SANDBOX_KEY = "_code_interpreter_interception_sandbox_key"
class FakeHandle:
def __init__(self, sandbox_id="sbx_fake"):
@ -51,6 +55,13 @@ class FakeLogging:
def __init__(self, litellm_call_id="k1"):
self.litellm_call_id = litellm_call_id
self.model_call_details = {}
self.dynamic_success_callbacks = []
def pre_call(self, *args, **kwargs):
return None
def post_call(self, *args, **kwargs):
return None
def _function_call_item(call_id="c1", name=LITELLM_CODE_EXECUTION_TOOL_NAME):
@ -62,6 +73,17 @@ def _function_call_item(call_id="c1", name=LITELLM_CODE_EXECUTION_TOOL_NAME):
}
def _chat_function_call_item(call_id="call_1", name=LITELLM_CODE_EXECUTION_TOOL_NAME):
return {
"id": call_id,
"type": "function",
"function": {
"name": name,
"arguments": '{"code":"print(40 + 2)"}',
},
}
class FakeResponse:
def __init__(self, output):
self.output = output
@ -74,6 +96,18 @@ def _iter_messages(plan):
return patch.messages
def test_interception_internal_key_prefix_sets_preserve_code_interpreter_state():
assert is_interception_internal_key("_code_interpreter_interception_active")
assert not is_interception_internal_key(
"_code_interpreter_interception_active",
prefixes=NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES,
)
assert is_interception_internal_key(
"_websearch_interception_converted_stream",
prefixes=NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES,
)
@pytest.mark.asyncio
async def test_build_plan_runs_code_and_feeds_output_back():
sandbox = FakeSandbox(stdout="42")
@ -133,6 +167,30 @@ async def test_pre_call_converts_code_interpreter_tool():
assert LITELLM_CODE_EXECUTION_TOOL_NAME in names
@pytest.mark.asyncio
async def test_pre_call_converts_code_interpreter_tool_for_chat_completions():
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
kwargs = {
"tools": [{"type": "code_interpreter", "container": {"type": "auto"}}],
"tool_choice": {"type": "code_interpreter"},
"custom_llm_provider": "openai",
}
result = await logger.async_pre_call_deployment_hook(kwargs, CallTypes.acompletion)
assert result is not None
tool = result["tools"][0]
assert tool["type"] == "function"
assert tool["function"]["name"] == LITELLM_CODE_EXECUTION_TOOL_NAME
assert tool["function"]["parameters"]["required"] == ["code"]
assert result["tool_choice"] == {
"type": "function",
"function": {"name": LITELLM_CODE_EXECUTION_TOOL_NAME},
}
assert result["litellm_metadata"][_ACTIVE_KEY] is True
assert result["litellm_metadata"][_SANDBOX_KEY] == result[_SANDBOX_KEY]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"tool_choice",
@ -184,9 +242,24 @@ async def test_pre_call_noop_on_non_responses():
"custom_llm_provider": "openai",
}
result = await logger.async_pre_call_deployment_hook(kwargs, CallTypes.aembedding)
assert result is None
@pytest.mark.asyncio
async def test_pre_call_noop_on_chat_completion_without_code_interpreter_tool():
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
kwargs = {
"tools": [{"type": "web_search"}],
"custom_llm_provider": "openai",
}
result = await logger.async_pre_call_deployment_hook(kwargs, CallTypes.acompletion)
assert result is None
assert _ACTIVE_KEY not in kwargs
assert _SANDBOX_KEY not in kwargs
@pytest.mark.asyncio
@ -524,14 +597,142 @@ async def test_gate_rechecks_provider_scope():
assert should_run is False
@pytest.mark.asyncio
async def test_chat_completion_gate_detects_code_execution_tool_call():
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
response = {
"choices": [
{"message": {"tool_calls": [_chat_function_call_item(call_id="call_123")]}}
]
}
should_run, payload = await logger.async_should_run_agentic_loop(
response=response,
model="gpt-5",
messages=[{"role": "user", "content": "x"}],
tools=[],
stream=False,
custom_llm_provider="openai",
kwargs={
_ACTIVE_KEY: True,
"_agentic_loop_api_surface": CHAT_COMPLETION_AGENTIC_SURFACE,
},
)
assert should_run is True
assert payload["tool_calls"][0]["id"] == "call_123"
assert payload["tool_calls"][0]["arguments"] == '{"code":"print(40 + 2)"}'
@pytest.mark.asyncio
async def test_chat_completion_gate_refuses_without_server_active_marker():
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
response = {"choices": [{"message": {"tool_calls": [_chat_function_call_item()]}}]}
should_run, payload = await logger.async_should_run_agentic_loop(
response=response,
model="gpt-5",
messages=[{"role": "user", "content": "x"}],
tools=[],
stream=False,
custom_llm_provider="openai",
kwargs={"_agentic_loop_api_surface": CHAT_COMPLETION_AGENTIC_SURFACE},
)
assert should_run is False
assert payload == {}
@pytest.mark.asyncio
async def test_chat_completion_build_plan_runs_code_and_appends_tool_message():
sandbox = FakeSandbox(stdout="42")
logger = CodeInterpreterInterceptionLogger(sandbox_config=sandbox)
native_chat_tool = {"type": "code_interpreter", "container": {"type": "auto"}}
plan = await logger.async_build_agentic_loop_plan(
tools={
"tool_calls": [
{
"id": "call_1",
"name": LITELLM_CODE_EXECUTION_TOOL_NAME,
"arguments": '{"code":"print(40 + 2)"}',
}
]
},
model="gpt-5",
messages=[{"role": "user", "content": "x"}],
response={
"choices": [{"message": {"tool_calls": [_chat_function_call_item()]}}]
},
anthropic_messages_provider_config=None,
anthropic_messages_optional_request_params={
"tools": [native_chat_tool],
"tool_choice": {"type": "code_interpreter", "container": {"type": "auto"}},
"temperature": 0,
},
logging_obj=FakeLogging(litellm_call_id="k1"),
stream=False,
kwargs={
"acompletion": True,
"litellm_call_id": "k1",
_ACTIVE_KEY: True,
_SANDBOX_KEY: "sbxkey1",
"_code_interpreter_interception_converted_stream": True,
"_agentic_loop_api_surface": CHAT_COMPLETION_AGENTIC_SURFACE,
},
)
assert sandbox.run_calls[0]["code"] == "print(40 + 2)"
patch = plan.request_patch
assert patch is not None
assert patch.tools == [
{
"type": "function",
"function": {
"name": LITELLM_CODE_EXECUTION_TOOL_NAME,
"description": "Execute python code in a sandbox and return stdout.",
"parameters": {
"type": "object",
"properties": {"code": {"type": "string"}},
"required": ["code"],
},
},
}
]
assert patch.optional_params == {"temperature": 0}
assert patch.kwargs == {
"litellm_call_id": "k1",
_ACTIVE_KEY: True,
_SANDBOX_KEY: "sbxkey1",
"_code_interpreter_interception_converted_stream": True,
"_agentic_loop_api_surface": CHAT_COMPLETION_AGENTIC_SURFACE,
}
assert patch.messages is not None
assert patch.messages[-2]["role"] == "assistant"
assert patch.messages[-2]["tool_calls"][0]["id"] == "call_1"
assert patch.messages[-1] == {
"role": "tool",
"tool_call_id": "call_1",
"content": "42",
}
assert plan.metadata["code_interpreter_calls"][0]["code"] == "print(40 + 2)"
@pytest.mark.asyncio
async def test_pre_call_strips_client_forged_marker_on_initial_request():
"""A client cannot pre-set the active marker on the original request."""
"""A client cannot pre-set the active marker on the original request: with no
native code_interpreter tool, any client-supplied interception markers in
litellm_metadata are scrubbed and the active flag in kwargs is cleared."""
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
kwargs = {
"tools": [{"type": "web_search"}],
"custom_llm_provider": "openai",
_ACTIVE_KEY: True,
"litellm_metadata": {
_ACTIVE_KEY: True,
_SANDBOX_KEY: "client-forged",
"safe_user_value": "kept",
},
}
await logger.async_pre_call_deployment_hook(kwargs, CallTypes.aresponses)
@ -540,6 +741,42 @@ async def test_pre_call_strips_client_forged_marker_on_initial_request():
"no native code_interpreter tool was present, so a client-supplied "
"active marker must be cleared"
)
assert kwargs["litellm_metadata"] == {"safe_user_value": "kept"}
@pytest.mark.asyncio
async def test_pre_call_strips_forged_loop_controls_then_mints_own_markers():
"""On an INITIAL request (no server-set _agentic_loop_depth) a client cannot
smuggle loop-control state: forged _agentic_loop_depth / max_agentic_loops and
interception markers in litellm_metadata are stripped before the interceptor
activates, so the only interception markers that survive are the ones the
server mints for the converted code_interpreter tool."""
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
kwargs = {
"tools": [{"type": "code_interpreter", "container": {"type": "auto"}}],
"custom_llm_provider": "openai",
"litellm_metadata": {
_ACTIVE_KEY: True,
_SANDBOX_KEY: "client-forged",
"_agentic_loop_depth": 99,
"max_agentic_loops": 999,
"safe_user_value": "kept",
},
}
result = await logger.async_pre_call_deployment_hook(kwargs, CallTypes.acompletion)
assert result is not None
metadata = result["litellm_metadata"]
assert metadata["safe_user_value"] == "kept"
assert "_agentic_loop_depth" not in metadata, "forged loop depth must be stripped"
assert "max_agentic_loops" not in metadata, "forged loop cap must be stripped"
assert metadata[_ACTIVE_KEY] is True
assert metadata[_SANDBOX_KEY] == result[_SANDBOX_KEY]
assert metadata[_SANDBOX_KEY] != "client-forged", (
"the surviving sandbox key must be the server-minted one, not the forged "
"value the client supplied"
)
@pytest.mark.asyncio

View file

@ -35,6 +35,13 @@ from litellm.integrations.otel.mount import ( # noqa: E402
)
@pytest.fixture(autouse=True)
def _clear_otel_v2_flag_cache():
is_otel_v2_enabled.cache_clear()
yield
is_otel_v2_enabled.cache_clear()
class _FakeSpan:
"""Minimal recording span capturing what the hook writes."""
@ -70,8 +77,10 @@ def _instrumented_app():
def test_gate_toggles_with_env(monkeypatch):
"""The startup mount is guarded by this flag."""
monkeypatch.delenv("LITELLM_OTEL_V2", raising=False)
is_otel_v2_enabled.cache_clear()
assert is_otel_v2_enabled() is False
monkeypatch.setenv("LITELLM_OTEL_V2", "1")
is_otel_v2_enabled.cache_clear()
assert is_otel_v2_enabled() is True

View file

@ -1,6 +1,8 @@
"""Tests for the OTel v2 sources of truth: span registry, semconv keys, config,
and the typed StandardLoggingPayload adapter. These need no OTel SDK."""
import pytest
from litellm.integrations.otel import (
BAGGAGE_PROMOTED_KEYS,
DB,
@ -28,6 +30,13 @@ from litellm.integrations.otel.model.spans import (
)
@pytest.fixture(autouse=True)
def _clear_otel_v2_flag_cache():
is_otel_v2_enabled.cache_clear()
yield
is_otel_v2_enabled.cache_clear()
def _sample_payload(**overrides):
payload = {
"call_type": "acompletion",
@ -561,11 +570,37 @@ def test_capture_message_content_normalizer_only_touches_strings():
def test_v2_flag_is_off_by_default(monkeypatch):
monkeypatch.delenv("LITELLM_OTEL_V2", raising=False)
is_otel_v2_enabled.cache_clear()
assert is_otel_v2_enabled() is False
monkeypatch.setenv("LITELLM_OTEL_V2", "true")
is_otel_v2_enabled.cache_clear()
assert is_otel_v2_enabled() is True
def test_v2_flag_resolved_once_not_per_call(monkeypatch):
"""Regression for LIT-3895: ``is_otel_v2_enabled`` sits on the proxy hot path
(auth, logging-callback setup). Building the pydantic-settings model on every
call re-scanned the environment at ~28us a pop and dropped throughput, so the
flag must be resolved once and cached rather than reconstructed per call."""
from litellm.integrations.otel.model import config as config_mod
constructions = 0
real_flag = config_mod._OTelV2Flag
def _counting_flag(*args, **kwargs):
nonlocal constructions
constructions += 1
return real_flag(*args, **kwargs)
monkeypatch.setattr(config_mod, "_OTelV2Flag", _counting_flag)
config_mod.is_otel_v2_enabled.cache_clear()
for _ in range(50):
config_mod.is_otel_v2_enabled()
assert constructions == 1
def test_config_from_env(monkeypatch):
for var in (
"OTEL_EXPORTER",

View file

@ -299,7 +299,6 @@ class TestGoogleInteractionsResponseStructure:
assert hasattr(response, "outputs")
assert hasattr(response, "usage")
assert hasattr(response, "model") or hasattr(response, "agent")
assert hasattr(response, "role")
assert hasattr(response, "created")
assert hasattr(response, "updated")

View file

@ -156,16 +156,19 @@ class TestResponseCompliance:
# The response is the dedicated `Interaction` schema. Google moved the
# output-only fields (notably the `steps` array, formerly `outputs`)
# off `CreateModelInteractionParams` and onto `Interaction`; the request
# schema no longer carries `steps`. Keep this aligned with the live spec.
# schema no longer carries `steps`. Google later moved `role` off
# `Interaction` onto the per-turn `Turn` schema (asserted in
# test_turn_schema), so it is no longer a top-level output field here.
# Keep this aligned with the live spec.
schema = spec_dict["components"]["schemas"]["Interaction"]
# Output fields (readOnly).
# Output fields (readOnly). `role` was removed from the `Interaction`
# schema by Google; it now lives only on `Turn`.
output_fields = [
"id",
"status",
"created",
"updated",
"role",
"steps",
"usage",
]

View file

@ -0,0 +1,422 @@
"""
Tests for the provider-agnostic chat completion agentic loop dispatcher
(`litellm/litellm_core_utils/chat_completion_agentic_loop.py`) and the
code-interpreter interception integration that drives it.
The load-bearing regression here protects a reviewer requirement: the internal
agentic/interception control fields must NEVER reach the outbound provider HTTP
request body. The relevant fields are:
_agentic_loop_depth
_agentic_loop_fingerprints
_agentic_loop_api_surface
max_agentic_loops
_code_interpreter_interception_active
_code_interpreter_interception_sandbox_key
_code_interpreter_interception_converted_stream
A scrubber in gpt_transformation.py used to strip these. That scrubber was
removed, so `test_internal_control_fields_never_leak_into_provider_body` proves
they stay out of the body even without it.
"""
import os
import sys
from typing import Any, Dict, List, Optional, Tuple
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
sys.path.insert(0, os.path.abspath("../../../.."))
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.code_interpreter_interception.handler import (
CodeInterpreterInterceptionLogger,
)
from litellm.litellm_core_utils.chat_completion_agentic_loop import (
maybe_run_chat_completion_agentic_loop,
)
from litellm.types.integrations.custom_logger import (
AgenticLoopPlan,
AgenticLoopRequestPatch,
)
from litellm.types.utils import (
Choices,
Function,
ChatCompletionMessageToolCall,
Message,
ModelResponse,
)
# The internal control fields that must never reach a provider request body.
_INTERNAL_CONTROL_FIELDS = (
"_agentic_loop_depth",
"_agentic_loop_fingerprints",
"_agentic_loop_api_surface",
"max_agentic_loops",
"_code_interpreter_interception_active",
"_code_interpreter_interception_sandbox_key",
"_code_interpreter_interception_converted_stream",
"litellm_metadata",
)
@pytest.fixture
def restore_callbacks():
"""Save/restore litellm.callbacks so a registered fake logger never pollutes
other tests in the suite."""
saved = list(litellm.callbacks)
try:
yield
finally:
litellm.callbacks = saved
class _SandboxResult:
def __init__(self, stdout: str) -> None:
self.stdout = stdout
self.error = None
class FakeSandboxConfig:
"""Injected sandbox so the interception loop runs no real network / E2B."""
def __init__(self) -> None:
self.created = 0
self.deleted = 0
self.run_codes: List[str] = []
async def acreate_sandbox(self) -> Any:
self.created += 1
return MagicMock(id="sandbox-123")
async def arun_code(self, container: Any, code: str) -> _SandboxResult:
self.run_codes.append(code)
return _SandboxResult(stdout="42\n")
async def adelete_sandbox(self, container: Any) -> None:
self.deleted += 1
def _tool_call_model_response() -> ModelResponse:
return ModelResponse(
choices=[
Choices(
finish_reason="tool_calls",
message=Message(
role="assistant",
content=None,
tool_calls=[
ChatCompletionMessageToolCall(
id="call_abc",
type="function",
function=Function(
name="litellm_code_execution",
arguments='{"code": "print(6*7)"}',
),
)
],
),
)
]
)
def _plain_model_response(content: str = "The answer is 42") -> ModelResponse:
return ModelResponse(
choices=[
Choices(
finish_reason="stop",
message=Message(role="assistant", content=content),
)
]
)
def _raw_response_for(model_response: ModelResponse) -> MagicMock:
"""Wrap a ModelResponse as the OpenAI `with_raw_response.create` return value
(an object exposing `.headers` and `.parse()` -> something with model_dump)."""
parsed = MagicMock()
parsed.model_dump.return_value = model_response.model_dump()
raw = MagicMock()
raw.headers = {}
raw.parse.return_value = parsed
return raw
# ---------------------------------------------------------------------------
# A) PROVIDER-PAYLOAD REGRESSION
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_internal_control_fields_never_leak_into_provider_body(restore_callbacks):
"""Drive a real acompletion with a native code_interpreter tool through the
interception logger + agentic loop, capturing every outbound OpenAI request
body. None of the internal control fields may appear at top-level or inside
extra_body on ANY of the captured calls."""
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandboxConfig())
litellm.callbacks = [logger]
# First create -> model emits a code_execution tool call (triggers the loop).
# Second create -> model returns a plain answer (loop terminates).
create = AsyncMock(
side_effect=[
_raw_response_for(_tool_call_model_response()),
_raw_response_for(_plain_model_response()),
]
)
mock_client = MagicMock()
mock_client.chat.completions.with_raw_response.create = create
response = await litellm.acompletion(
model="openai/gpt-4o-mini",
messages=[{"role": "user", "content": "what is 6*7?"}],
tools=[{"type": "code_interpreter"}],
tool_choice={"type": "code_interpreter"},
api_key="sk-test",
client=mock_client,
)
# The loop must have actually fired (sanity: two provider calls).
assert create.await_count == 2, (
"expected the agentic loop to issue a follow-up provider call; "
f"got {create.await_count} call(s)"
)
for idx, call in enumerate(create.await_args_list):
body = call.kwargs
extra_body = body.get("extra_body") or {}
for field in _INTERNAL_CONTROL_FIELDS:
assert field not in body, (
f"provider call #{idx}: internal field {field!r} leaked into "
f"top-level request body: {sorted(body.keys())}"
)
assert field not in extra_body, (
f"provider call #{idx}: internal field {field!r} leaked into "
f"extra_body: {sorted(extra_body.keys())}"
)
# The native code_interpreter tool must have been swapped for the
# function tool, never sent raw to OpenAI as a chat-completions request.
for tool in body.get("tools") or []:
assert tool.get("type") != "code_interpreter"
# The final response is the post-loop answer, not the tool-call turn.
assert response.choices[0].message.content == "The answer is 42"
# ---------------------------------------------------------------------------
# B) DISPATCHER UNIT TESTS
# ---------------------------------------------------------------------------
class _LoggingStub:
"""Minimal logging_obj: dispatcher only reads dynamic_success_callbacks and
litellm_call_id off it."""
litellm_call_id = "call-test"
dynamic_success_callbacks: List[Any] = []
class _GateOnlyLogger(CustomLogger):
"""Overrides the gate to fire, but builds a plan from request_patch."""
def __init__(self, plan: AgenticLoopPlan, tool_calls: Dict[str, Any]) -> None:
super().__init__()
self._plan = plan
self._tool_calls = tool_calls
self.cleanup_calls = 0
async def async_should_run_agentic_loop(
self,
response: Any,
model: str,
messages: List[Dict[str, Any]],
tools: Optional[List[Dict[str, Any]]],
stream: bool,
custom_llm_provider: str,
kwargs: Dict[str, Any],
) -> Tuple[bool, Dict[str, Any]]:
return True, self._tool_calls
async def async_build_agentic_loop_plan(
self,
tools: Dict[str, Any],
model: str,
messages: List[Dict[str, Any]],
response: Any,
anthropic_messages_provider_config: Any,
anthropic_messages_optional_request_params: Dict[str, Any],
logging_obj: Any,
stream: bool,
kwargs: Dict[str, Any],
) -> AgenticLoopPlan:
return self._plan
async def async_agentic_loop_cleanup_hook(
self, plan: AgenticLoopPlan, kwargs: Dict[str, Any]
) -> None:
self.cleanup_calls += 1
def _patched_messages() -> List[Dict[str, Any]]:
return [
{"role": "user", "content": "what is 6*7?"},
{
"role": "assistant",
"tool_calls": [
{
"id": "call_abc",
"type": "function",
"function": {
"name": "litellm_code_execution",
"arguments": '{"code": "print(6*7)"}',
},
}
],
},
{"role": "tool", "tool_call_id": "call_abc", "content": "42\n"},
]
@pytest.mark.asyncio
async def test_dispatcher_returns_none_when_no_callback_gates(restore_callbacks):
"""No callback overrides the gate -> dispatcher returns None so the caller
keeps the original response untouched."""
litellm.callbacks = []
result = await maybe_run_chat_completion_agentic_loop(
response=_plain_model_response(),
model="gpt-4o-mini",
messages=[{"role": "user", "content": "hi"}],
optional_params={},
kwargs={},
logging_obj=_LoggingStub(),
custom_llm_provider="openai",
stream=False,
)
assert result is None
@pytest.mark.asyncio
async def test_dispatcher_runs_followup_with_incremented_depth_and_patched_messages(
restore_callbacks,
):
"""A gating logger with a request_patch -> the dispatcher calls
litellm.acompletion exactly once with _agentic_loop_depth == 1 and the
patched messages. Loop-control state rides as litellm-level kwargs and is
mirrored into litellm_metadata; the provider-surface transient
_agentic_loop_api_surface is never forwarded. (Provider-body stripping of
these litellm-level kwargs is asserted separately in test A.)"""
followup = _plain_model_response("done")
plan = AgenticLoopPlan(
run_agentic_loop=True,
request_patch=AgenticLoopRequestPatch(messages=_patched_messages()),
)
logger = _GateOnlyLogger(plan=plan, tool_calls={"tool_calls": [{"id": "call_abc"}]})
litellm.callbacks = [logger]
acompletion_mock = AsyncMock(return_value=followup)
with patch.object(litellm, "acompletion", acompletion_mock):
result = await maybe_run_chat_completion_agentic_loop(
response=_tool_call_model_response(),
model="gpt-4o-mini",
messages=[{"role": "user", "content": "what is 6*7?"}],
optional_params={"temperature": 0.1},
kwargs={"_code_interpreter_interception_active": True},
logging_obj=_LoggingStub(),
custom_llm_provider="openai",
stream=False,
)
assert result is followup
acompletion_mock.assert_awaited_once()
call_kwargs = acompletion_mock.await_args.kwargs
assert call_kwargs["_agentic_loop_depth"] == 1
assert call_kwargs["messages"] == _patched_messages()
# Preserved non-internal optional param survives the rerun.
assert call_kwargs["temperature"] == 0.1
# Loop-control state is carried at the litellm level for the follow-up.
assert call_kwargs["max_agentic_loops"] >= 1
assert "_agentic_loop_fingerprints" in call_kwargs
# Interception markers are mirrored into litellm_metadata for the follow-up.
assert (
call_kwargs["litellm_metadata"]["_code_interpreter_interception_active"] is True
)
# The transient surface marker is NOT forwarded to the follow-up call.
assert "_agentic_loop_api_surface" not in call_kwargs
# Cleanup hook always runs.
assert logger.cleanup_calls == 1
@pytest.mark.asyncio
async def test_dispatcher_raises_when_depth_reaches_max_agentic_loops(
restore_callbacks,
):
"""depth >= max_agentic_loops -> ValueError mentioning max_agentic_loops,
before any follow-up call is attempted."""
logger = _GateOnlyLogger(
plan=AgenticLoopPlan(run_agentic_loop=True),
tool_calls={"tool_calls": [{"id": "call_abc"}]},
)
litellm.callbacks = [logger]
acompletion_mock = AsyncMock()
with patch.object(litellm, "acompletion", acompletion_mock):
with pytest.raises(ValueError, match="max_agentic_loops"):
await maybe_run_chat_completion_agentic_loop(
response=_tool_call_model_response(),
model="gpt-4o-mini",
messages=[{"role": "user", "content": "hi"}],
optional_params={},
kwargs={"_agentic_loop_depth": 3, "max_agentic_loops": 3},
logging_obj=_LoggingStub(),
custom_llm_provider="openai",
stream=False,
)
acompletion_mock.assert_not_awaited()
@pytest.mark.asyncio
async def test_dispatcher_raises_on_repeated_tool_call_fingerprint(restore_callbacks):
"""A tool_calls fingerprint already present in _agentic_loop_fingerprints ->
ValueError about the repeated fingerprint (cycle guard), with no follow-up
call."""
import json
# The dispatcher fingerprints the whole value the gate returns as its second
# tuple element, so the seeded fingerprint must mirror that dict exactly.
gate_tool_calls = {
"tool_calls": [{"id": "call_abc", "name": "litellm_code_execution"}]
}
fingerprint = json.dumps(gate_tool_calls, sort_keys=True, default=str)
logger = _GateOnlyLogger(
plan=AgenticLoopPlan(run_agentic_loop=True),
tool_calls=gate_tool_calls,
)
litellm.callbacks = [logger]
acompletion_mock = AsyncMock()
with patch.object(litellm, "acompletion", acompletion_mock):
with pytest.raises(ValueError, match="fingerprint"):
await maybe_run_chat_completion_agentic_loop(
response=_tool_call_model_response(),
model="gpt-4o-mini",
messages=[{"role": "user", "content": "hi"}],
optional_params={},
kwargs={
"_agentic_loop_depth": 0,
"max_agentic_loops": 3,
"_agentic_loop_fingerprints": [fingerprint],
},
logging_obj=_LoggingStub(),
custom_llm_provider="openai",
stream=False,
)
acompletion_mock.assert_not_awaited()

View file

@ -132,3 +132,17 @@ def test_azure_base_model_detection_preserved():
assert params is not None
assert "reasoning_effort" in params
assert "tools" in params
def test_sambanova_embeddings_request_returns_list_not_none():
"""The sambanova embeddings branch resolved the config but dropped the result,
so embedding requests got ``None`` instead of the supported-params list while the
chat branch returned correctly. A list (the sambanova embeddings config exposes no
extra params, hence ``[]``) must reach the caller."""
embedding_params = get_supported_openai_params(
model="E5-Mistral-7B-Instruct",
custom_llm_provider="sambanova",
request_type="embeddings",
)
assert embedding_params == []

View file

@ -126,6 +126,49 @@ def test_lists_with_sensitive_keys_are_masked():
assert masked["tags"] == ["prod", "test"]
def test_short_secrets_are_fully_masked():
"""
Regression test: secrets at or below the reveal threshold (visible_prefix +
visible_suffix, 8 by default) were returned verbatim instead of masked.
An exactly-8-char value hit masked_length == 0 and round-tripped unchanged;
anything shorter hit the early return. Both leaked short credentials (e.g. an
8-char redis password) in plaintext through mask_dict.
"""
masker = SensitiveDataMasker()
# Boundary: exactly 8 chars previously returned verbatim.
assert masker._mask_value("abcd1234") == "********"
# Below threshold previously hit the early return and leaked verbatim.
assert masker._mask_value("sk-12") == "*****"
# Values above the threshold must still partially reveal, not over-mask.
assert masker._mask_value("abcd12345") == "abcd*2345"
masked = masker.mask_dict({"redis_password": "pass1234", "api_key": "sk-7a"})
assert masked["redis_password"] == "********"
assert masked["api_key"] == "*****"
def test_mask_short_values_false_keeps_short_values_readable():
"""
mask_short_values=False opts out of full masking so short values are returned
as-is. This preserves the truncation use (e.g. CooldownCache shows the first 50
chars of an exception and only masks longer tails), while longer values are still
partially masked.
"""
masker = SensitiveDataMasker(
visible_prefix=50, visible_suffix=0, mask_short_values=False
)
short = "Test exception for structure validation"
assert masker._mask_value(short) == short
long_value = "x" * 60
masked = masker._mask_value(long_value)
assert masked.startswith("x" * 50)
assert masked.endswith("*" * 10)
assert len(masked) == 60
def test_cost_per_token_fields_not_masked():
"""
Regression test: cost fields like input_cost_per_token contain "token" in their name

View file

@ -878,6 +878,114 @@ def test_sync_streaming_bad_request_not_midstream(logging_obj: Logging):
assert "invalid maxOutputTokens" in str(excinfo.value)
def _bedrock_error_event(exception_type: str):
"""A mocked botocore event-stream error event: status_code is botocore's
hard-coded 400, with the real type in the :exception-type header."""
event = Mock()
event.to_response_dict = Mock(
return_value={
"status_code": 400,
"headers": {
":exception-type": exception_type,
":content-type": "application/json",
":message-type": "exception",
},
"body": b'{"message":"Bedrock had an internal error."}',
}
)
return event
@pytest.mark.asyncio
async def test_bedrock_midstream_internal_server_error_wraps_for_fallback(
logging_obj: Logging,
):
"""End-to-end regression for https://github.com/BerriAI/litellm/issues/24608:
a Bedrock mid-stream internalServerException event (botocore stamps it 400)
must flow through the real decoder, gain its modeled 500 status, and wrap
into MidStreamFallbackError so the Router can run streaming fallback.
Calls the real AWSEventStreamDecoder, so reverting the decoder status fix
makes the decoder raise BedrockError(400) and the gate raises BadRequestError
directly -> this test fails without the fix."""
from litellm.exceptions import MidStreamFallbackError
from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder
decoder = AWSEventStreamDecoder(model="anthropic.claude-3-sonnet-20240229-v1:0")
async def _bedrock_stream():
decoder._parse_message_from_event(
_bedrock_error_event("internalServerException")
)
yield # unreachable; the line above raises
async def _make_call(**kwargs):
return _bedrock_stream()
response = CustomStreamWrapper(
completion_stream=None,
model="anthropic.claude-3-sonnet-20240229-v1:0",
logging_obj=logging_obj,
custom_llm_provider="bedrock",
make_call=_make_call,
)
with pytest.raises(MidStreamFallbackError):
await response.__anext__()
@pytest.mark.asyncio
async def test_bedrock_5xx_wraps_for_midstream_fallback(logging_obj: Logging):
"""Gate contract: a Bedrock 5xx (here 503 serviceUnavailableException) wraps
into MidStreamFallbackError so the Router can run streaming fallback."""
from litellm.exceptions import MidStreamFallbackError
from litellm.llms.bedrock.chat.invoke_handler import BedrockError
async def _raise_503(**kwargs):
raise BedrockError(
status_code=503,
message="serviceUnavailableException Bedrock is unavailable.",
)
response = CustomStreamWrapper(
completion_stream=None,
model="anthropic.claude-3-sonnet-20240229-v1:0",
logging_obj=logging_obj,
custom_llm_provider="bedrock",
make_call=_raise_503,
)
with pytest.raises(MidStreamFallbackError):
await response.__anext__()
@pytest.mark.asyncio
async def test_bedrock_validation_error_raises_directly(logging_obj: Logging):
"""Gate contract: a Bedrock validationException (400) is a client error and
must surface directly, never wrapped into MidStreamFallbackError."""
from litellm.exceptions import MidStreamFallbackError
from litellm.llms.bedrock.chat.invoke_handler import BedrockError
async def _raise_400(**kwargs):
raise BedrockError(
status_code=400,
message="validationException malformed input.",
)
response = CustomStreamWrapper(
completion_stream=None,
model="anthropic.claude-3-sonnet-20240229-v1:0",
logging_obj=logging_obj,
custom_llm_provider="bedrock",
make_call=_raise_400,
)
with pytest.raises(Exception) as excinfo:
await response.__anext__()
assert not isinstance(excinfo.value, MidStreamFallbackError)
assert getattr(excinfo.value, "status_code", None) == 400
@pytest.mark.asyncio
async def test_async_streaming_read_timeout_triggers_midstream_fallback(
logging_obj: Logging,
@ -2646,7 +2754,9 @@ def test_chunk_creator_tool_calls_not_dropped_on_finish(
tool_calls=[
ChatCompletionDeltaToolCall(
id="call_abc",
function=Function(name="get_weather", arguments='{"city":"NYC"}'),
function=Function(
name="get_weather", arguments='{"city":"NYC"}'
),
type="function",
index=0,
)
@ -2741,3 +2851,131 @@ def test_record_partial_usage_for_failure_noop_without_chunks():
wrapper._record_partial_usage_for_failure()
assert "combined_usage_object" not in logging_obj.model_call_details
@pytest.mark.parametrize("sync_mode", [True, False])
@pytest.mark.asyncio
async def test_stream_chunk_builder_raise_at_end_of_stream_still_recovers_usage(
sync_mode,
):
"""stream_chunk_builder re-raises (as APIError) on large agentic tool-use
streams. That raise originates inside the except-StopIteration handler, so
before the fix it escaped __next__/__anext__ and the request was dropped from
SpendLogs while the provider billed the tokens. The wrapper must catch it and
recover usage from the raw chunks so cost is still tracked."""
final_usage_block = Usage(
completion_tokens=392, prompt_tokens=1799, total_tokens=2191
)
final_chunk = ModelResponseStream(
id="chatcmpl-raise-test",
created=1742056047,
model=None,
object="chat.completion.chunk",
choices=[
StreamingChoices(
finish_reason="stop",
index=0,
delta=Delta(content="", role="assistant"),
)
],
usage=final_usage_block,
)
test_chunks = bedrock_chunks + [final_chunk]
logging_obj = Logging(
model="bedrock/claude-haiku-4-5-20251001-v1:0",
messages=[{"role": "user", "content": "Hey"}],
stream=True,
call_type="completion",
start_time=time.time(),
litellm_call_id="raise-test",
function_id="1245",
)
response = CustomStreamWrapper(
completion_stream=ModelResponseListIterator(model_responses=test_chunks),
model="bedrock/claude-haiku-4-5-20251001-v1:0",
custom_llm_provider="bedrock",
logging_obj=logging_obj,
stream_options={"include_usage": True},
)
seen_usage = []
with patch.object(
litellm,
"stream_chunk_builder",
side_effect=Exception("simulated assembly failure"),
):
# before the fix this raised and dropped the request; it must not raise now
if sync_mode:
for chunk in response:
if getattr(chunk, "usage", None) is not None:
seen_usage.append(chunk.usage)
else:
async for chunk in response:
if getattr(chunk, "usage", None) is not None:
seen_usage.append(chunk.usage)
assert any(
u.total_tokens == final_usage_block.total_tokens for u in seen_usage
), "usage recovered from raw chunks was not emitted after stream_chunk_builder raised"
@pytest.mark.parametrize("sync_mode", [True, False])
@pytest.mark.asyncio
async def test_stream_chunk_builder_raise_and_usage_recovery_failure_does_not_crash(
sync_mode,
):
"""If end-of-stream assembly raises AND best-effort usage recovery from the raw
chunks also fails, the stream must still complete cleanly rather than propagate
the exception to the consumer."""
from litellm.litellm_core_utils import streaming_handler as sh_module
final_chunk = ModelResponseStream(
id="chatcmpl-raise-recover-fail",
created=1742056047,
model=None,
object="chat.completion.chunk",
choices=[
StreamingChoices(
finish_reason="stop",
index=0,
delta=Delta(content="", role="assistant"),
)
],
usage=Usage(completion_tokens=1, prompt_tokens=1, total_tokens=2),
)
response = CustomStreamWrapper(
completion_stream=ModelResponseListIterator(
model_responses=bedrock_chunks + [final_chunk]
),
model="bedrock/claude-haiku-4-5-20251001-v1:0",
custom_llm_provider="bedrock",
logging_obj=Logging(
model="bedrock/claude-haiku-4-5-20251001-v1:0",
messages=[{"role": "user", "content": "Hey"}],
stream=True,
call_type="completion",
start_time=time.time(),
litellm_call_id="raise-recover-fail",
function_id="1245",
),
stream_options={"include_usage": True},
)
with (
patch.object(
litellm, "stream_chunk_builder", side_effect=Exception("assembly failed")
),
patch.object(
sh_module, "calculate_total_usage", side_effect=Exception("recovery failed")
),
):
# must not raise even though both assembly and recovery fail
if sync_mode:
chunks = [c for c in response]
else:
chunks = [c async for c in response]
assert len(chunks) > 0

View file

@ -516,3 +516,217 @@ async def test_max_uses_none_falls_back_to_default():
)
assert str(_c.ADVISOR_MAX_USES) in str(exc_info.value)
# ---------------------------------------------------------------------------
# 12. Defense-in-depth: client-supplied advisor api_base/api_key are dropped
# unless the proxy admin opted into clientside credentials
# ---------------------------------------------------------------------------
ADVISOR_TOOL_WITH_CREDS = {
"type": "advisor_20260301",
"name": "advisor",
"model": "claude-opus-4-6",
"api_base": "https://other.example",
"api_key": "sk-other",
}
async def _run_advisor_and_capture_subcall_kwargs():
"""Run one advisor turn and return the kwargs of the advisor sub-call."""
from litellm.llms.anthropic.experimental_pass_through.messages.interceptors.advisor import (
AdvisorOrchestrationHandler,
)
advisor_tool_use_resp = _make_advisor_tool_use_response(tool_id="toolu_01")
advisor_advice_resp = _make_text_response("advice", model="claude-opus-4-6")
final_resp = _make_text_response("final answer")
captured = {}
call_count = 0
async def mock_call(model, messages, tools, stream, max_tokens, **kwargs):
nonlocal call_count
call_count += 1
if call_count == 1:
return advisor_tool_use_resp
if call_count == 2:
# The advisor sub-call — capture its routing kwargs.
captured["api_key"] = kwargs.get("api_key")
captured["api_base"] = kwargs.get("api_base")
return advisor_advice_resp
return final_resp
with patch(
"litellm.llms.anthropic.experimental_pass_through.messages.interceptors.advisor._call_messages_handler",
side_effect=mock_call,
):
h = AdvisorOrchestrationHandler()
await h.handle(
model="openai/gpt-4o-mini",
messages=MESSAGES,
tools=[ADVISOR_TOOL_WITH_CREDS],
stream=False,
max_tokens=512,
custom_llm_provider="openai",
)
return captured
@pytest.mark.asyncio
async def test_advisor_creds_dropped_when_proxy_opt_in_disabled():
"""On the proxy without opt-in, the caller's advisor api_base/api_key must
NOT reach the sub-call (would redirect it / leak the server key)."""
with patch(
"litellm.llms.anthropic.experimental_pass_through.messages.interceptors.advisor._allow_client_side_advisor_credentials",
return_value=False,
):
captured = await _run_advisor_and_capture_subcall_kwargs()
assert captured["api_key"] is None
assert captured["api_base"] is None
@pytest.mark.asyncio
async def test_advisor_creds_honored_when_proxy_opt_in_enabled():
"""With the admin opt-in, the documented clientside routing still works."""
with patch(
"litellm.llms.anthropic.experimental_pass_through.messages.interceptors.advisor._allow_client_side_advisor_credentials",
return_value=True,
):
captured = await _run_advisor_and_capture_subcall_kwargs()
assert captured["api_key"] == "sk-other"
assert captured["api_base"] == "https://other.example"
# ---------------------------------------------------------------------------
# 13. The proxy gate itself: _allow_client_side_advisor_credentials() and the
# full handle() driven by the real proxy general_settings flag.
# ---------------------------------------------------------------------------
def _fake_proxy_server(general_settings: Dict):
"""A stand-in litellm.proxy.proxy_server module exposing general_settings.
The real proxy_server pulls in heavy optional deps that may be absent in a
unit-test environment, so the gate's
``from litellm.proxy.proxy_server import general_settings`` is satisfied by
injecting this lightweight module into sys.modules.
"""
import types
module = types.ModuleType("litellm.proxy.proxy_server")
module.general_settings = general_settings # type: ignore[attr-defined]
return module
def test_allow_client_side_advisor_credentials_reads_proxy_flag():
"""The gate mirrors the proxy's allow_client_side_credentials opt-in."""
import sys
from litellm.llms.anthropic.experimental_pass_through.messages.interceptors.advisor import (
_allow_client_side_advisor_credentials,
)
cases = (
({"allow_client_side_credentials": True}, True),
({"allow_client_side_credentials": False}, False),
# Flag absent entirely -> default deny on the proxy.
({}, False),
)
for settings, expected in cases:
with patch.dict(
sys.modules,
{"litellm.proxy.proxy_server": _fake_proxy_server(settings)},
):
assert _allow_client_side_advisor_credentials() is expected
def test_allow_client_side_advisor_credentials_defaults_true_outside_proxy():
"""Outside the proxy (proxy_server import unavailable), there is no admin
boundary, so the gate permits client-supplied routing."""
import builtins
import sys
from litellm.llms.anthropic.experimental_pass_through.messages.interceptors.advisor import (
_allow_client_side_advisor_credentials,
)
real_import = builtins.__import__
def _blocked_import(name, *args, **kwargs):
if name == "litellm.proxy.proxy_server":
raise ImportError("proxy server unavailable")
return real_import(name, *args, **kwargs)
with patch.dict(sys.modules):
sys.modules.pop("litellm.proxy.proxy_server", None)
with patch.object(builtins, "__import__", _blocked_import):
assert _allow_client_side_advisor_credentials() is True
def test_advisor_gate_propagates_non_import_errors():
"""Non-ImportError failures during the proxy module probe must not
default permissive. If the proxy is partially loaded and raises
RuntimeError, the gate should surface that rather than silently
returning True."""
import sys
from litellm.llms.anthropic.experimental_pass_through.messages.interceptors import (
advisor,
)
original = sys.modules.get("litellm.proxy.proxy_server")
class _Broken:
def __getattr__(self, _name):
raise RuntimeError("partial proxy boot")
sys.modules["litellm.proxy.proxy_server"] = _Broken()
try:
with pytest.raises(RuntimeError, match="partial proxy boot"):
advisor._allow_client_side_advisor_credentials()
finally:
if original is None:
sys.modules.pop("litellm.proxy.proxy_server", None)
else:
sys.modules["litellm.proxy.proxy_server"] = original
@pytest.mark.asyncio
async def test_advisor_ignores_tool_credentials_when_clientside_disabled():
"""Driven by the real proxy flag (not a patched gate): with
allow_client_side_credentials False, the tool-supplied api_base/api_key must
not reach the advisor sub-call."""
import sys
with patch.dict(
sys.modules,
{
"litellm.proxy.proxy_server": _fake_proxy_server(
{"allow_client_side_credentials": False}
)
},
):
captured = await _run_advisor_and_capture_subcall_kwargs()
assert captured["api_key"] is None
assert captured["api_base"] is None
@pytest.mark.asyncio
async def test_advisor_uses_tool_credentials_when_clientside_enabled():
"""Driven by the real proxy flag: with allow_client_side_credentials True,
the tool-supplied api_base/api_key flow through to the advisor sub-call."""
import sys
with patch.dict(
sys.modules,
{
"litellm.proxy.proxy_server": _fake_proxy_server(
{"allow_client_side_credentials": True}
)
},
):
captured = await _run_advisor_and_capture_subcall_kwargs()
assert captured["api_key"] == "sk-other"
assert captured["api_base"] == "https://other.example"

View file

@ -163,6 +163,134 @@ def test_aws_profile_path_not_cached_in_iam_cache():
assert mock_profile.call_count == 2
def test_get_credentials_does_not_expand_request_env_reference():
"""
A parameter of the form os.environ/<VAR> reaching get_credentials is left as-is
rather than expanded against the process environment, so the downstream auth
helper only ever receives the literal value.
"""
env = _os_environ_without_aws_keys()
env["SERVER_ONLY_VALUE"] = "config-managed-value"
base = BaseAWSLLM()
with patch.dict(os.environ, env, clear=True), patch.object(
base,
"_auth_with_aws_profile",
return_value=(Credentials("ak", "sk", None), None),
) as mock_profile:
base.get_credentials(aws_profile_name="os.environ/SERVER_ONLY_VALUE")
assert mock_profile.call_args.args[0] == "os.environ/SERVER_ONLY_VALUE"
assert "config-managed-value" not in str(mock_profile.call_args)
def test_get_credentials_falls_back_to_ambient_aws_profile_name_env():
"""
The fixed AWS_* ambient fallback keeps working: an unset aws_profile_name
resolves from the AWS_PROFILE_NAME environment variable.
"""
env = _os_environ_without_aws_keys()
env["AWS_PROFILE_NAME"] = "ambient-profile"
base = BaseAWSLLM()
with patch.dict(os.environ, env, clear=True), patch.object(
base,
"_auth_with_aws_profile",
return_value=(Credentials("ak", "sk", None), None),
) as mock_profile:
base.get_credentials(aws_profile_name=None)
assert mock_profile.call_args.args[0] == "ambient-profile"
def test_get_credentials_ambient_fallback_resolves_aws_external_id():
"""
Each unset param falls back to its own AWS_* env var. Regression for an index
misalignment between the value list and the env-name list, which left
AWS_EXTERNAL_ID unresolved.
"""
env = _os_environ_without_aws_keys()
env["AWS_EXTERNAL_ID"] = "ext-from-env"
base = BaseAWSLLM()
with patch.dict(os.environ, env, clear=True), patch.object(
base,
"_auth_with_aws_role",
return_value=(Credentials("ak", "sk", "tok"), None),
) as mock_role:
base.get_credentials(
aws_role_name="arn:aws:iam::123456789012:role/x",
aws_session_name="s",
)
assert mock_role.call_args.kwargs["aws_external_id"] == "ext-from-env"
def _capturing_sts_client(captured: Dict[str, Any]) -> MagicMock:
sts = MagicMock()
def _assume(**params):
captured["WebIdentityToken"] = params.get("WebIdentityToken")
return {
"Credentials": {
"AccessKeyId": "AKIA",
"SecretAccessKey": "sk",
"SessionToken": "tok",
},
"PackedPolicySize": 10,
}
sts.assume_role_with_web_identity.side_effect = _assume
return sts
@pytest.mark.parametrize(
"token_ref",
["os.environ/SERVER_ONLY_VALUE", "SERVER_ONLY_VALUE"],
ids=["os_environ_prefix", "bare_env_name"],
)
def test_web_identity_token_env_reference_not_expanded(token_ref):
"""
A web-identity token that is an environment-variable reference (an os.environ/
prefix, or a bare name matching an env var) is rejected rather than expanded, so
the process-environment value is never used as the token.
"""
env = _os_environ_without_aws_keys()
env["SERVER_ONLY_VALUE"] = "server-only-value"
captured: Dict[str, Any] = {}
base = BaseAWSLLM()
with patch.dict(os.environ, env, clear=True), patch(
"boto3.client", side_effect=lambda *a, **k: _capturing_sts_client(captured)
), patch("boto3.Session", return_value=MagicMock()):
with pytest.raises(AwsAuthError):
base.get_credentials(
aws_web_identity_token=token_ref,
aws_role_name="arn:aws:iam::123456789012:role/x",
aws_session_name="s",
aws_sts_endpoint="https://custom-sts.example",
)
assert "server-only-value" not in str(captured)
def test_web_identity_token_oidc_reference_still_resolved():
"""
The env-reference guard does not over-reject: an oidc/ reference still flows to
get_secret (mocked to None here), surfacing the existing 401 rather than the 400
used for rejected env-var references.
"""
base = BaseAWSLLM()
env = _os_environ_without_aws_keys()
with patch.dict(os.environ, env, clear=True), patch(
"litellm.llms.bedrock.base_aws_llm.get_secret", return_value=None
):
with pytest.raises(AwsAuthError) as exc:
base.get_credentials(
aws_web_identity_token="oidc/circleci/",
aws_role_name="arn:aws:iam::123456789012:role/x",
aws_session_name="s",
)
assert exc.value.status_code == 401
def test_web_identity_path_not_cached_in_iam_cache():
base = BaseAWSLLM()
with patch.object(

View file

@ -3,25 +3,190 @@ import pytest
from litellm.llms.cloudflare.chat.transformation import CloudflareChatConfig
def test_get_complete_url_encodes_model_path_segment():
def test_supported_params_include_tools_and_tool_choice():
config = CloudflareChatConfig()
assert (
config.get_complete_url(
api_base="https://api.cloudflare.com/client/v4/accounts/acct/ai/run/",
api_key="cf-key",
model="@cf/meta/llama?x=1#frag",
optional_params={},
litellm_params={},
)
== "https://api.cloudflare.com/client/v4/accounts/acct/ai/run/%40cf/meta/llama%3Fx%3D1%23frag"
params = config.get_supported_openai_params(model="@cf/meta/llama-2-7b-chat-int8")
assert "tools" in params
assert "tool_choice" in params
assert "stream" in params
assert "max_tokens" in params
def test_get_complete_url_defaults_to_openai_compatible_endpoint(monkeypatch):
monkeypatch.setenv("CLOUDFLARE_ACCOUNT_ID", "acct")
config = CloudflareChatConfig()
url = config.get_complete_url(
api_base=None,
api_key="cf-key",
model="@cf/meta/llama-2-7b-chat-int8",
optional_params={},
litellm_params={},
)
with pytest.raises(ValueError, match="dot path segment"):
assert (
url
== "https://api.cloudflare.com/client/v4/accounts/acct/ai/v1/chat/completions"
)
assert "/ai/run/" not in url
def test_get_complete_url_appends_chat_completions_to_explicit_base():
config = CloudflareChatConfig()
url = config.get_complete_url(
api_base="https://api.cloudflare.com/client/v4/accounts/acct/ai/v1",
api_key="cf-key",
model="@cf/meta/llama-2-7b-chat-int8",
optional_params={},
litellm_params={},
)
assert (
url
== "https://api.cloudflare.com/client/v4/accounts/acct/ai/v1/chat/completions"
)
assert "/ai/run/" not in url
def test_get_complete_url_is_idempotent_for_full_base():
config = CloudflareChatConfig()
url = config.get_complete_url(
api_base="https://api.cloudflare.com/client/v4/accounts/acct/ai/v1/chat/completions",
api_key="cf-key",
model="@cf/meta/llama-2-7b-chat-int8",
optional_params={},
litellm_params={},
)
assert (
url
== "https://api.cloudflare.com/client/v4/accounts/acct/ai/v1/chat/completions"
)
def test_get_complete_url_falls_back_to_account_id_when_base_is_empty(monkeypatch):
monkeypatch.setenv("CLOUDFLARE_ACCOUNT_ID", "acct")
config = CloudflareChatConfig()
url = config.get_complete_url(
api_base="",
api_key="cf-key",
model="@cf/meta/llama-2-7b-chat-int8",
optional_params={},
litellm_params={},
)
assert (
url
== "https://api.cloudflare.com/client/v4/accounts/acct/ai/v1/chat/completions"
)
def test_get_complete_url_raises_when_account_id_and_base_missing(monkeypatch):
monkeypatch.delenv("CLOUDFLARE_ACCOUNT_ID", raising=False)
config = CloudflareChatConfig()
with pytest.raises(ValueError, match="Missing CLOUDFLARE_ACCOUNT_ID"):
config.get_complete_url(
api_base="https://api.cloudflare.com/client/v4/accounts/acct/ai/run/",
api_base=None,
api_key="cf-key",
model="../../accounts/other",
model="@cf/meta/llama-2-7b-chat-int8",
optional_params={},
litellm_params={},
)
def test_get_complete_url_raises_when_account_id_is_empty(monkeypatch):
monkeypatch.setenv("CLOUDFLARE_ACCOUNT_ID", " ")
config = CloudflareChatConfig()
with pytest.raises(ValueError, match="Missing CLOUDFLARE_ACCOUNT_ID"):
config.get_complete_url(
api_base=None,
api_key="cf-key",
model="@cf/meta/llama-2-7b-chat-int8",
optional_params={},
litellm_params={},
)
def test_get_complete_url_migrates_legacy_ai_run_base():
config = CloudflareChatConfig()
url = config.get_complete_url(
api_base="https://api.cloudflare.com/client/v4/accounts/acct/ai/run/",
api_key="cf-key",
model="@cf/meta/llama-2-7b-chat-int8",
optional_params={},
litellm_params={},
)
assert (
url
== "https://api.cloudflare.com/client/v4/accounts/acct/ai/v1/chat/completions"
)
assert "/ai/run" not in url
def test_transform_request_passes_tools_through_in_openai_format():
config = CloudflareChatConfig()
tools = [
{
"type": "function",
"function": {
"name": "get_weather",
"parameters": {
"type": "object",
"properties": {"city": {"type": "string"}},
},
},
}
]
messages = [{"role": "user", "content": "weather in nyc?"}]
body = config.transform_request(
model="@cf/meta/llama-2-7b-chat-int8",
messages=messages,
optional_params={"tools": tools, "tool_choice": "auto"},
litellm_params={},
headers={},
)
assert body["messages"] == messages
assert body["model"] == "@cf/meta/llama-2-7b-chat-int8"
assert body["tools"] == tools
assert body["tool_choice"] == "auto"
def test_validate_environment_requires_api_key():
config = CloudflareChatConfig()
with pytest.raises(ValueError, match="Missing Cloudflare API Key"):
config.validate_environment(
headers={},
model="@cf/meta/llama-2-7b-chat-int8",
messages=[],
optional_params={},
litellm_params={},
api_key=None,
)
def test_validate_environment_sets_bearer_and_content_type():
config = CloudflareChatConfig()
headers = config.validate_environment(
headers={},
model="@cf/meta/llama-2-7b-chat-int8",
messages=[],
optional_params={},
litellm_params={},
api_key="cf-key",
)
assert headers["Authorization"] == "Bearer cf-key"
assert headers["Content-Type"] == "application/json"

View file

@ -10,13 +10,21 @@ sys.path.insert(
0, os.path.abspath("../../../..")
) # Adds the parent directory to the system path
import litellm
from litellm.integrations.code_interpreter_interception.handler import (
CodeInterpreterInterceptionLogger,
LITELLM_CODE_EXECUTION_TOOL_NAME,
)
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import (
BaseLLMHTTPHandler,
_google_genai_streaming_hidden_params,
)
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.router import GenericLiteLLMParams
_ACTIVE_KEY = "_code_interpreter_interception_active"
_SANDBOX_KEY = "_code_interpreter_interception_sandbox_key"
def test_prepare_fake_stream_request():
# Initialize the BaseLLMHTTPHandler
@ -116,6 +124,117 @@ def test_response_api_handler_streams_when_provider_transform_adds_stream():
assert client.post.call_args.kwargs["json"]["stream"] is True
def test_response_api_handler_runs_agentic_hooks_in_sync_path(monkeypatch):
handler = BaseLLMHTTPHandler()
config = Mock()
config.validate_environment.return_value = {}
config.get_complete_url.return_value = "https://chatgpt.example.com/responses"
config.transform_responses_api_request.return_value = {
"model": "gpt-5",
"input": "hi",
}
config.sign_request.return_value = ({}, None)
initial_response = Mock()
final_response = Mock()
config.transform_response_api_response.return_value = initial_response
client = HTTPHandler(client=httpx.Client())
client.post = Mock(
return_value=httpx.Response(
200,
request=httpx.Request("POST", "https://chatgpt.example.com/responses"),
)
)
logging_obj = Mock()
monkeypatch.setattr(handler, "_has_agentic_completion_hook", Mock(return_value=True))
hook_mock = AsyncMock(return_value=final_response)
monkeypatch.setattr(handler, "_call_agentic_completion_hooks", hook_mock)
response = handler.response_api_handler(
model="gpt-5",
input="hi",
responses_api_provider_config=config,
response_api_optional_request_params={},
custom_llm_provider="openai",
litellm_params=GenericLiteLLMParams(),
logging_obj=logging_obj,
client=client,
)
assert response is final_response
hook_mock.assert_awaited_once()
assert hook_mock.call_args.kwargs["api_surface"] == "responses"
assert hook_mock.call_args.kwargs["messages"] == [
{"role": "user", "content": "hi"}
]
def test_response_api_handler_runs_responses_pre_call_hook_before_transform():
handler = BaseLLMHTTPHandler()
config = Mock()
config.validate_environment.return_value = {}
config.get_complete_url.return_value = "https://api.openai.com/v1/responses"
config.sign_request.return_value = ({}, None)
initial_response = ResponsesAPIResponse(
id="resp_1",
created_at=0,
output=[],
status="completed",
model="gpt-5",
)
config.transform_response_api_response.return_value = initial_response
def transform_responses_api_request(**kwargs):
return {
"model": kwargs["model"],
"input": kwargs["input"],
**kwargs["response_api_optional_request_params"],
}
config.transform_responses_api_request.side_effect = transform_responses_api_request
client = HTTPHandler(client=httpx.Client())
client.post = Mock(
return_value=httpx.Response(
200,
request=httpx.Request("POST", "https://api.openai.com/v1/responses"),
)
)
logging_obj = Mock()
logging_obj.dynamic_success_callbacks = []
old_callbacks = list(litellm.callbacks)
litellm.callbacks = [CodeInterpreterInterceptionLogger()]
try:
response = handler.response_api_handler(
model="gpt-5",
input="use code",
responses_api_provider_config=config,
response_api_optional_request_params={
"tools": [{"type": "code_interpreter", "container": {"type": "auto"}}]
},
custom_llm_provider="openai",
litellm_params=GenericLiteLLMParams(api_key="sk-test"),
logging_obj=logging_obj,
client=client,
)
finally:
litellm.callbacks = old_callbacks
assert response is initial_response
transform_kwargs = config.transform_responses_api_request.call_args.kwargs
tools = transform_kwargs["response_api_optional_request_params"]["tools"]
assert not any(tool.get("type") == "code_interpreter" for tool in tools)
assert any(
tool.get("type") == "function"
and tool.get("name") == LITELLM_CODE_EXECUTION_TOOL_NAME
for tool in tools
)
hook_litellm_params = transform_kwargs["litellm_params"]
assert hook_litellm_params.get(_ACTIVE_KEY) is True
assert hook_litellm_params.get(_SANDBOX_KEY)
@pytest.mark.asyncio
async def test_async_response_api_handler_streams_when_provider_transform_adds_stream():
handler = BaseLLMHTTPHandler()

Some files were not shown because too many files have changed in this diff Show more