Merge branch 'litellm_internal_staging' of https://github.com/BerriAI/litellm into pr30587

This commit is contained in:
yucheng-berri 2026-06-23 12:14:15 -07:00
commit 458f086df8
134 changed files with 14045 additions and 3555 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

@ -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

@ -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

@ -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
@ -3352,6 +3378,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)
@ -4006,6 +4045,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
@ -4123,6 +4168,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

@ -2954,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],
@ -2971,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
@ -3663,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

@ -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

@ -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

@ -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:

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

@ -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

@ -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

@ -162,7 +162,8 @@ class TestResponseCompliance:
# 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",

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

@ -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()

View file

@ -719,3 +719,93 @@ class TestMistralFileHandling:
# Check that file_ids are modified to match Mistral's expected format
assert result[0]["content"][1]["file_id"] == "file-12345" # type: ignore
assert result[0]["content"][2]["file_id"] == "file-67890" # type: ignore
class TestMistralStripsOutputOnlyFields:
"""Mistral rejects unknown input fields with a 422 ``extra_forbidden``.
LiteLLM attaches ``reasoning_content`` / ``thinking_blocks`` to assistant
responses, so replaying an assistant turn verbatim must not forward them.
Regression for https://github.com/BerriAI/litellm/issues/30835.
"""
def test_assistant_reasoning_content_is_dropped(self):
messages = cast(
List[AllMessageValues],
[
{"role": "user", "content": "Question?"},
{
"role": "assistant",
"content": "Follow-up",
"reasoning_content": "Some internal reasoning text.",
"thinking_blocks": [
{"type": "thinking", "thinking": "step", "signature": "mistral"}
],
},
],
)
result = cast(
List[AllMessageValues],
MistralConfig()._transform_messages(
messages=messages, model="mistral-medium-3-5"
),
)
assistant_message = result[-1]
assert "reasoning_content" not in assistant_message
assert "thinking_blocks" not in assistant_message
assert assistant_message["content"] == "Follow-up"
assert assistant_message["role"] == "assistant"
def test_non_assistant_messages_are_untouched(self):
messages = cast(
List[AllMessageValues],
[{"role": "user", "content": "Question?", "reasoning_content": "noise"}],
)
result = cast(
List[AllMessageValues],
MistralConfig()._transform_messages(
messages=messages, model="mistral-medium-3-5"
),
)
assert result[0].get("reasoning_content") == "noise"
def test_reasoning_content_dropped_when_image_present(self):
"""The image branch returns early, so stripping must run before it."""
messages = cast(
List[AllMessageValues],
[
{
"role": "user",
"content": [
{"type": "text", "text": "Describe this"},
{
"type": "image_url",
"image_url": {"url": "https://example.com/cat.png"},
},
],
},
{
"role": "assistant",
"content": "A cat.",
"reasoning_content": "leaked reasoning",
},
],
)
with patch.object(
MistralConfig,
"_transform_messages_sync",
side_effect=lambda transformed, model: transformed,
):
result = cast(
List[AllMessageValues],
MistralConfig()._transform_messages(
messages=messages, model="mistral-medium-3-5", is_async=False
),
)
assert "reasoning_content" not in result[-1]

View file

