Merge branch 'litellm_internal_staging' into litellm_cloudflare_native_openai_compatible

Resolve main.py conflict: the base refactored the cloudflare completion path
into _complete_cloudflare(ctx); reapply the removal of the /ai/run/ default so
the OpenAI-compatible api_base default lives only in get_complete_url.
This commit is contained in:
mateo-berri 2026-06-23 18:32:34 +00:00
commit b7367cf8e7
No known key found for this signature in database
97 changed files with 10786 additions and 3199 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

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

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

@ -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,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,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 SpecialMCPServerNames, 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
@ -4019,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
@ -4136,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

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

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

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

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

@ -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,6 +3148,19 @@ 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
all_litellm_params = (
@ -3448,6 +3463,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 +3521,7 @@ class SandboxProviders(str, Enum):
"""
E2B = "e2b"
OPENSANDBOX = "opensandbox"
class LiteLLMLoggingBaseClass:
@ -3610,10 +3627,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

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

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

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

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

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

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

View file

@ -543,3 +543,77 @@ async def test_get_available_models_for_user_expands_query_team_wildcard(
)
assert "openai/gpt-4o-mini" in result
def test_get_key_models_all_team_models_recursive_team():
"""GH#30619: when key and team both have all-team-models,
the sentinel should expand to proxy_model_list."""
from litellm.proxy.auth.model_checks import get_key_models
from litellm.proxy._types import SpecialModelNames
user_api_key_dict = type(
"obj", (object,),
{
"models": [SpecialModelNames.all_team_models.value],
"team_id": "team-1",
"team_models": [SpecialModelNames.all_team_models.value],
},
)()
proxy_model_list = ["model-a", "model-b"]
result = get_key_models(user_api_key_dict, proxy_model_list, {})
assert SpecialModelNames.all_team_models.value not in result
assert set(result) == {"model-a", "model-b"}
def test_get_key_models_all_team_models_keeps_mixed_team_entries():
from litellm.proxy.auth.model_checks import get_key_models
from litellm.proxy._types import SpecialModelNames
user_api_key_dict = type(
"obj",
(object,),
{
"models": [SpecialModelNames.all_team_models.value],
"team_id": "team-1",
"team_models": [
SpecialModelNames.all_team_models.value,
"restricted-model",
],
},
)()
result = get_key_models(user_api_key_dict, ["model-a", "model-b"], {})
assert SpecialModelNames.all_team_models.value not in result
assert set(result) == {"model-a", "model-b", "restricted-model"}
def test_get_team_models_all_team_models_expands():
"""GH#30619: all-team-models in team_models should expand."""
from litellm.proxy.auth.model_checks import get_team_models
from litellm.proxy._types import SpecialModelNames
result = get_team_models(
[SpecialModelNames.all_team_models.value],
["model-a", "model-b"],
{},
)
assert SpecialModelNames.all_team_models.value not in result
assert set(result) == {"model-a", "model-b"}
def test_get_team_models_all_team_models_expands_with_access_groups():
"""GH#30619: all-team-models with include_model_access_groups
should include access group keys."""
from litellm.proxy.auth.model_checks import get_team_models
from litellm.proxy._types import SpecialModelNames
result = get_team_models(
[SpecialModelNames.all_team_models.value],
["model-a", "model-b"],
{"group-1": ["g1-model"], "group-2": ["g2-model"]},
include_model_access_groups=True,
)
assert SpecialModelNames.all_team_models.value not in result
assert "model-a" in result
assert "model-b" in result
assert "group-1" in result
assert "group-2" in result

View file

@ -16,7 +16,11 @@ from unittest.mock import patch
import pytest
from litellm.proxy.db.db_url_settings import DatabaseURLSettings
from litellm.proxy.db.db_url_settings import (
DatabaseURLSettings,
unsupported_db_scheme,
unsupported_db_scheme_message,
)
def _apply() -> bool:
@ -27,6 +31,7 @@ def _apply() -> bool:
_MANAGED_DB_ENV_VARS = (
"IAM_TOKEN_DB_AUTH",
"DATABASE_URL",
"DIRECT_URL",
"DATABASE_URL_READ_REPLICA",
"DATABASE_HOST",
"DATABASE_PORT",
@ -287,3 +292,87 @@ def test_password_reader_uses_own_credentials(monkeypatch):
os.environ["DATABASE_URL_READ_REPLICA"]
== "postgresql://litellm_ro:ro_pw@reader.example.com:5432/litellm_db"
)
@pytest.mark.parametrize(
"url",
[
"postgresql://u:p@host:5432/db",
"postgres://u:p@host:5432/db",
"POSTGRESQL://u:p@host:5432/db",
"postgresql://host/db?schema=public",
],
)
def test_unsupported_db_scheme_accepts_postgres(url):
assert unsupported_db_scheme(url) is None
@pytest.mark.parametrize(
"url,scheme",
[
("sqlite:///data/litellm.db", "sqlite"),
("sqlite:///./local.db", "sqlite"),
("mysql://u:p@host:3306/db", "mysql"),
("mssql://host/db", "mssql"),
],
)
def test_unsupported_db_scheme_rejects_non_postgres(url, scheme):
assert unsupported_db_scheme(url) == scheme
def test_unsupported_db_scheme_does_not_echo_schemeless_credentials():
"""A malformed schemeless DSN must not leak its embedded credentials
through the return value (which callers log)."""
leaky = "litellm:s3cr3t_password@db.internal:5432/litellm"
result = unsupported_db_scheme(leaky)
assert result is not None
assert "s3cr3t_password" not in result
assert "db.internal" not in result
def test_apply_to_env_rejects_pinned_sqlite_writer(monkeypatch):
"""Componentized entrypoints pin DATABASE_URL and call apply_to_env; a
sqlite writer must raise here rather than reach Prisma."""
monkeypatch.setenv("DATABASE_URL", "sqlite:///data/litellm.db")
with pytest.raises(RuntimeError, match="sqlite"):
_apply()
# The bad URL must not have been propagated as a usable connection string.
assert os.environ["DATABASE_URL"] == "sqlite:///data/litellm.db"
def test_apply_to_env_rejects_pinned_sqlite_direct_url(monkeypatch):
"""DIRECT_URL reaches Prisma the same way DATABASE_URL does; a non-postgres
direct URL must be rejected in apply_to_env, matching the CLI startup guard."""
monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@writer.example.com:5432/db")
monkeypatch.setenv("DIRECT_URL", "sqlite:///data/litellm.db")
with pytest.raises(RuntimeError, match="DIRECT_URL.*sqlite"):
_apply()
def test_apply_to_env_rejects_pinned_non_postgres_reader(monkeypatch):
monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@writer.example.com:5432/db")
monkeypatch.setenv(
"DATABASE_URL_READ_REPLICA", "mysql://u:p@reader.example.com:3306/db"
)
with pytest.raises(RuntimeError, match="DATABASE_URL_READ_REPLICA.*mysql"):
_apply()
def test_apply_to_env_accepts_pinned_postgres(monkeypatch):
monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@host:5432/db")
# Operator-pinned URL: nothing reassembled, no error.
assert _apply() is False
def test_unsupported_db_scheme_message_names_var_and_scheme():
msg = unsupported_db_scheme_message("DIRECT_URL", "sqlite")
assert "DIRECT_URL" in msg
assert "sqlite" in msg
assert "postgresql://" in msg

View file

@ -16,8 +16,6 @@ Pins covered:
from __future__ import annotations
import json
from typing import Any, AsyncIterator
from unittest.mock import AsyncMock, MagicMock
import pytest
@ -26,8 +24,11 @@ from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.proxy_server import (
_apply_streaming_chunk_hooks,
_fast_serialize_simple_model_response_stream,
_format_fallback_metadata_sse_event,
_format_streaming_sse_chunk,
_get_client_requested_model_for_streaming,
_get_streaming_fallback_metadata,
_is_positive_int_like,
_restamp_streaming_chunk_model,
_serialize_streaming_chunk,
async_assistants_data_generator,
@ -71,6 +72,15 @@ async def _async_iter_raises(exc: Exception):
raise exc
class _FakeStream:
def __init__(self, chunks, hidden_params=None):
self._chunks = chunks
self._hidden_params = hidden_params or {}
def __aiter__(self):
return _async_iter(self._chunks)
# ---------------------------------------------------------------------------
# data_generator
# ---------------------------------------------------------------------------
@ -274,6 +284,34 @@ def test_restamp_streaming_chunk_model_overrides_model_on_dict():
assert logged is True
def test_restamp_streaming_chunk_model_uses_fallback_model_from_metadata():
chunk = _simple_chunk(model="openai/internal-fallback")
new_chunk, logged = _restamp_streaming_chunk_model(
chunk=chunk,
requested_model_from_client="primary-model",
request_data={"litellm_call_id": "id-1"},
model_mismatch_logged=False,
fallback_was_attempted=True,
fallback_model_from_metadata="fallback-model",
)
assert new_chunk.model == "fallback-model"
assert logged is True
def test_restamp_streaming_chunk_model_preserves_fallback_model_without_group():
chunk = _simple_chunk(model="openai/internal-fallback")
new_chunk, logged = _restamp_streaming_chunk_model(
chunk=chunk,
requested_model_from_client="primary-model",
request_data={},
model_mismatch_logged=False,
fallback_was_attempted=True,
fallback_model_from_metadata=None,
)
assert new_chunk.model == "openai/internal-fallback"
assert logged is False
def test_restamp_streaming_chunk_model_invalid_chunk_type_unchanged():
"""For a non-BaseModel, non-dict chunk the helper returns it as-is
along with the original ``model_mismatch_logged`` flag."""
@ -288,6 +326,147 @@ def test_restamp_streaming_chunk_model_invalid_chunk_type_unchanged():
assert logged is False
def test_is_positive_int_like_invalid_and_edge_values():
assert _is_positive_int_like(None) is False
assert _is_positive_int_like("not-a-number") is False
assert _is_positive_int_like(0) is False
assert _is_positive_int_like(-1) is False
assert _is_positive_int_like("1") is True
assert _is_positive_int_like(2) is True
def test_get_streaming_fallback_metadata_reads_headers():
fallback_errors = [
{
"message": "litellm.RateLimitError: upstream limited request",
"type": "RateLimitError",
"param": None,
"code": "429",
}
]
stream = _FakeStream(
[],
hidden_params={
"additional_headers": {
"x-litellm-attempted-fallbacks": "1",
"x-litellm-model-group": "fallback-model",
"x-litellm-fallback-errors": json.dumps(fallback_errors),
}
},
)
assert _get_streaming_fallback_metadata(stream) == (
True,
"fallback-model",
fallback_errors,
)
def test_get_streaming_fallback_metadata_no_additional_headers():
stream = _FakeStream([], hidden_params={})
assert _get_streaming_fallback_metadata(stream) == (False, None, [])
def test_get_streaming_fallback_metadata_zero_fallback_count():
stream = _FakeStream(
[],
hidden_params={
"additional_headers": {"x-litellm-attempted-fallbacks": 0}
},
)
assert _get_streaming_fallback_metadata(stream) == (False, None, [])
def test_get_streaming_fallback_metadata_no_model_group_returns_none_model():
stream = _FakeStream(
[],
hidden_params={
"additional_headers": {
"x-litellm-attempted-fallbacks": 1,
}
},
)
was_attempted, fallback_model, errors = _get_streaming_fallback_metadata(stream)
assert was_attempted is True
assert fallback_model is None
assert errors == []
def test_restamp_streaming_chunk_model_azure_router_preserves_model():
chunk = _simple_chunk(model="azure_ai/internal-deployment")
new_chunk, logged = _restamp_streaming_chunk_model(
chunk=chunk,
requested_model_from_client="azure_ai/model-router",
request_data={},
model_mismatch_logged=False,
)
assert new_chunk.model == "azure_ai/internal-deployment"
assert logged is False
def test_restamp_streaming_chunk_model_fastest_response_preserves_model():
chunk = _simple_chunk(model="winning-model")
new_chunk, logged = _restamp_streaming_chunk_model(
chunk=chunk,
requested_model_from_client="gpt-4,claude-3",
request_data={"fastest_response": True},
model_mismatch_logged=False,
)
assert new_chunk.model == "winning-model"
assert logged is False
def test_restamp_streaming_chunk_model_setattr_exception_logs_and_returns():
from pydantic import ConfigDict
class FrozenChunk(_simple_chunk().__class__):
model_config = ConfigDict(frozen=True)
chunk = FrozenChunk(
id="chatcmpl-test",
choices=[],
created=0,
model="openai/internal-x",
object="chat.completion.chunk",
)
new_chunk, logged = _restamp_streaming_chunk_model(
chunk=chunk,
requested_model_from_client="gpt-4",
request_data={"litellm_call_id": "test-id"},
model_mismatch_logged=False,
)
assert new_chunk.model == "openai/internal-x"
assert logged is True
def test_format_fallback_metadata_sse_event():
fallback_errors = [
{
"message": "litellm.RateLimitError: upstream limited request",
"type": "RateLimitError",
"param": None,
"code": "429",
}
]
event = _format_fallback_metadata_sse_event(
fallback_model="fallback-model",
fallback_errors=fallback_errors,
)
assert isinstance(event, str)
assert event.startswith("data: ")
payload = json.loads(event.removeprefix("data: ").removesuffix("\n\n"))
assert payload["choices"] == []
assert payload["litellm_fallback"] == {
"fallback_model": "fallback-model",
"errors": fallback_errors,
}
assert payload["id"] == "litellm-fallback-metadata"
assert payload["object"] == "chat.completion.chunk"
assert payload["model"] == "fallback-model"
assert isinstance(payload["created"], int)
# ---------------------------------------------------------------------------
# _fast_serialize_simple_model_response_stream
# ---------------------------------------------------------------------------
@ -473,7 +652,7 @@ async def test_async_data_generator_yields_sse_chunks_and_done(monkeypatch):
# First chunk is bytes (fast path) wrapped via _format_streaming_sse_chunk.
first = out[0]
assert isinstance(first, bytes)
payload = json.loads(first.removeprefix(b"data: ").rstrip(b"\n\n"))
payload = json.loads(first.removeprefix(b"data: ").removesuffix(b"\n\n"))
assert normalize(payload) == {
"id": "<VOLATILE>",
"object": "chat.completion.chunk",
@ -488,6 +667,172 @@ async def test_async_data_generator_yields_sse_chunks_and_done(monkeypatch):
}
@pytest.mark.asyncio
async def test_async_data_generator_uses_response_fallback_metadata(monkeypatch):
_patch_logging_flags(monkeypatch)
response = _FakeStream(
[_simple_chunk(model="openai/internal-fallback", content="hello")],
hidden_params={
"additional_headers": {
"x-litellm-attempted-fallbacks": 1,
"x-litellm-model-group": "fallback-model",
}
},
)
out = []
async for line in async_data_generator(
response=response,
user_api_key_dict=_user_auth(),
request_data={"model": "primary-model", "include_fallback_errors": True},
):
out.append(line)
first = out[0]
assert isinstance(first, bytes)
payload = json.loads(first.removeprefix(b"data: ").removesuffix(b"\n\n"))
assert payload["model"] == "fallback-model"
@pytest.mark.asyncio
async def test_async_data_generator_uses_chunk_fallback_metadata(monkeypatch):
_patch_logging_flags(monkeypatch)
chunk = _simple_chunk(model="openai/internal-fallback", content="hello")
chunk._hidden_params = {
"additional_headers": {
"x-litellm-attempted-fallbacks": 1,
"x-litellm-model-group": "fallback-model",
}
}
out = []
async for line in async_data_generator(
response=_async_iter([chunk]),
user_api_key_dict=_user_auth(),
request_data={"model": "primary-model"},
):
out.append(line)
first = out[0]
assert isinstance(first, bytes)
payload = json.loads(first.removeprefix(b"data: ").removesuffix(b"\n\n"))
assert payload["model"] == "fallback-model"
@pytest.mark.asyncio
async def test_async_data_generator_switches_model_mid_stream_on_fallback(monkeypatch):
"""Pre-fallback chunks keep the client-requested model; once a chunk carries
fallback metadata the model latches to the fallback group for the rest of the
stream. This pins the client-visible mid-stream model change."""
_patch_logging_flags(monkeypatch)
primary_chunk = _simple_chunk(model="openai/internal-primary", content="hi")
fallback_chunk = _simple_chunk(model="openai/internal-fallback", content="there")
fallback_chunk._hidden_params = {
"additional_headers": {
"x-litellm-attempted-fallbacks": 1,
"x-litellm-model-group": "fallback-model",
}
}
out = []
async for line in async_data_generator(
response=_async_iter([primary_chunk, fallback_chunk]),
user_api_key_dict=_user_auth(),
request_data={"model": "primary-model"},
):
out.append(line)
first_payload = json.loads(out[0].removeprefix(b"data: ").removesuffix(b"\n\n"))
second_payload = json.loads(out[1].removeprefix(b"data: ").removesuffix(b"\n\n"))
assert first_payload["model"] == "primary-model"
assert second_payload["model"] == "fallback-model"
@pytest.mark.asyncio
async def test_async_data_generator_emits_fallback_error_metadata_event(monkeypatch):
_patch_logging_flags(monkeypatch)
monkeypatch.setitem(ps.general_settings, "expose_fallback_errors_to_caller", True)
fallback_errors = [
{
"message": "litellm.RateLimitError: upstream limited request",
"type": "RateLimitError",
"param": None,
"code": "429",
}
]
response = _FakeStream(
[_simple_chunk(model="openai/internal-fallback", content="hello")],
hidden_params={
"additional_headers": {
"x-litellm-attempted-fallbacks": 1,
"x-litellm-model-group": "fallback-model",
"x-litellm-fallback-errors": json.dumps(fallback_errors),
}
},
)
out = []
async for line in async_data_generator(
response=response,
user_api_key_dict=_user_auth(),
request_data={"model": "primary-model", "include_fallback_errors": True},
):
out.append(line)
assert isinstance(out[0], bytes)
chunk_payload = json.loads(out[0].removeprefix(b"data: ").removesuffix(b"\n\n"))
assert chunk_payload["model"] == "fallback-model"
assert isinstance(out[1], str)
assert out[1].startswith("data: ")
metadata_payload = json.loads(out[1].removeprefix("data: ").removesuffix("\n\n"))
assert metadata_payload["choices"] == []
assert metadata_payload["litellm_fallback"] == {
"fallback_model": "fallback-model",
"errors": fallback_errors,
}
assert metadata_payload["id"] == "litellm-fallback-metadata"
assert metadata_payload["object"] == "chat.completion.chunk"
assert metadata_payload["model"] == "fallback-model"
assert isinstance(metadata_payload["created"], int)
@pytest.mark.asyncio
async def test_async_data_generator_skips_fallback_error_event_without_opt_in(
monkeypatch,
):
_patch_logging_flags(monkeypatch)
fallback_errors = [
{
"message": "litellm.RateLimitError: upstream limited request",
"type": "RateLimitError",
"param": None,
"code": "429",
}
]
response = _FakeStream(
[_simple_chunk(model="openai/internal-fallback", content="hello")],
hidden_params={
"additional_headers": {
"x-litellm-attempted-fallbacks": 1,
"x-litellm-model-group": "fallback-model",
"x-litellm-fallback-errors": json.dumps(fallback_errors),
}
},
)
out = []
async for line in async_data_generator(
response=response,
user_api_key_dict=_user_auth(),
request_data={"model": "primary-model"},
):
out.append(line)
assert isinstance(out[0], bytes)
payload = json.loads(out[0].removeprefix(b"data: ").removesuffix(b"\n\n"))
assert payload["model"] == "fallback-model"
@pytest.mark.asyncio
async def test_async_data_generator_mid_stream_exception_yields_error_payload(
monkeypatch,

View file

@ -1708,6 +1708,57 @@ class TestRunServerDbSetup:
use_migrate=True, use_v2_resolver=False
)
@patch("subprocess.run")
@patch("atexit.register")
@patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database")
@patch("litellm.proxy.db.check_migration.check_prisma_schema_diff")
@patch("litellm.proxy.db.prisma_client.should_update_prisma_schema")
def test_startup_exits_on_non_postgres_database_url(
self,
mock_should_update_schema,
mock_check_schema_diff,
mock_setup_database,
mock_atexit_register,
mock_subprocess_run,
):
"""A sqlite DATABASE_URL must exit immediately, before any prisma call,
instead of stalling on a migration against the postgresql-only schema."""
from litellm.proxy.proxy_cli import run_server
mock_subprocess_run.return_value = MagicMock(returncode=0)
mock_should_update_schema.return_value = True
mock_proxy_module = MagicMock(
app=MagicMock(),
ProxyConfig=MagicMock(),
KeyManagementSettings=MagicMock(),
save_worker_config=MagicMock(),
)
clean_env = {
k: v
for k, v in os.environ.items()
if k not in ("DATABASE_URL", "DIRECT_URL")
}
clean_env["DATABASE_URL"] = "sqlite:///data/litellm.db"
with (
patch.dict(os.environ, clean_env, clear=True),
patch.dict(
"sys.modules",
{
"proxy_server": mock_proxy_module,
"litellm.proxy.proxy_server": mock_proxy_module,
},
),
):
with pytest.raises(SystemExit) as exc_info:
run_server.main(
["--local", "--skip_server_startup"], standalone_mode=False
)
assert exc_info.value.code == 1
mock_setup_database.assert_not_called()
# --- Module-level helpers for worker startup hook tests ---

View file

@ -0,0 +1,142 @@
import json
from pydantic import BaseModel
from litellm.router_utils.add_retry_fallback_headers import (
add_fallback_headers_to_response,
add_retry_headers_to_response,
get_fallback_errors_from_headers,
get_hidden_params_dict,
)
class StreamingWrapper:
def __init__(self):
self._hidden_params = {"additional_headers": {"x-existing": "keep"}}
def test_add_fallback_headers_to_streaming_wrapper():
response = StreamingWrapper()
result = add_fallback_headers_to_response(
response=response,
attempted_fallbacks=1,
)
assert result is response
assert response._hidden_params["additional_headers"] == {
"x-existing": "keep",
"x-litellm-attempted-fallbacks": 1,
}
def test_add_fallback_headers_serializes_fallback_errors():
response = StreamingWrapper()
fallback_errors = [
{
"message": "litellm.RateLimitError: upstream limited request",
"type": "RateLimitError",
"param": None,
"code": "429",
}
]
result = add_fallback_headers_to_response(
response=response,
attempted_fallbacks=1,
fallback_errors=fallback_errors,
)
assert result is response
assert response._hidden_params["additional_headers"][
"x-litellm-attempted-fallbacks"
] == 1
assert (
json.loads(
response._hidden_params["additional_headers"]["x-litellm-fallback-errors"]
)
== fallback_errors
)
def test_add_retry_headers_to_streaming_wrapper():
response = StreamingWrapper()
result = add_retry_headers_to_response(
response=response,
attempted_retries=2,
max_retries=3,
)
assert result is response
assert response._hidden_params["additional_headers"] == {
"x-existing": "keep",
"x-litellm-attempted-retries": 2,
"x-litellm-max-retries": 3,
}
def test_get_hidden_params_dict_with_pydantic_model_hidden_params():
class InnerHiddenParams(BaseModel):
additional_headers: dict = {}
class Response:
def __init__(self):
self._hidden_params = InnerHiddenParams(
additional_headers={"x-custom": "value"}
)
result = get_hidden_params_dict(Response())
assert result == {"additional_headers": {"x-custom": "value"}}
def test_get_hidden_params_dict_with_no_hidden_params():
class PlainResponse:
pass
assert get_hidden_params_dict(PlainResponse()) == {}
def test_add_fallback_headers_when_no_existing_additional_headers():
class NoHeadersWrapper:
def __init__(self):
self._hidden_params = {}
response = NoHeadersWrapper()
result = add_fallback_headers_to_response(response=response, attempted_fallbacks=2)
assert result is response
assert response._hidden_params["additional_headers"]["x-litellm-attempted-fallbacks"] == 2
def test_add_fallback_headers_returns_none_when_response_is_none():
result = add_fallback_headers_to_response(response=None, attempted_fallbacks=1)
assert result is None
def test_add_fallback_headers_returns_unchanged_when_response_has_no_hidden_params():
class PlainObject:
pass
obj = PlainObject()
result = add_fallback_headers_to_response(response=obj, attempted_fallbacks=1)
assert result is obj
assert not hasattr(obj, "_hidden_params")
def test_get_fallback_errors_from_headers_existing_list_passthrough():
errors = [{"message": "err", "type": "T", "param": None, "code": "400"}]
result = get_fallback_errors_from_headers({"x-litellm-fallback-errors": errors})
assert result == errors
def test_get_fallback_errors_from_headers_invalid_json_returns_empty():
result = get_fallback_errors_from_headers(
{"x-litellm-fallback-errors": "not-valid-json-{"}
)
assert result == []
def test_get_fallback_errors_from_headers_missing_key_returns_empty():
result = get_fallback_errors_from_headers({})
assert result == []

View file

@ -0,0 +1,139 @@
import json
import pytest
from litellm.router_utils.fallback_event_handlers import run_async_fallback
class StreamingWrapper:
def __init__(self):
self._hidden_params = {"additional_headers": {}}
class FakeRouter:
def log_retry(self, kwargs, e):
return kwargs
async def async_function_with_fallbacks(self, *args, **kwargs):
return StreamingWrapper()
class AlwaysFailRouter:
def log_retry(self, kwargs, e):
return kwargs
async def async_function_with_fallbacks(self, *args, **kwargs):
raise RuntimeError("fallback model also failed")
@pytest.mark.asyncio
async def test_run_async_fallback_adds_errors_when_opted_in():
response = await run_async_fallback(
litellm_router=FakeRouter(),
fallback_model_group=["fallback-model"],
original_model_group="primary-model",
original_exception=RuntimeError("upstream limited request"),
max_fallbacks=3,
fallback_depth=0,
include_fallback_errors=True,
)
additional_headers = response._hidden_params["additional_headers"]
assert additional_headers["x-litellm-attempted-fallbacks"] == 1
assert json.loads(additional_headers["x-litellm-fallback-errors"]) == [
{
"message": "upstream limited request",
"type": "RuntimeError",
"param": None,
"code": None,
}
]
@pytest.mark.asyncio
async def test_run_async_fallback_omits_errors_without_opt_in():
response = await run_async_fallback(
litellm_router=FakeRouter(),
fallback_model_group=["fallback-model"],
original_model_group="primary-model",
original_exception=RuntimeError("upstream limited request"),
max_fallbacks=3,
fallback_depth=0,
)
additional_headers = response._hidden_params["additional_headers"]
assert additional_headers["x-litellm-attempted-fallbacks"] == 1
assert "x-litellm-fallback-errors" not in additional_headers
@pytest.mark.asyncio
async def test_run_async_fallback_raises_when_all_fallbacks_fail():
with pytest.raises(RuntimeError, match="fallback model also failed"):
await run_async_fallback(
litellm_router=AlwaysFailRouter(),
fallback_model_group=["fallback-model"],
original_model_group="primary-model",
original_exception=RuntimeError("original request failed"),
max_fallbacks=3,
fallback_depth=0,
include_fallback_errors=True,
)
class RecordingRouter:
def __init__(self):
self.received_kwargs = None
def log_retry(self, kwargs, e):
return kwargs
async def async_function_with_fallbacks(self, *args, **kwargs):
self.received_kwargs = kwargs
return StreamingWrapper()
@pytest.mark.asyncio
async def test_run_async_fallback_forwards_include_fallback_errors_to_nested_call():
"""A nested fallback (multi-hop) must keep collecting errors, so the opt-in
flag has to reach the nested async_function_with_fallbacks call."""
router = RecordingRouter()
await run_async_fallback(
litellm_router=router,
fallback_model_group=["fallback-model"],
original_model_group="primary-model",
original_exception=RuntimeError("upstream limited request"),
max_fallbacks=3,
fallback_depth=0,
include_fallback_errors=True,
)
assert router.received_kwargs.get("include_fallback_errors") is True
@pytest.mark.asyncio
async def test_run_async_fallback_does_not_forward_flag_without_opt_in():
router = RecordingRouter()
await run_async_fallback(
litellm_router=router,
fallback_model_group=["fallback-model"],
original_model_group="primary-model",
original_exception=RuntimeError("upstream limited request"),
max_fallbacks=3,
fallback_depth=0,
)
assert "include_fallback_errors" not in router.received_kwargs
@pytest.mark.asyncio
async def test_run_async_fallback_skips_original_model_group():
response = await run_async_fallback(
litellm_router=FakeRouter(),
fallback_model_group=["primary-model", "fallback-model"],
original_model_group="primary-model",
original_exception=RuntimeError("original failed"),
max_fallbacks=3,
fallback_depth=0,
)
assert response._hidden_params["additional_headers"]["x-litellm-attempted-fallbacks"] == 1

View file

@ -0,0 +1,647 @@
import json
import httpx
import pytest
import litellm
from litellm.llms.base_llm.sandbox.transformation import ContainerHandle
from litellm.llms.opensandbox.sandbox.transformation import (
MAX_OUTPUT_BYTES,
OPEN_SANDBOX_DEFAULT_TEMPLATE,
OpenSandboxSandboxConfig,
)
from litellm.utils import ProviderConfigManager
TEST_API_BASE = "https://sandbox.test/v1"
def http_status_error(status_code, url="http://test"):
return httpx.HTTPStatusError(
f"status {status_code}",
request=httpx.Request("GET", url),
response=httpx.Response(status_code),
)
def sse(data):
return f"data: {json.dumps(data)}"
class FakeResponse:
def __init__(self, *, json_data=None, lines=None, status_code=200):
self._json = json_data
self._lines = lines or []
self.status_code = status_code
def json(self):
return self._json
def raise_for_status(self):
if self.status_code >= 400:
raise http_status_error(self.status_code)
async def aiter_lines(self):
for line in self._lines:
yield line
class FakeHTTPClient:
def __init__(
self,
*,
create_json=None,
sandbox_states=None,
endpoint_json=None,
endpoint_responses=None,
execute_lines=None,
delete_status=204,
execute_raises=None,
):
self.create_json = create_json or {
"id": "osb_123",
"status": {"state": "Running"},
"createdAt": "2026-01-01T00:00:00Z",
"entrypoint": ["/opt/code-interpreter/code-interpreter.sh"],
}
self.sandbox_states = list(
sandbox_states
or [
{
"id": "osb_123",
"status": {"state": "Running"},
"createdAt": "2026-01-01T00:00:00Z",
"entrypoint": ["/opt/code-interpreter/code-interpreter.sh"],
}
]
)
self.endpoint_json = endpoint_json or {
"endpoint": "execd.local:44772",
"headers": {"X-EXECD-ACCESS-TOKEN": "execd-token"},
}
self.endpoint_responses = (
list(endpoint_responses) if endpoint_responses is not None else None
)
self.execute_lines = execute_lines or []
self.delete_status = delete_status
self.execute_raises = execute_raises
self.calls = []
async def post(self, url, headers=None, json=None, stream=False, **kwargs):
self.calls.append(("POST", url, headers, json, {"stream": stream}))
if url.endswith("/sandboxes"):
return FakeResponse(json_data=self.create_json)
if url.endswith("/code"):
if self.execute_raises is not None:
raise self.execute_raises
return FakeResponse(lines=self.execute_lines)
raise AssertionError(f"unexpected POST {url}")
async def get(self, url, headers=None, params=None, **kwargs):
self.calls.append(("GET", url, headers, None, params))
if "/endpoints/44772" in url:
if self.endpoint_responses is not None and self.endpoint_responses:
response = self.endpoint_responses.pop(0)
if isinstance(response, Exception):
raise response
if isinstance(response, FakeResponse):
return response
return FakeResponse(json_data=response)
return FakeResponse(json_data=self.endpoint_json)
if "/sandboxes/" in url:
state = self.sandbox_states.pop(0)
return FakeResponse(json_data=state)
raise AssertionError(f"unexpected GET {url}")
async def delete(self, url, headers=None, **kwargs):
self.calls.append(("DELETE", url, headers, None, None))
if not (200 <= self.delete_status < 300):
raise http_status_error(self.delete_status, url)
return FakeResponse(status_code=self.delete_status)
def test_parse_sse_lines_maps_output_result_count_and_error():
lines = [
sse({"type": "stdout", "text": "hello\n"}),
sse({"type": "stderr", "text": "warn\n"}),
sse({"type": "result", "results": {"text/plain": "4"}}),
sse({"type": "execution_count", "execution_count": 7}),
sse(
{
"type": "error",
"error": {
"ename": "ValueError",
"evalue": "bad",
"traceback": ["Traceback"],
},
}
),
]
result = OpenSandboxSandboxConfig._parse_lines(lines)
assert result.stdout == "hello\n"
assert result.stderr == "warn\n"
assert result.results == [{"text/plain": "4"}]
assert result.execution_count == 7
assert result.error == {
"name": "ValueError",
"value": "bad",
"traceback": ["Traceback"],
}
def test_parse_sse_lines_skips_non_json_and_control_lines():
lines = [
"event: message",
"not-json",
"",
sse({"type": "stdout", "text": "ok\n"}),
]
result = OpenSandboxSandboxConfig._parse_lines(lines)
assert result.stdout == "ok\n"
assert result.error is None
def test_parse_sse_lines_maps_fallback_shapes():
lines = [
"data:",
sse(["not-a-dict"]),
sse({"code": "BadRequest", "message": "nope"}),
sse({"type": "result", "text/plain": "4"}),
sse({"type": "error", "name": "RuntimeError", "text": "boom"}),
sse({"type": "execution_count", "execution_count": "8"}),
]
result = OpenSandboxSandboxConfig._parse_lines(lines)
assert result.results == [{"text/plain": "4"}]
assert result.execution_count == 8
assert result.error == {
"name": "BadRequest",
"value": "nope",
"traceback": [],
}
fallback_error = OpenSandboxSandboxConfig._parse_lines(
[sse({"type": "error", "name": "RuntimeError", "text": "boom"})]
)
assert fallback_error.error == {
"name": "RuntimeError",
"value": "boom",
"traceback": [],
}
empty_string_error = OpenSandboxSandboxConfig._parse_lines(
[
sse(
{
"type": "error",
"error": {
"ename": "",
"name": "FallbackName",
"evalue": "",
"value": "fallback value",
"traceback": [],
},
}
)
]
)
assert empty_string_error.error == {
"name": "",
"value": "",
"traceback": [],
}
def test_static_helpers_cover_defaults_and_fallbacks(monkeypatch):
def fake_secret(key):
if key == "OPEN_SANDBOX_API_KEY":
return "env-key"
if key == "OPEN_SANDBOX_API_BASE":
return TEST_API_BASE
return None
monkeypatch.setattr(
"litellm.llms.opensandbox.sandbox.transformation.get_secret_str",
fake_secret,
)
config = OpenSandboxSandboxConfig()
handle = ContainerHandle(id="osb", provider="opensandbox", domain="http://x/v1")
assert config.validate_environment() == "env-key"
assert config.validate_environment(api_key="") == ""
assert config._api_key(api_key=None, handle=handle) == "env-key"
handle._hidden_params = {"api_key": "stored-key"}
assert config._api_key(api_key=None, handle=handle) == "stored-key"
assert config._http(None) is not None
body = config._create_body(
template=None,
timeout=None,
allow_internet_access=False,
metadata=None,
env_vars=None,
resource_limits=None,
resource_requests=None,
entrypoint=None,
network_policy={"egress": [{"domain": "example.com"}]},
secure_access=True,
)
assert body["networkPolicy"] == {"egress": [{"domain": "example.com"}]}
assert body["secureAccess"] is True
other_body = config._create_body(
template=None,
timeout=None,
allow_internet_access=False,
metadata=None,
env_vars=None,
resource_limits=None,
resource_requests=None,
entrypoint=None,
network_policy=None,
secure_access=False,
)
assert body["resourceLimits"] is not other_body["resourceLimits"]
assert config._sandbox_state(None) is None
assert config._sandbox_state({"status": "Running"}) is None
assert config._as_str_dict(None) == {}
assert config._endpoint_base_url("http://execd.local", "https://api/v1") == (
"http://execd.local"
)
assert config._api_base(None) == TEST_API_BASE
assert config._api_base("https://direct.test/v1/") == "https://direct.test/v1"
assert config._as_int("9") == 9
assert config._as_int("nope") is None
assert config._as_int(None) is None
assert isinstance(
ProviderConfigManager.get_provider_sandbox_config("opensandbox"),
OpenSandboxSandboxConfig,
)
def test_api_base_requires_kwarg_or_env(monkeypatch):
monkeypatch.setattr(
"litellm.llms.opensandbox.sandbox.transformation.get_secret_str",
lambda key: None,
)
with pytest.raises(ValueError, match="api_base is required"):
OpenSandboxSandboxConfig._api_base(None)
@pytest.mark.asyncio
async def test_create_posts_default_body_and_omits_empty_api_key():
client = FakeHTTPClient()
handle = await OpenSandboxSandboxConfig().acreate_sandbox(
api_key="", api_base=TEST_API_BASE, client=client
)
method, url, headers, body, _ = client.calls[0]
assert method == "POST"
assert url == f"{TEST_API_BASE}/sandboxes"
assert "OPEN-SANDBOX-API-KEY" not in headers
assert body["image"] == {"uri": OPEN_SANDBOX_DEFAULT_TEMPLATE}
assert body["entrypoint"] == ["/opt/code-interpreter/code-interpreter.sh"]
assert body["timeout"] == 300
assert body["resourceLimits"] == {"cpu": "1", "memory": "2Gi"}
assert body["networkPolicy"] == {"defaultAction": "deny", "egress": []}
assert handle.id == "osb_123"
assert handle._hidden_params["execd_endpoint"] == "execd.local:44772"
@pytest.mark.asyncio
async def test_create_can_opt_into_internet_access():
client = FakeHTTPClient()
await OpenSandboxSandboxConfig().acreate_sandbox(
api_key="",
api_base=TEST_API_BASE,
allow_internet_access=True,
client=client,
)
_, _, _, body, _ = client.calls[0]
assert "networkPolicy" not in body
@pytest.mark.asyncio
async def test_create_custom_options_poll_and_endpoint_resolution():
client = FakeHTTPClient(
create_json={
"id": "osb_pending",
"status": {"state": "Pending"},
"createdAt": "2026-01-01T00:00:00Z",
"entrypoint": ["/bin/sh"],
},
sandbox_states=[
{
"id": "osb_pending",
"status": {"state": "Running"},
"createdAt": "2026-01-01T00:00:00Z",
"entrypoint": ["/bin/sh"],
}
],
)
handle = await OpenSandboxSandboxConfig().acreate_sandbox(
template="custom/image:latest",
timeout=600,
allow_internet_access=False,
api_key="osb-key",
api_base="https://sandbox.example/v1",
metadata={"suite": "unit"},
env_vars={"PYTHONUNBUFFERED": "1"},
resource_limits={"cpu": "500m", "memory": "512Mi"},
resource_requests={"cpu": "250m", "memory": "256Mi"},
entrypoint=["/bin/sh", "-lc", "sleep 3600"],
use_server_proxy=True,
client=client,
)
_, create_url, create_headers, body, _ = client.calls[0]
_, poll_url, poll_headers, _, _ = client.calls[1]
_, endpoint_url, endpoint_headers, _, endpoint_params = client.calls[2]
assert create_url == "https://sandbox.example/v1/sandboxes"
assert create_headers["OPEN-SANDBOX-API-KEY"] == "osb-key"
assert body["image"] == {"uri": "custom/image:latest"}
assert body["entrypoint"] == ["/bin/sh", "-lc", "sleep 3600"]
assert body["metadata"] == {"suite": "unit"}
assert body["env"] == {"PYTHONUNBUFFERED": "1"}
assert body["resourceLimits"] == {"cpu": "500m", "memory": "512Mi"}
assert body["resourceRequests"] == {"cpu": "250m", "memory": "256Mi"}
assert body["networkPolicy"] == {"defaultAction": "deny", "egress": []}
assert poll_url == "https://sandbox.example/v1/sandboxes/osb_pending"
assert poll_headers["OPEN-SANDBOX-API-KEY"] == "osb-key"
assert endpoint_url.endswith("/sandboxes/osb_pending/endpoints/44772")
assert endpoint_headers["OPEN-SANDBOX-API-KEY"] == "osb-key"
assert endpoint_params == {"use_server_proxy": True}
assert handle.id == "osb_pending"
@pytest.mark.asyncio
async def test_create_waits_across_pending_state(monkeypatch):
client = FakeHTTPClient(
create_json={
"id": "osb_pending",
"status": {"state": "Pending"},
"createdAt": "2026-01-01T00:00:00Z",
},
sandbox_states=[
{"id": "osb_pending", "status": {"state": "Pending"}},
{"id": "osb_pending", "status": {"state": "Running"}},
],
)
sleeps = []
async def fake_sleep(interval):
sleeps.append(interval)
monkeypatch.setattr(
"litellm.llms.opensandbox.sandbox.transformation.asyncio.sleep", fake_sleep
)
handle = await OpenSandboxSandboxConfig().acreate_sandbox(
api_key="",
api_base=TEST_API_BASE,
ready_timeout=1,
poll_interval=0.01,
client=client,
)
assert handle.id == "osb_pending"
assert sleeps == [0.01]
@pytest.mark.asyncio
async def test_create_raises_for_terminal_state():
client = FakeHTTPClient(
create_json={"id": "osb_failed", "status": {"state": "Pending"}},
sandbox_states=[
{"id": "osb_failed", "status": {"state": "Failed"}},
],
)
with pytest.raises(ValueError, match="entered Failed"):
await OpenSandboxSandboxConfig().acreate_sandbox(
api_key="", api_base=TEST_API_BASE, client=client
)
@pytest.mark.asyncio
async def test_create_times_out_waiting_for_running():
client = FakeHTTPClient(
create_json={"id": "osb_slow", "status": {"state": "Pending"}},
sandbox_states=[
{"id": "osb_slow", "status": {"state": "Pending"}},
],
)
with pytest.raises(TimeoutError, match="was not Running"):
await OpenSandboxSandboxConfig().acreate_sandbox(
api_key="",
api_base=TEST_API_BASE,
ready_timeout=0,
poll_interval=0,
client=client,
)
@pytest.mark.asyncio
async def test_create_waits_for_endpoint_resolution(monkeypatch):
client = FakeHTTPClient(
endpoint_responses=[
http_status_error(404, f"{TEST_API_BASE}/sandboxes/osb_123"),
{
"endpoint": "execd.local:44772",
"headers": {"X-EXECD-ACCESS-TOKEN": "execd-token"},
},
],
)
sleeps = []
async def fake_sleep(interval):
sleeps.append(interval)
monkeypatch.setattr(
"litellm.llms.opensandbox.sandbox.transformation.asyncio.sleep", fake_sleep
)
handle = await OpenSandboxSandboxConfig().acreate_sandbox(
api_key="",
api_base=TEST_API_BASE,
ready_timeout=1,
poll_interval=0.01,
client=client,
)
endpoint_calls = [call for call in client.calls if "/endpoints/44772" in call[1]]
assert handle._hidden_params["execd_endpoint"] == "execd.local:44772"
assert len(endpoint_calls) == 2
assert sleeps == [0.01]
@pytest.mark.asyncio
async def test_create_raises_when_endpoint_is_missing():
client = FakeHTTPClient(endpoint_json={"headers": {"X": "y"}})
with pytest.raises(TimeoutError, match="execd endpoint.*not ready"):
await OpenSandboxSandboxConfig().acreate_sandbox(
api_key="", api_base=TEST_API_BASE, ready_timeout=0, client=client
)
@pytest.mark.asyncio
async def test_create_reraises_non_404_endpoint_error():
client = FakeHTTPClient(endpoint_responses=[http_status_error(500)])
with pytest.raises(httpx.HTTPStatusError):
await OpenSandboxSandboxConfig().acreate_sandbox(
api_key="", api_base=TEST_API_BASE, client=client
)
@pytest.mark.asyncio
async def test_run_code_resolves_bare_id_and_posts_sse_request():
client = FakeHTTPClient(
execute_lines=[
sse({"type": "stdout", "text": "42\n"}),
]
)
result = await OpenSandboxSandboxConfig().arun_code(
container="osb_bare",
code="print(6*7)",
language="python",
api_key="",
api_base="http://sandbox.local/v1",
client=client,
)
endpoint_call = client.calls[0]
run_call = client.calls[1]
assert endpoint_call[0] == "GET"
assert (
endpoint_call[1] == "http://sandbox.local/v1/sandboxes/osb_bare/endpoints/44772"
)
assert run_call[0] == "POST"
assert run_call[1] == "http://execd.local:44772/code"
assert run_call[2]["X-EXECD-ACCESS-TOKEN"] == "execd-token"
assert run_call[3] == {
"code": "print(6*7)",
"context": {"language": "python"},
}
assert run_call[4] == {"stream": True}
assert result.stdout == "42\n"
@pytest.mark.asyncio
async def test_run_code_uses_https_for_scheme_less_endpoint_when_api_base_is_https():
client = FakeHTTPClient()
handle = ContainerHandle(
id="osb_https", provider="opensandbox", domain="https://sandbox.example/v1"
)
handle._hidden_params = {
"execd_endpoint": "execd.example/route/44772",
"execd_headers": {},
}
await OpenSandboxSandboxConfig().arun_code(
container=handle, code="print(1)", client=client
)
assert client.calls[0][1] == "https://execd.example/route/44772/code"
@pytest.mark.asyncio
async def test_run_code_aborts_on_output_over_cap():
client = FakeHTTPClient(execute_lines=["x" * (MAX_OUTPUT_BYTES + 1)])
handle = ContainerHandle(id="osb_big", provider="opensandbox", domain="http://x/v1")
handle._hidden_params = {"execd_endpoint": "execd.local:44772", "execd_headers": {}}
with pytest.raises(ValueError, match="exceeded"):
await OpenSandboxSandboxConfig().arun_code(
container=handle, code="print('x')", client=client
)
@pytest.mark.asyncio
async def test_delete_returns_false_on_404():
client = FakeHTTPClient(delete_status=404)
ok = await OpenSandboxSandboxConfig().adelete_sandbox(
container="osb_gone",
api_key="",
api_base="http://sandbox.local/v1",
client=client,
)
assert ok is False
@pytest.mark.asyncio
async def test_delete_reraises_non_404_http_error():
client = FakeHTTPClient(delete_status=500)
with pytest.raises(httpx.HTTPStatusError):
await OpenSandboxSandboxConfig().adelete_sandbox(
container="osb_err",
api_key="",
api_base="http://sandbox.local/v1",
client=client,
)
@pytest.mark.asyncio
async def test_public_lifecycle_create_run_delete():
client = FakeHTTPClient(
execute_lines=[
sse({"type": "stdout", "text": "42\n"}),
]
)
container = await litellm.acreate_sandbox(
provider="opensandbox", api_key="", api_base=TEST_API_BASE, client=client
)
result = await litellm.arun_code(
provider="opensandbox",
container=container,
code="print(6*7)",
api_key="",
client=client,
)
ok = await litellm.adelete_sandbox(
provider="opensandbox",
container=container,
api_key="",
client=client,
)
assert container.id == "osb_123"
assert result.stdout == "42\n"
assert ok is True
@pytest.mark.asyncio
async def test_code_interpreter_tool_deletes_even_when_run_raises():
client = FakeHTTPClient(execute_raises=RuntimeError("boom"))
with pytest.raises(RuntimeError, match="boom"):
await litellm.acode_interpreter_tool(
provider="opensandbox",
code="1/0",
api_key="",
api_base=TEST_API_BASE,
client=client,
)
assert [call[0] for call in client.calls] == ["POST", "GET", "POST", "DELETE"]
assert client.calls[0][1].endswith("/sandboxes")
assert client.calls[1][1].endswith("/endpoints/44772")
assert client.calls[2][1].endswith("/code")
assert client.calls[3][1].endswith("/sandboxes/osb_123")

View file

@ -0,0 +1,89 @@
"""
Regression tests for the Cloudflare Workers AI text-generation catalog in the
model-cost map.
The Cloudflare list was badly stale (only 4 ancient entries). These tests pin
the newly added current Workers AI models (sourced from Cloudflare's live
``/ai/models/search?task=Text Generation`` catalog) and guard against the root
``model_prices_and_context_window.json`` and the bundled
``litellm/model_prices_and_context_window_backup.json`` drifting out of sync for
the ``cloudflare/`` namespace.
"""
import json
import os
import pytest
import litellm
ROOT_MAP = os.path.join(
os.path.dirname(os.path.dirname(litellm.__file__)),
"model_prices_and_context_window.json",
)
BACKUP_MAP = os.path.join(
os.path.dirname(litellm.__file__),
"model_prices_and_context_window_backup.json",
)
@pytest.fixture(autouse=True)
def _use_local_model_cost_map(monkeypatch):
original_model_cost = litellm.model_cost
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
litellm.model_cost = litellm.get_model_cost_map(url="")
try:
yield
finally:
litellm.model_cost = original_model_cost
def _load(path: str) -> dict:
with open(path, encoding="utf-8") as f:
return json.load(f)
def _cloudflare_keys(data: dict) -> set:
return {k for k in data if k.startswith("cloudflare/")}
def test_glm_5_2_entry_is_present_and_well_formed():
entry = litellm.model_cost["cloudflare/@cf/zai-org/glm-5.2"]
assert entry["litellm_provider"] == "cloudflare"
assert entry["mode"] == "chat"
assert entry["supports_function_calling"] is True
assert entry["input_cost_per_token"] > 0
assert entry["output_cost_per_token"] > 0
def test_vision_model_is_flagged_supports_vision():
entry = litellm.model_cost["cloudflare/@cf/meta/llama-3.2-11b-vision-instruct"]
assert entry["litellm_provider"] == "cloudflare"
assert entry.get("supports_vision") is True
def test_additional_current_models_are_present():
for key in (
"cloudflare/@cf/openai/gpt-oss-120b",
"cloudflare/@cf/meta/llama-3.3-70b-instruct-fp8-fast",
):
entry = litellm.model_cost[key]
assert entry["litellm_provider"] == "cloudflare"
assert entry["mode"] == "chat"
assert entry["supports_function_calling"] is True
assert entry["input_cost_per_token"] > 0
assert entry["output_cost_per_token"] > 0
def test_root_and_backup_have_identical_cloudflare_keys():
if not os.path.exists(ROOT_MAP):
pytest.skip("root cost map only ships in source checkouts")
assert _cloudflare_keys(_load(ROOT_MAP)) == _cloudflare_keys(_load(BACKUP_MAP))
def test_root_and_backup_cloudflare_entries_are_byte_for_byte_equal():
if not os.path.exists(ROOT_MAP):
pytest.skip("root cost map only ships in source checkouts")
root = {k: v for k, v in _load(ROOT_MAP).items() if k.startswith("cloudflare/")}
backup = {k: v for k, v in _load(BACKUP_MAP).items() if k.startswith("cloudflare/")}
assert root == backup

View file

@ -505,6 +505,74 @@ def test_realtime_logging_object_allows_null_transcript_in_conversation_item_add
assert logging_result.results[0]["item"]["content"][0]["transcript"] is None
def test_realtime_logging_object_does_not_validate_unknown_event_types():
"""
A realtime session emits events outside the OpenAIRealtimeEvents union (e.g.
rate_limits.updated, response.function_call_arguments.delta). Building the
logging object must not revalidate every event against the union; doing so
produces thousands of Pydantic ValidationErrors per session, blocks the event
loop, and the raised error discards the session's usage. The events must
survive verbatim, the combined usage must be preserved, and serialization
must stay clean.
"""
import warnings
results: OpenAIRealtimeStreamList = [
{"type": "session.created", "event_id": "ev0", "session": {"id": "s"}},
]
for i in range(50):
results += [
{
"type": "rate_limits.updated",
"event_id": f"rl{i}",
"rate_limits": [{"name": "requests", "limit": 1000, "remaining": 900}],
},
{
"type": "response.function_call_arguments.delta",
"event_id": f"fc{i}",
"delta": "{}",
},
{
"type": "response.done",
"event_id": f"rd{i}",
"response": {
"usage": {
"input_tokens": 4,
"output_tokens": 6,
"total_tokens": 10,
}
},
},
]
usage = RealtimeAPITokenUsageProcessor.collect_and_combine_usage_from_realtime_stream_results(
results=results
)
# On unfixed code this raises pydantic ValidationError instead of returning.
logging_result = RealtimeAPITokenUsageProcessor.create_logging_realtime_object(
usage=usage,
results=results,
)
assert logging_result.usage.total_tokens == 500
assert len(logging_result.results) == len(results)
unknown_types = {
r["type"]
for r in logging_result.results
if r["type"]
in ("rate_limits.updated", "response.function_call_arguments.delta")
}
assert unknown_types == {
"rate_limits.updated",
"response.function_call_arguments.delta",
}
with warnings.catch_warnings():
warnings.simplefilter("error")
dumped = logging_result.model_dump()
assert len(dumped["results"]) == len(results)
def test_realtime_transcription_duration_cost(monkeypatch):
"""
gpt-realtime-whisper transcription sessions are billed by input audio duration

View file

@ -537,3 +537,147 @@ def test_inherit_builtin_cache_pricing_noop_for_unknown_backend():
)
assert model_info == {"input_cost_per_token": 0.000003}
def test_custom_pricing_field_denylist_covers_all_builtin_pricing_fields():
"""The shared-backend-key stripping in Router relies on
CustomPricingLiteLLMParams enumerating every per-deployment pricing field.
If a new pricing field is added to ModelInfoBase but not mirrored here, a
deployment override on that field leaks into the shared backend key and
every sibling deployment reads the wrong rate (LIT-3897). This guard fails
fast when the two drift apart.
"""
import typing
from litellm.types.utils import CustomPricingLiteLLMParams, ModelInfoBase
pricing_markers = ("cost", "price", "uplift", "vector_size", "tiered_pricing")
builtin_pricing_fields = {
name
for name in typing.get_type_hints(ModelInfoBase)
if any(marker in name for marker in pricing_markers)
}
denylisted_fields = set(CustomPricingLiteLLMParams.model_fields.keys())
uncovered = sorted(builtin_pricing_fields - denylisted_fields)
assert not uncovered, (
"ModelInfoBase pricing fields missing from CustomPricingLiteLLMParams; "
f"these would leak into shared backend keys: {uncovered}"
)
def test_tiered_pricing_override_isolated_from_sibling_via_model_info_lookup():
"""LIT-3897: a deployment that overrides a tiered pricing field
(input_cost_per_token_above_272k_tokens) must not pollute the shared
backend key, so a sibling sharing the same backend resolves its pricing
via litellm.get_model_info (the path /model/info uses) without seeing the
override.
"""
backend_model = "gemini/gemini-2.5-flash"
override = 0.000999
builtin_info = litellm.get_model_info(model=backend_model)
assert builtin_info.get("input_cost_per_token_above_272k_tokens") != override
model_keys = {
"lit3897-tiered-custom": litellm.model_cost.get("lit3897-tiered-custom"),
"lit3897-tiered-sibling": litellm.model_cost.get("lit3897-tiered-sibling"),
backend_model: copy.deepcopy(litellm.model_cost.get(backend_model)),
}
try:
Router(
model_list=[
{
"model_name": "custom-priced-flash",
"litellm_params": {
"model": backend_model,
"api_key": "fake-key-tiered-1",
},
"model_info": {
"id": "lit3897-tiered-custom",
"input_cost_per_token_above_272k_tokens": override,
"cache_read_input_token_cost_above_272k_tokens": override,
},
},
{
"model_name": "gemini-2.5-flash",
"litellm_params": {
"model": backend_model,
"api_key": "fake-key-tiered-2",
},
"model_info": {"id": "lit3897-tiered-sibling"},
},
],
)
shared = litellm.get_model_info(model=backend_model)
assert shared.get("input_cost_per_token_above_272k_tokens") != override, (
"Tiered override leaked into the shared backend key; siblings read "
"the wrong rate via /model/info"
)
assert shared.get("cache_read_input_token_cost_above_272k_tokens") != override
custom_entry = litellm.model_cost["lit3897-tiered-custom"]
assert custom_entry["input_cost_per_token_above_272k_tokens"] == override
assert custom_entry["cache_read_input_token_cost_above_272k_tokens"] == override
finally:
_restore_model_cost_entries(model_keys)
def test_custom_pricing_isolated_from_sibling_via_proxy_model_info_path():
"""LIT-3897 end to end through the proxy resolution helper: the override
deployment reports its custom input rate while the sibling keeps the
canonical gemini rate when /model/info resolves each deployment. Mirrors the
ticket config where the override is set on litellm_params.
"""
from litellm.proxy.proxy_server import _get_proxy_model_info
backend_model = "gemini/gemini-2.5-flash"
override_input = 5e-05
override_output = 1e-04
builtin_info = litellm.get_model_info(model=backend_model)
builtin_input = builtin_info["input_cost_per_token"]
assert builtin_input != override_input
model_keys = {
"lit3897-proxy-custom": litellm.model_cost.get("lit3897-proxy-custom"),
"lit3897-proxy-sibling": litellm.model_cost.get("lit3897-proxy-sibling"),
backend_model: copy.deepcopy(litellm.model_cost.get(backend_model)),
}
try:
router = Router(
model_list=[
{
"model_name": "custom-priced-flash",
"litellm_params": {
"model": backend_model,
"api_key": "fake-key-proxy-1",
"input_cost_per_token": override_input,
"output_cost_per_token": override_output,
},
"model_info": {"id": "lit3897-proxy-custom"},
},
{
"model_name": "gemini-2.5-flash",
"litellm_params": {
"model": backend_model,
"api_key": "fake-key-proxy-2",
},
"model_info": {"id": "lit3897-proxy-sibling"},
},
],
)
resolved = {
m["model_name"]: _get_proxy_model_info(model=copy.deepcopy(m))[
"model_info"
]["input_cost_per_token"]
for m in router.model_list
}
assert resolved["custom-priced-flash"] == override_input
assert resolved["gemini-2.5-flash"] == builtin_input
assert resolved["gemini-2.5-flash"] != resolved["custom-priced-flash"]
finally:
_restore_model_cost_entries(model_keys)

View file

@ -4,8 +4,9 @@ GitHub Issue: #18968 - Per-deployment max_retries/num_retries in litellm_params
"""
import pytest
from unittest.mock import MagicMock, patch
from unittest.mock import patch
import litellm
from litellm import Router
@ -188,3 +189,133 @@ class TestPerDeploymentNumRetries:
# Verify num_retries was converted from string to int
assert exc.num_retries == 6
class TestNumRetriesNoneGuard:
"""
Regression tests for the num_retries=None TypeError in async_function_with_retries.
When num_retries reaches async_function_with_retries as None - e.g. a caller passes
num_retries=None explicitly (dict.get() does not fall back on an existing None value),
an auto_router/complexity_router path does not propagate it, or
Router.update_settings(num_retries=None) is used - AND the underlying call fails with a
retryable error, the comparison `if num_retries > 0:` raised:
TypeError: '>' not supported between instances of 'NoneType' and 'int'
This masked the real upstream error (rate limit / connection / 5xx) behind a TypeError.
Related issues: #23316, #25889, #23699, #28126.
"""
@staticmethod
def _mock_router(num_retries=2):
return Router(
model_list=[
{
"model_name": "mock-model",
"litellm_params": {
"model": "gpt-4o-mini",
"mock_response": "ok",
},
}
],
num_retries=num_retries,
)
def test_update_kwargs_normalises_explicit_none_to_router_default(self):
"""
_update_kwargs_before_fallbacks must normalise an explicit num_retries=None to
the router default (not leave it as None), while preserving an explicit 0.
"""
router = self._mock_router(num_retries=4)
# explicit None -> router default
kwargs = {"num_retries": None}
router._update_kwargs_before_fallbacks(model="mock-model", kwargs=kwargs)
assert kwargs["num_retries"] == 4
# explicit 0 is preserved (retries stay disabled)
kwargs = {"num_retries": 0}
router._update_kwargs_before_fallbacks(model="mock-model", kwargs=kwargs)
assert kwargs["num_retries"] == 0
# absent -> router default (unchanged behaviour)
kwargs = {}
router._update_kwargs_before_fallbacks(model="mock-model", kwargs=kwargs)
assert kwargs["num_retries"] == 4
# explicit None with router default also None -> 0 (mirrors the downstream guard)
router.num_retries = None # simulate update_settings(num_retries=None) (#28126)
kwargs = {"num_retries": None}
router._update_kwargs_before_fallbacks(model="mock-model", kwargs=kwargs)
assert kwargs["num_retries"] == 0
@pytest.mark.asyncio
async def test_acompletion_num_retries_none_does_not_raise_typeerror(self):
"""
Per-request num_retries=None + a retryable error must NOT raise TypeError.
The router falls back to its configured num_retries and retries the (transient)
error, so the request succeeds.
"""
router = self._mock_router(num_retries=2)
with patch("asyncio.sleep", return_value=None):
response = await router.acompletion(
model="mock-model",
messages=[{"role": "user", "content": "hi"}],
num_retries=None, # the trigger
mock_testing_rate_limit_error=True, # retryable error path
)
assert response.choices[0].message.content == "ok"
@pytest.mark.asyncio
async def test_async_function_with_retries_none_falls_back_to_zero(self):
"""
When both the per-request value AND the router-level setting are None
(e.g. after Router.update_settings(num_retries=None), #28126), num_retries must
fall back to 0 and the real retryable error must surface - not a TypeError.
"""
router = self._mock_router(num_retries=0)
router.num_retries = None # simulate update_settings(num_retries=None)
async def failing_fn(*args, **kwargs):
raise litellm.RateLimitError(
message="boom", model="mock-model", llm_provider="openai"
)
with patch("asyncio.sleep", return_value=None):
with pytest.raises(litellm.RateLimitError):
await router.async_function_with_retries(
original_function=failing_fn,
model="mock-model",
messages=[{"role": "user", "content": "hi"}],
num_retries=None,
)
@pytest.mark.asyncio
async def test_async_function_with_retries_none_falls_back_to_router_default(self):
"""
A None per-request num_retries falls back to the router-level setting, so retries
still happen (original_function is invoked more than once) before the real error
is raised - proving None did not silently disable retries or crash.
"""
router = self._mock_router(num_retries=3)
calls = {"n": 0}
async def failing_fn(*args, **kwargs):
calls["n"] += 1
raise litellm.InternalServerError(
message="boom", model="mock-model", llm_provider="openai"
)
with patch("asyncio.sleep", return_value=None):
with pytest.raises(litellm.InternalServerError):
await router.async_function_with_retries(
original_function=failing_fn,
model="mock-model",
messages=[{"role": "user", "content": "hi"}],
metadata={}, # populated by acompletion in the real path; log_retry needs it
num_retries=None,
)
# 1 initial attempt + at least 1 retry -> proves None fell back to a positive int
assert calls["n"] >= 2

View file

@ -0,0 +1,187 @@
import json
from unittest.mock import MagicMock
import pytest
import litellm
from litellm.proxy.proxy_server import _should_include_fallback_errors
from litellm.router import Router
from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict
def test_apply_fallback_hidden_params_copies_from_fallback_response():
fallback_errors = [
{
"message": "litellm.RateLimitError: upstream limited request",
"type": "RateLimitError",
"param": None,
"code": "429",
}
]
chunk = litellm.ModelResponseStream(
id="test",
model="openai/internal-fallback",
choices=[],
)
chunk._hidden_params = {
"additional_headers": {"x-existing-chunk-header": "keep"},
"model_id": "chunk-model-id",
}
fallback_response = MagicMock()
fallback_response._hidden_params = {
"additional_headers": {
"x-litellm-attempted-fallbacks": 1,
"x-litellm-model-group": "fallback-model",
"x-litellm-fallback-errors": json.dumps(fallback_errors),
},
"api_base": "https://fallback.example",
}
Router._apply_fallback_hidden_params_to_item(
fallback_item=chunk,
prepared_fallback_hidden_params=Router._prepare_fallback_hidden_params(
fallback_response
),
)
assert chunk._hidden_params["api_base"] == "https://fallback.example"
assert chunk._hidden_params["model_id"] == "chunk-model-id"
assert chunk._hidden_params["additional_headers"] == {
"x-existing-chunk-header": "keep",
"x-litellm-attempted-fallbacks": 1,
"x-litellm-model-group": "fallback-model",
"x-litellm-fallback-errors": json.dumps(fallback_errors),
}
def _two_group_fallback_router() -> Router:
return litellm.Router(
model_list=[
{
"model_name": "primary-model",
"litellm_params": {"model": "openai/gpt-fake", "api_key": "sk-fake"},
},
{
"model_name": "fallback-model",
"litellm_params": {"model": "openai/gpt-fake-2", "api_key": "sk-fake"},
},
],
fallbacks=[{"primary-model": ["fallback-model"]}],
)
def _additional_headers(response: object) -> dict:
return get_hidden_params_dict(response).get("additional_headers", {})
@pytest.mark.asyncio
async def test_include_fallback_errors_propagates_through_router():
router = _two_group_fallback_router()
response = await router.acompletion(
model="primary-model",
messages=[{"role": "user", "content": "Hello"}],
mock_testing_fallbacks=True,
mock_response="fallback success",
include_fallback_errors=True,
)
headers = _additional_headers(response)
assert headers["x-litellm-attempted-fallbacks"] == 1
errors = json.loads(headers["x-litellm-fallback-errors"])
assert isinstance(errors, list) and len(errors) >= 1
assert set(errors[0].keys()) == {"message", "type", "param", "code"}
@pytest.mark.asyncio
async def test_router_omits_fallback_errors_without_opt_in():
router = _two_group_fallback_router()
response = await router.acompletion(
model="primary-model",
messages=[{"role": "user", "content": "Hello"}],
mock_testing_fallbacks=True,
mock_response="fallback success",
)
headers = _additional_headers(response)
assert headers["x-litellm-attempted-fallbacks"] == 1
assert "x-litellm-fallback-errors" not in headers
def test_prepare_fallback_hidden_params_no_additional_headers():
class FakeResponse:
_hidden_params = {"api_base": "http://example.com"}
hidden_params, headers = Router._prepare_fallback_hidden_params(FakeResponse())
assert hidden_params == {"api_base": "http://example.com"}
assert headers == {}
def test_apply_fallback_hidden_params_to_item_none_item():
Router._apply_fallback_hidden_params_to_item(
None, ({"api_base": "http://fallback.example"}, {"x-custom": "value"})
)
def test_apply_fallback_hidden_params_to_item_no_existing_additional_headers():
class FakeChunk:
_hidden_params = {"model_id": "test-id"}
chunk = FakeChunk()
Router._apply_fallback_hidden_params_to_item(
chunk,
(
{"api_base": "http://fallback.example"},
{"x-litellm-attempted-fallbacks": 1},
),
)
assert chunk._hidden_params["api_base"] == "http://fallback.example"
assert chunk._hidden_params["model_id"] == "test-id"
assert chunk._hidden_params["additional_headers"] == {
"x-litellm-attempted-fallbacks": 1
}
@pytest.mark.asyncio
async def test_set_response_headers_adds_model_group_to_streaming_wrapper():
class StreamingWrapper:
def __init__(self):
self._hidden_params = {"additional_headers": {"x-existing": "keep"}}
router = litellm.Router(model_list=[])
response = StreamingWrapper()
result = await router.set_response_headers(
response=response,
model_group="fallback-model",
)
assert result is response
assert response._hidden_params["additional_headers"] == {
"x-existing": "keep",
"x-litellm-model-group": "fallback-model",
}
def test_should_include_fallback_errors_gated_by_operator_setting():
request_data: dict = {"include_fallback_errors": True}
import litellm.proxy.proxy_server as ps
original = ps.general_settings.copy() if isinstance(ps.general_settings, dict) else {}
try:
ps.general_settings = {}
assert _should_include_fallback_errors(request_data) is False
ps.general_settings = {"expose_fallback_errors_to_caller": False}
assert _should_include_fallback_errors(request_data) is False
ps.general_settings = {"expose_fallback_errors_to_caller": True}
assert _should_include_fallback_errors(request_data) is True
ps.general_settings = {"expose_fallback_errors_to_caller": True}
assert _should_include_fallback_errors({}) is False
finally:
ps.general_settings = original

View file

@ -56,29 +56,58 @@ def test_paths_outside_repo_are_skipped():
def test_at_or_under_ceiling_passes():
budget = {"no-any-return": {"baseline": 5, "slack": 0}}
assert gate.evaluate({"no-any-return": 5}, budget) == []
assert gate.evaluate({"no-any-return": 5}, {}, budget) == []
def test_one_more_error_than_ceiling_fails():
budget = {"no-any-return": {"baseline": 5, "slack": 0}}
assert gate.evaluate({"no-any-return": 6}, budget) == [
gate.Breach("no-any-return", 6, 5)
assert gate.evaluate({"no-any-return": 6}, {}, budget) == [
gate.Breach("no-any-return", 6, 5, 6)
]
def test_slack_absorbs_small_increase_then_fails_past_it():
budget = {"arg-type": {"baseline": 5, "slack": 5}}
assert gate.evaluate({"arg-type": 10}, budget) == []
assert gate.evaluate({"arg-type": 11}, budget) == [gate.Breach("arg-type", 11, 10)]
assert gate.evaluate({"arg-type": 10}, {}, budget) == []
assert gate.evaluate({"arg-type": 11}, {}, budget) == [
gate.Breach("arg-type", 11, 10, 11)
]
def test_unbudgeted_new_code_uses_default_slack():
assert gate.evaluate({"brand-new": gate.DEFAULT_SLACK}, {}) == []
assert gate.evaluate({"brand-new": gate.DEFAULT_SLACK + 1}, {}) == [
gate.Breach("brand-new", gate.DEFAULT_SLACK + 1, gate.DEFAULT_SLACK)
assert gate.evaluate({"brand-new": gate.DEFAULT_SLACK}, {}, {}) == []
assert gate.evaluate({"brand-new": gate.DEFAULT_SLACK + 1}, {}, {}) == [
gate.Breach(
"brand-new",
gate.DEFAULT_SLACK + 1,
gate.DEFAULT_SLACK,
gate.DEFAULT_SLACK + 1,
)
]
def test_drift_already_over_cap_in_base_is_not_blamed_on_a_flat_change():
# The bystander case: a rule sits over its ceiling because two earlier PRs
# summed past it. A PR that branches off that base and adds nothing must pass
# -- total > cap but total == base, so the `> base` guard spares it.
budget = {"arg-type": {"baseline": 5, "slack": 5}}
assert gate.evaluate({"arg-type": 12}, {"arg-type": 12}, budget) == []
def test_change_that_grows_an_over_cap_rule_is_blamed_for_only_what_it_added():
# Over cap AND above base: blamed, and `added` is the delta vs base, not the
# whole overage, so the message points at this change's contribution.
budget = {"arg-type": {"baseline": 5, "slack": 5}}
assert gate.evaluate({"arg-type": 14}, {"arg-type": 12}, budget) == [
gate.Breach("arg-type", 14, 10, 2)
]
def test_reducing_an_over_cap_rule_below_base_passes():
budget = {"arg-type": {"baseline": 5, "slack": 5}}
assert gate.evaluate({"arg-type": 11}, {"arg-type": 12}, budget) == []
def test_no_output_against_a_nonempty_budget_is_a_vacuous_run():
# A crashed type checker emits nothing; the gate must not certify it as clean.
budget = {"no-untyped-def": {"baseline": 4888, "slack": 10}}

View file

@ -8,9 +8,16 @@ Usage:
pytest tests/test_litellm/types/test_completion.py -v
"""
import dataclasses
from typing import List
from litellm.types.completion import CompletionRequest, ChatCompletionMessageParam
import pytest
from litellm.types.completion import (
ChatCompletionMessageParam,
CompletionRequest,
_CompletionDispatchContext,
)
def test_completion_request_messages_type_validation():
@ -146,3 +153,55 @@ def test_completion_request_with_all_params():
assert request.presence_penalty == 0.0
assert request.stream is False
assert request.n == 1
def _build_dispatch_context() -> _CompletionDispatchContext:
return _CompletionDispatchContext(
_azure_detection_model="gpt-4o",
acompletion=False,
api_base=None,
api_key=None,
api_version=None,
client=None,
custom_llm_provider="openai",
custom_prompt_dict={},
extra_headers=None,
headers={},
hf_model_name=None,
kwargs={},
litellm_params={},
logger_fn=None,
logging=None, # type: ignore[arg-type]
max_retries=None,
max_tokens=None,
messages=[],
metadata=None,
model="gpt-4o",
model_response=None, # type: ignore[arg-type]
optional_params={},
organization=None,
provider_config=None,
shared_session=None,
stream=None,
temperature=None,
text_completion=False,
timeout=None,
top_p=None,
)
def test_dispatch_context_is_frozen():
"""A helper must not be able to re-route the call by rebinding a dispatch
input mid-flight; this pins the frozen invariant the dispatch shape relies on."""
ctx = _build_dispatch_context()
with pytest.raises(dataclasses.FrozenInstanceError):
ctx.model = "claude-haiku-4-5" # type: ignore[misc]
with pytest.raises(dataclasses.FrozenInstanceError):
ctx.custom_llm_provider = "anthropic" # type: ignore[misc]
def test_dispatch_context_uses_slots():
"""slots=True keeps the per-call context lightweight (no per-instance __dict__)."""
ctx = _build_dispatch_context()
assert not hasattr(ctx, "__dict__")
assert hasattr(type(ctx), "__slots__")

View file

@ -1097,3 +1097,52 @@ describe("OldTeams - delete team warning copy", () => {
);
});
});
describe("OldTeams - LIT-2530 organization stays optional for proxy admin with a single org", () => {
beforeEach(() => {
vi.clearAllMocks();
mockTeamInfoView.mockClear();
vi.mocked(fetchAvailableModelsForTeamOrKey).mockResolvedValue(["gpt-4"]);
vi.mocked(fetchMCPAccessGroups).mockResolvedValue([]);
vi.mocked(getGuardrailsList).mockResolvedValue({ guardrails: [] });
vi.mocked(teamListCall).mockResolvedValue({ teams: [], total: 0, page: 1, page_size: 100, total_pages: 1 });
vi.mocked(teamCreateCall).mockResolvedValue({
team_id: "new-team-1",
team_alias: "No Org Team",
models: ["gpt-4"],
organization_id: null,
keys: [],
members_with_roles: [],
spend: 0,
});
mockUseOrganizations.mockReturnValue({
data: [{ organization_id: "org-1", organization_alias: "Org 1", models: [], members: [] }],
});
});
it("creates a team with no organization when exactly one organization exists", async () => {
renderWithQueryClient(<OldTeams accessToken="test-token" userID="user-123" userRole="Admin" />);
const createButton = screen.getAllByRole("button", { name: /create team/i })[0];
act(() => {
fireEvent.click(createButton);
});
await waitFor(() => {
expect(screen.getByLabelText(/team name/i)).toBeInTheDocument();
});
fireEvent.change(screen.getByLabelText(/team name/i), { target: { value: "No Org Team" } });
fireEvent.change(screen.getByTestId("create-team-models-select"), { target: { value: "gpt-4" } });
const submitButtons = screen.getAllByRole("button", { name: /create team/i });
fireEvent.click(submitButtons[submitButtons.length - 1]);
await waitFor(() => {
expect(teamCreateCall).toHaveBeenCalledWith(
"test-token",
expect.objectContaining({ team_alias: "No Org Team", organization_id: null }),
);
});
});
});

View file

@ -262,14 +262,15 @@ const Teams: React.FC<TeamProps> = ({ accessToken, userID, userRole, premiumUser
useEffect(() => {
if (isTeamModalVisible) {
const adminOrgs = getAdminOrganizations(userRole, userID, organizations);
const isOrgAdmin = userRole !== "Admin";
// If there's exactly one organization the user is admin for, preselect it
if (adminOrgs.length === 1) {
// Org admins must scope a team to an org, so with exactly one we preselect it.
// Proxy admins can create org-less teams, so the field stays optional regardless of org count.
if (isOrgAdmin && adminOrgs.length === 1) {
const org = adminOrgs[0];
form.setFieldValue("organization_id", org.organization_id);
setCurrentOrgForCreateTeam(org);
} else {
// Reset the organization selection for multiple orgs
form.setFieldValue("organization_id", currentOrg?.organization_id || null);
setCurrentOrgForCreateTeam(currentOrg);
}
@ -1132,7 +1133,7 @@ const Teams: React.FC<TeamProps> = ({ accessToken, userID, userRole, premiumUser
: []
}
help={
isSingleOrg
isOrgAdmin && isSingleOrg
? "You can only create teams within this organization"
: isOrgAdmin
? "required"
@ -1142,7 +1143,7 @@ const Teams: React.FC<TeamProps> = ({ accessToken, userID, userRole, premiumUser
<Select
showSearch
allowClear={!isOrgAdmin}
disabled={isSingleOrg}
disabled={isOrgAdmin && isSingleOrg}
placeholder={hasNoOrgs ? "No organizations available" : "Search or select an Organization"}
onChange={(value) => {
form.setFieldValue("organization_id", value);

View file

@ -25112,6 +25112,8 @@ export interface components {
} | null;
/** Adaptive Router Default Model */
adaptive_router_default_model?: string | null;
/** Annotation Cost Per Page */
annotation_cost_per_page?: number | null;
/** Api Base */
api_base?: string | null;
/** Api Key */
@ -25156,8 +25158,12 @@ export interface components {
cache_read_input_token_cost_above_200k_tokens?: number | null;
/** Cache Read Input Token Cost Above 200K Tokens Priority */
cache_read_input_token_cost_above_200k_tokens_priority?: number | null;
/** Cache Read Input Token Cost Above 272K Tokens */
cache_read_input_token_cost_above_272k_tokens?: number | null;
/** Cache Read Input Token Cost Above 272K Tokens Priority */
cache_read_input_token_cost_above_272k_tokens_priority?: number | null;
/** Cache Read Input Token Cost Above 512K Tokens */
cache_read_input_token_cost_above_512k_tokens?: number | null;
/** Cache Read Input Token Cost Flex */
cache_read_input_token_cost_flex?: number | null;
/** Cache Read Input Token Cost Priority */
@ -25194,6 +25200,8 @@ export interface components {
input_cost_per_image?: number | null;
/** Input Cost Per Image Above 128K Tokens */
input_cost_per_image_above_128k_tokens?: number | null;
/** Input Cost Per Image Token */
input_cost_per_image_token?: number | null;
/** Input Cost Per Pixel */
input_cost_per_pixel?: number | null;
/** Input Cost Per Query */
@ -25208,8 +25216,12 @@ export interface components {
input_cost_per_token_above_200k_tokens?: number | null;
/** Input Cost Per Token Above 200K Tokens Priority */
input_cost_per_token_above_200k_tokens_priority?: number | null;
/** Input Cost Per Token Above 272K Tokens */
input_cost_per_token_above_272k_tokens?: number | null;
/** Input Cost Per Token Above 272K Tokens Priority */
input_cost_per_token_above_272k_tokens_priority?: number | null;
/** Input Cost Per Token Above 512K Tokens */
input_cost_per_token_above_512k_tokens?: number | null;
/** Input Cost Per Token Batches */
input_cost_per_token_batches?: number | null;
/** Input Cost Per Token Cache Hit */
@ -25255,6 +25267,10 @@ export interface components {
model_info?: {
[key: string]: unknown;
} | null;
/** Ocr Cost Per Credit */
ocr_cost_per_credit?: number | null;
/** Ocr Cost Per Page */
ocr_cost_per_page?: number | null;
/** Organization */
organization?: string | null;
/** Output Cost Per Audio Per Second */
@ -25285,8 +25301,12 @@ export interface components {
output_cost_per_token_above_200k_tokens?: number | null;
/** Output Cost Per Token Above 200K Tokens Priority */
output_cost_per_token_above_200k_tokens_priority?: number | null;
/** Output Cost Per Token Above 272K Tokens */
output_cost_per_token_above_272k_tokens?: number | null;
/** Output Cost Per Token Above 272K Tokens Priority */
output_cost_per_token_above_272k_tokens_priority?: number | null;
/** Output Cost Per Token Above 512K Tokens */
output_cost_per_token_above_512k_tokens?: number | null;
/** Output Cost Per Token Batches */
output_cost_per_token_batches?: number | null;
/** Output Cost Per Token Flex */
@ -25295,6 +25315,8 @@ export interface components {
output_cost_per_token_priority?: number | null;
/** Output Cost Per Video Per Second */
output_cost_per_video_per_second?: number | null;
/** Output Vector Size */
output_vector_size?: number | null;
/** Quality Router Config */
quality_router_config?: {
[key: string]: unknown;
@ -25303,6 +25325,10 @@ export interface components {
quality_router_default_model?: string | null;
/** Region Name */
region_name?: string | null;
/** Regional Processing Uplift Multiplier Eu */
regional_processing_uplift_multiplier_eu?: number | null;
/** Regional Processing Uplift Multiplier Us */
regional_processing_uplift_multiplier_us?: number | null;
/** Rpm */
rpm?: number | null;
/** S3 Bucket Name */
@ -32794,6 +32820,8 @@ export interface components {
} | null;
/** Adaptive Router Default Model */
adaptive_router_default_model?: string | null;
/** Annotation Cost Per Page */
annotation_cost_per_page?: number | null;
/** Api Base */
api_base?: string | null;
/** Api Key */
@ -32838,8 +32866,12 @@ export interface components {
cache_read_input_token_cost_above_200k_tokens?: number | null;
/** Cache Read Input Token Cost Above 200K Tokens Priority */
cache_read_input_token_cost_above_200k_tokens_priority?: number | null;
/** Cache Read Input Token Cost Above 272K Tokens */
cache_read_input_token_cost_above_272k_tokens?: number | null;
/** Cache Read Input Token Cost Above 272K Tokens Priority */
cache_read_input_token_cost_above_272k_tokens_priority?: number | null;
/** Cache Read Input Token Cost Above 512K Tokens */
cache_read_input_token_cost_above_512k_tokens?: number | null;
/** Cache Read Input Token Cost Flex */
cache_read_input_token_cost_flex?: number | null;
/** Cache Read Input Token Cost Priority */
@ -32876,6 +32908,8 @@ export interface components {
input_cost_per_image?: number | null;
/** Input Cost Per Image Above 128K Tokens */
input_cost_per_image_above_128k_tokens?: number | null;
/** Input Cost Per Image Token */
input_cost_per_image_token?: number | null;
/** Input Cost Per Pixel */
input_cost_per_pixel?: number | null;
/** Input Cost Per Query */
@ -32890,8 +32924,12 @@ export interface components {
input_cost_per_token_above_200k_tokens?: number | null;
/** Input Cost Per Token Above 200K Tokens Priority */
input_cost_per_token_above_200k_tokens_priority?: number | null;
/** Input Cost Per Token Above 272K Tokens */
input_cost_per_token_above_272k_tokens?: number | null;
/** Input Cost Per Token Above 272K Tokens Priority */
input_cost_per_token_above_272k_tokens_priority?: number | null;
/** Input Cost Per Token Above 512K Tokens */
input_cost_per_token_above_512k_tokens?: number | null;
/** Input Cost Per Token Batches */
input_cost_per_token_batches?: number | null;
/** Input Cost Per Token Cache Hit */
@ -32937,6 +32975,10 @@ export interface components {
model_info?: {
[key: string]: unknown;
} | null;
/** Ocr Cost Per Credit */
ocr_cost_per_credit?: number | null;
/** Ocr Cost Per Page */
ocr_cost_per_page?: number | null;
/** Organization */
organization?: string | null;
/** Output Cost Per Audio Per Second */
@ -32967,8 +33009,12 @@ export interface components {
output_cost_per_token_above_200k_tokens?: number | null;
/** Output Cost Per Token Above 200K Tokens Priority */
output_cost_per_token_above_200k_tokens_priority?: number | null;
/** Output Cost Per Token Above 272K Tokens */
output_cost_per_token_above_272k_tokens?: number | null;
/** Output Cost Per Token Above 272K Tokens Priority */
output_cost_per_token_above_272k_tokens_priority?: number | null;
/** Output Cost Per Token Above 512K Tokens */
output_cost_per_token_above_512k_tokens?: number | null;
/** Output Cost Per Token Batches */
output_cost_per_token_batches?: number | null;
/** Output Cost Per Token Flex */
@ -32977,6 +33023,8 @@ export interface components {
output_cost_per_token_priority?: number | null;
/** Output Cost Per Video Per Second */
output_cost_per_video_per_second?: number | null;
/** Output Vector Size */
output_vector_size?: number | null;
/** Quality Router Config */
quality_router_config?: {
[key: string]: unknown;
@ -32985,6 +33033,10 @@ export interface components {
quality_router_default_model?: string | null;
/** Region Name */
region_name?: string | null;
/** Regional Processing Uplift Multiplier Eu */
regional_processing_uplift_multiplier_eu?: number | null;
/** Regional Processing Uplift Multiplier Us */
regional_processing_uplift_multiplier_us?: number | null;
/** Rpm */
rpm?: number | null;
/** S3 Bucket Name */

16
uv.lock generated
View file

@ -9,7 +9,7 @@ resolution-markers = [
]
[options]
exclude-newer = "2026-06-17T22:13:32.966924Z"
exclude-newer = "0001-01-01T00:00:00Z" # This has no effect and is included for backwards compatibility when using relative exclude-newer values.
exclude-newer-span = "P3D"
[manifest]
@ -1489,6 +1489,18 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/c1/ea/53f2148663b321f21b5a606bd5f191517cf40b7072c0497d3c92c4a13b1e/executing-2.2.1-py2.py3-none-any.whl", hash = "sha256:760643d3452b4d777d295bb167ccc74c64a81df23fb5e08eff250c425a4b2017", size = 28317, upload-time = "2025-09-01T09:48:08.5Z" },
]
[[package]]
name = "expression"
version = "5.6.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "typing-extensions" },
]
sdist = { url = "https://files.pythonhosted.org/packages/43/c7/bb061623b5815566bda69f5e9d156e38a97ebb383b8db3d2dedb26415466/expression-5.6.0.tar.gz", hash = "sha256:454f6fe138347194a43c7f878d958efe9b84b9cc770e462010c7a52e18058065", size = 59147, upload-time = "2025-02-19T09:37:37.432Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/6b/a2/656b8bebe495117342a8676ccabf52b3885ce11a856c8dfe1fbbdc250d2d/expression-5.6.0-py3-none-any.whl", hash = "sha256:f5c62e38186c9287e088dee9cf3939b0bbde21cb4c59571872154a53d33dd7c0", size = 69673, upload-time = "2025-02-19T09:37:35.476Z" },
]
[[package]]
name = "fakeredis"
version = "2.34.1"
@ -3297,6 +3309,7 @@ proxy = [
{ name = "backoff" },
{ name = "boto3" },
{ name = "cryptography" },
{ name = "expression" },
{ name = "fastapi" },
{ name = "fastapi-sso" },
{ name = "granian" },
@ -3460,6 +3473,7 @@ requires-dist = [
{ name = "ddtrace", marker = "extra == 'proxy-runtime'", specifier = ">=2.19.0,<3.0" },
{ name = "detect-secrets", marker = "extra == 'proxy-runtime'", specifier = ">=1.5.0,<2.0" },
{ name = "diskcache", marker = "extra == 'caching'", specifier = ">=5.6.3,<6.0" },
{ name = "expression", marker = "extra == 'proxy'", specifier = ">=5.6.0,<6.0" },
{ name = "fastapi", marker = "extra == 'proxy'", specifier = ">=0.136.3,<1.0" },
{ name = "fastapi-sso", marker = "extra == 'proxy'", specifier = ">=0.19.0,<1.0" },
{ name = "fastuuid", specifier = ">=0.14.0,<1.0" },