mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
commit
39d132f2a6
132 changed files with 13895 additions and 3113 deletions
|
|
@ -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
141
docs/plugin_architecture.md
Normal 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
|
||||
|
|
@ -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",
|
||||
]
|
||||
473
litellm/integrations/code_interpreter_interception/handler.py
Normal file
473
litellm/integrations/code_interpreter_interception/handler.py
Normal 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)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
95
litellm/proxy/_experimental/mcp_server/AGENTS.md
Normal file
95
litellm/proxy/_experimental/mcp_server/AGENTS.md
Normal 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.
|
||||
|
|
@ -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 []
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}",
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
13
litellm/proxy/auth/auth_method.py
Normal file
13
litellm/proxy/auth/auth_method.py
Normal 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"
|
||||
|
|
@ -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))
|
||||
|
|
|
|||
101
litellm/proxy/auth/network.py
Normal file
101
litellm/proxy/auth/network.py
Normal 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,
|
||||
)
|
||||
33
litellm/proxy/auth/resolvers/__init__.py
Normal file
33
litellm/proxy/auth/resolvers/__init__.py
Normal 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",
|
||||
]
|
||||
50
litellm/proxy/auth/resolvers/exceptions.py
Normal file
50
litellm/proxy/auth/resolvers/exceptions.py
Normal 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"
|
||||
)
|
||||
82
litellm/proxy/auth/resolvers/models.py
Normal file
82
litellm/proxy/auth/resolvers/models.py
Normal 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)
|
||||
206
litellm/proxy/auth/resolvers/store.py
Normal file
206
litellm/proxy/auth/resolvers/store.py
Normal 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,
|
||||
)
|
||||
35
litellm/proxy/auth/roles.py
Normal file
35
litellm/proxy/auth/roles.py
Normal 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
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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()]
|
||||
|
|
|
|||
344
litellm/proxy/plugin_routes.py
Normal file
344
litellm/proxy/plugin_routes.py
Normal 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),
|
||||
)
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
60
litellm/sandbox/sandbox_tools.py
Normal file
60
litellm/sandbox/sandbox_tools.py
Normal 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([])
|
||||
22
litellm/types/integrations/code_interpreter_interception.py
Normal file
22
litellm/types/integrations/code_interpreter_interception.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)"}'
|
||||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
104
tests/test_litellm/proxy/auth/test_network.py
Normal file
104
tests/test_litellm/proxy/auth/test_network.py
Normal 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
|
||||
|
|
@ -18,7 +18,6 @@ from fastapi import HTTPException
|
|||
import litellm
|
||||
from litellm.proxy._types import InvitationClaim
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
35
tests/test_litellm/proxy/auth/test_resolvers_exceptions.py
Normal file
35
tests/test_litellm/proxy/auth/test_resolvers_exceptions.py
Normal 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"))
|
||||
95
tests/test_litellm/proxy/auth/test_resolvers_models.py
Normal file
95
tests/test_litellm/proxy/auth/test_resolvers_models.py
Normal 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()
|
||||
96
tests/test_litellm/proxy/auth/test_resolvers_seam.py
Normal file
96
tests/test_litellm/proxy/auth/test_resolvers_seam.py
Normal 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
|
||||
70
tests/test_litellm/proxy/auth/test_resolvers_store.py
Normal file
70
tests/test_litellm/proxy/auth/test_resolvers_store.py
Normal 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")
|
||||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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 ----
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
232
tests/test_litellm/proxy/test_plugin_routes.py
Normal file
232
tests/test_litellm/proxy/test_plugin_routes.py
Normal 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
|
||||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
181
tests/test_litellm/sandbox/test_sandbox_tools.py
Normal file
181
tests/test_litellm/sandbox/test_sandbox_tools.py
Normal 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
Loading…
Add table
Reference in a new issue