@ -2,6 +2,7 @@
Tests for JSON-based provider configuration system.
"""
import json
import os
import sys
from unittest.mock import MagicMock, patch
@ -244,6 +245,99 @@ class TestPinstripes:
assert result["temperature"] == 0.7
class TestDarkbloom:
def test_darkbloom_json_config_exists(self):
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
darkbloom = JSONProviderRegistry.get("darkbloom")
assert darkbloom is not None
assert darkbloom.base_url == "https://api.darkbloom.dev/v1"
assert darkbloom.api_key_env == "DARKBLOOM_API_KEY"
assert darkbloom.api_base_env == "DARKBLOOM_API_BASE"
assert darkbloom.param_mappings.get("max_completion_tokens") == "max_tokens"
def test_darkbloom_provider_resolution(self):
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
model, provider, api_key, api_base = get_llm_provider(
model="darkbloom/gemma-4-26b",
custom_llm_provider=None,
api_base=None,
api_key=None,
)
assert model == "gemma-4-26b"
assert provider == "darkbloom"
assert api_key is None
assert api_base == "https://api.darkbloom.dev/v1"
def test_darkbloom_dynamic_config(self):
from litellm.llms.openai_like.dynamic_config import create_config_class
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
provider = JSONProviderRegistry.get("darkbloom")
config_class = create_config_class(provider)
config = config_class()
api_base, api_key = config._get_openai_compatible_provider_info(None, None)
assert api_base == "https://api.darkbloom.dev/v1"
api_base, api_key = config._get_openai_compatible_provider_info(
"https://custom.darkbloom.dev/v1", "test-key"
)
assert api_base == "https://custom.darkbloom.dev/v1"
assert api_key == "test-key"
def test_darkbloom_complete_url_appends_endpoint(self):
from litellm.llms.openai_like.dynamic_config import create_config_class
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
provider = JSONProviderRegistry.get("darkbloom")
config_class = create_config_class(provider)
config = config_class()
url = config.get_complete_url(
api_base="https://api.darkbloom.dev/v1",
api_key="test-key",
model="darkbloom/gemma-4-26b",
optional_params={},
litellm_params={},
stream=True,
)
assert url == "https://api.darkbloom.dev/v1/chat/completions"
def test_darkbloom_provider_config_manager(self):
from litellm import LlmProviders
from litellm.utils import ProviderConfigManager
config = ProviderConfigManager.get_provider_chat_config(
model="gemma-4-26b", provider=LlmProviders.DARKBLOOM
)
assert config is not None
assert config.custom_llm_provider == "darkbloom"
def test_darkbloom_model_cost_map(self):
with open(
os.path.join(workspace_path, "model_prices_and_context_window.json")
) as f:
model_cost = json.load(f)
expected_models = {
"darkbloom/gemma-4-26b": (3e-08, 1.65e-07),
"darkbloom/gpt-oss-20b": (1.45e-08, 7e-08),
}
for model, (input_cost, output_cost) in expected_models.items():
assert model in model_cost
assert model_cost[model]["litellm_provider"] == "darkbloom"
assert model_cost[model]["max_output_tokens"] == 32768
assert model_cost[model]["supports_function_calling"] is True
assert model_cost[model]["supports_tool_choice"] is True
assert model_cost[model]["input_cost_per_token"] == input_cost
assert model_cost[model]["output_cost_per_token"] == output_cost
class TestPublicAIIntegration:
"""Integration tests for PublicAI provider"""

View file

@ -9,7 +9,7 @@ import json
import math
import os
import sys
from unittest.mock import Mock, patch
from unittest.mock import patch
import pytest
@ -120,10 +120,10 @@ class TestPerplexityCostCalculator:
# Expected costs:
# Input: 100 tokens * $2e-6 = $0.0002
# Output: 50 tokens * $8e-6 = $0.0004
# Search: 3 queries * ($0.005 / 1000) = $0.000015
# Total completion cost: $0.000415
# Search: 3 queries * $0.005 per request = $0.015
# Total completion cost: $0.0154
expected_prompt_cost = 100 * 2e-6
expected_completion_cost = (50 * 8e-6) + (3 / 1000 * 0.005)
expected_completion_cost = (50 * 8e-6) + (3 * 0.005)
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6)
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-6)
@ -195,10 +195,10 @@ class TestPerplexityCostCalculator:
# Total prompt cost = $0.00026
# Output (text): (50 - 15) tokens * $8e-6 = $0.00028
# Reasoning: 15 tokens * $3e-6 = $0.000045
# Search: 2 queries * ($0.005 / 1000) = $0.00001
# Total completion cost = $0.000335
# Search: 2 queries * $0.005 per request = $0.01
# Total completion cost = $0.010325
expected_prompt_cost = (100 * 2e-6) + (30 * 2e-6)
expected_completion_cost = ((50 - 15) * 8e-6) + (15 * 3e-6) + (2 / 1000 * 0.005)
expected_completion_cost = ((50 - 15) * 8e-6) + (15 * 3e-6) + (2 * 0.005)
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6)
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-6)
@ -311,7 +311,7 @@ class TestPerplexityCostCalculator:
# Calculate expected total cost (reasoning is a subset of completion_tokens)
expected_prompt_cost = (100 * 2e-6) + (15 * 2e-6) # Input + citation
expected_completion_cost = (
((50 - 10) * 8e-6) + (10 * 3e-6) + (1 / 1000 * 0.005)
((50 - 10) * 8e-6) + (10 * 3e-6) + (1 * 0.005)
) # Output (text) + reasoning + search
expected_total = expected_prompt_cost + expected_completion_cost
@ -361,7 +361,7 @@ class TestPerplexityCostCalculator:
expected_completion_cost = (
((50 - reasoning_tokens) * 8e-6)
+ (reasoning_tokens * 3e-6)
+ (search_queries / 1000 * 0.005)
+ (search_queries * 0.005)
)
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6)

View file

@ -9,7 +9,6 @@ import json
import math
import os
import sys
from unittest.mock import Mock, patch
import pytest
@ -106,8 +105,8 @@ class TestPerplexityIntegration:
expected_prompt_cost = (100 * 2e-6) + (citation_tokens * 2e-6)
expected_completion_cost = (
((50 - 10) * 8e-6) + (10 * 3e-6) + (2 / 1000 * 0.005)
)
((50 - 10) * 8e-6) + (10 * 3e-6) + (2 * 0.005)
) # Output (text) + reasoning + search
expected_total = expected_prompt_cost + expected_completion_cost
assert math.isclose(total_cost, expected_total, rel_tol=1e-6)
@ -152,8 +151,8 @@ class TestPerplexityIntegration:
expected_prompt_cost = (200 * 2e-6) + (40 * 2e-6)
expected_completion_cost = (
((100 - 25) * 8e-6) + (25 * 3e-6) + (3 / 1000 * 0.005)
)
((100 - 25) * 8e-6) + (25 * 3e-6) + (3 * 0.005)
) # Output (text) + reasoning + search
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6)
assert math.isclose(completion_cost_val, expected_completion_cost, rel_tol=1e-6)
@ -262,9 +261,9 @@ class TestPerplexityIntegration:
expected_prompt_cost = (50000 * 2e-6) + (5000 * 2e-6)
expected_completion_cost = (
((25000 - 10000) * 8e-6) + (10000 * 3e-6) + (100 / 1000 * 0.005)
)
expected_total = expected_prompt_cost + expected_completion_cost
((25000 - 10000) * 8e-6) + (10000 * 3e-6) + (100 * 0.005)
) # $0.65
expected_total = expected_prompt_cost + expected_completion_cost # $0.76
assert math.isclose(total_cost, expected_total, rel_tol=1e-6)
assert total_cost > 0.25
@ -326,7 +325,7 @@ class TestPerplexityIntegration:
# Should calculate costs correctly
expected_prompt_cost = (100 * 2e-6) + (10 * 2e-6)
expected_completion_cost = (50 * 8e-6) + (1 / 1000 * 0.005)
expected_completion_cost = (50 * 8e-6) + (1 * 0.005)
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6)
assert math.isclose(completion_cost_val, expected_completion_cost, rel_tol=1e-6)

View file

@ -346,3 +346,74 @@ def test_vertex_does_not_warn_when_dropping_non_guardrail_session_update(caplog)
"Vertex AI Realtime" in record.message and "session.update" in record.message
for record in caplog.records
)
async def test_async_realtime_does_not_forward_client_query_params_to_vertex_backend(
monkeypatch,
):
"""Regression: forwarding client ?model=/?intent= to the Vertex Live WSS URL causes 1007 errors.
Exercises ``async_realtime`` end-to-end so that re-adding ``_append_query_params``
(the reverted bug) would push ``model=``/``intent=`` onto the backend URL and fail here.
"""
import websockets
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
cfg = VertexAIRealtimeConfig(
access_token="tok", project="my-proj", location="us-central1"
)
captured = {}
def fake_connect(url, *args, **kwargs):
captured["url"] = url
raise RuntimeError("stop before establishing the backend connection")
monkeypatch.setattr(websockets, "connect", fake_connect)
await BaseLLMHTTPHandler().async_realtime(
model="gemini-live-2.5-flash-native-audio",
websocket=AsyncMock(),
logging_obj=MagicMock(),
provider_config=cfg,
headers={},
query_params={
"model": "gemini-live-2.5-flash-native-audio",
"intent": "chat",
},
)
assert "?" not in captured["url"]
assert "model=" not in captured["url"]
assert "intent=" not in captured["url"]
def test_vertex_function_call_output_omits_id():
"""Regression: Vertex Live rejects ``id`` on toolResponse.functionResponses (1007)."""
cfg = VertexAIRealtimeConfig(
access_token="tok", project="my-proj", location="us-central1"
)
cfg._tool_call_id_to_name["call_abc123"] = "terminate_call"
messages = cfg.transform_realtime_request(
json.dumps(
{
"type": "conversation.item.create",
"item": {
"type": "function_call_output",
"call_id": "call_abc123",
"output": '{"status": "ok"}',
},
}
),
"gemini-live-2.5-flash-native-audio",
session_configuration_request="existing",
)
assert len(messages) == 1
payload = json.loads(messages[0])
function_response = payload["toolResponse"]["functionResponses"][0]
assert "id" not in function_response
assert function_response["name"] == "terminate_call"
assert function_response["response"] == {"status": "ok"}

View file

@ -15,7 +15,11 @@ from starlette.datastructures import Headers
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
)
from litellm.proxy._types import SpecialHeaders, UserAPIKeyAuth
from litellm.proxy._types import (
SpecialHeaders,
SpecialMCPServerNames,
UserAPIKeyAuth,
)
@pytest.mark.asyncio
@ -166,6 +170,53 @@ class TestMCPRequestHandler:
mock_key_servers.assert_called_once_with(user_api_key_auth)
mock_team_servers.assert_called_once_with(user_api_key_auth)
@pytest.mark.parametrize("team_servers", [[], ["team_server1", "team_server2"]])
async def test_no_mcp_servers_sentinel_returns_empty(self, team_servers):
"""A key scoped to the no-mcp-servers sentinel resolves to zero servers,
overriding team inheritance and never leaking the sentinel marker."""
user_api_key_auth = UserAPIKeyAuth(
api_key="test-key", user_id="test-user", team_id="test-team"
)
key_object_permission = MagicMock()
key_object_permission.mcp_servers = [
SpecialMCPServerNames.no_mcp_servers.value
]
with patch.object(
MCPRequestHandler,
"_get_key_object_permission",
return_value=key_object_permission,
), patch.object(
MCPRequestHandler,
"_get_allowed_mcp_servers_for_team",
new_callable=AsyncMock,
return_value=team_servers,
):
result = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth)
assert result == []
async def test_get_allowed_mcp_servers_for_key_returns_sentinel_marker(self):
"""_get_allowed_mcp_servers_for_key surfaces the sentinel unexpanded so the
caller can short-circuit, ignoring any other entries on the key."""
user_api_key_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user")
key_object_permission = MagicMock()
key_object_permission.mcp_servers = [
SpecialMCPServerNames.no_mcp_servers.value,
"some-other-server",
]
with patch.object(
MCPRequestHandler,
"_get_key_object_permission",
return_value=key_object_permission,
):
result = await MCPRequestHandler._get_allowed_mcp_servers_for_key(
user_api_key_auth
)
assert result == [SpecialMCPServerNames.no_mcp_servers.value]
async def test_permission_inheritance_edge_cases(self):
"""Test edge cases in permission inheritance"""

View file

@ -0,0 +1,44 @@
"""Tests for the concrete httpx.Auth objects the resolver returns.
NoOpAuth must attach nothing; StaticHeaderAuth must set exactly the configured header. These
pin the header emission the api_key family and passthrough depend on.
"""
import httpx
from litellm.proxy._experimental.mcp_server.outbound_credentials import (
NoOpAuth,
StaticHeaderAuth,
)
def _apply(auth: httpx.Auth, request: httpx.Request) -> httpx.Request:
flow = auth.auth_flow(request)
sent = next(flow)
flow.close()
return sent
def test_noop_auth_attaches_no_authorization_header():
request = httpx.Request("GET", "https://upstream.example.com/mcp")
_apply(NoOpAuth(), request)
assert "authorization" not in request.headers
def test_static_header_auth_defaults_to_authorization():
request = httpx.Request("GET", "https://upstream.example.com/mcp")
_apply(StaticHeaderAuth("Bearer abc"), request)
assert request.headers["Authorization"] == "Bearer abc"
def test_static_header_auth_honors_custom_header_name():
request = httpx.Request("GET", "https://upstream.example.com/mcp")
_apply(StaticHeaderAuth("raw-key", header_name="X-API-Key"), request)
assert request.headers["X-API-Key"] == "raw-key"
assert "authorization" not in request.headers
def test_static_header_auth_masks_credential_from_introspection():
auth = StaticHeaderAuth("Bearer super-secret-token")
assert "super-secret-token" not in repr(auth)
assert "super-secret-token" not in str(vars(auth))

View file

@ -0,0 +1,57 @@
"""Tests for the resolver dispatch skeleton.
Every mode must reach its own arm and, until that arm is built, return a typed
`not_implemented` CredError rather than silently producing no credential. Parametrizing over
one config per mode also guards reachability: if a `case` were dropped, that mode would fall to
the `assert_never` tail and raise here instead of returning the stub.
"""
import pytest
from pydantic import SecretStr
from litellm.proxy._experimental.mcp_server.outbound_credentials import (
ApiKeyConfig,
AuthorizationCodeConfig,
AuthSpecKind,
AwsSigV4Config,
ClientCredentialsConfig,
Error,
NoneConfig,
PassthroughConfig,
ServerSpec,
SharedKey,
Subject,
TokenExchangeConfig,
UpstreamCredentialProvider,
)
_ONE_CONFIG_PER_MODE = [
(AuthSpecKind.none, NoneConfig()),
(AuthSpecKind.api_key, ApiKeyConfig(key_source=SharedKey(value=SecretStr("k")))),
(AuthSpecKind.passthrough, PassthroughConfig()),
(AuthSpecKind.client_credentials, ClientCredentialsConfig()),
(AuthSpecKind.token_exchange, TokenExchangeConfig()),
(AuthSpecKind.authorization_code, AuthorizationCodeConfig()),
(AuthSpecKind.aws_sigv4, AwsSigV4Config(region="us-east-1")),
]
@pytest.mark.asyncio
@pytest.mark.parametrize("kind, config", _ONE_CONFIG_PER_MODE)
async def test_every_mode_reaches_its_arm_and_returns_not_implemented(kind, config):
spec = ServerSpec(
server_id="s", resource="https://upstream.example.com", config=config
)
subject = Subject(tenant_id="", subject_id="")
result = await UpstreamCredentialProvider().resolve_credentials(subject, spec)
assert isinstance(result, Error)
assert result.error.tag == "not_implemented"
assert kind.value in result.error.summary
def test_all_seven_modes_are_covered():
# Guards that the parametrization (and therefore the dispatch) spans every AuthSpecKind, so a
# newly added mode without a test row is caught here rather than slipping through.
assert {kind for kind, _ in _ONE_CONFIG_PER_MODE} == set(AuthSpecKind)

View file

@ -0,0 +1,20 @@
"""Smoke test for the outbound_credentials Result union.
Result is trivial frozen dataclasses; its load-bearing guarantee (no `.ok` access before
the Error arm is eliminated) is a type-checker property, not a runtime one. This pins only
the runtime contract consumers rely on: each arm carries its payload and discriminates by
type. The union is exercised for real where it is used (see PR2's parse_auth_spec_kind).
"""
from litellm.proxy._experimental.mcp_server.outbound_credentials import (
Error,
Ok,
Result,
)
def test_ok_and_error_carry_payload_and_discriminate():
ok: Result[int, str] = Ok(5)
err: Result[int, str] = Error("boom")
assert isinstance(ok, Ok) and ok.ok == 5
assert isinstance(err, Error) and err.error == "boom"

View file

@ -0,0 +1,150 @@
"""Construction-time tests for the outbound_credentials vocabulary.
The point of the typed seam is that illegal mode/field combinations are unrepresentable:
a config missing a required field, an unknown mode, or a mismatched discriminated-union
source must fail at construction, not at resolve time. These tests pin that, plus the
CredError tag/summary surface and the derived auth_spec_kind. Each assertion fails if the
corresponding guarantee is mutated away.
"""
import pytest
from pydantic import SecretStr, TypeAdapter, ValidationError
from litellm.proxy._experimental.mcp_server.outbound_credentials import (
Ambient,
ApiKeyConfig,
AuthConfig,
AuthSpecKind,
AwsSigV4Config,
Byok,
CredError,
Error,
NoneConfig,
Ok,
ServerSpec,
SharedKey,
StaticKeys,
parse_auth_spec_kind,
)
_AUTH_CONFIG = TypeAdapter(AuthConfig)
def test_parse_auth_spec_kind_accepts_known_mode():
result = parse_auth_spec_kind("token_exchange")
assert isinstance(result, Ok)
assert result.ok is AuthSpecKind.token_exchange
def test_parse_auth_spec_kind_rejects_unknown_mode():
result = parse_auth_spec_kind("totally_made_up")
assert isinstance(result, Error)
assert result.error.tag == "unsupported_mode"
assert "totally_made_up" in result.error.summary
@pytest.mark.parametrize(
"factory, expected_tag",
[
(CredError.of_unauthorized, "unauthorized"),
(CredError.of_misconfigured, "misconfigured"),
(CredError.of_upstream_unavailable, "upstream_unavailable"),
(CredError.of_unsupported_mode, "unsupported_mode"),
(CredError.of_precondition_required, "precondition_required"),
(CredError.of_not_implemented, "not_implemented"),
],
)
def test_crederror_factory_sets_the_matching_tag(factory, expected_tag):
err = factory("detail text")
assert err.tag == expected_tag
assert "detail text" in err.summary
def test_apikeyconfig_requires_a_key_source():
with pytest.raises(ValidationError):
ApiKeyConfig() # type: ignore[call-arg]
def test_sharedkey_requires_a_value():
with pytest.raises(ValidationError):
SharedKey() # type: ignore[call-arg]
def test_static_keys_require_id_and_secret():
with pytest.raises(ValidationError):
StaticKeys(access_key_id="AKIA") # type: ignore[call-arg]
def test_aws_sigv4_requires_a_region():
with pytest.raises(ValidationError):
AwsSigV4Config() # type: ignore[call-arg]
def test_aws_sigv4_defaults_to_the_ambient_credential_chain():
cfg = AwsSigV4Config(region="us-east-1")
assert isinstance(cfg.credentials, Ambient)
assert cfg.service == "bedrock-agentcore"
def test_authconfig_discriminates_on_kind():
api_key = _AUTH_CONFIG.validate_python(
{"kind": "api_key", "key_source": {"source": "shared", "value": "k"}}
)
assert isinstance(api_key, ApiKeyConfig)
assert isinstance(api_key.key_source, SharedKey)
none = _AUTH_CONFIG.validate_python({"kind": "none"})
assert isinstance(none, NoneConfig)
def test_authconfig_rejects_unknown_kind():
with pytest.raises(ValidationError):
_AUTH_CONFIG.validate_python({"kind": "not_a_mode"})
def test_apikeysource_discriminates_and_rejects_unknown_source():
byok = ApiKeyConfig.model_validate({"key_source": {"source": "byok"}})
assert isinstance(byok.key_source, Byok)
with pytest.raises(ValidationError):
ApiKeyConfig.model_validate({"key_source": {"source": "mystery"}})
def test_server_spec_derives_auth_spec_kind_from_config():
spec = ServerSpec(
server_id="s1",
resource="https://api.example.com",
config=NoneConfig(),
)
assert spec.auth_spec_kind is AuthSpecKind.none
api_spec = ServerSpec(
server_id="s2",
resource="https://api.example.com",
config=ApiKeyConfig(key_source=SharedKey(value=SecretStr("k"))),
)
assert api_spec.auth_spec_kind is AuthSpecKind.api_key
def test_api_key_header_placement():
default = ApiKeyConfig(key_source=SharedKey(value=SecretStr("tok")))
assert default.header("tok") == ("Authorization", "Bearer tok")
raw = ApiKeyConfig(
header_name="X-API-Key",
value_prefix="",
key_source=SharedKey(value=SecretStr("tok")),
)
assert raw.header("tok") == ("X-API-Key", "tok")
def test_configs_are_frozen():
cfg = NoneConfig()
with pytest.raises(ValidationError):
cfg.kind = AuthSpecKind.api_key # type: ignore[misc]
def test_secrets_do_not_leak_in_repr():
key = SharedKey(value=SecretStr("super-secret"))
assert "super-secret" not in repr(key)
assert key.value.get_secret_value() == "super-secret"

View file

@ -41,9 +41,13 @@ class TestMask:
def test_empty_returns_none_label(self):
assert MCPDebug._mask("") == "(none)"
def test_short_value_unchanged(self):
# visible_prefix=6 + visible_suffix=4 = 10, so <= 10 chars unchanged
assert MCPDebug._mask("sk-1234") == "sk-1234"
def test_short_value_masked(self):
# Short auth values must not be echoed verbatim in debug headers, even though
# visible_prefix + visible_suffix would otherwise reveal the whole value.
masked = MCPDebug._mask("sk-1234")
assert "sk-1234" not in masked
assert set(masked) == {"*"}
assert len(masked) == len("sk-1234")
def test_long_value_masked(self):
result = MCPDebug._mask("Bearer eyJhbGciOiJSUzI1NiIsInR5cCI6IkpXVCJ9")

View file

@ -6293,3 +6293,153 @@ async def test_get_allowed_mcp_servers_from_mcp_server_names_empty_list_fails_cl
)
assert result == []
class TestProxyExceptionToHttpException:
"""Auth failures reach the MCP ASGI handlers as ProxyException, not
HTTPException. The handlers must map them back to their real status and
headers; otherwise they fall through to the generic 500 handler, dropping
the 401 + WWW-Authenticate challenge an OAuth client needs to re-authenticate
and surfacing the tool call as a cancelled/terminated session.
"""
def test_preserves_401_status_and_www_authenticate_header(self):
from litellm.proxy._experimental.mcp_server.server import (
_proxy_exception_to_http_exception,
)
from litellm.proxy._types import ProxyException
exc = ProxyException(
message="Authentication Error, invalid token",
type="auth_error",
param="key",
code=401,
headers={"WWW-Authenticate": 'Bearer resource_metadata="/x"'},
)
http_exc = _proxy_exception_to_http_exception(exc)
assert http_exc.status_code == 401
assert http_exc.detail == "Authentication Error, invalid token"
assert http_exc.headers["WWW-Authenticate"] == 'Bearer resource_metadata="/x"'
def test_preserves_403_status(self):
from litellm.proxy._experimental.mcp_server.server import (
_proxy_exception_to_http_exception,
)
from litellm.proxy._types import ProxyException
http_exc = _proxy_exception_to_http_exception(
ProxyException(
message="Forbidden", type="auth_error", param="key", code=403
)
)
assert http_exc.status_code == 403
def test_non_numeric_code_falls_back_to_500(self):
from litellm.proxy._experimental.mcp_server.server import (
_proxy_exception_to_http_exception,
)
from litellm.proxy._types import ProxyException
# ProxyException normalises code to the string "None" when unset.
http_exc = _proxy_exception_to_http_exception(
ProxyException(message="boom", type="server_error", param=None, code=None)
)
assert http_exc.status_code == 500
class TestStreamableHttpAuthErrorMapping:
"""End-to-end guard for the handler wiring: a ProxyException from auth must
propagate as the real HTTPException (401 + WWW-Authenticate), not be
flattened to a generic 500 by the catch-all handler.
"""
@pytest.mark.asyncio
async def test_streamable_http_propagates_proxy_exception_as_401(self):
from litellm.proxy._experimental.mcp_server import server as mcp_module
from litellm.proxy._types import ProxyException
scope = {
"type": "http",
"method": "POST",
"path": "/mcp/some_server",
"headers": [(b"x-litellm-api-key", b"sk-bad")],
}
async def receive():
return {"type": "http.request", "body": b"{}", "more_body": False}
sent = []
async def send(message):
sent.append(message)
auth_failure = ProxyException(
message="Authentication Error, invalid token",
type="auth_error",
param="key",
code=401,
headers={"WWW-Authenticate": "Bearer"},
)
with patch.object(
mcp_module,
"extract_mcp_auth_context",
new=AsyncMock(side_effect=auth_failure),
):
with pytest.raises(HTTPException) as exc_info:
await mcp_module.handle_streamable_http_mcp(scope, receive, send)
assert exc_info.value.status_code == 401
assert exc_info.value.headers["WWW-Authenticate"] == "Bearer"
# Must not have emitted a 500 body via the generic catch-all.
assert not any(
m.get("type") == "http.response.start" and m.get("status") == 500
for m in sent
)
@pytest.mark.asyncio
async def test_sse_propagates_proxy_exception_as_401(self):
from litellm.proxy._experimental.mcp_server import server as mcp_module
from litellm.proxy._types import ProxyException
scope = {
"type": "http",
"method": "GET",
"path": "/mcp/some_server",
"headers": [(b"x-litellm-api-key", b"sk-bad")],
}
async def receive():
return {"type": "http.request", "body": b"", "more_body": False}
sent = []
async def send(message):
sent.append(message)
auth_failure = ProxyException(
message="Authentication Error, invalid token",
type="auth_error",
param="key",
code=401,
headers={"WWW-Authenticate": "Bearer"},
)
with patch.object(
mcp_module,
"extract_mcp_auth_context",
new=AsyncMock(side_effect=auth_failure),
):
with pytest.raises(HTTPException) as exc_info:
await mcp_module.handle_sse_mcp(scope, receive, send)
assert exc_info.value.status_code == 401
assert exc_info.value.headers["WWW-Authenticate"] == "Bearer"
assert not any(
m.get("type") == "http.response.start" and m.get("status") == 500
for m in sent
)

View file

@ -2948,6 +2948,41 @@ class TestMCPServerManager:
assert "test_server_1" in result
assert "test_server_2" in result
@pytest.mark.asyncio
async def test_no_mcp_servers_sentinel_blocks_allow_all_keys(self):
"""A key scoped to no-mcp-servers gets zero servers even when allow_all_keys
servers exist, and the inner resolver is never consulted."""
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
)
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth
manager = MCPServerManager()
object_permission = LiteLLM_ObjectPermissionTable(
object_permission_id="perm_no_mcp",
mcp_servers=["no-mcp-servers"],
mcp_access_groups=[],
)
user_api_key_auth = UserAPIKeyAuth(
api_key="sk-test",
user_id="user-123",
object_permission=object_permission,
object_permission_id="perm_no_mcp",
)
with patch.object(
manager, "get_allow_all_keys_server_ids", return_value=["global-server"]
), patch.object(
MCPRequestHandler,
"get_allowed_mcp_servers",
new_callable=AsyncMock,
return_value=["leaked-server"],
) as mock_inner:
result = await manager.get_allowed_mcp_servers(user_api_key_auth)
assert result == []
mock_inner.assert_not_called()
@pytest.mark.asyncio
async def test_get_allowed_mcp_servers_anonymous_delegate_requires_oauth2(self):
"""Anonymous delegated auth listing should only include oauth2 servers."""

View file

@ -93,6 +93,37 @@ class TestApplyToolsetScope:
await _apply_toolset_scope(auth, "toolset-123")
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
@pytest.mark.parametrize("user_role", [None, LitellmUserRoles.PROXY_ADMIN.value])
async def test_no_mcp_servers_sentinel_denies_toolset_access(self, user_role):
"""A key scoped to the no-mcp-servers sentinel cannot reach a toolset it
would otherwise be granted (even as admin); the opt-out covers the
toolset path, which replaces mcp_servers and would drop the sentinel."""
from starlette.exceptions import HTTPException
from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope
op = LiteLLM_ObjectPermissionTable(
object_permission_id="test",
mcp_servers=["no-mcp-servers"],
mcp_toolsets=["toolset-123"],
)
auth = UserAPIKeyAuth(
api_key="sk-test", object_permission=op, user_role=user_role
)
resolve = AsyncMock(return_value={"server-a": ["tool1"]})
with patch(
"litellm.proxy._experimental.mcp_server.server."
"global_mcp_server_manager.resolve_toolset_tool_permissions",
new=resolve,
):
with pytest.raises(HTTPException) as exc_info:
await _apply_toolset_scope(auth, "toolset-123")
assert exc_info.value.status_code == 403
resolve.assert_not_awaited()
class TestFetchMCPToolsetsAccess:
"""Tests for GET /v1/mcp/toolset access control."""

View file

@ -452,6 +452,430 @@ async def test_semantic_filter_hook_skips_no_tools():
print("✅ Hook correctly skips requests without tools")
@pytest.mark.asyncio
async def test_semantic_filter_hook_preserves_native_tools():
"""
Regression test: mixed MCP + native tools.
Given: 5 MCP tools (registered in _tool_map) + 2 native OpenAI-format
function tools (not in _tool_map)
When: The hook filters tools
Then: The native tools must survive unconditionally, and only MCP
tools go through the semantic filter.
"""
from litellm.proxy._experimental.mcp_server.semantic_tool_filter import (
SemanticMCPToolFilter,
)
from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook
from litellm.types.utils import Embedding, EmbeddingResponse
mock_router = Mock()
def mock_embedding_sync(*args, **kwargs):
return EmbeddingResponse(
data=[Embedding(embedding=[0.1] * 1536, index=0, object="embedding")],
model="text-embedding-3-small",
object="list",
usage={"prompt_tokens": 10, "total_tokens": 10},
)
async def mock_embedding_async(*args, **kwargs):
return mock_embedding_sync()
mock_router.embedding = mock_embedding_sync
mock_router.aembedding = mock_embedding_async
filter_instance = SemanticMCPToolFilter(
embedding_model="text-embedding-3-small",
litellm_router_instance=mock_router,
top_k=2,
similarity_threshold=0.3,
enabled=True,
)
# --- MCP tools (registered in the semantic router) ---
mcp_tools = [
MCPTool(
name=f"mcp_tool_{i}",
description=f"MCP tool {i}",
inputSchema={"type": "object"},
)
for i in range(5)
]
filter_instance._build_router(mcp_tools)
# --- Native OpenAI-format function tools (NOT in _tool_map) ---
native_tools = [
{
"type": "function",
"function": {
"name": "get_current_weather",
"description": "Get the current weather",
"parameters": {"type": "object", "properties": {}},
},
},
{
"type": "function",
"function": {
"name": "search_web",
"description": "Search the web",
"parameters": {"type": "object", "properties": {}},
},
},
]
# Combine: MCP tools + native tools
all_tools = list(mcp_tools) + native_tools
hook = SemanticToolFilterHook(filter_instance)
data = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "What is the weather?"}],
"tools": all_tools,
"metadata": {},
}
result = await hook.async_pre_call_hook(
user_api_key_dict=Mock(),
cache=Mock(),
data=data,
call_type="completion",
)
assert result is not None, "Hook should return modified data"
filtered = result["tools"]
# Native tools must survive
native_in_result = [
t for t in filtered if isinstance(t, dict) and t.get("type") == "function"
]
assert (
len(native_in_result) == 2
), f"Both native tools must survive, got {len(native_in_result)}"
# MCP tools should be filtered (top_k=2)
mcp_in_result = [t for t in filtered if not isinstance(t, dict)]
assert (
len(mcp_in_result) <= 2
), f"MCP tools should be filtered to top_k=2, got {len(mcp_in_result)}"
# Total should be native + filtered MCP
assert len(filtered) <= 4, f"Expected at most 4 tools, got {len(filtered)}"
# Filter stats should be emitted (MCP tools were present)
assert "litellm_semantic_filter_stats" in result["metadata"]
# Stats should report MCP-only counts, not inflated with native tools
stats = result["metadata"]["litellm_semantic_filter_stats"]
mcp_before, mcp_after = stats.split("->")
assert (
int(mcp_before) == 5
), f"Stats 'from' should be MCP count (5), got {mcp_before}"
assert int(mcp_after) == len(
mcp_in_result
), f"Stats 'to' should match filtered MCP count, got {mcp_after}"
print(
f"✅ Hook preserves native tools: {len(all_tools)} -> {len(filtered)} "
f"({len(native_in_result)} native + {len(mcp_in_result)} MCP), "
f"stats={stats}"
)
@pytest.mark.asyncio
async def test_semantic_filter_hook_all_native_tools():
"""
Regression test: all-native request.
Given: Only native OpenAI-format function tools (none registered in
the MCP semantic router)
When: The hook processes the request
Then: All tools pass through, and NO spurious semantic filter response
headers are emitted (no litellm_semantic_filter_stats in metadata).
"""
from litellm.proxy._experimental.mcp_server.semantic_tool_filter import (
SemanticMCPToolFilter,
)
from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook
mock_router = Mock()
filter_instance = SemanticMCPToolFilter(
embedding_model="text-embedding-3-small",
litellm_router_instance=mock_router,
top_k=3,
similarity_threshold=0.3,
enabled=True,
)
# Build router with some MCP tools (so tool_router is not None)
mcp_tools = [
MCPTool(
name="some_mcp_tool",
description="An MCP tool",
inputSchema={"type": "object"},
)
]
from litellm.types.utils import Embedding, EmbeddingResponse
def mock_embedding_sync(*args, **kwargs):
return EmbeddingResponse(
data=[Embedding(embedding=[0.1] * 1536, index=0, object="embedding")],
model="text-embedding-3-small",
object="list",
usage={"prompt_tokens": 10, "total_tokens": 10},
)
async def mock_embedding_async(*args, **kwargs):
return mock_embedding_sync()
mock_router.embedding = mock_embedding_sync
mock_router.aembedding = mock_embedding_async
filter_instance._build_router(mcp_tools)
# --- Only native tools in the request ---
native_tools = [
{
"type": "function",
"function": {
"name": f"native_func_{i}",
"description": f"Native function {i}",
"parameters": {"type": "object", "properties": {}},
},
}
for i in range(3)
]
hook = SemanticToolFilterHook(filter_instance)
data = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "Hello"}],
"tools": native_tools,
"metadata": {},
}
result = await hook.async_pre_call_hook(
user_api_key_dict=Mock(),
cache=Mock(),
data=data,
call_type="completion",
)
assert result is not None, "Hook should return modified data"
filtered = result["tools"]
# All native tools must pass through
assert (
len(filtered) == 3
), f"All 3 native tools must pass through, got {len(filtered)}"
# No spurious semantic filter stats (P2 fix)
assert (
"litellm_semantic_filter_stats" not in result["metadata"]
), "Should NOT emit semantic filter stats for all-native-tool requests"
print(
f"✅ Hook passes through all {len(filtered)} native tools, "
f"no spurious filter headers emitted"
)
@pytest.mark.asyncio
async def test_semantic_filter_hook_responses_api_name_collision():
"""
Regression test: Responses API native tool with MCP-matching name.
Given: A Responses-API native tool whose top-level ``name`` collides
with an MCP canonical name in ``_tool_map``
When: The hook classifies tools
Then: The native tool must NOT be sent to the semantic filter, even
though its name matches an MCP canonical.
"""
from litellm.proxy._experimental.mcp_server.semantic_tool_filter import (
SemanticMCPToolFilter,
)
from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook
from litellm.types.utils import Embedding, EmbeddingResponse
mock_router = Mock()
def mock_embedding_sync(*args, **kwargs):
return EmbeddingResponse(
data=[Embedding(embedding=[0.1] * 1536, index=0, object="embedding")],
model="text-embedding-3-small",
object="list",
usage={"prompt_tokens": 10, "total_tokens": 10},
)
async def mock_embedding_async(*args, **kwargs):
return mock_embedding_sync()
mock_router.embedding = mock_embedding_sync
mock_router.aembedding = mock_embedding_async
filter_instance = SemanticMCPToolFilter(
embedding_model="text-embedding-3-small",
litellm_router_instance=mock_router,
top_k=2,
similarity_threshold=0.3,
enabled=True,
)
# Register an MCP tool with name "github-search"
mcp_tools = [
MCPTool(
name="github-search",
description="Search GitHub repos",
inputSchema={"type": "object"},
)
]
filter_instance._build_router(mcp_tools)
# Responses API native tool with SAME name as MCP canonical
responses_api_tool = {
"type": "function",
"name": "github-search",
"description": "Caller-owned search tool",
"parameters": {"type": "object"},
}
hook = SemanticToolFilterHook(filter_instance)
# Verify classification: should be native, not MCP
assert not hook._is_mcp_tool(responses_api_tool), (
"Responses API tool with type=function + top-level name "
"should be classified as native, not MCP"
)
# Full hook test: all-native request should preserve tools
data = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "Search GitHub"}],
"tools": [responses_api_tool],
"metadata": {},
}
result = await hook.async_pre_call_hook(
user_api_key_dict=Mock(),
cache=Mock(),
data=data,
call_type="completion",
)
# All tools are native → hook returns data with all tools preserved
filtered = (result or data)["tools"]
assert len(filtered) == 1, f"Native tool must survive, got {len(filtered)}"
assert filtered[0]["name"] == "github-search"
print("✅ Responses API tool with MCP-matching name correctly classified as native")
@pytest.mark.asyncio
async def test_semantic_filter_hook_preserves_tool_order():
"""
Regression test: tool ordering preservation.
Given: An interleaved request [mcp_tool_A, native_tool, mcp_tool_B]
When: The hook filters tools (all MCP tools survive)
Then: The output order must match the original request order,
NOT native-first.
"""
from litellm.proxy._experimental.mcp_server.semantic_tool_filter import (
SemanticMCPToolFilter,
)
from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook
from litellm.types.utils import Embedding, EmbeddingResponse
mock_router = Mock()
def mock_embedding_sync(*args, **kwargs):
return EmbeddingResponse(
data=[Embedding(embedding=[0.1] * 1536, index=0, object="embedding")],
model="text-embedding-3-small",
object="list",
usage={"prompt_tokens": 10, "total_tokens": 10},
)
async def mock_embedding_async(*args, **kwargs):
return mock_embedding_sync()
mock_router.embedding = mock_embedding_sync
mock_router.aembedding = mock_embedding_async
filter_instance = SemanticMCPToolFilter(
embedding_model="text-embedding-3-small",
litellm_router_instance=mock_router,
top_k=5,
similarity_threshold=0.3,
enabled=True,
)
# Register MCP tools
mcp_tool_a = MCPTool(
name="github-search",
description="Search GitHub",
inputSchema={"type": "object"},
)
mcp_tool_b = MCPTool(
name="github-issue",
description="Create GitHub issue",
inputSchema={"type": "object"},
)
filter_instance._build_router([mcp_tool_a, mcp_tool_b])
# Mock filter_tools to return both MCP tools (deterministic)
filter_instance.filter_tools = AsyncMock( # type: ignore[method-assign]
return_value=[mcp_tool_a, mcp_tool_b]
)
# Native tool (interleaved between MCP tools)
native_tool = {
"type": "function",
"function": {
"name": "weather_lookup",
"description": "Look up weather",
"parameters": {"type": "object", "properties": {}},
},
}
# Original order: [mcp_A, native, mcp_B]
original_tools = [mcp_tool_a, native_tool, mcp_tool_b]
hook = SemanticToolFilterHook(filter_instance)
data = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "Search GitHub and check weather"}],
"tools": original_tools,
"metadata": {},
}
result = await hook.async_pre_call_hook(
user_api_key_dict=Mock(),
cache=Mock(),
data=data,
call_type="completion",
)
assert result is not None, "Hook should return modified data"
filtered = result["tools"]
# All tools should survive
assert len(filtered) == 3, f"Expected 3 tools, got {len(filtered)}"
# Order must be preserved: [mcp_A, native, mcp_B]
assert filtered[0] is mcp_tool_a, "First tool should be mcp_tool_a"
assert filtered[1] is native_tool, "Second tool should be native_tool"
assert filtered[2] is mcp_tool_b, "Third tool should be mcp_tool_b"
print(
"✅ Tool ordering preserved: [mcp_A, native, mcp_B] maintained after filtering"
)
class TestGetToolsByNames:
"""
Regression coverage for SemanticMCPToolFilter._get_tools_by_names
@ -489,9 +913,7 @@ class TestGetToolsByNames:
{"name": "send_email", "description": "send mail"},
]
matched = filter_instance._get_tools_by_names(
["send_email"], available_tools
)
matched = filter_instance._get_tools_by_names(["send_email"], available_tools)
assert len(matched) == 1
assert matched[0]["name"] == "send_email"
@ -503,9 +925,7 @@ class TestGetToolsByNames:
client_name = "litellm_" + canonical
available_tools = [{"name": client_name, "description": "scrape"}]
matched = filter_instance._get_tools_by_names(
[canonical], available_tools
)
matched = filter_instance._get_tools_by_names([canonical], available_tools)
assert len(matched) == 1
# Must return the incoming tool unchanged so the client-facing
@ -516,13 +936,9 @@ class TestGetToolsByNames:
"""Some clients use dash as alias separator; accept that too."""
filter_instance = self._make_filter()
canonical = "weather_svc-get_weather"
available_tools = [
{"name": "mcp-" + canonical, "description": "weather"}
]
available_tools = [{"name": "mcp-" + canonical, "description": "weather"}]
matched = filter_instance._get_tools_by_names(
[canonical], available_tools
)
matched = filter_instance._get_tools_by_names([canonical], available_tools)
assert len(matched) == 1
assert matched[0]["name"] == "mcp-" + canonical
@ -552,9 +968,7 @@ class TestGetToolsByNames:
{"name": "litellm_" + canonical, "description": "wrapped"},
]
matched = filter_instance._get_tools_by_names(
[canonical], available_tools
)
matched = filter_instance._get_tools_by_names([canonical], available_tools)
assert len(matched) == 1
assert matched[0]["name"] == canonical
@ -567,9 +981,7 @@ class TestGetToolsByNames:
separator-anchored suffixes of ``litellm_api-fs-read_file``.
"""
filter_instance = self._make_filter()
available_tools = [
{"name": "litellm_api-fs-read_file", "description": "read"}
]
available_tools = [{"name": "litellm_api-fs-read_file", "description": "read"}]
matched = filter_instance._get_tools_by_names(
["fs-read_file", "api-fs-read_file"], available_tools
@ -590,9 +1002,7 @@ class TestGetToolsByNames:
{"name": "my_" + canonical, "description": "plain search"},
]
matched = filter_instance._get_tools_by_names(
[canonical], available_tools
)
matched = filter_instance._get_tools_by_names([canonical], available_tools)
assert len(matched) == 1
assert matched[0]["name"] == "my_" + canonical

View file

@ -351,6 +351,43 @@ async def test_can_key_call_model_all_team_models_no_team_id_is_denied():
assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied
@pytest.mark.asyncio
async def test_can_team_access_model_all_team_models_expands_router_models():
from litellm import Router
from litellm.proxy._types import SpecialModelNames
from litellm.proxy.auth.auth_checks import can_team_access_model
team_object = LiteLLM_TeamTable(
team_id="team-123",
models=[SpecialModelNames.all_team_models.value],
)
router = Router(
model_list=[
{
"model_name": "allowed-model",
"litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test"},
}
]
)
assert (
await can_team_access_model(
model="allowed-model",
team_object=team_object,
llm_router=router,
)
is True
)
with pytest.raises(ProxyException) as exc_info:
await can_team_access_model(
model="blocked-model",
team_object=team_object,
llm_router=router,
)
assert exc_info.value.type == ProxyErrorTypes.team_model_access_denied
@pytest.mark.asyncio
async def test_get_key_object_should_reconnect_once_on_db_connection_error():
mock_prisma_client = MagicMock()
@ -3792,3 +3829,54 @@ async def test_inference_route_still_enforces_team_budget():
valid_token=UserAPIKeyAuth(token="test-token", team_id="test-team"),
request=MagicMock(),
)
@pytest.mark.asyncio
async def test_virtual_key_max_budget_error_names_the_key():
"""BudgetExceededError for a virtual key must name the key (alias + masked key)
so operators don't have to reverse-map a spend figure back to a key."""
valid_token = UserAPIKeyAuth(
token="hashed-token",
key_alias="payments-prod",
key_name="sk-...um_g",
max_budget=10.0,
spend=0.0,
)
proxy_logging_obj = MagicMock()
proxy_logging_obj.budget_alerts = AsyncMock()
with patch(
"litellm.proxy.proxy_server.get_current_spend",
new=AsyncMock(return_value=25.0),
):
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await _virtual_key_max_budget_check(
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
)
message = str(exc_info.value)
assert "payments-prod" in message
assert "sk-...um_g" in message
@pytest.mark.asyncio
async def test_virtual_key_max_budget_not_exceeded_does_not_raise():
"""Spend below the configured budget must not raise."""
valid_token = UserAPIKeyAuth(
token="hashed-token",
key_alias="payments-prod",
max_budget=10.0,
spend=0.0,
)
proxy_logging_obj = MagicMock()
proxy_logging_obj.budget_alerts = AsyncMock()
with patch(
"litellm.proxy.proxy_server.get_current_spend",
new=AsyncMock(return_value=1.0),
):
await _virtual_key_max_budget_check(
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
)

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