Merge remote-tracking branch 'upstream/litellm_internal_staging' into litellm_hotfix_rust_mistral_ocr

# Conflicts:
#	tests/test_litellm/interactions/test_openapi_compliance.py
This commit is contained in:
Ishaan Jaffer 2026-06-22 21:14:26 -07:00
commit 39d132f2a6
No known key found for this signature in database
132 changed files with 13895 additions and 3113 deletions

View file

@ -120,6 +120,9 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = (
"/robots.txt",
# Health (k8s probes)
"/health",
# Plugin system
"/api/plugins",
"/plugin-proxy/",
)
BACKEND_EXACT_PATHS: frozenset[str] = frozenset(

141
docs/plugin_architecture.md Normal file
View file

@ -0,0 +1,141 @@
# LiteLLM Plugin Architecture
Plugins let external services appear as selectable modes in the litellm UI sidebar alongside the AI Gateway.
---
## Quick start
### 1. Configure the plugin
Add a `plugins` block to your litellm `config.yaml`:
```yaml
general_settings:
master_key: sk-...
plugins:
- name: my-plugin # unique identifier (no spaces)
display_name: My Plugin # shown in the UI dropdown
url: "https://my-plugin.example.com"
plugin_key: "sk-..." # plugin's own auth credential
```
`plugin_key` is injected as `Authorization: Bearer <plugin_key>` on every
request proxied through `/plugin-proxy/my-plugin/*`. The caller's litellm
credential is stripped before forwarding so the plugin never receives a live
litellm API key.
### 2. Implement two endpoints on your service
| Endpoint | Method | Purpose |
|---|---|---|
| `GET /api/plugin-manifest` | public | Returns plugin metadata for the UI |
| `POST /api/plugin-auth` | public | Decrypts the identity claim for seamless sign-in |
#### `GET /api/plugin-manifest`
```json
{
"name": "my-plugin",
"display_name": "My Plugin",
"version": "1.0.0",
"nav_items": [
{ "key": "home", "label": "Home", "icon": "HomeOutlined", "path": "/" },
{ "key": "reports", "label": "Reports", "icon": "BarChartOutlined", "path": "/reports" }
],
"capabilities": ["reports", "data"]
}
```
#### `POST /api/plugin-auth`
Receives `{ "session_claim": "<fernet-ciphertext>" }`.
The proxy never shares `LITELLM_SALT_KEY` with your plugin. Each plugin is
provisioned with its own dedicated key, derived as
`HMAC-SHA256(LITELLM_SALT_KEY, plugin_name)`. Compute it once on the proxy
host and hand the result to your plugin as a secret (e.g. `PLUGIN_AUTH_KEY`):
```bash
python -c 'import base64,hmac,hashlib,os; \
print(base64.urlsafe_b64encode(hmac.new(os.environ["LITELLM_SALT_KEY"].encode(), b"my-plugin", hashlib.sha256).digest()).decode())'
```
A compromised plugin holding only this scoped key cannot recover
`LITELLM_SALT_KEY` or decrypt any other litellm secret.
Decrypt and validate the claim with that key:
```python
import json, os, time
from cryptography.fernet import Fernet
_CLAIM_TTL_SECONDS = 30
def plugin_auth(session_claim: str) -> dict:
cipher = Fernet(os.environ["PLUGIN_AUTH_KEY"].encode())
claim = json.loads(cipher.decrypt(session_claim.encode(), ttl=_CLAIM_TTL_SECONDS))
if claim.get("plugin") != "my-plugin":
raise ValueError("claim audience mismatch")
if int(claim.get("exp", 0)) < int(time.time()):
raise ValueError("claim expired")
return claim
```
The claim is `{ "plugin", "user_id", "user_role", "exp" }`; it carries no
litellm bearer token. Establish the plugin's own session from `user_id` /
`user_role` and authenticate API calls back to litellm through the
`/plugin-proxy/my-plugin/*` reverse proxy, which injects `plugin_key` for you.
---
## How iframe auth works
```
litellm UI
├─ GET /api/plugins/auth-token -> { session_claim }
└─ postMessage({ type:"litellm-auth", session_claim }, pluginOrigin)
│
▼
Plugin iframe browser
└─ POST /api/plugin-auth { session_claim }
│
▼
Plugin server
├─ decrypt(session_claim, PLUGIN_AUTH_KEY) -> { user_id, user_role, exp }
└─ establish plugin session -> stored in sessionStorage
```
No litellm bearer token ever leaves the proxy; the claim only conveys the
caller's identity and expires after 30 seconds. A postMessage intercept
yields ciphertext that is useless without the plugin's scoped key.
---
## Proxy routes
- `GET /api/plugins` — list registered plugins (`name`, `display_name`, `url`). `plugin_key` is **never** returned; it stays server-side. Requires an authenticated caller.
- `GET /api/plugins/auth-token?plugin_name=<name>` — short-lived encrypted identity claim for the named plugin. Requires `LITELLM_SALT_KEY` to be set (503 otherwise) and the plugin to be registered (404 otherwise).
- `ANY /plugin-proxy/{name}/{path}` — authenticated reverse proxy to the plugin backend. Restricted to `proxy_admin`.
---
## Reverse proxy behaviour
When an admin (or server-to-server caller) hits `/plugin-proxy/<name>/<path>`, the proxy authenticates the caller locally, then rewrites the request before forwarding it to the plugin's `url`:
- **Every litellm credential header is stripped** — `Authorization`, `x-api-key`, `API-Key`, `x-goog-api-key`, `Ocp-Apim-Subscription-Key`, `x-litellm-api-key`, any configured `litellm_key_header_name`, plus `Cookie`. The plugin can never be handed the caller's live litellm key.
- **`plugin_key` is injected** as `Authorization: Bearer <plugin_key>` — the only credential the plugin receives.
- **Caller identity is forwarded** as `x-litellm-user-id` and `x-litellm-user-role` so the plugin can run its own authorization. These are informational, not credentials.
- **Responses are sandboxed** — `Content-Security-Policy: sandbox` and `X-Content-Type-Options: nosniff` are set so plugin-controlled bytes served from the litellm origin cannot execute against the dashboard.
---
## Security checklist
- [ ] `LITELLM_SALT_KEY` is set on the proxy and never shared with the plugin
- [ ] The plugin holds only its derived `HMAC(LITELLM_SALT_KEY, plugin_name)` key, provisioned as a dedicated secret
- [ ] `plugin_key` is a dedicated credential scoped to the plugin (not your litellm master key)
- [ ] Plugin's `POST /api/plugin-auth` enforces the claim's `plugin` audience and `exp` (30s TTL)
- [ ] Plugin treats `x-litellm-user-id` / `x-litellm-user-role` as identity hints, not as proof of authentication
- [ ] Plugin service URL uses HTTPS in production

View file

@ -0,0 +1,15 @@
"""
Code Interpreter Interception Module
Converts the native OpenAI Responses ``code_interpreter`` tool into a function
tool, runs the model-emitted code in a sandbox, and feeds the result back into
the agentic loop.
"""
from litellm.integrations.code_interpreter_interception.handler import (
CodeInterpreterInterceptionLogger,
)
__all__ = [
"CodeInterpreterInterceptionLogger",
]

View file

@ -0,0 +1,473 @@
"""
Code Interpreter Interception Handler
CustomLogger that swaps the native OpenAI Responses ``code_interpreter`` tool for
a function tool, executes the code the model emits inside a sandbox, and feeds the
captured stdout back through the typed agentic loop plan.
"""
import json
import time
import uuid
from typing import Any, cast
import litellm
from litellm._logging import verbose_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.integrations.code_interpreter_interception import (
CodeInterpreterInterceptionConfig,
)
from litellm.types.integrations.custom_logger import (
AgenticLoopPlan,
AgenticLoopRequestPatch,
)
from litellm.types.utils import CallTypes
LITELLM_CODE_EXECUTION_TOOL_NAME = "litellm_code_execution"
_INTERCEPTION_ACTIVE_KEY = "_code_interpreter_interception_active"
_SANDBOX_KEY = "_code_interpreter_interception_sandbox_key"
_CACHE_TTL_SECONDS = 15 * 60
def _resolve_sandbox_tool(sandbox_tool_name: str | None) -> dict[str, Any] | None:
try:
from litellm.sandbox.sandbox_tools import resolve_sandbox_tool
except ImportError:
return None
return resolve_sandbox_tool(sandbox_tool_name)
class CodeInterpreterInterceptionLogger(CustomLogger):
"""
CustomLogger that implements transparent code-interpreter execution loops.
Flow:
1. Replace the native ``code_interpreter`` tool with a function tool in the
pre-call hook so the model emits code as function-call arguments.
2. Detect ``litellm_code_execution`` function calls in the model response.
3. Run the emitted code in a sandbox (reused per request via a server-minted
sandbox key) and build a typed rerun plan that appends the
function_call_output.
"""
def __init__(
self,
enabled: bool = True,
enabled_providers: list[str] | None = None,
sandbox_tool_name: str | None = None,
sandbox_config: Any | None = None,
):
super().__init__()
self.enabled = enabled
self.enabled_providers = enabled_providers
self.sandbox_tool_name = sandbox_tool_name
self.sandbox_config = sandbox_config
self._container_cache: dict[str, tuple[Any, dict[str, Any] | None, float]] = {}
@classmethod
def from_config_yaml(
cls, config: CodeInterpreterInterceptionConfig
) -> "CodeInterpreterInterceptionLogger":
return cls(
enabled=bool(config.get("enabled", True)),
enabled_providers=config.get("enabled_providers"),
sandbox_tool_name=config.get("sandbox_tool_name"),
)
@staticmethod
def initialize_from_proxy_config(
litellm_settings: dict[str, Any],
callback_specific_params: dict[str, Any],
) -> "CodeInterpreterInterceptionLogger":
params: CodeInterpreterInterceptionConfig = {}
if "code_interpreter_interception_params" in litellm_settings:
params = litellm_settings["code_interpreter_interception_params"]
elif "code_interpreter_interception" in callback_specific_params and isinstance(
callback_specific_params["code_interpreter_interception"], dict
):
params = cast(
CodeInterpreterInterceptionConfig,
callback_specific_params["code_interpreter_interception"],
)
return CodeInterpreterInterceptionLogger.from_config_yaml(params)
async def async_pre_call_deployment_hook(
self, kwargs: dict[str, Any], call_type: CallTypes | None
) -> dict | None:
if not kwargs.get("_agentic_loop_depth"):
kwargs.pop(_INTERCEPTION_ACTIVE_KEY, None)
kwargs.pop(_SANDBOX_KEY, None)
if not self.enabled:
return None
if call_type not in (CallTypes.responses, CallTypes.aresponses):
return None
if (
self.enabled_providers is not None
and self._resolve_provider(kwargs) not in self.enabled_providers
):
return None
tools = kwargs.get("tools")
if not isinstance(tools, list):
return None
if not any(
isinstance(tool, dict) and tool.get("type") == "code_interpreter"
for tool in tools
):
return None
kwargs[_INTERCEPTION_ACTIVE_KEY] = True
kwargs[_SANDBOX_KEY] = uuid.uuid4().hex
if kwargs.get("stream"):
kwargs["stream"] = False
kwargs["_code_interpreter_interception_converted_stream"] = True
function_tool = {
"type": "function",
"name": LITELLM_CODE_EXECUTION_TOOL_NAME,
"description": "Execute python code in a sandbox and return stdout.",
"parameters": {
"type": "object",
"properties": {"code": {"type": "string"}},
"required": ["code"],
},
}
kwargs["tools"] = [
(
function_tool
if isinstance(tool, dict) and tool.get("type") == "code_interpreter"
else tool
)
for tool in tools
]
if self._tool_choice_targets_code_interpreter(kwargs.get("tool_choice")):
kwargs["tool_choice"] = {
"type": "function",
"name": LITELLM_CODE_EXECUTION_TOOL_NAME,
}
return kwargs
@staticmethod
def _tool_choice_targets_code_interpreter(tool_choice: Any) -> bool:
if not isinstance(tool_choice, dict):
return False
return (
tool_choice.get("type") == "code_interpreter"
or tool_choice.get("name") == "code_interpreter"
)
def _resolve_provider(self, kwargs: dict[str, Any]) -> str | None:
provider = kwargs.get("custom_llm_provider")
if provider:
return provider
model = kwargs.get("model")
if not isinstance(model, str):
return None
try:
return litellm.get_llm_provider(model=model)[1]
except Exception:
return None
async def async_should_run_agentic_loop(
self,
response: Any,
model: str,
messages: list[dict],
tools: list[dict] | None,
stream: bool,
custom_llm_provider: str,
kwargs: dict,
) -> tuple[bool, dict]:
if not self.enabled:
return False, {}
if not kwargs.get(_INTERCEPTION_ACTIVE_KEY):
return False, {}
if (
self.enabled_providers is not None
and custom_llm_provider not in self.enabled_providers
):
return False, {}
tool_calls = self._extract_code_execution_tool_calls(response=response)
if not tool_calls:
return False, {}
return True, {"tool_calls": tool_calls}
async def async_build_agentic_loop_plan(
self,
tools: dict,
model: str,
messages: list[dict],
response: Any,
anthropic_messages_provider_config: Any,
anthropic_messages_optional_request_params: dict,
logging_obj: Any,
stream: bool,
kwargs: dict,
) -> AgenticLoopPlan:
await self._prune_expired_cache()
tool_calls = cast(list[dict[str, Any]], tools.get("tool_calls", []))
sandbox_key = kwargs.get(_SANDBOX_KEY)
container, params = await self._get_or_create_container(cache_key=sandbox_key)
try:
container_id = getattr(container, "id", None)
input_list = self._normalize_messages(messages)
code_interpreter_calls = []
for tool_call in tool_calls:
arguments = tool_call.get("arguments", "")
code = self._parse_code(arguments)
stdout = await self._run_tool_call(
container=container, params=params, arguments=arguments
)
input_list.append(
{
"type": "function_call",
"call_id": tool_call.get("call_id"),
"name": LITELLM_CODE_EXECUTION_TOOL_NAME,
"arguments": arguments,
}
)
input_list.append(
{
"type": "function_call_output",
"call_id": tool_call.get("call_id"),
"output": stdout,
}
)
code_interpreter_calls.append(
{
"id": f"ci_{uuid.uuid4().hex}",
"type": "code_interpreter_call",
"status": "completed",
"code": code,
"container_id": container_id,
"outputs": (
[{"type": "logs", "logs": stdout}] if stdout else []
),
}
)
except Exception:
await self._delete_container_for_cache_key(sandbox_key)
raise
optional_params = anthropic_messages_optional_request_params
request_patch = AgenticLoopRequestPatch(
model=model,
messages=input_list,
tools=optional_params.get("tools"),
optional_params={k: v for k, v in optional_params.items() if k != "tools"},
kwargs={k: v for k, v in kwargs.items() if k != "litellm_logging_obj"},
)
return AgenticLoopPlan(
run_agentic_loop=True,
request_patch=request_patch,
metadata={
"tool_type": "code_interpreter",
"sandbox_key": sandbox_key or "",
"code_interpreter_calls": code_interpreter_calls,
},
)
async def async_agentic_loop_cleanup_hook(
self, plan: AgenticLoopPlan, kwargs: dict
) -> None:
metadata = plan.metadata or {} if plan else {}
await self._delete_container_for_cache_key(metadata.get("sandbox_key"))
async def async_post_agentic_loop_response_hook(
self, response: Any, plan: AgenticLoopPlan, kwargs: dict
) -> Any:
metadata = plan.metadata or {} if plan else {}
await self._delete_container_for_cache_key(metadata.get("sandbox_key"))
calls = metadata.get("code_interpreter_calls")
if not calls:
return response
is_dict = isinstance(response, dict)
output = (
response.get("output") if is_dict else getattr(response, "output", None)
)
if not isinstance(output, list):
return response
def _item_type(item: Any) -> Any:
return (
item.get("type")
if isinstance(item, dict)
else getattr(item, "type", None)
)
insert_at = next(
(i for i, item in enumerate(output) if _item_type(item) == "message"),
len(output),
)
new_output = output[:insert_at] + list(calls) + output[insert_at:]
if is_dict:
response["output"] = new_output
else:
response.output = new_output
return response
@staticmethod
def _parse_code(arguments: str) -> str:
try:
return json.loads(arguments).get("code", "") if arguments else ""
except (json.JSONDecodeError, TypeError, AttributeError):
return ""
async def _run_tool_call(
self, container: Any, params: dict[str, Any] | None, arguments: str
) -> str:
try:
code = json.loads(arguments).get("code", "") if arguments else ""
except (json.JSONDecodeError, TypeError):
return "[invalid tool arguments: could not parse code]"
result = await self._run_code(container=container, params=params, code=code)
if getattr(result, "error", None):
error = result.error
message = (
error.get("value") or error.get("name")
if isinstance(error, dict)
else str(error)
)
return f"[execution error] {message}"
return getattr(result, "stdout", "") or ""
async def _get_or_create_container(
self, cache_key: str | None
) -> tuple[Any, dict[str, Any] | None]:
if cache_key:
cached = self._container_cache.get(cache_key)
if cached is not None:
return cached[0], cached[1]
container, params = await self._create_container()
if cache_key:
self._container_cache[cache_key] = (container, params, time.time())
return container, params
async def _create_container(self) -> tuple[Any, dict[str, Any] | None]:
if self.sandbox_config is not None:
return await self.sandbox_config.acreate_sandbox(), None
params = _resolve_sandbox_tool(self.sandbox_tool_name)
if params is None:
raise ValueError(
"CodeInterpreterInterception: no sandbox available. Provide a "
"sandbox_config or configure a sandbox tool resolvable via "
"sandbox_tool_name."
)
container = await litellm.acreate_sandbox(
provider=params["sandbox_provider"],
api_key=params.get("api_key"),
api_base=params.get("api_base"),
)
return container, params
async def _run_code(
self, container: Any, params: dict[str, Any] | None, code: str
) -> Any:
if self.sandbox_config is not None:
return await self.sandbox_config.arun_code(container=container, code=code)
if params is None:
raise ValueError(
"CodeInterpreterInterception: no sandbox available to run code."
)
return await litellm.arun_code(
provider=params["sandbox_provider"],
container=container,
code=code,
api_key=params.get("api_key"),
)
async def _delete_container(
self, container: Any, params: dict[str, Any] | None
) -> None:
try:
if self.sandbox_config is not None:
await self.sandbox_config.adelete_sandbox(container=container)
return
if params is None:
return
await litellm.adelete_sandbox(
provider=params["sandbox_provider"],
container=container,
api_key=params.get("api_key"),
api_base=params.get("api_base"),
)
except Exception:
verbose_logger.exception(
"CodeInterpreterInterception: failed to delete sandbox container"
)
async def _delete_container_for_cache_key(self, cache_key: str | None) -> None:
if not cache_key:
return
cached = self._container_cache.pop(cache_key, None)
if cached is None:
return
await self._delete_container(container=cached[0], params=cached[1])
def _normalize_messages(self, messages: Any) -> list[dict[str, Any]]:
if isinstance(messages, str):
return [{"role": "user", "content": messages}]
if isinstance(messages, list):
return list(messages)
return []
def _extract_code_execution_tool_calls(self, response: Any) -> list[dict[str, Any]]:
if isinstance(response, dict):
output = response.get("output", [])
else:
output = getattr(response, "output", []) or []
if not isinstance(output, list):
return []
return [
{
"call_id": (
item.get("call_id")
if isinstance(item, dict)
else getattr(item, "call_id", None)
),
"name": LITELLM_CODE_EXECUTION_TOOL_NAME,
"arguments": (
item.get("arguments")
if isinstance(item, dict)
else getattr(item, "arguments", "")
),
}
for item in output
if self._is_code_execution_call(item)
]
def _is_code_execution_call(self, item: Any) -> bool:
if isinstance(item, dict):
return (
item.get("type") == "function_call"
and item.get("name") == LITELLM_CODE_EXECUTION_TOOL_NAME
)
return (
getattr(item, "type", None) == "function_call"
and getattr(item, "name", None) == LITELLM_CODE_EXECUTION_TOOL_NAME
)
async def _prune_expired_cache(self) -> None:
now = time.time()
expired = [
(cache_key, container, params)
for cache_key, (
container,
params,
created_at,
) in self._container_cache.items()
if now - created_at > _CACHE_TTL_SECONDS
]
for cache_key, container, params in expired:
self._container_cache.pop(cache_key, None)
await self._delete_container(container=container, params=params)

View file

@ -718,6 +718,24 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
"""
return response
async def async_agentic_loop_cleanup_hook(
self,
plan: AgenticLoopPlan,
kwargs: dict,
) -> None:
"""
Release resources held for an agentic-loop iteration.
Runs in a ``finally`` around the follow-up provider call, so it fires
whether the rerun returns normally, hits a loop safety abort, or raises
an upstream error. Implementations must be idempotent because the
post-response hook may already have released the same resource on the
success path. Use ``plan.metadata`` to locate what to clean up.
Default does nothing.
"""
return None
async def async_should_run_chat_completion_agentic_loop(
self,
response: Any,

View file

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

File diff suppressed because it is too large Load diff

View file

@ -6,6 +6,7 @@ import logging
import threading
import time
import traceback
from dataclasses import dataclass
from typing import (
Any,
AsyncIterator,
@ -97,6 +98,19 @@ def print_verbose(print_statement):
pass
@dataclass(frozen=True, slots=True)
class _ProviderChunkParsed:
response_obj: dict[str, Any]
@dataclass(frozen=True, slots=True)
class _ProviderChunkEarlyReturn:
value: Any
_ProviderChunkResult = Union[_ProviderChunkParsed, _ProviderChunkEarlyReturn]
class CustomStreamWrapper:
def __init__(
self,
@ -1145,381 +1159,392 @@ class CustomStreamWrapper:
del model_response.choices[0].delta.reasoning_content
return
def _dispatch_provider_chunk(
self,
chunk: Any,
model_response: ModelResponseStream,
completion_obj: dict[str, Any],
) -> _ProviderChunkResult:
response_obj: dict[str, Any] = {}
if (
isinstance(chunk, ModelResponseStream)
and self.custom_llm_provider is not None
and self.custom_llm_provider in litellm._custom_providers
):
_has_content = bool(
chunk.choices
and chunk.choices[0].delta is not None
and (
chunk.choices[0].delta.content or chunk.choices[0].delta.tool_calls
)
)
if self.received_finish_reason is not None:
if not _has_content:
raise StopIteration
if chunk.choices and chunk.choices[0].finish_reason:
self.received_finish_reason = chunk.choices[0].finish_reason
if not _has_content:
return _ProviderChunkEarlyReturn(None)
# Strip finish_reason from the content chunk so it appears
# only on the trailing empty-delta chunk (OpenAI spec).
# finish_reason_handler() will emit the proper terminal chunk.
chunk.choices[0].finish_reason = None # type: ignore[assignment]
return _ProviderChunkEarlyReturn(chunk)
if (
isinstance(chunk, dict)
and generic_chunk_has_all_required_fields(
chunk=chunk
) # check if chunk is a generic streaming chunk
) or (
self.custom_llm_provider
and self.custom_llm_provider in litellm._custom_providers
):
if self.received_finish_reason is not None:
_chunk_has_content = isinstance(chunk, dict) and (
bool(chunk.get("text", ""))
or chunk.get("tool_use") is not None
# Usage-only final chunks are valid and needed to surface
# finish_reason/usage to downstream translators.
or chunk.get("usage") is not None
)
if not _chunk_has_content and (
not isinstance(chunk, dict)
or "provider_specific_fields" not in chunk
):
raise StopIteration
anthropic_response_obj: GChunk = cast(GChunk, chunk)
completion_obj["content"] = anthropic_response_obj["text"]
if anthropic_response_obj["is_finished"]:
self.received_finish_reason = anthropic_response_obj["finish_reason"]
if anthropic_response_obj["finish_reason"]:
self.intermittent_finish_reason = anthropic_response_obj[
"finish_reason"
]
if anthropic_response_obj["usage"] is not None:
setattr(
model_response,
"usage",
litellm.Usage(**anthropic_response_obj["usage"]),
)
if (
"tool_use" in anthropic_response_obj
and anthropic_response_obj["tool_use"] is not None
):
completion_obj["tool_calls"] = [anthropic_response_obj["tool_use"]]
if (
"provider_specific_fields" in anthropic_response_obj
and anthropic_response_obj["provider_specific_fields"] is not None
):
for key, value in anthropic_response_obj[
"provider_specific_fields"
].items():
setattr(model_response, key, value)
response_obj = cast(dict[str, Any], anthropic_response_obj)
elif self.model == "replicate" or self.custom_llm_provider == "replicate":
response_obj = self.handle_replicate_chunk(chunk)
completion_obj["content"] = response_obj["text"]
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
elif self.custom_llm_provider and self.custom_llm_provider == "predibase":
response_obj = self.handle_predibase_chunk(chunk)
completion_obj["content"] = response_obj["text"]
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
elif (
self.custom_llm_provider and self.custom_llm_provider == "baseten"
): # baseten doesn't provide streaming
completion_obj["content"] = self.handle_baseten_chunk(chunk)
elif (
self.custom_llm_provider and self.custom_llm_provider == "ai21"
): # ai21 doesn't provide streaming
response_obj = self.handle_ai21_chunk(chunk)
completion_obj["content"] = response_obj["text"]
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
elif self.custom_llm_provider and self.custom_llm_provider == "maritalk":
response_obj = self.handle_maritalk_chunk(chunk)
completion_obj["content"] = response_obj["text"]
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
elif self.custom_llm_provider and self.custom_llm_provider == "vllm":
completion_obj["content"] = chunk[0].outputs[0].text
elif (
self.custom_llm_provider and self.custom_llm_provider == "aleph_alpha"
): # aleph alpha doesn't provide streaming
response_obj = self.handle_aleph_alpha_chunk(chunk)
completion_obj["content"] = response_obj["text"]
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
elif self.custom_llm_provider == "nlp_cloud":
try:
response_obj = self.handle_nlp_cloud_chunk(chunk)
completion_obj["content"] = response_obj["text"]
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
except Exception as e:
if self.received_finish_reason:
raise e
else:
if self.sent_first_chunk is False:
raise Exception("An unknown error occurred with the stream")
self.received_finish_reason = "stop"
elif self.custom_llm_provider == "vertex_ai" and not isinstance(
chunk, ModelResponseStream
):
chunk = cast(Any, chunk)
import proto # type: ignore
if hasattr(chunk, "candidates") is True:
try:
try:
completion_obj["content"] = chunk.text # type: ignore
except Exception as e:
original_exception = e
if "Part has no text." in str(e):
## check for function calling
function_call = (
chunk.candidates[0].content.parts[0].function_call # type: ignore
)
args_dict = {}
# Check if it's a RepeatedComposite instance
for key, val in function_call.args.items():
if isinstance(
val,
proto.marshal.collections.repeated.RepeatedComposite, # type: ignore
):
# If so, convert to list
args_dict[key] = [v for v in val]
else:
args_dict[key] = val
try:
args_str = json.dumps(args_dict)
except Exception as e:
raise e
_delta_obj = litellm.utils.Delta(
content=None,
tool_calls=[
{
"id": f"call_{str(uuid.uuid4())}",
"function": {
"arguments": args_str,
"name": function_call.name,
},
"type": "function",
}
],
)
_streaming_response = StreamingChoices(delta=_delta_obj)
_model_response = ModelResponseStream()
_model_response.choices = [_streaming_response]
response_obj = {"original_chunk": _model_response}
else:
raise original_exception
if (
hasattr(chunk.candidates[0], "finish_reason") # type: ignore
and chunk.candidates[0].finish_reason.name # type: ignore
!= "FINISH_REASON_UNSPECIFIED"
): # every non-final chunk in vertex ai has this
self.received_finish_reason = map_finish_reason( # type: ignore
chunk.candidates[0].finish_reason.name
)
except Exception:
if chunk.candidates[0].finish_reason.name == "SAFETY": # type: ignore
raise Exception(
f"The response was blocked by VertexAI. {str(chunk)}"
)
else:
completion_obj["content"] = str(chunk)
elif self.custom_llm_provider == "petals":
if self.completion_stream is None or len(self.completion_stream) == 0:
if self.received_finish_reason is not None:
raise StopIteration
else:
self.received_finish_reason = "stop"
chunk_size = 30
stream = cast(Any, self.completion_stream)
new_chunk = stream[:chunk_size]
completion_obj["content"] = new_chunk
self.completion_stream = stream[chunk_size:]
elif self.custom_llm_provider == "palm":
# fake streaming
response_obj = {}
if self.completion_stream is None or len(self.completion_stream) == 0:
if self.received_finish_reason is not None:
raise StopIteration
else:
self.received_finish_reason = "stop"
chunk_size = 30
stream = cast(Any, self.completion_stream)
new_chunk = stream[:chunk_size]
completion_obj["content"] = new_chunk
self.completion_stream = stream[chunk_size:]
elif self.custom_llm_provider == "triton":
response_obj = self.handle_triton_stream(chunk)
completion_obj["content"] = response_obj["text"]
print_verbose(f"completion obj content: {completion_obj['content']}")
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
elif self.custom_llm_provider == "text-completion-openai":
response_obj = self.handle_openai_text_completion_chunk(chunk)
completion_obj["content"] = response_obj["text"]
print_verbose(f"completion obj content: {completion_obj['content']}")
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
if response_obj["usage"] is not None:
setattr(
model_response,
"usage",
litellm.Usage(
prompt_tokens=response_obj["usage"].prompt_tokens,
completion_tokens=response_obj["usage"].completion_tokens,
total_tokens=response_obj["usage"].total_tokens,
),
)
elif self.custom_llm_provider == "text-completion-codestral":
if not isinstance(chunk, str):
raise ValueError(f"chunk is not a string: {chunk}")
response_obj = cast(
dict[str, Any],
litellm.CodestralTextCompletionConfig()._chunk_parser(chunk),
)
completion_obj["content"] = response_obj["text"]
print_verbose(f"completion obj content: {completion_obj['content']}")
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
if "usage" in response_obj is not None:
setattr(
model_response,
"usage",
litellm.Usage(
prompt_tokens=response_obj["usage"].prompt_tokens,
completion_tokens=response_obj["usage"].completion_tokens,
total_tokens=response_obj["usage"].total_tokens,
),
)
elif self.custom_llm_provider == "azure_text":
response_obj = self.handle_azure_text_completion_chunk(chunk)
completion_obj["content"] = response_obj["text"]
print_verbose(f"completion obj content: {completion_obj['content']}")
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
elif self.custom_llm_provider == "cached_response":
chunk = cast(ModelResponseStream, chunk)
response_obj = {
"text": chunk.choices[0].delta.content,
"is_finished": True,
"finish_reason": chunk.choices[0].finish_reason,
"original_chunk": chunk,
"tool_calls": (
chunk.choices[0].delta.tool_calls
if hasattr(chunk.choices[0].delta, "tool_calls")
else None
),
}
completion_obj["content"] = response_obj["text"]
if response_obj["tool_calls"] is not None:
completion_obj["tool_calls"] = response_obj["tool_calls"]
print_verbose(f"completion obj content: {completion_obj['content']}")
if hasattr(chunk, "id"):
model_response.id = chunk.id
self.response_id = chunk.id
if hasattr(chunk, "system_fingerprint"):
self.system_fingerprint = chunk.system_fingerprint
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
else: # openai / azure chat model
if self.custom_llm_provider in [
LlmProviders.AZURE.value,
LlmProviders.AZURE_AI.value,
]:
if isinstance(chunk, BaseModel) and hasattr(chunk, "model"):
# for azure, we need to pass the model from the original chunk
self.model = getattr(chunk, "model", self.model)
response_obj = self.handle_openai_chat_completion_chunk(chunk)
if response_obj is None:
return _ProviderChunkEarlyReturn(None)
completion_obj["content"] = response_obj["text"]
self.intermittent_finish_reason = response_obj.get("finish_reason", None)
if response_obj["is_finished"]:
if response_obj["finish_reason"] == "error":
raise Exception(
"{} raised a streaming error - finish_reason: error, no content string given. Received Chunk={}".format(
self.custom_llm_provider, response_obj
)
)
self.received_finish_reason = response_obj["finish_reason"]
if response_obj.get("original_chunk", None) is not None:
if hasattr(response_obj["original_chunk"], "id"):
model_response = self.set_model_id(
response_obj["original_chunk"].id, model_response
)
if hasattr(response_obj["original_chunk"], "system_fingerprint"):
model_response.system_fingerprint = response_obj[
"original_chunk"
].system_fingerprint
self.system_fingerprint = response_obj[
"original_chunk"
].system_fingerprint
if response_obj["logprobs"] is not None:
model_response.choices[0].logprobs = response_obj["logprobs"]
if response_obj["usage"] is not None:
if isinstance(response_obj["usage"], dict):
setattr(
model_response,
"usage",
litellm.Usage(
prompt_tokens=response_obj["usage"].get(
"prompt_tokens", None
)
or None,
completion_tokens=response_obj["usage"].get(
"completion_tokens", None
)
or None,
total_tokens=response_obj["usage"].get("total_tokens", None)
or None,
),
)
elif isinstance(response_obj["usage"], Usage):
setattr(
model_response,
"usage",
response_obj["usage"],
)
elif isinstance(response_obj["usage"], BaseModel):
setattr(
model_response,
"usage",
litellm.Usage(**response_obj["usage"].model_dump()),
)
return _ProviderChunkParsed(response_obj)
def chunk_creator(self, chunk: Any): # type: ignore
if hasattr(chunk, "id"):
self.response_id = chunk.id
model_response = self.model_response_creator()
response_obj: Dict[str, Any] = {}
response_obj: dict[str, Any] = {}
try:
# return this for all models
completion_obj: Dict[str, Any] = {"content": ""}
from litellm.types.utils import GenericStreamingChunk as GChunk
if (
isinstance(chunk, ModelResponseStream)
and self.custom_llm_provider is not None
and self.custom_llm_provider in litellm._custom_providers
):
_has_content = bool(
chunk.choices
and chunk.choices[0].delta is not None
and (
chunk.choices[0].delta.content
or chunk.choices[0].delta.tool_calls
)
)
if self.received_finish_reason is not None:
if not _has_content:
raise StopIteration
if chunk.choices and chunk.choices[0].finish_reason:
self.received_finish_reason = chunk.choices[0].finish_reason
if not _has_content:
return None
# Strip finish_reason from the content chunk so it appears
# only on the trailing empty-delta chunk (OpenAI spec).
# finish_reason_handler() will emit the proper terminal chunk.
chunk.choices[0].finish_reason = None # type: ignore[assignment]
return chunk
if (
isinstance(chunk, dict)
and generic_chunk_has_all_required_fields(
chunk=chunk
) # check if chunk is a generic streaming chunk
) or (
self.custom_llm_provider
and self.custom_llm_provider in litellm._custom_providers
):
if self.received_finish_reason is not None:
_chunk_has_content = isinstance(chunk, dict) and (
bool(chunk.get("text", ""))
or chunk.get("tool_use") is not None
# Usage-only final chunks are valid and needed to surface
# finish_reason/usage to downstream translators.
or chunk.get("usage") is not None
)
if not _chunk_has_content and (
not isinstance(chunk, dict)
or "provider_specific_fields" not in chunk
):
raise StopIteration
anthropic_response_obj: GChunk = cast(GChunk, chunk)
completion_obj["content"] = anthropic_response_obj["text"]
if anthropic_response_obj["is_finished"]:
self.received_finish_reason = anthropic_response_obj[
"finish_reason"
]
if anthropic_response_obj["finish_reason"]:
self.intermittent_finish_reason = anthropic_response_obj[
"finish_reason"
]
if anthropic_response_obj["usage"] is not None:
setattr(
model_response,
"usage",
litellm.Usage(**anthropic_response_obj["usage"]),
)
if (
"tool_use" in anthropic_response_obj
and anthropic_response_obj["tool_use"] is not None
):
completion_obj["tool_calls"] = [anthropic_response_obj["tool_use"]]
if (
"provider_specific_fields" in anthropic_response_obj
and anthropic_response_obj["provider_specific_fields"] is not None
):
for key, value in anthropic_response_obj[
"provider_specific_fields"
].items():
setattr(model_response, key, value)
response_obj = cast(Dict[str, Any], anthropic_response_obj)
elif self.model == "replicate" or self.custom_llm_provider == "replicate":
response_obj = self.handle_replicate_chunk(chunk)
completion_obj["content"] = response_obj["text"]
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
elif self.custom_llm_provider and self.custom_llm_provider == "predibase":
response_obj = self.handle_predibase_chunk(chunk)
completion_obj["content"] = response_obj["text"]
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
elif (
self.custom_llm_provider and self.custom_llm_provider == "baseten"
): # baseten doesn't provide streaming
completion_obj["content"] = self.handle_baseten_chunk(chunk)
elif (
self.custom_llm_provider and self.custom_llm_provider == "ai21"
): # ai21 doesn't provide streaming
response_obj = self.handle_ai21_chunk(chunk)
completion_obj["content"] = response_obj["text"]
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
elif self.custom_llm_provider and self.custom_llm_provider == "maritalk":
response_obj = self.handle_maritalk_chunk(chunk)
completion_obj["content"] = response_obj["text"]
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
elif self.custom_llm_provider and self.custom_llm_provider == "vllm":
completion_obj["content"] = chunk[0].outputs[0].text
elif (
self.custom_llm_provider and self.custom_llm_provider == "aleph_alpha"
): # aleph alpha doesn't provide streaming
response_obj = self.handle_aleph_alpha_chunk(chunk)
completion_obj["content"] = response_obj["text"]
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
elif self.custom_llm_provider == "nlp_cloud":
try:
response_obj = self.handle_nlp_cloud_chunk(chunk)
completion_obj["content"] = response_obj["text"]
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
except Exception as e:
if self.received_finish_reason:
raise e
else:
if self.sent_first_chunk is False:
raise Exception("An unknown error occurred with the stream")
self.received_finish_reason = "stop"
elif self.custom_llm_provider == "vertex_ai" and not isinstance(
chunk, ModelResponseStream
):
import proto # type: ignore
if hasattr(chunk, "candidates") is True:
try:
try:
completion_obj["content"] = chunk.text # type: ignore
except Exception as e:
original_exception = e
if "Part has no text." in str(e):
## check for function calling
function_call = (
chunk.candidates[0].content.parts[0].function_call # type: ignore
)
args_dict = {}
# Check if it's a RepeatedComposite instance
for key, val in function_call.args.items():
if isinstance(
val,
proto.marshal.collections.repeated.RepeatedComposite, # type: ignore
):
# If so, convert to list
args_dict[key] = [v for v in val]
else:
args_dict[key] = val
try:
args_str = json.dumps(args_dict)
except Exception as e:
raise e
_delta_obj = litellm.utils.Delta(
content=None,
tool_calls=[
{
"id": f"call_{str(uuid.uuid4())}",
"function": {
"arguments": args_str,
"name": function_call.name,
},
"type": "function",
}
],
)
_streaming_response = StreamingChoices(delta=_delta_obj)
_model_response = ModelResponseStream()
_model_response.choices = [_streaming_response]
response_obj = {"original_chunk": _model_response}
else:
raise original_exception
if (
hasattr(chunk.candidates[0], "finish_reason") # type: ignore
and chunk.candidates[0].finish_reason.name # type: ignore
!= "FINISH_REASON_UNSPECIFIED"
): # every non-final chunk in vertex ai has this
self.received_finish_reason = map_finish_reason( # type: ignore
chunk.candidates[0].finish_reason.name
)
except Exception:
if chunk.candidates[0].finish_reason.name == "SAFETY": # type: ignore
raise Exception(
f"The response was blocked by VertexAI. {str(chunk)}"
)
else:
completion_obj["content"] = str(chunk)
elif self.custom_llm_provider == "petals":
if self.completion_stream is None or len(self.completion_stream) == 0:
if self.received_finish_reason is not None:
raise StopIteration
else:
self.received_finish_reason = "stop"
chunk_size = 30
new_chunk = self.completion_stream[:chunk_size] # type: ignore[index]
completion_obj["content"] = new_chunk
self.completion_stream = self.completion_stream[chunk_size:] # type: ignore[index]
elif self.custom_llm_provider == "palm":
# fake streaming
response_obj = {}
if self.completion_stream is None or len(self.completion_stream) == 0:
if self.received_finish_reason is not None:
raise StopIteration
else:
self.received_finish_reason = "stop"
chunk_size = 30
new_chunk = self.completion_stream[:chunk_size] # type: ignore[index]
completion_obj["content"] = new_chunk
self.completion_stream = self.completion_stream[chunk_size:] # type: ignore[index]
elif self.custom_llm_provider == "triton":
response_obj = self.handle_triton_stream(chunk)
completion_obj["content"] = response_obj["text"]
print_verbose(f"completion obj content: {completion_obj['content']}")
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
elif self.custom_llm_provider == "text-completion-openai":
response_obj = self.handle_openai_text_completion_chunk(chunk)
completion_obj["content"] = response_obj["text"]
print_verbose(f"completion obj content: {completion_obj['content']}")
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
if response_obj["usage"] is not None:
setattr(
model_response,
"usage",
litellm.Usage(
prompt_tokens=response_obj["usage"].prompt_tokens,
completion_tokens=response_obj["usage"].completion_tokens,
total_tokens=response_obj["usage"].total_tokens,
),
)
elif self.custom_llm_provider == "text-completion-codestral":
if not isinstance(chunk, str):
raise ValueError(f"chunk is not a string: {chunk}")
response_obj = cast(
Dict[str, Any],
litellm.CodestralTextCompletionConfig()._chunk_parser(chunk),
)
completion_obj["content"] = response_obj["text"]
print_verbose(f"completion obj content: {completion_obj['content']}")
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
if "usage" in response_obj is not None:
setattr(
model_response,
"usage",
litellm.Usage(
prompt_tokens=response_obj["usage"].prompt_tokens,
completion_tokens=response_obj["usage"].completion_tokens,
total_tokens=response_obj["usage"].total_tokens,
),
)
elif self.custom_llm_provider == "azure_text":
response_obj = self.handle_azure_text_completion_chunk(chunk)
completion_obj["content"] = response_obj["text"]
print_verbose(f"completion obj content: {completion_obj['content']}")
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
elif self.custom_llm_provider == "cached_response":
chunk = cast(ModelResponseStream, chunk)
response_obj = {
"text": chunk.choices[0].delta.content,
"is_finished": True,
"finish_reason": chunk.choices[0].finish_reason,
"original_chunk": chunk,
"tool_calls": (
chunk.choices[0].delta.tool_calls
if hasattr(chunk.choices[0].delta, "tool_calls")
else None
),
}
completion_obj["content"] = response_obj["text"]
if response_obj["tool_calls"] is not None:
completion_obj["tool_calls"] = response_obj["tool_calls"]
print_verbose(f"completion obj content: {completion_obj['content']}")
if hasattr(chunk, "id"):
model_response.id = chunk.id
self.response_id = chunk.id
if hasattr(chunk, "system_fingerprint"):
self.system_fingerprint = chunk.system_fingerprint
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
else: # openai / azure chat model
if self.custom_llm_provider in [
LlmProviders.AZURE.value,
LlmProviders.AZURE_AI.value,
]:
if isinstance(chunk, BaseModel) and hasattr(chunk, "model"):
# for azure, we need to pass the model from the original chunk
self.model = getattr(chunk, "model", self.model)
response_obj = self.handle_openai_chat_completion_chunk(chunk)
if response_obj is None:
return
completion_obj["content"] = response_obj["text"]
self.intermittent_finish_reason = response_obj.get(
"finish_reason", None
)
if response_obj["is_finished"]:
if response_obj["finish_reason"] == "error":
raise Exception(
"{} raised a streaming error - finish_reason: error, no content string given. Received Chunk={}".format(
self.custom_llm_provider, response_obj
)
)
self.received_finish_reason = response_obj["finish_reason"]
if response_obj.get("original_chunk", None) is not None:
if hasattr(response_obj["original_chunk"], "id"):
model_response = self.set_model_id(
response_obj["original_chunk"].id, model_response
)
if hasattr(response_obj["original_chunk"], "system_fingerprint"):
model_response.system_fingerprint = response_obj[
"original_chunk"
].system_fingerprint
self.system_fingerprint = response_obj[
"original_chunk"
].system_fingerprint
if response_obj["logprobs"] is not None:
model_response.choices[0].logprobs = response_obj["logprobs"]
if response_obj["usage"] is not None:
if isinstance(response_obj["usage"], dict):
setattr(
model_response,
"usage",
litellm.Usage(
prompt_tokens=response_obj["usage"].get(
"prompt_tokens", None
)
or None,
completion_tokens=response_obj["usage"].get(
"completion_tokens", None
)
or None,
total_tokens=response_obj["usage"].get(
"total_tokens", None
)
or None,
),
)
elif isinstance(response_obj["usage"], Usage):
setattr(
model_response,
"usage",
response_obj["usage"],
)
elif isinstance(response_obj["usage"], BaseModel):
setattr(
model_response,
"usage",
litellm.Usage(**response_obj["usage"].model_dump()),
)
completion_obj: dict[str, Any] = {"content": ""}
dispatch_result = self._dispatch_provider_chunk(
chunk=chunk,
model_response=model_response,
completion_obj=completion_obj,
)
if isinstance(dispatch_result, _ProviderChunkEarlyReturn):
return dispatch_result.value
response_obj = dispatch_result.response_obj
model_response.model = self.model
## FUNCTION CALL PARSING
@ -1980,11 +2005,29 @@ class CustomStreamWrapper:
except StopIteration:
if self.sent_last_chunk is True:
complete_streaming_response = litellm.stream_chunk_builder(
chunks=self.chunks,
messages=self.messages,
logging_obj=self.logging_obj,
)
try:
complete_streaming_response = litellm.stream_chunk_builder(
chunks=self.chunks,
messages=self.messages,
logging_obj=self.logging_obj,
)
except Exception as e:
# stream_chunk_builder can re-raise (as APIError) on large agentic
# streams. The raise originates inside this except-StopIteration block,
# so the sibling `except Exception` below does not catch it; it would
# escape __next__ and drop the request from SpendLogs. Recover
# best-effort usage from the raw chunks so cost is still tracked
verbose_logger.warning(
"stream_chunk_builder raised at end-of-stream (%s); logging "
"best-effort usage from chunks.",
str(e),
)
try:
complete_streaming_response = self.model_response_creator(
chunk={"usage": calculate_total_usage(chunks=self.chunks)}
)
except Exception:
complete_streaming_response = None
response = self.model_response_creator()
if complete_streaming_response is not None:
@ -2209,11 +2252,27 @@ class CustomStreamWrapper:
except (StopAsyncIteration, StopIteration):
if self.sent_last_chunk is True:
# log the final chunk with accurate streaming values
complete_streaming_response = litellm.stream_chunk_builder(
chunks=self.chunks,
messages=self.messages,
logging_obj=self.logging_obj,
)
try:
complete_streaming_response = litellm.stream_chunk_builder(
chunks=self.chunks,
messages=self.messages,
logging_obj=self.logging_obj,
)
except Exception as e:
# see sync __next__: a raise from stream_chunk_builder inside this
# except handler escapes __anext__ and drops the request from SpendLogs.
# Recover best-effort usage from the raw chunks so cost is still tracked
verbose_logger.warning(
"stream_chunk_builder raised at end-of-stream (%s); logging "
"best-effort usage from chunks.",
str(e),
)
try:
complete_streaming_response = self.model_response_creator(
chunk={"usage": calculate_total_usage(chunks=self.chunks)}
)
except Exception:
complete_streaming_response = None
response = self.model_response_creator()
if complete_streaming_response is not None:

View file

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

View file

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

View file

@ -2407,11 +2407,17 @@ class BaseLLMHTTPHandler:
provider_config=responses_api_provider_config,
)
return responses_api_provider_config.transform_response_api_response(
model=model,
raw_response=response,
logging_obj=logging_obj,
initial_response = (
responses_api_provider_config.transform_response_api_response(
model=model,
raw_response=response,
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.
return initial_response
async def async_response_api_handler(
self,
@ -2570,12 +2576,44 @@ class BaseLLMHTTPHandler:
provider_config=responses_api_provider_config,
)
return responses_api_provider_config.transform_response_api_response(
model=model,
raw_response=response,
logging_obj=logging_obj,
initial_response = (
responses_api_provider_config.transform_response_api_response(
model=model,
raw_response=response,
logging_obj=logging_obj,
)
)
final_response = await 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",
)
result = final_response if final_response is not None else initial_response
if litellm_params.get(
"_code_interpreter_interception_converted_stream"
) and not litellm_params.get("_agentic_loop_depth"):
return self._wrap_responses_response_as_fake_stream(
result=result,
model=model,
responses_api_provider_config=responses_api_provider_config,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
)
return result
async def async_delete_response_api_handler(
self,
response_id: str,
@ -4875,6 +4913,132 @@ class BaseLLMHTTPHandler:
return response
async def _execute_responses_agentic_plan(
self,
plan: AgenticLoopPlan,
model: str,
response_api_optional_request_params: dict,
logging_obj: "LiteLLMLoggingObj",
kwargs: dict,
depth: int,
max_loops: int,
fingerprints: list[str],
fingerprint: str,
callback: Any | None = None,
) -> Any:
patch = plan.request_patch or AgenticLoopRequestPatch()
if patch.messages is None:
raise ValueError("Agentic loop plan missing patched responses input")
optional_params = dict(response_api_optional_request_params)
optional_params.update(patch.optional_params)
if patch.tools is not None:
optional_params["tools"] = patch.tools
optional_params = {
k: v
for k, v in optional_params.items()
if k != "stream" and k != "_code_interpreter_interception_converted_stream"
}
internal_keys = {"litellm_logging_obj"}
kwargs_for_followup = {
k: v
for k, v in kwargs.items()
if not k.startswith("_websearch_interception")
and not k.startswith("_compression_interception")
and k != "_code_interpreter_interception_converted_stream"
and k not in internal_keys
and k not in optional_params
}
kwargs_for_followup.update(patch.kwargs)
kwargs_for_followup["_agentic_loop_depth"] = depth + 1
kwargs_for_followup["max_agentic_loops"] = max_loops
kwargs_for_followup["_agentic_loop_fingerprints"] = fingerprints + [fingerprint]
try:
response = await litellm.aresponses(
model=patch.model or model,
input=patch.messages,
**optional_params,
**kwargs_for_followup,
)
if callback is not None:
try:
response = await callback.async_post_agentic_loop_response_hook(
response=response, plan=plan, kwargs=kwargs
)
except Exception as e:
_call_id = getattr(logging_obj, "litellm_call_id", "unknown")
verbose_logger.exception(
"LiteLLM.AgenticHookError: Exception in "
"async_post_agentic_loop_response_hook [call_id=%s model=%s]: %s",
_call_id,
model,
str(e),
)
return response
finally:
if callback is not None:
await self._run_agentic_loop_cleanup(
callback=callback,
plan=plan,
kwargs=kwargs,
logging_obj=logging_obj,
model=model,
)
@staticmethod
async def _run_agentic_loop_cleanup(
callback: Any,
plan: AgenticLoopPlan,
kwargs: dict,
logging_obj: "LiteLLMLoggingObj",
model: str,
) -> None:
try:
await callback.async_agentic_loop_cleanup_hook(plan=plan, kwargs=kwargs)
except Exception as e:
_call_id = getattr(logging_obj, "litellm_call_id", "unknown")
verbose_logger.exception(
"LiteLLM.AgenticHookError: Exception in "
"async_agentic_loop_cleanup_hook [call_id=%s model=%s]: %s",
_call_id,
model,
str(e),
)
def _wrap_responses_response_as_fake_stream(
self,
result: Any,
model: str,
responses_api_provider_config: Any,
logging_obj: "LiteLLMLoggingObj",
custom_llm_provider: str,
) -> Any:
"""
Wrap a completed responses result as a synthetic stream.
Used when an interceptor forced stream=False to run the agentic loop on
the non-streaming path, but the caller originally asked for streaming.
"""
import httpx
from litellm.responses.streaming_iterator import (
MockResponsesAPIStreamingIterator,
)
payload = result.model_dump() if hasattr(result, "model_dump") else result
raw_response = httpx.Response(status_code=200, json=payload)
return MockResponsesAPIStreamingIterator(
response=raw_response,
model=model,
responses_api_provider_config=responses_api_provider_config,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
)
async def _execute_chat_completion_agentic_plan(
self,
plan: AgenticLoopPlan,
@ -4940,6 +5104,7 @@ class BaseLLMHTTPHandler:
stream: bool,
custom_llm_provider: str,
kwargs: Dict,
api_surface: str = "anthropic_messages",
) -> Optional[Any]:
"""
Call agentic completion hooks for all custom loggers (Anthropic Messages API).
@ -5046,6 +5211,20 @@ class BaseLLMHTTPHandler:
if not plan.run_agentic_loop:
continue
if api_surface == "responses":
return await self._execute_responses_agentic_plan(
plan=plan,
model=model,
response_api_optional_request_params=anthropic_messages_optional_request_params,
logging_obj=logging_obj,
kwargs=kwargs_with_provider,
depth=depth,
max_loops=max_loops,
fingerprints=fingerprints,
fingerprint=fingerprint,
callback=callback,
)
return await self._execute_anthropic_agentic_plan(
plan=plan,
model=model,
@ -5083,7 +5262,7 @@ class BaseLLMHTTPHandler:
else False
)
if websearch_converted_stream:
if api_surface == "anthropic_messages" and websearch_converted_stream:
from typing import cast
from litellm._logging import verbose_logger

View file

@ -51,11 +51,13 @@ class E2BSandboxConfig(BaseSandboxConfig):
timeout: int | None = None,
allow_internet_access: bool = True,
api_key: str | None = None,
api_base: str | None = None,
metadata: dict | None = None,
client: AsyncHTTPHandler | None = None,
**kwargs,
) -> ContainerHandle:
key = self.validate_environment(api_key=api_key)
base = api_base or E2B_API_BASE
body = {
"templateID": template or E2B_DEFAULT_TEMPLATE,
"timeout": timeout if timeout is not None else DEFAULT_SANDBOX_TIMEOUT,
@ -68,7 +70,7 @@ class E2BSandboxConfig(BaseSandboxConfig):
response = cast(
httpx.Response,
await self._http(client).post(
url=f"{E2B_API_BASE}/sandboxes",
url=f"{base}/sandboxes",
headers={"X-API-Key": key, "Content-Type": "application/json"},
json=body,
),
@ -84,6 +86,7 @@ class E2BSandboxConfig(BaseSandboxConfig):
"envd_access_token": data.get("envdAccessToken"),
"traffic_access_token": data.get("trafficAccessToken"),
"api_key": key,
"api_base": base,
}
return handle
@ -130,6 +133,7 @@ class E2BSandboxConfig(BaseSandboxConfig):
*,
container: Union[ContainerHandle, str],
api_key: str | None = None,
api_base: str | None = None,
client: AsyncHTTPHandler | None = None,
**kwargs,
) -> bool:
@ -139,11 +143,12 @@ class E2BSandboxConfig(BaseSandboxConfig):
or handle._hidden_params.get("api_key")
or self.validate_environment()
)
base = api_base or handle._hidden_params.get("api_base") or E2B_API_BASE
try:
response = cast(
httpx.Response,
await self._http(client).delete(
url=f"{E2B_API_BASE}/sandboxes/{handle.id}",
url=f"{base}/sandboxes/{handle.id}",
headers={"X-API-Key": key},
),
)

View file

@ -1,5 +1,15 @@
import json
from typing import Any, List, Literal, Optional, Tuple, Union, cast
from typing import (
Any,
AsyncIterator,
Iterator,
List,
Literal,
Optional,
Tuple,
Union,
cast,
)
import httpx
@ -15,7 +25,6 @@ from litellm.litellm_core_utils.llm_response_utils.get_headers import (
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionImageObject,
ChatCompletionToolParam,
OpenAIChatCompletionToolParam,
)
@ -25,6 +34,7 @@ from litellm.types.utils import (
Function,
Message,
ModelResponse,
ModelResponseStream,
ProviderSpecificModelInfo,
)
from litellm.utils import (
@ -34,10 +44,34 @@ from litellm.utils import (
supports_tool_choice,
)
from ...openai.chat.gpt_transformation import OpenAIGPTConfig
from ...openai.chat.gpt_transformation import (
OpenAIChatCompletionStreamingHandler,
OpenAIGPTConfig,
)
from ..common_utils import FireworksAIException
def _extract_fireworks_hidden_params(payload: dict) -> dict:
"""
Collect Fireworks-specific response fields (perf_metrics, prompt_token_ids,
per-choice raw_output and token_ids) from a non-streaming completion payload
or a single streaming chunk, so the same data lands in ``_hidden_params`` on
both response paths.
"""
choices = [c for c in (payload.get("choices") or []) if isinstance(c, dict)]
top_level = {
f"fireworks_{field}": payload[field]
for field in ("perf_metrics", "prompt_token_ids")
if field in payload
}
per_choice = {
f"fireworks_{dest}": [c[field] for c in choices if field in c]
for field, dest in (("raw_output", "raw_outputs"), ("token_ids", "token_ids"))
if any(field in c for c in choices)
}
return {**top_level, **per_choice}
class FireworksAIConfig(OpenAIGPTConfig):
"""
Reference: https://docs.fireworks.ai/api-reference/post-chatcompletions
@ -60,8 +94,7 @@ class FireworksAIConfig(OpenAIGPTConfig):
logprobs: Optional[int] = None
reasoning_effort: Optional[str] = None
# Non OpenAI parameters - Fireworks AI only params
prompt_truncate_length: Optional[int] = None
prompt_truncate_len: Optional[int] = None
context_length_exceeded_behavior: Optional[Literal["error", "truncate"]] = None
def __init__(
@ -80,7 +113,7 @@ class FireworksAIConfig(OpenAIGPTConfig):
user: Optional[str] = None,
logprobs: Optional[int] = None,
reasoning_effort: Optional[str] = None,
prompt_truncate_length: Optional[int] = None,
prompt_truncate_len: Optional[int] = None,
context_length_exceeded_behavior: Optional[Literal["error", "truncate"]] = None,
) -> None:
locals_ = locals().copy()
@ -108,8 +141,30 @@ class FireworksAIConfig(OpenAIGPTConfig):
"response_format",
"user",
"logprobs",
"prompt_truncate_length",
"prompt_truncate_len",
"context_length_exceeded_behavior",
"seed",
"top_logprobs",
"min_p",
"typical_p",
"repetition_penalty",
"mirostat_target",
"mirostat_lr",
"logit_bias",
"echo",
"echo_last",
"ignore_eos",
"prompt_cache_key",
"prompt_cache_isolation_key",
"raw_output",
"perf_metrics_in_response",
"return_token_ids",
"safe_tokenization",
"service_tier",
"speculation",
"prediction",
"stream_options",
"sampling_mask",
]
# Only add tools for models that support function calling
@ -133,9 +188,11 @@ class FireworksAIConfig(OpenAIGPTConfig):
if supports_tool_choice(model=model, custom_llm_provider="fireworks_ai"):
supported_params.append("tool_choice")
# Only add reasoning_effort for models that support it
# Only add reasoning params for models that support it
if supports_reasoning(model=model, custom_llm_provider="fireworks_ai"):
supported_params.append("reasoning_effort")
supported_params.append("reasoning_history")
supported_params.append("thinking")
return supported_params
@ -151,6 +208,18 @@ class FireworksAIConfig(OpenAIGPTConfig):
param == "tools" and value is not None
for param, value in non_default_params.items()
)
if (
non_default_params.get("thinking") is not None
and non_default_params.get("reasoning_effort") is not None
):
raise litellm.BadRequestError(
message=(
"Fireworks AI chat completions does not support specifying both "
"`thinking` and `reasoning_effort` in the same request."
),
model=model,
llm_provider="fireworks_ai",
)
for param, value in non_default_params.items():
if param == "tool_choice":
@ -174,40 +243,19 @@ class FireworksAIConfig(OpenAIGPTConfig):
optional_params["response_format"] = value
elif param == "max_completion_tokens":
optional_params["max_tokens"] = value
elif param == "reasoning_effort":
if value is True:
optional_params["reasoning_effort"] = "medium"
elif value is False:
optional_params["reasoning_effort"] = "none"
else:
optional_params["reasoning_effort"] = value
elif param in supported_openai_params:
if value is not None:
optional_params[param] = value
return optional_params
def _add_transform_inline_image_block(
self,
content: ChatCompletionImageObject,
model: str,
disable_add_transform_inline_image_block: Optional[bool],
) -> ChatCompletionImageObject:
"""
Add transform_inline to the image_url (allows non-vision models to parse documents/images/etc.)
- ignore if model is a vision model
- ignore if user has disabled this feature
"""
if (
"vision" in model or disable_add_transform_inline_image_block
): # allow user to toggle this feature.
return content
if isinstance(content["image_url"], str):
# Skip base64 data URLs — appending #transform=inline corrupts the
# base64 payload and causes an "Incorrect padding" decode error on
# the Fireworks side. Data URLs are already inlined by definition.
# Lower-case before checking: URI schemes are case-insensitive (RFC 3986).
if not content["image_url"].lower().startswith("data:"):
content["image_url"] = f"{content['image_url']}#transform=inline"
elif isinstance(content["image_url"], dict):
url = content["image_url"]["url"]
if not url.lower().startswith("data:"):
content["image_url"]["url"] = f"{url}#transform=inline"
return content
def _transform_tools(
self, tools: List[OpenAIChatCompletionToolParam]
) -> List[OpenAIChatCompletionToolParam]:
@ -225,36 +273,46 @@ class FireworksAIConfig(OpenAIGPTConfig):
self, messages: List[AllMessageValues], model: str, litellm_params: dict
) -> List[AllMessageValues]:
"""
Add 'transform=inline' to the url of the image_url
Strip fields not permitted by FireworksAI from messages.
"""
from litellm.litellm_core_utils.prompt_templates.common_utils import (
filter_value_from_dict,
migrate_file_to_image_url,
)
disable_add_transform_inline_image_block = cast(
Optional[bool],
litellm_params.get("disable_add_transform_inline_image_block")
or litellm.disable_add_transform_inline_image_block,
supports_vision_value = self._get_model_cost_capability_exact(
model=model, capability="supports_vision"
)
## For any 'file' message type with pdf content, move to 'image_url' message type
for message in messages:
if message["role"] == "user":
_message_content = message.get("content")
if _message_content is not None and isinstance(_message_content, list):
for idx, content in enumerate(_message_content):
if content["type"] == "file":
_message_content[idx] = migrate_file_to_image_url(content)
for message in messages:
if message["role"] == "user":
_message_content = message.get("content")
if _message_content is not None and isinstance(_message_content, list):
for content in _message_content:
if content["type"] == "image_url":
content = self._add_transform_inline_image_block(
content=content,
if not isinstance(content, dict):
continue
if content.get("type") == "file":
raise litellm.BadRequestError(
message=(
"Fireworks AI chat completions does not support "
"file content blocks. For PDFs, convert pages to "
"images and send image_url blocks to a Fireworks "
"vision model, or extract text before calling a "
"text-only model."
),
model=model,
disable_add_transform_inline_image_block=disable_add_transform_inline_image_block,
llm_provider="fireworks_ai",
)
if (
content.get("type") == "image_url"
and supports_vision_value is False
):
raise litellm.BadRequestError(
message=(
f"Fireworks AI model {model} does not support "
"image inputs. Use a Fireworks vision model or "
"remove image_url content blocks."
),
model=model,
llm_provider="fireworks_ai",
)
filter_value_from_dict(cast(dict, message), "cache_control")
# Remove fields not permitted by FireworksAI (additionalProperties: false
@ -317,43 +375,55 @@ class FireworksAIConfig(OpenAIGPTConfig):
return True
return ("-" + key_short + "-") in short_name
def _get_model_cost_capability(self, model: str, capability: str) -> Optional[bool]:
@staticmethod
def _short_model_name(model: str) -> str:
short_name = model
if short_name.startswith("fireworks_ai/"):
short_name = short_name[len("fireworks_ai/") :]
if short_name.startswith("accounts/fireworks/models/"):
short_name = short_name[len("accounts/fireworks/models/") :]
return short_name
candidate_keys = [
def _get_model_cost_capability_exact(
self, model: str, capability: str
) -> Optional[bool]:
short_name = self._short_model_name(model)
candidate_keys = (
model,
f"fireworks_ai/{short_name}",
f"fireworks_ai/accounts/fireworks/models/{short_name}",
]
)
for candidate_key in candidate_keys:
model_info = litellm.model_cost.get(candidate_key)
if model_info is not None and model_info.get(capability) is not None:
return cast(Optional[bool], model_info.get(capability))
return None
# Fallback: preserve historical substring matching for model name
# variants (e.g. fine-tuned or regionally-suffixed versions of a
# known model). Pick the *longest* matching entry so a more specific
# known model (e.g. "qwen3-8b-instruct") wins over a less specific
# one (e.g. "qwen3-8b") when the query model is more specific still.
# Use hyphen-aligned matching to avoid false positives where a short
# known model name is an unrelated substring of a longer one.
best_match_short: Optional[str] = None
best_match_value: Optional[bool] = None
for key_short, model_info in self._get_fireworks_index():
if model_info.get(capability) is None:
continue
if not self._matches_on_hyphen_boundary(short_name, key_short):
continue
if best_match_short is None or len(key_short) > len(best_match_short):
best_match_short = key_short
best_match_value = cast(Optional[bool], model_info.get(capability))
def _get_model_cost_capability(self, model: str, capability: str) -> Optional[bool]:
exact = self._get_model_cost_capability_exact(
model=model, capability=capability
)
if exact is not None:
return exact
return best_match_value
# Fallback: substring matching for model name variants (e.g. fine-tuned
# or regionally-suffixed versions of a known model). Pick the *longest*
# matching entry so a more specific known model (e.g. "qwen3-8b-instruct")
# wins over a less specific one (e.g. "qwen3-8b"). Hyphen-aligned matching
# avoids false positives where a short known name is an unrelated
# substring of a longer one. This stays a soft signal: capability-gated
# hard rejections use the exact lookup so a fuzzy match never blocks a
# custom deployment.
short_name = self._short_model_name(model)
matches = [
(key_short, cast(Optional[bool], model_info.get(capability)))
for key_short, model_info in self._get_fireworks_index()
if model_info.get(capability) is not None
and self._matches_on_hyphen_boundary(short_name, key_short)
]
if not matches:
return None
return max(matches, key=lambda match: len(match[0]))[1]
def get_provider_info(self, model: str) -> ProviderSpecificModelInfo:
supports_function_calling_value = self._get_model_cost_capability(
@ -362,12 +432,16 @@ class FireworksAIConfig(OpenAIGPTConfig):
supports_reasoning_value = self._get_model_cost_capability(
model=model, capability="supports_reasoning"
)
supports_vision_value = self._get_model_cost_capability(
model=model, capability="supports_vision"
)
supports_pdf_input_value = self._get_model_cost_capability(
model=model, capability="supports_pdf_input"
)
provider_specific_model_info: ProviderSpecificModelInfo = {
"supports_function_calling": True,
"supports_prompt_caching": True, # https://docs.fireworks.ai/guides/prompt-caching
"supports_pdf_input": True, # via document inlining
"supports_vision": True, # via document inlining
}
if supports_function_calling_value is not None:
@ -381,6 +455,14 @@ class FireworksAIConfig(OpenAIGPTConfig):
supports_reasoning_value
)
if supports_vision_value is not None:
provider_specific_model_info["supports_vision"] = supports_vision_value
if supports_pdf_input_value is not None:
provider_specific_model_info["supports_pdf_input"] = (
supports_pdf_input_value
)
return provider_specific_model_info
def transform_request(
@ -402,6 +484,15 @@ class FireworksAIConfig(OpenAIGPTConfig):
if "tools" in optional_params and optional_params["tools"] is not None:
tools = self._transform_tools(tools=optional_params["tools"])
optional_params["tools"] = tools
if optional_params.get("stream"):
stream_options = optional_params.get("stream_options")
if stream_options is None:
optional_params["stream_options"] = {"include_usage": True}
elif stream_options.get("include_usage") is not False:
optional_params["stream_options"] = {
**stream_options,
"include_usage": True,
}
return super().transform_request(
model=model,
messages=messages,
@ -494,10 +585,25 @@ class FireworksAIConfig(OpenAIGPTConfig):
)
)
response._hidden_params = {"additional_headers": additional_headers}
response._hidden_params = {
"additional_headers": additional_headers,
**_extract_fireworks_hidden_params(completion_response),
}
return response
def get_model_response_iterator(
self,
streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse],
sync_stream: bool,
json_mode: Optional[bool] = False,
) -> Any:
return FireworksAIChatCompletionStreamingHandler(
streaming_response=streaming_response,
sync_stream=sync_stream,
json_mode=json_mode,
)
def _get_openai_compatible_provider_info(
self, api_base: Optional[str], api_key: Optional[str]
) -> Tuple[Optional[str], Optional[str]]:
@ -554,3 +660,15 @@ class FireworksAIConfig(OpenAIGPTConfig):
or get_secret_str("FIREWORKSAI_API_KEY")
or get_secret_str("FIREWORKS_AI_TOKEN")
)
class FireworksAIChatCompletionStreamingHandler(OpenAIChatCompletionStreamingHandler):
def chunk_parser(self, chunk: dict) -> ModelResponseStream:
parsed = super().chunk_parser(chunk)
fireworks_fields = _extract_fireworks_hidden_params(chunk)
if fireworks_fields:
parsed.provider_specific_fields = {
**(getattr(parsed, "provider_specific_fields", None) or {}),
**fireworks_fields,
}
return parsed

View file

@ -15011,7 +15011,7 @@
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": false
"supports_vision": true
},
"fireworks_ai/accounts/fireworks/models/mixtral-8x22b-instruct-hf": {
"input_cost_per_token": 1.2e-06,
@ -15314,7 +15314,7 @@
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": false
"supports_vision": true
},
"fireworks_ai/qwen3p7-plus": {
"cache_read_input_token_cost": 8e-08,

View file

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

View file

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

View file

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

View file

@ -63,7 +63,7 @@ from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy._types import SpecialMCPServerNames, UserAPIKeyAuth
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
from litellm.proxy.litellm_pre_call_utils import (
LiteLLMProxyRequestSetup,
@ -3352,6 +3352,19 @@ if MCP_AVAILABLE:
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
# A key scoped to no MCP servers opts out of every MCP path. Enforce it
# here too, since toolset scoping replaces mcp_servers and would otherwise
# drop the sentinel. Checked before the admin branch, mirroring
# get_allowed_mcp_servers.
original_op = user_api_key_auth.object_permission
if original_op is not None and SpecialMCPServerNames.no_mcp_servers.value in (
original_op.mcp_servers or []
):
raise HTTPException(
status_code=403,
detail="API key is scoped to no MCP servers; toolset access is denied.",
)
# Access control: non-admin keys must have this toolset in their grant list.
# Use _user_has_admin_view so that PROXY_ADMIN_VIEW_ONLY is also treated as admin.
is_admin = _user_has_admin_view(user_api_key_auth)

View file

@ -461,6 +461,7 @@ class LiteLLMRoutes(enum.Enum):
"/mcp/tools/call",
"/mcp-rest/tools/list",
"/mcp-rest/tools/call",
"/v1/mcp/tools",
]
# MCP server CRUD routes — control-plane. Gated by DISABLE_ADMIN_ENDPOINTS.
@ -2142,6 +2143,20 @@ class UserHeaderMapping(LiteLLMPydanticObjectBase):
UserMCPManagementMode = Literal["restricted", "view_all"]
class PluginConfig(LiteLLMPydanticObjectBase):
"""A single external service registered as an embeddable UI plugin."""
name: str = Field(description="unique plugin identifier (kebab-case)")
display_name: str | None = Field(
None, description="human-readable label shown in the UI view switcher"
)
url: str = Field(description="base URL of the plugin service")
plugin_key: str | None = Field(
None,
description="plugin's own credential, injected as Bearer auth only on /plugin-proxy/<name>/* reverse-proxy calls",
)
class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
"""
Documents all the fields supported by `general_settings` in config.yaml
@ -2150,6 +2165,9 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
completion_model: Optional[str] = Field(
None, description="proxy level default model for all chat completion calls"
)
plugins: list[PluginConfig] | None = Field(
None, description="external services registered as embeddable UI plugins"
)
key_management_system: Optional[KeyManagementSystem] = Field(
None, description="key manager to load keys from / decrypt keys with"
)
@ -2960,6 +2978,10 @@ class SpecialModelNames(enum.Enum):
no_default_models = "no-default-models"
class SpecialMCPServerNames(enum.Enum):
no_mcp_servers = "no-mcp-servers"
class SpecialProxyStrings(enum.Enum):
default_user_id = "default_user_id" # global proxy admin
@ -3808,6 +3830,28 @@ class SpecialHeaders(enum.Enum):
mcp_servers = "x-mcp-servers"
mcp_access_groups = "x-mcp-access-groups"
@classmethod
def litellm_credential_header_names(cls) -> "frozenset[str]":
"""Lowercased header names user_api_key_auth accepts as a litellm key.
Every header here authenticates the caller, so any code that forwards a
request onward (e.g. the plugin reverse proxy) must strip all of them to
avoid leaking the caller's litellm credential downstream. The static
custom-key header (general_settings.litellm_key_header_name) is runtime
config and must be added on top of this set by the caller.
"""
return frozenset(
header.value.lower()
for header in (
cls.openai_authorization,
cls.azure_authorization,
cls.anthropic_authorization,
cls.google_ai_studio_authorization,
cls.azure_apim_authorization,
cls.custom_litellm_api_key,
)
)
class LitellmDataForBackendLLMCall(TypedDict, total=False):
headers: dict

View file

@ -699,6 +699,11 @@ async def common_checks(
if valid_token is not None:
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
LiteLLMProxyRequestSetup.pre_seed_litellm_metadata_for_route(
request_data=request_body,
route=route,
)
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
request_data=request_body,
user_api_key_dict=valid_token,
@ -3658,9 +3663,18 @@ async def _virtual_key_max_budget_check(
# so a NaN max_budget would silently disable enforcement. Treat a
# non-finite max_budget as "no configured limit" rather than as a bypass.
if math.isfinite(valid_token.max_budget) and spend >= valid_token.max_budget:
# name the key in the error so operators don't have to reverse-map
# spend back to a key; key_name is the masked form (last 4 chars)
key_label = valid_token.key_alias or "key"
key_descriptor = (
f"{key_label} ({valid_token.key_name})"
if valid_token.key_name
else key_label
)
raise litellm.BudgetExceededError(
current_cost=spend,
max_budget=valid_token.max_budget,
message=f"Budget has been exceeded! Key={key_descriptor} Current cost: {spend}, Max budget: {valid_token.max_budget}",
)

View file

@ -0,0 +1,13 @@
from __future__ import annotations
from enum import Enum
class AuthMethod(str, Enum):
API_KEY = "api_key"
HTTP_BASIC = "http_basic"
BEARER_JWT = "bearer_jwt"
OAUTH2_INTROSPECTION = "oauth2_introspection"
OIDC = "oidc"
SAML = "saml"
MUTUAL_TLS = "mutual_tls"

View file

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

View file

@ -0,0 +1,101 @@
from __future__ import annotations
import ipaddress
from typing import Any, Union
from fastapi import Request
from pydantic import BaseModel, Field
from litellm._logging import verbose_proxy_logger
TrustedProxyNetwork = Union[ipaddress.IPv4Network, ipaddress.IPv6Network]
class NetworkContext(BaseModel):
client_ip: str | None = None
host: str | None = None
via_trusted_proxy: bool = False
class TrustedProxyConfig(BaseModel):
use_forwarded_for: bool = False
trusted_proxy_cidrs: list[str] = Field(default_factory=list)
def normalize_cidr_ranges(
configured_ranges: Any, *, setting_name: str = "trusted_proxy_cidrs"
) -> list[str]:
if not configured_ranges:
return []
if isinstance(configured_ranges, str):
return [r.strip() for r in configured_ranges.split(",") if r.strip()]
if isinstance(configured_ranges, (list, tuple, set)):
return [str(r).strip() for r in configured_ranges if str(r).strip()]
verbose_proxy_logger.warning(
"Invalid %s value: expected a list of CIDR ranges, got %s",
setting_name,
type(configured_ranges).__name__,
)
return []
def parse_trusted_proxy_ranges(
configured_ranges: Any, *, setting_name: str = "trusted_proxy_cidrs"
) -> list[TrustedProxyNetwork]:
networks: list[TrustedProxyNetwork] = []
for cidr in normalize_cidr_ranges(configured_ranges, setting_name=setting_name):
try:
networks.append(ipaddress.ip_network(cidr, strict=False))
except ValueError:
verbose_proxy_logger.warning(
"Invalid CIDR in %s: %s, skipping", setting_name, cidr
)
return networks
def ip_in_networks(client_ip: str | None, networks: list[TrustedProxyNetwork]) -> bool:
if not client_ip or not networks:
return False
try:
addr = ipaddress.ip_address(client_ip.strip())
except ValueError:
return False
return any(addr in network for network in networks)
def _is_valid_ip(value: str) -> bool:
try:
ipaddress.ip_address(value)
return True
except ValueError:
return False
def resolve_client_ip(
request: Request, config: TrustedProxyConfig
) -> tuple[str | None, bool]:
"""Resolve the real client IP, trusting X-Forwarded-For only when the direct
peer is itself a configured trusted proxy. Walks the header right-to-left and
returns the first hop that is not a trusted proxy, so a forged left-most entry
cannot spoof the client."""
peer = request.client.host if request.client else None
networks = parse_trusted_proxy_ranges(config.trusted_proxy_cidrs)
if not config.use_forwarded_for or not ip_in_networks(peer, networks):
return peer, False
forwarded = request.headers.get("x-forwarded-for", "")
hops = [h.strip() for h in forwarded.split(",") if h.strip()]
for hop in reversed(hops):
if _is_valid_ip(hop) and not ip_in_networks(hop, networks):
return hop, True
return peer, True
def resolve_network_context(
request: Request, config: TrustedProxyConfig
) -> NetworkContext:
ip, via_proxy = resolve_client_ip(request, config)
return NetworkContext(
client_ip=ip,
host=request.headers.get("host"),
via_trusted_proxy=via_proxy,
)

View file

@ -0,0 +1,33 @@
from litellm.proxy.auth.resolvers.exceptions import (
IdentityResolutionError,
KeyNotFoundError,
KeyNotInCacheError,
NoDatabaseConnectionError,
PrincipalMissingSourceKeyError,
)
from litellm.proxy.auth.resolvers.models import (
CredentialRef,
EndUserIdentity,
OrganizationIdentity,
Principal,
PrincipalType,
ProjectIdentity,
TeamIdentity,
UserIdentity,
)
__all__ = [
"CredentialRef",
"EndUserIdentity",
"IdentityResolutionError",
"KeyNotFoundError",
"KeyNotInCacheError",
"NoDatabaseConnectionError",
"OrganizationIdentity",
"Principal",
"PrincipalMissingSourceKeyError",
"PrincipalType",
"ProjectIdentity",
"TeamIdentity",
"UserIdentity",
]

View file

@ -0,0 +1,50 @@
from __future__ import annotations
from fastapi import status
from litellm.proxy._types import ProxyErrorTypes, ProxyException
class IdentityResolutionError(Exception):
"""Base for every failure raised while resolving a caller's identity."""
class NoDatabaseConnectionError(IdentityResolutionError):
def __init__(self) -> None:
super().__init__(
"No DB Connected. See - https://docs.litellm.ai/docs/proxy/virtual_keys"
)
class KeyNotInCacheError(IdentityResolutionError):
def __init__(self, hashed_token: str) -> None:
super().__init__(
f"Key doesn't exist in cache + check_cache_only=True. key={hashed_token}."
)
class KeyNotFoundError(IdentityResolutionError, ProxyException):
"""The token matched nothing in the cache or the verification token table.
Also a ``ProxyException`` so the auth flow keeps mapping a missing key to the
OpenAI 401 contract unchanged while callers migrate onto the typed hierarchy.
"""
def __init__(self, hashed_token: str) -> None:
ProxyException.__init__(
self,
message="Authentication Error, Invalid proxy server token passed. key={}, not found in db. Create key via `/key/generate` call.".format(
hashed_token
),
type=ProxyErrorTypes.token_not_found_in_db,
param="key",
code=status.HTTP_401_UNAUTHORIZED,
)
class PrincipalMissingSourceKeyError(IdentityResolutionError):
def __init__(self) -> None:
super().__init__(
"Principal carries no source key; it was not produced by "
"IdentityStore.resolve"
)

View file

@ -0,0 +1,82 @@
from __future__ import annotations
from enum import Enum
from pydantic import BaseModel, ConfigDict, Field
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_method import AuthMethod
from litellm.proxy.auth.network import NetworkContext
from litellm.proxy.auth.roles import Role, TeamRole
class PrincipalType(str, Enum):
HUMAN = "human"
SERVICE_ACCOUNT = "service_account"
class UserIdentity(BaseModel):
id: str
external_id: str | None = None
user_name: str | None = None
email: str | None = None
display_name: str | None = None
class OrganizationIdentity(BaseModel):
id: str
name: str | None = None
class TeamIdentity(BaseModel):
id: str
name: str | None = None
role: TeamRole = TeamRole.MEMBER
class ProjectIdentity(BaseModel):
id: str
name: str | None = None
class EndUserIdentity(BaseModel):
id: str
class CredentialRef(BaseModel):
key_id: str | None = None
token_id: str | None = None
class Principal(BaseModel):
"""Normalized caller identity, resolved once per request at the auth seam.
Frozen and constructed fresh per request, never cached or shared. The identity
fields carry no policy, budget, or rate-limit state. ``source_key`` is a
transitional carrier for the resolved key object so ``key_from_principal`` can
hand it to the request flow that still consumes ``UserAPIKeyAuth``; it is
excluded from serialization and repr and goes away once those consumers read
identity off the Principal directly.
"""
model_config = ConfigDict(frozen=True)
principal_type: PrincipalType
subject: str
issuer: str | None = None
audience: list[str] = Field(default_factory=list)
user: UserIdentity | None = None
organization: OrganizationIdentity | None = None
teams: list[TeamIdentity] = Field(default_factory=list)
project: ProjectIdentity | None = None
end_user: EndUserIdentity | None = None
roles: list[Role] = Field(default_factory=list)
scopes: list[str] = Field(default_factory=list)
auth_method: AuthMethod
credential_ref: CredentialRef = Field(default_factory=CredentialRef)
network: NetworkContext = Field(default_factory=NetworkContext)
source_key: UserAPIKeyAuth | None = Field(default=None, exclude=True, repr=False)

View file

@ -0,0 +1,206 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Sequence
from pydantic import BaseModel
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_checks import (
_cache_key_object,
_copy_user_api_key_auth_for_cache,
_fetch_key_object_from_db_with_reconnect,
get_object_permission,
)
from litellm.proxy.auth.resolvers.exceptions import (
KeyNotFoundError,
KeyNotInCacheError,
NoDatabaseConnectionError,
PrincipalMissingSourceKeyError,
)
from litellm.proxy.auth.auth_method import AuthMethod
from litellm.proxy.auth.network import NetworkContext
from litellm.proxy.auth.resolvers.models import (
CredentialRef,
EndUserIdentity,
OrganizationIdentity,
Principal,
PrincipalType,
ProjectIdentity,
TeamIdentity,
UserIdentity,
)
from litellm.proxy.auth.roles import TeamRole, map_role, team_role
if TYPE_CHECKING:
from litellm.caching.caching import DualCache
from litellm.integrations.opentelemetry import Span
from litellm.proxy.utils import PrismaClient, ProxyLogging
class IdentityStore:
"""The auth flow's resolver: one combined_view lookup, projected into a Principal.
``resolve`` does the lookup (cache, then DB via the shared lower-level helpers,
then write-back) and returns the per-caller Principal. The Principal carries the
source key object so ``key_from_principal`` can hand it back to the parts of the
request flow that still consume ``UserAPIKeyAuth`` (budget, rate limits, policy);
that carrier is a stopgap until those consumers read identity off the Principal.
The Prisma client, key cache, the request's tracing span / logging sink, and
whether this store may only read the cache are injected so the composition root
can build the store once the proxy DB is connected; the span and logging sink
are infra the DB call is instrumented with and ``check_cache_only`` is a store
mode, none of them inputs to resolving identity. ``auth_checks.get_key_object``
stays as the legacy entrypoint for its other callers until they migrate onto
this store.
"""
def __init__(
self,
prisma_client: PrismaClient | None,
cache: DualCache,
*,
parent_otel_span: Span | None = None,
proxy_logging_obj: ProxyLogging | None = None,
check_cache_only: bool = False,
) -> None:
self._prisma = prisma_client
self._cache = cache
self._parent_otel_span = parent_otel_span
self._proxy_logging_obj = proxy_logging_obj
self._check_cache_only = check_cache_only
async def resolve(
self,
hashed_token: str,
*,
auth_method: AuthMethod = AuthMethod.API_KEY,
network: NetworkContext | None = None,
) -> Principal:
key = await self._resolve_key(hashed_token)
return self._principal_from_key(
key,
auth_method=auth_method,
network=network,
subject_fallback=key.token,
credential_ref=CredentialRef(token_id=key.token),
)
@staticmethod
def key_from_principal(principal: Principal) -> UserAPIKeyAuth:
"""Hand back the resolved key object carried on the Principal.
Stopgap for the request flow that still consumes ``UserAPIKeyAuth`` for
budget, rate-limit, and policy state. Only Principals produced by
``resolve`` carry a source key.
"""
if principal.source_key is None:
raise PrincipalMissingSourceKeyError()
return principal.source_key
async def _resolve_key(self, hashed_token: str) -> UserAPIKeyAuth:
if self._prisma is None:
raise NoDatabaseConnectionError()
cached = await self._cache.async_get_cache(
key=hashed_token, model_type=UserAPIKeyAuth
)
if cached is not None:
return _copy_user_api_key_auth_for_cache(user_api_key_obj=cached)
if self._check_cache_only:
raise KeyNotInCacheError(hashed_token)
from_db: BaseModel | None = await _fetch_key_object_from_db_with_reconnect(
hashed_token=hashed_token,
prisma_client=self._prisma,
parent_otel_span=self._parent_otel_span,
proxy_logging_obj=self._proxy_logging_obj,
)
if from_db is None:
raise KeyNotFoundError(hashed_token)
key = UserAPIKeyAuth(**from_db.model_dump(exclude_none=True))
if key.object_permission_id and not key.object_permission:
try:
key.object_permission = await get_object_permission(
object_permission_id=key.object_permission_id,
prisma_client=self._prisma,
user_api_key_cache=self._cache,
parent_otel_span=self._parent_otel_span,
proxy_logging_obj=self._proxy_logging_obj,
)
except Exception as e:
verbose_proxy_logger.debug(
f"Failed to load object_permission for key with object_permission_id={key.object_permission_id}: {e}"
)
await _cache_key_object(
hashed_token=hashed_token,
user_api_key_obj=key,
user_api_key_cache=self._cache,
proxy_logging_obj=self._proxy_logging_obj,
)
return key
@staticmethod
def _principal_from_key(
key: UserAPIKeyAuth,
*,
auth_method: AuthMethod,
issuer: str | None = None,
subject_fallback: str | None = None,
scopes: Sequence[str] = (),
credential_ref: CredentialRef | None = None,
network: NetworkContext | None = None,
) -> Principal:
"""Project the identity slice off an already-resolved key object and carry
the key on the Principal so ``key_from_principal`` can recover it.
Pure: issues no lookup. Both ``resolve`` and the auth seam call this so
identity is projected once off whichever key object they already hold.
"""
teams: list[TeamIdentity] = []
if key.team_id is not None:
role = (
team_role(key.team_member.role) if key.team_member else TeamRole.MEMBER
)
teams.append(TeamIdentity(id=key.team_id, name=key.team_alias, role=role))
organization = (
OrganizationIdentity(id=key.org_id, name=key.organization_alias)
if key.org_id is not None
else None
)
user = (
UserIdentity(id=key.user_id, email=key.user_email)
if key.user_id is not None
else None
)
project = (
ProjectIdentity(id=key.project_id, name=key.project_alias)
if key.project_id is not None
else None
)
end_user = (
EndUserIdentity(id=key.end_user_id) if key.end_user_id is not None else None
)
mapped = map_role(key.user_role)
return Principal(
principal_type=(
PrincipalType.HUMAN if key.user_id else PrincipalType.SERVICE_ACCOUNT
),
subject=key.user_id or key.key_alias or subject_fallback or "",
issuer=issuer,
user=user,
organization=organization,
teams=teams,
project=project,
end_user=end_user,
roles=[mapped] if mapped else [],
scopes=list(scopes),
auth_method=auth_method,
credential_ref=credential_ref or CredentialRef(),
network=network or NetworkContext(),
source_key=key,
)

View file

@ -0,0 +1,35 @@
from __future__ import annotations
from enum import Enum
class Role(str, Enum):
PLATFORM_ADMIN = "platform_admin"
PLATFORM_VIEWER = "platform_viewer"
ORG_ADMIN = "org_admin"
ORG_VIEWER = "org_viewer"
TEAM_ADMIN = "team_admin"
TEAM_MEMBER = "team_member"
class TeamRole(str, Enum):
ADMIN = "admin"
MEMBER = "member"
_ROLE_MAP: dict[str, Role] = {
"proxy_admin": Role.PLATFORM_ADMIN,
"proxy_admin_viewer": Role.PLATFORM_VIEWER,
"org_admin": Role.ORG_ADMIN,
}
def map_role(value: str | None) -> Role | None:
"""Map a LiteLLM ``user_role`` string to a platform Role."""
if value is None:
return None
return _ROLE_MAP.get(value)
def team_role(role: str | None) -> TeamRole:
return TeamRole.ADMIN if role == "admin" else TeamRole.MEMBER

View file

@ -1,12 +1,15 @@
import ipaddress
from typing import Any, Dict, List, Optional, Union
from typing import Any, Dict, Optional
from fastapi import Request
from litellm._logging import verbose_proxy_logger
from litellm.proxy.auth.network import (
ip_in_networks,
normalize_cidr_ranges,
parse_trusted_proxy_ranges,
)
TRUSTED_PROXY_RANGES_KEY = "trusted_proxy_ranges"
TrustedProxyNetwork = Union[ipaddress.IPv4Network, ipaddress.IPv6Network]
def _get_proxy_general_settings() -> Dict[str, Any]:
@ -18,43 +21,20 @@ def _get_proxy_general_settings() -> Dict[str, Any]:
return {}
def _normalize_cidr_ranges(configured_ranges: Any, *, setting_name: str) -> List[str]:
if not configured_ranges:
return []
if isinstance(configured_ranges, str):
return [
raw_range.strip()
for raw_range in configured_ranges.split(",")
if raw_range.strip()
]
if isinstance(configured_ranges, (list, tuple, set)):
return [
str(raw_range).strip()
for raw_range in configured_ranges
if str(raw_range).strip()
]
verbose_proxy_logger.warning(
"Invalid %s value: expected a list of CIDR ranges, got %s",
setting_name,
type(configured_ranges).__name__,
def get_trusted_proxy_cidrs(
general_settings: dict[str, Any] | None = None,
) -> list[str]:
"""Operator-configured trusted reverse-proxy CIDRs, normalized to strings.
Empty when none are configured, in which case X-Forwarded-For must not be
trusted and only the direct peer is authoritative.
"""
if general_settings is None:
general_settings = _get_proxy_general_settings()
return normalize_cidr_ranges(
general_settings.get(TRUSTED_PROXY_RANGES_KEY),
setting_name=TRUSTED_PROXY_RANGES_KEY,
)
return []
def parse_trusted_proxy_ranges(
configured_ranges: Any,
*,
setting_name: str = TRUSTED_PROXY_RANGES_KEY,
) -> List[TrustedProxyNetwork]:
networks: List[TrustedProxyNetwork] = []
for cidr in _normalize_cidr_ranges(configured_ranges, setting_name=setting_name):
try:
networks.append(ipaddress.ip_network(cidr, strict=False))
except ValueError:
verbose_proxy_logger.warning(
"Invalid CIDR in %s: %s, skipping", setting_name, cidr
)
return networks
def _get_direct_client_ip(request: Request) -> Optional[str]:
@ -65,18 +45,6 @@ def _get_direct_client_ip(request: Request) -> Optional[str]:
return None
def _is_ip_in_networks(
client_ip: Optional[str], networks: List[TrustedProxyNetwork]
) -> bool:
if not client_ip or not networks:
return False
try:
addr = ipaddress.ip_address(client_ip.strip())
except ValueError:
return False
return any(addr in network for network in networks)
def require_trusted_proxy_request(
*,
request: Request,
@ -105,7 +73,7 @@ def require_trusted_proxy_request(
)
direct_client_ip = _get_direct_client_ip(request)
if not _is_ip_in_networks(direct_client_ip, trusted_networks):
if not ip_in_networks(direct_client_ip, trusted_networks):
verbose_proxy_logger.warning(
"%s rejected identity headers from untrusted direct client IP %r",
feature_name,

View file

@ -42,7 +42,6 @@ from litellm.proxy.auth.auth_checks import (
common_checks,
get_end_user_object,
get_jwt_key_mapping_object,
get_key_object,
get_project_object,
get_team_object,
get_user_object,
@ -63,7 +62,12 @@ from litellm.proxy.auth.auth_utils import (
from litellm.proxy.auth.handle_jwt import JWTAuthManager, JWTHandler
from litellm.proxy.auth.oauth2_check import Oauth2Handler
from litellm.proxy.auth.oauth2_proxy_hook import handle_oauth2_proxy_request
from litellm.proxy.auth.auth_method import AuthMethod
from litellm.proxy.auth.network import TrustedProxyConfig, resolve_network_context
from litellm.proxy.auth.resolvers import CredentialRef, Principal
from litellm.proxy.auth.resolvers.store import IdentityStore
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.auth.trusted_proxy_utils import get_trusted_proxy_cidrs
from litellm.proxy.common_utils.cache_coordinator import EventDrivenCacheCoordinator
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
@ -758,12 +762,13 @@ async def _auto_register_jwt_mapping(
claim_value,
)
auto_registered_key = await get_key_object(
hashed_token=token_hash,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
auto_registered_key = IdentityStore.key_from_principal(
await IdentityStore(
prisma_client,
user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
).resolve(hashed_token=token_hash)
)
if auto_registered_key is not None:
auto_registered_key.org_id = org_id
@ -865,12 +870,13 @@ async def _resolve_jwt_to_virtual_key(
)
return None
elif cached_mapping is not None:
return await get_key_object(
hashed_token=cached_mapping,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
return IdentityStore.key_from_principal(
await IdentityStore(
prisma_client,
user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
).resolve(hashed_token=cached_mapping)
)
# Resolve the mapping from DB, or treat prisma_client=None as a definitive
@ -889,12 +895,13 @@ async def _resolve_jwt_to_virtual_key(
value=token_hash,
ttl=jwt_handler.litellm_jwtauth.virtual_key_mapping_cache_ttl,
)
return await get_key_object(
hashed_token=token_hash,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
return IdentityStore.key_from_principal(
await IdentityStore(
prisma_client,
user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
).resolve(hashed_token=token_hash)
)
# No mapping found (DB miss or no DB) — apply no-match policy.
@ -1493,13 +1500,14 @@ async def _user_api_key_auth_builder(
## Check CACHE
try:
with tracer.trace("litellm.proxy.auth.get_key_object_check_cache"):
valid_token = await get_key_object(
hashed_token=hash_token(api_key),
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_cache_only=True,
valid_token = IdentityStore.key_from_principal(
await IdentityStore(
prisma_client,
user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_cache_only=True,
).resolve(hashed_token=hash_token(api_key))
)
except Exception:
verbose_logger.debug("api key not found in cache.")
@ -1679,12 +1687,13 @@ async def _user_api_key_auth_builder(
try:
with tracer.trace("litellm.proxy.auth.get_key_object_from_db"):
valid_token = await get_key_object(
hashed_token=api_key,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
valid_token = IdentityStore.key_from_principal(
await IdentityStore(
prisma_client,
user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
).resolve(hashed_token=api_key)
)
except ProxyException as e:
if e.code == 401 or e.code == "401":
@ -2387,6 +2396,17 @@ async def _run_centralized_common_checks(
llm_router=llm_router,
)
# Pin the metadata variable name (litellm_metadata vs metadata) before
# any tag merge runs. Without this, header tags from
# apply_client_tag_policy_pre_auth would land in `metadata` while the
# later seed in common_checks pushes key tags and the
# _tag_max_budget_check read into `litellm_metadata`, hiding header
# tags from per-tag budget enforcement on LITELLM_METADATA_ROUTES.
LiteLLMProxyRequestSetup.pre_seed_litellm_metadata_for_route(
request_data=request_data,
route=route,
)
# Merge x-litellm-tags into request_data BEFORE common_checks runs.
# _tag_max_budget_check inside common_checks only inspects request_data;
# without this pre-merge, header-supplied tags bypass tag-budget
@ -2501,6 +2521,34 @@ def _should_skip_budget_checks(
return False
def _resolve_request_principal(
request: Request, valid_token: UserAPIKeyAuth
) -> Principal:
"""Project the resolved identity into one per-request Principal, off the key
object the builder already fetched, and stamp the request network context
onto it once. X-Forwarded-For is only trusted when the operator configured
``trusted_proxy_ranges``; otherwise the direct peer is authoritative.
credential_ref and a stable subject fallback are always set off the token so
the Principal can never be anonymous, even for a keyless service-account key
with no user or alias."""
cidrs = get_trusted_proxy_cidrs()
network = resolve_network_context(
request,
TrustedProxyConfig(use_forwarded_for=bool(cidrs), trusted_proxy_cidrs=cidrs),
)
auth_method = (
AuthMethod.BEARER_JWT if valid_token.jwt_claims else AuthMethod.API_KEY
)
return IdentityStore._principal_from_key(
valid_token,
auth_method=auth_method,
network=network,
subject_fallback=valid_token.token,
credential_ref=CredentialRef(token_id=valid_token.token),
)
@tracer.wrap()
async def user_api_key_auth(
request: Request,
@ -2615,6 +2663,22 @@ async def user_api_key_auth(
model=request_data.get("model") if isinstance(request_data, dict) else None,
)
user_api_key_auth_obj.request_route = normalize_request_route(route)
# Resolve caller identity once, here at the seam, into a single per-request
# Principal projected off the key object the builder already fetched (no
# second lookup). Downstream consumers read identity off this instead of
# re-resolving it. Additive and defensive: a projection failure must never
# reject an already-authenticated request, so it is left unset on failure;
# any future consumer must treat a missing principal as deny, not allow.
try:
request.state.principal = _resolve_request_principal(
request, user_api_key_auth_obj
)
except Exception as e:
verbose_proxy_logger.warning(
"Principal projection at auth seam failed (non-fatal): %s", e
)
return user_api_key_auth_obj

View file

@ -71,6 +71,23 @@ def initialize_callbacks_on_proxy(
imported_list.append(compression_interception_obj)
continue
if (
isinstance(callback, str)
and callback == "code_interpreter_interception"
):
from litellm.integrations.code_interpreter_interception.handler import (
CodeInterpreterInterceptionLogger,
)
code_interpreter_interception_obj = (
CodeInterpreterInterceptionLogger.initialize_from_proxy_config(
litellm_settings=litellm_settings,
callback_specific_params=callback_specific_params,
)
)
imported_list.append(code_interpreter_interception_obj)
continue
# check if callback is a custom logger compatible callback
if isinstance(callback, str):
callback = LoggingCallbackManager._add_custom_callback_generic_api_str(

View file

@ -0,0 +1,17 @@
model_list:
- model_name: gpt-5
litellm_params:
model: openai/gpt-5
# Sandbox tools configuration
sandbox_tools:
- sandbox_tool_name: "my-e2b"
litellm_params:
sandbox_provider: "e2b"
api_key: os.environ/E2B_API_KEY
litellm_settings:
callbacks: ["code_interpreter_interception"]
code_interpreter_interception_params:
enabled_providers: ["openai"]
sandbox_tool_name: "my-e2b"

View file

@ -108,6 +108,7 @@ def parse_cache_control(cache_control):
LITELLM_METADATA_ROUTES = (
"batches",
"bedrock",
"/v1/messages",
"responses",
"files",
@ -141,6 +142,20 @@ _UNTRUSTED_ROOT_CONTROL_FIELDS = (
"service_callback",
"logger_fn",
"litellm_disabled_callbacks",
# Agentic-loop control fields. These bound or drive an interceptor's agentic
# loop (web search, compression, code interpreter) and are server-controlled.
# A client-supplied value would forge loop depth/cycle state, mark an
# interception as active (triggering sandbox code execution without the
# native tool ever being present), force the completed response to be
# re-wrapped as a synthetic stream the caller never asked for, or raise the
# loop ceiling to drive many upstream model calls and sandbox executions
# from a single request.
"_agentic_loop_depth",
"_agentic_loop_fingerprints",
"_code_interpreter_interception_active",
"_code_interpreter_interception_converted_stream",
"_code_interpreter_interception_sandbox_key",
"max_agentic_loops",
)
_UNTRUSTED_METADATA_CONTROL_FIELDS = (
@ -1223,6 +1238,27 @@ class LiteLLMProxyRequestSetup:
return tags
@staticmethod
def pre_seed_litellm_metadata_for_route(
request_data: dict,
route: str,
) -> None:
"""Pre-seed ``litellm_metadata`` for routes that track tags there.
Routes in ``LITELLM_METADATA_ROUTES`` (e.g. Bedrock, ``/v1/messages``,
responses, batches, files) store request-scoped tag metadata in
``litellm_metadata`` rather than the provider-facing ``metadata``
field. ``get_metadata_variable_name_from_kwargs`` picks the target
based on whether ``litellm_metadata`` is present, so it must be
seeded BEFORE any tag merge runs; otherwise header tags from
``apply_client_tag_policy_pre_auth`` land in ``metadata`` while
key tags from ``apply_key_tags_pre_auth`` and the read in
``_tag_max_budget_check`` resolve to ``litellm_metadata``, leaving
header tags invisible to per-tag budget enforcement.
"""
if any(metadata_route in route for metadata_route in LITELLM_METADATA_ROUTES):
request_data.setdefault("litellm_metadata", {})
@staticmethod
def apply_key_tags_pre_auth(
request_data: dict,
@ -1454,8 +1490,7 @@ async def add_litellm_data_to_request(
_metadata_variable_name=_metadata_variable_name,
)
# Add headers to metadata for guardrails to access (fixes #17477)
# Guardrails use metadata["headers"] to access request headers (e.g., User-Agent)
# Expose request headers under the metadata field for guardrails (fixes #17477)
if _metadata_variable_name in data and isinstance(
data[_metadata_variable_name], dict
):

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -0,0 +1,344 @@
"""
Plugin proxy routes for litellm.
Enables external services to register as plugins and be proxied through
the litellm proxy server.
Config (in litellm config.yaml general_settings):
plugins:
- name: my-plugin
url: "http://localhost:3210"
display_name: "My Plugin"
plugin_key: "sk-..." # optional: plugin's own auth key
Plugin iframe auth:
The UI calls GET /api/plugins/auth-token to receive a short-lived identity
claim ({user_id, user_role, plugin, exp}) encrypted with a per-plugin key
derived as HMAC-SHA256(LITELLM_SALT_KEY, plugin_name). The claim carries no
litellm bearer token, so a compromised plugin learns only the caller's
identity, never their credential. LITELLM_SALT_KEY itself is never shared
with plugins — each plugin holds only its own derived key.
"""
import base64
import hashlib
import hmac as _hmac
import json
import os
import time
from collections.abc import Mapping
from cryptography.fernet import Fernet, InvalidToken
from fastapi import APIRouter, Depends, HTTPException, Request, Response
from litellm.proxy._types import PluginConfig, SpecialHeaders, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.custom_http import httpxSpecialProvider
router = APIRouter()
# Hop-by-hop headers (RFC 7230) and the litellm session cookie — never forwarded
# to a plugin backend. Credential headers are added on top per-request from the
# canonical SpecialHeaders set so the plugin only ever authenticates via its own
# injected plugin_key.
_HOP_BY_HOP_STRIP = frozenset(
{
"host",
"connection",
"transfer-encoding",
"te",
"trailers",
"upgrade",
"cookie",
}
)
def _configured_key_header_names() -> frozenset[str]:
"""The lowercased general_settings.litellm_key_header_name, if configured.
Read live from the proxy module (not import-time) so a custom key header set
via config is honoured without a restart. Returns empty when unset.
"""
try:
from litellm.proxy import proxy_server
except Exception:
return frozenset()
general_settings = getattr(proxy_server, "general_settings", None)
if not isinstance(general_settings, dict):
return frozenset()
name: object = general_settings.get("litellm_key_header_name")
return frozenset({name.lower()}) if isinstance(name, str) and name else frozenset()
def _request_strip_headers() -> frozenset[str]:
"""Headers to drop before forwarding a request to a plugin backend.
Every header user_api_key_auth accepts as a litellm credential is stripped —
Authorization, x-api-key, API-Key, x-goog-api-key, Ocp-Apim-Subscription-Key,
x-litellm-api-key, and any configured custom key header — so a plugin can
never be handed the caller's live litellm key (confused-deputy escalation).
"""
return (
_HOP_BY_HOP_STRIP
| SpecialHeaders.litellm_credential_header_names()
| _configured_key_header_names()
)
# Headers to strip from plugin RESPONSES before returning to the browser.
# httpx already decompresses and de-chunks the body, so forwarding the wire
# encoding headers causes clients to attempt double-decompression (garbage) or
# incorrect length checks. set-cookie is removed so plugins cannot overwrite
# litellm session cookies.
_RESPONSE_STRIP = {
"content-encoding",
"transfer-encoding",
"content-length",
"set-cookie",
}
def _safe_response_headers(raw: "Mapping[str, str]") -> dict[str, str]:
"""Strip wire-encoding/cookie headers and force proxied responses inert.
Plugin-controlled bytes are served from the litellm dashboard origin, so a
compromised plugin could return an HTML/JS document that executes with the
admin's session against same-origin management APIs. A sandbox CSP forces
the response into an opaque origin with scripts disabled, and nosniff stops
content-type confusion from re-enabling execution. Both are set last so a
plugin cannot override them with its own headers.
"""
return {
**{k: v for k, v in raw.items() if k.lower() not in _RESPONSE_STRIP},
"content-security-policy": "sandbox",
"x-content-type-options": "nosniff",
}
# In-memory plugin registry — populated from general_settings at startup
_plugin_registry: dict[str, PluginConfig] = {}
# ---------------------------------------------------------------------------
# Key derivation — audience-scoped per plugin so compromising one plugin
# cannot be used to forge claims for another. LITELLM_SALT_KEY is NEVER
# shared with plugins; each plugin only receives a key derived from
# HMAC(LITELLM_SALT_KEY, plugin_name) which reveals nothing about the master.
# ---------------------------------------------------------------------------
def _plugin_fernet(plugin_name: str) -> Fernet:
"""Return a Fernet cipher whose key is scoped to a specific plugin.
Key material: HMAC-SHA256(LITELLM_SALT_KEY, plugin_name).
A plugin possessing its own key cannot derive the master salt or
forge claims intended for a different plugin.
"""
salt = os.getenv("LITELLM_SALT_KEY", "").encode()
derived = _hmac.new(salt, plugin_name.encode(), hashlib.sha256).digest()
return Fernet(base64.urlsafe_b64encode(derived))
_CLAIM_TTL_SECONDS = 30 # identity claims expire after 30 s
def issue_plugin_session_claim(
plugin_name: str, user_id: str | None, user_role: str | None
) -> str:
"""Issue a short-lived, audience-scoped identity claim for the plugin.
The claim contains {user_id, user_role, plugin, exp}. Crucially it
contains NO litellm bearer token — the plugin can only derive the
caller's identity, not act as them against the proxy.
"""
claim = {
"plugin": plugin_name,
"user_id": user_id or "",
"user_role": user_role or "",
"exp": int(time.time()) + _CLAIM_TTL_SECONDS,
}
return _plugin_fernet(plugin_name).encrypt(json.dumps(claim).encode()).decode()
def verify_plugin_session_claim(plugin_name: str, ciphertext: str) -> dict:
"""Verify and decode a plugin session claim.
Raises ValueError if the HMAC is invalid, the audience is wrong, or
the claim is expired. Returns the decoded claim dict on success.
"""
try:
raw = _plugin_fernet(plugin_name).decrypt(
ciphertext.encode(), ttl=_CLAIM_TTL_SECONDS
)
claim = json.loads(raw)
except (InvalidToken, Exception) as exc:
raise ValueError("Invalid, tampered, or expired plugin session claim") from exc
if claim.get("plugin") != plugin_name:
raise ValueError("Plugin claim audience mismatch")
if int(claim.get("exp", 0)) < int(time.time()):
raise ValueError("Plugin session claim expired")
return claim
# ---------------------------------------------------------------------------
# Config
# ---------------------------------------------------------------------------
def register_plugins_from_config(general_settings: dict[str, object]) -> None:
"""Replace the plugin registry from general_settings.
Replaces (not merges) so plugins removed from config are immediately
unreachable without requiring a process restart.
"""
raw = general_settings.get("plugins")
entries: list[object] = raw if isinstance(raw, list) else []
new_registry = {
p.name: p for p in (PluginConfig.model_validate(entry) for entry in entries)
}
_plugin_registry.clear()
_plugin_registry.update(new_registry)
# ---------------------------------------------------------------------------
# Routes
# ---------------------------------------------------------------------------
@router.get("/api/plugins", tags=["plugins"])
async def list_plugins(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> list[dict[str, str]]:
"""Return registered plugins for authenticated UI callers.
plugin_key is never returned — the browser never needs it (the proxy injects
it server-side from the registry), and exposing it here would leak the
credential into React state and DevTools. Admin key management goes through
the redacted /config/field/info path instead.
"""
return [
{
"name": plugin.name,
"display_name": plugin.display_name or plugin.name,
"url": plugin.url,
}
for plugin in _plugin_registry.values()
]
@router.get("/api/plugins/auth-token", tags=["plugins"])
async def plugin_auth_token(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
plugin_name: str = "litellm-platform-plugin",
) -> dict:
"""Issue a short-lived, audience-scoped plugin session claim.
The claim contains {user_id, user_role, plugin, exp}. It does NOT
contain the caller's litellm bearer token — a compromised plugin can
only learn the caller's identity, not impersonate them against the proxy.
Encrypted with a key derived from HMAC(LITELLM_SALT_KEY, plugin_name),
so each plugin holds only its own key and cannot forge claims for others.
Requires LITELLM_SALT_KEY to be set; returns 503 otherwise.
"""
if not os.getenv("LITELLM_SALT_KEY"):
raise HTTPException(
status_code=503,
detail="LITELLM_SALT_KEY is not configured; plugin iframe auth unavailable.",
)
if plugin_name not in _plugin_registry:
raise HTTPException(
status_code=404, detail=f"Plugin '{plugin_name}' is not registered."
)
user_id = getattr(user_api_key_dict, "user_id", None)
user_role = getattr(user_api_key_dict, "user_role", None)
return {
"session_claim": issue_plugin_session_claim(plugin_name, user_id, user_role)
}
@router.api_route(
"/plugin-proxy/{plugin_name}/{path:path}",
methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS", "HEAD"],
tags=["plugins"],
include_in_schema=False,
)
async def plugin_proxy(
plugin_name: str,
path: str,
request: Request,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> Response:
"""Authenticated reverse-proxy to a registered plugin backend.
Restricted to proxy_admin callers — the shared plugin_key must not be
usable as a confused-deputy credential by regular users. Plugin UIs
talk to the plugin service directly via the iframe; this route is for
administrative and server-to-server access only.
The caller's litellm credential is stripped and replaced with the
plugin's own plugin_key so plugins never receive a live litellm API key.
"""
if getattr(user_api_key_dict, "user_role", None) != "proxy_admin":
return Response(
content="Plugin proxy access requires proxy_admin role.",
status_code=403,
)
plugin = _plugin_registry.get(plugin_name)
if not plugin:
return Response(
content=f"Plugin '{plugin_name}' not registered",
status_code=404,
)
target_url = f"{plugin.url.rstrip('/')}/{path}"
query = request.url.query
if query:
target_url = f"{target_url}?{query}"
body = await request.body()
# Strip caller credentials and hop-by-hop headers from forwarded request
strip = _request_strip_headers()
forward_headers = {
k: v for k, v in request.headers.items() if k.lower() not in strip
}
# Inject plugin's own credential as upstream auth (if configured)
plugin_key = plugin.plugin_key
if plugin_key:
forward_headers["authorization"] = f"Bearer {plugin_key}"
# Forward caller identity so the plugin can enforce its own access control.
# The plugin MUST NOT trust these as credentials — they are informational.
# The plugin_key above is the only authentication mechanism.
user_id = getattr(user_api_key_dict, "user_id", None)
user_role = getattr(user_api_key_dict, "user_role", None)
if user_id:
forward_headers["x-litellm-user-id"] = str(user_id)
if user_role:
forward_headers["x-litellm-user-role"] = str(user_role)
handler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.PassThroughEndpoint
)
try:
req = handler.client.build_request(
method=request.method,
url=target_url,
headers=forward_headers,
content=body,
)
# Do not follow redirects — a redirect to an internal URL would allow
# the plugin to SSRF the proxy into fetching arbitrary internal services.
resp = await handler.client.send(req, follow_redirects=False)
except Exception:
return Response(
content=f"Cannot connect to plugin '{plugin_name}' at {plugin.url}",
status_code=502,
)
return Response(
content=resp.content,
status_code=resp.status_code,
headers=_safe_response_headers(resp.headers),
)

View file

@ -419,6 +419,10 @@ from litellm.proxy.management_endpoints.workflow_management_endpoints import (
)
from litellm.proxy.management_helpers.audit_logs import create_audit_log_for_update
from litellm.proxy.memory.memory_endpoints import router as memory_router
from litellm.proxy.plugin_routes import (
router as plugin_router,
register_plugins_from_config,
)
from litellm.proxy.middleware.in_flight_requests_middleware import (
InFlightRequestsMiddleware,
)
@ -4502,6 +4506,8 @@ class ProxyConfig:
load_from_azure_key_vault(use_azure_key_vault=use_azure_key_vault)
### ALERTING ###
self._load_alerting_settings(general_settings=general_settings)
### PLUGINS ###
register_plugins_from_config(general_settings)
### CONNECT TO DATABASE ###
database_url = general_settings.get("database_url", None)
if database_url and database_url.startswith("os.environ/"):
@ -4773,6 +4779,11 @@ class ProxyConfig:
config
)
## SANDBOX TOOLS SETTINGS
from litellm.sandbox.sandbox_tools import register_sandbox_tools
register_sandbox_tools(config.get("sandbox_tools") or [])
## /fine_tuning/jobs endpoints config
finetuning_config = config.get("finetune_settings", None)
set_fine_tuning_config(config=finetuning_config)
@ -5663,6 +5674,10 @@ class ProxyConfig:
llm_router=llm_router,
)
if _general_settings is not None and "plugins" in _general_settings:
general_settings["plugins"] = _general_settings["plugins"]
register_plugins_from_config(general_settings)
async def _reschedule_spend_log_cleanup_job(self):
"""
Reschedule the spend log cleanup job based on current general_settings.
@ -11566,9 +11581,12 @@ async def _get_caller_byok_team_scope(
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
):
return None
key_team_scope: set[str] = (
{user_api_key_dict.team_id} if user_api_key_dict.team_id else set()
)
user_id = user_api_key_dict.user_id
if user_id is None:
return set()
return key_team_scope
try:
user_row = await UserRepository(prisma_client).table.find_unique(
where={"user_id": user_id}
@ -11576,12 +11594,12 @@ async def _get_caller_byok_team_scope(
except Exception:
verbose_proxy_logger.exception(
"Failed to look up caller teams while scoping BYOK search; "
"defaulting to no team access."
"defaulting to key team scope only."
)
return set()
return key_team_scope
if user_row is None:
return set()
return set(user_row.teams or [])
return key_team_scope
return key_team_scope | set(user_row.teams or [])
def _byok_row_outside_caller_teams(
@ -14931,6 +14949,41 @@ async def update_config(
Keep it more precise, to prevent overwrite other values unintentially
"""
_PLUGIN_KEY_REDACTED = "***"
def _preserve_redacted_plugin_keys(incoming: object, existing: object) -> object:
"""Restore real plugin_key values the client never sees.
/config/field/info redacts every plugin_key to ``"***"``, so an admin
editing a plugin posts that placeholder (or a blank, when the UI clears the
field) straight back. Treat a blank or redacted plugin_key as "keep the
stored credential" by sourcing it from the existing config; only a real,
non-redacted value replaces it, and a blank with no stored key drops the
field entirely instead of persisting the placeholder.
"""
if not isinstance(incoming, list):
return incoming
stored_keys = {
p["name"]: p["plugin_key"]
for p in (existing if isinstance(existing, list) else [])
if isinstance(p, dict) and p.get("name") and p.get("plugin_key")
}
def resolve(plugin: object) -> object:
if not isinstance(plugin, dict):
return plugin
key = plugin.get("plugin_key")
if key not in (None, "", _PLUGIN_KEY_REDACTED):
return plugin
name = plugin.get("name")
if name in stored_keys:
return {**plugin, "plugin_key": stored_keys[name]}
return {k: v for k, v in plugin.items() if k != "plugin_key"}
return [resolve(p) for p in incoming]
@router.post(
"/config/field/update",
@ -14997,7 +15050,13 @@ async def update_config_general_settings(
## update db
general_settings[data.field_name] = data.field_value
field_value = data.field_value
if data.field_name == "plugins":
field_value = _preserve_redacted_plugin_keys(
field_value, general_settings.get("plugins")
)
general_settings[data.field_name] = field_value
response = await ConfigRepository(prisma_client).table.upsert(
where={"param_name": "general_settings"},
@ -15008,6 +15067,9 @@ async def update_config_general_settings(
)
await invalidate_config_param("general_settings")
if data.field_name == "plugins":
register_plugins_from_config(general_settings)
return response
@ -15063,9 +15125,19 @@ async def get_config_general_settings(
general_settings = dict(db_general_settings.param_value)
if field_name in general_settings:
return ConfigFieldInfo(
field_name=field_name, field_value=general_settings[field_name]
)
field_value = general_settings[field_name]
# Redact plugin_key from plugin configs so the shared credential
# is never returned even to admin-viewer callers.
if field_name == "plugins" and isinstance(field_value, list):
field_value = [
(
{k: ("***" if k == "plugin_key" else v) for k, v in p.items()}
if isinstance(p, dict)
else p
)
for p in field_value
]
return ConfigFieldInfo(field_name=field_name, field_value=field_value)
else:
raise HTTPException(
status_code=400,
@ -16387,6 +16459,7 @@ app.include_router(model_access_group_management_router)
app.include_router(tag_management_router)
app.include_router(workflow_management_router)
app.include_router(memory_router)
app.include_router(plugin_router)
app.include_router(cost_tracking_settings_router)
app.include_router(router_settings_router)
app.include_router(fallback_management_router)

View file

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

View file

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

View file

@ -70,6 +70,7 @@ async def acreate_sandbox(
timeout: int | None = None,
allow_internet_access: bool = True,
api_key: str | None = None,
api_base: str | None = None,
**kwargs,
) -> ContainerHandle:
_update_logging(kwargs, provider, "create_sandbox")
@ -78,6 +79,7 @@ async def acreate_sandbox(
timeout=timeout,
allow_internet_access=allow_internet_access,
api_key=api_key,
api_base=api_base,
**_forward_kwargs(kwargs),
)
@ -104,12 +106,14 @@ async def adelete_sandbox(
provider: str,
container: Union[ContainerHandle, str],
api_key: str | None = None,
api_base: str | None = None,
**kwargs,
) -> bool:
_update_logging(kwargs, provider, "delete_sandbox")
return await _get_config(provider).adelete_sandbox(
container=container,
api_key=api_key,
api_base=api_base,
**_forward_kwargs(kwargs),
)
@ -121,6 +125,7 @@ async def acode_interpreter_tool(
template: str | None = None,
timeout: int | None = None,
api_key: str | None = None,
api_base: str | None = None,
**kwargs,
) -> CodeExecutionResult:
_update_logging(kwargs, provider, "code_interpreter_tool")
@ -128,7 +133,11 @@ async def acode_interpreter_tool(
forwarded = _forward_kwargs(kwargs)
container = await config.acreate_sandbox(
template=template, timeout=timeout, api_key=api_key, **forwarded
template=template,
timeout=timeout,
api_key=api_key,
api_base=api_base,
**forwarded,
)
try:
return await config.arun_code(
@ -137,7 +146,7 @@ async def acode_interpreter_tool(
finally:
try:
await config.adelete_sandbox(
container=container, api_key=api_key, **forwarded
container=container, api_key=api_key, api_base=api_base, **forwarded
)
except Exception as e:
litellm._logging.verbose_logger.debug(

View file

@ -0,0 +1,60 @@
"""
Registry for sandbox tools configured via the proxy's top-level `sandbox_tools`.
A sandbox tool maps a name to a sandbox provider plus its credentials, so the
code interpreter interceptor can resolve a tool by name to provider/key/base.
"""
from collections.abc import Iterator
from litellm._logging import verbose_logger
_SANDBOX_TOOL_REGISTRY: dict[str, dict] = {}
def _resolve_secret_value(value: str | None) -> str | None:
if not isinstance(value, str):
return None
if value.startswith("os.environ/"):
from litellm.secret_managers.main import get_secret_str
return get_secret_str(value)
return value
def _iter_valid_tools(tools: list[dict]) -> Iterator[tuple[str, dict]]:
for tool in tools:
if not isinstance(tool, dict):
verbose_logger.warning("sandbox_tools: skipping non-dict entry %r", tool)
continue
name = tool.get("sandbox_tool_name")
if not name:
verbose_logger.warning(
"sandbox_tools: skipping entry missing 'sandbox_tool_name': %r", tool
)
continue
params = tool.get("litellm_params") or {}
provider = params.get("sandbox_provider")
if not provider:
verbose_logger.warning(
"sandbox_tools: skipping entry missing 'sandbox_provider': %r", tool
)
continue
yield name, {
"sandbox_provider": provider,
"api_key": _resolve_secret_value(params.get("api_key")),
"api_base": _resolve_secret_value(params.get("api_base")),
}
def register_sandbox_tools(tools: list[dict]) -> None:
global _SANDBOX_TOOL_REGISTRY
_SANDBOX_TOOL_REGISTRY = dict(_iter_valid_tools(tools))
def resolve_sandbox_tool(name: str) -> dict | None:
return _SANDBOX_TOOL_REGISTRY.get(name)
def clear_sandbox_tools() -> None:
register_sandbox_tools([])

View file

@ -0,0 +1,22 @@
"""
Type definitions for Code Interpreter Interception integration.
"""
from typing import List, TypedDict
class CodeInterpreterInterceptionConfig(TypedDict, total=False):
"""
Configuration parameters for CodeInterpreterInterceptionLogger.
Used in proxy_config.yaml under litellm_settings:
litellm_settings:
code_interpreter_interception_params:
enabled: true
enabled_providers: ["openai"]
sandbox_tool_name: "my_e2b_sandbox"
"""
enabled: bool
enabled_providers: List[str]
sandbox_tool_name: str

View file

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

View file

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

View file

@ -37,6 +37,8 @@ from pydantic import (
ConfigDict,
Field,
PrivateAttr,
SkipValidation,
field_serializer,
field_validator,
)
from typing_extensions import Required, TypedDict
@ -3610,10 +3612,20 @@ class LiteLLMBatch(Batch):
class LiteLLMRealtimeStreamLoggingObject(LiteLLMPydanticObjectBase):
results: OpenAIRealtimeStreamList
# Events are already well-formed provider dicts. Validating them against the
# OpenAIRealtimeEvents union makes Pydantic try every member per event, which
# floods thousands of ValidationErrors for events outside the union (e.g.
# rate_limits.updated), blocks the event loop, and discards the session usage.
results: SkipValidation[OpenAIRealtimeStreamList]
usage: Usage
_hidden_params: dict = {}
@field_serializer("results")
def _serialize_results(
self, results: OpenAIRealtimeStreamList
) -> List[Dict[str, Any]]:
return [dict(event) for event in results]
def __contains__(self, key):
# Define custom behavior for the 'in' operator
return hasattr(self, key)

View file

@ -15019,7 +15019,7 @@
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": false
"supports_vision": true
},
"fireworks_ai/accounts/fireworks/models/mixtral-8x22b-instruct-hf": {
"input_cost_per_token": 1.2e-06,
@ -15322,7 +15322,7 @@
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": false
"supports_vision": true
},
"fireworks_ai/qwen3p7-plus": {
"cache_read_input_token_cost": 8e-08,

View file

@ -5,13 +5,9 @@ Each rule has a hard ceiling (baseline + slack) in ruff-strict-budget.json. The
gate counts each rule across the whole tree and fails when a rule is both over
its ceiling and higher than the base it merges into, so a change is blamed for
the violations it adds, never for drift that already exists in the base.
The base is the merge-base of the current branch with --base; this matches CI,
which checks out the PR head sha and runs the gate against the PR's base sha.
"""
import argparse
import contextlib
import json
import re
import shutil
@ -19,7 +15,6 @@ import subprocess
import sys
import tempfile
from collections import Counter
from collections.abc import Iterator
from pathlib import Path
from typing import NamedTuple
@ -45,12 +40,6 @@ class Breach(NamedTuple):
added: int
class GateInputs(NamedTuple):
head: list[Violation]
base: dict[str, int]
changed: dict[str, set[int]]
def _run(cmd: list, cwd: Path = REPO_ROOT) -> str:
proc = subprocess.run(cmd, cwd=cwd, capture_output=True, text=True)
if proc.returncode not in (0, 1):
@ -67,14 +56,14 @@ def _ruff_json(cwd: Path, config: Path) -> list:
return json.loads(raw or "[]")
def collect_violations(root: Path, config: Path) -> list:
def head_violations() -> list:
out = []
for item in _ruff_json(root, config):
for item in _ruff_json(REPO_ROOT, STRICT_CONFIG):
name = Path(item["filename"])
rel = (
(name if name.is_absolute() else root / name)
(name if name.is_absolute() else REPO_ROOT / name)
.resolve()
.relative_to(root)
.relative_to(REPO_ROOT)
.as_posix()
)
out.append(Violation(rel, item["location"]["row"], item["code"]))
@ -85,29 +74,27 @@ def count_by_rule(violations: list) -> dict:
return dict(Counter(v.code for v in violations))
@contextlib.contextmanager
def _temp_worktree(ref: str) -> Iterator[Path]:
parent = Path(tempfile.mkdtemp(prefix="ruff_wt_"))
def base_counts(ref: str) -> dict:
parent = Path(tempfile.mkdtemp(prefix="ruff_base_"))
worktree = parent / "wt"
try:
_run(["git", "worktree", "add", "--detach", str(worktree), ref])
yield worktree
shutil.copy(STRICT_CONFIG, worktree / "ruff-strict.toml")
items = _ruff_json(worktree, worktree / "ruff-strict.toml")
return dict(Counter(item["code"] for item in items))
finally:
subprocess.run(
["git", "worktree", "remove", "--force", str(worktree)],
cwd=REPO_ROOT,
capture_output=True,
text=True,
)
_run(["git", "worktree", "remove", "--force", str(worktree)])
shutil.rmtree(parent, ignore_errors=True)
def base_counts(ref: str) -> dict:
with _temp_worktree(ref) as worktree:
shutil.copy(STRICT_CONFIG, worktree / "ruff-strict.toml")
return count_by_rule(
collect_violations(worktree, worktree / "ruff-strict.toml")
)
def evaluate(head: dict, base: dict, budget: dict) -> list:
breaches = []
for rule, spec in budget.items():
cap = spec["baseline"] + spec["slack"]
total = head.get(rule, 0)
if total > cap and total > base.get(rule, 0):
breaches.append(Breach(rule, total, cap, total - base.get(rule, 0)))
return sorted(breaches)
def parse_changed_lines(diff_text: str) -> dict:
@ -123,31 +110,24 @@ def parse_changed_lines(diff_text: str) -> dict:
return changed
def evaluate(head: dict, base: dict, budget: dict) -> list:
breaches = []
for rule, spec in budget.items():
cap = spec["baseline"] + spec["slack"]
total = head.get(rule, 0)
if total > cap and total > base.get(rule, 0):
breaches.append(Breach(rule, total, cap, total - base.get(rule, 0)))
return sorted(breaches)
def introduced(violations: list, changed: dict) -> list:
return [v for v in violations if v.line in changed.get(v.file, set())]
def gather(base: str) -> GateInputs:
def cmd_check(base: str) -> None:
budget = json.loads(BUDGET_PATH.read_text())
head = head_violations()
base_point = _run(["git", "merge-base", base, "HEAD"]).strip() or base
diff = _run(["git", "diff", base_point, "--unified=0", "--no-color", "--", TARGET])
return GateInputs(
collect_violations(REPO_ROOT, STRICT_CONFIG),
base_counts(base_point),
parse_changed_lines(diff),
breaches = evaluate(count_by_rule(head), base_counts(base_point), budget)
if not breaches:
print(f"OK: every strict rule is within its codebase ceiling (base {base})")
return
new = introduced(
head,
parse_changed_lines(
_run(["git", "diff", base_point, "--unified=0", "--no-color", "--", TARGET])
),
)
def report(breaches: list, new: list, base: str) -> None:
print(f"FAIL: strict-rule totals exceed their ceiling (base {base}):")
for breach in breaches:
print(
@ -158,24 +138,12 @@ def report(breaches: list, new: list, base: str) -> None:
print(
"Reduce the new violations or remove an equal number elsewhere; the ceiling is baseline + slack in ruff-strict-budget.json."
)
summary = "; ".join(f"{b.rule} {b.total}/{b.cap} (+{b.added})" for b in breaches)
print(f"BREACHED RULES: {summary}")
def cmd_check(base: str) -> None:
budget = json.loads(BUDGET_PATH.read_text())
inputs = gather(base)
breaches = evaluate(count_by_rule(inputs.head), inputs.base, budget)
if not breaches:
print(f"OK: every strict rule is within its codebase ceiling (base {base})")
return
report(breaches, introduced(inputs.head, inputs.changed), base)
raise SystemExit(1)
def cmd_update() -> None:
budget = json.loads(BUDGET_PATH.read_text())
head = count_by_rule(collect_violations(REPO_ROOT, STRICT_CONFIG))
head = count_by_rule(head_violations())
for rule in budget:
budget[rule]["baseline"] = head.get(rule, 0)
BUDGET_PATH.write_text(json.dumps(budget, indent=2, sort_keys=True) + "\n")

View file

@ -4,6 +4,8 @@ import pytest
from fastapi import Request
from litellm_enterprise.proxy.auth.user_api_key_auth import enterprise_custom_auth
from litellm.proxy._types import UserAPIKeyAuth
@pytest.mark.asyncio
async def test_enterprise_custom_auth_none_user_auth():
@ -49,16 +51,19 @@ async def test_enterprise_custom_auth_returns_string():
mock_user_auth = AsyncMock(return_value="sk-test-key")
request = MagicMock(spec=Request)
with patch(
"litellm.proxy.auth.user_api_key_auth.enterprise_custom_auth", mock_user_auth
), patch("litellm.proxy.proxy_server.master_key", "sk-1234"), patch(
"litellm.proxy.proxy_server.prisma_client", MagicMock()
with (
patch(
"litellm.proxy.auth.user_api_key_auth.enterprise_custom_auth",
mock_user_auth,
),
patch("litellm.proxy.proxy_server.master_key", "sk-1234"),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
):
# Verify the key is correctly handled in _user_api_key_auth_builder
with patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object"
"litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key"
) as mock_get_key_object:
mock_get_key_object.return_value = MagicMock(
mock_get_key_object.return_value = UserAPIKeyAuth(
token="sk-test-key",
user_role="internal_user",
team_id=None,
@ -82,9 +87,7 @@ async def test_enterprise_custom_auth_returns_string():
except Exception as e:
print("error:", e)
# Verify get_key_object was called with the correct key
# Verify the key lookup was called with the correct hashed key
mock_get_key_object.assert_called_once()
# The key should be hashed before being passed to get_key_object
assert mock_get_key_object.call_args[1]["hashed_token"] == hash_token(
"sk-test-key"
)
# The key should be hashed before being passed to the resolver
assert mock_get_key_object.call_args[0][0] == hash_token("sk-test-key")

View file

@ -1588,18 +1588,21 @@ def test_token_counter_with_image_url_with_detail_high():
assert _tokens == DEFAULT_IMAGE_TOKEN_COUNT + 7
def test_fireworks_ai_document_inlining():
def test_fireworks_ai_vision_capability_from_cost_map(monkeypatch):
"""
With document inlining, all fireworks ai models are now:
- supports_pdf
- supports_vision
Fireworks deprecated document inlining on 2025-06-30, so vision/PDF support is
no longer hardcoded to True for every Fireworks model. Capabilities are read
from the model cost map: unmapped models no longer advertise vision or PDF
support, while mapped VLMs still do.
"""
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
from litellm.utils import supports_pdf_input, supports_vision
litellm._turn_on_debug()
assert supports_vision("fireworks_ai/llama-3.1-8b-instruct") is False
assert supports_pdf_input("fireworks_ai/llama-3.1-8b-instruct") is False
assert supports_pdf_input("fireworks_ai/llama-3.1-8b-instruct") is True
assert supports_vision("fireworks_ai/llama-3.1-8b-instruct") is True
assert supports_vision("fireworks_ai/minimax-m3") is True
def test_logprobs_type():

View file

@ -9,7 +9,6 @@ sys.path.insert(
import litellm
from litellm import transcription
from litellm.llms.fireworks_ai.chat.transformation import FireworksAIConfig
from base_llm_unit_tests import BaseLLMChatTest
from base_audio_transcription_unit_tests import BaseLLMAudioTranscriptionTest
fireworks = FireworksAIConfig()
@ -146,11 +145,8 @@ class TestFireworksAIAudioTranscription(BaseLLMAudioTranscriptionTest):
)
def test_document_inlining_example(disable_add_transform_inline_image_block):
"""
Document inlining appends ``#transform=inline`` to image/PDF URLs in the
outgoing request unless explicitly disabled. Assert the transform on the
serialized payload rather than making a live Fireworks call — the live
call only proved the model responded and broke whenever Fireworks rotated
its serverless model catalog.
Fireworks document inlining has been removed from the platform. LiteLLM
must not append ``#transform=inline`` regardless of the legacy disable flag.
"""
from unittest.mock import patch
@ -163,7 +159,7 @@ def test_document_inlining_example(disable_add_transform_inline_image_block):
with patch.object(client, "post") as mock_post:
try:
completion(
model="fireworks_ai/accounts/fireworks/models/deepseek-v3p1",
model="fireworks_ai/accounts/fireworks/models/minimax-m3",
messages=[
{
"role": "user",
@ -182,89 +178,80 @@ def test_document_inlining_example(disable_add_transform_inline_image_block):
disable_add_transform_inline_image_block=disable_add_transform_inline_image_block,
client=client,
)
except Exception as e:
print(e)
except Exception:
pass
mock_post.assert_called_once()
json_data = json.loads(mock_post.call_args.kwargs["data"])
sent_url = json_data["messages"][0]["content"][0]["image_url"]["url"]
if disable_add_transform_inline_image_block is True:
assert sent_url == pdf_url
assert "#transform=inline" not in sent_url
else:
assert sent_url == pdf_url + "#transform=inline"
assert sent_url == pdf_url
assert "#transform=inline" not in sent_url
@pytest.mark.parametrize(
"content, model, expected_url",
"content, expected_url",
[
(
{"image_url": "http://example.com/image.png"},
"gpt-4",
"http://example.com/image.png#transform=inline",
"http://example.com/image.png",
),
(
{"image_url": {"url": "http://example.com/image.png"}},
"gpt-4",
{"url": "http://example.com/image.png#transform=inline"},
{"url": "http://example.com/image.png"},
),
(
{"image_url": "http://example.com/image.png"},
"vision-gpt",
"http://example.com/image.png",
),
# data: URLs must never have #transform=inline appended — doing so
# corrupts the base64 payload (fixes #23583).
# URI schemes are case-insensitive (RFC 3986) so check all variants.
(
{"image_url": "data:image/png;base64,iVBORw0KGgo="},
"gpt-4",
"data:image/png;base64,iVBORw0KGgo=",
),
(
{"image_url": {"url": "data:image/jpeg;base64,/9j/4AAQ=="}},
"gpt-4",
{"url": "data:image/jpeg;base64,/9j/4AAQ=="},
),
(
{"image_url": "Data:image/png;base64,iVBORw0KGgo="},
"gpt-4",
"Data:image/png;base64,iVBORw0KGgo=",
),
],
)
def test_transform_inline(content, model, expected_url):
def test_transform_inline_no_longer_added(content, expected_url):
image_block = {"type": "image_url", **content}
messages = [{"role": "user", "content": [image_block]}]
result = litellm.FireworksAIConfig()._add_transform_inline_image_block(
content=content, model=model, disable_add_transform_inline_image_block=False
result = litellm.FireworksAIConfig()._transform_messages_helper(
messages=messages,
model="accounts/fireworks/models/minimax-m3",
litellm_params={},
)
result_image_block = result[0]["content"][0]
if isinstance(expected_url, str):
assert result["image_url"] == expected_url
assert result_image_block["image_url"] == expected_url
else:
assert result["image_url"]["url"] == expected_url["url"]
assert result_image_block["image_url"]["url"] == expected_url["url"]
@pytest.mark.parametrize(
"model, is_disabled, expected_url",
[
("gpt-4", True, "http://example.com/image.png"),
("vision-gpt", False, "http://example.com/image.png"),
("gpt-4", False, "http://example.com/image.png#transform=inline"),
],
"is_disabled",
[True, False],
)
def test_global_disable_flag(model, is_disabled, expected_url):
content = {"image_url": "http://example.com/image.png"}
result = litellm.FireworksAIConfig()._add_transform_inline_image_block(
content=content,
model=model,
disable_add_transform_inline_image_block=is_disabled,
def test_global_disable_flag_no_longer_adds_transform_inline(is_disabled):
url = "http://example.com/image.png"
litellm.disable_add_transform_inline_image_block = is_disabled
messages = [
{
"role": "user",
"content": [{"type": "image_url", "image_url": url}],
}
]
result = litellm.FireworksAIConfig()._transform_messages_helper(
messages=messages,
model="accounts/fireworks/models/minimax-m3",
litellm_params={},
)
assert result["image_url"] == expected_url
assert result[0]["content"][0]["image_url"] == url
litellm.disable_add_transform_inline_image_block = False # Reset for other tests
def test_global_disable_flag_with_transform_messages_helper(monkeypatch):
from openai import OpenAI
from unittest.mock import patch
from litellm import completion
from litellm.llms.custom_httpx.http_handler import HTTPHandler
@ -279,7 +266,7 @@ def test_global_disable_flag_with_transform_messages_helper(monkeypatch):
) as mock_post:
try:
completion(
model="fireworks_ai/accounts/fireworks/models/deepseek-v3p1",
model="fireworks_ai/accounts/fireworks/models/minimax-m3",
messages=[
{
"role": "user",
@ -296,11 +283,10 @@ def test_global_disable_flag_with_transform_messages_helper(monkeypatch):
],
client=client,
)
except Exception as e:
print(e)
except Exception:
pass
mock_post.assert_called_once()
print(mock_post.call_args.kwargs)
json_data = json.loads(mock_post.call_args.kwargs["data"])
assert (
"#transform=inline"

View file

@ -60,7 +60,8 @@ async def test_jwt_to_virtual_key_mapping_resolution():
# Use patch to mock get_key_object in the module where it's used
with patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object", new_callable=AsyncMock
"litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key",
new_callable=AsyncMock,
) as mock_get_key:
mock_get_key.return_value = mock_key_obj
@ -105,7 +106,8 @@ async def test_jwt_to_virtual_key_mapping_no_mapping():
# Mock get_key_object just in case
with patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object", new_callable=AsyncMock
"litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key",
new_callable=AsyncMock,
):
user_api_key_cache = DualCache()
@ -481,7 +483,8 @@ async def test_reject_behavior_raises_403_on_no_mapping():
user_api_key_cache = DualCache()
with patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object", new_callable=AsyncMock
"litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key",
new_callable=AsyncMock,
):
with pytest.raises(HTTPException) as exc_info:
await _resolve_jwt_to_virtual_key(
@ -519,7 +522,8 @@ async def test_reject_behavior_caches_sentinel_after_db_miss():
user_api_key_cache = DualCache()
with patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object", new_callable=AsyncMock
"litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key",
new_callable=AsyncMock,
):
# First call — DB miss, should raise 403 and write sentinel
with pytest.raises(HTTPException) as exc_info:
@ -578,7 +582,8 @@ async def test_reject_behavior_raises_403_on_cached_no_mapping():
await user_api_key_cache.async_set_cache(cache_key, "__NO_MAPPING__")
with patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object", new_callable=AsyncMock
"litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key",
new_callable=AsyncMock,
):
with pytest.raises(HTTPException) as exc_info:
await _resolve_jwt_to_virtual_key(
@ -671,7 +676,7 @@ async def test_auto_register_creates_key_and_mapping_when_helper_invoked():
with (
patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object",
"litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key",
new_callable=AsyncMock,
) as mock_get_key,
patch(
@ -803,7 +808,7 @@ async def test_auto_register_race_condition_unique_conflict():
with (
patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object",
"litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key",
new_callable=AsyncMock,
) as mock_get_key,
patch(
@ -833,13 +838,7 @@ async def test_auto_register_race_condition_unique_conflict():
# Cache should hold the winner's token, not the loser's
cached = await user_api_key_cache.async_get_cache("jwt_key_mapping:sub:user-42")
assert cached == "winner_token_hash"
mock_get_key.assert_called_once_with(
hashed_token="winner_token_hash",
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=None,
)
mock_get_key.assert_called_once_with("winner_token_hash")
# ──────────────────────────────────────────────
@ -1093,7 +1092,7 @@ async def test_auto_register_race_conflict_tolerates_delete_failure():
with (
patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object",
"litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key",
new_callable=AsyncMock,
) as mock_get_key,
patch(
@ -1249,7 +1248,7 @@ async def test_auto_register_helper_stamps_validated_identity_context():
with (
patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object",
"litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key",
new_callable=AsyncMock,
) as mock_get_key,
patch(

View file

@ -0,0 +1,982 @@
"""
Unit tests for CodeInterpreterInterceptionLogger.
All sandbox dependencies are injected (dependency injection, no monkeypatch):
a FakeSandbox stands in for the real e2b config and records how it is called.
"""
import time
import pytest
from litellm.integrations.code_interpreter_interception.handler import (
CodeInterpreterInterceptionLogger,
LITELLM_CODE_EXECUTION_TOOL_NAME,
)
from litellm.llms.base_llm.sandbox.transformation import CodeExecutionResult
from litellm.types.utils import CallTypes
_ACTIVE_KEY = "_code_interpreter_interception_active"
_SANDBOX_KEY = "_code_interpreter_interception_sandbox_key"
class FakeHandle:
def __init__(self, sandbox_id="sbx_fake"):
self.id = sandbox_id
class FakeSandbox:
"""Records acreate_sandbox / arun_code / adelete_sandbox calls."""
def __init__(self, stdout="42"):
self.stdout = stdout
self.create_calls = []
self.run_calls = []
self.delete_calls = []
async def acreate_sandbox(self, **kwargs):
self.create_calls.append(kwargs)
return FakeHandle()
async def arun_code(self, *, container, code, **kwargs):
self.run_calls.append({"container": container, "code": code})
return CodeExecutionResult(stdout=self.stdout)
async def adelete_sandbox(self, *, container, **kwargs):
self.delete_calls.append({"container": container})
return True
class FakeLogging:
def __init__(self, litellm_call_id="k1"):
self.litellm_call_id = litellm_call_id
self.model_call_details = {}
def _function_call_item(call_id="c1", name=LITELLM_CODE_EXECUTION_TOOL_NAME):
return {
"type": "function_call",
"call_id": call_id,
"name": name,
"arguments": '{"code":"print(40 + 2)"}',
}
class FakeResponse:
def __init__(self, output):
self.output = output
def _iter_messages(plan):
patch = plan.request_patch
assert patch is not None, "plan.request_patch must be set"
assert patch.messages is not None, "plan.request_patch.messages must be set"
return patch.messages
@pytest.mark.asyncio
async def test_build_plan_runs_code_and_feeds_output_back():
sandbox = FakeSandbox(stdout="42")
logger = CodeInterpreterInterceptionLogger(sandbox_config=sandbox)
response = FakeResponse(output=[_function_call_item()])
plan = await logger.async_build_agentic_loop_plan(
tools={
"tool_calls": [
{
"call_id": "c1",
"name": LITELLM_CODE_EXECUTION_TOOL_NAME,
"arguments": '{"code":"print(40 + 2)"}',
}
]
},
model="gpt-5",
messages=[{"role": "user", "content": "x"}],
response=response,
anthropic_messages_provider_config=None,
anthropic_messages_optional_request_params={"tools": []},
logging_obj=FakeLogging(litellm_call_id="k1"),
stream=False,
kwargs={"litellm_call_id": "k1", _SANDBOX_KEY: "sbxkey1"},
)
assert sandbox.run_calls, "sandbox.arun_code must be invoked"
assert sandbox.run_calls[0]["code"] == "print(40 + 2)"
messages = _iter_messages(plan)
outputs = [
m
for m in messages
if isinstance(m, dict) and m.get("type") == "function_call_output"
]
assert outputs, "expected a function_call_output item appended"
output_item = next(m for m in outputs if m.get("call_id") == "c1")
assert "42" in str(output_item["output"])
@pytest.mark.asyncio
async def test_pre_call_converts_code_interpreter_tool():
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
kwargs = {
"tools": [{"type": "code_interpreter", "container": {"type": "auto"}}],
"custom_llm_provider": "openai",
}
result = await logger.async_pre_call_deployment_hook(kwargs, CallTypes.aresponses)
assert result is not None
tools = result["tools"]
assert not any(
t.get("type") == "code_interpreter" for t in tools
), "code_interpreter tool must be removed"
names = [t.get("name") or (t.get("function") or {}).get("name") for t in tools]
assert LITELLM_CODE_EXECUTION_TOOL_NAME in names
@pytest.mark.asyncio
@pytest.mark.parametrize(
"tool_choice",
[
{"type": "code_interpreter"},
{"type": "hosted_tool", "name": "code_interpreter"},
],
)
async def test_pre_call_rewrites_forced_code_interpreter_tool_choice(tool_choice):
"""A forced tool_choice targeting the native code_interpreter tool must be
rewritten to the generated function tool; otherwise the outbound request
references a tool that no longer exists and the provider rejects it."""
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
kwargs = {
"tools": [{"type": "code_interpreter", "container": {"type": "auto"}}],
"tool_choice": tool_choice,
"custom_llm_provider": "openai",
}
result = await logger.async_pre_call_deployment_hook(kwargs, CallTypes.aresponses)
assert result is not None
assert result["tool_choice"] == {
"type": "function",
"name": LITELLM_CODE_EXECUTION_TOOL_NAME,
}
@pytest.mark.asyncio
async def test_pre_call_leaves_unrelated_tool_choice_untouched():
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
kwargs = {
"tools": [{"type": "code_interpreter", "container": {"type": "auto"}}],
"tool_choice": "auto",
"custom_llm_provider": "openai",
}
result = await logger.async_pre_call_deployment_hook(kwargs, CallTypes.aresponses)
assert result is not None
assert result["tool_choice"] == "auto"
@pytest.mark.asyncio
async def test_pre_call_noop_on_non_responses():
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
kwargs = {
"tools": [{"type": "code_interpreter", "container": {"type": "auto"}}],
"custom_llm_provider": "openai",
}
result = await logger.async_pre_call_deployment_hook(kwargs, CallTypes.acompletion)
assert result is None
@pytest.mark.asyncio
async def test_should_run_detects_only_matching_function_call():
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
active_kwargs = {"_code_interpreter_interception_active": True}
match = FakeResponse(
output=[_function_call_item(name=LITELLM_CODE_EXECUTION_TOOL_NAME)]
)
should_run, payload = await logger.async_should_run_agentic_loop(
response=match,
model="gpt-5",
messages=[{"role": "user", "content": "x"}],
tools=[],
stream=False,
custom_llm_provider="openai",
kwargs=active_kwargs,
)
assert should_run is True
assert payload.get("tool_calls")
no_match = FakeResponse(output=[_function_call_item(name="something_else")])
should_run2, payload2 = await logger.async_should_run_agentic_loop(
response=no_match,
model="gpt-5",
messages=[{"role": "user", "content": "x"}],
tools=[],
stream=False,
custom_llm_provider="openai",
kwargs=active_kwargs,
)
assert should_run2 is False
assert payload2 == {}
@pytest.mark.asyncio
async def test_container_reused_within_request_via_server_sandbox_key():
sandbox = FakeSandbox(stdout="42")
logger = CodeInterpreterInterceptionLogger(sandbox_config=sandbox)
response = FakeResponse(output=[_function_call_item()])
common = dict(
tools={
"tool_calls": [
{
"call_id": "c1",
"name": LITELLM_CODE_EXECUTION_TOOL_NAME,
"arguments": '{"code":"print(40 + 2)"}',
}
]
},
model="gpt-5",
messages=[{"role": "user", "content": "x"}],
response=response,
anthropic_messages_provider_config=None,
anthropic_messages_optional_request_params={"tools": []},
stream=False,
)
await logger.async_build_agentic_loop_plan(
logging_obj=FakeLogging(litellm_call_id="k1"),
kwargs={"litellm_call_id": "k1", _SANDBOX_KEY: "server-nonce-1"},
**common,
)
await logger.async_build_agentic_loop_plan(
logging_obj=FakeLogging(litellm_call_id="k1"),
kwargs={"litellm_call_id": "k1", _SANDBOX_KEY: "server-nonce-1"},
**common,
)
assert (
len(sandbox.create_calls) == 1
), "the sandbox is reused across loop iterations sharing one server sandbox key"
@pytest.mark.asyncio
async def test_colliding_caller_call_id_does_not_share_sandbox():
"""Two requests with the same caller-controlled litellm_call_id but distinct
server-minted sandbox keys must NOT share a container; otherwise one user's
code could read another in-flight request's sandbox state."""
sandbox = FakeSandbox(stdout="42")
logger = CodeInterpreterInterceptionLogger(sandbox_config=sandbox)
common = dict(
tools={
"tool_calls": [
{
"call_id": "c1",
"name": LITELLM_CODE_EXECUTION_TOOL_NAME,
"arguments": '{"code":"print(1)"}',
}
]
},
model="gpt-5",
messages=[{"role": "user", "content": "x"}],
response=FakeResponse(output=[_function_call_item()]),
anthropic_messages_provider_config=None,
anthropic_messages_optional_request_params={"tools": []},
stream=False,
)
await logger.async_build_agentic_loop_plan(
logging_obj=FakeLogging(litellm_call_id="shared"),
kwargs={"litellm_call_id": "shared", _SANDBOX_KEY: "nonce-A"},
**common,
)
await logger.async_build_agentic_loop_plan(
logging_obj=FakeLogging(litellm_call_id="shared"),
kwargs={"litellm_call_id": "shared", _SANDBOX_KEY: "nonce-B"},
**common,
)
assert (
len(sandbox.create_calls) == 2
), "distinct server sandbox keys must isolate sandboxes despite a colliding call id"
@pytest.mark.asyncio
async def test_pre_call_mints_server_sandbox_key():
"""The interceptor mints a server-side sandbox key (not derived from the
caller-controlled call id) when it activates."""
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
kwargs = {
"tools": [{"type": "code_interpreter", "container": {"type": "auto"}}],
"custom_llm_provider": "openai",
"litellm_call_id": "caller-supplied",
_SANDBOX_KEY: "caller-forged",
}
result = await logger.async_pre_call_deployment_hook(kwargs, CallTypes.aresponses)
assert result is not None
assert result[_SANDBOX_KEY] not in ("caller-forged", "caller-supplied")
assert len(result[_SANDBOX_KEY]) >= 16
@pytest.mark.asyncio
async def test_build_plan_records_code_interpreter_call_metadata():
sandbox = FakeSandbox(stdout="42")
logger = CodeInterpreterInterceptionLogger(sandbox_config=sandbox)
plan = await logger.async_build_agentic_loop_plan(
tools={
"tool_calls": [
{
"call_id": "c1",
"name": LITELLM_CODE_EXECUTION_TOOL_NAME,
"arguments": '{"code":"print(40 + 2)"}',
}
]
},
model="gpt-5",
messages=[{"role": "user", "content": "x"}],
response=FakeResponse(output=[_function_call_item()]),
anthropic_messages_provider_config=None,
anthropic_messages_optional_request_params={"tools": []},
logging_obj=FakeLogging(litellm_call_id="k1"),
stream=False,
kwargs={"litellm_call_id": "k1", _SANDBOX_KEY: "sbxkey1"},
)
calls = plan.metadata["code_interpreter_calls"]
assert calls, "build_plan must record a code_interpreter_call for re-injection"
assert calls[0]["code"] == "print(40 + 2)"
assert calls[0]["container_id"] == "sbx_fake"
assert calls[0]["type"] == "code_interpreter_call"
assert calls[0]["status"] == "completed"
assert calls[0]["outputs"] == [{"type": "logs", "logs": "42"}], (
"outputs must be an OpenAI-shaped logs array (not None) so clients that "
"iterate over code_interpreter_call.outputs do not break"
)
@pytest.mark.asyncio
async def test_build_plan_outputs_empty_array_when_no_stdout():
"""No stdout must still yield an iteration-safe empty array, never None."""
sandbox = FakeSandbox(stdout="")
logger = CodeInterpreterInterceptionLogger(sandbox_config=sandbox)
plan = await logger.async_build_agentic_loop_plan(
tools={
"tool_calls": [
{
"call_id": "c1",
"name": LITELLM_CODE_EXECUTION_TOOL_NAME,
"arguments": '{"code":"pass"}',
}
]
},
model="gpt-5",
messages=[{"role": "user", "content": "x"}],
response=FakeResponse(output=[_function_call_item()]),
anthropic_messages_provider_config=None,
anthropic_messages_optional_request_params={"tools": []},
logging_obj=FakeLogging(litellm_call_id="k1"),
stream=False,
kwargs={"litellm_call_id": "k1", _SANDBOX_KEY: "sbxkey1"},
)
assert plan.metadata["code_interpreter_calls"][0]["outputs"] == []
@pytest.mark.asyncio
async def test_post_hook_injects_code_interpreter_call_matching_openai_shape():
from litellm.types.integrations.custom_logger import AgenticLoopPlan
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
ci_item = {
"id": "ci_x",
"type": "code_interpreter_call",
"status": "completed",
"code": "print(1)",
"container_id": "sbx_fake",
"outputs": [{"type": "logs", "logs": "1"}],
}
plan = AgenticLoopPlan(
run_agentic_loop=True,
metadata={"code_interpreter_calls": [ci_item]},
)
response = FakeResponse(output=[{"type": "message", "content": []}])
out = await logger.async_post_agentic_loop_response_hook(
response=response, plan=plan, kwargs={}
)
types = [item.get("type") for item in out.output]
assert types == ["code_interpreter_call", "message"], (
"code_interpreter_call must be re-injected before the message, matching "
"OpenAI's native output ordering"
)
assert set(out.output[0].keys()) == {
"id",
"type",
"status",
"code",
"container_id",
"outputs",
}, "injected item must match OpenAI's code_interpreter_call keys exactly"
@pytest.mark.asyncio
async def test_post_hook_noop_without_recorded_calls():
from litellm.types.integrations.custom_logger import AgenticLoopPlan
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
response = FakeResponse(output=[{"type": "message", "content": []}])
out = await logger.async_post_agentic_loop_response_hook(
response=response, plan=AgenticLoopPlan(run_agentic_loop=True), kwargs={}
)
assert [item.get("type") for item in out.output] == ["message"]
@pytest.mark.asyncio
async def test_pre_call_forces_non_stream_for_loop():
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
kwargs = {
"tools": [{"type": "code_interpreter", "container": {"type": "auto"}}],
"custom_llm_provider": "openai",
"stream": True,
}
out = await logger.async_pre_call_deployment_hook(kwargs, CallTypes.aresponses)
assert out is not None
assert out["stream"] is False, "loop requires a non-streaming upstream call"
assert out["_code_interpreter_interception_converted_stream"] is True, (
"the converted-stream flag must be set so the final response is wrapped "
"back into a stream for the caller"
)
async def _build_plan(logger, sandbox, call_id="k1", provider="openai"):
return await logger.async_build_agentic_loop_plan(
tools={
"tool_calls": [
{
"call_id": "c1",
"name": LITELLM_CODE_EXECUTION_TOOL_NAME,
"arguments": '{"code":"print(40 + 2)"}',
}
]
},
model="gpt-5",
messages=[{"role": "user", "content": "x"}],
response=FakeResponse(output=[_function_call_item()]),
anthropic_messages_provider_config=None,
anthropic_messages_optional_request_params={"tools": []},
logging_obj=FakeLogging(litellm_call_id=call_id),
stream=False,
kwargs={"litellm_call_id": call_id, _SANDBOX_KEY: "sbxkey1"},
)
@pytest.mark.asyncio
async def test_gate_refuses_without_server_active_marker():
"""A forged litellm_code_execution call must not trigger the loop unless the
pre-call hook actually converted a native code_interpreter tool."""
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
forged = FakeResponse(
output=[_function_call_item(name=LITELLM_CODE_EXECUTION_TOOL_NAME)]
)
should_run, payload = await logger.async_should_run_agentic_loop(
response=forged,
model="gpt-5",
messages=[{"role": "user", "content": "x"}],
tools=[],
stream=False,
custom_llm_provider="openai",
kwargs={},
)
assert should_run is False
assert payload == {}
@pytest.mark.asyncio
async def test_gate_rechecks_provider_scope():
"""enabled_providers must be re-enforced at the gate, not only in pre-call."""
logger = CodeInterpreterInterceptionLogger(
sandbox_config=FakeSandbox(), enabled_providers=["openai"]
)
response = FakeResponse(
output=[_function_call_item(name=LITELLM_CODE_EXECUTION_TOOL_NAME)]
)
should_run, _ = await logger.async_should_run_agentic_loop(
response=response,
model="claude-x",
messages=[{"role": "user", "content": "x"}],
tools=[],
stream=False,
custom_llm_provider="anthropic",
kwargs={_ACTIVE_KEY: True},
)
assert should_run is False
@pytest.mark.asyncio
async def test_pre_call_strips_client_forged_marker_on_initial_request():
"""A client cannot pre-set the active marker on the original request."""
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
kwargs = {
"tools": [{"type": "web_search"}],
"custom_llm_provider": "openai",
_ACTIVE_KEY: True,
}
await logger.async_pre_call_deployment_hook(kwargs, CallTypes.aresponses)
assert _ACTIVE_KEY not in kwargs, (
"no native code_interpreter tool was present, so a client-supplied "
"active marker must be cleared"
)
@pytest.mark.asyncio
async def test_pre_call_preserves_marker_on_server_followup():
"""On a server-driven followup (depth>0) the marker is trusted and kept."""
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
kwargs = {
"tools": [{"type": "function", "name": LITELLM_CODE_EXECUTION_TOOL_NAME}],
"custom_llm_provider": "openai",
"_agentic_loop_depth": 1,
_ACTIVE_KEY: True,
}
await logger.async_pre_call_deployment_hook(kwargs, CallTypes.aresponses)
assert kwargs.get(_ACTIVE_KEY) is True, (
"the server-set marker must survive followup requests so multi-round "
"code execution keeps working"
)
@pytest.mark.asyncio
async def test_sandbox_deleted_after_loop_completes():
sandbox = FakeSandbox(stdout="42")
logger = CodeInterpreterInterceptionLogger(sandbox_config=sandbox)
plan = await _build_plan(logger, sandbox, call_id="k1")
assert sandbox.create_calls, "sandbox must be created during the loop"
assert (
not sandbox.delete_calls
), "sandbox must outlive the loop until the final hook"
await logger.async_post_agentic_loop_response_hook(
response=FakeResponse(output=[{"type": "message", "content": []}]),
plan=plan,
kwargs={},
)
assert len(sandbox.delete_calls) == 1, (
"the sandbox must be deleted once the final response is assembled, "
"otherwise it keeps running and billing"
)
assert "sbxkey1" not in logger._container_cache
@pytest.mark.asyncio
async def test_post_hook_delete_is_idempotent_across_loop_levels():
sandbox = FakeSandbox(stdout="42")
logger = CodeInterpreterInterceptionLogger(sandbox_config=sandbox)
plan = await _build_plan(logger, sandbox, call_id="k1")
response = FakeResponse(output=[{"type": "message", "content": []}])
await logger.async_post_agentic_loop_response_hook(
response=response, plan=plan, kwargs={}
)
await logger.async_post_agentic_loop_response_hook(
response=response, plan=plan, kwargs={}
)
assert len(sandbox.delete_calls) == 1, (
"deleting an already-removed container must be a no-op so unwinding "
"loop levels do not double-delete"
)
@pytest.mark.asyncio
async def test_build_plan_deletes_sandbox_when_execution_raises():
"""If sandbox execution raises before a plan is built (e.g. E2B aborts
output over its cap), the cached sandbox must be deleted before re-raising,
otherwise a caller can leak paid containers until the prune TTL."""
class RaisingSandbox(FakeSandbox):
async def arun_code(self, *, container, code, **kwargs):
raise ValueError("output exceeded cap")
sandbox = RaisingSandbox()
logger = CodeInterpreterInterceptionLogger(sandbox_config=sandbox)
with pytest.raises(ValueError, match="exceeded cap"):
await _build_plan(logger, sandbox, call_id="k1")
assert len(sandbox.create_calls) == 1, "the sandbox must have been created"
assert len(sandbox.delete_calls) == 1, (
"a build failure must delete the cached sandbox so it does not keep "
"running and billing"
)
assert "sbxkey1" not in logger._container_cache
@pytest.mark.asyncio
async def test_cleanup_hook_deletes_sandbox():
sandbox = FakeSandbox(stdout="42")
logger = CodeInterpreterInterceptionLogger(sandbox_config=sandbox)
plan = await _build_plan(logger, sandbox, call_id="k1")
await logger.async_agentic_loop_cleanup_hook(plan=plan, kwargs={})
assert len(sandbox.delete_calls) == 1, (
"the cleanup hook must delete the sandbox so a rerun failure cannot "
"leak a running container"
)
assert "sbxkey1" not in logger._container_cache
@pytest.mark.asyncio
async def test_cleanup_hook_is_idempotent_with_post_hook():
sandbox = FakeSandbox(stdout="42")
logger = CodeInterpreterInterceptionLogger(sandbox_config=sandbox)
plan = await _build_plan(logger, sandbox, call_id="k1")
await logger.async_post_agentic_loop_response_hook(
response=FakeResponse(output=[{"type": "message", "content": []}]),
plan=plan,
kwargs={},
)
await logger.async_agentic_loop_cleanup_hook(plan=plan, kwargs={})
assert len(sandbox.delete_calls) == 1, (
"cleanup running in finally after the success-path post hook already "
"deleted the sandbox must not double-delete"
)
@pytest.mark.asyncio
async def test_responses_plan_cleans_up_sandbox_when_followup_raises():
"""If the agentic rerun fails, _execute_responses_agentic_plan must still
invoke the cleanup hook so the sandbox is not left running."""
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.types.integrations.custom_logger import (
AgenticLoopPlan,
AgenticLoopRequestPatch,
)
cleanup_calls = []
class CleanupCallback(CustomLogger):
async def async_post_agentic_loop_response_hook(self, response, plan, kwargs):
return response
async def async_agentic_loop_cleanup_hook(self, plan, kwargs):
cleanup_calls.append(plan)
plan = AgenticLoopPlan(
run_agentic_loop=True,
request_patch=AgenticLoopRequestPatch(
model="gpt-5", messages=[{"role": "user", "content": "x"}]
),
metadata={"sandbox_key": "sbxkey1"},
)
original = litellm.aresponses
async def _boom(*args, **kwargs):
raise RuntimeError("upstream blew up")
litellm.aresponses = _boom
try:
with pytest.raises(RuntimeError, match="upstream blew up"):
await BaseLLMHTTPHandler()._execute_responses_agentic_plan(
plan=plan,
model="gpt-5",
response_api_optional_request_params={},
logging_obj=FakeLogging(litellm_call_id="k1"),
kwargs={},
depth=0,
max_loops=3,
fingerprints=[],
fingerprint="fp",
callback=CleanupCallback(),
)
finally:
litellm.aresponses = original
assert cleanup_calls == [plan], (
"cleanup hook must run in finally even when the rerun raises, otherwise "
"the sandbox keeps running until the prune TTL"
)
@pytest.mark.asyncio
async def test_run_code_does_not_re_resolve_registry(monkeypatch):
"""Params resolved once at create time must be reused for running code, so a
registry clear between create and run cannot turn into a create-then-fail."""
import litellm
from litellm.sandbox import sandbox_tools
sandbox_tools.register_sandbox_tools(
[
{
"sandbox_tool_name": "e2b_default",
"litellm_params": {"sandbox_provider": "e2b", "api_key": "sk-x"},
}
]
)
create_kwargs = {}
run_kwargs = {}
async def fake_acreate_sandbox(**kwargs):
create_kwargs.update(kwargs)
return FakeHandle()
async def fake_arun_code(**kwargs):
run_kwargs.update(kwargs)
return CodeExecutionResult(stdout="ok")
monkeypatch.setattr(litellm, "acreate_sandbox", fake_acreate_sandbox)
monkeypatch.setattr(litellm, "arun_code", fake_arun_code)
logger = CodeInterpreterInterceptionLogger(sandbox_tool_name="e2b_default")
try:
container, params = await logger._get_or_create_container(cache_key="k1")
assert params is not None and params["sandbox_provider"] == "e2b"
sandbox_tools.clear_sandbox_tools()
stdout = await logger._run_tool_call(
container=container, params=params, arguments='{"code":"print(1)"}'
)
finally:
sandbox_tools.clear_sandbox_tools()
assert stdout == "ok", "run must succeed using the params captured at create time"
assert run_kwargs["provider"] == "e2b"
@pytest.mark.asyncio
async def test_run_tool_call_surfaces_execution_error():
"""A sandbox execution error must be fed back to the model as a labelled
string, not raised, so the agentic loop can react to it."""
class ErroringSandbox(FakeSandbox):
async def arun_code(self, *, container, code, **kwargs):
self.run_calls.append({"container": container, "code": code})
return CodeExecutionResult(
stdout="", error={"name": "ValueError", "value": "boom"}
)
sandbox = ErroringSandbox()
logger = CodeInterpreterInterceptionLogger(sandbox_config=sandbox)
container = await logger._create_container()
stdout = await logger._run_tool_call(
container=container[0], params=None, arguments='{"code":"raise ValueError(1)"}'
)
assert stdout == "[execution error] boom"
@pytest.mark.asyncio
async def test_run_tool_call_reports_unparseable_arguments():
"""Malformed tool arguments must produce a parse error string the model can
see rather than crashing the interceptor."""
sandbox = FakeSandbox()
logger = CodeInterpreterInterceptionLogger(sandbox_config=sandbox)
container = await logger._create_container()
stdout = await logger._run_tool_call(
container=container[0], params=None, arguments="not-json"
)
assert stdout == "[invalid tool arguments: could not parse code]"
assert not sandbox.run_calls, "code must not run when arguments cannot be parsed"
@pytest.mark.asyncio
async def test_pre_call_skips_provider_outside_scope():
"""enabled_providers must filter the pre-call conversion so a request to an
out-of-scope provider is left untouched."""
logger = CodeInterpreterInterceptionLogger(
sandbox_config=FakeSandbox(), enabled_providers=["openai"]
)
kwargs = {
"tools": [{"type": "code_interpreter", "container": {"type": "auto"}}],
"custom_llm_provider": "anthropic",
}
result = await logger.async_pre_call_deployment_hook(kwargs, CallTypes.aresponses)
assert result is None
assert kwargs["tools"][0]["type"] == "code_interpreter", "tool must be untouched"
assert _ACTIVE_KEY not in kwargs
@pytest.mark.asyncio
async def test_resolve_provider_falls_back_to_model_lookup():
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
assert logger._resolve_provider({"custom_llm_provider": "openai"}) == "openai"
assert logger._resolve_provider({"model": "gpt-5"}) == "openai"
assert logger._resolve_provider({"model": 123}) is None
assert logger._resolve_provider({"model": "no-such-provider-xyz"}) is None
@pytest.mark.asyncio
async def test_create_container_without_sandbox_raises():
"""The registry path must raise a clear error when no sandbox is resolvable
instead of silently creating nothing."""
logger = CodeInterpreterInterceptionLogger(sandbox_tool_name="missing")
with pytest.raises(ValueError, match="no sandbox available"):
await logger._create_container()
@pytest.mark.asyncio
async def test_run_code_without_params_raises():
logger = CodeInterpreterInterceptionLogger(sandbox_tool_name="missing")
with pytest.raises(ValueError, match="no sandbox available to run code"):
await logger._run_code(container=FakeHandle(), params=None, code="print(1)")
@pytest.mark.asyncio
async def test_delete_container_swallows_errors():
"""A delete failure must not propagate; the request already succeeded."""
class FailingDeleteSandbox(FakeSandbox):
async def adelete_sandbox(self, *, container, **kwargs):
raise RuntimeError("e2b unreachable")
sandbox = FailingDeleteSandbox()
logger = CodeInterpreterInterceptionLogger(sandbox_config=sandbox)
container, params = await logger._create_container()
await logger._delete_container(container=container, params=params)
@pytest.mark.asyncio
async def test_prune_expired_cache_deletes_underlying_container():
"""Expired cache entries must have their sandbox deleted, not just dropped,
otherwise an orphaned sandbox keeps running."""
import litellm.integrations.code_interpreter_interception.handler as handler_mod
sandbox = FakeSandbox()
logger = CodeInterpreterInterceptionLogger(sandbox_config=sandbox)
container, params = await logger._create_container()
logger._container_cache["old"] = (
container,
params,
time.time() - handler_mod._CACHE_TTL_SECONDS - 1,
)
await logger._prune_expired_cache()
assert "old" not in logger._container_cache
assert len(sandbox.delete_calls) == 1, "expired sandbox must be deleted"
@pytest.mark.asyncio
async def test_normalize_messages_handles_str_and_unknown():
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
assert logger._normalize_messages("hi") == [{"role": "user", "content": "hi"}]
assert logger._normalize_messages([{"role": "user"}]) == [{"role": "user"}]
assert logger._normalize_messages(42) == []
def test_from_config_yaml_reads_fields():
cfg = {
"enabled": False,
"enabled_providers": ["openai"],
"sandbox_tool_name": "e2b_default",
}
logger = CodeInterpreterInterceptionLogger.from_config_yaml(cfg)
assert logger.enabled is False
assert logger.enabled_providers == ["openai"]
assert logger.sandbox_tool_name == "e2b_default"
def test_initialize_from_proxy_config_prefers_litellm_settings():
logger = CodeInterpreterInterceptionLogger.initialize_from_proxy_config(
litellm_settings={
"code_interpreter_interception_params": {
"enabled_providers": ["openai"],
"sandbox_tool_name": "e2b_default",
}
},
callback_specific_params={},
)
assert logger.enabled_providers == ["openai"]
assert logger.sandbox_tool_name == "e2b_default"
@pytest.mark.asyncio
async def test_build_plan_handles_dict_shaped_response():
"""A responses payload delivered as a plain dict (not an object) must flow
through detection, execution, and re-injection the same as the typed form."""
sandbox = FakeSandbox(stdout="42")
logger = CodeInterpreterInterceptionLogger(sandbox_config=sandbox)
dict_response = {"output": [_function_call_item()]}
plan = await logger.async_build_agentic_loop_plan(
tools={"tool_calls": logger._extract_code_execution_tool_calls(dict_response)},
model="gpt-5",
messages=[{"role": "user", "content": "x"}],
response=dict_response,
anthropic_messages_provider_config=None,
anthropic_messages_optional_request_params={"tools": []},
logging_obj=FakeLogging(litellm_call_id="k1"),
stream=False,
kwargs={"litellm_call_id": "k1", _SANDBOX_KEY: "sbxkey1"},
)
assert sandbox.run_calls, "code must run for a dict-shaped response"
assert plan.metadata["code_interpreter_calls"][0]["code"] == "print(40 + 2)"
out = await logger.async_post_agentic_loop_response_hook(
response={"output": [{"type": "message", "content": []}]},
plan=plan,
kwargs={},
)
assert [item.get("type") for item in out["output"]] == [
"code_interpreter_call",
"message",
], "the dict-shaped response must get the code_interpreter_call re-injected"
@pytest.mark.asyncio
async def test_extract_tool_calls_reads_object_attributes():
"""Detection must work when output items are objects with attributes, not
only dicts."""
class Item:
def __init__(self):
self.type = "function_call"
self.name = LITELLM_CODE_EXECUTION_TOOL_NAME
self.call_id = "c9"
self.arguments = '{"code":"print(1)"}'
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
calls = logger._extract_code_execution_tool_calls(FakeResponse(output=[Item()]))
assert len(calls) == 1
assert calls[0]["call_id"] == "c9"
assert calls[0]["arguments"] == '{"code":"print(1)"}'

View file

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

View file

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

View file

@ -153,18 +153,22 @@ class TestResponseCompliance:
def test_interaction_response_fields(self, spec_dict):
"""Verify our InteractionsAPIResponse has correct fields."""
# The response is the Interaction schema
# Check CreateModelInteractionParams which includes output fields
schema = spec_dict["components"]["schemas"]["CreateModelInteractionParams"]
# The response is the dedicated `Interaction` schema. Google moved the
# output-only fields (notably the `steps` array, formerly `outputs`)
# off `CreateModelInteractionParams` and onto `Interaction`; the request
# schema no longer carries `steps`. Google later moved `role` off
# `Interaction` onto the per-turn `Turn` schema (asserted in
# test_turn_schema), so it is no longer a top-level output field here.
# Keep this aligned with the live spec.
schema = spec_dict["components"]["schemas"]["Interaction"]
# Output fields (readOnly)
# Output fields (readOnly).
output_fields = [
"id",
"status",
"created",
"updated",
"role",
"outputs",
"steps",
"usage",
]
@ -174,10 +178,13 @@ class TestResponseCompliance:
def test_status_enum_values(self, spec_dict):
"""Verify status enum values match spec."""
schema = spec_dict["components"]["schemas"]["CreateModelInteractionParams"]
# `status` is an output-only field; validate against the response schema.
schema = spec_dict["components"]["schemas"]["Interaction"]
status_prop = schema["properties"]["status"]
# Google Interactions API uses lowercase status values (updated Feb 2026;
# "budget_exceeded" added to the published spec mid-2026).
# Google Interactions API uses lowercase status values (updated Feb 2026).
# Keep this an exact match: this test intentionally breaks CI when
# Google changes the live spec — that breakage is how we get notified
# to review the change.
expected_statuses = [
"in_progress",
"requires_action",

View file

@ -1,6 +1,7 @@
import os
import sys
import httpx
import pytest
import litellm
@ -550,3 +551,78 @@ class TestExtractAndRaiseLitellmException:
)
assert result is None
class ModelError(Exception):
"""Mimics replicate's SDK exception, whose mapping keys on the class name."""
class CohereConnectionError(Exception):
"""Mimics cohere's SDK exception, whose mapping keys on the class name."""
def test_replicate_model_error_maps_to_bad_request():
"""The replicate branch keys on ``type(original_exception).__name__ ==
"ModelError"`` rather than on the error string. The dispatch now lives in a
per-provider helper, so this class-name value has to be threaded into the
helper; if it is not, the bare name ``exception_type`` resolves to the
module-level function and the comparison is always False, silently mismapping
to APIConnectionError."""
original_exception = ModelError("the deployed model failed to return a prediction")
with pytest.raises(litellm.BadRequestError) as excinfo:
exception_type(
model="replicate/meta/llama-2-70b-chat",
original_exception=original_exception,
custom_llm_provider="replicate",
)
assert excinfo.value.llm_provider == "replicate"
def test_cohere_connection_error_maps_to_rate_limit():
"""The cohere branch keys on ``"CohereConnectionError" in
type(original_exception).__name__``. With the dispatch extracted into a helper
the class-name value must be passed in; otherwise ``in`` runs against the
module-level ``exception_type`` function object and raises TypeError, which the
outer handler swallows into a generic APIConnectionError."""
original_exception = CohereConnectionError("connection reset by peer")
original_exception.message = "connection reset by peer"
with pytest.raises(litellm.RateLimitError) as excinfo:
exception_type(
model="command-r",
original_exception=original_exception,
custom_llm_provider="cohere",
)
assert excinfo.value.llm_provider == "cohere"
class ReplicateError(Exception):
"""Mimics a replicate HTTP error carrying a status_code and response."""
def __init__(self, message: str, status_code: int) -> None:
super().__init__(message)
self.message = message
self.status_code = status_code
self.response = httpx.Response(
status_code=status_code,
request=httpx.Request("POST", "https://api.replicate.com/v1/predictions"),
)
def test_replicate_422_maps_to_unprocessable_entity():
"""The replicate status-code ladder carried two identical ``status_code == 422``
branches; the second was unreachable dead code. After dropping the duplicate the
surviving branch must still map 422 to UnprocessableEntityError."""
original_exception = ReplicateError("validation failed for the input", 422)
with pytest.raises(litellm.UnprocessableEntityError) as excinfo:
exception_type(
model="replicate/meta/llama-2-70b-chat",
original_exception=original_exception,
custom_llm_provider="replicate",
)
assert excinfo.value.llm_provider == "replicate"

View file

@ -18,6 +18,8 @@ from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.streaming_handler import (
AUDIO_ATTRIBUTE,
CustomStreamWrapper,
_ProviderChunkEarlyReturn,
_ProviderChunkParsed,
)
from litellm.types.utils import (
CompletionTokensDetailsWrapper,
@ -1801,6 +1803,344 @@ def test_usage_only_chunk_not_dropped_when_finish_reason_already_set(
assert result.usage is not None
def _run_dispatch(wrapper: CustomStreamWrapper, chunk):
model_response = wrapper.model_response_creator()
completion_obj = {"content": ""}
result = wrapper._dispatch_provider_chunk(
chunk=chunk,
model_response=model_response,
completion_obj=completion_obj,
)
return result, model_response, completion_obj
def test_dispatch_vllm_extracts_output_text(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
"""vllm chunks expose text at chunk[0].outputs[0].text; the dispatch must
surface that as the content and report a parsed result."""
initialized_custom_stream_wrapper.custom_llm_provider = "vllm"
class _Output:
text = "hello from vllm"
class _VLLMChunk:
outputs = [_Output()]
result, _, completion_obj = _run_dispatch(
initialized_custom_stream_wrapper, [_VLLMChunk()]
)
assert isinstance(result, _ProviderChunkParsed)
assert completion_obj["content"] == "hello from vllm"
def test_dispatch_petals_slices_completion_stream(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
"""petals fakes streaming by slicing 30 chars off the buffered completion
stream each call, leaving the remainder for the next chunk."""
initialized_custom_stream_wrapper.custom_llm_provider = "petals"
initialized_custom_stream_wrapper.completion_stream = "A" * 50
result, _, completion_obj = _run_dispatch(
initialized_custom_stream_wrapper, chunk=None
)
assert isinstance(result, _ProviderChunkParsed)
assert completion_obj["content"] == "A" * 30
assert initialized_custom_stream_wrapper.completion_stream == "A" * 20
assert initialized_custom_stream_wrapper.received_finish_reason is None
def test_dispatch_petals_empty_stream_sets_stop(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
"""An exhausted petals stream marks the turn finished with a stop reason."""
initialized_custom_stream_wrapper.custom_llm_provider = "petals"
initialized_custom_stream_wrapper.completion_stream = ""
result, _, completion_obj = _run_dispatch(
initialized_custom_stream_wrapper, chunk=None
)
assert isinstance(result, _ProviderChunkParsed)
assert completion_obj["content"] == ""
assert initialized_custom_stream_wrapper.received_finish_reason == "stop"
def test_dispatch_petals_empty_stream_after_finish_raises(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
"""Once petals has already finished, an empty stream signals end-of-iteration."""
initialized_custom_stream_wrapper.custom_llm_provider = "petals"
initialized_custom_stream_wrapper.completion_stream = ""
initialized_custom_stream_wrapper.received_finish_reason = "stop"
with pytest.raises(StopIteration):
_run_dispatch(initialized_custom_stream_wrapper, chunk=None)
def test_dispatch_palm_slices_completion_stream(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
"""palm uses the same fake-streaming slice strategy as petals."""
initialized_custom_stream_wrapper.custom_llm_provider = "palm"
initialized_custom_stream_wrapper.completion_stream = "B" * 40
result, _, completion_obj = _run_dispatch(
initialized_custom_stream_wrapper, chunk=None
)
assert isinstance(result, _ProviderChunkParsed)
assert completion_obj["content"] == "B" * 30
assert initialized_custom_stream_wrapper.completion_stream == "B" * 10
def test_dispatch_cached_response_extracts_delta(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
"""cached_response replays a stored ModelResponseStream; the dispatch lifts
its delta content, finish_reason and id back onto the live response."""
initialized_custom_stream_wrapper.custom_llm_provider = "cached_response"
chunk = ModelResponseStream(
id="chatcmpl-cache-1",
choices=[
StreamingChoices(
index=0,
delta=Delta(content="cached text"),
finish_reason="stop",
)
],
)
result, model_response, completion_obj = _run_dispatch(
initialized_custom_stream_wrapper, chunk
)
assert isinstance(result, _ProviderChunkParsed)
assert completion_obj["content"] == "cached text"
assert initialized_custom_stream_wrapper.received_finish_reason == "stop"
assert model_response.id == "chatcmpl-cache-1"
assert initialized_custom_stream_wrapper.response_id == "chatcmpl-cache-1"
def test_dispatch_vertex_ai_legacy_text_and_finish_reason(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
"""Legacy vertex_ai chunks (non-ModelResponseStream) expose .text and a
candidate finish_reason enum that must be normalised to an OpenAI reason."""
initialized_custom_stream_wrapper.custom_llm_provider = "vertex_ai"
class _FinishReason:
name = "STOP"
class _Candidate:
finish_reason = _FinishReason()
class _VertexChunk:
candidates = [_Candidate()]
text = "vertex content"
result, _, completion_obj = _run_dispatch(
initialized_custom_stream_wrapper, _VertexChunk()
)
assert isinstance(result, _ProviderChunkParsed)
assert completion_obj["content"] == "vertex content"
assert initialized_custom_stream_wrapper.received_finish_reason == "stop"
def test_dispatch_vertex_ai_legacy_without_candidates_stringifies_chunk(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
"""A legacy vertex_ai chunk with no candidates falls back to str(chunk)."""
initialized_custom_stream_wrapper.custom_llm_provider = "vertex_ai"
class _RawChunk:
def __str__(self) -> str:
return "raw vertex blob"
result, _, completion_obj = _run_dispatch(
initialized_custom_stream_wrapper, _RawChunk()
)
assert isinstance(result, _ProviderChunkParsed)
assert completion_obj["content"] == "raw vertex blob"
def test_dispatch_vertex_ai_legacy_function_call(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
"""A legacy vertex_ai chunk whose part has no text but carries a
function_call is converted into an OpenAI tool-call delta."""
initialized_custom_stream_wrapper.custom_llm_provider = "vertex_ai"
class _FunctionCall:
name = "get_weather"
args = {"location": "SF"}
class _Part:
function_call = _FunctionCall()
class _Content:
parts = [_Part()]
class _FinishReason:
name = "STOP"
class _Candidate:
content = _Content()
finish_reason = _FinishReason()
class _VertexFunctionChunk:
candidates = [_Candidate()]
@property
def text(self):
raise RuntimeError("Part has no text.")
result, _, _ = _run_dispatch(
initialized_custom_stream_wrapper, _VertexFunctionChunk()
)
assert isinstance(result, _ProviderChunkParsed)
tool_calls = result.response_obj["original_chunk"].choices[0].delta.tool_calls
assert tool_calls[0].function.name == "get_weather"
assert json.loads(tool_calls[0].function.arguments) == {"location": "SF"}
assert initialized_custom_stream_wrapper.received_finish_reason == "stop"
def test_dispatch_custom_provider_returns_chunk_early(
monkeypatch: pytest.MonkeyPatch,
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
"""A registered custom provider passes its already-OpenAI-shaped chunk
straight through as an early return rather than re-parsing it."""
monkeypatch.setattr(litellm, "_custom_providers", ["my-custom-llm"])
initialized_custom_stream_wrapper.custom_llm_provider = "my-custom-llm"
chunk = ModelResponseStream(
choices=[
StreamingChoices(index=0, delta=Delta(content="hi"), finish_reason=None)
]
)
result, _, _ = _run_dispatch(initialized_custom_stream_wrapper, chunk)
assert isinstance(result, _ProviderChunkEarlyReturn)
assert result.value is chunk
def test_dispatch_custom_provider_finish_only_returns_none_early(
monkeypatch: pytest.MonkeyPatch,
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
"""A custom-provider chunk that carries only a finish_reason (no content)
records the reason and returns None so no empty delta is emitted."""
monkeypatch.setattr(litellm, "_custom_providers", ["my-custom-llm"])
initialized_custom_stream_wrapper.custom_llm_provider = "my-custom-llm"
chunk = ModelResponseStream(
choices=[
StreamingChoices(index=0, delta=Delta(content=None), finish_reason="stop")
]
)
result, _, _ = _run_dispatch(initialized_custom_stream_wrapper, chunk)
assert isinstance(result, _ProviderChunkEarlyReturn)
assert result.value is None
assert initialized_custom_stream_wrapper.received_finish_reason == "stop"
def test_dispatch_text_completion_codestral_parses_chunk(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
"""text-completion-codestral streams raw SSE JSON strings that the dispatch
routes through CodestralTextCompletionConfig to extract content/finish."""
initialized_custom_stream_wrapper.custom_llm_provider = "text-completion-codestral"
chunk = json.dumps(
{"choices": [{"delta": {"content": "codestral text"}, "finish_reason": "stop"}]}
)
result, _, completion_obj = _run_dispatch(initialized_custom_stream_wrapper, chunk)
assert isinstance(result, _ProviderChunkParsed)
assert completion_obj["content"] == "codestral text"
assert initialized_custom_stream_wrapper.received_finish_reason == "stop"
def test_dispatch_text_completion_codestral_requires_string(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
"""The codestral branch only knows how to parse raw strings; anything else
is a programming error and must surface loudly."""
initialized_custom_stream_wrapper.custom_llm_provider = "text-completion-codestral"
with pytest.raises(ValueError):
_run_dispatch(initialized_custom_stream_wrapper, {"not": "a string"})
def test_dispatch_triton_stream(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
"""triton stream chunks arrive as dicts keyed by text_output/stop_reason."""
initialized_custom_stream_wrapper.custom_llm_provider = "triton"
chunk = {"text_output": "triton text", "is_finished": True, "stop_reason": "stop"}
result, _, completion_obj = _run_dispatch(initialized_custom_stream_wrapper, chunk)
assert isinstance(result, _ProviderChunkParsed)
assert completion_obj["content"] == "triton text"
assert initialized_custom_stream_wrapper.received_finish_reason == "stop"
def test_dispatch_ai21_decodes_completion(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
"""ai21 does fake streaming over a single byte-encoded JSON completion."""
initialized_custom_stream_wrapper.custom_llm_provider = "ai21"
chunk = json.dumps({"completions": [{"data": {"text": "ai21 text"}}]}).encode(
"utf-8"
)
result, _, completion_obj = _run_dispatch(initialized_custom_stream_wrapper, chunk)
assert isinstance(result, _ProviderChunkParsed)
assert completion_obj["content"] == "ai21 text"
assert initialized_custom_stream_wrapper.received_finish_reason == "stop"
def test_dispatch_text_completion_openai_with_usage(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
"""text-completion-openai chunks expose choices[].text and an optional usage
block that the dispatch lifts onto the model response."""
initialized_custom_stream_wrapper.custom_llm_provider = "text-completion-openai"
class _Choice:
text = "oai text"
finish_reason = "stop"
class _Usage:
prompt_tokens = 5
completion_tokens = 3
total_tokens = 8
class _TextChunk:
choices = [_Choice()]
usage = _Usage()
result, model_response, completion_obj = _run_dispatch(
initialized_custom_stream_wrapper, _TextChunk()
)
assert isinstance(result, _ProviderChunkParsed)
assert completion_obj["content"] == "oai text"
assert initialized_custom_stream_wrapper.received_finish_reason == "stop"
assert model_response.usage.prompt_tokens == 5
assert model_response.usage.total_tokens == 8
@pytest.mark.asyncio
async def test_custom_stream_wrapper_anext_does_not_block_event_loop_for_sync_iterators(
logging_obj: Logging,
@ -2306,7 +2646,9 @@ def test_chunk_creator_tool_calls_not_dropped_on_finish(
tool_calls=[
ChatCompletionDeltaToolCall(
id="call_abc",
function=Function(name="get_weather", arguments='{"city":"NYC"}'),
function=Function(
name="get_weather", arguments='{"city":"NYC"}'
),
type="function",
index=0,
)
@ -2401,3 +2743,131 @@ def test_record_partial_usage_for_failure_noop_without_chunks():
wrapper._record_partial_usage_for_failure()
assert "combined_usage_object" not in logging_obj.model_call_details
@pytest.mark.parametrize("sync_mode", [True, False])
@pytest.mark.asyncio
async def test_stream_chunk_builder_raise_at_end_of_stream_still_recovers_usage(
sync_mode,
):
"""stream_chunk_builder re-raises (as APIError) on large agentic tool-use
streams. That raise originates inside the except-StopIteration handler, so
before the fix it escaped __next__/__anext__ and the request was dropped from
SpendLogs while the provider billed the tokens. The wrapper must catch it and
recover usage from the raw chunks so cost is still tracked."""
final_usage_block = Usage(
completion_tokens=392, prompt_tokens=1799, total_tokens=2191
)
final_chunk = ModelResponseStream(
id="chatcmpl-raise-test",
created=1742056047,
model=None,
object="chat.completion.chunk",
choices=[
StreamingChoices(
finish_reason="stop",
index=0,
delta=Delta(content="", role="assistant"),
)
],
usage=final_usage_block,
)
test_chunks = bedrock_chunks + [final_chunk]
logging_obj = Logging(
model="bedrock/claude-haiku-4-5-20251001-v1:0",
messages=[{"role": "user", "content": "Hey"}],
stream=True,
call_type="completion",
start_time=time.time(),
litellm_call_id="raise-test",
function_id="1245",
)
response = CustomStreamWrapper(
completion_stream=ModelResponseListIterator(model_responses=test_chunks),
model="bedrock/claude-haiku-4-5-20251001-v1:0",
custom_llm_provider="bedrock",
logging_obj=logging_obj,
stream_options={"include_usage": True},
)
seen_usage = []
with patch.object(
litellm,
"stream_chunk_builder",
side_effect=Exception("simulated assembly failure"),
):
# before the fix this raised and dropped the request; it must not raise now
if sync_mode:
for chunk in response:
if getattr(chunk, "usage", None) is not None:
seen_usage.append(chunk.usage)
else:
async for chunk in response:
if getattr(chunk, "usage", None) is not None:
seen_usage.append(chunk.usage)
assert any(
u.total_tokens == final_usage_block.total_tokens for u in seen_usage
), "usage recovered from raw chunks was not emitted after stream_chunk_builder raised"
@pytest.mark.parametrize("sync_mode", [True, False])
@pytest.mark.asyncio
async def test_stream_chunk_builder_raise_and_usage_recovery_failure_does_not_crash(
sync_mode,
):
"""If end-of-stream assembly raises AND best-effort usage recovery from the raw
chunks also fails, the stream must still complete cleanly rather than propagate
the exception to the consumer."""
from litellm.litellm_core_utils import streaming_handler as sh_module
final_chunk = ModelResponseStream(
id="chatcmpl-raise-recover-fail",
created=1742056047,
model=None,
object="chat.completion.chunk",
choices=[
StreamingChoices(
finish_reason="stop",
index=0,
delta=Delta(content="", role="assistant"),
)
],
usage=Usage(completion_tokens=1, prompt_tokens=1, total_tokens=2),
)
response = CustomStreamWrapper(
completion_stream=ModelResponseListIterator(
model_responses=bedrock_chunks + [final_chunk]
),
model="bedrock/claude-haiku-4-5-20251001-v1:0",
custom_llm_provider="bedrock",
logging_obj=Logging(
model="bedrock/claude-haiku-4-5-20251001-v1:0",
messages=[{"role": "user", "content": "Hey"}],
stream=True,
call_type="completion",
start_time=time.time(),
litellm_call_id="raise-recover-fail",
function_id="1245",
),
stream_options={"include_usage": True},
)
with (
patch.object(
litellm, "stream_chunk_builder", side_effect=Exception("assembly failed")
),
patch.object(
sh_module, "calculate_total_usage", side_effect=Exception("recovery failed")
),
):
# must not raise even though both assembly and recovery fail
if sync_mode:
chunks = [c for c in response]
else:
chunks = [c async for c in response]
assert len(chunks) > 0

View file

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

View file

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

View file

@ -1,9 +1,8 @@
import json
import os
import sys
from unittest.mock import AsyncMock, MagicMock, patch
from unittest.mock import MagicMock, patch
import httpx
import pytest
import litellm
@ -12,10 +11,14 @@ sys.path.insert(
0, os.path.abspath("../../../../..")
) # Adds the parent directory to the system path
from litellm import get_model_info, supports_reasoning
from litellm import get_model_info, supports_reasoning, supports_vision
from litellm.llms.fireworks_ai.chat.transformation import FireworksAIConfig
from litellm.types.llms.openai import ChatCompletionToolCallFunctionChunk
from litellm.types.utils import ChatCompletionMessageToolCall, Function, Message
from litellm.types.utils import (
ChatCompletionMessageToolCall,
Function,
Message,
ModelResponse,
)
@pytest.fixture(autouse=True)
@ -98,12 +101,12 @@ def test_supports_reasoning_effort():
for model in supported_models:
assert (
supports_reasoning(model=model, custom_llm_provider="fireworks_ai") == True
supports_reasoning(model=model, custom_llm_provider="fireworks_ai") is True
), f"{model} should support reasoning_effort"
for model in unsupported_models:
assert (
supports_reasoning(model=model, custom_llm_provider="fireworks_ai") == False
supports_reasoning(model=model, custom_llm_provider="fireworks_ai") is False
), f"{model} should not support reasoning_effort"
@ -115,11 +118,13 @@ def test_get_supported_openai_params_reasoning_effort():
"fireworks_ai/accounts/fireworks/models/glm-5p1"
)
assert "reasoning_effort" in supported_params
assert "thinking" in supported_params
unsupported_params = config.get_supported_openai_params(
"fireworks_ai/accounts/fireworks/models/llama-v3-70b-instruct"
)
assert "reasoning_effort" not in unsupported_params
assert "thinking" not in unsupported_params
def test_get_supported_openai_params_parallel_tool_calls():
@ -181,41 +186,6 @@ def test_get_provider_info_omits_false_supports_reasoning(monkeypatch):
assert "supports_reasoning" not in info
def test_add_transform_inline_image_block_skips_data_urls():
"""
data: URLs must not have #transform=inline appended — doing so corrupts the
base64 payload and raises binascii.Error: Incorrect padding on the Fireworks side.
Regression test for https://github.com/BerriAI/litellm/issues/23583
"""
config = FireworksAIConfig()
data_url = "data:image/jpeg;base64,/9j/4AAQSkZJRgAB"
# str branch
str_content = {"type": "image_url", "image_url": data_url}
result = config._add_transform_inline_image_block(
str_content, model="gpt-4", disable_add_transform_inline_image_block=False
)
assert result["image_url"] == data_url, "data URL must not be modified (str branch)"
# dict branch
dict_content = {"type": "image_url", "image_url": {"url": data_url}}
result = config._add_transform_inline_image_block(
dict_content, model="gpt-4", disable_add_transform_inline_image_block=False
)
assert (
result["image_url"]["url"] == data_url
), "data URL must not be modified (dict branch)"
# regular https URL should still get the suffix
https_content = {"type": "image_url", "image_url": "https://example.com/image.jpg"}
result = config._add_transform_inline_image_block(
https_content, model="gpt-4", disable_add_transform_inline_image_block=False
)
assert result["image_url"].endswith(
"#transform=inline"
), "https URL should get #transform=inline"
@pytest.mark.parametrize(
"api_base, expected_url_prefix",
[
@ -582,3 +552,549 @@ def test_transform_request_routes_short_form_model_to_models_path():
headers={},
)
assert result["model"] == "accounts/fireworks/models/glm-5p2"
def _make_fireworks_raw_response(body: dict) -> MagicMock:
mock = MagicMock()
mock.status_code = 200
mock.json.return_value = body
mock.text = json.dumps(body)
mock.headers = {}
return mock
_BASE_CHAT_COMPLETION_RESPONSE: dict = {
"id": "resp-test",
"object": "chat.completion",
"created": 1234567890,
"model": "glm-5p1",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "Hello"},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
}
def _run_transform_response(response_body: dict) -> ModelResponse:
config = FireworksAIConfig()
raw_response = _make_fireworks_raw_response(response_body)
logging_obj = MagicMock()
return config.transform_response(
model="accounts/fireworks/models/glm-5p1",
raw_response=raw_response,
model_response=ModelResponse(),
logging_obj=logging_obj,
request_data={},
messages=[{"role": "user", "content": "Hi"}],
optional_params={},
litellm_params={},
encoding=None,
api_key="test-key",
)
_REASONING_MODEL = "fireworks_ai/accounts/fireworks/models/glm-5p1"
_NON_REASONING_MODEL = "fireworks_ai/accounts/fireworks/models/llama-v3-70b-instruct"
def test_get_supported_openai_params_includes_all_fireworks_params():
config = FireworksAIConfig()
params = config.get_supported_openai_params(_REASONING_MODEL)
required = [
"seed",
"top_logprobs",
"min_p",
"typical_p",
"repetition_penalty",
"mirostat_target",
"mirostat_lr",
"logit_bias",
"echo",
"echo_last",
"ignore_eos",
"prompt_cache_key",
"prompt_cache_isolation_key",
"raw_output",
"perf_metrics_in_response",
"return_token_ids",
"safe_tokenization",
"service_tier",
"speculation",
"prediction",
"stream_options",
"sampling_mask",
"thinking",
"reasoning_history",
]
missing = [p for p in required if p not in params]
assert missing == [], f"Missing params: {missing}"
def test_native_openai_params_flow_end_to_end_with_drop_params_false():
"""
The OpenAI-native params Fireworks supports (seed, top_logprobs, logit_bias,
prompt_cache_key, service_tier, prediction) previously hit
``UnsupportedParamsError`` with ``drop_params=False`` because they were absent
from ``get_supported_openai_params``. Listing them must let them survive the
``get_optional_params`` gate and reach the request, not just appear in the
supported list. Asserting via ``get_optional_params`` (the real gate) rather
than ``map_openai_params`` catches a revert of the supported-params additions,
which a list-membership check would not.
"""
native = {
"seed": 42,
"top_logprobs": 3,
"logit_bias": {"1": 1},
"prompt_cache_key": "cache-key",
"service_tier": "auto",
"prediction": {"type": "content", "content": "x"},
}
optional_params = litellm.get_optional_params(
model="accounts/fireworks/models/llama-v3-70b-instruct",
custom_llm_provider="fireworks_ai",
drop_params=False,
**native,
)
for key, value in native.items():
assert optional_params.get(key) == value
def test_prompt_truncate_len_correct_name():
config = FireworksAIConfig()
params = config.get_supported_openai_params(_REASONING_MODEL)
assert "prompt_truncate_len" in params
assert "prompt_truncate_length" not in params
result = config.map_openai_params(
{"prompt_truncate_len": 4096},
{},
_REASONING_MODEL,
drop_params=False,
)
assert result == {"prompt_truncate_len": 4096}
def test_stream_options_include_usage_auto_injected():
config = FireworksAIConfig()
result = config.transform_request(
model="accounts/fireworks/models/glm-5p1",
messages=[{"role": "user", "content": "Hi"}],
optional_params={"stream": True},
litellm_params={},
headers={},
)
assert result["stream_options"] == {"include_usage": True}
def test_stream_options_not_injected_when_not_streaming():
config = FireworksAIConfig()
result = config.transform_request(
model="accounts/fireworks/models/glm-5p1",
messages=[{"role": "user", "content": "Hi"}],
optional_params={},
litellm_params={},
headers={},
)
assert "stream_options" not in result
def test_stream_options_preserves_user_override():
config = FireworksAIConfig()
result = config.transform_request(
model="accounts/fireworks/models/glm-5p1",
messages=[{"role": "user", "content": "Hi"}],
optional_params={"stream": True, "stream_options": {"include_usage": False}},
litellm_params={},
headers={},
)
assert result["stream_options"]["include_usage"] is False
def test_reasoning_history_in_supported_params():
config = FireworksAIConfig()
reasoning_params = config.get_supported_openai_params(_REASONING_MODEL)
assert "reasoning_history" in reasoning_params
non_reasoning_params = config.get_supported_openai_params(_NON_REASONING_MODEL)
assert "reasoning_history" not in non_reasoning_params
def test_thinking_param_passthrough():
config = FireworksAIConfig()
thinking = {"type": "disabled"}
result = config.map_openai_params(
{"thinking": thinking},
{},
_REASONING_MODEL,
drop_params=False,
)
assert result == {"thinking": thinking}
def test_thinking_and_reasoning_effort_conflict_rejected():
config = FireworksAIConfig()
with pytest.raises(
litellm.BadRequestError,
match="does not support specifying both `thinking` and `reasoning_effort`",
):
config.map_openai_params(
{
"thinking": {"type": "enabled", "budget_tokens": 4096},
"reasoning_effort": "medium",
},
{},
_REASONING_MODEL,
drop_params=False,
)
def test_minimax_m3_supports_vision_from_model_map():
config = FireworksAIConfig()
for model in [
"fireworks_ai/accounts/fireworks/models/minimax-m3",
"fireworks_ai/minimax-m3",
]:
assert supports_vision(model=model, custom_llm_provider="fireworks_ai") is True
assert config.get_provider_info(model)["supports_vision"] is True
def test_transform_messages_helper_rejects_file_blocks():
config = FireworksAIConfig()
messages = [
{
"role": "user",
"content": [
{
"type": "file",
"file": {
"file_data": "data:application/pdf;base64,JVBERi0xLjQKJSVFT0YK",
"filename": "tiny.pdf",
},
},
{"type": "text", "text": "Describe this"},
],
}
]
with pytest.raises(
litellm.BadRequestError,
match="Fireworks AI chat completions does not support file content blocks",
):
config._transform_messages_helper(
messages, model="accounts/fireworks/models/kimi-k2p6", litellm_params={}
)
def test_transform_messages_helper_rejects_non_vision_image_inputs():
config = FireworksAIConfig()
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "Describe this"},
{
"type": "image_url",
"image_url": {
"url": "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAE="
},
},
],
}
]
with pytest.raises(litellm.BadRequestError, match="does not support image inputs"):
config._transform_messages_helper(
messages, model="accounts/fireworks/models/glm-5p2", litellm_params={}
)
def test_transform_messages_helper_allows_vision_image_inputs():
config = FireworksAIConfig()
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "Describe this"},
{
"type": "image_url",
"image_url": {
"url": "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAE="
},
},
],
}
]
out = config._transform_messages_helper(
messages, model="accounts/fireworks/models/minimax-m3", litellm_params={}
)
assert out == messages
def test_image_inputs_not_rejected_for_fuzzy_non_vision_match():
"""
A custom/fine-tuned model id that hyphen-matches a known non-vision model
(glm-5p2 has supports_vision=False) must not inherit that False via the
substring fallback and hard-reject valid image_url blocks. The capability
gate for rejection uses an exact cost-map match; the fuzzy fallback stays a
soft signal only, so an unmapped vision-capable deployment is not blocked.
"""
config = FireworksAIConfig()
custom_model = "accounts/myorg/models/custom-glm-5p2"
assert config._get_model_cost_capability(custom_model, "supports_vision") is False
assert (
config._get_model_cost_capability_exact(custom_model, "supports_vision") is None
)
messages = [
{
"role": "user",
"content": [
{
"type": "image_url",
"image_url": {
"url": "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAE="
},
},
],
}
]
out = config._transform_messages_helper(
messages, model=custom_model, litellm_params={}
)
assert out == messages
def test_transform_messages_helper_skips_non_dict_content():
config = FireworksAIConfig()
messages = [
{
"role": "user",
"content": ["just a string", {"type": "text", "text": "hello"}],
}
]
out = config._transform_messages_helper(
messages, model="accounts/fireworks/models/glm-5p2", litellm_params={}
)
assert out == messages
def test_transform_messages_helper_no_transform_inline():
config = FireworksAIConfig()
url = "https://example.com/image.jpg"
messages = [
{
"role": "user",
"content": [{"type": "image_url", "image_url": url}],
}
]
out = config._transform_messages_helper(
messages, model="accounts/fireworks/models/minimax-m3", litellm_params={}
)
block = out[0]["content"][0]
assert block["image_url"] == url
assert "#transform=inline" not in block["image_url"]
def test_get_provider_info_vision_from_model_cost(monkeypatch):
config = FireworksAIConfig()
vision_model = "fireworks_ai/test-vision-from-cost"
monkeypatch.setitem(
litellm.model_cost,
vision_model,
{"supports_vision": True, "supports_pdf_input": True},
)
info = config.get_provider_info(vision_model)
assert info["supports_vision"] is True
assert info["supports_pdf_input"] is True
no_vision_model = "fireworks_ai/test-no-vision-from-cost"
monkeypatch.setitem(litellm.model_cost, no_vision_model, {})
info_no_vision = config.get_provider_info(no_vision_model)
assert info_no_vision.get("supports_vision") is not True
assert "supports_pdf_input" not in info_no_vision
def test_reasoning_effort_boolean_true_to_medium():
config = FireworksAIConfig()
result = config.map_openai_params(
{"reasoning_effort": True},
{},
_REASONING_MODEL,
drop_params=False,
)
assert result["reasoning_effort"] == "medium"
def test_reasoning_effort_boolean_false_to_none():
config = FireworksAIConfig()
result = config.map_openai_params(
{"reasoning_effort": False},
{},
_REASONING_MODEL,
drop_params=False,
)
assert result["reasoning_effort"] == "none"
def test_reasoning_effort_string_passthrough():
config = FireworksAIConfig()
result = config.map_openai_params(
{"reasoning_effort": "high"},
{},
_REASONING_MODEL,
drop_params=False,
)
assert result["reasoning_effort"] == "high"
def test_reasoning_effort_integer_passthrough():
config = FireworksAIConfig()
result = config.map_openai_params(
{"reasoning_effort": 1000},
{},
_REASONING_MODEL,
drop_params=False,
)
assert result["reasoning_effort"] == 1000
assert isinstance(result["reasoning_effort"], int)
def test_transform_response_captures_perf_metrics():
body = {
**_BASE_CHAT_COMPLETION_RESPONSE,
"perf_metrics": {"prompt-tokens": 10},
}
result = _run_transform_response(body)
assert result._hidden_params["fireworks_perf_metrics"] == {"prompt-tokens": 10}
def test_transform_response_captures_prompt_token_ids():
body = {
**_BASE_CHAT_COMPLETION_RESPONSE,
"prompt_token_ids": [1, 2, 3],
}
result = _run_transform_response(body)
assert result._hidden_params["fireworks_prompt_token_ids"] == [1, 2, 3]
def test_transform_response_captures_raw_output():
raw_output = {
"prompt_fragments": [],
"prompt_token_ids": [],
"completion": "test",
}
body = {
**_BASE_CHAT_COMPLETION_RESPONSE,
"choices": [
{
**_BASE_CHAT_COMPLETION_RESPONSE["choices"][0],
"raw_output": raw_output,
}
],
}
result = _run_transform_response(body)
assert result._hidden_params["fireworks_raw_outputs"] == [raw_output]
def test_transform_response_captures_token_ids():
body = {
**_BASE_CHAT_COMPLETION_RESPONSE,
"choices": [
{
**_BASE_CHAT_COMPLETION_RESPONSE["choices"][0],
"token_ids": [4, 5, 6],
}
],
}
result = _run_transform_response(body)
assert result._hidden_params["fireworks_token_ids"] == [[4, 5, 6]]
def test_streaming_surfaces_fireworks_response_fields():
"""
The Fireworks-specific response fields captured into _hidden_params for
non-streaming calls must also reach streamed responses. They ride the
streamed chunks' provider_specific_fields (litellm rebuilds each streamed
chunk, so per-chunk _hidden_params does not survive): per-choice
token_ids/raw_output on the content chunk, response-level
perf_metrics/prompt_token_ids on the final usage chunk. Driving the real
litellm.completion(stream=True) path also covers the get_model_response_iterator
wiring; dropping the Fireworks iterator would leave these fields unset.
"""
from litellm.llms.custom_httpx.http_handler import HTTPHandler
model = "accounts/fireworks/models/llama-v3p1-8b-instruct"
raw_output = {"completion": "Hi"}
sse_lines = [
"data: "
+ json.dumps(
{
"id": "stream-1",
"object": "chat.completion.chunk",
"created": 1,
"model": model,
"choices": [
{
"index": 0,
"delta": {"role": "assistant", "content": "Hi"},
"token_ids": [123],
"raw_output": raw_output,
}
],
}
),
"data: "
+ json.dumps(
{
"id": "stream-1",
"object": "chat.completion.chunk",
"created": 1,
"model": model,
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}],
"usage": {
"prompt_tokens": 5,
"completion_tokens": 1,
"total_tokens": 6,
},
"perf_metrics": {"prompt-tokens": 5},
"prompt_token_ids": [1, 2, 3],
}
),
"data: [DONE]",
]
raw_response = MagicMock()
raw_response.status_code = 200
raw_response.headers = {}
raw_response.iter_lines = lambda: iter(sse_lines)
client = HTTPHandler()
with patch.object(client, "post", return_value=raw_response):
stream = litellm.completion(
model=f"fireworks_ai/{model}",
messages=[{"role": "user", "content": "hi"}],
stream=True,
api_key="fw-test-key",
client=client,
)
surfaced: dict = {}
for chunk in stream:
fields = getattr(chunk, "provider_specific_fields", None) or {}
surfaced.update(
{k: v for k, v in fields.items() if k.startswith("fireworks_")}
)
assert surfaced["fireworks_token_ids"] == [[123]]
assert surfaced["fireworks_raw_outputs"] == [raw_output]
assert surfaced["fireworks_perf_metrics"] == {"prompt-tokens": 5}
assert surfaced["fireworks_prompt_token_ids"] == [1, 2, 3]

View file

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

View file

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

View file

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

View file

@ -1750,6 +1750,60 @@ async def test_reject_clientside_metadata_tags_allows_key_tags_without_client_ta
assert request_body["metadata"]["tags"] == ["engineering"]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"route",
[
"/bedrock/model/us.anthropic.claude-sonnet-4-6/invoke",
"/v1/messages",
],
)
async def test_common_checks_metadata_route_keeps_key_tags_out_of_provider_metadata(
route,
):
"""GH#30629: on routes that track tags in litellm_metadata (bedrock, /v1/messages,
responses, ...) key-level tags must land in litellm_metadata, never in the
provider-facing metadata field (Bedrock rejects non-user_id metadata with HTTP 400).
The auth-time pre-seed keys off LITELLM_METADATA_ROUTES, so hardcoding a single route
or dropping the pre-seed makes apply_key_tags_pre_auth fall back to metadata; this
guards that regression.
"""
from fastapi import Request
from litellm.proxy.auth.auth_checks import common_checks
request_body = {"messages": [{"role": "user", "content": "test"}]}
mock_request = MagicMock(spec=Request)
valid_token = UserAPIKeyAuth(
token="test-token",
metadata={"tags": ["engineering"]},
)
with patch(
"litellm.proxy.auth.auth_checks.get_tag_objects_batch",
new_callable=AsyncMock,
return_value={},
):
result = await common_checks(
request_body=request_body,
team_object=None,
user_object=None,
end_user_object=None,
global_proxy_spend=None,
general_settings={},
route=route,
llm_router=None,
proxy_logging_obj=MagicMock(),
valid_token=valid_token,
request=mock_request,
)
assert result is True
assert request_body["litellm_metadata"]["tags"] == ["engineering"]
assert "metadata" not in request_body
@pytest.mark.asyncio
async def test_virtual_key_soft_budget_check_with_user_obj():
"""Test _virtual_key_soft_budget_check includes user_email when user_obj is provided"""
@ -2365,7 +2419,9 @@ async def test_virtual_key_budget_check_reads_from_spend_counter():
proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
proxy_logging_obj.budget_alerts = AsyncMock()
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
async def mock_get_current_spend(
counter_key, fallback_spend, max_budget=None, **kwargs
):
if counter_key == "spend:key:test-hashed-token":
return 1.5
return fallback_spend
@ -2397,7 +2453,9 @@ async def test_virtual_key_budget_check_fallback_no_counter():
proxy_logging_obj.budget_alerts = AsyncMock()
# get_current_spend returns fallback_spend when no counter exists
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
async def mock_get_current_spend(
counter_key, fallback_spend, max_budget=None, **kwargs
):
return fallback_spend
with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend):
@ -2424,7 +2482,9 @@ async def test_team_budget_check_reads_from_spend_counter():
proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
proxy_logging_obj.budget_alerts = AsyncMock()
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
async def mock_get_current_spend(
counter_key, fallback_spend, max_budget=None, **kwargs
):
if counter_key == "spend:team:test-team":
return 1.5
return fallback_spend
@ -2449,7 +2509,9 @@ async def test_end_user_budget_check_reads_from_spend_counter():
litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0),
)
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
async def mock_get_current_spend(
counter_key, fallback_spend, max_budget=None, **kwargs
):
if counter_key == "spend:end_user:customer-1":
return 1.5
return fallback_spend
@ -2475,7 +2537,9 @@ async def test_tag_budget_check_reads_from_spend_counter():
litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0),
)
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
async def mock_get_current_spend(
counter_key, fallback_spend, max_budget=None, **kwargs
):
if counter_key == "spend:tag:paid-tag":
return 1.5
return fallback_spend
@ -2523,7 +2587,9 @@ async def test_team_member_budget_check_reads_from_spend_counter():
proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
async def mock_get_current_spend(
counter_key, fallback_spend, max_budget=None, **kwargs
):
if counter_key == "spend:team_member:test-user:test-team":
return 1.5
return fallback_spend
@ -2756,7 +2822,9 @@ async def test_team_member_budget_check_falls_back_to_team_default_budget_id():
return_value=fake_budget_row
)
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
async def mock_get_current_spend(
counter_key, fallback_spend, max_budget=None, **kwargs
):
if counter_key == "spend:team_member:test-user:test-team":
return 70.0
return fallback_spend
@ -2853,7 +2921,9 @@ async def test_team_member_budget_check_per_member_override_wins_over_team_defau
mocked_spend = 70.0
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
async def mock_get_current_spend(
counter_key, fallback_spend, max_budget=None, **kwargs
):
if counter_key == "spend:team_member:test-user:test-team":
return mocked_spend
return fallback_spend
@ -2943,7 +3013,9 @@ async def test_team_member_budget_check_null_clone_falls_back_to_team_default():
return_value=fake_default_row
)
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
async def mock_get_current_spend(
counter_key, fallback_spend, max_budget=None, **kwargs
):
if counter_key == "spend:team_member:test-user:test-team":
return 500.0
return fallback_spend
@ -3010,7 +3082,9 @@ async def test_team_member_budget_check_null_clone_with_null_default_skips_enfor
return_value=fake_default_row
)
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
async def mock_get_current_spend(
counter_key, fallback_spend, max_budget=None, **kwargs
):
if counter_key == "spend:team_member:test-user:test-team":
return 1000.0
return fallback_spend
@ -3077,7 +3151,9 @@ async def test_team_member_budget_check_zero_team_default_treated_as_no_cap():
return_value=fake_default_row
)
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
async def mock_get_current_spend(
counter_key, fallback_spend, max_budget=None, **kwargs
):
if counter_key == "spend:team_member:test-user:test-team":
return 0.0
return fallback_spend
@ -3135,7 +3211,9 @@ async def test_team_member_budget_check_zero_per_member_row_still_blocks():
prisma_client = MagicMock()
prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=None)
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
async def mock_get_current_spend(
counter_key, fallback_spend, max_budget=None, **kwargs
):
if counter_key == "spend:team_member:test-user:test-team":
return 0.0
return fallback_spend
@ -3714,3 +3792,54 @@ async def test_inference_route_still_enforces_team_budget():
valid_token=UserAPIKeyAuth(token="test-token", team_id="test-team"),
request=MagicMock(),
)
@pytest.mark.asyncio
async def test_virtual_key_max_budget_error_names_the_key():
"""BudgetExceededError for a virtual key must name the key (alias + masked key)
so operators don't have to reverse-map a spend figure back to a key."""
valid_token = UserAPIKeyAuth(
token="hashed-token",
key_alias="payments-prod",
key_name="sk-...um_g",
max_budget=10.0,
spend=0.0,
)
proxy_logging_obj = MagicMock()
proxy_logging_obj.budget_alerts = AsyncMock()
with patch(
"litellm.proxy.proxy_server.get_current_spend",
new=AsyncMock(return_value=25.0),
):
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await _virtual_key_max_budget_check(
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
)
message = str(exc_info.value)
assert "payments-prod" in message
assert "sk-...um_g" in message
@pytest.mark.asyncio
async def test_virtual_key_max_budget_not_exceeded_does_not_raise():
"""Spend below the configured budget must not raise."""
valid_token = UserAPIKeyAuth(
token="hashed-token",
key_alias="payments-prod",
max_budget=10.0,
spend=0.0,
)
proxy_logging_obj = MagicMock()
proxy_logging_obj.budget_alerts = AsyncMock()
with patch(
"litellm.proxy.proxy_server.get_current_spend",
new=AsyncMock(return_value=1.0),
):
await _virtual_key_max_budget_check(
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
)

View file

@ -2160,3 +2160,48 @@ class TestGetRequestRouteTemplate:
lambda self: (_ for _ in ()).throw(RuntimeError("boom"))
)
assert get_request_route_template(req) is None
class TestIsRequestBodySafeBlocksModelList:
"""model_list is an SDK-only field with no proxy API meaning; it must
be rejected from the request body regardless of any opt-in."""
def test_model_list_rejected_with_no_opt_in(self):
with pytest.raises(ValueError, match="model_list is not allowed"):
is_request_body_safe(
request_body={
"model": "gpt-4",
"messages": [{"role": "user", "content": "hi"}],
"model_list": [{"model_name": "x", "litellm_params": {}}],
},
general_settings={},
llm_router=None,
model="gpt-4",
)
def test_model_list_rejected_even_with_proxy_wide_opt_in(self):
with pytest.raises(ValueError, match="model_list is not allowed"):
is_request_body_safe(
request_body={
"model": "gpt-4",
"messages": [{"role": "user", "content": "hi"}],
"model_list": [],
},
general_settings={"allow_client_side_credentials": True},
llm_router=None,
model="gpt-4",
)
def test_normal_body_still_passes(self):
assert (
is_request_body_safe(
request_body={
"model": "gpt-4",
"messages": [{"role": "user", "content": "hi"}],
},
general_settings={},
llm_router=None,
model="gpt-4",
)
is True
)

View file

@ -106,7 +106,9 @@ async def test_custom_auth_enforces_end_user_budget_when_common_checks_skipped()
litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0),
)
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
async def mock_get_current_spend(
counter_key, fallback_spend, max_budget=None, **kwargs
):
if counter_key == "spend:end_user:customer-1":
return 5.0
return fallback_spend

View file

@ -1206,7 +1206,13 @@ async def test_auth_builder_returns_team_membership_object():
JWTAuthManager,
"get_objects",
new_callable=AsyncMock,
return_value=(user_object, None, None, mock_team_membership, user_object.user_id),
return_value=(
user_object,
None,
None,
mock_team_membership,
user_object.user_id,
),
) as mock_get_objects,
patch.object(
JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock
@ -3509,9 +3515,7 @@ def test_canonical_user_id_no_change_when_ids_match():
user_object = LiteLLM_UserTable(user_id=same, user_email=same)
assert (
JWTAuthManager._canonical_user_id_from_db(
user_id=same, user_object=user_object
)
JWTAuthManager._canonical_user_id_from_db(user_id=same, user_object=user_object)
== same
)
@ -3802,12 +3806,15 @@ async def test_get_objects_team_membership_uses_rebound_user_id():
user_id_jwt_field="email", user_id_upsert=True
)
with patch(
"litellm.proxy.auth.handle_jwt.get_user_object",
side_effect=fake_get_user_object,
), patch(
"litellm.proxy.auth.handle_jwt.get_team_membership",
side_effect=fake_get_team_membership,
with (
patch(
"litellm.proxy.auth.handle_jwt.get_user_object",
side_effect=fake_get_user_object,
),
patch(
"litellm.proxy.auth.handle_jwt.get_team_membership",
side_effect=fake_get_team_membership,
),
):
(
user_object,

View file

@ -436,7 +436,9 @@ def test_wildcard_custom_prefix_keeps_org_segment_for_non_provider_first_segment
result = get_known_models_from_wildcard(
wildcard_model="my_hf/*",
litellm_params=LiteLLM_Params(model="huggingface/*", custom_llm_provider="huggingface"),
litellm_params=LiteLLM_Params(
model="huggingface/*", custom_llm_provider="huggingface"
),
)
assert result == ["my_hf/meta-llama/Llama-3-8B"]

View file

@ -0,0 +1,104 @@
from __future__ import annotations
from typing import Any, Dict, List, Optional, Tuple
from fastapi import Request
from litellm.proxy.auth.network import (
TrustedProxyConfig,
resolve_client_ip,
resolve_network_context,
)
TRUSTED = TrustedProxyConfig(use_forwarded_for=True, trusted_proxy_cidrs=["10.0.0.0/8"])
def make_request(
*,
headers: Optional[Dict[str, str]] = None,
client: Optional[Tuple[str, int]] = ("203.0.113.7", 5555),
) -> Request:
raw_headers: List[Tuple[bytes, bytes]] = [
(key.lower().encode(), value.encode()) for key, value in (headers or {}).items()
]
scope: Dict[str, Any] = {
"type": "http",
"http_version": "1.1",
"method": "GET",
"path": "/",
"raw_path": b"/",
"query_string": b"",
"headers": raw_headers,
"client": client,
"server": ("testserver", 80),
"scheme": "http",
}
return Request(scope)
def test_xff_ignored_when_forwarding_disabled():
config = TrustedProxyConfig(
use_forwarded_for=False, trusted_proxy_cidrs=["10.0.0.0/8"]
)
request = make_request(
headers={"x-forwarded-for": "203.0.113.9"}, client=("10.0.0.1", 1)
)
ip, via_proxy = resolve_client_ip(request, config)
assert ip == "10.0.0.1"
assert via_proxy is False
def test_xff_honored_from_trusted_peer():
request = make_request(
headers={"x-forwarded-for": "203.0.113.9, 10.0.0.5"}, client=("10.0.0.1", 1)
)
ip, via_proxy = resolve_client_ip(request, TRUSTED)
assert ip == "203.0.113.9"
assert via_proxy is True
def test_spoofed_xff_from_untrusted_peer_is_ignored():
request = make_request(
headers={"x-forwarded-for": "203.0.113.9"}, client=("8.8.8.8", 1)
)
ip, via_proxy = resolve_client_ip(request, TRUSTED)
assert ip == "8.8.8.8"
assert via_proxy is False
def test_right_to_left_parse_skips_chained_trusted_proxies():
request = make_request(
headers={"x-forwarded-for": "198.51.100.4, 10.1.1.1, 10.0.0.9"},
client=("10.0.0.1", 1),
)
ip, via_proxy = resolve_client_ip(request, TRUSTED)
assert ip == "198.51.100.4"
assert via_proxy is True
def test_all_trusted_hops_fall_back_to_peer():
request = make_request(
headers={"x-forwarded-for": "10.1.1.1, 10.0.0.9"}, client=("10.0.0.1", 1)
)
ip, via_proxy = resolve_client_ip(request, TRUSTED)
assert ip == "10.0.0.1"
assert via_proxy is True
def test_invalid_xff_token_is_skipped():
request = make_request(
headers={"x-forwarded-for": "not-an-ip, 203.0.113.50"}, client=("10.0.0.1", 1)
)
ip, _ = resolve_client_ip(request, TRUSTED)
assert ip == "203.0.113.50"
def test_network_context_captures_host_and_proxy_flag():
request = make_request(
headers={"x-forwarded-for": "203.0.113.9", "host": "proxy.litellm.ai"},
client=("10.0.0.1", 1),
)
ctx = resolve_network_context(request, TRUSTED)
assert ctx.client_ip == "203.0.113.9"
assert ctx.host == "proxy.litellm.ai"
assert ctx.via_trusted_proxy is True

View file

@ -18,7 +18,6 @@ from fastapi import HTTPException
import litellm
from litellm.proxy._types import InvitationClaim
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------

View file

@ -0,0 +1,35 @@
from __future__ import annotations
from litellm.proxy._types import ProxyErrorTypes, ProxyException
from litellm.proxy.auth.resolvers.exceptions import (
IdentityResolutionError,
KeyNotFoundError,
KeyNotInCacheError,
NoDatabaseConnectionError,
PrincipalMissingSourceKeyError,
)
def test_all_resolution_errors_share_one_base():
errors = [
NoDatabaseConnectionError(),
KeyNotInCacheError("hashed"),
KeyNotFoundError("hashed"),
PrincipalMissingSourceKeyError(),
]
assert all(isinstance(e, IdentityResolutionError) for e in errors)
def test_key_not_found_preserves_the_public_401_contract():
# The auth seam catches ProxyException and rewrites the 401 message, so a
# missing key must keep mapping to that exact contract.
error = KeyNotFoundError("hashed-token")
assert isinstance(error, ProxyException)
assert error.code == "401"
assert error.type == ProxyErrorTypes.token_not_found_in_db.value
assert error.param == "key"
def test_key_not_in_cache_names_the_token():
assert "hashed-token" in str(KeyNotInCacheError("hashed-token"))

View file

@ -0,0 +1,95 @@
from __future__ import annotations
import pytest
from pydantic import ValidationError
from litellm.proxy.auth.auth_method import AuthMethod
from litellm.proxy.auth.resolvers.models import (
EndUserIdentity,
Principal,
PrincipalType,
ProjectIdentity,
TeamIdentity,
UserIdentity,
)
from litellm.proxy.auth.roles import Role, TeamRole
def _principal() -> Principal:
return Principal(
principal_type=PrincipalType.HUMAN,
subject="u1",
auth_method=AuthMethod.OIDC,
)
def test_principal_is_frozen():
principal = _principal()
with pytest.raises(ValidationError):
principal.subject = "mutated"
def test_principal_defaults_are_independent_instances():
a = _principal()
b = _principal()
assert a.teams == [] and a.scopes == [] and a.audience == []
assert a.teams is not b.teams
def test_principal_requires_identity_core_fields():
with pytest.raises(ValidationError):
Principal(subject="u1") # missing principal_type + auth_method
def test_principal_roles_are_validated_against_role_enum():
principal = Principal(
principal_type=PrincipalType.HUMAN,
subject="u1",
auth_method=AuthMethod.OIDC,
roles=["org_admin"],
)
assert principal.roles == [Role.ORG_ADMIN]
assert isinstance(principal.roles[0], Role)
with pytest.raises(ValidationError):
Principal(
principal_type=PrincipalType.HUMAN,
subject="u1",
auth_method=AuthMethod.OIDC,
roles=["not_a_real_role"],
)
def test_principal_default_network_and_collections():
principal = Principal(
principal_type=PrincipalType.SERVICE_ACCOUNT,
subject="svc",
auth_method=AuthMethod.MUTUAL_TLS,
)
assert principal.teams == []
assert principal.scopes == []
assert principal.project is None
assert principal.end_user is None
assert principal.network.client_ip is None
assert principal.network.via_trusted_proxy is False
def test_team_identity_defaults_to_member_role():
team = TeamIdentity(id="g1")
assert team.role == TeamRole.MEMBER
def test_user_identity_optional_fields_default_none():
user = UserIdentity(id="u1")
assert user.email is None
assert user.external_id is None
def test_project_identity_name_is_optional():
assert ProjectIdentity(id="p1").name is None
assert ProjectIdentity(id="p1", name="Acme").name == "Acme"
def test_end_user_identity_requires_id():
with pytest.raises(ValidationError):
EndUserIdentity()

View file

@ -0,0 +1,96 @@
from __future__ import annotations
from typing import Any, Dict, List, Optional, Tuple
from fastapi import Request
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_method import AuthMethod
from litellm.proxy.auth.resolvers.models import PrincipalType
from litellm.proxy.auth.roles import Role
from litellm.proxy.auth.user_api_key_auth import _resolve_request_principal
def _request(
*,
headers: Optional[Dict[str, str]] = None,
client: Optional[Tuple[str, int]] = ("203.0.113.7", 5555),
) -> Request:
raw: List[Tuple[bytes, bytes]] = [
(k.lower().encode(), v.encode()) for k, v in (headers or {}).items()
]
scope: Dict[str, Any] = {
"type": "http",
"http_version": "1.1",
"method": "POST",
"path": "/v1/chat/completions",
"raw_path": b"/v1/chat/completions",
"query_string": b"",
"headers": raw,
"client": client,
"server": ("testserver", 80),
"scheme": "http",
}
return Request(scope)
def test_seam_projects_full_identity_from_key_object():
token = UserAPIKeyAuth(
token="hashed-token",
user_id="u-1",
user_role="org_admin",
team_id="t-1",
team_alias="Eng",
org_id="o-1",
organization_alias="Acme",
end_user_id="cust-9",
)
principal = _resolve_request_principal(_request(), token)
assert principal.principal_type == PrincipalType.HUMAN
assert principal.auth_method == AuthMethod.API_KEY
assert principal.user is not None and principal.user.id == "u-1"
assert principal.roles == [Role.ORG_ADMIN]
assert [t.id for t in principal.teams] == ["t-1"]
assert principal.teams[0].name == "Eng"
assert principal.organization is not None and principal.organization.id == "o-1"
assert principal.organization.name == "Acme"
assert principal.end_user is not None and principal.end_user.id == "cust-9"
# the key is always identifiable via credential_ref, even when other ids exist
assert principal.credential_ref.token_id == "hashed-token"
def test_seam_principal_is_never_anonymous_for_keyless_service_account():
# no user_id and no key_alias -> subject must still identify the key
token = UserAPIKeyAuth(token="hashed-token")
principal = _resolve_request_principal(_request(), token)
assert principal.principal_type == PrincipalType.SERVICE_ACCOUNT
assert principal.user is None
assert principal.subject == "hashed-token"
assert principal.credential_ref.token_id == "hashed-token"
def test_seam_stamps_direct_peer_when_no_trusted_proxy_configured():
token = UserAPIKeyAuth(token="hashed-token", user_id="u-1")
# No trusted_proxy_ranges configured -> XFF is not trusted, direct peer wins.
principal = _resolve_request_principal(
_request(headers={"x-forwarded-for": "10.9.9.9"}, client=("203.0.113.7", 5555)),
token,
)
assert principal.network.client_ip == "203.0.113.7"
assert principal.network.via_trusted_proxy is False
def test_seam_detects_jwt_auth_method():
token = UserAPIKeyAuth(
token="hashed-token", user_id="u-2", jwt_claims={"sub": "u-2"}
)
principal = _resolve_request_principal(_request(), token)
assert principal.auth_method == AuthMethod.BEARER_JWT

View file

@ -0,0 +1,70 @@
from __future__ import annotations
from typing import Dict, Optional
import pytest
from litellm.proxy._types import UserAPIKeyAuth, hash_token
from litellm.proxy.auth.resolvers.exceptions import (
NoDatabaseConnectionError,
PrincipalMissingSourceKeyError,
)
from litellm.proxy.auth.auth_method import AuthMethod
from litellm.proxy.auth.resolvers.models import Principal, PrincipalType
from litellm.proxy.auth.resolvers.store import IdentityStore
class _FakeCache:
"""Stands in for the DualCache get_key_object reads. It returns a cache hit
before the DB is touched, so seeding it exercises resolve without a database
(a non-None prisma client is still required; it is never reached on a hit)."""
def __init__(self, entries: Optional[Dict[str, object]] = None) -> None:
self._entries = entries or {}
async def async_get_cache(self, key, *args, **kwargs):
return self._entries.get(key)
async def async_set_cache(self, *args, **kwargs):
return None
async def test_resolve_returns_a_principal_projected_from_the_looked_up_key():
raw = "sk-live-abc"
key = UserAPIKeyAuth(token=hash_token(raw), user_id="u-1", team_id="t-1")
store = IdentityStore(object(), _FakeCache({hash_token(raw): key}))
principal = await store.resolve(hashed_token=hash_token(raw))
assert isinstance(principal, Principal)
assert principal.principal_type == PrincipalType.HUMAN
assert principal.user is not None and principal.user.id == "u-1"
assert [t.id for t in principal.teams] == ["t-1"]
async def test_resolve_carries_the_key_for_key_from_principal():
raw = "sk-live-abc"
key = UserAPIKeyAuth(token=hash_token(raw), user_id="u-1", team_id="t-1")
store = IdentityStore(object(), _FakeCache({hash_token(raw): key}))
principal = await store.resolve(hashed_token=hash_token(raw))
recovered = IdentityStore.key_from_principal(principal)
assert recovered.user_id == "u-1"
assert recovered.team_id == "t-1"
def test_key_from_principal_raises_when_no_source_key_is_carried():
bare = Principal(
principal_type=PrincipalType.SERVICE_ACCOUNT,
subject="svc",
auth_method=AuthMethod.API_KEY,
)
with pytest.raises(PrincipalMissingSourceKeyError):
IdentityStore.key_from_principal(bare)
async def test_resolve_raises_without_a_db_connection():
store = IdentityStore(None, _FakeCache())
with pytest.raises(NoDatabaseConnectionError):
await store.resolve(hashed_token="missing")

View file

@ -403,6 +403,48 @@ def test_virtual_key_llm_api_routes_rejects_mcp_multi_segment_admin_subpaths(
assert exc_info.value.status_code == 403
@pytest.mark.parametrize(
"route, method",
[
("/mcp", "POST"),
("/mcp/", "POST"),
("/mcp/my-server", "POST"), # matches the /mcp/{subpath} pattern
("/mcp/tools", "GET"),
("/mcp/tools/list", "POST"),
("/mcp/tools/call", "POST"),
("/mcp-rest/tools/list", "GET"),
("/mcp-rest/tools/call", "POST"),
("/v1/mcp/tools", "GET"),
],
)
def test_virtual_key_llm_api_routes_allows_mcp_inference_endpoints(route, method):
"""Every MCP inference/discovery endpoint must be reachable by virtual keys
scoped to allowed_routes=["llm_api_routes"], the default the Create Key UI
applies.
/v1/mcp/tools is the most recent addition: before it joined this group a key
could list tools via /mcp/tools/list and /mcp-rest/tools/list but got a 403
on the equivalent /v1/mcp/tools. Unlike /v1/mcp/server, none of these paths
have a management write counterpart, so they live directly in
`mcp_inference_routes` rather than behind a method-aware carve-out.
"""
assert RouteChecks.is_llm_api_route(route=route) is True
valid_token = UserAPIKeyAuth(
user_id="test_user",
allowed_routes=["llm_api_routes"],
)
result = RouteChecks.is_virtual_key_allowed_to_call_route(
route=route,
valid_token=valid_token,
request=_mock_request(method),
)
assert result is True
def test_spend_logs_v2_classified_as_management_not_llm_api():
"""Paginated spend logs are a management/spend read route, not an LLM API."""

View file

@ -1106,7 +1106,7 @@ async def test_proxy_admin_expired_key_from_cache():
# Mock get_key_object to return expired token from cache
with (
patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object",
"litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key",
new_callable=AsyncMock,
) as mock_get_key_object,
patch(
@ -1261,7 +1261,7 @@ async def test_scim_deactivated_user_key_is_rejected():
with (
patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object",
"litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key",
new_callable=AsyncMock,
return_value=valid_token,
),
@ -2484,7 +2484,7 @@ async def test_user_api_key_auth_builder_no_blocking_calls():
stack.enter_context(p)
stack.enter_context(
patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object",
"litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key",
new_callable=AsyncMock,
return_value=valid_token,
)
@ -2598,7 +2598,7 @@ async def test_team_metadata_refreshed_from_team_object_during_auth():
with (
patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object",
"litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key",
new_callable=AsyncMock,
return_value=valid_token,
),
@ -2689,6 +2689,67 @@ async def test_centralized_common_checks_runs_for_standard_auth():
setattr(_proxy_server_mod, k, v)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"route",
[
"/bedrock/model/us.anthropic.claude-sonnet-4-6/invoke",
"/v1/messages",
],
)
async def test_centralized_common_checks_routes_header_tags_to_litellm_metadata(route):
"""GH#30629: on LITELLM_METADATA_ROUTES the tag-budget read resolves to
litellm_metadata, so the litellm_metadata pre-seed must run before
apply_client_tag_policy_pre_auth merges x-litellm-tags. Otherwise header tags
land in metadata and silently escape _tag_max_budget_check. This guards the
pre-seed call site in _run_centralized_common_checks; dropping it routes header
tags back into metadata.
"""
import litellm.proxy.proxy_server as _proxy_server_mod
from fastapi import Request
from starlette.datastructures import URL
token = UserAPIKeyAuth(api_key="sk-test", user_id="u1")
request = Request(
scope={
"type": "http",
"method": "POST",
"headers": [(b"x-litellm-tags", b"tenant:acme")],
"query_string": b"",
}
)
request._url = URL(url=route)
request_data: dict = {"model": "us.anthropic.claude-sonnet-4-6"}
attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=None)
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
try:
for k, v in attrs.items():
setattr(_proxy_server_mod, k, v)
with (
patch(
"litellm.proxy.auth.user_api_key_auth.common_checks",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.auth.user_api_key_auth._reserve_budget_after_common_checks",
new_callable=AsyncMock,
),
):
await _run_centralized_common_checks(
user_api_key_auth_obj=token,
request=request,
request_data=request_data,
route=route,
)
finally:
for k, v in originals.items():
setattr(_proxy_server_mod, k, v)
assert request_data["litellm_metadata"]["tags"] == ["tenant:acme"]
assert "metadata" not in request_data
@pytest.mark.asyncio
async def test_centralized_common_checks_skipped_for_custom_auth_without_flag():
"""Existing RPS guarantee: custom-auth deployments without
@ -3652,7 +3713,7 @@ async def _run_builder_with_key_lookup(get_key_object_mock):
request._url = URL(url="/chat/completions")
with (
patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object",
"litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key",
get_key_object_mock,
),
patch(

View file

@ -13,7 +13,10 @@ from litellm.proxy.management_endpoints.scim.scim_transformations import (
ScimTransformations,
)
from litellm.types.proxy.management_endpoints.scim_v2 import (
SCIM_ENTERPRISE_USER_SCHEMA,
SCIMEnterpriseUser,
SCIMPatchOperation,
SCIMUser,
)
@ -149,6 +152,77 @@ class TestScimTransformations:
assert scim_user.name.givenName == "Test"
assert scim_user.name.familyName == "User"
@pytest.mark.asyncio
async def test_transform_user_with_enterprise_metadata(self, mock_prisma_client):
mock_client, mock_find_unique = mock_prisma_client
mock_find_unique.return_value = None
user = LiteLLM_UserTable(
user_id="user-ent",
user_email="ent@example.com",
user_alias=None,
teams=[],
created_at=None,
updated_at=None,
metadata={
"scim_enterprise": {"costCenter": "CC-42", "department": "Platform"}
},
)
with patch("litellm.proxy.proxy_server.prisma_client", mock_client):
scim_user = await ScimTransformations.transform_litellm_user_to_scim_user(
user
)
assert scim_user.enterprise_user is not None
assert scim_user.enterprise_user.costCenter == "CC-42"
assert scim_user.enterprise_user.department == "Platform"
assert SCIM_ENTERPRISE_USER_SCHEMA in scim_user.schemas
@pytest.mark.asyncio
async def test_transform_user_without_enterprise_metadata_omits_schema(
self, mock_user, mock_prisma_client
):
mock_client, mock_find_unique = mock_prisma_client
team1 = LiteLLM_TeamTable(
team_id="team-1", team_alias="Team One", members_with_roles=[]
)
team2 = LiteLLM_TeamTable(
team_id="team-2", team_alias="Team Two", members_with_roles=[]
)
mock_find_unique.side_effect = [team1, team2]
with patch("litellm.proxy.proxy_server.prisma_client", mock_client):
scim_user = await ScimTransformations.transform_litellm_user_to_scim_user(
mock_user
)
assert scim_user.enterprise_user is None
assert SCIM_ENTERPRISE_USER_SCHEMA not in scim_user.schemas
def test_scim_user_serialization_omits_absent_enterprise_urn(self):
without_enterprise = SCIMUser(
schemas=["urn:ietf:params:scim:schemas:core:2.0:User"],
id="user-1",
userName="user@example.com",
)
dumped = without_enterprise.model_dump(by_alias=True)
assert SCIM_ENTERPRISE_USER_SCHEMA not in dumped
assert "enterprise_user" not in dumped
assert SCIM_ENTERPRISE_USER_SCHEMA not in dumped["schemas"]
with_enterprise = SCIMUser(
schemas=[
"urn:ietf:params:scim:schemas:core:2.0:User",
SCIM_ENTERPRISE_USER_SCHEMA,
],
id="user-2",
userName="ent@example.com",
enterprise_user=SCIMEnterpriseUser(costCenter="CC-42"),
)
dumped_ent = with_enterprise.model_dump(by_alias=True)
assert dumped_ent[SCIM_ENTERPRISE_USER_SCHEMA]["costCenter"] == "CC-42"
@pytest.mark.asyncio
async def test_transform_litellm_team_to_scim_group(
self, mock_team, mock_prisma_client

View file

@ -638,6 +638,84 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys():
assert key_data.metrics.spend == 10.0
def _daily_user_spend_record(*, user_id, api_key, spend):
"""A LiteLLM_DailyUserSpend row as the per-user breakdown reads it."""
return SimpleNamespace(
date="2024-01-01",
user_id=user_id,
api_key=api_key,
model="gpt-4",
model_group="gpt-4",
custom_llm_provider="openai",
mcp_namespaced_tool_name=None,
endpoint="/chat/completions",
spend=spend,
prompt_tokens=10,
completion_tokens=5,
cache_read_input_tokens=0,
cache_creation_input_tokens=0,
api_requests=1,
successful_requests=1,
failed_requests=0,
)
@pytest.mark.asyncio
async def test_get_daily_activity_applies_resolve_entity_metadata_to_breakdown():
"""Regression for LIT-3889: the Spend Per User chart showed raw UUIDs.
/user/daily/activity used to pass entity_metadata_field=None, so every
user entity in the breakdown carried empty metadata and the dashboard had
nothing to render but the user_id UUID. The page-scoped resolver must put
the resolved email/alias onto the entity metadata so the UI can label it,
while a spender with no email on file still falls back to the raw UUID.
"""
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
records = [
_daily_user_spend_record(user_id="user-with-email", api_key="key-1", spend=7.0),
_daily_user_spend_record(user_id="user-no-email", api_key="key-2", spend=3.0),
]
mock_table = MagicMock()
mock_table.count = AsyncMock(return_value=len(records))
mock_table.find_many = AsyncMock(return_value=records)
mock_prisma.db.litellm_dailyuserspend = mock_table
mock_prisma.db.litellm_verificationtoken = MagicMock()
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
seen_user_ids = {}
async def resolver(page_records):
seen_user_ids["ids"] = {r.user_id for r in page_records}
return {"user-with-email": {"user_email": "spender@example.com"}}
result = await get_daily_activity(
prisma_client=mock_prisma,
table_name="litellm_dailyuserspend",
entity_id_field="user_id",
entity_id=None,
entity_metadata_field=None,
start_date="2024-01-01",
end_date="2024-01-01",
model=None,
api_key=None,
page=1,
page_size=1000,
resolve_entity_metadata=resolver,
)
# Resolver is driven by the user_ids actually on the page
assert seen_user_ids["ids"] == {"user-with-email", "user-no-email"}
entities = result.results[0].breakdown.entities
# Email is on the entity metadata so the UI labels the chart with it
assert entities["user-with-email"].metadata["user_email"] == "spender@example.com"
# No email on file -> empty metadata -> UI falls back to the UUID
assert entities["user-no-email"].metadata == {}
class TestAdjustDatesForTimezone:
"""
Regression tests for the timezone double-counting bug.
@ -758,6 +836,8 @@ class TestBuildAggregatedSqlQuery:
]
assert "model = $4" in sql
assert "api_key = $5" in sql
@pytest.mark.asyncio
async def test_get_daily_activity_aggregated_empty_result_set():
"""Regression test for the empty-range 500.

View file

@ -2,6 +2,7 @@ import json
import os
import sys
from datetime import datetime, timezone
from types import SimpleNamespace
import pytest
from fastapi.testclient import TestClient
@ -20,6 +21,7 @@ from litellm.proxy._types import (
)
from litellm.proxy.management_endpoints.internal_user_endpoints import (
LiteLLM_UserTableWithKeyCount,
_resolve_user_email_metadata,
_update_internal_user_params,
get_user_key_counts,
get_users,
@ -657,6 +659,56 @@ async def test_get_users_includes_timestamps(mocker):
assert user_response.key_count == 0
@pytest.mark.asyncio
async def test_get_users_redacts_scim_enterprise_metadata(mocker):
"""
/user/list must strip scim_enterprise from each user's metadata while leaving
the rest of the metadata intact, matching the user-info endpoints.
"""
mock_prisma_client = mocker.MagicMock()
mock_user_row = mocker.MagicMock()
mock_user_row.user_id = "listed-user"
mock_user_row.model_dump.return_value = {
"user_id": "listed-user",
"user_email": "listed@example.com",
"user_role": "internal_user",
"metadata": {
"scim_metadata": {"givenName": "Jane", "familyName": "Doe"},
"scim_enterprise": {"costCenter": "CC-42", "department": "Platform"},
},
}
async def mock_find_many(*args, **kwargs):
return [mock_user_row]
async def mock_count(*args, **kwargs):
return 1
mock_prisma_client.db.litellm_usertable.find_many = mock_find_many
mock_prisma_client.db.litellm_usertable.count = mock_count
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
async def mock_get_user_key_counts(*args, **kwargs):
return {"listed-user": 0}
mocker.patch(
"litellm.proxy.management_endpoints.internal_user_endpoints.get_user_key_counts",
mock_get_user_key_counts,
)
admin_key = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
response = await get_users(
page=1, page_size=1, user_api_key_dict=admin_key, organization_ids=None
)
listed = response["users"][0]
assert listed.metadata == {
"scim_metadata": {"givenName": "Jane", "familyName": "Doe"}
}
assert "scim_enterprise" not in (listed.metadata or {})
def test_validate_sort_params():
"""
Test that validate_sort_params returns None if sort_by is None
@ -2167,6 +2219,94 @@ async def test_user_info_v2_proxy_admin_can_query_any_user(mocker):
assert response.metadata == {"team": "engineering"}
@pytest.mark.asyncio
async def test_user_info_v2_redacts_scim_enterprise_metadata(mocker):
"""
SCIM enterprise attributes are persisted in metadata for reporting, but
/v2/user/info must not surface them; the rest of metadata is preserved.
"""
from fastapi import Request
from litellm.proxy._types import UserInfoV2Response
from litellm.proxy.management_endpoints.internal_user_endpoints import user_info_v2
mock_prisma_client = mocker.MagicMock()
mock_user_row = mocker.MagicMock()
mock_user_row.model_dump.return_value = {
"user_id": "target-user-123",
"user_email": "target@example.com",
"metadata": {
"scim_metadata": {"givenName": "Jane", "familyName": "Doe"},
"scim_enterprise": {
"costCenter": "CC-42",
"department": "Platform",
"employeeNumber": "E-1001",
},
},
"teams": ["team-1"],
}
async def mock_find_unique(*args, **kwargs):
if kwargs.get("where", {}).get("user_id") == "target-user-123":
return mock_user_row
return None
mock_prisma_client.db.litellm_usertable.find_unique = mocker.AsyncMock(
side_effect=mock_find_unique
)
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
mock_request = mocker.MagicMock(spec=Request)
admin_key = UserAPIKeyAuth(
user_id="admin-user", user_role=LitellmUserRoles.PROXY_ADMIN
)
response = await user_info_v2(
request=mock_request,
user_id="target-user-123",
user_api_key_dict=admin_key,
)
assert isinstance(response, UserInfoV2Response)
assert response.metadata == {
"scim_metadata": {"givenName": "Jane", "familyName": "Doe"}
}
assert "scim_enterprise" not in (response.metadata or {})
def test_build_user_info_response_redacts_scim_enterprise_metadata():
"""
The shared /user/info builder strips scim_enterprise from the returned user row
while leaving every other metadata key intact.
"""
from litellm.proxy.management_endpoints.internal_user_endpoints import (
_build_user_info_response,
)
user_row = {
"user_id": "target-user-123",
"metadata": {
"scim_metadata": {"givenName": "Jane"},
"scim_enterprise": {"costCenter": "CC-42"},
},
}
response = _build_user_info_response(
user_id="target-user-123",
user_info=user_row,
keys=None,
team_list=[],
teams_1=None,
)
assert response.user_info is not None
assert response.user_info["metadata"] == {"scim_metadata": {"givenName": "Jane"}}
assert "scim_enterprise" not in response.user_info["metadata"]
@pytest.mark.asyncio
async def test_user_info_v2_internal_user_can_query_self(mocker):
"""
@ -2960,3 +3100,56 @@ async def test_ghsa_wvg4_proxy_admin_can_update_user_budget(mocker):
user_request=user_request, user_api_key_dict=admin_caller
)
assert result is not None
@pytest.mark.asyncio
async def test_resolve_user_email_metadata_maps_page_user_ids_to_email(mocker):
"""Regression for LIT-3889.
The Spend Per User chart rendered raw UUIDs because the per-user activity
breakdown carried no email. This resolver must turn the user_ids on the
page into {user_id: {user_email, user_alias}} so the chart can label each
spender, and it must only look up the user_ids actually present (not the
whole user table).
"""
mock_prisma_client = mocker.MagicMock()
find_many = mocker.AsyncMock(
return_value=[
SimpleNamespace(
user_id="u1", user_email="alice@example.com", user_alias="Alice"
),
SimpleNamespace(user_id="u2", user_email=None, user_alias="bob-alias"),
]
)
mock_prisma_client.db.litellm_usertable.find_many = find_many
records = [
SimpleNamespace(user_id="u1"),
SimpleNamespace(user_id="u1"), # duplicate -> deduped
SimpleNamespace(user_id="u2"),
]
result = await _resolve_user_email_metadata(mock_prisma_client, records)
assert result == {
"u1": {"user_email": "alice@example.com", "user_alias": "Alice"},
"u2": {"user_email": None, "user_alias": "bob-alias"},
}
where_arg = find_many.call_args.kwargs["where"]
assert set(where_arg["user_id"]["in"]) == {"u1", "u2"}
@pytest.mark.asyncio
async def test_resolve_user_email_metadata_skips_db_when_no_user_ids(mocker):
"""No user_ids on the page (e.g. all spend is unattributed) means no query."""
mock_prisma_client = mocker.MagicMock()
find_many = mocker.AsyncMock(return_value=[])
mock_prisma_client.db.litellm_usertable.find_many = find_many
records = [SimpleNamespace(user_id=None), SimpleNamespace(user_id="")]
result = await _resolve_user_email_metadata(mock_prisma_client, records)
assert result == {}
find_many.assert_not_called()

View file

@ -14,6 +14,7 @@ from litellm.proxy.management_helpers.object_permission_utils import (
_extract_requested_mcp_access_groups,
_extract_requested_mcp_server_ids,
_resolve_team_allowed_mcp_servers,
_rewrite_object_permission_mcp_servers,
_set_object_permission,
validate_key_mcp_servers_against_team,
validate_key_search_tools_against_team,
@ -111,6 +112,31 @@ def test_extract_requested_mcp_server_ids_none():
assert _extract_requested_mcp_server_ids({}) == set()
def test_extract_requested_mcp_server_ids_excludes_no_mcp_servers_sentinel():
obj_perm = {"mcp_servers": ["no-mcp-servers", "server-1"]}
assert _extract_requested_mcp_server_ids(obj_perm) == {"server-1"}
def test_rewrite_object_permission_mcp_servers_preserves_sentinel():
obj_perm = {"mcp_servers": ["no-mcp-servers", "alias-1"]}
_rewrite_object_permission_mcp_servers(obj_perm, {"alias-1": {"server-1"}})
assert obj_perm["mcp_servers"] == ["no-mcp-servers", "server-1"]
@pytest.mark.asyncio
async def test_validate_no_mcp_servers_sentinel_passes_and_preserved():
"""A key scoped to no-mcp-servers passes team validation untouched, keeping the
sentinel so it is not mistaken for an unknown server and rejected."""
team_obj = _make_team_obj(mcp_servers=["server-1"])
obj_perm = {"mcp_servers": ["no-mcp-servers"]}
result = await validate_key_mcp_servers_against_team(
object_permission=obj_perm,
team_obj=team_obj,
)
assert result == obj_perm
assert obj_perm["mcp_servers"] == ["no-mcp-servers"]
# ---- Tests for _extract_requested_mcp_access_groups ----

View file

@ -1885,3 +1885,310 @@ class TestNonStreamingResponseRedaction:
leaked = logging_obj.model_call_details.get("complete_streaming_response")
assert leaked is None
assert redacted.choices[0].message.content == "redacted-by-litellm"
def _sse_bytes(data: dict) -> bytes:
return f"event: {data['type']}\ndata: {json.dumps(data)}\n\n".encode()
class TestAnthropicUsageOnlyFallback:
"""When stream_chunk_builder cannot reassemble a large/agentic stream (returns
None or raises), Anthropic still emits token usage in the message_start /
message_delta SSE events. The handler must recover usage-only so the request is
priced instead of being dropped from SpendLogs while Anthropic billed the tokens."""
_CHUNKS = [
_sse_bytes(
{
"type": "message_start",
"message": {
"model": "claude-3-5-haiku-20241022",
"usage": {
"input_tokens": 100,
"cache_read_input_tokens": 40,
"cache_creation_input_tokens": 20,
"output_tokens": 1,
},
},
}
),
_sse_bytes(
{
"type": "message_delta",
"usage": {
"output_tokens": 55,
"server_tool_use": {"web_search_requests": 2},
},
}
),
]
def test_build_usage_only_recovers_cache_inclusive_usage(self):
response = (
AnthropicPassthroughLoggingHandler._build_usage_only_response_from_chunks(
all_chunks=self._CHUNKS, model="claude-3-5-haiku-20241022"
)
)
assert response is not None
usage = response.usage
# prompt_tokens must be cache-inclusive (input + cache_read + cache_creation)
assert usage.prompt_tokens == 160
assert usage.completion_tokens == 55
assert usage._cache_read_input_tokens == 40
assert usage._cache_creation_input_tokens == 20
assert usage.prompt_tokens_details.cached_tokens == 40
assert usage.server_tool_use.web_search_requests == 2
def test_build_usage_only_returns_none_without_usage_events(self):
chunks = [_sse_bytes({"type": "content_block_delta", "delta": {"text": "hi"}})]
assert (
AnthropicPassthroughLoggingHandler._build_usage_only_response_from_chunks(
all_chunks=chunks, model="claude-3-5-haiku-20241022"
)
is None
)
def test_build_usage_only_recovers_cache_split_server_tools_and_model(self):
# the model is "unknown" up-front and only the 5m/1h cache split is sent
# (no flat cache_creation_input_tokens); web/tool-search and geo arrive in
# message_delta. All must be recovered and priced, not left at $0.
chunks = [
"event: ping\ndata: [DONE]\n\n", # ignored sentinel between real events
_sse_bytes(
{
"type": "message_start",
"message": {
"model": "claude-opus-4-6",
"usage": {
"input_tokens": 80,
"output_tokens": 1,
"cache_creation": {
"ephemeral_5m_input_tokens": 12,
"ephemeral_1h_input_tokens": 8,
},
"inference_geo": "us",
},
},
}
),
_sse_bytes(
{
"type": "message_delta",
"delta": {"stop_reason": "tool_use"},
"usage": {
"output_tokens": 40,
"cache_read_input_tokens": 5,
"inference_geo": "us",
"server_tool_use": {
"web_search_requests": 1,
"tool_search_requests": 3,
},
},
}
),
]
response = (
AnthropicPassthroughLoggingHandler._build_usage_only_response_from_chunks(
all_chunks=chunks, model="unknown"
)
)
assert response is not None
assert response.model == "claude-opus-4-6"
# the real stop_reason is surfaced, not a hardcoded "stop"
assert response.choices[0].finish_reason == "tool_calls"
usage = response.usage
# 80 input + 20 cache_creation (derived from 12+8) + 5 cache_read
assert usage.prompt_tokens == 105
assert usage.completion_tokens == 40
assert usage._cache_creation_input_tokens == 20
assert usage._cache_read_input_tokens == 5
assert usage.server_tool_use.web_search_requests == 1
assert usage.server_tool_use.tool_search_requests == 3
@pytest.mark.parametrize(
"event_str,expected",
[
("data: [DONE]", None),
("data: ", None),
("data: {not-json", None),
("event: ping", None),
('data: {"a": 1}', {"a": 1}),
],
)
def test_extract_sse_data_handles_malformed_and_sentinel_lines(
self, event_str, expected
):
assert (
AnthropicPassthroughLoggingHandler._extract_sse_data(event_str) == expected
)
def _real_logging_obj(self):
from litellm.litellm_core_utils.litellm_logging import Logging as RealLoggingObj
logging_obj = RealLoggingObj(
model="claude-3-5-haiku-20241022",
messages=[{"role": "user", "content": "hi"}],
stream=True,
call_type="pass_through_endpoint",
start_time=datetime.now(),
litellm_call_id="test-call-id",
function_id="1",
)
logging_obj.model_call_details["litellm_params"] = {}
logging_obj.litellm_params = {}
return logging_obj
@patch("litellm.completion_cost")
@patch.object(
AnthropicPassthroughLoggingHandler, "_build_complete_streaming_response"
)
def test_handler_falls_back_when_assembly_returns_none(
self, mock_assemble, mock_cost
):
mock_assemble.return_value = None
mock_cost.return_value = 0.0021
logging_obj = self._real_logging_obj()
result = AnthropicPassthroughLoggingHandler._handle_logging_anthropic_collected_chunks(
litellm_logging_obj=logging_obj,
passthrough_success_handler_obj=MagicMock(),
url_route="/anthropic/v1/messages",
request_body={"model": "claude-3-5-haiku-20241022", "stream": True},
endpoint_type="messages",
start_time=datetime.now(),
all_chunks=list(self._CHUNKS),
end_time=datetime.now(),
)
assert result["result"] is not None
assert result["result"].usage.completion_tokens == 55
assert result["kwargs"]["response_cost"] == 0.0021
@patch("litellm.completion_cost")
@patch.object(
AnthropicPassthroughLoggingHandler, "_build_complete_streaming_response"
)
def test_handler_falls_back_when_assembly_raises(self, mock_assemble, mock_cost):
import litellm
mock_assemble.side_effect = litellm.APIError(
status_code=500,
message="boom",
llm_provider="anthropic",
model="claude-3-5-haiku-20241022",
)
mock_cost.return_value = 0.0021
logging_obj = self._real_logging_obj()
result = AnthropicPassthroughLoggingHandler._handle_logging_anthropic_collected_chunks(
litellm_logging_obj=logging_obj,
passthrough_success_handler_obj=MagicMock(),
url_route="/anthropic/v1/messages",
request_body={"model": "claude-3-5-haiku-20241022", "stream": True},
endpoint_type="messages",
start_time=datetime.now(),
all_chunks=list(self._CHUNKS),
end_time=datetime.now(),
)
# a raise from stream_chunk_builder must be treated like a None result,
# not propagate out and drop the request from SpendLogs
assert result["result"] is not None
assert result["result"].usage.completion_tokens == 55
assert result["kwargs"]["response_cost"] == 0.0021
@patch.object(
AnthropicPassthroughLoggingHandler, "_build_complete_streaming_response"
)
def test_handler_returns_none_when_no_usage_recoverable(self, mock_assemble):
# assembly fails AND the chunks carry no usage event, so there is nothing
# to price; the handler must return None rather than fabricate a response
mock_assemble.return_value = None
logging_obj = self._real_logging_obj()
chunks = [_sse_bytes({"type": "content_block_delta", "delta": {"text": "hi"}})]
result = AnthropicPassthroughLoggingHandler._handle_logging_anthropic_collected_chunks(
litellm_logging_obj=logging_obj,
passthrough_success_handler_obj=MagicMock(),
url_route="/anthropic/v1/messages",
request_body={"model": "claude-3-5-haiku-20241022", "stream": True},
endpoint_type="messages",
start_time=datetime.now(),
all_chunks=chunks,
end_time=datetime.now(),
)
assert result["result"] is None
assert result["kwargs"] == {}
@patch.object(
AnthropicPassthroughLoggingHandler, "_build_usage_only_response_from_chunks"
)
@patch.object(
AnthropicPassthroughLoggingHandler, "_build_complete_streaming_response"
)
def test_handler_does_not_crash_when_usage_only_fallback_raises(
self, mock_assemble, mock_fallback
):
# if the usage-only fallback itself raises, it must be treated as None and
# drop gracefully, not propagate out and crash the success handler
mock_assemble.return_value = None
mock_fallback.side_effect = Exception("fallback boom")
logging_obj = self._real_logging_obj()
result = AnthropicPassthroughLoggingHandler._handle_logging_anthropic_collected_chunks(
litellm_logging_obj=logging_obj,
passthrough_success_handler_obj=MagicMock(),
url_route="/anthropic/v1/messages",
request_body={"model": "claude-3-5-haiku-20241022", "stream": True},
endpoint_type="messages",
start_time=datetime.now(),
all_chunks=list(self._CHUNKS),
end_time=datetime.now(),
)
assert result["result"] is None
assert result["kwargs"] == {}
class TestAnthropicResponseCostRecordedOnModelCallDetails:
"""The pass-through success path reads spend from
model_call_details["response_cost"], not from kwargs, so the streaming payload
builder must record it there or streaming pass-through logs $0."""
def test_create_payload_records_response_cost_on_model_call_details(self):
from litellm.types.utils import Choices, Message, ModelResponse
logging_obj = MagicMock()
logging_obj.model_call_details = {}
logging_obj.get_router_model_id.return_value = None
logging_obj.litellm_params = {}
logging_obj.litellm_call_id = "test-call-id"
response = ModelResponse(
id="test-id",
choices=[
Choices(
finish_reason="stop",
index=0,
message=Message(content="hello", role="assistant"),
)
],
created=1234567890,
model="claude-3-7-sonnet-20250219",
usage={"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
)
kwargs = AnthropicPassthroughLoggingHandler._create_anthropic_response_logging_payload(
litellm_model_response=response,
model="claude-3-7-sonnet-20250219",
kwargs={},
start_time=datetime.now(),
end_time=datetime.now(),
logging_obj=logging_obj,
)
assert (
logging_obj.model_call_details["response_cost"] == kwargs["response_cost"]
)
assert logging_obj.model_call_details["response_cost"] > 0

View file

@ -2682,15 +2682,19 @@ async def test_add_litellm_data_to_request_adds_headers_to_metadata():
version="1.0",
)
# Verify headers are added to metadata for guardrails
assert "metadata" in result, "metadata should be present in result"
assert "headers" in result["metadata"], "headers should be present in metadata"
# Verify headers are added to litellm_metadata for guardrails.
# Bedrock passthrough uses litellm_metadata to prevent key-level
# tags from leaking into the provider payload (GH#30629).
assert "litellm_metadata" in result, "litellm_metadata should be present in result"
assert (
"headers" in result["litellm_metadata"]
), "headers should be present in litellm_metadata"
assert isinstance(
result["metadata"]["headers"], dict
result["litellm_metadata"]["headers"], dict
), "headers should be a dictionary"
# Verify specific headers are accessible (important for guardrails)
headers = result["metadata"]["headers"]
headers = result["litellm_metadata"]["headers"]
assert (
"user-agent" in headers or "User-Agent" in headers
), "User-Agent header should be accessible in metadata"

View file

@ -118,3 +118,20 @@ async def test_chunk_processor_does_not_schedule_logging_when_no_chunks():
assert received == []
mock_route.assert_not_called()
def test_convert_raw_bytes_survives_truncated_multibyte_sequence():
"""A stream cut mid-multibyte-sequence (client disconnect) must still decode
via errors="replace" so the usage events already received are logged, instead
of raising UnicodeDecodeError and dropping the whole request from SpendLogs."""
# the 3-byte "☃" (E2 98 83) is cut after 2 bytes, leaving an invalid sequence
# that strict utf-8 decode would raise on, discarding the message_delta line too
truncated_codepoint = "☃".encode("utf-8")[:2]
raw_bytes = [
b'data: {"text": "' + truncated_codepoint,
b'\ndata: {"type": "message_delta"}\n',
]
lines = PassThroughStreamingHandler._convert_raw_bytes_to_str_lines(raw_bytes)
assert any('"type": "message_delta"' in line for line in lines)

View file

@ -9,7 +9,7 @@ Pins (PR2):
from __future__ import annotations
from unittest.mock import AsyncMock, MagicMock
from unittest.mock import MagicMock
import pytest
@ -128,6 +128,104 @@ def test_v1_model_info_no_model_list_error(client, auth_as, null_router, path):
assert "LLM Model List not loaded" in response.text
# ---------------------------------------------------------------------------
# GET /model/info — team BYOK scoping (issue #30983)
# ---------------------------------------------------------------------------
_BYOK_TEAM_ID = "team-abc"
_BYOK_PUBLIC_NAME = "my-byok-gpt-4"
_BYOK_INTERNAL_NAME = f"model_name_{_BYOK_TEAM_ID}_0123456789abcdef"
@pytest.fixture
def byok_team_router(monkeypatch):
"""Router holding one team-scoped BYOK deployment for team `team-abc`.
Mirrors how a team's own-key BYOK model lives in the router: the routing
key is an internal mangled name while the public name lives in
`model_info.team_public_model_name`.
"""
byok_deployment = {
"model_name": _BYOK_INTERNAL_NAME,
"litellm_params": {"model": "openai/gpt-4"},
"model_info": {
"id": "byok-deployment-id",
"db_model": True,
"team_id": _BYOK_TEAM_ID,
"team_public_model_name": _BYOK_PUBLIC_NAME,
},
}
router = MagicMock()
router.model_list = [byok_deployment]
router.get_model_list_from_model_alias = MagicMock(return_value=[])
router.get_model_names = MagicMock(return_value=[])
router.get_model_access_groups = MagicMock(return_value={})
router.get_model_ids = MagicMock(return_value=[])
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(proxy_server, "llm_model_list", [byok_deployment])
monkeypatch.setattr(proxy_server, "user_model", None)
yield router
@pytest.mark.parametrize("path", ["/v1/model/info", "/model/info"])
def test_model_info_team_key_sees_own_byok_model(client, auth_as, byok_team_router, mock_prisma, monkeypatch, path):
"""Regression for #30983: a team key (user_id=None) must see its own
team's BYOK model under the public name.
Before the fix `_get_caller_byok_team_scope` keyed only off the bound
user's team memberships, returned an empty set for a team key, and the
BYOK row was dropped -> `{"data": []}`.
"""
from litellm.proxy._types import LitellmUserRoles
monkeypatch.setattr(proxy_server, "prisma_client", mock_prisma)
mock_prisma.db.litellm_usertable.find_unique.return_value = None
with auth_as(
role=LitellmUserRoles.INTERNAL_USER,
user_id=None,
team_id=_BYOK_TEAM_ID,
team_models=[_BYOK_PUBLIC_NAME],
):
response = client.get(path)
assert response.status_code == 200
data = response.json()["data"]
surfaced_names = [m.get("model_name") for m in data]
assert _BYOK_PUBLIC_NAME in surfaced_names
assert _BYOK_INTERNAL_NAME not in surfaced_names
@pytest.mark.parametrize("path", ["/v1/model/info", "/model/info"])
def test_model_info_team_key_cannot_see_other_teams_byok_model(
client, auth_as, byok_team_router, mock_prisma, monkeypatch, path
):
"""A team key for a different team must NOT see team-abc's BYOK row.
Guards the fix from over-broadening into a cross-team metadata leak.
"""
from litellm.proxy._types import LitellmUserRoles
monkeypatch.setattr(proxy_server, "prisma_client", mock_prisma)
mock_prisma.db.litellm_usertable.find_unique.return_value = None
with auth_as(
role=LitellmUserRoles.INTERNAL_USER,
user_id=None,
team_id="other-team",
team_models=[_BYOK_PUBLIC_NAME],
):
response = client.get(path)
assert response.status_code == 200
data = response.json()["data"]
surfaced_names = [m.get("model_name") for m in data]
assert _BYOK_PUBLIC_NAME not in surfaced_names
assert _BYOK_INTERNAL_NAME not in surfaced_names
# ---------------------------------------------------------------------------
# GET /model_group/info
# ---------------------------------------------------------------------------

View file

@ -316,8 +316,10 @@ async def test_model_info_v1_unrestricted_key_hides_other_team_byok(monkeypatch)
@pytest.mark.asyncio
async def test_model_info_v1_service_key_hides_all_team_byok(monkeypatch):
"""A key without a resolvable user (e.g. CI/service token) sees only
global deployments, never any team-scoped BYOK rows."""
"""A key with no resolvable user and no team (e.g. a CI/service token
created outside any team) sees only global deployments, never team-scoped
BYOK rows. A team-scoped key does see its own team's rows (issue #30983),
pinned by the /model/info route tests."""
team_row = _team_row()
other_team_row = _other_team_row()
global_row = {
@ -343,7 +345,7 @@ async def test_model_info_v1_service_key_hides_all_team_byok(monkeypatch):
caller = UserAPIKeyAuth(
user_id=None,
user_role=LitellmUserRoles.INTERNAL_USER,
team_id="team-abc-123",
team_id=None,
models=[],
team_models=[],
)
@ -352,6 +354,106 @@ async def test_model_info_v1_service_key_hides_all_team_byok(monkeypatch):
assert [m["model_info"]["id"] for m in resp["data"]] == ["global-id-1"]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"find_unique",
[
AsyncMock(return_value=MagicMock(teams=[])),
AsyncMock(return_value=None),
AsyncMock(side_effect=RuntimeError("db down")),
],
ids=["user-not-in-team", "user-row-missing", "user-lookup-error"],
)
async def test_model_info_v1_team_key_sees_own_byok_regardless_of_user_lookup(
monkeypatch, find_unique
):
"""A team-scoped key sees its own team's BYOK rows even when the bound user
is not a member of that team, has no DB row, or the lookup errors; the
key's team_id is authoritative (issue #30983). Other teams' rows stay
hidden."""
global_row = {
"model_name": "gpt-4",
"litellm_params": {"model": "gpt-4"},
"model_info": {"id": "global-id-1", "db_model": False},
}
router = MagicMock()
router.model_list = [_team_row(), _other_team_row(), global_row]
router.get_model_names.return_value = ["gpt-4"]
router.get_model_access_groups.return_value = {}
prisma_client = MagicMock()
prisma_client.db.litellm_usertable.find_unique = find_unique
async def _populate(**kwargs):
return kwargs["all_models"]
monkeypatch.setattr(ps, "user_model", None)
monkeypatch.setattr(ps, "llm_model_list", router.model_list)
monkeypatch.setattr(ps, "llm_router", router)
monkeypatch.setattr(ps, "prisma_client", prisma_client)
monkeypatch.setattr(ps, "_populate_team_access_on_models", _populate)
monkeypatch.setattr(
ps, "_enrich_model_info_with_litellm_data", lambda model, **kw: model
)
caller = UserAPIKeyAuth(
user_id="user-1",
user_role=LitellmUserRoles.INTERNAL_USER,
team_id="team-abc-123",
models=[],
team_models=[],
)
resp = await ps.model_info_v1(user_api_key_dict=caller, litellm_model_id=None)
assert [m["model_info"]["id"] for m in resp["data"]] == ["byok-id-1", "global-id-1"]
@pytest.mark.asyncio
async def test_model_info_v1_user_team_membership_grants_byok(monkeypatch):
"""A user's own team memberships still grant that team's BYOK rows, unioned
with any team the key itself is scoped to."""
global_row = {
"model_name": "gpt-4",
"litellm_params": {"model": "gpt-4"},
"model_info": {"id": "global-id-1", "db_model": False},
}
router = MagicMock()
router.model_list = [_team_row(), _other_team_row(), global_row]
router.get_model_names.return_value = ["gpt-4"]
router.get_model_access_groups.return_value = {}
prisma_client = MagicMock()
prisma_client.db.litellm_usertable.find_unique = AsyncMock(
return_value=MagicMock(teams=["team-other"])
)
async def _populate(**kwargs):
return kwargs["all_models"]
monkeypatch.setattr(ps, "user_model", None)
monkeypatch.setattr(ps, "llm_model_list", router.model_list)
monkeypatch.setattr(ps, "llm_router", router)
monkeypatch.setattr(ps, "prisma_client", prisma_client)
monkeypatch.setattr(ps, "_populate_team_access_on_models", _populate)
monkeypatch.setattr(
ps, "_enrich_model_info_with_litellm_data", lambda model, **kw: model
)
caller = UserAPIKeyAuth(
user_id="user-2",
user_role=LitellmUserRoles.INTERNAL_USER,
team_id=None,
models=[],
team_models=[],
)
resp = await ps.model_info_v1(user_api_key_dict=caller, litellm_model_id=None)
assert [m["model_info"]["id"] for m in resp["data"]] == [
"byok-id-other",
"global-id-1",
]
@pytest.mark.asyncio
async def test_model_info_v1_populates_access_via_team_ids(monkeypatch):
"""`/v1/model/info` must populate access_via_team_ids when the DB is connected."""

View file

@ -201,6 +201,48 @@ def test_anthropic_provider_fields_support_byok():
), "api_base must appear before api_key in credential_fields (matches AI21 and ANTHROPIC_TEXT convention)."
def test_bedrock_mantle_provider_fields():
"""Amazon Bedrock Mantle must be a selectable provider in the Add Model flow.
The dropdown is driven entirely by /public/providers/fields, so a missing
entry means Mantle cannot be added through the UI at all (regression guard
for LIT-3885). The credential fields must match what the backend actually
honors: an optional bearer api_key (BYOK), the AWS SigV4 chain, a region,
and an api_base override.
"""
app_instance = FastAPI()
app_instance.include_router(router)
test_client = TestClient(app_instance)
response = test_client.get("/public/providers/fields")
assert response.status_code == 200
providers = response.json()
mantle = next((p for p in providers if p["provider"] == "BedrockMantle"), None)
assert mantle is not None, "Bedrock Mantle provider entry not found"
# provider must equal the UI provider_map key so the model dropdown resolves
# bedrock_mantle models; litellm_provider must be the backend slug.
assert mantle["provider_display_name"] == "Amazon Bedrock Mantle"
assert mantle["litellm_provider"] == "bedrock_mantle"
assert mantle["default_model_placeholder"].startswith("bedrock_mantle/")
fields_by_key = {f["key"]: f for f in mantle["credential_fields"]}
# Bearer-token auth is BYOK: optional and masked.
assert "api_key" in fields_by_key
assert fields_by_key["api_key"]["required"] is False
assert fields_by_key["api_key"]["field_type"] == "password"
# AWS SigV4 fallback credentials.
assert fields_by_key["aws_access_key_id"]["field_type"] == "password"
assert fields_by_key["aws_secret_access_key"]["field_type"] == "password"
assert "aws_region_name" in fields_by_key
# api_base override so admins can target a custom Mantle host without env access.
assert fields_by_key["api_base"]["field_type"] == "text"
def test_google_ai_studio_provider_fields_expose_api_base():
"""The Google AI Studio (gemini) credential form must let admins set a custom
api_base so they can point at a Gemini-compatible gateway (e.g. a self-hosted
@ -819,3 +861,44 @@ def test_public_mcp_hub_returns_empty_when_whitelist_unset():
assert response.status_code == 200
assert response.json() == []
app.dependency_overrides.clear()
def test_public_mcp_hub_does_not_expose_upstream_url():
"""Regression: /public/mcp_hub is unauthenticated, so the gateway-internal
upstream url must never appear in its response even when the server has one."""
from litellm.types.mcp_server.mcp_server_manager import MCPServer
from litellm.proxy._types import MCPTransport
app = FastAPI()
app.include_router(router)
app.dependency_overrides[user_api_key_auth] = lambda: MagicMock()
client = TestClient(app)
secret_url = "https://internal-only.example.com/mcp"
server = MCPServer(
server_id="listed",
name="listed",
server_name="listed",
url=secret_url,
transport=MCPTransport.http,
available_on_public_internet=True,
)
mock_manager = MagicMock()
mock_manager.get_public_mcp_servers.return_value = [server]
with (
patch("litellm.public_mcp_servers", ["listed"]),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
mock_manager,
),
):
response = client.get("/public/mcp_hub")
assert response.status_code == 200
data = response.json()
assert [item["server_id"] for item in data] == ["listed"]
assert all("url" not in item for item in data)
assert secret_url not in response.text
app.dependency_overrides.clear()

View file

@ -96,6 +96,20 @@ class TestGetMetadataVariableName:
request = self._make_request("/v1/embeddings")
assert _get_metadata_variable_name(request) == "metadata"
def test_returns_litellm_metadata_for_bedrock_invoke(self):
# GH#30629: bedrock passthrough must use litellm_metadata
# to prevent key-level tags from leaking into provider body
request = self._make_request(
"/bedrock/model/us.anthropic.claude-sonnet-4-6/invoke"
)
assert _get_metadata_variable_name(request) == "litellm_metadata"
def test_returns_litellm_metadata_for_bedrock_converse(self):
request = self._make_request(
"/bedrock/model/us.anthropic.claude-sonnet-4-6/converse"
)
assert _get_metadata_variable_name(request) == "litellm_metadata"
def test_get_enforced_params_for_service_account_settings():
"""
@ -718,7 +732,18 @@ async def test_add_litellm_data_to_request_strips_user_control_fields():
@pytest.mark.asyncio
@pytest.mark.parametrize(
"control_field",
["callbacks", "service_callback", "logger_fn", "litellm_disabled_callbacks"],
[
"callbacks",
"service_callback",
"logger_fn",
"litellm_disabled_callbacks",
"_agentic_loop_depth",
"_agentic_loop_fingerprints",
"_code_interpreter_interception_active",
"_code_interpreter_interception_converted_stream",
"_code_interpreter_interception_sandbox_key",
"max_agentic_loops",
],
)
async def test_add_litellm_data_to_request_strips_callback_control_fields(
control_field,
@ -741,12 +766,19 @@ async def test_add_litellm_data_to_request_strips_callback_control_fields(
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
sample_value = (
["langfuse"]
if control_field
in ("callbacks", "service_callback", "litellm_disabled_callbacks")
else "module.func"
)
sample_values = {
"callbacks": ["langfuse"],
"service_callback": ["langfuse"],
"litellm_disabled_callbacks": ["langfuse"],
"logger_fn": "module.func",
"_agentic_loop_depth": 5,
"_agentic_loop_fingerprints": ["forged"],
"_code_interpreter_interception_active": True,
"_code_interpreter_interception_converted_stream": True,
"_code_interpreter_interception_sandbox_key": "forged-key",
"max_agentic_loops": 9999,
}
sample_value = sample_values[control_field]
updated = await add_litellm_data_to_request(
data={
@ -4150,7 +4182,9 @@ class TestApplyClientTagPolicyPreAuth:
litellm_budget_table=LiteLLM_BudgetTable(max_budget=0.10),
)
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
async def mock_get_current_spend(
counter_key, fallback_spend, max_budget=None, **kwargs
):
if counter_key == "spend:tag:paid":
return 0.50
return fallback_spend
@ -4207,7 +4241,9 @@ class TestApplyClientTagPolicyPreAuth:
litellm_budget_table=LiteLLM_BudgetTable(max_budget=0.10),
)
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
async def mock_get_current_spend(
counter_key, fallback_spend, max_budget=None, **kwargs
):
if counter_key == "spend:tag:tenant:acme":
return 0.50
return fallback_spend
@ -4234,6 +4270,99 @@ class TestApplyClientTagPolicyPreAuth:
assert exc_info.value.current_cost == 0.50
assert exc_info.value.max_budget == 0.10
@pytest.mark.asyncio
@pytest.mark.parametrize(
"route",
[
"/bedrock/model/us.anthropic.claude-sonnet-4-6/invoke",
"/v1/messages",
],
)
async def test_header_tags_visible_to_tag_max_budget_check_on_metadata_route(
self, route
):
"""Regression: on LITELLM_METADATA_ROUTES (bedrock, /v1/messages, ...),
common_checks pre-seeds ``litellm_metadata`` and writes key tags there
before ``_tag_max_budget_check`` reads from the same key. The auth wrapper
calls ``apply_client_tag_policy_pre_auth`` first, so without an earlier
pre-seed header tags land in ``metadata`` and the budget check (now
resolving to ``litellm_metadata``) silently ignores them. This test mirrors
the actual auth-time call order and verifies that an over-budget
header-supplied tag still trips ``_tag_max_budget_check``.
"""
from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_TagTable
from litellm.proxy.auth.auth_checks import common_checks
from litellm.proxy.utils import ProxyLogging
request_mock = _build_request_mock_with_headers(
{"x-litellm-tags": "tenant:acme"}
)
data = {"model": "us.anthropic.claude-sonnet-4-6"}
valid_token = UserAPIKeyAuth(
token="test-token",
api_key="hashed-key",
metadata={},
team_metadata={},
)
LiteLLMProxyRequestSetup.pre_seed_litellm_metadata_for_route(
request_data=data,
route=route,
)
LiteLLMProxyRequestSetup.apply_client_tag_policy_pre_auth(
request=request_mock,
request_data=data,
user_api_key_dict=valid_token,
)
tag_object = LiteLLM_TagTable(
tag_name="tenant:acme",
spend=0.0,
litellm_budget_table=LiteLLM_BudgetTable(max_budget=0.10),
)
async def mock_get_current_spend(
counter_key, fallback_spend, max_budget=None, **kwargs
):
if counter_key == "spend:tag:tenant:acme":
return 0.50
return fallback_spend
with (
patch(
"litellm.proxy.proxy_server.prisma_client",
MagicMock(),
),
patch(
"litellm.proxy.proxy_server.get_current_spend",
mock_get_current_spend,
),
patch(
"litellm.proxy.auth.auth_checks.get_tag_objects_batch",
new_callable=AsyncMock,
return_value={"tenant:acme": tag_object},
),
):
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await common_checks(
request_body=data,
team_object=None,
user_object=None,
end_user_object=None,
global_proxy_spend=None,
general_settings={},
route=route,
llm_router=None,
proxy_logging_obj=ProxyLogging(user_api_key_cache=None),
valid_token=valid_token,
request=request_mock,
)
assert exc_info.value.current_cost == 0.50
assert exc_info.value.max_budget == 0.10
assert "metadata" not in data
assert data["litellm_metadata"]["tags"] == ["tenant:acme"]
class TestApplyKeyTagsPreAuth:
def test_merges_key_tags_into_metadata(self):
@ -4362,7 +4491,9 @@ class TestApplyKeyTagsPreAuth:
litellm_budget_table=LiteLLM_BudgetTable(max_budget=0.10),
)
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
async def mock_get_current_spend(
counter_key, fallback_spend, max_budget=None, **kwargs
):
if counter_key == "spend:tag:engineering":
return 0.50
return fallback_spend
@ -4413,7 +4544,9 @@ class TestApplyKeyTagsPreAuth:
litellm_budget_table=LiteLLM_BudgetTable(max_budget=0.10),
)
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
async def mock_get_current_spend(
counter_key, fallback_spend, max_budget=None, **kwargs
):
if counter_key == "spend:tag:engineering":
return 0.05
return fallback_spend

View file

@ -0,0 +1,232 @@
"""Regression tests for UI-registered embed plugins.
Covers three bugs:
1. `general_settings.plugins` was not a field on ConfigGeneralSettings, so the
admin UI's POST /config/field/update with field_name="plugins" was rejected
with "Invalid field=plugins passed in."
2. The in-memory plugin registry only refreshed at startup, so a plugin added
via the UI did not appear in /api/plugins until a restart.
3. Plugins persisted to DB general_settings were not loaded on startup (the
registry only initialised from the YAML config), so UI-added plugins vanished
after a restart.
"""
import asyncio
from unittest.mock import MagicMock
from litellm.proxy._types import (
ConfigGeneralSettings,
LitellmUserRoles,
PluginConfig,
UserAPIKeyAuth,
)
from litellm.proxy.plugin_routes import list_plugins, register_plugins_from_config
def _admin() -> UserAPIKeyAuth:
return UserAPIKeyAuth(api_key="sk-admin", user_role=LitellmUserRoles.PROXY_ADMIN)
def _non_admin() -> UserAPIKeyAuth:
return UserAPIKeyAuth(api_key="sk-user", user_role=LitellmUserRoles.INTERNAL_USER)
def test_plugins_is_a_valid_general_setting() -> None:
"""The config-update endpoint gates on this exact membership check."""
assert "plugins" in ConfigGeneralSettings.model_fields
def test_config_general_settings_parses_plugin_list() -> None:
"""A list of plugin dicts (what the UI sends) coerces into PluginConfig."""
settings = ConfigGeneralSettings.model_validate(
{
"plugins": [
{
"name": "chat-ui",
"display_name": "Chat UI",
"url": "http://localhost:3300",
},
{
"name": "agent-builder",
"url": "http://127.0.0.1:4010",
"plugin_key": "sk-secret",
},
]
}
)
plugins = settings.plugins
assert plugins is not None
assert [p.name for p in plugins] == ["chat-ui", "agent-builder"]
assert isinstance(plugins[0], PluginConfig)
assert plugins[1].display_name is None
assert plugins[1].plugin_key == "sk-secret"
def test_registered_plugins_appear_in_list_without_restart() -> None:
"""register_plugins_from_config makes UI-added plugins visible immediately,
and replaces (not merges) so removed plugins disappear."""
register_plugins_from_config(
{
"plugins": [
{
"name": "chat-ui",
"display_name": "Chat UI",
"url": "http://localhost:3300",
}
]
}
)
names = [p["name"] for p in asyncio.run(list_plugins(user_api_key_dict=_admin()))]
assert names == ["chat-ui"]
register_plugins_from_config(
{
"plugins": [
{
"name": "chat-ui",
"display_name": "Chat UI",
"url": "http://localhost:3300",
},
{
"name": "agent-builder",
"display_name": "Agent Builder",
"url": "http://127.0.0.1:4010",
},
]
}
)
names = sorted(
p["name"] for p in asyncio.run(list_plugins(user_api_key_dict=_admin()))
)
assert names == ["agent-builder", "chat-ui"]
# Removing a plugin from config drops it from the live list.
register_plugins_from_config({})
assert asyncio.run(list_plugins(user_api_key_dict=_admin())) == []
def test_plugin_key_is_never_returned_to_the_browser() -> None:
"""plugin_key is a credential the UI never needs; /api/plugins must omit it
for every caller, admin included, so it never lands in browser state."""
register_plugins_from_config(
{
"plugins": [
{
"name": "p",
"display_name": "P",
"url": "http://localhost:9",
"plugin_key": "sk-secret",
}
]
}
)
admin_entry = asyncio.run(list_plugins(user_api_key_dict=_admin()))[0]
user_entry = asyncio.run(list_plugins(user_api_key_dict=_non_admin()))[0]
assert "plugin_key" not in admin_entry
assert "plugin_key" not in user_entry
assert admin_entry["url"] == "http://localhost:9"
register_plugins_from_config({})
def test_db_persisted_plugins_load_on_startup() -> None:
"""Plugins saved to DB general_settings must register when the DB config is
merged at startup, not just when present in the YAML file."""
from litellm.proxy.proxy_server import ProxyConfig
register_plugins_from_config({}) # start empty (as if YAML had no plugins)
ProxyConfig()._add_general_settings_from_db_config(
config_data={
"general_settings": {
"plugins": [
{
"name": "db-plugin",
"display_name": "DB Plugin",
"url": "http://localhost:5000",
}
]
}
},
general_settings={},
proxy_logging_obj=MagicMock(),
)
names = [p["name"] for p in asyncio.run(list_plugins(user_api_key_dict=_admin()))]
assert names == ["db-plugin"]
register_plugins_from_config({})
def test_safe_response_headers_sandbox_and_strips_wire_headers() -> None:
"""Proxied plugin responses must be inert and shed wire/cookie headers."""
from litellm.proxy.plugin_routes import _safe_response_headers
out = _safe_response_headers(
{
"content-type": "text/html",
"content-encoding": "gzip",
"content-length": "123",
"set-cookie": "session=abc",
"content-security-policy": "default-src *",
}
)
assert out["content-security-policy"] == "sandbox"
assert out["x-content-type-options"] == "nosniff"
assert out["content-type"] == "text/html"
for stripped in ("content-encoding", "content-length", "set-cookie"):
assert stripped not in out
def test_litellm_credential_header_names_covers_every_auth_header() -> None:
"""The canonical strip set must list every header user_api_key_auth accepts
as a litellm key, so a new auth header can't silently start leaking."""
from litellm.proxy._types import SpecialHeaders
assert SpecialHeaders.litellm_credential_header_names() == {
"authorization",
"api-key",
"x-api-key",
"x-goog-api-key",
"ocp-apim-subscription-key",
"x-litellm-api-key",
}
def test_every_litellm_auth_header_is_stripped_before_forwarding() -> None:
"""A plugin must never receive any header that authenticates against litellm,
only the hop-by-hop set and benign headers are forwarded."""
from litellm.proxy.plugin_routes import _request_strip_headers
strip = _request_strip_headers()
incoming = {
"Authorization": "Bearer sk-litellm",
"API-Key": "sk-litellm",
"X-Api-Key": "sk-litellm",
"X-Goog-Api-Key": "sk-litellm",
"Ocp-Apim-Subscription-Key": "sk-litellm",
"X-Litellm-Api-Key": "sk-litellm",
"Cookie": "litellm_session=abc",
"Accept": "application/json",
"X-Trace-Id": "t-1",
}
forwarded = {k: v for k, v in incoming.items() if k.lower() not in strip}
assert forwarded == {"Accept": "application/json", "X-Trace-Id": "t-1"}
def test_configured_custom_key_header_is_stripped() -> None:
"""A custom general_settings.litellm_key_header_name must also be stripped,
read live so config changes are honoured without a restart."""
from litellm.proxy import proxy_server
from litellm.proxy.plugin_routes import _request_strip_headers
original = getattr(proxy_server, "general_settings", None)
proxy_server.general_settings = {"litellm_key_header_name": "X-My-Tenant-Key"}
try:
assert "x-my-tenant-key" in _request_strip_headers()
finally:
proxy_server.general_settings = original

View file

@ -1348,6 +1348,7 @@ async def test_apply_search_filter_scopes_byok_to_caller_teams():
non_admin = MagicMock(spec=UserAPIKeyAuth)
non_admin.user_role = LitellmUserRoles.INTERNAL_USER
non_admin.user_id = "user-mine"
non_admin.team_id = None
filtered, total_count = await _apply_search_filter_to_models(
all_models=[caller_team_byok, other_team_byok, public_model],
@ -1381,6 +1382,7 @@ async def test_apply_search_filter_scopes_byok_to_caller_teams():
admin = MagicMock(spec=UserAPIKeyAuth)
admin.user_role = LitellmUserRoles.PROXY_ADMIN
admin.user_id = "admin-1"
admin.team_id = None
filtered_admin, _ = await _apply_search_filter_to_models(
all_models=[caller_team_byok, other_team_byok, public_model],
@ -8326,3 +8328,39 @@ def test_get_config_list_includes_cancel_on_disconnect(monkeypatch):
assert fields["cancel_on_disconnect"]["field_type"] == "Boolean"
finally:
app.dependency_overrides.clear()
def test_preserve_redacted_plugin_keys_keeps_stored_credential():
"""A redacted or blank plugin_key on update must not overwrite the real key."""
from litellm.proxy.proxy_server import _preserve_redacted_plugin_keys
existing = [{"name": "p1", "url": "https://p1", "plugin_key": "sk-real-1"}]
redacted = _preserve_redacted_plugin_keys(
[{"name": "p1", "url": "https://p1-new", "plugin_key": "***"}], existing
)
assert redacted == [
{"name": "p1", "url": "https://p1-new", "plugin_key": "sk-real-1"}
]
blanked = _preserve_redacted_plugin_keys(
[{"name": "p1", "url": "https://p1", "plugin_key": ""}], existing
)
assert blanked[0]["plugin_key"] == "sk-real-1"
def test_preserve_redacted_plugin_keys_sets_new_and_drops_orphan_placeholder():
"""A real new key replaces; a placeholder with no stored key is dropped, never persisted."""
from litellm.proxy.proxy_server import _preserve_redacted_plugin_keys
existing = [{"name": "p1", "url": "https://p1", "plugin_key": "sk-real-1"}]
rotated = _preserve_redacted_plugin_keys(
[{"name": "p1", "url": "https://p1", "plugin_key": "sk-new"}], existing
)
assert rotated[0]["plugin_key"] == "sk-new"
new_plugin = _preserve_redacted_plugin_keys(
[{"name": "p2", "url": "https://p2", "plugin_key": "***"}], existing
)
assert "plugin_key" not in new_plugin[0]

View file

@ -105,6 +105,33 @@ def test_jsonify_team_object_converts_members_to_json_string(
}
def test_jsonify_team_object_converts_budget_limits_to_json_string(
prisma_client: PrismaClient,
) -> None:
data = {
"team_id": "t1",
"budget_limits": [
{
"budget_duration": "1d",
"max_budget": 10.0,
"reset_at": "2026-01-01T00:00:00Z",
},
{
"budget_duration": "7d",
"max_budget": 50.0,
"reset_at": "2026-01-07T00:00:00Z",
},
],
"models": ["gpt-4"],
}
result = prisma_client.jsonify_team_object(data)
assert result == {
"team_id": "t1",
"budget_limits": json.dumps(data["budget_limits"]),
"models": ["gpt-4"],
}
def test_jsonify_team_object_error_on_non_dict(prisma_client: PrismaClient) -> None:
with pytest.raises(AttributeError):
prisma_client.jsonify_team_object(None) # type: ignore[arg-type]

View file

@ -295,3 +295,24 @@ async def test_public_lifecycle_create_run_delete():
async def test_unsupported_provider_raises():
with pytest.raises(ValueError):
await litellm.acreate_sandbox(provider="not-a-provider")
# ---------- api_base override ----------
@pytest.mark.asyncio
async def test_create_uses_api_base_override():
client = FakeHTTPClient()
await E2BSandboxConfig().acreate_sandbox(
api_base="http://my-sandbox:8080", api_key="k", client=client
)
_, url, _, _ = client.calls[0]
assert url == "http://my-sandbox:8080/sandboxes"
@pytest.mark.asyncio
async def test_create_defaults_to_e2b_api_base():
client = FakeHTTPClient()
await E2BSandboxConfig().acreate_sandbox(api_key="k", client=client)
_, url, _, _ = client.calls[0]
assert url == "https://api.e2b.app/sandboxes"

View file

@ -0,0 +1,181 @@
"""Unit tests for the sandbox-tool registry."""
from litellm.sandbox import sandbox_tools
def _reset():
sandbox_tools.clear_sandbox_tools()
def test_register_resolves_provider_key_and_base():
_reset()
sandbox_tools.register_sandbox_tools(
[
{
"sandbox_tool_name": "e2b_default",
"litellm_params": {
"sandbox_provider": "e2b",
"api_key": "sk-literal",
"api_base": "https://sandbox.internal",
},
}
]
)
resolved = sandbox_tools.resolve_sandbox_tool("e2b_default")
assert resolved == {
"sandbox_provider": "e2b",
"api_key": "sk-literal",
"api_base": "https://sandbox.internal",
}
_reset()
def test_register_clears_stale_entries_on_reload():
"""A tool removed from the config must not survive a re-registration."""
_reset()
sandbox_tools.register_sandbox_tools(
[
{
"sandbox_tool_name": "old",
"litellm_params": {"sandbox_provider": "e2b"},
}
]
)
assert sandbox_tools.resolve_sandbox_tool("old") is not None
sandbox_tools.register_sandbox_tools(
[
{
"sandbox_tool_name": "new",
"litellm_params": {"sandbox_provider": "e2b"},
}
]
)
assert sandbox_tools.resolve_sandbox_tool("new") is not None
assert (
sandbox_tools.resolve_sandbox_tool("old") is None
), "stale tool must be gone after the config is reloaded"
_reset()
def test_register_empty_list_clears_removed_tools():
"""Reloading a config with sandbox_tools removed (the proxy passes an empty
list) must drop previously registered credentials from the process."""
_reset()
sandbox_tools.register_sandbox_tools(
[
{
"sandbox_tool_name": "e2b_default",
"litellm_params": {"sandbox_provider": "e2b", "api_key": "sk-x"},
}
]
)
assert sandbox_tools.resolve_sandbox_tool("e2b_default") is not None
sandbox_tools.register_sandbox_tools([])
assert (
sandbox_tools.resolve_sandbox_tool("e2b_default") is None
), "removing sandbox_tools from config must clear stale credentials"
_reset()
def test_register_resolves_secret_from_env(monkeypatch):
_reset()
monkeypatch.setenv("MY_SANDBOX_KEY", "sk-from-env")
sandbox_tools.register_sandbox_tools(
[
{
"sandbox_tool_name": "e2b_default",
"litellm_params": {
"sandbox_provider": "e2b",
"api_key": "os.environ/MY_SANDBOX_KEY",
},
}
]
)
resolved = sandbox_tools.resolve_sandbox_tool("e2b_default")
assert resolved is not None
assert resolved["api_key"] == "sk-from-env"
assert resolved["api_base"] is None
_reset()
def test_resolve_unknown_returns_none():
_reset()
assert sandbox_tools.resolve_sandbox_tool("nope") is None
def test_register_skips_malformed_entries_without_crashing():
"""A single malformed entry (missing sandbox_tool_name, or not a dict) must
not crash registration during proxy startup/hot-reload; valid entries in the
same list must still register."""
_reset()
sandbox_tools.register_sandbox_tools(
[
{"litellm_params": {"sandbox_provider": "e2b"}}, # missing name
"not-a-dict", # wrong type
{"sandbox_tool_name": "", "litellm_params": {}}, # empty name
{
"sandbox_tool_name": "good",
"litellm_params": {"sandbox_provider": "e2b"},
},
]
)
assert sandbox_tools.resolve_sandbox_tool("good") is not None
assert sandbox_tools.resolve_sandbox_tool("") is None
assert set(sandbox_tools._SANDBOX_TOOL_REGISTRY) == {"good"}
_reset()
def test_register_skips_entry_missing_sandbox_provider():
"""An entry with a name but no sandbox_provider must be skipped at
registration so it cannot later resolve and call acreate_sandbox(provider=None),
which fails with a cryptic runtime error instead of a clear startup warning."""
_reset()
sandbox_tools.register_sandbox_tools(
[
{"sandbox_tool_name": "no_provider", "litellm_params": {"api_key": "sk-x"}},
{
"sandbox_tool_name": "null_provider",
"litellm_params": {"sandbox_provider": None},
},
{
"sandbox_tool_name": "good",
"litellm_params": {"sandbox_provider": "e2b"},
},
]
)
assert sandbox_tools.resolve_sandbox_tool("no_provider") is None
assert sandbox_tools.resolve_sandbox_tool("null_provider") is None
assert set(sandbox_tools._SANDBOX_TOOL_REGISTRY) == {"good"}
_reset()
def test_register_swaps_registry_atomically():
"""register_sandbox_tools must replace the registry in one rebind so a
concurrent resolve never observes a half-populated or transiently empty
registry between clearing and repopulating."""
_reset()
sandbox_tools.register_sandbox_tools(
[{"sandbox_tool_name": "a", "litellm_params": {"sandbox_provider": "e2b"}}]
)
before = sandbox_tools._SANDBOX_TOOL_REGISTRY
sandbox_tools.register_sandbox_tools(
[
{"sandbox_tool_name": "b", "litellm_params": {"sandbox_provider": "e2b"}},
{"sandbox_tool_name": "c", "litellm_params": {"sandbox_provider": "e2b"}},
]
)
after = sandbox_tools._SANDBOX_TOOL_REGISTRY
assert after is not before, "the registry must be replaced, not mutated in place"
assert set(after) == {"b", "c"}
assert "a" not in after
_reset()

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