mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
commit
b7367cf8e7
97 changed files with 10786 additions and 3199 deletions
8
.github/workflows/test-linting.yml
vendored
8
.github/workflows/test-linting.yml
vendored
|
|
@ -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: |
|
||||
|
|
|
|||
1
.github/workflows/test-unit-misc.yml
vendored
1
.github/workflows/test-unit-misc.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
3
Makefile
3
Makefile
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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`
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
1
litellm/llms/opensandbox/__init__.py
Normal file
1
litellm/llms/opensandbox/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
|||
|
||||
1
litellm/llms/opensandbox/sandbox/__init__.py
Normal file
1
litellm/llms/opensandbox/sandbox/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
|||
|
||||
598
litellm/llms/opensandbox/sandbox/transformation.py
Normal file
598
litellm/llms/opensandbox/sandbox/transformation.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -32,6 +32,9 @@ class VertexAIRealtimeConfig(GeminiRealtimeConfig):
|
|||
self._project = project
|
||||
self._location = location
|
||||
|
||||
def _include_function_response_id(self) -> bool:
|
||||
return False
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# URL
|
||||
# ------------------------------------------------------------------
|
||||
|
|
|
|||
6638
litellm/main.py
6638
litellm/main.py
File diff suppressed because it is too large
Load diff
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
@ -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
|
||||
|
|
@ -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")
|
||||
)
|
||||
|
|
@ -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]
|
||||
|
|
@ -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}"))
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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).
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -300,7 +300,7 @@
|
|||
"slack": 3
|
||||
},
|
||||
"RET504": {
|
||||
"baseline": 709,
|
||||
"baseline": 702,
|
||||
"slack": 20
|
||||
},
|
||||
"RUF010": {
|
||||
|
|
|
|||
|
|
@ -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__":
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
34
tests/llm_translation/test_bedrock_embedding_pricing.py
Normal file
34
tests/llm_translation/test_bedrock_embedding_pricing.py
Normal 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"
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
@ -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)
|
||||
|
|
@ -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"
|
||||
|
|
@ -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"
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 ---
|
||||
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
139
tests/test_litellm/router_utils/test_fallback_event_handlers.py
Normal file
139
tests/test_litellm/router_utils/test_fallback_event_handlers.py
Normal 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
|
||||
647
tests/test_litellm/sandbox/test_opensandbox_sandbox.py
Normal file
647
tests/test_litellm/sandbox/test_opensandbox_sandbox.py
Normal 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")
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
187
tests/test_litellm/test_router_streaming_fallback_metadata.py
Normal file
187
tests/test_litellm/test_router_streaming_fallback_metadata.py
Normal 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
|
||||
|
|
@ -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}}
|
||||
|
|
|
|||
|
|
@ -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__")
|
||||
|
|
|
|||
|
|
@ -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 }),
|
||||
);
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
52
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
52
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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
16
uv.lock
generated
|
|
@ -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" },
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue