mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge branch 'BerriAI:litellm_internal_staging' into vertex_gemma_client_handler_fix
This commit is contained in:
commit
3c2df9da95
64 changed files with 6003 additions and 434 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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -2142,6 +2142,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 +2164,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"
|
||||
)
|
||||
|
|
@ -3808,6 +3825,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
|
||||
|
|
|
|||
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"
|
||||
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":
|
||||
|
|
@ -2501,6 +2510,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 +2652,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"
|
||||
|
|
@ -141,6 +141,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 = (
|
||||
|
|
|
|||
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.
|
||||
|
|
@ -14931,6 +14946,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 +15047,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 +15064,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 +15122,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 +16456,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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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)"}'
|
||||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -2365,7 +2365,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 +2399,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 +2428,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 +2455,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 +2483,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 +2533,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 +2768,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 +2867,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 +2959,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 +3028,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 +3097,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 +3157,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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
@ -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,
|
||||
),
|
||||
|
|
@ -3652,7 +3652,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(
|
||||
|
|
|
|||
|
|
@ -718,7 +718,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 +752,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 +4168,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 +4227,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
|
||||
|
|
@ -4362,7 +4384,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 +4437,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
|
||||
|
|
@ -8326,3 +8326,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]
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -82,13 +82,3 @@ def test_introduced_keeps_only_violations_on_changed_lines():
|
|||
@pytest.mark.parametrize("hunk", ["@@ -1 +1 @@", "@@ -1,0 +1,2 @@"])
|
||||
def test_parse_changed_lines_handles_single_and_ranged_hunks(hunk):
|
||||
assert gate.parse_changed_lines(f"+++ b/litellm/a.py\n{hunk}\n")["litellm/a.py"]
|
||||
|
||||
|
||||
def test_report_emits_breached_rules_as_final_line(capsys):
|
||||
# CI surfaces only the tail of the log, so the breached-rule summary (rule,
|
||||
# total/cap, added) must be the last line or it gets truncated away.
|
||||
breaches = sorted([gate.Breach("UP045", 530, 529, 1), gate.Breach("ANN401", 12, 10, 2)])
|
||||
new = [gate.Violation("litellm/types/llms/bedrock.py", 16, "UP045")]
|
||||
gate.report(breaches, new, "origin/litellm_internal_staging")
|
||||
last = capsys.readouterr().out.strip().splitlines()[-1]
|
||||
assert last == "BREACHED RULES: ANN401 12/10 (+2); UP045 530/529 (+1)"
|
||||
|
|
|
|||
|
|
@ -4292,7 +4292,7 @@ _FIREWORKS_MODELS = [
|
|||
6e-08,
|
||||
512000,
|
||||
512000,
|
||||
False,
|
||||
True,
|
||||
True,
|
||||
),
|
||||
(
|
||||
|
|
|
|||
|
|
@ -0,0 +1,83 @@
|
|||
import { describe, it, expect, vi } from "vitest";
|
||||
import { render, waitFor } from "@testing-library/react";
|
||||
import { AgentControlPlaneView } from "./layout";
|
||||
|
||||
const { getMock } = vi.hoisted(() => ({ getMock: vi.fn(() => Promise.resolve({ session_claim: "claim" })) }));
|
||||
|
||||
const pluginModeValue = {
|
||||
mode: "litellm-platform-plugin" as string,
|
||||
setMode: vi.fn(),
|
||||
plugins: [{ name: "litellm-platform-plugin", display_name: "Chat UI", url: "http://localhost:3300" }],
|
||||
activePlugin: { name: "litellm-platform-plugin", display_name: "Chat UI", url: "http://localhost:3300" } as {
|
||||
name: string;
|
||||
display_name: string;
|
||||
url: string;
|
||||
} | null,
|
||||
};
|
||||
|
||||
vi.mock("@/contexts/PluginModeContext", () => ({ usePluginMode: () => pluginModeValue }));
|
||||
vi.mock("@/contexts/AuthContext", () => ({ useAuth: () => ({ accessToken: "sk-test-token" }) }));
|
||||
|
||||
vi.mock("@/lib/http/client", () => ({
|
||||
createApiClient: () => ({ get: getMock }),
|
||||
}));
|
||||
vi.mock("@/components/networking", () => ({ getProxyBaseUrl: () => "" }));
|
||||
|
||||
describe("AgentControlPlaneView iframe", () => {
|
||||
it("embeds the plugin at its ROOT url, never a hardcoded subpath like /sessions", () => {
|
||||
const { container } = render(<AgentControlPlaneView />);
|
||||
const iframe = container.querySelector("iframe");
|
||||
|
||||
expect(iframe).not.toBeNull();
|
||||
const src = iframe!.getAttribute("src")!;
|
||||
expect(src).toBe("http://localhost:3300/");
|
||||
expect(src).not.toContain("/sessions");
|
||||
// title comes from the plugin's display_name, not a hardcoded label
|
||||
expect(iframe!.getAttribute("title")).toBe("Chat UI");
|
||||
});
|
||||
|
||||
it("does not double the slash when the plugin url has a trailing slash", () => {
|
||||
pluginModeValue.activePlugin = {
|
||||
name: "litellm-platform-plugin",
|
||||
display_name: "Chat UI",
|
||||
url: "http://localhost:3300/",
|
||||
};
|
||||
const { container } = render(<AgentControlPlaneView />);
|
||||
expect(container.querySelector("iframe")!.getAttribute("src")).toBe("http://localhost:3300/");
|
||||
pluginModeValue.activePlugin = {
|
||||
name: "litellm-platform-plugin",
|
||||
display_name: "Chat UI",
|
||||
url: "http://localhost:3300",
|
||||
};
|
||||
});
|
||||
|
||||
it("does not leak the raw token in the iframe src (token goes via encrypted postMessage)", () => {
|
||||
const { container } = render(<AgentControlPlaneView />);
|
||||
expect(container.querySelector("iframe")!.getAttribute("src")).not.toContain("token");
|
||||
});
|
||||
|
||||
it("does not delegate clipboard-read to the untrusted plugin iframe", () => {
|
||||
const { container } = render(<AgentControlPlaneView />);
|
||||
const allow = container.querySelector("iframe")!.getAttribute("allow") ?? "";
|
||||
|
||||
expect(allow).not.toContain("clipboard-read");
|
||||
expect(allow).toContain("clipboard-write");
|
||||
});
|
||||
|
||||
it("requests the auth-token claim scoped to the active plugin, not a hardcoded default", async () => {
|
||||
getMock.mockClear();
|
||||
pluginModeValue.activePlugin = { name: "reports-plugin", display_name: "Reports", url: "http://localhost:3300" };
|
||||
render(<AgentControlPlaneView />);
|
||||
|
||||
await waitFor(() => expect(getMock).toHaveBeenCalled());
|
||||
const [path, opts] = getMock.mock.calls[0];
|
||||
expect(path).toBe("/api/plugins/auth-token");
|
||||
expect(opts.query).toEqual({ plugin_name: "reports-plugin" });
|
||||
|
||||
pluginModeValue.activePlugin = {
|
||||
name: "litellm-platform-plugin",
|
||||
display_name: "Chat UI",
|
||||
url: "http://localhost:3300",
|
||||
};
|
||||
});
|
||||
});
|
||||
|
|
@ -1,6 +1,6 @@
|
|||
"use client";
|
||||
|
||||
import React, { Suspense, useState } from "react";
|
||||
import React, { Suspense, useState, useRef, useEffect } from "react";
|
||||
import Navbar from "@/components/navbar";
|
||||
import LoadingScreen from "@/components/common_components/LoadingScreen";
|
||||
import { ThemeProvider } from "@/contexts/ThemeContext";
|
||||
|
|
@ -9,6 +9,88 @@ import SidebarProvider from "@/app/(dashboard)/components/SidebarProvider";
|
|||
import { useRouter, useSearchParams, usePathname } from "next/navigation";
|
||||
import { DebugWarningBanner } from "@/components/DebugWarningBanner";
|
||||
import { MIGRATED_PAGES, migratedHref, legacyPageHref, legacyKeyForPathname } from "@/utils/migratedPages";
|
||||
import { PluginModeProvider, usePluginMode } from "@/contexts/PluginModeContext";
|
||||
import { createApiClient } from "@/lib/http/client";
|
||||
import { getProxyBaseUrl } from "@/components/networking";
|
||||
|
||||
const pluginApiClient = createApiClient({ getBaseUrl: () => getProxyBaseUrl() ?? "" });
|
||||
|
||||
// Wrapper so PluginModeProvider receives the live accessToken from auth context,
|
||||
// which means plugin data refreshes on login/logout without stale cookie reads.
|
||||
function PluginModeProviderWithAuth({ children }: { children: React.ReactNode }) {
|
||||
const { accessToken } = useAuth();
|
||||
return <PluginModeProvider accessToken={accessToken}>{children}</PluginModeProvider>;
|
||||
}
|
||||
|
||||
export function AgentControlPlaneView() {
|
||||
const { activePlugin } = usePluginMode();
|
||||
const activePluginName = activePlugin?.name;
|
||||
const agentPlatformUrl = activePlugin?.url ?? "";
|
||||
const { accessToken } = useAuth();
|
||||
const iframeRef = useRef<HTMLIFrameElement>(null);
|
||||
const [auth, setAuth] = useState<{ plugin: string; claim: string } | null>(null);
|
||||
|
||||
// Fetch a short-lived identity claim scoped to the *active* plugin. The claim
|
||||
// is encrypted under that plugin's own per-plugin key, so it must be requested
|
||||
// per plugin and re-fetched when the user switches plugins.
|
||||
useEffect(() => {
|
||||
if (!accessToken || !activePluginName) return;
|
||||
let cancelled = false;
|
||||
pluginApiClient
|
||||
.get("/api/plugins/auth-token", { accessToken, query: { plugin_name: activePluginName } })
|
||||
.then((data: { session_claim?: string }) => {
|
||||
if (!cancelled && data?.session_claim) setAuth({ plugin: activePluginName, claim: data.session_claim });
|
||||
})
|
||||
.catch(() => {});
|
||||
return () => {
|
||||
cancelled = true;
|
||||
};
|
||||
}, [accessToken, activePluginName]);
|
||||
|
||||
// Deliver the claim to the iframe via postMessage, but only while it was issued
|
||||
// for the plugin currently mounted — never replay one plugin's claim to another.
|
||||
// targetOrigin is the configured plugin URL — no other origin receives it.
|
||||
useEffect(() => {
|
||||
const iframe = iframeRef.current;
|
||||
if (!iframe || !auth || auth.plugin !== activePluginName || !agentPlatformUrl) return;
|
||||
const send = () => {
|
||||
iframe.contentWindow?.postMessage({ type: "litellm-auth", session_claim: auth.claim }, agentPlatformUrl);
|
||||
};
|
||||
// Cover both orderings: the iframe may have already fired `load` before the
|
||||
// claim arrived (send now), or it may load/reload later (send on the event).
|
||||
send();
|
||||
iframe.addEventListener("load", send);
|
||||
return () => iframe.removeEventListener("load", send);
|
||||
}, [auth, activePluginName, agentPlatformUrl]);
|
||||
|
||||
if (!agentPlatformUrl) {
|
||||
return (
|
||||
<div className="flex flex-1 items-center justify-center text-gray-500">
|
||||
<div className="text-center">
|
||||
<p className="text-lg font-medium mb-2">Plugin</p>
|
||||
<p className="text-sm">Configure the plugin URL in settings</p>
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
// Embed the plugin at its root; the plugin renders its own full UI (incl. nav) inside.
|
||||
return (
|
||||
<iframe
|
||||
ref={iframeRef}
|
||||
src={`${agentPlatformUrl.replace(/\/$/, "")}/`}
|
||||
style={{
|
||||
width: "100%",
|
||||
height: "100%",
|
||||
border: "none",
|
||||
flex: 1,
|
||||
minHeight: "calc(100vh - 56px)",
|
||||
}}
|
||||
title={activePlugin?.display_name ?? "Plugin"}
|
||||
allow="clipboard-write"
|
||||
/>
|
||||
);
|
||||
}
|
||||
|
||||
function DashboardShell({ children }: { children: React.ReactNode }) {
|
||||
const router = useRouter();
|
||||
|
|
@ -16,6 +98,7 @@ function DashboardShell({ children }: { children: React.ReactNode }) {
|
|||
const pathname = usePathname();
|
||||
const { accessToken } = useAuth();
|
||||
const [sidebarCollapsed, setSidebarCollapsed] = useState(false);
|
||||
const { mode } = usePluginMode();
|
||||
|
||||
const page = legacyKeyForPathname(pathname) || searchParams.get("page") || "api-keys";
|
||||
|
||||
|
|
@ -34,10 +117,18 @@ function DashboardShell({ children }: { children: React.ReactNode }) {
|
|||
/>
|
||||
<DebugWarningBanner accessToken={accessToken} />
|
||||
<div className="flex flex-1">
|
||||
<div className="mt-2">
|
||||
<SidebarProvider setPage={navigateToPage} defaultSelectedKey={page} sidebarCollapsed={sidebarCollapsed} />
|
||||
</div>
|
||||
<main className="flex-1">{children}</main>
|
||||
{mode !== "ai-gateway" ? (
|
||||
<div className="flex-1 flex">
|
||||
<AgentControlPlaneView />
|
||||
</div>
|
||||
) : (
|
||||
<>
|
||||
<div className="mt-2">
|
||||
<SidebarProvider setPage={navigateToPage} defaultSelectedKey={page} sidebarCollapsed={sidebarCollapsed} />
|
||||
</div>
|
||||
<main className="flex-1">{children}</main>
|
||||
</>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
|
|
@ -62,7 +153,9 @@ function LayoutContent({ children }: { children: React.ReactNode }) {
|
|||
export default function Layout({ children }: { children: React.ReactNode }) {
|
||||
return (
|
||||
<Suspense fallback={<LoadingScreen />}>
|
||||
<LayoutContent>{children}</LayoutContent>
|
||||
<PluginModeProviderWithAuth>
|
||||
<LayoutContent>{children}</LayoutContent>
|
||||
</PluginModeProviderWithAuth>
|
||||
</Suspense>
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ import LoggingSettings from "./Settings/AdminSettings/LoggingSettings/LoggingSet
|
|||
import SSOSettings from "./Settings/AdminSettings/SSOSettings/SSOSettings";
|
||||
import UISettings from "./Settings/AdminSettings/UISettings/UISettings";
|
||||
import HashicorpVault from "./Settings/AdminSettings/HashicorpVault/HashicorpVault";
|
||||
import PluginSettings from "./Settings/AdminSettings/PluginSettings/PluginSettings";
|
||||
import SSOModals from "./SSOModals";
|
||||
import UIAccessControlForm from "./UIAccessControlForm";
|
||||
|
||||
|
|
@ -373,6 +374,11 @@ const AdminPanel: React.FC<AdminPanelProps> = ({ proxySettings }) => {
|
|||
label: "Hashicorp Vault",
|
||||
children: <HashicorpVault />,
|
||||
},
|
||||
{
|
||||
key: "plugins",
|
||||
label: "Plugins",
|
||||
children: <PluginSettings />,
|
||||
},
|
||||
];
|
||||
|
||||
return (
|
||||
|
|
|
|||
|
|
@ -0,0 +1,76 @@
|
|||
import { afterEach, describe, expect, it, vi } from "vitest";
|
||||
import { act, fireEvent, render, screen, waitFor } from "@testing-library/react";
|
||||
import ViewSwitcher from "./ViewSwitcher";
|
||||
|
||||
const { mockUsePluginMode, state } = vi.hoisted(() => {
|
||||
const state = {
|
||||
mode: "ai-gateway" as string,
|
||||
setMode: vi.fn(),
|
||||
plugins: [] as { name: string; display_name: string; url: string }[],
|
||||
activePlugin: null as { name: string; display_name: string; url: string } | null,
|
||||
};
|
||||
return { mockUsePluginMode: vi.fn(() => state), state };
|
||||
});
|
||||
|
||||
vi.mock("@/contexts/PluginModeContext", () => ({ usePluginMode: mockUsePluginMode }));
|
||||
|
||||
describe("ViewSwitcher", () => {
|
||||
afterEach(() => {
|
||||
state.mode = "ai-gateway";
|
||||
state.plugins = [];
|
||||
state.setMode.mockClear();
|
||||
});
|
||||
|
||||
it("renders nothing when there are no plugins", () => {
|
||||
const { container } = render(<ViewSwitcher />);
|
||||
expect(container.firstChild).toBeNull();
|
||||
});
|
||||
|
||||
it("labels the button from the active plugin's display_name and lists AI Gateway + each plugin", async () => {
|
||||
state.plugins = [
|
||||
{ name: "litellm-platform-plugin", display_name: "Chat UI", url: "http://localhost:3300" },
|
||||
{ name: "obs", display_name: "Observability", url: "http://localhost:9000" },
|
||||
];
|
||||
state.mode = "litellm-platform-plugin";
|
||||
render(<ViewSwitcher />);
|
||||
|
||||
expect(screen.getByRole("button")).toHaveTextContent("Chat UI");
|
||||
expect(screen.queryByText("Agent Control Plane")).not.toBeInTheDocument();
|
||||
|
||||
act(() => {
|
||||
fireEvent.click(screen.getByRole("button"));
|
||||
});
|
||||
await waitFor(() => expect(screen.getByText("AI Gateway")).toBeInTheDocument());
|
||||
expect(screen.getByText("Observability")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("switches to AI Gateway when that entry is picked", async () => {
|
||||
state.plugins = [{ name: "litellm-platform-plugin", display_name: "Chat UI", url: "http://localhost:3300" }];
|
||||
state.mode = "litellm-platform-plugin";
|
||||
render(<ViewSwitcher />);
|
||||
|
||||
act(() => {
|
||||
fireEvent.click(screen.getByRole("button"));
|
||||
});
|
||||
await waitFor(() => expect(screen.getByText("AI Gateway")).toBeInTheDocument());
|
||||
act(() => {
|
||||
fireEvent.click(screen.getByText("AI Gateway"));
|
||||
});
|
||||
expect(state.setMode).toHaveBeenCalledWith("ai-gateway");
|
||||
});
|
||||
|
||||
it("switches to a plugin by name when its entry is picked", async () => {
|
||||
state.plugins = [{ name: "litellm-platform-plugin", display_name: "Chat UI", url: "http://localhost:3300" }];
|
||||
state.mode = "ai-gateway";
|
||||
render(<ViewSwitcher />);
|
||||
|
||||
act(() => {
|
||||
fireEvent.click(screen.getByRole("button"));
|
||||
});
|
||||
await waitFor(() => expect(screen.getByText("Chat UI")).toBeInTheDocument());
|
||||
act(() => {
|
||||
fireEvent.click(screen.getByText("Chat UI"));
|
||||
});
|
||||
expect(state.setMode).toHaveBeenCalledWith("litellm-platform-plugin");
|
||||
});
|
||||
});
|
||||
46
ui/litellm-dashboard/src/components/Navbar/ViewSwitcher.tsx
Normal file
46
ui/litellm-dashboard/src/components/Navbar/ViewSwitcher.tsx
Normal file
|
|
@ -0,0 +1,46 @@
|
|||
import React from "react";
|
||||
import { Dropdown } from "antd";
|
||||
import { AppstoreOutlined, CheckOutlined, DownOutlined } from "@ant-design/icons";
|
||||
import type { MenuProps } from "antd";
|
||||
import { usePluginMode } from "@/contexts/PluginModeContext";
|
||||
|
||||
const GATEWAY = "ai-gateway";
|
||||
|
||||
export default function ViewSwitcher() {
|
||||
const { mode, setMode, plugins } = usePluginMode();
|
||||
|
||||
// Only a switcher when there is at least one plugin to switch to.
|
||||
if (plugins.length === 0) return null;
|
||||
|
||||
const activeLabel = plugins.find((p) => p.name === mode)?.display_name ?? "AI Gateway";
|
||||
|
||||
const entries = [
|
||||
{ value: GATEWAY, label: "AI Gateway" },
|
||||
...plugins.map((p) => ({ value: p.name, label: p.display_name })),
|
||||
];
|
||||
|
||||
const items: MenuProps["items"] = entries.map((e) => ({
|
||||
key: e.value,
|
||||
label: (
|
||||
<div className="flex items-center justify-between gap-6 py-0.5">
|
||||
<span className="font-medium">{e.label}</span>
|
||||
{e.value === mode && <CheckOutlined className="text-blue-600" />}
|
||||
</div>
|
||||
),
|
||||
}));
|
||||
|
||||
const onClick: MenuProps["onClick"] = ({ key }) => setMode(key);
|
||||
|
||||
return (
|
||||
<Dropdown menu={{ items, onClick, selectedKeys: [mode] }} trigger={["click"]}>
|
||||
<button
|
||||
type="button"
|
||||
className="flex items-center gap-2 rounded-md border border-gray-200 px-2.5 py-1.5 text-sm font-medium text-gray-700 transition-colors hover:bg-gray-50"
|
||||
>
|
||||
<AppstoreOutlined className="text-gray-500" />
|
||||
<span>{activeLabel}</span>
|
||||
<DownOutlined className="text-[10px] text-gray-400" />
|
||||
</button>
|
||||
</Dropdown>
|
||||
);
|
||||
}
|
||||
|
|
@ -0,0 +1,173 @@
|
|||
"use client";
|
||||
|
||||
import { useState, useEffect } from "react";
|
||||
import { Button, Card, Form, Input, Modal, Space, Table, Typography } from "antd";
|
||||
import { DeleteOutlined, EditOutlined, PlusOutlined } from "@ant-design/icons";
|
||||
import { getConfigFieldSetting, updateConfigFieldSetting } from "@/components/networking";
|
||||
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
||||
|
||||
const { Title, Text, Paragraph } = Typography;
|
||||
|
||||
interface Plugin {
|
||||
name: string;
|
||||
display_name: string;
|
||||
url: string;
|
||||
plugin_key?: string;
|
||||
}
|
||||
|
||||
export default function PluginSettings() {
|
||||
const { accessToken } = useAuthorized();
|
||||
const [plugins, setPlugins] = useState<Plugin[]>([]);
|
||||
const [loading, setLoading] = useState(true);
|
||||
const [saving, setSaving] = useState(false);
|
||||
const [modalOpen, setModalOpen] = useState(false);
|
||||
const [editingIndex, setEditingIndex] = useState<number | null>(null);
|
||||
const [form] = Form.useForm<Plugin>();
|
||||
|
||||
useEffect(() => {
|
||||
if (!accessToken) return;
|
||||
getConfigFieldSetting(accessToken, "plugins")
|
||||
.then((data) => {
|
||||
const val = data?.field_value;
|
||||
setPlugins(Array.isArray(val) ? val : []);
|
||||
})
|
||||
.catch(() => setPlugins([]))
|
||||
.finally(() => setLoading(false));
|
||||
}, [accessToken]);
|
||||
|
||||
const save = async (updated: Plugin[]) => {
|
||||
if (!accessToken) return;
|
||||
setSaving(true);
|
||||
try {
|
||||
await updateConfigFieldSetting(accessToken, "plugins", updated);
|
||||
setPlugins(updated);
|
||||
} finally {
|
||||
setSaving(false);
|
||||
}
|
||||
};
|
||||
|
||||
const openAdd = () => {
|
||||
setEditingIndex(null);
|
||||
form.resetFields();
|
||||
setModalOpen(true);
|
||||
};
|
||||
|
||||
const openEdit = (idx: number) => {
|
||||
setEditingIndex(idx);
|
||||
// plugin_key arrives redacted ("***"); start it blank so an untouched save
|
||||
// keeps the stored credential instead of overwriting it with the placeholder.
|
||||
form.setFieldsValue({ ...plugins[idx], plugin_key: "" });
|
||||
setModalOpen(true);
|
||||
};
|
||||
|
||||
const handleDelete = (idx: number) => {
|
||||
const updated = plugins.filter((_, i) => i !== idx);
|
||||
save(updated);
|
||||
};
|
||||
|
||||
const handleOk = async () => {
|
||||
const values = await form.validateFields();
|
||||
const updated =
|
||||
editingIndex !== null ? plugins.map((p, i) => (i === editingIndex ? values : p)) : [...plugins, values];
|
||||
await save(updated);
|
||||
setModalOpen(false);
|
||||
};
|
||||
|
||||
const columns = [
|
||||
{
|
||||
title: "Name",
|
||||
dataIndex: "name",
|
||||
key: "name",
|
||||
render: (v: string) => <Text code>{v}</Text>,
|
||||
},
|
||||
{ title: "Display Name", dataIndex: "display_name", key: "display_name" },
|
||||
{
|
||||
title: "URL",
|
||||
dataIndex: "url",
|
||||
key: "url",
|
||||
render: (v: string) => (
|
||||
<a href={v} target="_blank" rel="noopener noreferrer">
|
||||
{v}
|
||||
</a>
|
||||
),
|
||||
},
|
||||
{
|
||||
title: "Plugin Key",
|
||||
dataIndex: "plugin_key",
|
||||
key: "plugin_key",
|
||||
render: (v?: string) => (v ? <Text code>{"•".repeat(8)}</Text> : <Text type="secondary">—</Text>),
|
||||
},
|
||||
{
|
||||
title: "Actions",
|
||||
key: "actions",
|
||||
render: (_: unknown, __: Plugin, idx: number) => (
|
||||
<Space>
|
||||
<Button icon={<EditOutlined />} size="small" onClick={() => openEdit(idx)} />
|
||||
<Button icon={<DeleteOutlined />} size="small" danger onClick={() => handleDelete(idx)} />
|
||||
</Space>
|
||||
),
|
||||
},
|
||||
];
|
||||
|
||||
return (
|
||||
<Card>
|
||||
<Title level={4}>Plugins</Title>
|
||||
<Paragraph>
|
||||
Register external services as plugins. Once added, users can toggle to the plugin from the mode switcher in the
|
||||
top-left of the sidebar.
|
||||
</Paragraph>
|
||||
<Paragraph type="secondary" style={{ fontSize: 12 }}>
|
||||
Each plugin must expose <Text code>GET /api/plugin-manifest</Text> returning nav items and capabilities.
|
||||
</Paragraph>
|
||||
|
||||
<Button type="primary" icon={<PlusOutlined />} onClick={openAdd} style={{ marginBottom: 16 }}>
|
||||
Add Plugin
|
||||
</Button>
|
||||
|
||||
<Table dataSource={plugins} columns={columns} rowKey="name" loading={loading} pagination={false} size="small" />
|
||||
|
||||
<Modal
|
||||
title={editingIndex !== null ? "Edit Plugin" : "Add Plugin"}
|
||||
open={modalOpen}
|
||||
onOk={handleOk}
|
||||
onCancel={() => setModalOpen(false)}
|
||||
confirmLoading={saving}
|
||||
okText="Save"
|
||||
>
|
||||
<Form form={form} layout="vertical" style={{ marginTop: 16 }}>
|
||||
<Form.Item
|
||||
name="name"
|
||||
label="Name (identifier)"
|
||||
rules={[{ required: true, message: "Required" }]}
|
||||
extra="Used in URLs and config. No spaces. E.g. litellm-platform-plugin"
|
||||
>
|
||||
<Input placeholder="litellm-platform-plugin" />
|
||||
</Form.Item>
|
||||
<Form.Item name="display_name" label="Display Name" rules={[{ required: true, message: "Required" }]}>
|
||||
<Input placeholder="Agent Control Plane" />
|
||||
</Form.Item>
|
||||
<Form.Item
|
||||
name="url"
|
||||
label="URL"
|
||||
rules={[
|
||||
{ required: true, message: "Required" },
|
||||
{ type: "url", message: "Must be a valid URL" },
|
||||
]}
|
||||
extra="Base URL of the plugin service"
|
||||
>
|
||||
<Input placeholder="https://your-plugin.example.com" />
|
||||
</Form.Item>
|
||||
<Form.Item
|
||||
name="plugin_key"
|
||||
label="Plugin Key"
|
||||
extra="Optional. The plugin's own credential, injected as Authorization: Bearer <key> only when litellm reverse-proxies API calls to the plugin's backend (/plugin-proxy/<name>/*). Leave blank for plugins that use the forwarded litellm user token (e.g. iframe plugins) — that path uses the user's token, not this key."
|
||||
>
|
||||
<Input.Password
|
||||
placeholder={editingIndex !== null ? "Leave blank to keep current key" : "sk-... (optional)"}
|
||||
/>
|
||||
</Form.Item>
|
||||
</Form>
|
||||
</Modal>
|
||||
</Card>
|
||||
);
|
||||
}
|
||||
|
|
@ -16,6 +16,7 @@ import { CommunityEngagementButtons } from "./Navbar/CommunityEngagementButtons/
|
|||
import { NAV_PRODUCT_LINK_CLASS } from "./Navbar/navProductLinkClass";
|
||||
import { NotificationsBell } from "./Navbar/NotificationsBell/NotificationsBell";
|
||||
import UserDropdown from "./Navbar/UserDropdown/UserDropdown";
|
||||
import ViewSwitcher from "./Navbar/ViewSwitcher";
|
||||
import WorkerDropdown from "./Navbar/WorkerDropdown/WorkerDropdown";
|
||||
|
||||
interface NavbarProps {
|
||||
|
|
@ -111,6 +112,12 @@ const Navbar: React.FC<NavbarProps> = ({
|
|||
</div>
|
||||
</div>
|
||||
|
||||
{!isPublicPage && (
|
||||
<div className="ml-4 flex shrink-0 items-center border-l border-gray-200 pl-4">
|
||||
<ViewSwitcher />
|
||||
</div>
|
||||
)}
|
||||
|
||||
<div className="ml-auto flex min-w-0 flex-1 items-center justify-end gap-4">
|
||||
{showWorkerSwitch && (
|
||||
<div className="flex shrink-0 items-center">
|
||||
|
|
|
|||
64
ui/litellm-dashboard/src/contexts/PluginModeContext.test.tsx
Normal file
64
ui/litellm-dashboard/src/contexts/PluginModeContext.test.tsx
Normal file
|
|
@ -0,0 +1,64 @@
|
|||
import { describe, it, expect, vi, beforeEach } from "vitest";
|
||||
import { render, screen, waitFor } from "@testing-library/react";
|
||||
import { PluginModeProvider, usePluginMode } from "./PluginModeContext";
|
||||
import type { Plugin } from "./PluginModeContext";
|
||||
|
||||
const { getMock } = vi.hoisted(() => ({ getMock: vi.fn() }));
|
||||
|
||||
vi.mock("@/lib/http/client", () => ({
|
||||
createApiClient: () => ({ get: getMock }),
|
||||
}));
|
||||
vi.mock("@/components/networking", () => ({ getProxyBaseUrl: () => "" }));
|
||||
|
||||
function ModeProbe() {
|
||||
const { mode, activePlugin } = usePluginMode();
|
||||
return (
|
||||
<div>
|
||||
<span data-testid="mode">{mode}</span>
|
||||
<span data-testid="active">{activePlugin?.name ?? "none"}</span>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
const renderWithPlugins = (plugins: Plugin[]) => {
|
||||
getMock.mockResolvedValueOnce(plugins);
|
||||
return render(
|
||||
<PluginModeProvider accessToken="sk-test">
|
||||
<ModeProbe />
|
||||
</PluginModeProvider>,
|
||||
);
|
||||
};
|
||||
|
||||
describe("PluginModeProvider effectiveMode fallback", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
localStorage.setItem("litellm_plugin_mode", "my-plugin");
|
||||
});
|
||||
|
||||
it("falls back to ai-gateway once an empty plugins list loads", async () => {
|
||||
renderWithPlugins([]);
|
||||
|
||||
await waitFor(() => expect(getMock).toHaveBeenCalled());
|
||||
await waitFor(() => expect(screen.getByTestId("mode").textContent).toBe("ai-gateway"));
|
||||
expect(screen.getByTestId("active").textContent).toBe("none");
|
||||
});
|
||||
|
||||
it("keeps the stored mode when it is still registered", async () => {
|
||||
renderWithPlugins([{ name: "my-plugin", display_name: "My Plugin", url: "https://p.example.com" }]);
|
||||
|
||||
await waitFor(() => expect(screen.getByTestId("active").textContent).toBe("my-plugin"));
|
||||
expect(screen.getByTestId("mode").textContent).toBe("my-plugin");
|
||||
});
|
||||
|
||||
it("falls back to ai-gateway when the plugins fetch fails, never stranding the user", async () => {
|
||||
getMock.mockRejectedValueOnce(new Error("network down"));
|
||||
render(
|
||||
<PluginModeProvider accessToken="sk-test">
|
||||
<ModeProbe />
|
||||
</PluginModeProvider>,
|
||||
);
|
||||
|
||||
await waitFor(() => expect(getMock).toHaveBeenCalled());
|
||||
await waitFor(() => expect(screen.getByTestId("mode").textContent).toBe("ai-gateway"));
|
||||
});
|
||||
});
|
||||
95
ui/litellm-dashboard/src/contexts/PluginModeContext.tsx
Normal file
95
ui/litellm-dashboard/src/contexts/PluginModeContext.tsx
Normal file
|
|
@ -0,0 +1,95 @@
|
|||
"use client";
|
||||
|
||||
import React, { createContext, useContext, useState, useEffect } from "react";
|
||||
import { createApiClient } from "@/lib/http/client";
|
||||
import { getProxyBaseUrl } from "@/components/networking";
|
||||
|
||||
export type PluginMode = "ai-gateway" | string; // "ai-gateway" or a registered plugin name
|
||||
|
||||
export interface PluginNavItem {
|
||||
key: string;
|
||||
label: string;
|
||||
icon?: string;
|
||||
path: string;
|
||||
badge?: boolean;
|
||||
}
|
||||
|
||||
export interface Plugin {
|
||||
name: string;
|
||||
display_name: string;
|
||||
url: string;
|
||||
plugin_key?: string;
|
||||
nav_items?: PluginNavItem[];
|
||||
capabilities?: string[];
|
||||
}
|
||||
|
||||
interface PluginModeContextValue {
|
||||
mode: PluginMode;
|
||||
setMode: (mode: PluginMode) => void;
|
||||
plugins: Plugin[];
|
||||
activePlugin: Plugin | null;
|
||||
}
|
||||
|
||||
const PluginModeContext = createContext<PluginModeContextValue>({
|
||||
mode: "ai-gateway",
|
||||
setMode: () => {},
|
||||
plugins: [],
|
||||
activePlugin: null,
|
||||
});
|
||||
|
||||
const STORAGE_KEY = "litellm_plugin_mode";
|
||||
const pluginApiClient = createApiClient({ getBaseUrl: () => getProxyBaseUrl() ?? "" });
|
||||
|
||||
function readStoredMode(): PluginMode {
|
||||
if (typeof window === "undefined") return "ai-gateway";
|
||||
return localStorage.getItem(STORAGE_KEY) ?? "ai-gateway";
|
||||
}
|
||||
|
||||
interface PluginModeProviderProps {
|
||||
children: React.ReactNode;
|
||||
/** Pass the current access token from the app's auth context. */
|
||||
accessToken?: string | null;
|
||||
}
|
||||
|
||||
export function PluginModeProvider({ children, accessToken }: PluginModeProviderProps) {
|
||||
const [mode, setModeState] = useState<PluginMode>(readStoredMode);
|
||||
const [plugins, setPlugins] = useState<Plugin[]>([]);
|
||||
const [loaded, setLoaded] = useState(false);
|
||||
|
||||
useEffect(() => {
|
||||
// Re-fetch whenever the auth token changes (handles login/logout cycles)
|
||||
if (!accessToken) return;
|
||||
pluginApiClient
|
||||
.get("/api/plugins", { accessToken })
|
||||
.then((data: Plugin[]) => {
|
||||
setPlugins(Array.isArray(data) ? data : []);
|
||||
})
|
||||
.catch(() => {})
|
||||
// Mark loaded even on failure so a stored plugin mode still falls back to
|
||||
// ai-gateway; otherwise a failed fetch would strand the user on a blank
|
||||
// plugin view with no switcher to escape.
|
||||
.finally(() => setLoaded(true));
|
||||
}, [accessToken]);
|
||||
|
||||
// Once plugins have loaded, fall back to ai-gateway if the persisted mode is
|
||||
// no longer registered — including when the list came back empty (all plugins
|
||||
// removed). Derived rather than setState-in-effect to avoid cascading renders.
|
||||
const effectiveMode = mode !== "ai-gateway" && loaded && !plugins.some((p) => p.name === mode) ? "ai-gateway" : mode;
|
||||
|
||||
const setMode = (m: PluginMode) => {
|
||||
setModeState(m);
|
||||
localStorage.setItem(STORAGE_KEY, m);
|
||||
};
|
||||
|
||||
const activePlugin = plugins.find((p) => p.name === effectiveMode) ?? null;
|
||||
|
||||
return (
|
||||
<PluginModeContext.Provider value={{ mode: effectiveMode, setMode, plugins, activePlugin }}>
|
||||
{children}
|
||||
</PluginModeContext.Provider>
|
||||
);
|
||||
}
|
||||
|
||||
export function usePluginMode() {
|
||||
return useContext(PluginModeContext);
|
||||
}
|
||||
464
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
464
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -515,6 +515,60 @@ export interface paths {
|
|||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/api/plugins": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
/**
|
||||
* List Plugins
|
||||
* @description 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.
|
||||
*/
|
||||
get: operations["list_plugins_api_plugins_get"];
|
||||
put?: never;
|
||||
post?: never;
|
||||
delete?: never;
|
||||
options?: never;
|
||||
head?: never;
|
||||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/api/plugins/auth-token": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
/**
|
||||
* Plugin Auth Token
|
||||
* @description 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.
|
||||
*/
|
||||
get: operations["plugin_auth_token_api_plugins_auth_token_get"];
|
||||
put?: never;
|
||||
post?: never;
|
||||
delete?: never;
|
||||
options?: never;
|
||||
head?: never;
|
||||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/apply_guardrail": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
@ -8945,6 +8999,106 @@ export interface paths {
|
|||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/plugin-proxy/{plugin_name}/{path}": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
/**
|
||||
* Plugin Proxy
|
||||
* @description 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.
|
||||
*/
|
||||
get: operations["plugin_proxy_plugin_proxy__plugin_name___path__get"];
|
||||
/**
|
||||
* Plugin Proxy
|
||||
* @description 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.
|
||||
*/
|
||||
put: operations["plugin_proxy_plugin_proxy__plugin_name___path__put"];
|
||||
/**
|
||||
* Plugin Proxy
|
||||
* @description 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.
|
||||
*/
|
||||
post: operations["plugin_proxy_plugin_proxy__plugin_name___path__post"];
|
||||
/**
|
||||
* Plugin Proxy
|
||||
* @description 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.
|
||||
*/
|
||||
delete: operations["plugin_proxy_plugin_proxy__plugin_name___path__delete"];
|
||||
/**
|
||||
* Plugin Proxy
|
||||
* @description 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.
|
||||
*/
|
||||
options: operations["plugin_proxy_plugin_proxy__plugin_name___path__options"];
|
||||
/**
|
||||
* Plugin Proxy
|
||||
* @description 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.
|
||||
*/
|
||||
head: operations["plugin_proxy_plugin_proxy__plugin_name___path__head"];
|
||||
/**
|
||||
* Plugin Proxy
|
||||
* @description 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.
|
||||
*/
|
||||
patch: operations["plugin_proxy_plugin_proxy__plugin_name___path__patch"];
|
||||
trace?: never;
|
||||
};
|
||||
"/policies": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
@ -22216,6 +22370,11 @@ export interface components {
|
|||
* @description Default upstream request timeout in seconds for native and custom pass-through endpoints that use pass_through_request. Defaults to 600 when unset.
|
||||
*/
|
||||
pass_through_request_timeout?: number | null;
|
||||
/**
|
||||
* Plugins
|
||||
* @description external services registered as embeddable UI plugins
|
||||
*/
|
||||
plugins?: components["schemas"]["PluginConfig"][] | null;
|
||||
/**
|
||||
* Reject Clientside Metadata Tags
|
||||
* @description When set to True, rejects requests that contain client-side 'metadata.tags' to prevent users from influencing budgets by sending different tags. Tags can only be inherited from the API key metadata.
|
||||
|
|
@ -28185,6 +28344,32 @@ export interface components {
|
|||
*/
|
||||
name: string;
|
||||
};
|
||||
/**
|
||||
* PluginConfig
|
||||
* @description A single external service registered as an embeddable UI plugin.
|
||||
*/
|
||||
PluginConfig: {
|
||||
/**
|
||||
* Display Name
|
||||
* @description human-readable label shown in the UI view switcher
|
||||
*/
|
||||
display_name?: string | null;
|
||||
/**
|
||||
* Name
|
||||
* @description unique plugin identifier (kebab-case)
|
||||
*/
|
||||
name: string;
|
||||
/**
|
||||
* Plugin Key
|
||||
* @description plugin's own credential, injected as Bearer auth only on /plugin-proxy/<name>/* reverse-proxy calls
|
||||
*/
|
||||
plugin_key?: string | null;
|
||||
/**
|
||||
* Url
|
||||
* @description base URL of the plugin service
|
||||
*/
|
||||
url: string;
|
||||
};
|
||||
/**
|
||||
* PluginListItem
|
||||
* @description Plugin item in list responses.
|
||||
|
|
@ -33488,6 +33673,61 @@ export interface operations {
|
|||
};
|
||||
};
|
||||
};
|
||||
list_plugins_api_plugins_get: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
requestBody?: never;
|
||||
responses: {
|
||||
/** @description Successful Response */
|
||||
200: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": {
|
||||
[key: string]: string;
|
||||
}[];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
plugin_auth_token_api_plugins_auth_token_get: {
|
||||
parameters: {
|
||||
query?: {
|
||||
plugin_name?: string;
|
||||
};
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
requestBody?: never;
|
||||
responses: {
|
||||
/** @description Successful Response */
|
||||
200: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": {
|
||||
[key: string]: unknown;
|
||||
};
|
||||
};
|
||||
};
|
||||
/** @description Validation Error */
|
||||
422: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["HTTPValidationError"];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
apply_guardrail_apply_guardrail_post: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
@ -44501,6 +44741,230 @@ export interface operations {
|
|||
};
|
||||
};
|
||||
};
|
||||
plugin_proxy_plugin_proxy__plugin_name___path__get: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path: {
|
||||
plugin_name: string;
|
||||
path: string;
|
||||
};
|
||||
cookie?: never;
|
||||
};
|
||||
requestBody?: never;
|
||||
responses: {
|
||||
/** @description Successful Response */
|
||||
200: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": unknown;
|
||||
};
|
||||
};
|
||||
/** @description Validation Error */
|
||||
422: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["HTTPValidationError"];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
plugin_proxy_plugin_proxy__plugin_name___path__put: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path: {
|
||||
plugin_name: string;
|
||||
path: string;
|
||||
};
|
||||
cookie?: never;
|
||||
};
|
||||
requestBody?: never;
|
||||
responses: {
|
||||
/** @description Successful Response */
|
||||
200: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": unknown;
|
||||
};
|
||||
};
|
||||
/** @description Validation Error */
|
||||
422: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["HTTPValidationError"];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
plugin_proxy_plugin_proxy__plugin_name___path__post: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path: {
|
||||
plugin_name: string;
|
||||
path: string;
|
||||
};
|
||||
cookie?: never;
|
||||
};
|
||||
requestBody?: never;
|
||||
responses: {
|
||||
/** @description Successful Response */
|
||||
200: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": unknown;
|
||||
};
|
||||
};
|
||||
/** @description Validation Error */
|
||||
422: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["HTTPValidationError"];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
plugin_proxy_plugin_proxy__plugin_name___path__delete: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path: {
|
||||
plugin_name: string;
|
||||
path: string;
|
||||
};
|
||||
cookie?: never;
|
||||
};
|
||||
requestBody?: never;
|
||||
responses: {
|
||||
/** @description Successful Response */
|
||||
200: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": unknown;
|
||||
};
|
||||
};
|
||||
/** @description Validation Error */
|
||||
422: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["HTTPValidationError"];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
plugin_proxy_plugin_proxy__plugin_name___path__options: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path: {
|
||||
plugin_name: string;
|
||||
path: string;
|
||||
};
|
||||
cookie?: never;
|
||||
};
|
||||
requestBody?: never;
|
||||
responses: {
|
||||
/** @description Successful Response */
|
||||
200: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": unknown;
|
||||
};
|
||||
};
|
||||
/** @description Validation Error */
|
||||
422: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["HTTPValidationError"];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
plugin_proxy_plugin_proxy__plugin_name___path__head: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path: {
|
||||
plugin_name: string;
|
||||
path: string;
|
||||
};
|
||||
cookie?: never;
|
||||
};
|
||||
requestBody?: never;
|
||||
responses: {
|
||||
/** @description Successful Response */
|
||||
200: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": unknown;
|
||||
};
|
||||
};
|
||||
/** @description Validation Error */
|
||||
422: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["HTTPValidationError"];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
plugin_proxy_plugin_proxy__plugin_name___path__patch: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path: {
|
||||
plugin_name: string;
|
||||
path: string;
|
||||
};
|
||||
cookie?: never;
|
||||
};
|
||||
requestBody?: never;
|
||||
responses: {
|
||||
/** @description Successful Response */
|
||||
200: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": unknown;
|
||||
};
|
||||
};
|
||||
/** @description Validation Error */
|
||||
422: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["HTTPValidationError"];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
create_policy_policies_post: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue