Merge remote-tracking branch 'origin/main' into litellm_batch_update_database_tasks

This commit is contained in:
Ryan Crabbe 2026-02-25 12:12:10 -08:00
commit e1e05d27f8
73 changed files with 5657 additions and 262 deletions

View file

@ -203,7 +203,7 @@ curl -X POST 'http://0.0.0.0:4000/v1/chat/completions' \
{
"mcpServers": {
"LiteLLM": {
"url": "http://localhost:4000/mcp",
"url": "http://localhost:4000/mcp/",
"headers": {
"x-litellm-api-key": "Bearer sk-1234"
}

View file

@ -0,0 +1,145 @@
---
slug: gpt_5_3_codex
title: "Day 0 Support: GPT-5.3-Codex"
date: 2026-02-24T10:00:00
authors:
- name: Sameer Kankute
title: SWE @ LiteLLM (LLM Translation)
url: https://www.linkedin.com/in/sameer-kankute/
image_url: https://pbs.twimg.com/profile_images/2001352686994907136/ONgNuSk5_400x400.jpg
- name: Krrish Dholakia
title: "CEO, LiteLLM"
url: https://www.linkedin.com/in/krish-d/
image_url: https://pbs.twimg.com/profile_images/1298587542745358340/DZv3Oj-h_400x400.jpg
- name: Ishaan Jaff
title: "CTO, LiteLLM"
url: https://www.linkedin.com/in/reffajnaahsi/
image_url: https://pbs.twimg.com/profile_images/1613813310264340481/lz54oEiB_400x400.jpg
description: "Day 0 support for GPT-5.3-Codex on LiteLLM, including phase parameter handling for Responses API."
tags: [openai, gpt-5.3-codex, codex, day 0 support]
hide_table_of_contents: false
---
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
LiteLLM now supports GPT-5.3-Codex on Day 0, including support for the new assistant `phase` metadata on Responses API output items.
## Why `phase` matters for GPT-5.3-Codex
`phase` appears on assistant output items and helps distinguish preamble/commentary turns from final closeout responses.
Reference: [Phase parameter docs](https://developers.openai.com/api/reference/overview)
Supported values:
- `null`
- `"commentary"`
- `"final_answer"`
Important:
- Persist assistant output items with `phase` exactly as returned.
- Send those assistant items back on the next turn.
- Do **not** add `phase` to user messages.
## Docker Image
```bash
docker pull ghcr.io/berriai/litellm:v1.81.12-stable.gpt-5.3
```
## Usage
<Tabs>
<TabItem value="proxy" label="LiteLLM Proxy">
**1. Setup config.yaml**
```yaml
model_list:
- model_name: gpt-5.3-codex
litellm_params:
model: openai/gpt-5.3-codex
```
**2. Start the proxy**
```bash
docker run -d \
-p 4000:4000 \
-e ANTHROPIC_API_KEY=$OPENAI_API_KEY \
-v $(pwd)/config.yaml:/app/config.yaml \
ghcr.io/berriai/litellm:v1.81.12-stable.gpt-5.3 \
--config /app/config.yaml
```
**3. Test it**
```bash
curl -X POST "http://0.0.0.0:4000/v1/responses" \
-H "Content-Type: application/json" \
-H "Authorization: Bearer $LITELLM_KEY" \
-d '{
"model": "gpt-5.3-codex",
"input": "Write a Python script that checks if a number is prime."
}'
```
</TabItem>
</Tabs>
## Python Example: Persist `phase` with OpenAI Client + LiteLLM Base URL
```python
from openai import OpenAI
client = OpenAI(
base_url="http://0.0.0.0:4000/v1", # LiteLLM Proxy
api_key="your-litellm-api-key",
)
items = [] # Persist this per conversation/thread
def _item_get(item, key, default=None):
if isinstance(item, dict):
return item.get(key, default)
return getattr(item, key, default)
def run_turn(user_text: str):
global items
# User message: no phase field
items.append(
{
"type": "message",
"role": "user",
"content": [{"type": "input_text", "text": user_text}],
}
)
resp = client.responses.create(
model="gpt-5.3-codex",
input=items,
)
# Persist assistant output items verbatim, including phase
for out_item in (resp.output or []):
items.append(out_item)
# Optional: inspect latest phase for UI/telemetry routing
latest_phase = None
for out_item in reversed(resp.output or []):
if _item_get(out_item, "type") == "output_item.done" and _item_get(out_item, "phase") is not None:
latest_phase = _item_get(out_item, "phase")
break
return resp, latest_phase
```
## Notes
- Use `/v1/responses` for GPT Codex models.
- Preserve full assistant output history for best multi-turn behavior.
- If `phase` metadata is dropped during history reconstruction, output quality can degrade on long-running tasks.

View file

@ -641,7 +641,7 @@ import asyncio
config = {
"mcpServers": {
"mcp_group": {
"url": "http://localhost:4000/mcp",
"url": "http://localhost:4000/mcp/",
"headers": {
"x-mcp-servers": "dev_group", # assume this gives access to github, zapier and deepwiki
"x-litellm-api-key": "Bearer sk-1234",

View file

@ -122,7 +122,7 @@ Use this to track overall LiteLLM Proxy usage.
| Metric Name | Description |
|----------------------|--------------------------------------|
| `litellm_proxy_failed_requests_metric` | Total number of failed responses from proxy - the client did not get a success response from litellm proxy. Labels: `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "user_email", "exception_status", "exception_class", "route", "model_id"` |
| `litellm_proxy_total_requests_metric` | Total number of requests made to the proxy server - track number of client side requests. Labels: `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "status_code", "user_email", "route", "model_id"` |
| `litellm_proxy_total_requests_metric` | Total number of requests made to the proxy server - track number of client side requests. Labels: `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "status_code", "user_email", "route", "model_id"`. Optionally includes `"stream"` — see [Emit Stream Label](#emit-stream-label). |
### Callback Logging Metrics
@ -214,9 +214,31 @@ litellm_settings:
```
### Emit Stream Label
Add a `stream` label to `litellm_proxy_total_requests_metric` to split requests by streaming vs. non-streaming. Disabled by default.
```yaml title="config.yaml"
litellm_settings:
callbacks: ["prometheus"]
prometheus_emit_stream_label: true
```
When enabled, `litellm_proxy_total_requests_metric` gains a `stream` label with values `"True"`, `"False"`, or `"None"`.
```
litellm_proxy_total_requests_metric{..., stream="True"} 42
litellm_proxy_total_requests_metric{..., stream="False"} 100
```
:::note
This label is opt-in because adding a new label to an existing metric changes its cardinality and breaks existing Prometheus queries / Grafana dashboards that target this metric. Enable it only on fresh deployments or when you are ready to update your dashboards.
:::
## [BETA] Custom Metrics
Track custom metrics on prometheus on all events mentioned above.
Track custom metrics on prometheus on all events mentioned above.
### Custom Metadata Labels

View file

@ -1,5 +1,5 @@
---
title: "v1.81.12-stable - Guardrail Policy Templates & Action Builder"
title: "v1.81.12-stable.1 - Guardrail Policy Templates & Action Builder"
slug: "v1-81-12"
date: 2026-02-14T00:00:00
authors:
@ -27,7 +27,7 @@ import Image from '@theme/IdealImage';
docker run \
-e STORE_MODEL_IN_DB=True \
-p 4000:4000 \
ghcr.io/berriai/litellm:main-v1.81.12-stable
ghcr.io/berriai/litellm:main-v1.81.12-stable.1
```
</TabItem>

View file

@ -374,6 +374,7 @@ enable_end_user_cost_tracking_prometheus_only: Optional[bool] = None
custom_prometheus_metadata_labels: List[str] = []
custom_prometheus_tags: List[str] = []
prometheus_metrics_config: Optional[List] = None
prometheus_emit_stream_label: bool = False
disable_add_prefix_to_prompt: bool = (
False # used by anthropic, to disable adding prefix to prompt
)

View file

@ -244,6 +244,7 @@ REDIS_DAILY_TAG_SPEND_UPDATE_BUFFER_KEY = "litellm_daily_tag_spend_update_buffer
MAX_REDIS_BUFFER_DEQUEUE_COUNT = int(os.getenv("MAX_REDIS_BUFFER_DEQUEUE_COUNT", 100))
# Bounds asyncio.Queue() instances (log queues, spend update queues, etc.) to prevent unbounded memory growth
LITELLM_ASYNCIO_QUEUE_MAXSIZE = int(os.getenv("LITELLM_ASYNCIO_QUEUE_MAXSIZE", 1000))
TOOL_POLICY_CACHE_TTL_SECONDS = int(os.getenv("TOOL_POLICY_CACHE_TTL_SECONDS", 60))
# Aggregation threshold: default to 80% of the asyncio queue maxsize so the check can always trigger.
# Must be < LITELLM_ASYNCIO_QUEUE_MAXSIZE; if set higher the aggregation logic will never fire.
MAX_SIZE_IN_MEMORY_QUEUE = int(

View file

@ -974,6 +974,9 @@ class PrometheusLogger(CustomLogger):
),
client_ip=standard_logging_payload["metadata"].get("requester_ip_address"),
user_agent=standard_logging_payload["metadata"].get("user_agent"),
stream=str(standard_logging_payload.get("stream"))
if litellm.prometheus_emit_stream_label
else None,
)
if (
@ -1624,6 +1627,9 @@ class PrometheusLogger(CustomLogger):
client_ip=_metadata.get("requester_ip_address"),
user_agent=_metadata.get("user_agent"),
model_id=model_id,
stream=str(request_data.get("stream"))
if litellm.prometheus_emit_stream_label
else None,
)
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(

View file

@ -11,6 +11,7 @@ from typing import Any, List, Optional
import httpx
from litellm.types.utils import Usage
from litellm.llms.bedrock.chat.invoke_transformations.amazon_qwen3_transformation import (
AmazonQwen3Config,
)
@ -79,10 +80,11 @@ class AmazonQwen2Config(AmazonQwen3Config):
# Set usage information if available in response
if "usage" in response_data:
usage_data = response_data["usage"]
if hasattr(model_response, 'usage'):
model_response.usage.prompt_tokens = usage_data.get("prompt_tokens", 0)
model_response.usage.completion_tokens = usage_data.get("completion_tokens", 0)
model_response.usage.total_tokens = usage_data.get("total_tokens", 0)
model_response.usage = Usage(
prompt_tokens=usage_data.get("prompt_tokens", 0),
completion_tokens=usage_data.get("completion_tokens", 0),
total_tokens=usage_data.get("total_tokens", 0),
)
return model_response

View file

@ -10,6 +10,7 @@ from typing import Any, List, Optional
import httpx
from litellm.types.utils import Usage
from litellm.llms.base_llm.chat.transformation import BaseConfig
from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import (
AmazonInvokeConfig,
@ -201,10 +202,11 @@ class AmazonQwen3Config(AmazonInvokeConfig, BaseConfig):
# Set usage information if available in response
if "usage" in response_data:
usage_data = response_data["usage"]
if hasattr(model_response, 'usage'):
model_response.usage.prompt_tokens = usage_data.get("prompt_tokens", 0)
model_response.usage.completion_tokens = usage_data.get("completion_tokens", 0)
model_response.usage.total_tokens = usage_data.get("total_tokens", 0)
model_response.usage = Usage(
prompt_tokens=usage_data.get("prompt_tokens", 0),
completion_tokens=usage_data.get("completion_tokens", 0),
total_tokens=usage_data.get("total_tokens", 0),
)
return model_response

View file

@ -131,9 +131,7 @@ class OpenAIOSeriesConfig(OpenAIGPTConfig):
def is_model_o_series_model(self, model: str) -> bool:
model = model.split("/")[-1] # could be "openai/o3" or "o3"
return model in litellm.open_ai_chat_completion_models and any(
model.startswith(pfx) for pfx in ("o1", "o3", "o4")
)
return model.startswith(("o1", "o3", "o4")) and model in litellm.open_ai_chat_completion_models
@overload
def _transform_messages(
@ -173,4 +171,4 @@ class OpenAIOSeriesConfig(OpenAIGPTConfig):
else:
return super()._transform_messages(
messages, model, is_async=cast(Literal[False], False)
)
)

View file

@ -20562,6 +20562,39 @@
"supports_tool_choice": true,
"supports_vision": true
},
"gpt-5.3-codex": {
"cache_read_input_token_cost": 1.75e-07,
"cache_read_input_token_cost_priority": 3.5e-07,
"input_cost_per_token": 1.75e-06,
"input_cost_per_token_priority": 3.5e-06,
"litellm_provider": "openai",
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "responses",
"output_cost_per_token": 1.4e-05,
"output_cost_per_token_priority": 2.8e-05,
"supported_endpoints": [
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": false,
"supports_tool_choice": true,
"supports_vision": true
},
"gpt-5-mini": {
"cache_read_input_token_cost": 2.5e-08,
"cache_read_input_token_cost_flex": 1.25e-08,

View file

@ -10,6 +10,7 @@ from litellm.proxy._experimental.mcp_server.ui_session_utils import (
)
from litellm.proxy._experimental.mcp_server.utils import merge_mcp_headers
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.types.mcp import MCPAuth
@ -685,7 +686,7 @@ if MCP_AVAILABLE:
return await _execute_with_mcp_client(
new_mcp_server_request,
_test_connection_operation,
raw_headers=dict(request.headers),
raw_headers=_safe_get_request_headers(request),
)
@router.post("/test/tools/list")
@ -744,5 +745,5 @@ if MCP_AVAILABLE:
_list_tools_operation,
mcp_auth_header=mcp_auth_header,
oauth2_headers=oauth2_headers,
raw_headers=dict(request.headers),
raw_headers=_safe_get_request_headers(request),
)

View file

@ -183,6 +183,7 @@ class LitellmTableNames(str, enum.Enum):
KEY_TABLE_NAME = "LiteLLM_VerificationToken"
PROXY_MODEL_TABLE_NAME = "LiteLLM_ProxyModelTable"
MANAGED_FILE_TABLE_NAME = "LiteLLM_ManagedFileTable"
TOOL_TABLE_NAME = "LiteLLM_ToolTable"
class Litellm_EntityType(enum.Enum):
@ -4123,6 +4124,15 @@ class SpendUpdateQueueItem(TypedDict, total=False):
response_cost: Optional[float]
class ToolDiscoveryQueueItem(TypedDict, total=False):
tool_name: str
origin: Optional[str] # MCP server name or "user_defined"
created_by: Optional[str]
key_hash: Optional[str] # hash of virtual key that triggered discovery
team_id: Optional[str] # team that triggered discovery
key_alias: Optional[str] # human-readable key alias
class LiteLLM_ManagedFileTable(LiteLLMPydanticObjectBase):
unified_file_id: str
file_object: Optional[OpenAIFileObject] = None

View file

@ -483,7 +483,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
parent_otel_span = (
open_telemetry_logger.create_litellm_proxy_request_started_span(
start_time=start_time,
headers=dict(request.headers),
headers=_safe_get_request_headers(request),
)
)
@ -562,7 +562,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
parent_otel_span=parent_otel_span,
request_headers=dict(request.headers),
request_headers=_safe_get_request_headers(request),
)
is_proxy_admin = result["is_proxy_admin"]

View file

@ -135,17 +135,29 @@ def _safe_set_request_parsed_body(
def _safe_get_request_headers(request: Optional[Request]) -> dict:
"""
[Non-Blocking] Safely get the request headers
[Non-Blocking] Safely get the request headers.
Caches the result on request.state to avoid re-creating dict(request.headers) per call.
Warning: Callers must NOT mutate the returned dict — it is shared across
all callers within the same request via the cache.
"""
if request is None:
return {}
cached = getattr(request.state, "_cached_headers", None)
if cached is not None:
return cached
try:
if request is None:
return {}
return dict(request.headers)
headers = dict(request.headers)
except Exception as e:
verbose_proxy_logger.debug(
"Unexpected error reading request headers - {}".format(e)
)
return {}
headers = {}
try:
request.state._cached_headers = headers
except Exception:
pass # request.state may not be available in all contexts
return headers
def check_file_size_under_limit(

View file

@ -3,6 +3,7 @@ from fastapi_sso.sso.base import OpenID
from litellm._logging import verbose_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
class CustomSSOLoginHandler(CustomLogger):
@ -18,7 +19,7 @@ class CustomSSOLoginHandler(CustomLogger):
self,
request: Request,
) -> OpenID:
request_headers_dict = dict(request.headers)
request_headers_dict = _safe_get_request_headers(request)
verbose_logger.debug("inside custom ui sso sign in hook...")
return OpenID(
id=request_headers_dict.get("x-litellm-user-id") or "123",

View file

@ -13,7 +13,17 @@ import random
import time
import traceback
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, cast, overload
from typing import (
TYPE_CHECKING,
Any,
Dict,
List,
Literal,
Optional,
Union,
cast,
overload,
)
import litellm
from litellm._logging import verbose_proxy_logger
@ -23,18 +33,19 @@ from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
from litellm.proxy._types import (
DB_CONNECTION_ERROR_TYPES,
BaseDailySpendTransaction,
DailyTagSpendTransaction,
DailyOrganizationSpendTransaction,
DailyTeamSpendTransaction,
DailyEndUserSpendTransaction,
DailyUserSpendTransaction,
DailyAgentSpendTransaction,
DailyEndUserSpendTransaction,
DailyOrganizationSpendTransaction,
DailyTagSpendTransaction,
DailyTeamSpendTransaction,
DailyUserSpendTransaction,
DBSpendUpdateTransactions,
Litellm_EntityType,
LiteLLM_UserTable,
SpendLogsMetadata,
SpendLogsPayload,
SpendUpdateQueueItem,
ToolDiscoveryQueueItem,
)
from litellm.proxy.db.db_transaction_queue.daily_spend_update_queue import (
DailySpendUpdateQueue,
@ -42,6 +53,9 @@ from litellm.proxy.db.db_transaction_queue.daily_spend_update_queue import (
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager
from litellm.proxy.db.db_transaction_queue.redis_update_buffer import RedisUpdateBuffer
from litellm.proxy.db.db_transaction_queue.spend_update_queue import SpendUpdateQueue
from litellm.proxy.db.db_transaction_queue.tool_discovery_queue import (
ToolDiscoveryQueue,
)
from litellm.proxy.route_llm_request import ROUTE_ENDPOINT_MAPPING
if TYPE_CHECKING:
@ -67,6 +81,7 @@ class DBSpendUpdateWriter:
self.redis_update_buffer = RedisUpdateBuffer(redis_cache=self.redis_cache)
self.pod_lock_manager = PodLockManager()
self.spend_update_queue = SpendUpdateQueue()
self.tool_discovery_queue = ToolDiscoveryQueue()
self.daily_spend_update_queue = DailySpendUpdateQueue()
self.daily_team_spend_update_queue = DailySpendUpdateQueue()
self.daily_end_user_spend_update_queue = DailySpendUpdateQueue()
@ -165,10 +180,124 @@ class DBSpendUpdateWriter:
)
)
self._enqueue_tool_registry_upsert(
kwargs=kwargs,
completion_response=completion_response,
hashed_token=hashed_token,
team_id=team_id,
)
verbose_proxy_logger.debug("Runs spend update on all tables")
except Exception:
verbose_proxy_logger.error(
"Spend tracking - update_database failed. Spend log insertion or daily transaction enqueue "
"may not have completed for this request. "
"response_cost=%s, token=%s, user_id=%s, team_id=%s, org_id=%s, end_user_id=%s - %s",
response_cost,
token,
user_id,
team_id,
org_id,
end_user_id,
traceback.format_exc(),
)
def _enqueue_tool_registry_upsert(
self,
kwargs: Optional[dict],
completion_response: Optional[Any],
hashed_token: Optional[str] = None,
team_id: Optional[str] = None,
) -> None:
"""
Extract tool names from the LLM request and response and enqueue them
for upsert into LiteLLM_ToolTable via ToolDiscoveryQueue.
Handles four sources:
- MCP tools: standard_logging_object.mcp_tool_call_metadata.namespaced_tool_name
- Response tool_calls (OpenAI / Anthropic pass-through converted to OpenAI format):
completion_response.choices[].message.tool_calls[].function.name
- Request tools array (OpenAI format): kwargs["tools"][].function.name
- Request tools array (Anthropic /messages format): kwargs["passthrough_logging_payload"]
["request_body"]["tools"][].name
"""
try:
if kwargs is None:
return
# Extract key_alias from kwargs metadata if available
key_alias: Optional[str] = None
_litellm_params = kwargs.get("litellm_params") or {}
_metadata = _litellm_params.get("metadata") or {}
key_alias = _metadata.get("user_api_key_alias") or None
def _enqueue(tool_name: str, origin: str = "user_defined") -> None:
self.tool_discovery_queue.add_update(
ToolDiscoveryQueueItem(
tool_name=tool_name,
origin=origin,
key_hash=hashed_token,
team_id=team_id,
key_alias=key_alias,
)
)
# --- MCP tool calls ---
sl_object = kwargs.get("standard_logging_object")
if sl_object is not None:
mcp_metadata = (
sl_object.get("metadata", {}) or {}
).get("mcp_tool_call_metadata")
if mcp_metadata and isinstance(mcp_metadata, dict):
tool_name = mcp_metadata.get("namespaced_tool_name") or mcp_metadata.get("name")
mcp_server_name = mcp_metadata.get("mcp_server_name")
if tool_name:
_enqueue(tool_name, origin=mcp_server_name or "user_defined")
# --- Tools from request body (OpenAI format: tools[].function.name) ---
request_tools = kwargs.get("tools") or []
for tool_def in request_tools:
if not isinstance(tool_def, dict):
continue
fn = tool_def.get("function") or {}
name = fn.get("name") if isinstance(fn, dict) else None
if name:
_enqueue(name)
# --- Tools from Anthropic /messages pass-through request body
# (Anthropic format: tools[].name, no "function" wrapper) ---
passthrough_payload = kwargs.get("passthrough_logging_payload") or {}
request_body = (
passthrough_payload.get("request_body")
if isinstance(passthrough_payload, dict)
else None
) or {}
for tool_def in request_body.get("tools") or []:
if not isinstance(tool_def, dict):
continue
name = tool_def.get("name")
if name:
_enqueue(name)
# --- Response tool_calls (OpenAI format; Anthropic pass-through converts tool_use here) ---
if completion_response is not None and hasattr(completion_response, "choices"):
for choice in completion_response.choices or []:
message = getattr(choice, "message", None)
if message is None:
continue
tool_calls = getattr(message, "tool_calls", None)
if not tool_calls:
continue
for tc in tool_calls:
fn = getattr(tc, "function", None)
if fn is None:
continue
tool_name = getattr(fn, "name", None)
if tool_name:
_enqueue(tool_name)
except Exception as e:
verbose_proxy_logger.debug(
f"Error updating Prisma database: {traceback.format_exc()}"
"_enqueue_tool_registry_upsert error (non-blocking): %s", e
)
async def _batch_database_updates(
@ -389,9 +518,14 @@ class DBSpendUpdateWriter:
)
)
except Exception as e:
verbose_proxy_logger.debug(
"\033[91m"
+ f"Update User DB call failed to execute {str(e)}\n{traceback.format_exc()}"
verbose_proxy_logger.error(
"Spend tracking - failed to enqueue user spend update. "
"user_id=%s, end_user_id=%s, response_cost=%s - %s\n%s",
user_id,
end_user_id,
response_cost,
str(e),
traceback.format_exc(),
)
async def _update_team_db(
@ -428,11 +562,24 @@ class DBSpendUpdateWriter:
response_cost=response_cost,
)
)
except Exception:
pass
except Exception as e:
verbose_proxy_logger.error(
"Spend tracking - failed to enqueue team member spend update. "
"team_id=%s, user_id=%s, response_cost=%s - %s\n%s",
team_id,
user_id,
response_cost,
str(e),
traceback.format_exc(),
)
except Exception as e:
verbose_proxy_logger.debug(
f"Update Team DB failed to execute - {str(e)}\n{traceback.format_exc()}"
verbose_proxy_logger.error(
"Spend tracking - failed to enqueue team spend update. "
"team_id=%s, response_cost=%s - %s\n%s",
team_id,
response_cost,
str(e),
traceback.format_exc(),
)
raise e
@ -457,8 +604,13 @@ class DBSpendUpdateWriter:
)
)
except Exception as e:
verbose_proxy_logger.debug(
f"Update Org DB failed to execute - {str(e)}\n{traceback.format_exc()}"
verbose_proxy_logger.error(
"Spend tracking - failed to enqueue org spend update. "
"org_id=%s, response_cost=%s - %s\n%s",
org_id,
response_cost,
str(e),
traceback.format_exc(),
)
raise e
@ -505,8 +657,13 @@ class DBSpendUpdateWriter:
)
)
except Exception as e:
verbose_proxy_logger.debug(
f"Update Tag DB failed to execute - {str(e)}\n{traceback.format_exc()}"
verbose_proxy_logger.error(
"Spend tracking - failed to enqueue tag spend update. "
"request_tags=%s, response_cost=%s - %s\n%s",
request_tags,
response_cost,
str(e),
traceback.format_exc(),
)
raise e
@ -607,6 +764,17 @@ class DBSpendUpdateWriter:
await self.redis_update_buffer.get_all_update_transactions_from_redis_buffer()
)
if db_spend_update_transactions is not None:
verbose_proxy_logger.info(
"Spend tracking - committing spend updates from Redis to DB: "
"keys=%d, users=%d, teams=%d, orgs=%d, end_users=%d, team_members=%d, tags=%d",
len(db_spend_update_transactions.get("key_list_transactions") or {}),
len(db_spend_update_transactions.get("user_list_transactions") or {}),
len(db_spend_update_transactions.get("team_list_transactions") or {}),
len(db_spend_update_transactions.get("org_list_transactions") or {}),
len(db_spend_update_transactions.get("end_user_list_transactions") or {}),
len(db_spend_update_transactions.get("team_member_list_transactions") or {}),
len(db_spend_update_transactions.get("tag_list_transactions") or {}),
)
await self._commit_spend_updates_to_db(
prisma_client=prisma_client,
n_retry_times=n_retry_times,
@ -677,7 +845,12 @@ class DBSpendUpdateWriter:
daily_spend_transactions=daily_agent_spend_update_transactions,
)
except Exception as e:
verbose_proxy_logger.error(f"Error committing spend updates: {e}")
verbose_proxy_logger.error(
"Spend tracking - failed to commit spend updates from Redis to DB. "
"Data already popped from Redis may be lost. Error: %s\n%s",
str(e),
traceback.format_exc(),
)
finally:
await self.pod_lock_manager.release_lock(
cronjob_id=DB_SPEND_UPDATE_JOB_NAME,
@ -793,6 +966,25 @@ class DBSpendUpdateWriter:
daily_spend_transactions=daily_agent_spend_update_transactions,
)
################## Tool Registry Upserts ##################
await self._flush_tool_discovery_queue(prisma_client=prisma_client)
async def _flush_tool_discovery_queue(
self,
prisma_client: PrismaClient,
) -> None:
"""Flush ToolDiscoveryQueue and batch-upsert new tools into LiteLLM_ToolTable."""
from litellm.proxy.db.tool_registry_writer import batch_upsert_tools
try:
items = self.tool_discovery_queue.flush()
if items:
await batch_upsert_tools(prisma_client=prisma_client, items=items)
except Exception as e:
verbose_proxy_logger.debug(
"_flush_tool_discovery_queue error (non-blocking): %s", e
)
async def _commit_spend_updates_to_db( # noqa: PLR0915
self,
prisma_client: PrismaClient,

View file

@ -86,6 +86,11 @@ class DailySpendUpdateQueue(BaseUpdateQueue):
) -> Dict[str, BaseDailySpendTransaction]:
"""Get all updates from the queue and return all updates aggregated by daily_transaction_key. Works for both user and team spend updates."""
updates = await self.flush_all_updates_from_in_memory_queue()
if len(updates) > 0:
verbose_proxy_logger.info(
"Spend tracking - flushed %d daily spend update items from in-memory queue",
len(updates),
)
aggregated_daily_spend_update_transactions = (
DailySpendUpdateQueue.get_aggregated_daily_spend_update_transactions(
updates

View file

@ -80,6 +80,14 @@ class PodLockManager:
)
self._emit_acquired_lock_event(cronjob_id, self.pod_id)
return True
else:
verbose_proxy_logger.info(
"Spend tracking - pod %s could not acquire lock for cronjob_id=%s, "
"held by pod %s. Spend updates in Redis will wait for the leader pod to commit.",
self.pod_id,
cronjob_id,
current_value,
)
return False
except Exception as e:
verbose_proxy_logger.error(
@ -124,10 +132,12 @@ class PodLockManager:
pod_id=self.pod_id,
)
else:
verbose_proxy_logger.debug(
"Pod %s failed to release Redis lock for cronjob_id=%s",
verbose_proxy_logger.warning(
"Spend tracking - pod %s failed to release Redis lock for cronjob_id=%s. "
"Lock will expire after TTL=%ds.",
self.pod_id,
cronjob_id,
DEFAULT_CRON_JOB_LOCK_TTL_SECONDS,
)
else:
verbose_proxy_logger.debug(

View file

@ -96,14 +96,28 @@ class RedisUpdateBuffer:
list_of_transactions = [safe_dumps(transactions)]
if self.redis_cache is None:
return
current_redis_buffer_size = await self.redis_cache.async_rpush(
key=redis_key,
values=list_of_transactions,
)
await self._emit_new_item_added_to_redis_buffer_event(
queue_size=current_redis_buffer_size,
service=service_type,
)
try:
current_redis_buffer_size = await self.redis_cache.async_rpush(
key=redis_key,
values=list_of_transactions,
)
verbose_proxy_logger.debug(
"Spend tracking - pushed spend updates to Redis buffer. "
"redis_key=%s, buffer_size=%s",
redis_key,
current_redis_buffer_size,
)
await self._emit_new_item_added_to_redis_buffer_event(
queue_size=current_redis_buffer_size,
service=service_type,
)
except Exception as e:
verbose_proxy_logger.error(
"Spend tracking - failed to push spend updates to Redis (redis_key=%s). "
"Error: %s",
redis_key,
str(e),
)
async def store_in_memory_spend_updates_in_redis(
self,
@ -305,6 +319,13 @@ class RedisUpdateBuffer:
if list_of_transactions is None:
return None
verbose_proxy_logger.info(
"Spend tracking - popped %d spend update batches from Redis buffer (key=%s). "
"These items are now removed from Redis and must be committed to DB.",
len(list_of_transactions) if isinstance(list_of_transactions, list) else 1,
REDIS_UPDATE_BUFFER_KEY,
)
# Parse the list of transactions from JSON strings
parsed_transactions = self._parse_list_of_transactions(list_of_transactions)

View file

@ -31,6 +31,11 @@ class SpendUpdateQueue(BaseUpdateQueue):
) -> DBSpendUpdateTransactions:
"""Flush all updates from the queue and return all updates aggregated by entity type."""
updates = await self.flush_all_updates_from_in_memory_queue()
if len(updates) > 0:
verbose_proxy_logger.info(
"Spend tracking - flushed %d spend update items from in-memory queue",
len(updates),
)
verbose_proxy_logger.debug("Aggregating updates by entity type: %s", updates)
return self.get_aggregated_db_spend_update_transactions(updates)

View file

@ -0,0 +1,54 @@
"""
In-memory buffer for tool registry upserts.
Unlike SpendUpdateQueue (which aggregates increments), ToolDiscoveryQueue
uses set-deduplication: each unique tool_name is only queued once per flush
cycle (~30s). The seen-set is cleared on every flush so that call_count
increments in subsequent cycles rather than stopping after the first flush.
"""
from typing import List, Set
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import ToolDiscoveryQueueItem
class ToolDiscoveryQueue:
"""
In-memory buffer for tool registry upserts.
Deduplicates by tool_name within each flush cycle: a tool is only queued
once per ~30s batch, so call_count increments once per flush cycle the
tool appears in (not once per invocation, but not once per pod lifetime
either). The seen-set is cleared on flush so subsequent batches can
re-count the same tool.
"""
def __init__(self) -> None:
self._seen_tool_names: Set[str] = set()
self._pending: List[ToolDiscoveryQueueItem] = []
def add_update(self, item: ToolDiscoveryQueueItem) -> None:
"""Enqueue a tool discovery item if tool_name has not been seen before."""
tool_name = item.get("tool_name", "")
if not tool_name:
return
if tool_name in self._seen_tool_names:
verbose_proxy_logger.debug(
"ToolDiscoveryQueue: skipping already-seen tool %s", tool_name
)
return
self._seen_tool_names.add(tool_name)
self._pending.append(item)
verbose_proxy_logger.debug(
"ToolDiscoveryQueue: queued new tool %s (origin=%s)",
tool_name,
item.get("origin"),
)
def flush(self) -> List[ToolDiscoveryQueueItem]:
"""Return and clear all pending items. Resets seen-set so the next
flush cycle can re-count the same tools."""
items, self._pending = self._pending, []
self._seen_tool_names.clear()
return items

View file

@ -0,0 +1,179 @@
"""
DB helpers for LiteLLM_ToolTable — the global tool registry.
Tools are auto-discovered from LLM responses and upserted here.
Admins use the management endpoints to read and update call_policy.
NOTE: Uses raw SQL (query_raw / execute_raw) instead of Prisma model methods
because the generated Prisma Python client may not have LiteLLM_ToolTable
when running against an older generated schema.
"""
import uuid
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Dict, List, Optional
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import ToolDiscoveryQueueItem
from litellm.types.tool_management import LiteLLM_ToolTableRow, ToolCallPolicy
if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient
def _row_to_model(row: dict) -> LiteLLM_ToolTableRow:
return LiteLLM_ToolTableRow(
tool_id=row.get("tool_id", ""),
tool_name=row.get("tool_name", ""),
origin=row.get("origin"),
call_policy=row.get("call_policy", "untrusted"),
call_count=int(row.get("call_count") or 0),
assignments=row.get("assignments"),
key_hash=row.get("key_hash"),
team_id=row.get("team_id"),
key_alias=row.get("key_alias"),
created_at=row.get("created_at"),
updated_at=row.get("updated_at"),
created_by=row.get("created_by"),
updated_by=row.get("updated_by"),
)
async def batch_upsert_tools(
prisma_client: "PrismaClient",
items: List[ToolDiscoveryQueueItem],
) -> None:
"""
Batch-upsert tool registry rows via raw SQL.
On first insert: sets call_policy = "untrusted" (schema default), call_count = 1.
On conflict: increments call_count; preserves existing call_policy.
"""
if not items:
return
try:
data = [item for item in items if item.get("tool_name")]
if not data:
return
for item in data:
tool_name = item.get("tool_name", "")
origin = item.get("origin") or "user_defined"
created_by = item.get("created_by") or "system"
key_hash = item.get("key_hash")
team_id = item.get("team_id")
key_alias = item.get("key_alias")
now = datetime.now(timezone.utc).isoformat()
await prisma_client.db.execute_raw(
'INSERT INTO "LiteLLM_ToolTable" '
"(tool_id, tool_name, origin, call_policy, call_count, created_by, updated_by, key_hash, team_id, key_alias, created_at, updated_at) "
"VALUES ($7, $1, $2, 'untrusted', 1, $3, $3, $4, $5, $6, $8, $8) "
"ON CONFLICT (tool_name) DO UPDATE SET "
"call_count = \"LiteLLM_ToolTable\".call_count + 1, "
"updated_at = $8",
tool_name,
origin,
created_by,
key_hash,
team_id,
key_alias,
str(uuid.uuid4()),
now,
)
verbose_proxy_logger.debug(
"tool_registry_writer: upserted %d tool(s)", len(data)
)
except Exception as e:
verbose_proxy_logger.error("tool_registry_writer batch_upsert_tools error: %s", e)
async def list_tools(
prisma_client: "PrismaClient",
call_policy: Optional[ToolCallPolicy] = None,
) -> List[LiteLLM_ToolTableRow]:
"""Return all tools, optionally filtered by call_policy."""
try:
if call_policy is not None:
rows = await prisma_client.db.query_raw(
'SELECT tool_id, tool_name, origin, call_policy, call_count, assignments, '
'key_hash, team_id, key_alias, created_at, updated_at, created_by, updated_by '
'FROM "LiteLLM_ToolTable" WHERE call_policy = $1 ORDER BY created_at DESC',
call_policy,
)
else:
rows = await prisma_client.db.query_raw(
'SELECT tool_id, tool_name, origin, call_policy, call_count, assignments, '
'key_hash, team_id, key_alias, created_at, updated_at, created_by, updated_by '
'FROM "LiteLLM_ToolTable" ORDER BY created_at DESC',
)
return [_row_to_model(row) for row in rows]
except Exception as e:
verbose_proxy_logger.error("tool_registry_writer list_tools error: %s", e)
return []
async def get_tool(
prisma_client: "PrismaClient",
tool_name: str,
) -> Optional[LiteLLM_ToolTableRow]:
"""Return a single tool row by tool_name."""
try:
rows = await prisma_client.db.query_raw(
'SELECT tool_id, tool_name, origin, call_policy, call_count, assignments, '
'key_hash, team_id, key_alias, created_at, updated_at, created_by, updated_by '
'FROM "LiteLLM_ToolTable" WHERE tool_name = $1',
tool_name,
)
if not rows:
return None
return _row_to_model(rows[0])
except Exception as e:
verbose_proxy_logger.error("tool_registry_writer get_tool error: %s", e)
return None
async def update_tool_policy(
prisma_client: "PrismaClient",
tool_name: str,
call_policy: ToolCallPolicy,
updated_by: Optional[str],
) -> Optional[LiteLLM_ToolTableRow]:
"""Update the call_policy for a tool. Upserts the row if it does not exist yet."""
try:
_updated_by = updated_by or "system"
now = datetime.now(timezone.utc).isoformat()
await prisma_client.db.execute_raw(
'INSERT INTO "LiteLLM_ToolTable" (tool_id, tool_name, call_policy, created_by, updated_by, created_at, updated_at) '
"VALUES ($4, $1, $2, $3, $3, $5, $5) "
"ON CONFLICT (tool_name) DO UPDATE SET call_policy = $2, updated_by = $3, updated_at = $5",
tool_name,
call_policy,
_updated_by,
str(uuid.uuid4()),
now,
)
return await get_tool(prisma_client, tool_name)
except Exception as e:
verbose_proxy_logger.error("tool_registry_writer update_tool_policy error: %s", e)
return None
async def get_tools_by_names(
prisma_client: "PrismaClient",
tool_names: List[str],
) -> Dict[str, str]:
"""
Return a {tool_name: call_policy} map for the given tool names.
Used by the policy enforcement guardrail — single batch query, never N+1.
"""
if not tool_names:
return {}
try:
placeholders = ", ".join(f"${i+1}" for i in range(len(tool_names)))
rows = await prisma_client.db.query_raw(
f'SELECT tool_name, call_policy FROM "LiteLLM_ToolTable" WHERE tool_name IN ({placeholders})',
*tool_names,
)
return {row["tool_name"]: row["call_policy"] for row in rows}
except Exception as e:
verbose_proxy_logger.error("tool_registry_writer get_tools_by_names error: %s", e)
return {}

View file

@ -0,0 +1,16 @@
import litellm
from litellm.types.guardrails import Guardrail, LitellmParams
def initialize_guardrail(litellm_params: LitellmParams, guardrail: Guardrail):
from litellm.proxy.guardrails.guardrail_hooks.tool_policy.tool_policy_guardrail import (
ToolPolicyGuardrail,
)
_callback = ToolPolicyGuardrail(
guardrail_name=guardrail.get("guardrail_name", "tool_policy"),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
)
litellm.logging_callback_manager.add_litellm_callback(_callback)
return _callback

View file

@ -0,0 +1,163 @@
"""
Tool Policy Guardrail
Reads call_policy from LiteLLM_ToolTable and enforces it on LLM requests/responses.
Policy values:
"trusted" - allow through (no action)
"untrusted" - allow through (no action; default for newly discovered tools)
"blocked" - raise HTTPException, preventing the tool call
"dual_llm" - (Phase 3) send to second LLM for verification; currently treated as allowed
Configuration in proxy config YAML:
guardrails:
- guardrail_name: "tool_policy"
litellm_params:
guardrail: tool_policy
mode: post_call
or both pre and post call:
- guardrail_name: "tool_policy"
litellm_params:
guardrail: tool_policy
mode: during_call # runs before LLM and on response
"""
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional
from fastapi import HTTPException
from litellm._logging import verbose_proxy_logger
from litellm.caching.dual_cache import DualCache
from litellm.constants import TOOL_POLICY_CACHE_TTL_SECONDS
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
GUARDRAIL_NAME = "tool_policy"
class ToolPolicyGuardrail(CustomGuardrail):
"""
Guardrail that enforces per-tool call policies stored in LiteLLM_ToolTable.
Tools with call_policy="blocked" are rejected before/after the LLM call.
Tools with call_policy="trusted" or "untrusted" pass through unchanged.
"""
def __init__(self, **kwargs: Any) -> None:
if "supported_event_hooks" not in kwargs:
kwargs["supported_event_hooks"] = [
GuardrailEventHooks.pre_call,
GuardrailEventHooks.post_call,
GuardrailEventHooks.during_call,
]
super().__init__(**kwargs)
self._policy_cache: DualCache = DualCache()
@log_guardrail_information
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional["LiteLLMLoggingObj"] = None,
) -> GenericGuardrailAPIInputs:
"""
Enforce tool policies on both request tools and response tool_calls.
- input_type="request": check inputs["tools"] (tool definitions in the LLM request)
- input_type="response": check inputs["tool_calls"] (tool_calls in the LLM response)
Raises HTTPException (400) if any tool is "blocked".
"""
if input_type == "request":
tools = inputs.get("tools") or []
tool_names = [
t["function"]["name"]
for t in tools
if isinstance(t, dict)
and isinstance(t.get("function"), dict)
and t["function"].get("name")
]
else: # response
tool_calls = inputs.get("tool_calls") or []
tool_names = []
for tc in tool_calls:
fn = None
if isinstance(tc, dict):
fn = (tc.get("function") or {}).get("name")
elif hasattr(tc, "function"):
fn = getattr(tc.function, "name", None)
if fn:
tool_names.append(fn)
if not tool_names:
return inputs
policy_map = await self._get_policies_cached(tool_names)
blocked = [name for name in tool_names if policy_map.get(name) == "blocked"]
if blocked:
verbose_proxy_logger.warning(
"ToolPolicyGuardrail: blocking tool(s) %s (policy=blocked)", blocked
)
raise HTTPException(
status_code=400,
detail={
"error": "Violated tool policy",
"blocked_tools": blocked,
"message": f"Tool(s) {blocked} are blocked by policy.",
},
)
return inputs
async def _get_policies_cached(self, tool_names: List[str]) -> Dict[str, str]:
"""
Batch-fetch call_policy for the given tool names.
Caches per individual tool name (not per combination) so that adding
a new tool to a request doesn't invalidate the cached policies for all
the other tools already in the cache.
"""
from litellm.proxy.db.tool_registry_writer import get_tools_by_names
from litellm.proxy.proxy_server import prisma_client
if not tool_names or prisma_client is None:
return {}
result: Dict[str, str] = {}
cache_misses: List[str] = []
for name in tool_names:
cached = await self._policy_cache.async_get_cache(f"tool_policy:{name}")
if cached is not None and isinstance(cached, str):
result[name] = cached
else:
cache_misses.append(name)
if cache_misses:
fetched = await get_tools_by_names(
prisma_client=prisma_client, tool_names=cache_misses
)
for name, policy in fetched.items():
result[name] = policy
await self._policy_cache.async_set_cache(
key=f"tool_policy:{name}",
value=policy,
ttl=TOOL_POLICY_CACHE_TTL_SECONDS,
)
verbose_proxy_logger.debug(
"ToolPolicyGuardrail: fetched %d policies from DB (cache hits: %d)",
len(cache_misses),
len(tool_names) - len(cache_misses),
)
return result

View file

@ -14,6 +14,7 @@ from litellm.proxy._types import (AddTeamCallback, CommonProxyErrors,
LitellmDataForBackendLLMCall,
LitellmUserRoles, SpecialHeaders,
TeamCallbackMetadata, UserAPIKeyAuth)
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
# Cache special headers as a frozenset for O(1) lookup performance
_SPECIAL_HEADERS_CACHE = frozenset(
@ -824,7 +825,7 @@ async def add_litellm_data_to_request( # noqa: PLR0915
from litellm.proxy.proxy_server import llm_router, premium_user
from litellm.types.proxy.litellm_pre_call_utils import SecretFields
_raw_headers: Dict[str, str] = dict(request.headers)
_raw_headers: Dict[str, str] = _safe_get_request_headers(request)
_headers: Dict[str, str] = clean_headers(
request.headers,
litellm_key_header_name=(

View file

@ -21,6 +21,7 @@ from typing_extensions import TypedDict
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
from litellm._uuid import uuid
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.proxy._types import (
@ -724,7 +725,7 @@ async def get_service_provider_config(request: Request):
"SCIM ServiceProviderConfig request: method=%s url=%s headers=%s",
request.method,
request.url,
dict(request.headers),
_safe_get_request_headers(request),
)
meta = {
"resourceType": "ServiceProviderConfig",

View file

@ -0,0 +1,149 @@
"""
TOOL POLICY MANAGEMENT
All /tool management endpoints
GET /v1/tool/list - List all discovered tools and their policies
GET /v1/tool/{tool_name} - Get a single tool's details
POST /v1/tool/policy - Update the call_policy for a tool
"""
from typing import Optional
from fastapi import APIRouter, Depends, HTTPException
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.types.tool_management import (
LiteLLM_ToolTableRow,
ToolCallPolicy,
ToolListResponse,
ToolPolicyUpdateRequest,
ToolPolicyUpdateResponse,
)
router = APIRouter()
@router.get(
"/v1/tool/list",
tags=["tool management"],
dependencies=[Depends(user_api_key_auth)],
response_model=ToolListResponse,
)
async def list_tools(
call_policy: Optional[ToolCallPolicy] = None,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
List all auto-discovered tools and their call policies.
Parameters:
- call_policy: Optional filter — one of "trusted", "untrusted", "dual_llm", "blocked"
"""
from litellm.proxy.db.tool_registry_writer import list_tools as db_list_tools
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
)
try:
tools = await db_list_tools(prisma_client=prisma_client, call_policy=call_policy)
return ToolListResponse(tools=tools, total=len(tools))
except Exception as e:
verbose_proxy_logger.exception("Error listing tools: %s", e)
raise HTTPException(status_code=500, detail=str(e))
@router.get(
"/v1/tool/{tool_name:path}",
tags=["tool management"],
dependencies=[Depends(user_api_key_auth)],
response_model=LiteLLM_ToolTableRow,
)
async def get_tool(
tool_name: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Get details for a single tool.
Parameters:
- tool_name: The tool name (supports namespaced names with slashes)
"""
from litellm.proxy.db.tool_registry_writer import get_tool as db_get_tool
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
)
try:
tool = await db_get_tool(prisma_client=prisma_client, tool_name=tool_name)
if tool is None:
raise HTTPException(
status_code=404, detail=f"Tool '{tool_name}' not found"
)
return tool
except HTTPException:
raise
except Exception as e:
verbose_proxy_logger.exception("Error getting tool: %s", e)
raise HTTPException(status_code=500, detail=str(e))
@router.post(
"/v1/tool/policy",
tags=["tool management"],
dependencies=[Depends(user_api_key_auth)],
response_model=ToolPolicyUpdateResponse,
)
async def update_tool_policy(
data: ToolPolicyUpdateRequest,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Set the call policy for a tool.
Parameters:
- tool_name: str - The tool to update
- call_policy: "trusted" | "untrusted" | "dual_llm" | "blocked"
Setting a tool to "blocked" will cause the ToolPolicyGuardrail to remove
that tool_call from LLM responses before returning them to the client.
"""
from litellm.proxy.db.tool_registry_writer import (
update_tool_policy as db_update_tool_policy,
)
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
)
try:
updated = await db_update_tool_policy(
prisma_client=prisma_client,
tool_name=data.tool_name,
call_policy=data.call_policy,
updated_by=user_api_key_dict.user_id,
)
if updated is None:
raise HTTPException(
status_code=500, detail=f"Failed to update policy for tool '{data.tool_name}'"
)
return ToolPolicyUpdateResponse(
tool_name=updated.tool_name,
call_policy=updated.call_policy,
updated=True,
)
except HTTPException:
raise
except Exception as e:
verbose_proxy_logger.exception("Error updating tool policy: %s", e)
raise HTTPException(status_code=500, detail=str(e))

View file

@ -0,0 +1,9 @@
"""
Usage endpoints package.
Re-exports the router from endpoints module.
"""
from litellm.proxy.management_endpoints.usage_endpoints.endpoints import ( # noqa: F401
router,
)

View file

@ -0,0 +1,578 @@
"""
AI Usage Chat - uses LLM tool calling to answer questions about
usage/spend data by querying the aggregated daily activity endpoints.
"""
import json
from datetime import date
from typing import Any, AsyncIterator, Callable, Dict, List, Literal, Optional
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import DEFAULT_COMPETITOR_DISCOVERY_MODEL
from litellm.types.proxy.management_endpoints.common_daily_activity import (
SpendAnalyticsPaginatedResponse,
)
from typing_extensions import TypedDict
# ---------------------------------------------------------------------------
# Constants
# ---------------------------------------------------------------------------
USAGE_AI_TEMPERATURE = 0.2
TABLE_DAILY_USER_SPEND = "litellm_dailyuserspend"
TABLE_DAILY_TEAM_SPEND = "litellm_dailyteamspend"
TABLE_DAILY_TAG_SPEND = "litellm_dailytagspend"
ENTITY_FIELD_USER = "user_id"
ENTITY_FIELD_TEAM = "team_id"
ENTITY_FIELD_TAG = "tag"
PAGINATED_PAGE_SIZE = 200
MAX_CHAT_MESSAGES = 20
TOP_N_MODELS = 15
TOP_N_PROVIDERS = 10
TOP_N_KEYS = 10
# ---------------------------------------------------------------------------
# Types
# ---------------------------------------------------------------------------
class SSEStatusEvent(TypedDict):
type: Literal["status"]
message: str
class SSEToolCallEvent(TypedDict, total=False):
type: Literal["tool_call"]
tool_name: str
tool_label: str
arguments: Dict[str, str]
status: Literal["running", "complete", "error"]
error: str
class SSEChunkEvent(TypedDict):
type: Literal["chunk"]
content: str
class SSEDoneEvent(TypedDict):
type: Literal["done"]
class SSEErrorEvent(TypedDict):
type: Literal["error"]
message: str
SSEEvent = (
SSEStatusEvent | SSEToolCallEvent | SSEChunkEvent | SSEDoneEvent | SSEErrorEvent
)
class ToolHandler(TypedDict):
fetch: Callable[..., Any]
summarise: Callable[[Dict[str, Any]], str]
label: str
# ---------------------------------------------------------------------------
# Tool definitions (OpenAI function-calling schema)
# ---------------------------------------------------------------------------
_DATE_PARAMS = {
"start_date": {"type": "string", "description": "Start date in YYYY-MM-DD format"},
"end_date": {"type": "string", "description": "End date in YYYY-MM-DD format"},
}
_TOOL_USAGE = {
"type": "function",
"function": {
"name": "get_usage_data",
"description": (
"Fetch aggregated global usage/spend data. Returns daily spend, "
"token counts, request counts, and breakdowns by model, provider, "
"and API key. Use for overall spend, top models, top providers."
),
"parameters": {
"type": "object",
"properties": {
**_DATE_PARAMS,
"user_id": {
"type": "string",
"description": "Optional user ID filter. Omit for global view.",
},
},
"required": ["start_date", "end_date"],
},
},
}
_TOOL_TEAM = {
"type": "function",
"function": {
"name": "get_team_usage_data",
"description": (
"Fetch usage/spend data broken down by team. Use for questions "
"like 'which team spends the most' or 'show me team X usage'."
),
"parameters": {
"type": "object",
"properties": {
**_DATE_PARAMS,
"team_ids": {
"type": "string",
"description": "Optional comma-separated team IDs. Omit for all teams.",
},
},
"required": ["start_date", "end_date"],
},
},
}
_TOOL_TAG = {
"type": "function",
"function": {
"name": "get_tag_usage_data",
"description": (
"Fetch usage/spend data broken down by tag. Tags are labels "
"attached to requests (features, environments, credentials)."
),
"parameters": {
"type": "object",
"properties": {
**_DATE_PARAMS,
"tags": {
"type": "string",
"description": "Optional comma-separated tag names. Omit for all tags.",
},
},
"required": ["start_date", "end_date"],
},
},
}
TOOLS_BASE = [_TOOL_USAGE]
TOOLS_ADMIN = [_TOOL_USAGE, _TOOL_TEAM, _TOOL_TAG]
def get_tools_for_role(is_admin: bool) -> List[Dict[str, Any]]:
"""Return the tool list appropriate for the user's role."""
return TOOLS_ADMIN if is_admin else TOOLS_BASE
_SYSTEM_PROMPT_BASE = (
"You are an AI assistant embedded in the LiteLLM Usage dashboard. "
"You help users understand their LLM API spend and usage data.\n\n"
"ALWAYS call the appropriate tool(s) first to fetch data before answering. "
"You may call multiple tools if the question spans different dimensions.\n\n"
"Guidelines:\n"
"- Be concise and specific. Use exact numbers from the data.\n"
"- Format costs as dollar amounts (e.g. $12.34).\n"
"- When comparing entities, show a ranked list.\n"
"- If data is empty or no results found, say so clearly.\n"
"- Do not hallucinate data — only use what the tools return.\n"
"- Today's date will be provided below. Use it to interpret relative dates "
"like 'this week', 'this month', 'last 7 days', etc."
)
_TOOL_DESCRIPTIONS_ADMIN = (
"You have access to these tools:\n"
"- `get_usage_data`: Global/user-level usage (spend, models, providers, API keys)\n"
"- `get_team_usage_data`: Team-level usage breakdown\n"
"- `get_tag_usage_data`: Tag-level usage breakdown\n\n"
)
_TOOL_DESCRIPTIONS_BASE = (
"You have access to this tool:\n"
"- `get_usage_data`: Your usage data (spend, models, providers, API keys)\n\n"
)
def _build_system_prompt(is_admin: bool) -> str:
"""Build role-appropriate system prompt with today's date."""
tool_desc = _TOOL_DESCRIPTIONS_ADMIN if is_admin else _TOOL_DESCRIPTIONS_BASE
return (
f"{_SYSTEM_PROMPT_BASE}\n\n{tool_desc}"
f"Today's date: {date.today().isoformat()}"
)
# keep a public reference for test assertions
SYSTEM_PROMPT = _SYSTEM_PROMPT_BASE
# ---------------------------------------------------------------------------
# Data fetchers
# ---------------------------------------------------------------------------
def _parse_csv_ids(raw: Optional[str]) -> Optional[List[str]]:
if not raw:
return None
return [t.strip() for t in raw.split(",") if t.strip()]
async def _query_activity(
table_name: str,
entity_id_field: str,
entity_id: Optional[Any],
start_date: str,
end_date: str,
*,
use_aggregated: bool = False,
) -> SpendAnalyticsPaginatedResponse:
"""Shared helper that calls the daily activity query layer."""
from litellm.proxy.management_endpoints.common_daily_activity import (
get_daily_activity,
get_daily_activity_aggregated,
)
from litellm.proxy.proxy_server import prisma_client
if use_aggregated:
return await get_daily_activity_aggregated(
prisma_client=prisma_client,
table_name=table_name,
entity_id_field=entity_id_field,
entity_id=entity_id,
entity_metadata_field=None,
start_date=start_date,
end_date=end_date,
model=None,
api_key=None,
)
return await get_daily_activity(
prisma_client=prisma_client,
table_name=table_name,
entity_id_field=entity_id_field,
entity_id=entity_id,
entity_metadata_field=None,
start_date=start_date,
end_date=end_date,
model=None,
api_key=None,
page=1,
page_size=PAGINATED_PAGE_SIZE,
)
async def _fetch_usage_data(
start_date: str, end_date: str, user_id: Optional[str] = None
) -> Dict[str, Any]:
resp = await _query_activity(
TABLE_DAILY_USER_SPEND,
ENTITY_FIELD_USER,
user_id,
start_date,
end_date,
use_aggregated=True,
)
return resp.model_dump(mode="json")
async def _fetch_team_usage_data(
start_date: str, end_date: str, team_ids: Optional[str] = None
) -> Dict[str, Any]:
resp = await _query_activity(
TABLE_DAILY_TEAM_SPEND,
ENTITY_FIELD_TEAM,
_parse_csv_ids(team_ids),
start_date,
end_date,
)
return resp.model_dump(mode="json")
async def _fetch_tag_usage_data(
start_date: str, end_date: str, tags: Optional[str] = None
) -> Dict[str, Any]:
resp = await _query_activity(
TABLE_DAILY_TAG_SPEND,
ENTITY_FIELD_TAG,
_parse_csv_ids(tags),
start_date,
end_date,
)
return resp.model_dump(mode="json")
# ---------------------------------------------------------------------------
# Summarisers — convert raw JSON to concise text the LLM can reason over
# ---------------------------------------------------------------------------
def _accumulate_breakdown(
results: List[Dict[str, Any]], dimension: str, fields: List[str]
) -> Dict[str, Dict[str, float]]:
"""Aggregate a single breakdown dimension across days."""
totals: Dict[str, Dict[str, float]] = {}
for day in results:
for key, entry in day.get("breakdown", {}).get(dimension, {}).items():
if key not in totals:
totals[key] = {f: 0.0 for f in fields}
m = entry.get("metrics", {})
for f in fields:
totals[key][f] += m.get(f, 0)
return totals
def _ranked_lines(
totals: Dict[str, Dict[str, float]],
fmt: Callable[[str, Dict[str, float]], str],
limit: int,
) -> List[str]:
"""Sort by spend descending, format each entry, and truncate."""
return [
fmt(name, vals)
for name, vals in sorted(totals.items(), key=lambda x: -x[1].get("spend", 0))[
:limit
]
]
def _summarise_usage_data(data: Dict[str, Any]) -> str:
meta = data.get("metadata", {})
results = data.get("results", [])
header = (
f"Total Spend: ${meta.get('total_spend', 0):.4f}\n"
f"Total Requests: {meta.get('total_api_requests', 0)}\n"
f"Successful: {meta.get('total_successful_requests', 0)} | "
f"Failed: {meta.get('total_failed_requests', 0)}\n"
f"Total Tokens: {meta.get('total_tokens', 0)}"
)
models = _accumulate_breakdown(
results, "models", ["spend", "api_requests", "total_tokens"]
)
providers = _accumulate_breakdown(results, "providers", ["spend", "api_requests"])
model_lines = _ranked_lines(
models,
lambda n, d: f" - {n}: ${d['spend']:.4f} ({int(d['api_requests'])} reqs, {int(d['total_tokens'])} tokens)",
TOP_N_MODELS,
)
provider_lines = _ranked_lines(
providers,
lambda n, d: f" - {n}: ${d['spend']:.4f} ({int(d['api_requests'])} reqs)",
TOP_N_PROVIDERS,
)
sections = [header, ""]
sections += ["Top Models by Spend:"] + (model_lines or [" (no data)"]) + [""]
sections += ["Top Providers by Spend:"] + (provider_lines or [" (no data)"])
return "\n".join(sections)
def _summarise_entity_data(data: Dict[str, Any], entity_label: str) -> str:
"""Summarise team/tag entity usage data."""
results = data.get("results", [])
if not results:
return f"No {entity_label} usage data found for the given date range."
totals: Dict[str, Dict[str, Any]] = {}
for day in results:
for eid, entry in day.get("breakdown", {}).get("entities", {}).items():
if eid not in totals:
alias = entry.get("metadata", {}).get("alias", eid)
totals[eid] = {"alias": alias, "spend": 0.0, "requests": 0, "tokens": 0}
m = entry.get("metrics", {})
totals[eid]["spend"] += m.get("spend", 0)
totals[eid]["requests"] += m.get("api_requests", 0)
totals[eid]["tokens"] += m.get("total_tokens", 0)
lines = [f"{entity_label} Usage ({len(totals)} {entity_label.lower()}s):", ""]
for eid, d in sorted(totals.items(), key=lambda x: -x[1]["spend"]):
label = d["alias"] if d["alias"] != eid else eid
lines.append(
f"- {label} (ID: {eid}): ${d['spend']:.4f} | "
f"{int(d['requests'])} reqs | {int(d['tokens'])} tokens"
)
return "\n".join(lines)
# ---------------------------------------------------------------------------
# Tool dispatch registry
# ---------------------------------------------------------------------------
TOOL_HANDLERS: Dict[str, ToolHandler] = {
"get_usage_data": ToolHandler(
fetch=_fetch_usage_data,
summarise=_summarise_usage_data,
label="global usage data",
),
"get_team_usage_data": ToolHandler(
fetch=_fetch_team_usage_data,
summarise=lambda data: _summarise_entity_data(data, "Team"),
label="team usage data",
),
"get_tag_usage_data": ToolHandler(
fetch=_fetch_tag_usage_data,
summarise=lambda data: _summarise_entity_data(data, "Tag"),
label="tag usage data",
),
}
# ---------------------------------------------------------------------------
# SSE streaming
# ---------------------------------------------------------------------------
def _sse(event: SSEEvent) -> str:
return f"data: {json.dumps(event)}\n\n"
def _resolve_fetch_kwargs(
fn_name: str,
fn_args: Dict[str, str],
user_id: Optional[str],
is_admin: bool,
) -> Dict[str, Any]:
"""Build keyword arguments for a tool's fetch function."""
start_date = fn_args.get("start_date", "")
end_date = fn_args.get("end_date", "")
if not start_date or not end_date:
raise ValueError("Missing required start_date or end_date from tool arguments")
kwargs: Dict[str, Any] = {"start_date": start_date, "end_date": end_date}
if fn_name == "get_usage_data":
if not is_admin:
kwargs["user_id"] = user_id
elif fn_args.get("user_id"):
kwargs["user_id"] = fn_args["user_id"]
elif fn_name == "get_team_usage_data" and fn_args.get("team_ids"):
kwargs["team_ids"] = fn_args["team_ids"]
elif fn_name == "get_tag_usage_data" and fn_args.get("tags"):
kwargs["tags"] = fn_args["tags"]
return kwargs
async def _execute_tool_call(
handler: ToolHandler,
fn_name: str,
fn_args: Dict[str, str],
user_id: Optional[str],
is_admin: bool,
) -> str:
"""Run a single tool and return the summarised result text."""
kwargs = _resolve_fetch_kwargs(fn_name, fn_args, user_id, is_admin)
raw_data = await handler["fetch"](**kwargs)
return handler["summarise"](raw_data)
async def _process_tool_call(
tc: Any,
chat_messages: List[Dict[str, Any]],
user_id: Optional[str],
is_admin: bool,
) -> AsyncIterator[str]:
"""Execute a single tool call, yielding SSE events for status."""
fn_name = tc.function.name
fn_args = json.loads(tc.function.arguments)
allowed_names = {t["function"]["name"] for t in get_tools_for_role(is_admin)}
handler = TOOL_HANDLERS.get(fn_name)
if fn_name not in allowed_names or not handler:
chat_messages.append(
{
"role": "tool",
"tool_call_id": tc.id,
"content": f"Tool not available: {fn_name}",
}
)
return
tool_event_base = {
"type": "tool_call",
"tool_name": fn_name,
"tool_label": handler["label"],
"arguments": fn_args,
}
yield _sse({**tool_event_base, "status": "running"})
try:
tool_result = await _execute_tool_call(
handler, fn_name, fn_args, user_id, is_admin
)
yield _sse({**tool_event_base, "status": "complete"})
except Exception as e:
verbose_proxy_logger.error("Tool %s failed: %s", fn_name, e)
tool_result = f"Error fetching {handler['label']}. Please try again."
yield _sse({**tool_event_base, "status": "error"})
chat_messages.append(
{"role": "tool", "tool_call_id": tc.id, "content": tool_result}
)
async def _stream_final_response(
model: str, chat_messages: List[Dict[str, Any]]
) -> AsyncIterator[str]:
"""Stream the final LLM response after tool results are appended."""
yield _sse({"type": "status", "message": "Analyzing results..."})
response = await litellm.acompletion(
model=model,
messages=chat_messages,
stream=True,
temperature=USAGE_AI_TEMPERATURE,
)
async for chunk in response:
delta = chunk.choices[0].delta.content
if delta:
yield _sse({"type": "chunk", "content": delta})
async def stream_usage_ai_chat(
messages: List[Dict[str, str]],
model: Optional[str] = None,
user_id: Optional[str] = None,
is_admin: bool = False,
) -> AsyncIterator[str]:
"""Stream SSE events: status → tool_call → chunk → done."""
resolved_model = (model or "").strip() or DEFAULT_COMPETITOR_DISCOVERY_MODEL
truncated = (
messages[-MAX_CHAT_MESSAGES:] if len(messages) > MAX_CHAT_MESSAGES else messages
)
chat_messages: List[Dict[str, Any]] = [
{"role": "system", "content": _build_system_prompt(is_admin)},
*truncated,
]
try:
yield _sse({"type": "status", "message": "Thinking..."})
tools = get_tools_for_role(is_admin)
response = await litellm.acompletion(
model=resolved_model,
messages=chat_messages,
tools=tools,
temperature=USAGE_AI_TEMPERATURE,
)
choice = response.choices[0] # type: ignore
if not choice.message.tool_calls:
if choice.message.content:
yield _sse({"type": "chunk", "content": choice.message.content})
yield _sse({"type": "done"})
return
chat_messages.append(choice.message.model_dump())
for tc in choice.message.tool_calls:
async for event in _process_tool_call(tc, chat_messages, user_id, is_admin):
yield event
async for event in _stream_final_response(resolved_model, chat_messages):
yield event
yield _sse({"type": "done"})
except Exception as e:
verbose_proxy_logger.error("AI usage chat failed: %s", e)
yield _sse(
{
"type": "error",
"message": "An internal error occurred. Please try again.",
}
)

View file

@ -0,0 +1,65 @@
"""
USAGE AI CHAT ENDPOINTS
/usage/ai/chat - Stream AI chat responses about usage data
"""
from typing import List, Literal, Optional
from fastapi import APIRouter, Depends, Request
from fastapi.responses import StreamingResponse
from pydantic import BaseModel, Field
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
router = APIRouter()
class ChatMessage(BaseModel):
role: Literal["user", "assistant"]
content: str
class UsageAIChatRequest(BaseModel):
messages: List[ChatMessage] = Field(
..., description="Chat messages (user/assistant history)"
)
model: Optional[str] = Field(default=None, description="Model to use for AI chat")
@router.post(
"/usage/ai/chat",
tags=["Budget & Spend Tracking"],
dependencies=[Depends(user_api_key_auth)],
)
async def usage_ai_chat(
data: UsageAIChatRequest,
request: Request,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
AI chat about usage data. Streams SSE events with the AI response.
The AI agent has access to tools that query aggregated daily activity data.
"""
from litellm.proxy.management_endpoints.common_utils import (
_user_has_admin_view,
)
from litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat import (
stream_usage_ai_chat,
)
is_admin = _user_has_admin_view(user_api_key_dict)
user_id = user_api_key_dict.user_id
messages = [{"role": m.role, "content": m.content} for m in data.messages]
return StreamingResponse(
stream_usage_ai_chat(
messages=messages,
model=data.model,
user_id=user_id,
is_admin=is_admin,
),
media_type="text/event-stream",
headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"},
)

View file

@ -28,6 +28,7 @@ from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
_safe_get_request_headers,
_safe_set_request_parsed_body,
get_form_data,
get_request_body,
@ -60,7 +61,7 @@ def create_request_copy(request: Request):
return {
"method": request.method,
"url": str(request.url),
"headers": dict(request.headers),
"headers": _safe_get_request_headers(request).copy(),
"cookies": request.cookies,
"query_params": dict(request.query_params),
}
@ -329,7 +330,7 @@ async def vllm_proxy_route(
method=request.method,
endpoint=endpoint,
request_query_params=request.query_params,
request_headers=dict(request.headers),
request_headers=_safe_get_request_headers(request),
stream=request_body.get("stream", False),
content=None,
data=None,
@ -1307,7 +1308,7 @@ async def azure_proxy_route(
method=request.method,
endpoint=endpoint,
request_query_params=request.query_params,
request_headers=dict(request.headers),
request_headers=_safe_get_request_headers(request),
stream=request_body.get("stream", False),
content=None,
data=None,
@ -1505,7 +1506,7 @@ def get_vertex_ai_allowed_incoming_headers(request: Request) -> dict:
Returns:
dict: Headers dictionary with only allowed headers
"""
incoming_headers = dict(request.headers) or {}
incoming_headers = _safe_get_request_headers(request)
headers = {}
for header_name in ALLOWED_VERTEX_AI_PASSTHROUGH_HEADERS:
if header_name in incoming_headers:
@ -1621,7 +1622,7 @@ async def _prepare_vertex_auth_headers(
if (
vertex_credentials is None or vertex_credentials.vertex_project is None
) and router_credentials is None:
headers = dict(request.headers) or {}
headers = _safe_get_request_headers(request).copy()
headers_passed_through = True
verbose_proxy_logger.debug(
"default_vertex_config not set, incoming request headers %s", headers

View file

@ -50,7 +50,10 @@ from litellm.proxy._types import (
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
_safe_get_request_headers,
)
from litellm.proxy.utils import get_server_root_path
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.custom_http import httpxSpecialProvider
@ -649,7 +652,7 @@ async def pass_through_request( # noqa: PLR0915
url = httpx.URL(target)
headers = custom_headers
headers = HttpPassThroughEndpointHelpers.forward_headers_from_request(
request_headers=dict(request.headers),
request_headers=_safe_get_request_headers(request).copy(),
headers=headers,
forward_headers=forward_headers,
)

View file

@ -294,6 +294,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
from litellm.proxy.common_utils.html_forms.ui_login import build_ui_login_form
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
_safe_get_request_headers,
check_file_size_under_limit,
get_form_data,
)
@ -391,6 +392,7 @@ from litellm.proxy.management_endpoints.organization_endpoints import (
router as organization_router,
)
from litellm.proxy.management_endpoints.policy_endpoints import router as policy_router
from litellm.proxy.management_endpoints.usage_endpoints import router as usage_ai_router
from litellm.proxy.management_endpoints.project_endpoints import (
router as project_router,
)
@ -409,6 +411,9 @@ from litellm.proxy.management_endpoints.team_endpoints import (
update_team,
validate_membership,
)
from litellm.proxy.management_endpoints.tool_management_endpoints import (
router as tool_management_router,
)
from litellm.proxy.management_endpoints.ui_sso import (
get_disabled_non_admin_personal_key_creation,
)
@ -1768,8 +1773,14 @@ async def update_cache( # noqa: PLR0915
("{}:spend".format(litellm_proxy_admin_name), increment)
)
except Exception as e:
verbose_proxy_logger.debug(
f"An error occurred updating user cache: {str(e)}\n\n{traceback.format_exc()}"
verbose_proxy_logger.warning(
"Spend tracking - failed to update user spend in cache. "
"Budget enforcement may use stale spend values. "
"user_id=%s, response_cost=%s - %s\n%s",
user_id,
response_cost,
str(e),
traceback.format_exc(),
)
### UPDATE END-USER SPEND ###
@ -1806,8 +1817,14 @@ async def update_cache( # noqa: PLR0915
existing_spend_obj.spend = new_spend
values_to_update_in_cache.append((_id, existing_spend_obj.json()))
except Exception as e:
verbose_proxy_logger.exception(
f"An error occurred updating end user cache: {str(e)}"
verbose_proxy_logger.warning(
"Spend tracking - failed to update end user spend in cache. "
"Budget enforcement may use stale spend values. "
"end_user_id=%s, response_cost=%s - %s\n%s",
end_user_id,
response_cost,
str(e),
traceback.format_exc(),
)
### UPDATE TEAM SPEND ###
@ -1848,8 +1865,14 @@ async def update_cache( # noqa: PLR0915
existing_spend_obj.spend = new_spend
values_to_update_in_cache.append((_id, existing_spend_obj))
except Exception as e:
verbose_proxy_logger.exception(
f"An error occurred updating end user cache: {str(e)}"
verbose_proxy_logger.warning(
"Spend tracking - failed to update team spend in cache. "
"Budget enforcement may use stale spend values. "
"team_id=%s, response_cost=%s - %s\n%s",
team_id,
response_cost,
str(e),
traceback.format_exc(),
)
### UPDATE TAG SPEND ###
@ -1894,8 +1917,14 @@ async def update_cache( # noqa: PLR0915
existing_tag_obj.spend = new_spend
values_to_update_in_cache.append((cache_key, existing_tag_obj))
except Exception as e:
verbose_proxy_logger.exception(
f"An error occurred updating tag cache: {str(e)}"
verbose_proxy_logger.warning(
"Spend tracking - failed to update tag spend in cache. "
"Budget enforcement may use stale spend values. "
"tags=%s, response_cost=%s - %s\n%s",
tags,
response_cost,
str(e),
traceback.format_exc(),
)
if token is not None and response_cost is not None:
@ -10424,7 +10453,7 @@ async def async_queue_request(
data["proxy_server_request"] = {
"url": str(request.url),
"method": request.method,
"headers": dict(request.headers),
"headers": _safe_get_request_headers(request).copy(),
"body": copy.copy(data), # use copy instead of deepcopy
}
@ -10445,7 +10474,7 @@ async def async_queue_request(
data["metadata"] = {}
data["metadata"]["user_api_key"] = user_api_key_dict.api_key
data["metadata"]["user_api_key_metadata"] = user_api_key_dict.metadata
_headers = dict(request.headers)
_headers = _safe_get_request_headers(request).copy()
_headers.pop(
"authorization", None
) # do not store the original `sk-..` api key in the db
@ -12844,6 +12873,7 @@ app.include_router(caching_router)
app.include_router(analytics_router)
app.include_router(guardrails_router)
app.include_router(policy_router)
app.include_router(usage_ai_router)
app.include_router(policy_crud_router)
app.include_router(policy_resolve_router)
app.include_router(search_tool_management_router)
@ -12857,6 +12887,7 @@ app.include_router(budget_management_router)
app.include_router(model_management_router)
app.include_router(model_access_group_management_router)
app.include_router(tag_management_router)
app.include_router(tool_management_router)
app.include_router(cost_tracking_settings_router)
app.include_router(router_settings_router)
app.include_router(fallback_management_router)

View file

@ -1051,6 +1051,26 @@ model LiteLLM_PolicyAttachmentTable {
updated_by String?
}
// Global tool registry - auto-discovered from LLM responses; admins set call_policy here
model LiteLLM_ToolTable {
tool_id String @id @default(uuid())
tool_name String @unique // e.g. "huggingface_remote-mcp__dynamic_space"
origin String? // MCP server name or "user_defined"
call_policy String @default("untrusted") // "trusted" | "untrusted" | "dual_llm" | "blocked"
call_count Int @default(0) // cumulative number of times this tool was seen
assignments Json? @default("{}")
key_hash String? // hash of the virtual key that first called this tool
team_id String? // team that first called this tool
key_alias String? // human-readable alias of the virtual key
created_at DateTime @default(now())
created_by String?
updated_at DateTime @default(now()) @updatedAt
updated_by String?
@@index([call_policy])
@@index([team_id])
}
//Unified Access Groups table for storing unified access groups
model LiteLLM_AccessGroupTable {
access_group_id String @id @default(uuid())

View file

@ -4457,6 +4457,11 @@ class ProxyUpdateSpend:
len(logs_to_process) :
]
popped_batch = True
if len(logs_to_process) > 0:
verbose_proxy_logger.info(
"Spend tracking - processing %d spend logs for DB write",
len(logs_to_process),
)
start_time = time.time()
try:
for i in range(n_retry_times + 1):
@ -4503,9 +4508,17 @@ class ProxyUpdateSpend:
f"{len(logs_to_process)} logs processed. Remaining in queue: {remaining_count}"
)
break
except DB_CONNECTION_ERROR_TYPES:
except DB_CONNECTION_ERROR_TYPES as e:
if i is None:
i = 0
verbose_proxy_logger.warning(
"Spend tracking - DB connection error writing spend logs, "
"retry %d/%d. logs_count=%d, error=%s",
i + 1,
n_retry_times,
len(logs_to_process),
str(e),
)
if i >= n_retry_times:
raise
await asyncio.sleep(2**i)
@ -4620,8 +4633,8 @@ async def update_spend_logs_job(
logs_to_process=logs_to_process,
)
except Exception as guardrail_tracking_err:
verbose_proxy_logger.debug(
"Guardrail usage tracking failed (non-fatal): %s",
verbose_proxy_logger.warning(
"Spend tracking - guardrail usage tracking failed (non-fatal): %s",
guardrail_tracking_err,
)

View file

@ -19,6 +19,7 @@ from fastapi import APIRouter, Request, Response
import litellm
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
from litellm.proxy.litellm_pre_call_utils import _get_dynamic_logging_metadata
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
create_pass_through_route,
@ -32,7 +33,7 @@ def create_request_copy(request: Request):
return {
"method": request.method,
"url": str(request.url),
"headers": dict(request.headers),
"headers": _safe_get_request_headers(request).copy(),
"cookies": request.cookies,
"query_params": dict(request.query_params),
}

View file

@ -163,7 +163,11 @@ from litellm.types.utils import (
)
from litellm.types.utils import ModelInfo
from litellm.types.utils import ModelInfo as ModelMapInfo
from litellm.types.utils import ModelResponseStream, StandardLoggingPayload, Usage
from litellm.types.utils import (
ModelResponseStream,
StandardLoggingPayload,
Usage,
)
from litellm.utils import (
CustomStreamWrapper,
EmbeddingResponse,
@ -1555,6 +1559,9 @@ class Router:
logging_obj=model_response.logging_obj,
)
self._async_generator = async_generator
# Preserve hidden params (including litellm_overhead_time_ms) from original response
if hasattr(model_response, "_hidden_params"):
self._hidden_params = model_response._hidden_params.copy()
def __aiter__(self):
return self
@ -6416,6 +6423,7 @@ class Router:
self.model_list = []
self.model_id_to_deployment_index_map = {} # Reset the index
self.model_name_to_deployment_indices = {} # Reset the model_name index
self._invalidate_model_group_info_cache()
self._invalidate_access_groups_cache()
# we add api_base/api_key each model so load balancing between azure/gpt on api_base1 and api_base2 works
@ -6726,6 +6734,7 @@ class Router:
"""
idx = len(self.model_list)
self.model_list.append(model)
self._invalidate_model_group_info_cache()
self._invalidate_access_groups_cache()
# Update model_id index for O(1) lookup
@ -6774,6 +6783,7 @@ class Router:
if removal_idx is not None:
self.model_list.pop(removal_idx)
self._invalidate_model_group_info_cache()
self._invalidate_access_groups_cache()
self._update_deployment_indices_after_removal(
model_id=deployment_id, removal_idx=removal_idx
@ -6808,6 +6818,7 @@ class Router:
if deployment_idx is not None:
# Pop the item from the list first
item = self.model_list.pop(deployment_idx)
self._invalidate_model_group_info_cache()
self._invalidate_access_groups_cache()
self._update_deployment_indices_after_removal(
model_id=id, removal_idx=deployment_idx
@ -6978,9 +6989,9 @@ class Router:
raise ValueError("Deployment not found")
## GET BASE MODEL
base_model = deployment.get("model_info", {}).get("base_model", None)
base_model = (deployment.get("model_info") or {}).get("base_model", None)
if base_model is None:
base_model = deployment.get("litellm_params", {}).get("base_model", None)
base_model = (deployment.get("litellm_params") or {}).get("base_model", None)
model = base_model
@ -6995,7 +7006,7 @@ class Router:
raise ValueError(
f"Deployment missing valid litellm_params. "
f"Got: {type(litellm_params_data).__name__}, "
f"deployment_id: {deployment.get('model_info', {}).get('id', 'unknown')}"
f"deployment_id: {(deployment.get('model_info') or {}).get('id', 'unknown')}"
)
_model, custom_llm_provider, _, _ = litellm.get_llm_provider(
model=litellm_params.model,
@ -7015,10 +7026,10 @@ class Router:
if potential_models is not None:
for potential_model in potential_models:
try:
if potential_model.get("model_info", {}).get(
if (potential_model.get("model_info") or {}).get(
"id"
) == deployment.get("model_info", {}).get("id"):
model = potential_model.get("litellm_params", {}).get(
) == (deployment.get("model_info") or {}).get("id"):
model = (potential_model.get("litellm_params") or {}).get(
"model"
)
break
@ -7039,9 +7050,10 @@ class Router:
model_info = litellm.get_model_info(model=model_info_name)
## CHECK USER SET MODEL INFO
user_model_info = deployment.get("model_info", {})
user_model_info = deployment.get("model_info") or {}
model_info.update(user_model_info)
if model_info is not None:
model_info.update(user_model_info)
return model_info
@ -7568,6 +7580,7 @@ class Router:
"""
# First populate the model_list
self.model_list = []
self._invalidate_model_group_info_cache()
self._invalidate_access_groups_cache()
for _, model in enumerate(model_list):
# Extract model_info from the model dict
@ -7912,6 +7925,13 @@ class Router:
return returned_models
def _invalidate_model_group_info_cache(self) -> None:
"""Invalidate the cached model group info.
Call this whenever self.model_list is modified to ensure the cache is rebuilt.
"""
self._cached_get_model_group_info.cache_clear()
def _invalidate_access_groups_cache(self) -> None:
"""Invalidate the cached access groups.

View file

@ -55,22 +55,21 @@ def _sanitize_prometheus_label_value(value: Optional[Any]) -> Optional[str]:
return None
# Coerce non-string values (int, bool, etc.) to str before sanitizing
if not isinstance(value, str):
value = str(value)
str_value: str = value if isinstance(value, str) else str(value)
# Remove Unicode line/paragraph separators that break text format
value = value.replace("\u2028", "").replace("\u2029", "")
str_value = str_value.replace("\u2028", "").replace("\u2029", "")
# Remove carriage returns
value = value.replace("\r", "")
str_value = str_value.replace("\r", "")
# Replace newlines with spaces
value = value.replace("\n", " ")
str_value = str_value.replace("\n", " ")
# Escape backslashes and double quotes per Prometheus exposition format
value = value.replace("\\", "\\\\").replace('"', '\\"')
str_value = str_value.replace("\\", "\\\\").replace('"', '\\"')
return value
return str_value
@dataclass
@ -185,6 +184,7 @@ class UserAPIKeyLabelNames(Enum):
CLIENT_IP = "client_ip"
USER_AGENT = "user_agent"
CALLBACK_NAME = "callback_name"
STREAM = "stream"
DEFINED_PROMETHEUS_METRICS = Literal[
@ -638,6 +638,14 @@ class PrometheusMetricLabels:
]
)
# Conditionally add stream label to litellm_proxy_total_requests_metric
if (
label_name == "litellm_proxy_total_requests_metric"
and litellm.prometheus_emit_stream_label is True
and UserAPIKeyLabelNames.STREAM.value not in default_labels
):
custom_labels.append(UserAPIKeyLabelNames.STREAM.value)
return default_labels + custom_labels
@ -709,6 +717,9 @@ class UserAPIKeyLabelValues(BaseModel):
user_agent: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.USER_AGENT.value)
] = None
stream: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.STREAM.value)
] = None
class PrometheusMetricsConfig(BaseModel):

View file

@ -6,6 +6,7 @@ from typing_extensions import Any, List, Optional, TypedDict
from litellm.types.llms.base import BaseLiteLLMOpenAIResponseObject
Phase = Optional[Literal["commentary", "final_answer"]] # TODO: Once openai sdk has updated, we can remove this and use the openai sdk type
class GenericResponseOutputItemContentAnnotation(BaseLiteLLMOpenAIResponseObject):
"""Annotation for content in a message"""
@ -35,6 +36,7 @@ class OutputFunctionToolCall(BaseLiteLLMOpenAIResponseObject):
type: Optional[str] # "function_call"
id: Optional[str]
status: Literal["in_progress", "completed", "incomplete"]
phase: Phase = None
class OutputImageGenerationCall(BaseLiteLLMOpenAIResponseObject):
@ -57,6 +59,7 @@ class GenericResponseOutputItem(BaseLiteLLMOpenAIResponseObject):
status: str # "completed", "in_progress", etc.
role: str # "assistant", "user", etc.
content: List[OutputText]
phase: Phase = None
class DeleteResponseResult(BaseLiteLLMOpenAIResponseObject):

View file

@ -0,0 +1,42 @@
"""
Pydantic models for Tool Policy management endpoints.
"""
from datetime import datetime
from typing import Dict, List, Literal, Optional
from pydantic import BaseModel
ToolCallPolicy = Literal["trusted", "untrusted", "dual_llm", "blocked"]
class LiteLLM_ToolTableRow(BaseModel):
tool_id: str
tool_name: str
origin: Optional[str] = None
call_policy: ToolCallPolicy = "untrusted"
call_count: int = 0
assignments: Optional[Dict] = None
key_hash: Optional[str] = None
team_id: Optional[str] = None
key_alias: Optional[str] = None
created_at: Optional[datetime] = None
updated_at: Optional[datetime] = None
created_by: Optional[str] = None
updated_by: Optional[str] = None
class ToolListResponse(BaseModel):
tools: List[LiteLLM_ToolTableRow]
total: int
class ToolPolicyUpdateRequest(BaseModel):
tool_name: str
call_policy: ToolCallPolicy
class ToolPolicyUpdateResponse(BaseModel):
tool_name: str
call_policy: ToolCallPolicy
updated: bool

View file

@ -1826,7 +1826,7 @@ class ModelResponse(ModelResponseBase):
else:
usage = usage
elif stream is None or stream is False:
usage = Usage()
usage = None # avoid constructing throwaway Usage; set by convert_to_model_response_object
if hidden_params:
self._hidden_params = hidden_params

View file

@ -20562,6 +20562,39 @@
"supports_tool_choice": true,
"supports_vision": true
},
"gpt-5.3-codex": {
"cache_read_input_token_cost": 1.75e-07,
"cache_read_input_token_cost_priority": 3.5e-07,
"input_cost_per_token": 1.75e-06,
"input_cost_per_token_priority": 3.5e-06,
"litellm_provider": "openai",
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "responses",
"output_cost_per_token": 1.4e-05,
"output_cost_per_token_priority": 2.8e-05,
"supported_endpoints": [
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": false,
"supports_tool_choice": true,
"supports_vision": true
},
"gpt-5-mini": {
"cache_read_input_token_cost": 2.5e-08,
"cache_read_input_token_cost_flex": 1.25e-08,

View file

@ -1453,3 +1453,47 @@ def test_convert_to_model_response_object_falsy_id_preserves_auto_generated(fals
)
assert result.id == original_id
assert result.id.startswith("chatcmpl-")
def test_convert_to_model_response_object_default_usage_overwritten():
"""
Regression test: convert_to_model_response_object must properly set Usage
on a ModelResponse that only has the default Usage from ModelResponse.__init__()
(i.e. no extra litellm.Usage() set via setattr beforehand).
This validates the optimization of removing the redundant
`setattr(model_response, "usage", litellm.Usage())` in completion().
"""
mr = ModelResponse()
# usage is not set by default (optimization: avoid constructing throwaway Usage)
assert not hasattr(mr, "usage")
response_object = {
"id": "chatcmpl-usage-test",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "Hello"},
"finish_reason": "stop",
}
],
"usage": {
"prompt_tokens": 15,
"completion_tokens": 7,
"total_tokens": 22,
},
"model": "gpt-4o",
}
result = convert_to_model_response_object(
model_response_object=mr,
response_object=response_object,
stream=False,
start_time=datetime.now(),
end_time=datetime.now(),
)
assert isinstance(result, ModelResponse)
assert result.usage.prompt_tokens == 15
assert result.usage.completion_tokens == 7
assert result.usage.total_tokens == 22

View file

@ -0,0 +1,81 @@
"""
Unit tests for prometheus_emit_stream_label opt-in setting.
Tests that:
- stream label is NOT added to litellm_proxy_total_requests_metric by default
- stream label IS added when litellm.prometheus_emit_stream_label = True
- stream value is populated correctly from standard_logging_payload
"""
import pytest
import litellm
from litellm.types.integrations.prometheus import (
PrometheusMetricLabels,
UserAPIKeyLabelNames,
)
def test_stream_label_not_present_by_default():
"""stream label should NOT appear in litellm_proxy_total_requests_metric unless opted in"""
litellm.prometheus_emit_stream_label = False
labels = PrometheusMetricLabels.get_labels("litellm_proxy_total_requests_metric")
assert UserAPIKeyLabelNames.STREAM.value not in labels
def test_stream_label_present_when_opted_in():
"""stream label SHOULD appear in litellm_proxy_total_requests_metric when opted in"""
litellm.prometheus_emit_stream_label = True
try:
labels = PrometheusMetricLabels.get_labels("litellm_proxy_total_requests_metric")
assert UserAPIKeyLabelNames.STREAM.value in labels
finally:
litellm.prometheus_emit_stream_label = False
def test_stream_label_not_in_other_metrics_when_opted_in():
"""stream label should NOT be added to other metrics even when opted in"""
litellm.prometheus_emit_stream_label = True
try:
other_metrics = [
"litellm_proxy_failed_requests_metric",
"litellm_spend_metric",
"litellm_input_tokens_metric",
"litellm_output_tokens_metric",
"litellm_llm_api_latency_metric",
]
for metric in other_metrics:
labels = PrometheusMetricLabels.get_labels(metric)
assert UserAPIKeyLabelNames.STREAM.value not in labels, (
f"stream label should not be in {metric}"
)
finally:
litellm.prometheus_emit_stream_label = False
def test_stream_label_name():
"""STREAM label name should be 'stream'"""
assert UserAPIKeyLabelNames.STREAM.value == "stream"
def test_user_api_key_label_values_has_stream_field():
"""UserAPIKeyLabelValues should accept stream field"""
from litellm.types.integrations.prometheus import UserAPIKeyLabelValues
values = UserAPIKeyLabelValues(stream="True")
assert values.stream == "True"
values_false = UserAPIKeyLabelValues(stream="False")
assert values_false.stream == "False"
values_none = UserAPIKeyLabelValues()
assert values_none.stream is None
def test_stream_label_in_model_dump():
"""stream field appears in model_dump() output for use in prometheus_label_factory"""
from litellm.types.integrations.prometheus import UserAPIKeyLabelValues
values = UserAPIKeyLabelValues(stream="True")
dumped = values.model_dump()
assert "stream" in dumped
assert dumped["stream"] == "True"

View file

@ -925,4 +925,318 @@ def test_get_supported_openai_params():
assert "temperature" in params
assert "stream" in params
assert "background" in params
assert "stream" in params
assert "stream" in params
class TestPhaseParameter:
"""Tests for the `phase` parameter on assistant output items (gpt-5.3-codex)."""
def setup_method(self):
self.config = OpenAIResponsesAPIConfig()
self.model = "gpt-5.3-codex"
self.logging_obj = MagicMock()
@staticmethod
def _make_output_text(text: str):
from litellm.types.responses.main import OutputText
return OutputText(type="output_text", text=text, annotations=[])
def test_generic_response_output_item_accepts_phase_commentary(self):
from litellm.types.responses.main import GenericResponseOutputItem
item = GenericResponseOutputItem(
type="message",
id="msg_001",
status="completed",
role="assistant",
content=[self._make_output_text("Thinking...")],
phase="commentary",
)
assert item.phase == "commentary"
def test_generic_response_output_item_accepts_phase_final_answer(self):
from litellm.types.responses.main import GenericResponseOutputItem
item = GenericResponseOutputItem(
type="message",
id="msg_002",
status="completed",
role="assistant",
content=[self._make_output_text("The answer is 42.")],
phase="final_answer",
)
assert item.phase == "final_answer"
def test_generic_response_output_item_phase_defaults_to_none(self):
from litellm.types.responses.main import GenericResponseOutputItem
item = GenericResponseOutputItem(
type="message",
id="msg_003",
status="completed",
role="assistant",
content=[self._make_output_text("Hello")],
)
assert item.phase is None
def test_output_function_tool_call_accepts_phase(self):
from litellm.types.responses.main import OutputFunctionToolCall
item = OutputFunctionToolCall(
type="function_call",
id="fc_001",
arguments='{"query": "test"}',
call_id="call_001",
name="search",
status="completed",
phase="commentary",
)
assert item.phase == "commentary"
def test_input_passthrough_dict_preserves_phase(self):
"""Dict input items (the normal HTTP flow) must preserve phase verbatim."""
input_items = [
{
"type": "message",
"role": "user",
"content": [{"type": "input_text", "text": "Hi"}],
},
{
"type": "message",
"role": "assistant",
"content": [{"type": "output_text", "text": "Preamble..."}],
"phase": "commentary",
},
{
"type": "message",
"role": "assistant",
"content": [{"type": "output_text", "text": "Done."}],
"phase": "final_answer",
},
{
"type": "message",
"role": "assistant",
"content": [{"type": "output_text", "text": "Neutral."}],
"phase": None,
},
]
result = self.config._validate_input_param(input_items)
assert isinstance(result, list)
assert "phase" not in result[0]
assert result[1]["phase"] == "commentary"
assert result[2]["phase"] == "final_answer"
assert result[3]["phase"] is None
def test_input_passthrough_pydantic_preserves_non_null_phase(self):
"""Pydantic input items must preserve non-null phase values."""
from litellm.types.responses.main import GenericResponseOutputItem
item = GenericResponseOutputItem(
type="message",
id="msg_010",
status="completed",
role="assistant",
content=[self._make_output_text("commentary")],
phase="commentary",
)
result = self.config._validate_input_param([item])
assert isinstance(result, list)
assert result[0]["phase"] == "commentary"
def test_response_parsing_preserves_phase_on_output(self):
"""Non-streaming response must preserve phase on output items."""
raw_json = {
"id": "resp_001",
"created_at": 1700000000,
"model": "gpt-5.3-codex",
"object": "response",
"status": "completed",
"output": [
{
"type": "message",
"id": "msg_001",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "preamble"}],
"phase": "commentary",
},
{
"type": "message",
"id": "msg_002",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "answer"}],
"phase": "final_answer",
},
],
"usage": {"input_tokens": 10, "output_tokens": 20, "total_tokens": 30},
}
response = ResponsesAPIResponse(**raw_json)
assert len(response.output) == 2
for idx, output_item in enumerate(response.output):
if isinstance(output_item, dict):
phase = output_item.get("phase")
else:
phase = getattr(output_item, "phase", None)
expected = "commentary" if idx == 0 else "final_answer"
assert phase == expected, (
f"output[{idx}] phase={phase!r}, expected {expected!r}"
)
def test_streaming_output_item_done_preserves_phase(self):
"""OutputItemDoneEvent must preserve phase on its item."""
from litellm.types.llms.openai import (
OutputItemDoneEvent,
ResponsesAPIStreamEvents,
)
chunk = {
"type": "response.output_item.done",
"output_index": 0,
"sequence_number": 3,
"item": {
"type": "message",
"id": "msg_100",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "done"}],
"phase": "final_answer",
},
}
result = self.config.transform_streaming_response(
model=self.model, parsed_chunk=chunk, logging_obj=self.logging_obj
)
assert isinstance(result, OutputItemDoneEvent)
assert result.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE
assert getattr(result.item, "phase", None) == "final_answer"
def test_streaming_output_item_added_preserves_phase(self):
"""OutputItemAddedEvent must preserve phase on its item."""
from litellm.types.llms.openai import (
OutputItemAddedEvent,
ResponsesAPIStreamEvents,
)
chunk = {
"type": "response.output_item.added",
"output_index": 0,
"item": {
"type": "message",
"id": "msg_200",
"role": "assistant",
"phase": "commentary",
},
}
result = self.config.transform_streaming_response(
model=self.model, parsed_chunk=chunk, logging_obj=self.logging_obj
)
assert isinstance(result, OutputItemAddedEvent)
assert result.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED
assert getattr(result.item, "phase", None) == "commentary"
def test_streaming_response_completed_preserves_phase(self):
"""ResponseCompletedEvent must preserve phase on output items inside the response."""
completed_chunk = {
"type": "response.completed",
"response": {
"id": "resp_300",
"created_at": 1700000000,
"model": "gpt-5.3-codex",
"object": "response",
"status": "completed",
"output": [
{
"type": "message",
"id": "msg_300",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "final"}],
"phase": "final_answer",
}
],
"usage": {
"input_tokens": 5,
"output_tokens": 10,
"total_tokens": 15,
},
},
}
result = self.config.transform_streaming_response(
model=self.model,
parsed_chunk=completed_chunk,
logging_obj=self.logging_obj,
)
assert result.type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED
output_item = result.response.output[0]
if isinstance(output_item, dict):
assert output_item["phase"] == "final_answer"
else:
assert getattr(output_item, "phase", None) == "final_answer"
def test_phase_roundtrip_output_to_input(self):
"""Simulate full round-trip: parse response output, then send items back as input."""
raw_json = {
"id": "resp_rt",
"created_at": 1700000000,
"model": "gpt-5.3-codex",
"object": "response",
"status": "completed",
"output": [
{
"type": "message",
"id": "msg_rt1",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "preamble"}],
"phase": "commentary",
},
{
"type": "message",
"id": "msg_rt2",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "answer"}],
"phase": "final_answer",
},
],
"usage": {"input_tokens": 10, "output_tokens": 20, "total_tokens": 30},
}
response = ResponsesAPIResponse(**raw_json)
input_items = []
for item in response.output:
if isinstance(item, dict):
input_items.append(item)
else:
input_items.append(
item.model_dump() if hasattr(item, "model_dump") else dict(item)
)
input_items.append(
{
"type": "message",
"role": "user",
"content": [{"type": "input_text", "text": "next question"}],
}
)
validated = self.config._validate_input_param(input_items)
assert isinstance(validated, list)
assert validated[0]["phase"] == "commentary"
assert validated[1]["phase"] == "final_answer"
assert "phase" not in validated[2]

View file

@ -761,3 +761,70 @@ async def test_request_body_with_html_script_tags():
f"Message content with HTML was modified during parsing: "
f"expected={msg['content']!r}, got={result['messages'][2]['content']!r}"
)
def test_safe_get_request_headers_caches_on_request_state():
"""
Test that _safe_get_request_headers caches the result on request.state
and returns the same object on subsequent calls.
"""
mock_request = MagicMock()
mock_request.headers = {"content-type": "application/json", "authorization": "Bearer sk-123"}
mock_request.state = MagicMock(spec=[]) # empty spec so getattr returns default
# First call — should create and cache
result1 = _safe_get_request_headers(mock_request)
assert result1 == {"content-type": "application/json", "authorization": "Bearer sk-123"}
assert mock_request.state._cached_headers is result1
# Second call — should return the cached object (same identity)
result2 = _safe_get_request_headers(mock_request)
assert result2 is result1
def test_safe_get_request_headers_none_request():
"""
Test that _safe_get_request_headers returns empty dict for None request.
"""
result = _safe_get_request_headers(None)
assert result == {}
def test_safe_get_request_headers_copy_protects_cache():
"""
Test that callers using .copy() before mutation do not corrupt the cache.
"""
mock_request = MagicMock()
mock_request.headers = {"authorization": "Bearer sk-123", "host": "localhost"}
mock_request.state = MagicMock(spec=[])
original = _safe_get_request_headers(mock_request)
# Simulate what mutation call sites do: copy then pop
mutable = _safe_get_request_headers(mock_request).copy()
mutable.pop("authorization", None)
# Cache must be unaffected
assert "authorization" in _safe_get_request_headers(mock_request)
assert _safe_get_request_headers(mock_request) is original
def test_safe_get_request_headers_state_unavailable():
"""
Test that _safe_get_request_headers still returns headers when
request.state rejects attribute writes (the except path on the cache-write).
"""
class ReadOnlyState:
"""State object that allows reads but raises on writes."""
def __setattr__(self, name, value):
raise AttributeError("read-only state")
def __getattr__(self, name):
return None # _cached_headers not found → triggers fresh read
mock_request = MagicMock()
mock_request.headers = {"content-type": "application/json"}
mock_request.state = ReadOnlyState()
result = _safe_get_request_headers(mock_request)
assert result == {"content-type": "application/json"}

View file

@ -0,0 +1,75 @@
"""
Unit tests for ToolDiscoveryQueue.
"""
import os
import sys
import pytest
sys.path.insert(0, os.path.abspath("../../.."))
from litellm.proxy.db.db_transaction_queue.tool_discovery_queue import (
ToolDiscoveryQueue,
)
@pytest.fixture
def queue():
return ToolDiscoveryQueue()
def test_add_single_tool(queue):
queue.add_update({"tool_name": "my_tool", "origin": "user_defined"})
items = queue.flush()
assert len(items) == 1
assert items[0]["tool_name"] == "my_tool"
assert items[0]["origin"] == "user_defined"
def test_deduplication_same_name(queue):
"""Adding the same tool_name twice should only keep the first."""
queue.add_update({"tool_name": "tool_a", "origin": "mcp_server"})
queue.add_update({"tool_name": "tool_a", "origin": "user_defined"})
items = queue.flush()
assert len(items) == 1
assert items[0]["origin"] == "mcp_server" # first wins
def test_deduplication_different_names(queue):
queue.add_update({"tool_name": "tool_a"})
queue.add_update({"tool_name": "tool_b"})
items = queue.flush()
assert len(items) == 2
names = {i["tool_name"] for i in items}
assert names == {"tool_a", "tool_b"}
def test_flush_clears_pending(queue):
queue.add_update({"tool_name": "tool_x"})
items1 = queue.flush()
assert len(items1) == 1
items2 = queue.flush()
assert len(items2) == 0
def test_seen_names_reset_after_flush(queue):
"""Seen-set is cleared on flush so the same tool can re-enter the next cycle."""
queue.add_update({"tool_name": "tool_a"})
queue.flush()
queue.add_update({"tool_name": "tool_a"}) # same tool, new cycle
items = queue.flush()
assert len(items) == 1
assert items[0]["tool_name"] == "tool_a"
def test_empty_tool_name_ignored(queue):
queue.add_update({"tool_name": ""})
queue.add_update({"tool_name": None}) # type: ignore[arg-type]
items = queue.flush()
assert len(items) == 0
def test_flush_returns_list(queue):
result = queue.flush()
assert isinstance(result, list)

View file

@ -0,0 +1,197 @@
"""
Unit tests for tool_registry_writer.py — uses a mock prisma client
that exposes execute_raw / query_raw (matching the actual raw-SQL implementation).
"""
import os
import sys
from datetime import datetime, timezone
from unittest.mock import AsyncMock, MagicMock
import pytest
sys.path.insert(0, os.path.abspath("../../.."))
from litellm.proxy.db.tool_registry_writer import (
batch_upsert_tools,
get_tool,
get_tools_by_names,
list_tools,
update_tool_policy,
)
def _make_prisma(query_rows=None):
"""Return a minimal mock prisma_client with execute_raw / query_raw."""
default_row = {
"tool_id": "uuid-1",
"tool_name": "my_tool",
"origin": "user_defined",
"call_policy": "untrusted",
"call_count": 1,
"assignments": {},
"key_hash": None,
"team_id": None,
"key_alias": None,
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
"created_by": None,
"updated_by": None,
}
rows = query_rows if query_rows is not None else [default_row]
prisma = MagicMock()
prisma.db.execute_raw = AsyncMock(return_value=None)
prisma.db.query_raw = AsyncMock(return_value=rows)
return prisma
@pytest.mark.asyncio
async def test_batch_upsert_tools_calls_execute_raw():
prisma = _make_prisma()
items = [{"tool_name": "tool_a", "origin": "mcp_server", "created_by": None}]
await batch_upsert_tools(prisma, items)
prisma.db.execute_raw.assert_awaited_once()
call_args = prisma.db.execute_raw.call_args
sql = call_args.args[0]
assert "LiteLLM_ToolTable" in sql
assert "ON CONFLICT" in sql
@pytest.mark.asyncio
async def test_batch_upsert_tools_empty_list():
prisma = _make_prisma()
await batch_upsert_tools(prisma, [])
prisma.db.execute_raw.assert_not_awaited()
@pytest.mark.asyncio
async def test_batch_upsert_tools_skips_empty_names():
prisma = _make_prisma()
items = [{"tool_name": "", "origin": None}, {"tool_name": None}] # type: ignore[list-item]
await batch_upsert_tools(prisma, items)
prisma.db.execute_raw.assert_not_awaited()
@pytest.mark.asyncio
async def test_batch_upsert_multiple_tools_calls_execute_raw_per_tool():
prisma = _make_prisma()
items = [
{"tool_name": "tool_a", "origin": "mcp_server", "created_by": None},
{"tool_name": "tool_b", "origin": "user_defined", "created_by": "alice"},
]
await batch_upsert_tools(prisma, items)
assert prisma.db.execute_raw.await_count == 2
@pytest.mark.asyncio
async def test_list_tools_no_filter():
row = {
"tool_id": "id1",
"tool_name": "tool_a",
"origin": "mcp",
"call_policy": "untrusted",
"call_count": 5,
"assignments": {},
"key_hash": None,
"team_id": None,
"key_alias": None,
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
"created_by": None,
"updated_by": None,
}
prisma = _make_prisma(query_rows=[row])
result = await list_tools(prisma)
assert len(result) == 1
assert result[0].tool_name == "tool_a"
assert result[0].call_count == 5
prisma.db.query_raw.assert_awaited_once()
@pytest.mark.asyncio
async def test_list_tools_with_policy_filter():
row = {
"tool_id": "id1",
"tool_name": "blocked_tool",
"origin": None,
"call_policy": "blocked",
"call_count": 2,
"assignments": None,
"key_hash": None,
"team_id": None,
"key_alias": None,
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
"created_by": None,
"updated_by": None,
}
prisma = _make_prisma(query_rows=[row])
result = await list_tools(prisma, call_policy="blocked")
assert result[0].call_policy == "blocked"
call_args = prisma.db.query_raw.call_args
sql = call_args.args[0]
assert "WHERE call_policy" in sql
@pytest.mark.asyncio
async def test_get_tool_found():
prisma = _make_prisma()
result = await get_tool(prisma, "my_tool")
assert result is not None
assert result.tool_name == "my_tool"
prisma.db.query_raw.assert_awaited_once()
@pytest.mark.asyncio
async def test_get_tool_not_found():
prisma = _make_prisma(query_rows=[])
result = await get_tool(prisma, "nonexistent")
assert result is None
@pytest.mark.asyncio
async def test_update_tool_policy_calls_execute_raw():
row = {
"tool_id": "uuid-1",
"tool_name": "my_tool",
"origin": "user_defined",
"call_policy": "blocked",
"call_count": 1,
"assignments": {},
"key_hash": None,
"team_id": None,
"key_alias": None,
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
"created_by": None,
"updated_by": "admin",
}
prisma = _make_prisma(query_rows=[row])
result = await update_tool_policy(prisma, "my_tool", "blocked", "admin")
assert result is not None
assert result.call_policy == "blocked"
prisma.db.execute_raw.assert_awaited_once()
call_args = prisma.db.execute_raw.call_args
sql = call_args.args[0]
assert "ON CONFLICT" in sql
assert "call_policy" in sql
@pytest.mark.asyncio
async def test_get_tools_by_names_returns_policy_map():
rows = [
{"tool_name": "tool_a", "call_policy": "trusted"},
{"tool_name": "tool_b", "call_policy": "blocked"},
]
prisma = _make_prisma(query_rows=rows)
result = await get_tools_by_names(prisma, ["tool_a", "tool_b"])
assert result == {"tool_a": "trusted", "tool_b": "blocked"}
@pytest.mark.asyncio
async def test_get_tools_by_names_empty_list():
prisma = _make_prisma()
result = await get_tools_by_names(prisma, [])
assert result == {}
prisma.db.query_raw.assert_not_awaited()

View file

@ -0,0 +1,181 @@
"""
Unit tests for ToolPolicyGuardrail.
"""
import os
import sys
from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
sys.path.insert(0, os.path.abspath("../../../../../.."))
from litellm.proxy.guardrails.guardrail_hooks.tool_policy.tool_policy_guardrail import (
ToolPolicyGuardrail,
)
from litellm.types.guardrails import GuardrailEventHooks
@pytest.fixture
def guardrail():
return ToolPolicyGuardrail()
# --- helpers ---
def _tool_request_inputs(tool_names: list) -> dict:
return {
"tools": [
{"type": "function", "function": {"name": name, "description": ""}}
for name in tool_names
]
}
def _tool_response_inputs(tool_names: list) -> dict:
return {
"tool_calls": [
{"type": "function", "function": {"name": name}}
for name in tool_names
]
}
# --- tests ---
def test_guardrail_supports_pre_and_post_call(guardrail):
hooks = guardrail.supported_event_hooks
assert GuardrailEventHooks.pre_call in hooks
assert GuardrailEventHooks.post_call in hooks
@pytest.mark.asyncio
async def test_no_tools_in_request_passes_through(guardrail):
inputs: Any = {"tools": []}
result = await guardrail.apply_guardrail(
inputs=inputs, request_data={}, input_type="request"
)
assert result is inputs
@pytest.mark.asyncio
async def test_no_tool_calls_in_response_passes_through(guardrail):
inputs: Any = {"tool_calls": []}
result = await guardrail.apply_guardrail(
inputs=inputs, request_data={}, input_type="response"
)
assert result is inputs
@pytest.mark.asyncio
async def test_untrusted_tools_pass_through(guardrail):
policy_map = {"search": "untrusted", "read_file": "trusted"}
with patch.object(guardrail, "_get_policies_cached", new=AsyncMock(return_value=policy_map)):
inputs: Any = _tool_request_inputs(["search", "read_file"])
result = await guardrail.apply_guardrail(
inputs=inputs, request_data={}, input_type="request"
)
assert result is inputs
@pytest.mark.asyncio
async def test_blocked_tool_in_request_raises_http_exception(guardrail):
policy_map = {"dangerous_tool": "blocked"}
with patch.object(guardrail, "_get_policies_cached", new=AsyncMock(return_value=policy_map)):
inputs: Any = _tool_request_inputs(["dangerous_tool"])
with pytest.raises(HTTPException) as exc_info:
await guardrail.apply_guardrail(
inputs=inputs, request_data={}, input_type="request"
)
assert exc_info.value.status_code == 400
assert "dangerous_tool" in exc_info.value.detail["blocked_tools"]
@pytest.mark.asyncio
async def test_blocked_tool_in_response_raises_http_exception(guardrail):
policy_map = {"exfil_tool": "blocked"}
with patch.object(guardrail, "_get_policies_cached", new=AsyncMock(return_value=policy_map)):
inputs: Any = _tool_response_inputs(["exfil_tool"])
with pytest.raises(HTTPException) as exc_info:
await guardrail.apply_guardrail(
inputs=inputs, request_data={}, input_type="response"
)
assert exc_info.value.status_code == 400
assert "exfil_tool" in exc_info.value.detail["blocked_tools"]
@pytest.mark.asyncio
async def test_mixed_blocked_and_allowed_raises_for_blocked(guardrail):
policy_map = {"safe_tool": "trusted", "bad_tool": "blocked"}
with patch.object(guardrail, "_get_policies_cached", new=AsyncMock(return_value=policy_map)):
inputs: Any = _tool_request_inputs(["safe_tool", "bad_tool"])
with pytest.raises(HTTPException) as exc_info:
await guardrail.apply_guardrail(
inputs=inputs, request_data={}, input_type="request"
)
blocked = exc_info.value.detail["blocked_tools"]
assert "bad_tool" in blocked
assert "safe_tool" not in blocked
@pytest.mark.asyncio
async def test_tool_not_in_db_passes_through(guardrail):
"""Tools not found in the DB (no entry) should not be blocked."""
with patch.object(guardrail, "_get_policies_cached", new=AsyncMock(return_value={})):
inputs: Any = _tool_request_inputs(["unknown_tool"])
result = await guardrail.apply_guardrail(
inputs=inputs, request_data={}, input_type="request"
)
assert result is inputs
@pytest.mark.asyncio
async def test_get_policies_cached_uses_cache(guardrail):
"""Second call with same tool names should return the cached result."""
policy_map = {"tool_a": "trusted"}
with patch(
"litellm.proxy.db.tool_registry_writer.get_tools_by_names",
new=AsyncMock(return_value=policy_map),
) as mock_db, patch(
"litellm.proxy.proxy_server.prisma_client",
new=MagicMock(),
):
# first call — should hit DB
result1 = await guardrail._get_policies_cached(["tool_a"])
assert result1 == policy_map
# second call — should hit cache, not DB again
result2 = await guardrail._get_policies_cached(["tool_a"])
assert result2 == policy_map
assert mock_db.call_count == 1
@pytest.mark.asyncio
async def test_get_policies_cached_no_prisma(guardrail):
"""Without a prisma client, returns empty dict."""
with patch(
"litellm.proxy.proxy_server.prisma_client",
None,
):
result = await guardrail._get_policies_cached(["tool_a"])
assert result == {}
@pytest.mark.asyncio
async def test_response_tool_calls_as_objects(guardrail):
"""tool_calls that are objects (not dicts) with .function.name should work."""
policy_map = {"obj_tool": "blocked"}
with patch.object(guardrail, "_get_policies_cached", new=AsyncMock(return_value=policy_map)):
fn = MagicMock()
fn.name = "obj_tool"
tc = MagicMock()
tc.function = fn
inputs: Any = {"tool_calls": [tc]}
with pytest.raises(HTTPException):
await guardrail.apply_guardrail(
inputs=inputs, request_data={}, input_type="response"
)

View file

@ -0,0 +1,149 @@
"""
Unit tests for tool management endpoints (/v1/tool/*).
Uses FastAPI TestClient with mocked DB functions.
Patches target the source modules (litellm.proxy.db.tool_registry_writer.*
and litellm.proxy.proxy_server.prisma_client) because the endpoint code
imports these inside function bodies to avoid circular imports.
"""
import os
import sys
from datetime import datetime, timezone
from typing import Optional
from unittest.mock import AsyncMock, MagicMock, patch
from fastapi import FastAPI
from fastapi.testclient import TestClient
sys.path.insert(0, os.path.abspath("../../.."))
from litellm.proxy.management_endpoints.tool_management_endpoints import router
from litellm.types.tool_management import LiteLLM_ToolTableRow
# --- helpers ---
def _make_tool_row(
tool_name: str = "my_tool",
call_policy: str = "untrusted",
origin: Optional[str] = None,
) -> LiteLLM_ToolTableRow:
now = datetime.now(timezone.utc)
return LiteLLM_ToolTableRow(
tool_id="uuid-1",
tool_name=tool_name,
origin=origin,
call_policy=call_policy, # type: ignore[arg-type]
assignments={},
created_at=now,
updated_at=now,
)
def _make_app() -> FastAPI:
"""Build a minimal FastAPI app with the tool management router."""
app = FastAPI()
app.include_router(router)
return app
# Stub the auth dependency so we don't need a real proxy running.
def _override_auth():
from litellm.proxy._types import UserAPIKeyAuth
return UserAPIKeyAuth(api_key="sk-test", user_id="admin")
# A real (non-None) prisma stub for truthiness checks.
_MOCK_PRISMA = MagicMock()
# --- test class ---
class TestToolManagementEndpoints:
def setup_method(self):
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
app = _make_app()
app.dependency_overrides[user_api_key_auth] = _override_auth
self.client = TestClient(app, raise_server_exceptions=True)
@patch(
"litellm.proxy.db.tool_registry_writer.list_tools",
new_callable=AsyncMock,
)
@patch("litellm.proxy.proxy_server.prisma_client", _MOCK_PRISMA)
def test_list_tools_returns_200(self, mock_db_list):
mock_db_list.return_value = [_make_tool_row()]
resp = self.client.get("/v1/tool/list")
assert resp.status_code == 200
body = resp.json()
assert body["total"] == 1
assert body["tools"][0]["tool_name"] == "my_tool"
@patch(
"litellm.proxy.db.tool_registry_writer.list_tools",
new_callable=AsyncMock,
)
@patch("litellm.proxy.proxy_server.prisma_client", _MOCK_PRISMA)
def test_list_tools_with_policy_filter(self, mock_db_list):
mock_db_list.return_value = [_make_tool_row(call_policy="blocked")]
resp = self.client.get("/v1/tool/list?call_policy=blocked")
assert resp.status_code == 200
assert resp.json()["tools"][0]["call_policy"] == "blocked"
@patch(
"litellm.proxy.db.tool_registry_writer.get_tool",
new_callable=AsyncMock,
)
@patch("litellm.proxy.proxy_server.prisma_client", _MOCK_PRISMA)
def test_get_tool_found(self, mock_db_get):
mock_db_get.return_value = _make_tool_row(tool_name="tool_a")
resp = self.client.get("/v1/tool/tool_a")
assert resp.status_code == 200
assert resp.json()["tool_name"] == "tool_a"
@patch(
"litellm.proxy.db.tool_registry_writer.get_tool",
new_callable=AsyncMock,
)
@patch("litellm.proxy.proxy_server.prisma_client", _MOCK_PRISMA)
def test_get_tool_not_found_returns_404(self, mock_db_get):
mock_db_get.return_value = None
resp = self.client.get("/v1/tool/nonexistent", follow_redirects=True)
assert resp.status_code == 404
@patch(
"litellm.proxy.db.tool_registry_writer.update_tool_policy",
new_callable=AsyncMock,
)
@patch("litellm.proxy.proxy_server.prisma_client", _MOCK_PRISMA)
def test_update_tool_policy_blocked(self, mock_db_update):
mock_db_update.return_value = _make_tool_row(call_policy="blocked")
resp = self.client.post(
"/v1/tool/policy",
json={"tool_name": "my_tool", "call_policy": "blocked"},
)
assert resp.status_code == 200
body = resp.json()
assert body["call_policy"] == "blocked"
assert body["updated"] is True
@patch("litellm.proxy.proxy_server.prisma_client", None)
def test_list_tools_no_db_returns_500(self):
resp = self.client.get("/v1/tool/list")
assert resp.status_code == 500
def test_update_tool_policy_invalid_policy_returns_422(self):
resp = self.client.post(
"/v1/tool/policy",
json={"tool_name": "my_tool", "call_policy": "invalid_value"},
)
assert resp.status_code == 422

View file

@ -0,0 +1,402 @@
"""
Tests for AI Usage Chat module.
"""
import json
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat import (
TOOL_HANDLERS,
TOOLS_ADMIN,
TOOLS_BASE,
_build_system_prompt,
_summarise_entity_data,
_summarise_usage_data,
stream_usage_ai_chat,
)
SAMPLE_AGGREGATED_RESPONSE = {
"results": [
{
"date": "2025-01-15",
"metrics": {
"spend": 50.25,
"prompt_tokens": 20000,
"completion_tokens": 10000,
"total_tokens": 30000,
"api_requests": 500,
"successful_requests": 480,
"failed_requests": 20,
"cache_read_input_tokens": 0,
"cache_creation_input_tokens": 0,
},
"breakdown": {
"models": {
"gpt-4": {
"metrics": {
"spend": 40.0,
"api_requests": 300,
"total_tokens": 25000,
},
"metadata": {},
"api_key_breakdown": {},
},
},
"providers": {
"openai": {
"metrics": {"spend": 50.25, "api_requests": 500},
"metadata": {},
"api_key_breakdown": {},
},
},
"api_keys": {
"sk-test123": {
"metrics": {"spend": 50.25},
"metadata": {"key_alias": "Production Key"},
},
},
"model_groups": {},
"mcp_servers": {},
"entities": {},
},
},
],
"metadata": {
"total_spend": 50.25,
"total_api_requests": 500,
"total_successful_requests": 480,
"total_failed_requests": 20,
"total_tokens": 30000,
},
}
SAMPLE_TEAM_RESPONSE = {
"results": [
{
"date": "2025-01-15",
"metrics": {"spend": 100.0, "api_requests": 1000, "total_tokens": 50000},
"breakdown": {
"entities": {
"team-1": {
"metrics": {
"spend": 60.0,
"api_requests": 600,
"total_tokens": 30000,
},
"metadata": {"alias": "Engineering"},
"api_key_breakdown": {},
},
"team-2": {
"metrics": {
"spend": 40.0,
"api_requests": 400,
"total_tokens": 20000,
},
"metadata": {"alias": "Marketing"},
"api_key_breakdown": {},
},
},
"models": {},
"providers": {},
"api_keys": {},
"model_groups": {},
"mcp_servers": {},
},
},
],
"metadata": {"total_spend": 100.0, "total_api_requests": 1000},
}
class TestToolSchemas:
def test_admin_tools_include_all(self):
assert len(TOOLS_ADMIN) == 3
names = {t["function"]["name"] for t in TOOLS_ADMIN}
assert "get_usage_data" in names
assert "get_team_usage_data" in names
assert "get_tag_usage_data" in names
def test_base_tools_restricted_to_usage_only(self):
assert len(TOOLS_BASE) == 1
assert TOOLS_BASE[0]["function"]["name"] == "get_usage_data"
def test_admin_prompt_mentions_all_tools(self):
prompt = _build_system_prompt(is_admin=True)
assert "get_usage_data" in prompt
assert "get_team_usage_data" in prompt
assert "get_tag_usage_data" in prompt
def test_non_admin_prompt_only_mentions_usage_tool(self):
prompt = _build_system_prompt(is_admin=False)
assert "get_usage_data" in prompt
assert "get_team_usage_data" not in prompt
assert "get_tag_usage_data" not in prompt
def test_system_prompt_includes_todays_date(self):
from datetime import date
prompt = _build_system_prompt(is_admin=True)
assert date.today().isoformat() in prompt
class TestSummariseUsageData:
def test_summarise_includes_totals(self):
summary = _summarise_usage_data(SAMPLE_AGGREGATED_RESPONSE)
assert "$50.25" in summary
assert "500" in summary
def test_summarise_includes_models(self):
summary = _summarise_usage_data(SAMPLE_AGGREGATED_RESPONSE)
assert "gpt-4" in summary
def test_summarise_includes_providers(self):
summary = _summarise_usage_data(SAMPLE_AGGREGATED_RESPONSE)
assert "openai" in summary
def test_summarise_handles_empty_data(self):
empty = {"results": [], "metadata": {}}
summary = _summarise_usage_data(empty)
assert "no data" in summary.lower()
class TestSummariseEntityData:
def test_team_summary_includes_teams(self):
summary = _summarise_entity_data(SAMPLE_TEAM_RESPONSE, "Team")
assert "Engineering" in summary
assert "Marketing" in summary
assert "$60.0" in summary
assert "$40.0" in summary
def test_team_summary_empty(self):
empty = {"results": [], "metadata": {}}
summary = _summarise_entity_data(empty, "Team")
assert "No Team usage data" in summary
class TestStreamUsageAiChat:
@pytest.mark.asyncio
async def test_stream_emits_status_events(self):
mock_tool_call = MagicMock()
mock_tool_call.id = "call_123"
mock_tool_call.function.name = "get_usage_data"
mock_tool_call.function.arguments = json.dumps(
{
"start_date": "2025-01-01",
"end_date": "2025-01-31",
}
)
mock_first_response = MagicMock()
mock_first_response.choices = [MagicMock()]
mock_first_response.choices[0].message.tool_calls = [mock_tool_call]
mock_first_response.choices[0].message.model_dump.return_value = {
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_123",
"type": "function",
"function": {
"name": "get_usage_data",
"arguments": '{"start_date":"2025-01-01","end_date":"2025-01-31"}',
},
}
],
}
async def mock_stream():
chunk = MagicMock()
chunk.choices = [MagicMock()]
chunk.choices[0].delta.content = "Total spend is $50.25"
yield chunk
with patch(
"litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.litellm"
) as mock_litellm, patch(
"litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat._fetch_usage_data",
new_callable=AsyncMock,
) as mock_fetch:
mock_litellm.acompletion = AsyncMock(
side_effect=[
mock_first_response,
mock_stream(),
]
)
mock_fetch.return_value = SAMPLE_AGGREGATED_RESPONSE
events = []
async for event in stream_usage_ai_chat(
messages=[{"role": "user", "content": "What is my total spend?"}],
model="gpt-4o-mini",
user_id="user-123",
is_admin=True,
):
events.append(json.loads(event.replace("data: ", "").strip()))
status_events = [e for e in events if e["type"] == "status"]
tool_call_events = [e for e in events if e["type"] == "tool_call"]
chunk_events = [e for e in events if e["type"] == "chunk"]
done_events = [e for e in events if e["type"] == "done"]
assert len(status_events) >= 1
assert "Thinking" in status_events[0]["message"]
assert len(tool_call_events) >= 1
assert tool_call_events[0]["tool_name"] == "get_usage_data"
assert tool_call_events[0]["status"] in ("running", "complete")
assert len(chunk_events) >= 1
assert len(done_events) == 1
@pytest.mark.asyncio
async def test_stream_handles_team_tool(self):
mock_tool_call = MagicMock()
mock_tool_call.id = "call_team"
mock_tool_call.function.name = "get_team_usage_data"
mock_tool_call.function.arguments = json.dumps(
{
"start_date": "2025-01-01",
"end_date": "2025-01-31",
}
)
mock_first_response = MagicMock()
mock_first_response.choices = [MagicMock()]
mock_first_response.choices[0].message.tool_calls = [mock_tool_call]
mock_first_response.choices[0].message.model_dump.return_value = {
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_team",
"type": "function",
"function": {
"name": "get_team_usage_data",
"arguments": '{"start_date":"2025-01-01","end_date":"2025-01-31"}',
},
}
],
}
async def mock_stream():
chunk = MagicMock()
chunk.choices = [MagicMock()]
chunk.choices[0].delta.content = "Engineering is the top team."
yield chunk
with patch(
"litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.litellm"
) as mock_litellm, patch(
"litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat._fetch_team_usage_data",
new_callable=AsyncMock,
) as mock_fetch:
mock_litellm.acompletion = AsyncMock(
side_effect=[
mock_first_response,
mock_stream(),
]
)
mock_fetch.return_value = SAMPLE_TEAM_RESPONSE
events = []
async for event in stream_usage_ai_chat(
messages=[{"role": "user", "content": "Which team spends the most?"}],
model="gpt-4o-mini",
is_admin=True,
):
events.append(json.loads(event.replace("data: ", "").strip()))
chunk_events = [e for e in events if e["type"] == "chunk"]
assert len(chunk_events) >= 1
assert "Engineering" in chunk_events[0]["content"]
@pytest.mark.asyncio
async def test_stream_handles_error(self):
with patch(
"litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.litellm"
) as mock_litellm:
mock_litellm.acompletion = AsyncMock(side_effect=Exception("LLM error"))
events = []
async for event in stream_usage_ai_chat(
messages=[{"role": "user", "content": "test"}],
):
events.append(json.loads(event.replace("data: ", "").strip()))
error_events = [e for e in events if e["type"] == "error"]
assert len(error_events) == 1
assert "internal error" in error_events[0]["message"].lower()
@pytest.mark.asyncio
async def test_non_admin_enforces_user_id(self):
mock_tool_call = MagicMock()
mock_tool_call.id = "call_456"
mock_tool_call.function.name = "get_usage_data"
mock_tool_call.function.arguments = json.dumps(
{
"start_date": "2025-01-01",
"end_date": "2025-01-31",
"user_id": "other-user",
}
)
mock_first_response = MagicMock()
mock_first_response.choices = [MagicMock()]
mock_first_response.choices[0].message.tool_calls = [mock_tool_call]
mock_first_response.choices[0].message.model_dump.return_value = {
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_456",
"type": "function",
"function": {
"name": "get_usage_data",
"arguments": '{"start_date":"2025-01-01","end_date":"2025-01-31","user_id":"other-user"}',
},
}
],
}
async def mock_stream():
chunk = MagicMock()
chunk.choices = [MagicMock()]
chunk.choices[0].delta.content = "Data."
yield chunk
mock_fetch = AsyncMock(return_value=SAMPLE_AGGREGATED_RESPONSE)
with patch(
"litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.litellm"
) as mock_litellm, patch.dict(
"litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.TOOL_HANDLERS",
{
"get_usage_data": {
"fetch": mock_fetch,
"summarise": _summarise_usage_data,
"label": "global usage data",
}
},
):
mock_litellm.acompletion = AsyncMock(
side_effect=[
mock_first_response,
mock_stream(),
]
)
events = []
async for event in stream_usage_ai_chat(
messages=[{"role": "user", "content": "Show data"}],
model="gpt-4o-mini",
user_id="my-user-id",
is_admin=False,
):
events.append(event)
mock_fetch.assert_called_once_with(
start_date="2025-01-01",
end_date="2025-01-31",
user_id="my-user-id",
)

View file

@ -1,4 +1,6 @@
import copy
import datetime
from typing import AsyncGenerator
from unittest.mock import AsyncMock, MagicMock
import pytest
@ -1348,3 +1350,202 @@ class TestOverrideOpenAIResponseModel:
# Verify the model was not changed
assert response_obj.model == fallback_model
class TestStreamingOverheadHeader:
"""
Tests that x-litellm-overhead-duration-ms is emitted in streaming responses.
Regression tests for: streaming requests not including overhead header.
"""
def test_get_custom_headers_includes_overhead_when_set(self):
"""
get_custom_headers() returns x-litellm-overhead-duration-ms
when litellm_overhead_time_ms is in hidden_params.
"""
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
mock_user_api_key_dict.tpm_limit = None
mock_user_api_key_dict.rpm_limit = None
mock_user_api_key_dict.max_budget = None
mock_user_api_key_dict.spend = 0.0
mock_user_api_key_dict.allowed_model_region = None
hidden_params = {
"litellm_overhead_time_ms": 42.5,
"_response_ms": 500.0,
"model_id": "test-model-id",
"api_base": "https://api.openai.com",
}
headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
user_api_key_dict=mock_user_api_key_dict,
call_id="test-call-id",
model_id="test-model-id",
cache_key="",
api_base="https://api.openai.com",
version="1.0.0",
response_cost=0.001,
model_region="",
hidden_params=hidden_params,
)
assert "x-litellm-overhead-duration-ms" in headers
assert headers["x-litellm-overhead-duration-ms"] == "42.5"
def test_get_custom_headers_omits_overhead_when_none(self):
"""
get_custom_headers() omits x-litellm-overhead-duration-ms
when litellm_overhead_time_ms is not in hidden_params.
"""
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
mock_user_api_key_dict.tpm_limit = None
mock_user_api_key_dict.rpm_limit = None
mock_user_api_key_dict.max_budget = None
mock_user_api_key_dict.spend = 0.0
mock_user_api_key_dict.allowed_model_region = None
hidden_params = {
"_response_ms": 500.0,
"model_id": "test-model-id",
}
headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
user_api_key_dict=mock_user_api_key_dict,
call_id="test-call-id",
model_id="test-model-id",
cache_key="",
api_base="https://api.openai.com",
version="1.0.0",
response_cost=0.001,
model_region="",
hidden_params=hidden_params,
)
# Should be absent (None gets filtered by exclude_values)
assert "x-litellm-overhead-duration-ms" not in headers
def test_update_response_metadata_sets_overhead_on_stream_wrapper(self):
"""
update_response_metadata() sets litellm_overhead_time_ms on
a streaming response's _hidden_params when llm_api_duration_ms is available.
"""
from litellm.litellm_core_utils.llm_response_utils.response_metadata import (
update_response_metadata,
)
# Mock the logging object with llm_api_duration_ms set
mock_logging_obj = MagicMock()
mock_logging_obj.model_call_details = {
"llm_api_duration_ms": 200.0,
"litellm_params": {},
}
mock_logging_obj.caching_details = None
mock_logging_obj.callback_duration_ms = None
mock_logging_obj.litellm_call_id = "test-call-id"
mock_logging_obj._response_cost_calculator = MagicMock(return_value=0.001)
# Simulate a streaming result object with _hidden_params (like CustomStreamWrapper)
stream_result = MagicMock()
stream_result._hidden_params = {
"model_id": "test-model-id",
"api_base": "https://api.openai.com",
"additional_headers": {},
}
start_time = datetime.datetime.now() - datetime.timedelta(milliseconds=300)
end_time = datetime.datetime.now()
update_response_metadata(
result=stream_result,
logging_obj=mock_logging_obj,
model="gpt-4o",
kwargs={},
start_time=start_time,
end_time=end_time,
)
assert "litellm_overhead_time_ms" in stream_result._hidden_params
overhead = stream_result._hidden_params["litellm_overhead_time_ms"]
assert overhead is not None
assert isinstance(overhead, float)
# overhead = total_response_ms (~300ms) - llm_api_duration_ms (200ms) = ~100ms
assert overhead > 0
@pytest.mark.asyncio
async def test_streaming_response_includes_overhead_header(self):
"""
StreamingResponse returned by create_response() includes
x-litellm-overhead-duration-ms in its headers.
"""
async def mock_generator() -> AsyncGenerator[str, None]:
yield 'data: {"id":"chatcmpl-test","choices":[{"delta":{"content":"hi"}}]}\n\n'
yield "data: [DONE]\n\n"
headers = {
"x-litellm-overhead-duration-ms": "42.5",
"x-litellm-call-id": "test-call-id",
"x-litellm-model-id": "test-model-id",
}
response = await create_response(
generator=mock_generator(),
media_type="text/event-stream",
headers=headers,
)
assert isinstance(response, StreamingResponse)
assert response.headers.get("x-litellm-overhead-duration-ms") == "42.5"
def test_streaming_overhead_header_in_custom_headers_from_stream_hidden_params(
self,
):
"""
Verifies that when get_custom_headers() is called with a streaming
response's hidden_params (containing litellm_overhead_time_ms),
the x-litellm-overhead-duration-ms header is correctly populated.
This tests the critical path: update_response_metadata sets the value
→ get_custom_headers reads it → StreamingResponse header is set.
"""
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
mock_user_api_key_dict.tpm_limit = None
mock_user_api_key_dict.rpm_limit = None
mock_user_api_key_dict.max_budget = None
mock_user_api_key_dict.spend = 0.0
mock_user_api_key_dict.allowed_model_region = None
# This is what CustomStreamWrapper._hidden_params looks like after
# update_response_metadata() has been called on it
hidden_params = {
"model_id": "openai-gpt4o-deployment",
"api_base": "https://api.openai.com",
"additional_headers": {},
"litellm_overhead_time_ms": 55.3, # set by update_response_metadata
"_response_ms": 280.0,
"litellm_call_id": "test-call-id",
"response_cost": 0.002,
"cache_key": None,
"fastest_response_batch_completion": None,
"callback_duration_ms": None,
}
custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
user_api_key_dict=mock_user_api_key_dict,
call_id="test-call-id",
model_id=hidden_params.get("model_id"),
cache_key=hidden_params.get("cache_key") or "",
api_base=hidden_params.get("api_base") or "",
version="1.0.0",
response_cost=hidden_params.get("response_cost"),
model_region="",
hidden_params=hidden_params,
)
# The overhead header must be present and correct
assert "x-litellm-overhead-duration-ms" in custom_headers, (
"x-litellm-overhead-duration-ms header must be emitted during streaming. "
"It was missing — this is the streaming overhead header regression."
)
assert custom_headers["x-litellm-overhead-duration-ms"] == "55.3"

View file

@ -925,6 +925,73 @@ def test_router_get_model_access_groups_team_only_models():
assert list(access_groups.keys()) == ["default-models"]
def test_cached_get_model_group_info():
"""
Test that _cached_get_model_group_info caches results and
invalidates on deployment changes.
"""
from litellm.types.router import Deployment, LiteLLM_Params
router = litellm.Router(
model_list=[
{
"model_name": "gpt-4",
"litellm_params": {"model": "gpt-4", "api_key": "fake"},
"model_info": {"tpm": 1000, "rpm": 100},
},
]
)
# First call should compute and cache
result1 = router._cached_get_model_group_info("gpt-4")
assert result1 is not None
assert result1.tpm == 1000
# Second call should hit cache (same object)
result2 = router._cached_get_model_group_info("gpt-4")
assert result1 is result2
# Add a deployment — cache should be invalidated
router.add_deployment(
Deployment(
model_name="gpt-4",
litellm_params=LiteLLM_Params(model="gpt-4", api_key="fake2"),
model_info={"tpm": 2000, "rpm": 200},
)
)
result3 = router._cached_get_model_group_info("gpt-4")
assert result3 is not result2
assert result3 is not None
assert result3.tpm == 3000 # 1000 + 2000
# Delete a deployment — cache should be invalidated
deployment_id = router.model_list[-1]["model_info"]["id"]
router.delete_deployment(id=deployment_id)
result4 = router._cached_get_model_group_info("gpt-4")
assert result4 is not result3
assert result4 is not None
assert result4.tpm == 1000
# set_model_list — cache should be invalidated
router.set_model_list(
[
{
"model_name": "gpt-4",
"litellm_params": {"model": "gpt-4", "api_key": "fake"},
"model_info": {"tpm": 5000},
},
]
)
result5 = router._cached_get_model_group_info("gpt-4")
assert result5 is not result4
assert result5 is not None
assert result5.tpm == 5000
# Verify cache still works after invalidation
result6 = router._cached_get_model_group_info("gpt-4")
assert result5 is result6
def test_get_model_access_groups_caching():
"""
Test that get_model_access_groups caches the no-args result
@ -1297,6 +1364,61 @@ async def test_acompletion_streaming_iterator_edge_cases():
print("✓ Edge case tests passed!")
@pytest.mark.asyncio
async def test_acompletion_streaming_iterator_preserves_hidden_params():
"""
Regression test: FallbackStreamWrapper must copy _hidden_params from the
original CustomStreamWrapper so that x-litellm-overhead-duration-ms (and
other hidden params) are present in the proxy response headers for streaming.
"""
from unittest.mock import MagicMock
router = litellm.Router(
model_list=[
{
"model_name": "gpt-4",
"litellm_params": {"model": "gpt-4", "api_key": "fake-key"},
}
],
)
# Simulate a CustomStreamWrapper that already has timing metadata set by
# update_response_metadata (litellm_overhead_time_ms, _response_ms, etc.)
mock_response = MagicMock()
mock_response.model = "gpt-4"
mock_response.custom_llm_provider = "openai"
mock_response.logging_obj = MagicMock()
mock_response._hidden_params = {
"litellm_overhead_time_ms": 12.34,
"_response_ms": 500.0,
"litellm_call_id": "test-call-id",
"api_base": "https://api.openai.com",
"additional_headers": {},
}
# Make the mock iterable (yields nothing — we only care about hidden_params)
async def _empty():
return
yield # make it an async generator
mock_response.__aiter__ = lambda self: _empty().__aiter__()
result = await router._acompletion_streaming_iterator(
model_response=mock_response,
messages=[{"role": "user", "content": "hi"}],
initial_kwargs={"model": "gpt-4", "stream": True},
)
# The returned FallbackStreamWrapper must carry the original _hidden_params
assert hasattr(result, "_hidden_params"), "result must have _hidden_params"
assert result._hidden_params.get("litellm_overhead_time_ms") == 12.34, (
"litellm_overhead_time_ms must be preserved — "
"this is what drives x-litellm-overhead-duration-ms in streaming responses"
)
assert result._hidden_params.get("litellm_call_id") == "test-call-id"
assert result._hidden_params.get("_response_ms") == 500.0
@pytest.mark.asyncio
async def test_async_function_with_fallbacks_common_utils():
"""Test the async_function_with_fallbacks_common_utils method"""
@ -1858,7 +1980,7 @@ def test_get_deployment_credentials_with_provider_resolves_credential_name():
litellm_credential_name to actual credential values (for UI-created models).
"""
from litellm.types.utils import CredentialItem
# Setup credential list with a test credential
litellm.credential_list = [
CredentialItem(

View file

@ -1757,29 +1757,6 @@
"url": "https://opencollective.com/libvips"
}
},
"node_modules/@isaacs/balanced-match": {
"version": "4.0.1",
"resolved": "https://registry.npmjs.org/@isaacs/balanced-match/-/balanced-match-4.0.1.tgz",
"integrity": "sha512-yzMTt9lEb8Gv7zRioUilSglI0c0smZ9k5D65677DLWLtWJaXIS3CqcGyUFByYKlnUj6TkjLVs54fBl6+TiGQDQ==",
"dev": true,
"license": "MIT",
"engines": {
"node": "20 || >=22"
}
},
"node_modules/@isaacs/brace-expansion": {
"version": "5.0.1",
"resolved": "https://registry.npmjs.org/@isaacs/brace-expansion/-/brace-expansion-5.0.1.tgz",
"integrity": "sha512-WMz71T1JS624nWj2n2fnYAuPovhv7EUhk69R6i9dsVyzxt5eM3bjwvgk9L+APE1TRscGysAVMANkB0jh0LQZrQ==",
"dev": true,
"license": "MIT",
"dependencies": {
"@isaacs/balanced-match": "^4.0.1"
},
"engines": {
"node": "20 || >=22"
}
},
"node_modules/@istanbuljs/schema": {
"version": "0.1.3",
"resolved": "https://registry.npmjs.org/@istanbuljs/schema/-/schema-0.1.3.tgz",
@ -3696,32 +3673,6 @@
"typescript": ">=4.8.4 <6.0.0"
}
},
"node_modules/@typescript-eslint/typescript-estree/node_modules/brace-expansion": {
"version": "2.0.2",
"resolved": "https://registry.npmjs.org/brace-expansion/-/brace-expansion-2.0.2.tgz",
"integrity": "sha512-Jt0vHyM+jmUBqojB7E1NIYadt0vI0Qxjxd2TErW94wDz+E2LAm5vKMXXwg6ZZBTHPuUlDgQHKXvjGBdfcF1ZDQ==",
"dev": true,
"license": "MIT",
"dependencies": {
"balanced-match": "^1.0.0"
}
},
"node_modules/@typescript-eslint/typescript-estree/node_modules/minimatch": {
"version": "9.0.5",
"resolved": "https://registry.npmjs.org/minimatch/-/minimatch-9.0.5.tgz",
"integrity": "sha512-G6T0ZX48xgozx7587koeX9Ys2NYy6Gmv//P89sEte9V9whIapMNF4idKxnW2QtCcLiTWlb/wfCabAtAFWhhBow==",
"dev": true,
"license": "ISC",
"dependencies": {
"brace-expansion": "^2.0.1"
},
"engines": {
"node": ">=16 || 14 >=14.17"
},
"funding": {
"url": "https://github.com/sponsors/isaacs"
}
},
"node_modules/@typescript-eslint/utils": {
"version": "8.54.0",
"resolved": "https://registry.npmjs.org/@typescript-eslint/utils/-/utils-8.54.0.tgz",
@ -4749,11 +4700,14 @@
}
},
"node_modules/balanced-match": {
"version": "1.0.2",
"resolved": "https://registry.npmjs.org/balanced-match/-/balanced-match-1.0.2.tgz",
"integrity": "sha512-3oSeUO0TMV67hN1AmbXsK4yaqU7tjiHlbxRDZOpH0KW9+CeX4bRAaX0Anxt0tx2MrpRpWwQaPwIlISEJhYU5Pw==",
"version": "4.0.4",
"resolved": "https://registry.npmjs.org/balanced-match/-/balanced-match-4.0.4.tgz",
"integrity": "sha512-BLrgEcRTwX2o6gGxGOCNyMvGSp35YofuYzw9h1IMTRmKqttAZZVU67bdb9Pr2vUHA8+j3i2tJfjO6C6+4myGTA==",
"dev": true,
"license": "MIT"
"license": "MIT",
"engines": {
"node": "18 || 20 || >=22"
}
},
"node_modules/baseline-browser-mapping": {
"version": "2.9.19",
@ -4787,14 +4741,16 @@
}
},
"node_modules/brace-expansion": {
"version": "1.1.12",
"resolved": "https://registry.npmjs.org/brace-expansion/-/brace-expansion-1.1.12.tgz",
"integrity": "sha512-9T9UjW3r0UW5c1Q7GTwllptXwhvYmEzFhzMfZ9H7FQWt+uZePjZPjBP/W1ZEyZ1twGWom5/56TF4lPcqjnDHcg==",
"version": "5.0.3",
"resolved": "https://registry.npmjs.org/brace-expansion/-/brace-expansion-5.0.3.tgz",
"integrity": "sha512-fy6KJm2RawA5RcHkLa1z/ScpBeA762UF9KmZQxwIbDtRJrgLzM10depAiEQ+CXYcoiqW1/m96OAAoke2nE9EeA==",
"dev": true,
"license": "MIT",
"dependencies": {
"balanced-match": "^1.0.0",
"concat-map": "0.0.1"
"balanced-match": "^4.0.2"
},
"engines": {
"node": "18 || 20 || >=22"
}
},
"node_modules/braces": {
@ -5149,13 +5105,6 @@
"integrity": "sha512-VRhuHOLoKYOy4UbilLbUzbYg93XLjv2PncJC50EuTWPA3gaja1UjBsUP/D/9/juV3vQFr6XBEzn9KCAHdUvOHw==",
"license": "MIT"
},
"node_modules/concat-map": {
"version": "0.0.1",
"resolved": "https://registry.npmjs.org/concat-map/-/concat-map-0.0.1.tgz",
"integrity": "sha512-/Srv4dswyQNBfohGpz9o6Yb3Gz3SrUDqBH5rTuhGR7ahtlbYKnVxw2bCFMRljaA7EXHaXZ8wsHdodFvbkhKmqg==",
"dev": true,
"license": "MIT"
},
"node_modules/copy-to-clipboard": {
"version": "3.3.3",
"resolved": "https://registry.npmjs.org/copy-to-clipboard/-/copy-to-clipboard-3.3.3.tgz",
@ -6924,22 +6873,6 @@
"node": ">=10.13.0"
}
},
"node_modules/glob/node_modules/minimatch": {
"version": "10.1.1",
"resolved": "https://registry.npmjs.org/minimatch/-/minimatch-10.1.1.tgz",
"integrity": "sha512-enIvLvRAFZYXJzkCYG5RKmPfrFArdLv+R+lbQ53BmIMLIry74bjKzX6iHAm8WYamJkhSSEabrWN5D97XnKObjQ==",
"dev": true,
"license": "BlueOak-1.0.0",
"dependencies": {
"@isaacs/brace-expansion": "^5.0.0"
},
"engines": {
"node": "20 || >=22"
},
"funding": {
"url": "https://github.com/sponsors/isaacs"
}
},
"node_modules/globals": {
"version": "14.0.0",
"resolved": "https://registry.npmjs.org/globals/-/globals-14.0.0.tgz",
@ -9035,16 +8968,19 @@
}
},
"node_modules/minimatch": {
"version": "3.1.2",
"resolved": "https://registry.npmjs.org/minimatch/-/minimatch-3.1.2.tgz",
"integrity": "sha512-J7p63hRiAjw1NDEww1W7i37+ByIrOWO5XQQAzZ3VOcL0PNybwpfmV/N05zFAzwQ9USyEcX6t3UO+K5aqBQOIHw==",
"version": "10.2.2",
"resolved": "https://registry.npmjs.org/minimatch/-/minimatch-10.2.2.tgz",
"integrity": "sha512-+G4CpNBxa5MprY+04MbgOw1v7So6n5JY166pFi9KfYwT78fxScCeSNQSNzp6dpPSW2rONOps6Ocam1wFhCgoVw==",
"dev": true,
"license": "ISC",
"license": "BlueOak-1.0.0",
"dependencies": {
"brace-expansion": "^1.1.7"
"brace-expansion": "^5.0.2"
},
"engines": {
"node": "*"
"node": "18 || 20 || >=22"
},
"funding": {
"url": "https://github.com/sponsors/isaacs"
}
},
"node_modules/minimist": {
@ -12004,32 +11940,6 @@
"node": ">=18"
}
},
"node_modules/test-exclude/node_modules/brace-expansion": {
"version": "2.0.2",
"resolved": "https://registry.npmjs.org/brace-expansion/-/brace-expansion-2.0.2.tgz",
"integrity": "sha512-Jt0vHyM+jmUBqojB7E1NIYadt0vI0Qxjxd2TErW94wDz+E2LAm5vKMXXwg6ZZBTHPuUlDgQHKXvjGBdfcF1ZDQ==",
"dev": true,
"license": "MIT",
"dependencies": {
"balanced-match": "^1.0.0"
}
},
"node_modules/test-exclude/node_modules/minimatch": {
"version": "9.0.5",
"resolved": "https://registry.npmjs.org/minimatch/-/minimatch-9.0.5.tgz",
"integrity": "sha512-G6T0ZX48xgozx7587koeX9Ys2NYy6Gmv//P89sEte9V9whIapMNF4idKxnW2QtCcLiTWlb/wfCabAtAFWhhBow==",
"dev": true,
"license": "ISC",
"dependencies": {
"brace-expansion": "^2.0.1"
},
"engines": {
"node": ">=16 || 14 >=14.17"
},
"funding": {
"url": "https://github.com/sponsors/isaacs"
}
},
"node_modules/thenify": {
"version": "3.3.1",
"resolved": "https://registry.npmjs.org/thenify/-/thenify-3.3.1.tgz",
@ -13085,21 +12995,6 @@
"type": "github",
"url": "https://github.com/sponsors/wooorm"
}
},
"node_modules/@next/swc-win32-ia32-msvc": {
"version": "14.2.33",
"resolved": "https://registry.npmjs.org/@next/swc-win32-ia32-msvc/-/swc-win32-ia32-msvc-14.2.33.tgz",
"integrity": "sha512-pc9LpGNKhJ0dXQhZ5QMmYxtARwwmWLpeocFmVG5Z0DzWq5Uf0izcI8tLc+qOpqxO1PWqZ5A7J1blrUIKrIFc7Q==",
"cpu": [
"ia32"
],
"optional": true,
"os": [
"win32"
],
"engines": {
"node": ">= 10"
}
}
}
}

View file

@ -38,6 +38,7 @@ import Usage from "@/components/usage";
import UserDashboard from "@/components/user_dashboard";
import { AccessGroupsPage } from "@/components/AccessGroups/AccessGroupsPage";
import VectorStoreManagement from "@/components/vector_store_management";
import ToolPolicies from "@/components/ToolPolicies";
import SpendLogsTable from "@/components/view_logs";
import ViewUserDashboard from "@/components/view_users";
import { ThemeProvider } from "@/contexts/ThemeContext";
@ -548,6 +549,8 @@ function CreateKeyPageContent() {
<AccessGroupsPage />
) : page == "vector-stores" ? (
<VectorStoreManagement accessToken={accessToken} userRole={userRole} userID={userID} />
) : page == "tool-policies" ? (
<ToolPolicies accessToken={accessToken} userRole={userRole} />
) : page == "guardrails-monitor" ? (
<GuardrailsMonitorView accessToken={accessToken} />
) : page == "new_usage" ? (

View file

@ -0,0 +1,415 @@
"use client";
import React, { useCallback, useDeferredValue, useEffect, useState } from "react";
import { Select, Switch, Tooltip } from "antd";
import { Select, Tooltip } from "antd";
import {
Table,
TableHead,
TableHeaderCell,
TableBody,
TableRow,
TableCell,
} from "@tremor/react";
import { TimeCell } from "./view_logs/time_cell";
import { TableHeaderSortDropdown } from "./common_components/TableHeaderSortDropdown/TableHeaderSortDropdown";
import type { SortState } from "./common_components/TableHeaderSortDropdown/TableHeaderSortDropdown";
import FilterComponent, { FilterOption } from "./molecules/filter";
import { fetchToolsList, updateToolPolicy, ToolRow } from "./networking";
const POLICY_OPTIONS = [
{ value: "trusted", label: "trusted", color: "#065f46", bg: "#d1fae5", border: "#6ee7b7" },
{ value: "blocked", label: "blocked", color: "#991b1b", bg: "#fee2e2", border: "#fca5a5" },
] as const;
type PolicyValue = "trusted" | "blocked";
const policyStyle = (p: string) =>
POLICY_OPTIONS.find((o) => o.value === p) ?? POLICY_OPTIONS[1];
type SortField = "tool_name" | "call_policy" | "team_id" | "key_alias" | "created_at" | "call_count";
interface FilterValues {
[key: string]: string;
}
interface ToolPoliciesProps {
accessToken: string | null;
userRole?: string;
}
const PolicySelect: React.FC<{
value: string;
toolName: string;
saving: boolean;
onChange: (toolName: string, policy: string) => void;
}> = ({ value, toolName, saving, onChange }) => {
const style = policyStyle(value);
return (
<Select
size="small"
value={value}
disabled={saving}
loading={saving}
onChange={(v) => onChange(toolName, v)}
onClick={(e) => e.stopPropagation()}
style={{
minWidth: 110,
fontWeight: 500,
}}
styles={{
selector: {
backgroundColor: style.bg,
borderColor: style.border,
color: style.color,
borderRadius: 999,
fontSize: 11,
fontWeight: 600,
paddingLeft: 8,
paddingRight: 4,
},
}}
popupMatchSelectWidth={false}
options={POLICY_OPTIONS.map((o) => ({
value: o.value,
label: (
<span
style={{
display: "inline-flex",
alignItems: "center",
gap: 6,
fontSize: 12,
fontWeight: 500,
color: o.color,
}}
>
<span
style={{
width: 8,
height: 8,
borderRadius: "50%",
backgroundColor: o.color,
display: "inline-block",
flexShrink: 0,
}}
/>
{o.label}
</span>
),
}))}
/>
);
};
export const ToolPolicies: React.FC<ToolPoliciesProps> = ({ accessToken }) => {
const [tools, setTools] = useState<ToolRow[]>([]);
const [loading, setLoading] = useState(true);
const [isFetching, setIsFetching] = useState(false);
const [error, setError] = useState<string | null>(null);
const [saving, setSaving] = useState<string | null>(null);
const [searchTerm, setSearchTerm] = useState("");
const [sortField, setSortField] = useState<SortField>("created_at");
const [sortOrder, setSortOrder] = useState<"asc" | "desc">("desc");
const [currentPage, setCurrentPage] = useState(1);
const [isLiveTail, setIsLiveTail] = useState(true);
const [activeFilters, setActiveFilters] = useState<FilterValues>({});
const pageSize = 50;
const isFetchingDeferred = useDeferredValue(isFetching);
const isButtonLoading = isFetching || isFetchingDeferred;
const load = useCallback(async () => {
if (!accessToken) return;
setIsFetching(true);
setError(null);
try {
const rows = await fetchToolsList(accessToken);
setTools(rows);
} catch (e: any) {
setError(e.message ?? "Failed to load tools");
} finally {
setIsFetching(false);
setLoading(false);
}
}, [accessToken]);
useEffect(() => { load(); }, [load]);
useEffect(() => {
if (!isLiveTail) return;
const id = setInterval(load, 15000);
return () => clearInterval(id);
}, [isLiveTail, load]);
const handlePolicyChange = async (toolName: string, newPolicy: string) => {
if (!accessToken) return;
setSaving(toolName);
try {
await updateToolPolicy(accessToken, toolName, newPolicy);
setTools((prev) =>
prev.map((t) => (t.tool_name === toolName ? { ...t, call_policy: newPolicy } : t))
);
} catch (e: any) {
alert(`Failed to update policy: ${e.message}`);
} finally {
setSaving(null);
}
};
const handleSortChange = (field: SortField, newState: SortState) => {
if (newState === false) {
setSortField("created_at");
setSortOrder("desc");
} else {
setSortField(field);
setSortOrder(newState);
}
setCurrentPage(1);
};
const handleApplyFilters = (filters: FilterValues) => {
setActiveFilters(filters);
setCurrentPage(1);
};
const handleResetFilters = () => {
setActiveFilters({});
setCurrentPage(1);
};
// Build unique team/key options from loaded data
const teamOptions = Array.from(new Set(tools.map((t) => t.team_id).filter(Boolean))).map(
(v) => ({ label: v as string, value: v as string })
);
const keyAliasOptions = Array.from(new Set(tools.map((t) => t.key_alias).filter(Boolean))).map(
(v) => ({ label: v as string, value: v as string })
);
const filterOptions: FilterOption[] = [
{
name: "Policy",
label: "Policy",
options: POLICY_OPTIONS.map((o) => ({ label: o.label, value: o.value })),
},
{
name: "Team Name",
label: "Team Name",
options: teamOptions,
},
{
name: "Key Name",
label: "Key Name",
options: keyAliasOptions,
},
];
const SortHeader = ({ label, field }: { label: string; field: SortField }) => (
<div className="flex items-center gap-1">
<span>{label}</span>
<TableHeaderSortDropdown
sortState={sortField === field ? sortOrder : false}
onSortChange={(s) => handleSortChange(field, s)}
/>
</div>
);
const filtered = tools.filter((t) => {
if (searchTerm) {
const q = searchTerm.toLowerCase();
const matchesSearch =
t.tool_name.toLowerCase().includes(q) ||
(t.team_id ?? "").toLowerCase().includes(q) ||
(t.key_alias ?? "").toLowerCase().includes(q) ||
(t.key_hash ?? "").toLowerCase().includes(q) ||
t.call_policy.toLowerCase().includes(q);
if (!matchesSearch) return false;
}
if (activeFilters["Policy"] && t.call_policy !== activeFilters["Policy"]) return false;
if (activeFilters["Team Name"] && t.team_id !== activeFilters["Team Name"]) return false;
if (activeFilters["Key Name"] && t.key_alias !== activeFilters["Key Name"]) return false;
return true;
});
const sorted = [...filtered].sort((a, b) => {
const av = (a as any)[sortField] ?? "";
const bv = (b as any)[sortField] ?? "";
if (av < bv) return sortOrder === "desc" ? 1 : -1;
if (av > bv) return sortOrder === "desc" ? -1 : 1;
return 0;
});
const totalPages = Math.max(1, Math.ceil(sorted.length / pageSize));
const paginated = sorted.slice((currentPage - 1) * pageSize, currentPage * pageSize);
return (
<div className="p-6 w-full">
<h1 className="text-2xl font-semibold text-gray-900 mb-6">Tool Policies</h1>
<div className="bg-white rounded-lg shadow w-full max-w-full box-border">
{/* Toolbar */}
<div className="border-b px-6 py-4 w-full max-w-full box-border">
<div className="flex flex-col md:flex-row items-start md:items-center justify-between space-y-4 md:space-y-0 w-full max-w-full box-border">
<div className="flex flex-wrap items-center gap-3">
<div className="relative w-64">
<input
type="text"
placeholder="Search by Tool Name"
className="w-full px-3 py-2 pl-8 border rounded-md text-sm focus:outline-none focus:ring-2 focus:ring-blue-500 focus:border-blue-500"
value={searchTerm}
onChange={(e) => { setSearchTerm(e.target.value); setCurrentPage(1); }}
/>
<svg className="absolute left-2.5 top-2.5 h-4 w-4 text-gray-500" fill="none" stroke="currentColor" viewBox="0 0 24 24">
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M21 21l-6-6m2-5a7 7 0 11-14 0 7 7 0 0114 0z" />
</svg>
</div>
<div className="flex items-center gap-2">
<span className="text-sm font-medium text-gray-900">Live Tail</span>
<Switch color="green" checked={isLiveTail} onChange={setIsLiveTail} />
</div>
<button
onClick={load}
disabled={isButtonLoading}
className="flex items-center gap-1.5 px-3 py-2 text-sm border rounded-md hover:bg-gray-50 disabled:opacity-60"
>
<svg className={`w-4 h-4 ${isButtonLoading ? "animate-spin" : ""}`} fill="none" stroke="currentColor" viewBox="0 0 24 24">
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M4 4v5h.582m15.356 2A8.001 8.001 0 004.582 9m0 0H9m11 11v-5h-.581m0 0a8.003 8.003 0 01-15.357-2m15.357 2H15" />
</svg>
{isButtonLoading ? "Fetching" : "Fetch"}
</button>
</div>
<div className="flex items-center gap-4 text-sm text-gray-600 whitespace-nowrap">
<span>
Showing {filtered.length === 0 ? 0 : (currentPage - 1) * pageSize + 1} - {Math.min(currentPage * pageSize, filtered.length)} of {filtered.length} results
</span>
<span>Page {currentPage} of {totalPages}</span>
<div className="flex gap-1">
<button onClick={() => setCurrentPage((p) => Math.max(1, p - 1))} disabled={currentPage === 1}
className="px-3 py-1.5 border rounded-md text-sm hover:bg-gray-50 disabled:opacity-40">Previous</button>
<button onClick={() => setCurrentPage((p) => Math.min(totalPages, p + 1))} disabled={currentPage === totalPages}
className="px-3 py-1.5 border rounded-md text-sm hover:bg-gray-50 disabled:opacity-40">Next</button>
</div>
</div>
</div>
{/* Filter row */}
<div className="mt-3">
<FilterComponent
options={filterOptions}
onApplyFilters={handleApplyFilters}
onResetFilters={handleResetFilters}
buttonLabel="Filters"
/>
</div>
</div>
{/* Auto-refresh banner */}
{isLiveTail && (
<div className="bg-green-50 border-b border-green-100 px-6 py-2 flex items-center justify-between">
<span className="text-sm text-green-700">Auto-refreshing every 15 seconds</span>
<button onClick={() => setIsLiveTail(false)} className="text-xs text-green-600 underline">Stop</button>
</div>
)}
{error && (
<div className="mx-6 mt-4 p-3 bg-red-50 border border-red-200 rounded text-sm text-red-700">{error}</div>
)}
{/* Table */}
<Table className="[&_td]:py-0.5 [&_th]:py-1 w-full">
<TableHead>
<TableRow>
<TableHeaderCell className="py-1 h-8"><SortHeader label="Discovered" field="created_at" /></TableHeaderCell>
<TableHeaderCell className="py-1 h-8"><SortHeader label="Tool Name" field="tool_name" /></TableHeaderCell>
<TableHeaderCell className="py-1 h-8"><SortHeader label="Policy" field="call_policy" /></TableHeaderCell>
<TableHeaderCell className="py-1 h-8"><SortHeader label="# Calls" field="call_count" /></TableHeaderCell>
<TableHeaderCell className="py-1 h-8"><SortHeader label="Team Name" field="team_id" /></TableHeaderCell>
<TableHeaderCell className="py-1 h-8">Key Hash</TableHeaderCell>
<TableHeaderCell className="py-1 h-8"><SortHeader label="Key Name" field="key_alias" /></TableHeaderCell>
<TableHeaderCell className="py-1 h-8">Origin</TableHeaderCell>
</TableRow>
</TableHead>
<TableBody>
{loading ? (
<TableRow>
<TableCell colSpan={8} className="h-8 text-center text-gray-500">Loading tools…</TableCell>
</TableRow>
) : paginated.length === 0 ? (
<TableRow>
<TableCell colSpan={8} className="h-8 text-center text-gray-500">
No tools discovered yet. Make a chat completion that returns tool_calls to start auto-discovery.
</TableCell>
</TableRow>
) : (
paginated.map((tool) => (
<TableRow key={tool.tool_id} className="h-8 hover:bg-gray-50">
<TableCell className="py-0.5 max-h-8 overflow-hidden whitespace-nowrap">
<TimeCell utcTime={tool.created_at ?? ""} />
</TableCell>
<TableCell className="py-0.5 max-h-8 overflow-hidden">
<Tooltip title={tool.tool_name}>
<span className="font-mono text-xs max-w-[20ch] truncate block font-medium">
{tool.tool_name}
</span>
</Tooltip>
</TableCell>
<TableCell className="py-0.5 max-h-8">
<PolicySelect
value={tool.call_policy}
toolName={tool.tool_name}
saving={saving === tool.tool_name}
onChange={handlePolicyChange}
/>
</TableCell>
<TableCell className="py-0.5 max-h-8 text-right tabular-nums text-sm font-mono text-gray-700">
{(tool.call_count ?? 0).toLocaleString()}
</TableCell>
<TableCell className="py-0.5 max-h-8 overflow-hidden whitespace-nowrap">
<Tooltip title={tool.team_id ?? "-"}>
<span className="max-w-[15ch] truncate block">{tool.team_id ?? "-"}</span>
</Tooltip>
</TableCell>
<TableCell className="py-0.5 max-h-8 overflow-hidden whitespace-nowrap">
<Tooltip title={tool.key_hash ?? "-"}>
<span className="font-mono max-w-[15ch] truncate block text-blue-600">
{tool.key_hash ?? "-"}
</span>
</Tooltip>
</TableCell>
<TableCell className="py-0.5 max-h-8 overflow-hidden whitespace-nowrap">
<Tooltip title={tool.key_alias ?? "-"}>
<span className="max-w-[15ch] truncate block">{tool.key_alias ?? "-"}</span>
</Tooltip>
</TableCell>
<TableCell className="py-0.5 max-h-8 overflow-hidden whitespace-nowrap">
<Tooltip title={tool.origin ?? "-"}>
<span className="max-w-[15ch] truncate block">{tool.origin ?? "-"}</span>
</Tooltip>
</TableCell>
</TableRow>
))
)}
</TableBody>
</Table>
{/* Bottom pagination (only when > 1 page) */}
{totalPages > 1 && (
<div className="border-t px-6 py-3 flex items-center justify-between text-sm text-gray-600">
<span>Showing {(currentPage - 1) * pageSize + 1} - {Math.min(currentPage * pageSize, sorted.length)} of {sorted.length}</span>
<div className="flex gap-1">
<button onClick={() => setCurrentPage((p) => Math.max(1, p - 1))} disabled={currentPage === 1}
className="px-3 py-1.5 border rounded-md hover:bg-gray-50 disabled:opacity-40">Previous</button>
<button onClick={() => setCurrentPage((p) => Math.min(totalPages, p + 1))} disabled={currentPage === totalPages}
className="px-3 py-1.5 border rounded-md hover:bg-gray-50 disabled:opacity-40">Next</button>
</div>
</div>
)}
</div>
</div>
);
};
export default ToolPolicies;

View file

@ -0,0 +1,85 @@
import { screen } from "@testing-library/react";
import { beforeAll, describe, expect, it, vi } from "vitest";
import { renderWithProviders } from "../../../../tests/test-utils";
import UsageAIChatPanel from "./UsageAIChatPanel";
beforeAll(() => {
if (typeof window !== "undefined" && !window.ResizeObserver) {
window.ResizeObserver = class ResizeObserver {
observe() {}
unobserve() {}
disconnect() {}
} as any;
}
});
vi.mock("../../networking", () => ({
modelHubCall: vi.fn().mockResolvedValue({
data: [
{ model_group: "gpt-4" },
{ model_group: "claude-3-opus" },
],
}),
usageAiChatStream: vi.fn(),
}));
const defaultProps = {
open: true,
onClose: vi.fn(),
accessToken: "test-token",
};
describe("UsageAIChatPanel", () => {
it("should render the panel when open", () => {
renderWithProviders(<UsageAIChatPanel {...defaultProps} />);
expect(screen.getByText("Ask AI")).toBeInTheDocument();
expect(
screen.getByText("Ask about your spend, models, keys, and trends")
).toBeInTheDocument();
});
it("should render model selector", () => {
renderWithProviders(<UsageAIChatPanel {...defaultProps} />);
expect(screen.getByText("Select a model (optional, defaults to gpt-4o-mini)")).toBeInTheDocument();
});
it("should render empty state message when no conversation", () => {
renderWithProviders(<UsageAIChatPanel {...defaultProps} />);
expect(screen.getByText("Ask a question about your usage")).toBeInTheDocument();
});
it("should render the send button", () => {
renderWithProviders(<UsageAIChatPanel {...defaultProps} />);
expect(screen.getByText("Send")).toBeInTheDocument();
});
it("should render input placeholder", () => {
renderWithProviders(<UsageAIChatPanel {...defaultProps} />);
expect(screen.getByPlaceholderText("Ask about your usage...")).toBeInTheDocument();
});
it("should render clear chat button", () => {
renderWithProviders(<UsageAIChatPanel {...defaultProps} />);
expect(screen.getByText("Clear chat")).toBeInTheDocument();
});
it("should have the panel element even when closed (just off-screen)", () => {
renderWithProviders(<UsageAIChatPanel {...defaultProps} open={false} />);
expect(screen.getByTestId("usage-ai-chat-panel")).toBeInTheDocument();
expect(screen.getByTestId("usage-ai-chat-panel")).toHaveClass("translate-x-full");
});
it("should not have translate-x-full class when open", () => {
renderWithProviders(<UsageAIChatPanel {...defaultProps} open={true} />);
expect(screen.getByTestId("usage-ai-chat-panel")).not.toHaveClass("translate-x-full");
expect(screen.getByTestId("usage-ai-chat-panel")).toHaveClass("translate-x-0");
});
});

View file

@ -0,0 +1,402 @@
import React, { useEffect, useRef, useState } from "react";
import { Button, Select, Input, Spin } from "antd";
import ReactMarkdown from "react-markdown";
import { modelHubCall, usageAiChatStream, UsageAiToolCallEvent } from "../../networking";
const { TextArea } = Input;
interface ToolCallStep {
tool_name: string;
tool_label: string;
arguments: Record<string, string>;
status: "running" | "complete" | "error";
error?: string;
}
interface ChatMessage {
role: "user" | "assistant";
content: string;
toolCalls?: ToolCallStep[];
}
interface UsageAIChatPanelProps {
open: boolean;
onClose: () => void;
accessToken: string | null;
}
const TOOL_ICONS: Record<string, string> = {
get_usage_data: "📊",
get_team_usage_data: "👥",
get_tag_usage_data: "🏷️",
};
const ToolCallDisplay: React.FC<{ step: ToolCallStep }> = ({ step }) => {
const icon = TOOL_ICONS[step.tool_name] || "🔧";
const args = step.arguments;
const dateRange = args.start_date && args.end_date
? `${args.start_date} → ${args.end_date}`
: "";
const filter = args.team_ids || args.tags || args.user_id || "";
return (
<div className="flex items-start gap-2 px-3 py-2 rounded-lg bg-gray-100 border border-gray-200 text-xs">
<span className="flex-shrink-0 mt-0.5">
{step.status === "running" ? (
<Spin size="small" />
) : step.status === "error" ? (
<span className="text-red-500">✗</span>
) : (
<span className="text-green-600">✓</span>
)}
</span>
<div className="min-w-0">
<div className="font-medium text-gray-700">
{icon} {step.tool_label}
</div>
{dateRange && (
<div className="text-gray-500 mt-0.5">{dateRange}</div>
)}
{filter && (
<div className="text-gray-500 mt-0.5">Filter: {filter}</div>
)}
{step.status === "error" && step.error && (
<div className="text-red-600 mt-0.5">{step.error}</div>
)}
</div>
</div>
);
};
const MarkdownContent: React.FC<{ content: string }> = ({ content }) => (
<ReactMarkdown
components={{
p: ({ children }) => <p className="mb-2 last:mb-0">{children}</p>,
strong: ({ children }) => <strong className="font-semibold">{children}</strong>,
ul: ({ children }) => <ul className="list-disc pl-4 mb-2 space-y-0.5">{children}</ul>,
ol: ({ children }) => <ol className="list-decimal pl-4 mb-2 space-y-0.5">{children}</ol>,
li: ({ children }) => <li>{children}</li>,
h1: ({ children }) => <h4 className="font-semibold text-sm mt-2 mb-1">{children}</h4>,
h2: ({ children }) => <h4 className="font-semibold text-sm mt-2 mb-1">{children}</h4>,
h3: ({ children }) => <h4 className="font-semibold text-sm mt-2 mb-1">{children}</h4>,
code: ({ children, className }) => {
const isBlock = className?.includes("language-");
return isBlock ? (
<pre className="bg-gray-100 rounded p-2 my-1 overflow-x-auto text-xs">
<code>{children}</code>
</pre>
) : (
<code className="px-1 py-0.5 rounded bg-gray-100 text-xs font-mono">{children}</code>
);
},
table: ({ children }) => (
<div className="overflow-x-auto my-2">
<table className="text-xs border-collapse w-full">{children}</table>
</div>
),
th: ({ children }) => <th className="border border-gray-200 px-2 py-1 bg-gray-50 font-medium text-left">{children}</th>,
td: ({ children }) => <td className="border border-gray-200 px-2 py-1">{children}</td>,
}}
>
{content}
</ReactMarkdown>
);
const UsageAIChatPanel: React.FC<UsageAIChatPanelProps> = ({
open,
onClose,
accessToken,
}) => {
const [messages, setMessages] = useState<ChatMessage[]>([]);
const [inputText, setInputText] = useState("");
const [isLoading, setIsLoading] = useState(false);
const [selectedModel, setSelectedModel] = useState<string | undefined>(undefined);
const [availableModels, setAvailableModels] = useState<string[]>([]);
const [isLoadingModels, setIsLoadingModels] = useState(false);
const [streamingContent, setStreamingContent] = useState("");
const [statusMessage, setStatusMessage] = useState<string | null>(null);
const [activeToolCalls, setActiveToolCalls] = useState<ToolCallStep[]>([]);
const messagesEndRef = useRef<HTMLDivElement>(null);
const abortControllerRef = useRef<AbortController | null>(null);
useEffect(() => {
if (open && availableModels.length === 0) {
loadModels();
}
}, [open]);
useEffect(() => {
if (typeof messagesEndRef.current?.scrollIntoView === "function") {
messagesEndRef.current.scrollIntoView({ behavior: "smooth" });
}
}, [messages, streamingContent, activeToolCalls, statusMessage]);
const loadModels = async () => {
if (!accessToken) return;
setIsLoadingModels(true);
try {
const fetchedModels = await modelHubCall(accessToken);
if (fetchedModels?.data?.length > 0) {
const models = fetchedModels.data
.map((item: any) => item.model_group as string)
.sort();
setAvailableModels(models);
}
} catch (error) {
console.error("Failed to load models:", error);
} finally {
setIsLoadingModels(false);
}
};
const handleSend = async () => {
if (!accessToken || !inputText.trim() || isLoading) return;
const userMessage: ChatMessage = { role: "user", content: inputText.trim() };
const updatedMessages = [...messages, userMessage];
setMessages(updatedMessages);
setInputText("");
setIsLoading(true);
setStreamingContent("");
setStatusMessage(null);
setActiveToolCalls([]);
const abortController = new AbortController();
abortControllerRef.current = abortController;
let accumulated = "";
const toolCalls: ToolCallStep[] = [];
try {
await usageAiChatStream(
accessToken,
updatedMessages.slice(-20).map((m) => ({ role: m.role, content: m.content })),
selectedModel || "",
(content: string) => {
setStatusMessage(null);
accumulated += content;
setStreamingContent(accumulated);
},
() => {
setStatusMessage(null);
setActiveToolCalls([]);
setMessages((prev) => [
...prev,
{ role: "assistant", content: accumulated, toolCalls: toolCalls.length > 0 ? [...toolCalls] : undefined },
]);
setStreamingContent("");
},
(errorMsg: string) => {
setStatusMessage(null);
setActiveToolCalls([]);
setMessages((prev) => [
...prev,
{ role: "assistant", content: `Error: ${errorMsg}` },
]);
setStreamingContent("");
},
(status: string) => {
setStatusMessage(status);
},
(event: UsageAiToolCallEvent) => {
const idx = toolCalls.findIndex((tc) => tc.tool_name === event.tool_name);
if (idx >= 0) {
toolCalls[idx] = { ...event };
} else {
toolCalls.push({ ...event });
}
setActiveToolCalls([...toolCalls]);
},
abortController.signal,
);
} catch (error: any) {
if (error?.name === "AbortError" || abortController.signal.aborted) {
return;
}
const errorMsg = error?.message || "Failed to get response. Please try again.";
setMessages((prev) => [
...prev,
{ role: "assistant", content: `Error: ${errorMsg}` },
]);
setStreamingContent("");
} finally {
setIsLoading(false);
abortControllerRef.current = null;
}
};
const handleKeyDown = (e: React.KeyboardEvent<HTMLTextAreaElement>) => {
if (e.key === "Enter" && !e.shiftKey) {
e.preventDefault();
handleSend();
}
};
const handleClose = () => {
if (abortControllerRef.current) {
abortControllerRef.current.abort();
}
onClose();
};
const handleClear = () => {
setMessages([]);
setStreamingContent("");
setActiveToolCalls([]);
setStatusMessage(null);
};
return (
<div
data-testid="usage-ai-chat-panel"
className={`fixed top-0 right-0 h-full bg-white border-l border-gray-200 shadow-2xl z-50 flex flex-col transition-transform duration-300 ease-in-out ${
open ? "translate-x-0" : "translate-x-full"
}`}
style={{ width: 420 }}
>
{/* Header */}
<div className="px-5 pt-5 pb-3 border-b border-gray-100 flex-shrink-0">
<div className="flex items-center justify-between mb-1">
<div className="flex items-center gap-2">
<svg className="w-5 h-5 text-blue-600" viewBox="0 0 16 16" fill="currentColor">
<path d="M8 1l1.5 3.5L13 6l-3.5 1.5L8 11 6.5 7.5 3 6l3.5-1.5L8 1zm4 7l.75 1.75L14.5 10.5l-1.75.75L12 13l-.75-1.75L9.5 10.5l1.75-.75L12 8zM4 9l.75 1.75L6.5 11.5l-1.75.75L4 14l-.75-1.75L1.5 11.5l1.75-.75L4 9z" />
</svg>
<h3 className="text-base font-semibold text-gray-900">Ask AI</h3>
</div>
<button
onClick={handleClose}
className="text-gray-400 hover:text-gray-600 transition-colors p-1 rounded-md hover:bg-gray-100"
>
<svg className="w-5 h-5" fill="none" stroke="currentColor" viewBox="0 0 24 24">
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M6 18L18 6M6 6l12 12" />
</svg>
</button>
</div>
<p className="text-xs text-gray-500">
Ask about your spend, models, keys, and trends
</p>
</div>
{/* Model selector */}
<div className="px-5 py-3 border-b border-gray-100 flex-shrink-0">
<Select
placeholder="Select a model (optional, defaults to gpt-4o-mini)"
value={selectedModel}
onChange={(value) => setSelectedModel(value)}
loading={isLoadingModels}
showSearch
allowClear
size="small"
className="w-full"
options={availableModels.map((m) => ({ label: m, value: m }))}
filterOption={(input, option) =>
(option?.label ?? "").toLowerCase().includes(input.toLowerCase())
}
/>
</div>
{/* Chat messages */}
<div className="flex-1 overflow-y-auto p-4 space-y-3 bg-gray-50">
{messages.length === 0 && !streamingContent && !isLoading && (
<div className="flex flex-col items-center justify-center h-full text-gray-400">
<svg className="w-8 h-8 mb-2" fill="none" stroke="currentColor" viewBox="0 0 24 24">
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={1.5} d="M8 10h.01M12 10h.01M16 10h.01M9 16H5a2 2 0 01-2-2V6a2 2 0 012-2h14a2 2 0 012 2v8a2 2 0 01-2 2h-5l-5 5v-5z" />
</svg>
<p className="text-sm font-medium">Ask a question about your usage</p>
<p className="text-xs mt-1">e.g. &quot;Which model costs me the most?&quot;</p>
</div>
)}
{messages.map((msg, idx) => (
<div key={idx}>
{msg.role === "user" ? (
<div className="flex justify-end">
<div className="max-w-[88%] rounded-xl px-3.5 py-2 text-sm leading-relaxed bg-blue-600 text-white">
{msg.content}
</div>
</div>
) : (
<div className="space-y-2">
{/* Tool calls for this message */}
{msg.toolCalls && msg.toolCalls.length > 0 && (
<div className="space-y-1.5">
{msg.toolCalls.map((tc, tcIdx) => (
<ToolCallDisplay key={tcIdx} step={tc} />
))}
</div>
)}
{/* Response */}
<div className="max-w-[95%] rounded-xl px-3.5 py-2.5 text-sm leading-relaxed bg-white border border-gray-200 text-gray-800">
<MarkdownContent content={msg.content} />
</div>
</div>
)}
</div>
))}
{/* Active tool calls (in-progress) */}
{isLoading && activeToolCalls.length > 0 && (
<div className="space-y-1.5">
{activeToolCalls.map((tc, idx) => (
<ToolCallDisplay key={idx} step={tc} />
))}
</div>
)}
{/* Status / spinner */}
{isLoading && !streamingContent && (
<div className="flex items-center gap-2 px-3 py-2 text-xs text-gray-500">
<Spin size="small" />
<span className="italic">{statusMessage || "Thinking..."}</span>
</div>
)}
{/* Streaming response */}
{streamingContent && (
<div className="max-w-[95%] rounded-xl px-3.5 py-2.5 text-sm leading-relaxed bg-white border border-gray-200 text-gray-800">
<MarkdownContent content={streamingContent} />
</div>
)}
<div ref={messagesEndRef} />
</div>
{/* Input area */}
<div className="px-4 py-3 border-t border-gray-200 bg-white flex-shrink-0">
<div className="flex gap-2">
<TextArea
value={inputText}
onChange={(e) => setInputText(e.target.value)}
onKeyDown={handleKeyDown}
placeholder="Ask about your usage..."
autoSize={{ minRows: 1, maxRows: 3 }}
className="flex-1"
disabled={isLoading}
/>
<Button
type="primary"
onClick={handleSend}
disabled={!inputText.trim() || isLoading}
loading={isLoading}
>
Send
</Button>
</div>
<div className="flex justify-between items-center mt-2">
<button
onClick={handleClear}
className="text-xs text-gray-400 hover:text-gray-600 transition-colors"
disabled={messages.length === 0}
>
Clear chat
</button>
<span className="text-xs text-gray-400">
Enter to send
</span>
</div>
</div>
</div>
);
};
export default UsageAIChatPanel;

View file

@ -100,6 +100,10 @@ vi.mock("../../EntityUsageExport", () => ({
default: () => <div>Entity Usage Export Modal</div>,
}));
vi.mock("./UsageAIChatPanel", () => ({
default: () => <div data-testid="usage-ai-chat-panel">Usage AI Chat Panel</div>,
}));
vi.mock("@/app/(dashboard)/hooks/customers/useCustomers", () => ({
useCustomers: vi.fn(),
}));
@ -990,6 +994,28 @@ describe("UsagePage", () => {
});
});
describe("Ask AI button", () => {
it("should render Ask AI button in global view", async () => {
renderWithProviders(<UsagePage {...defaultProps} />);
await waitFor(() => {
expect(mockUserDailyActivityAggregatedCall).toHaveBeenCalled();
});
expect(screen.getByText("Ask AI")).toBeInTheDocument();
});
it("should render AI chat panel component", async () => {
renderWithProviders(<UsagePage {...defaultProps} />);
await waitFor(() => {
expect(mockUserDailyActivityAggregatedCall).toHaveBeenCalled();
});
expect(screen.getByTestId("usage-ai-chat-panel")).toBeInTheDocument();
});
});
describe("model view toggle", () => {
it("should show Public Model Name view by default", async () => {
renderWithProviders(<UsagePage {...defaultProps} />);

View file

@ -50,6 +50,7 @@ import EntityUsage, { EntityList } from "./EntityUsage/EntityUsage";
import SpendByProvider from "./EntityUsage/SpendByProvider";
import TopKeyView from "./EntityUsage/TopKeyView";
import { UsageOption, UsageViewSelect } from "./UsageViewSelect/UsageViewSelect";
import UsageAIChatPanel from "./UsageAIChatPanel";
interface UsagePageProps {
teams: Team[];
@ -142,6 +143,7 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
const [modelViewType, setModelViewType] = useState<"groups" | "individual">("groups");
const [isCloudZeroModalOpen, setIsCloudZeroModalOpen] = useState(false);
const [isGlobalExportModalOpen, setIsGlobalExportModalOpen] = useState(false);
const [isAiChatOpen, setIsAiChatOpen] = useState(false);
const [usageView, setUsageView] = useState<UsageOption>("global");
const [showCredentialBanner, setShowCredentialBanner] = useState(true);
const [topKeysLimit, setTopKeysLimit] = useState<number>(5);
@ -505,21 +507,33 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
<Tab>MCP Server Activity</Tab>
<Tab>Endpoint Activity</Tab>
</TabList>
<Button
onClick={() => setIsGlobalExportModalOpen(true)}
icon={() => (
<svg className="w-4 h-4" fill="none" stroke="currentColor" viewBox="0 0 24 24">
<path
strokeLinecap="round"
strokeLinejoin="round"
strokeWidth={2}
d="M4 16v1a3 3 0 003 3h10a3 3 0 003-3v-1m-4-4l-4 4m0 0l-4-4m4 4V4"
/>
</svg>
)}
>
Export Data
</Button>
<div className="flex items-center gap-2">
<Button
onClick={() => setIsAiChatOpen(true)}
icon={() => (
<svg className="w-4 h-4" viewBox="0 0 16 16" fill="currentColor">
<path d="M8 1l1.5 3.5L13 6l-3.5 1.5L8 11 6.5 7.5 3 6l3.5-1.5L8 1zm4 7l.75 1.75L14.5 10.5l-1.75.75L12 13l-.75-1.75L9.5 10.5l1.75-.75L12 8zM4 9l.75 1.75L6.5 11.5l-1.75.75L4 14l-.75-1.75L1.5 11.5l1.75-.75L4 9z" />
</svg>
)}
>
Ask AI
</Button>
<Button
onClick={() => setIsGlobalExportModalOpen(true)}
icon={() => (
<svg className="w-4 h-4" fill="none" stroke="currentColor" viewBox="0 0 24 24">
<path
strokeLinecap="round"
strokeLinejoin="round"
strokeWidth={2}
d="M4 16v1a3 3 0 003 3h10a3 3 0 003-3v-1m-4-4l-4 4m0 0l-4-4m4 4V4"
/>
</svg>
)}
>
Export Data
</Button>
</div>
</div>
<TabPanels>
{/* Cost Panel */}
@ -925,6 +939,13 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
selectedFilters={[]}
customTitle="Export Usage Data"
/>
{/* AI Chat Panel */}
<UsageAIChatPanel
open={isAiChatOpen}
onClose={() => setIsAiChatOpen(false)}
accessToken={accessToken}
/>
</div>
);
};

View file

@ -134,6 +134,12 @@ const menuGroups: MenuGroup[] = [
label: "Vector Stores",
icon: <DatabaseOutlined />,
},
{
key: "tool-policies",
page: "tool-policies",
label: "Tool Policies",
icon: <SafetyOutlined />,
},
],
},
],

View file

@ -5907,6 +5907,82 @@ export const enrichPolicyTemplateStream = async (
}
};
export interface UsageAiToolCallEvent {
tool_name: string;
tool_label: string;
arguments: Record<string, string>;
status: "running" | "complete" | "error";
error?: string;
}
export const usageAiChatStream = async (
accessToken: string,
messages: { role: string; content: string }[],
model: string,
onChunk: (content: string) => void,
onDone: () => void,
onError?: (error: string) => void,
onStatus?: (message: string) => void,
onToolCall?: (event: UsageAiToolCallEvent) => void,
signal?: AbortSignal,
) => {
const url = proxyBaseUrl
? `${proxyBaseUrl}/usage/ai/chat`
: `/usage/ai/chat`;
const response = await fetch(url, {
method: "POST",
headers: {
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
"Content-Type": "application/json",
},
body: JSON.stringify({ messages, model }),
signal,
});
if (!response.ok) {
const errorData = await response.json();
const errorMessage = deriveErrorMessage(errorData);
handleError(errorMessage);
throw new Error(errorMessage);
}
const reader = response.body?.getReader();
if (!reader) throw new Error("No response body");
const decoder = new TextDecoder();
let buffer = "";
while (true) {
const { done, value } = await reader.read();
if (done) break;
buffer += decoder.decode(value, { stream: true });
const lines = buffer.split("\n");
buffer = lines.pop() || "";
for (const line of lines) {
if (!line.startsWith("data: ")) continue;
try {
const event = JSON.parse(line.slice(6));
if (event.type === "chunk") {
onChunk(event.content);
} else if (event.type === "status") {
onStatus?.(event.message);
} else if (event.type === "tool_call") {
onToolCall?.(event as UsageAiToolCallEvent);
} else if (event.type === "done") {
onDone();
} else if (event.type === "error") {
onError?.(event.message);
}
} catch {
// skip malformed events
}
}
}
};
export const createPolicyCall = async (accessToken: string, policyData: any) => {
try {
const url = proxyBaseUrl ? `${proxyBaseUrl}/policies` : `/policies`;
@ -9854,3 +9930,57 @@ export const checkGdprCompliance = async (
}
return response.json();
};
export interface ToolRow {
tool_id: string;
tool_name: string;
origin?: string;
call_policy: string;
call_count?: number;
assignments?: Record<string, any>;
key_hash?: string;
team_id?: string;
key_alias?: string;
created_at?: string;
updated_at?: string;
created_by?: string;
updated_by?: string;
}
export const fetchToolsList = async (accessToken: string): Promise<ToolRow[]> => {
const url = proxyBaseUrl ? `${proxyBaseUrl}/v1/tool/list` : `/v1/tool/list`;
const response = await fetch(url, {
method: "GET",
headers: {
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
"Content-Type": "application/json",
},
});
if (!response.ok) {
const errorData = await response.text();
throw new Error(errorData);
}
const data = await response.json();
return data.tools ?? [];
};
export const updateToolPolicy = async (
accessToken: string,
toolName: string,
callPolicy: string
): Promise<ToolRow> => {
const url = proxyBaseUrl ? `${proxyBaseUrl}/v1/tool/policy` : `/v1/tool/policy`;
const response = await fetch(url, {
method: "POST",
headers: {
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
"Content-Type": "application/json",
},
body: JSON.stringify({ tool_name: toolName, call_policy: callPolicy }),
});
if (!response.ok) {
const errorData = await response.text();
throw new Error(errorData);
}
return response.json();
};

View file

@ -0,0 +1,54 @@
import { describe, it, expect } from "vitest";
import { render, screen } from "@testing-library/react";
import LatencyBasedConfiguration from "./LatencyBasedConfiguration";
describe("LatencyBasedConfiguration", () => {
it("should render the section heading", () => {
render(<LatencyBasedConfiguration routingStrategyArgs={{}} />);
expect(screen.getByText("Latency-Based Configuration")).toBeInTheDocument();
});
it("should render default params when no args are provided", () => {
render(<LatencyBasedConfiguration routingStrategyArgs={null as any} />);
// Default: ttl=3600, lowest_latency_buffer=0
expect(screen.getByDisplayValue("3600")).toBeInTheDocument();
expect(screen.getByDisplayValue("0")).toBeInTheDocument();
});
it("should render the provided routing strategy args as inputs", () => {
const args = { ttl: 7200, lowest_latency_buffer: 0.1 };
render(<LatencyBasedConfiguration routingStrategyArgs={args} />);
expect(screen.getByDisplayValue("7200")).toBeInTheDocument();
expect(screen.getByDisplayValue("0.1")).toBeInTheDocument();
});
it("should render an input with the correct name attribute for each param", () => {
const args = { ttl: 3600, lowest_latency_buffer: 0 };
render(<LatencyBasedConfiguration routingStrategyArgs={args} />);
expect(screen.getByRole("textbox", { name: /ttl/i })).toBeInTheDocument();
expect(screen.getByRole("textbox", { name: /lowest latency buffer/i })).toBeInTheDocument();
});
it("should display the TTL parameter explanation", () => {
render(<LatencyBasedConfiguration routingStrategyArgs={null as any} />);
expect(
screen.getByText(/sliding window to look back over/i)
).toBeInTheDocument();
});
it("should display the lowest_latency_buffer parameter explanation", () => {
render(<LatencyBasedConfiguration routingStrategyArgs={null as any} />);
expect(
screen.getByText(/shuffle between deployments within this %/i)
).toBeInTheDocument();
});
it("should render object values stringified into the input", () => {
const args = { ttl: { nested: true } };
render(<LatencyBasedConfiguration routingStrategyArgs={args} />);
// HTML input type=text strips newlines, so check that the key/value appears
const input = screen.getByRole("textbox", { name: /ttl/i }) as HTMLInputElement;
expect(input.value).toContain('"nested"');
expect(input.value).toContain('true');
});
});

View file

@ -0,0 +1,79 @@
import { describe, it, expect } from "vitest";
import { render, screen } from "@testing-library/react";
import ReliabilityRetriesSection from "./ReliabilityRetriesSection";
const baseSettings = {
num_retries: 3,
timeout: 30,
allowed_fails: 2,
fallbacks: ["gpt-3.5"],
context_window_fallbacks: [],
routing_strategy_args: { ttl: 3600 },
routing_strategy: "simple-shuffle",
enable_tag_filtering: false,
};
describe("ReliabilityRetriesSection", () => {
it("should render the section heading", () => {
render(<ReliabilityRetriesSection routerSettings={{}} routerFieldsMetadata={{}} />);
expect(screen.getByText("Reliability & Retries")).toBeInTheDocument();
});
it("should render input fields for non-excluded settings", () => {
render(<ReliabilityRetriesSection routerSettings={baseSettings} routerFieldsMetadata={{}} />);
expect(screen.getByDisplayValue("3")).toBeInTheDocument(); // num_retries
expect(screen.getByDisplayValue("30")).toBeInTheDocument(); // timeout
expect(screen.getByDisplayValue("2")).toBeInTheDocument(); // allowed_fails
});
it("should not render inputs for excluded keys", () => {
render(<ReliabilityRetriesSection routerSettings={baseSettings} routerFieldsMetadata={{}} />);
// Each excluded key must not produce a visible input value
const inputs = screen.queryAllByRole("textbox");
const inputNames = inputs.map((el) => el.getAttribute("name"));
expect(inputNames).not.toContain("fallbacks");
expect(inputNames).not.toContain("context_window_fallbacks");
expect(inputNames).not.toContain("routing_strategy_args");
expect(inputNames).not.toContain("routing_strategy");
expect(inputNames).not.toContain("enable_tag_filtering");
});
it("should use ui_field_name from metadata as the label", () => {
const metadata = {
num_retries: { ui_field_name: "Number of Retries", field_description: "How many times to retry" },
};
render(
<ReliabilityRetriesSection routerSettings={{ num_retries: 3 }} routerFieldsMetadata={metadata} />
);
expect(screen.getByText("Number of Retries")).toBeInTheDocument();
});
it("should fall back to the raw param name when no metadata label is available", () => {
render(
<ReliabilityRetriesSection routerSettings={{ num_retries: 3 }} routerFieldsMetadata={{}} />
);
expect(screen.getByText("num_retries")).toBeInTheDocument();
});
it("should render null values as an empty input", () => {
render(
<ReliabilityRetriesSection routerSettings={{ timeout: null }} routerFieldsMetadata={{}} />
);
const input = screen.getByRole("textbox", { name: /timeout/i }) as HTMLInputElement;
expect(input.value).toBe("");
});
it("should render object values stringified into the input", () => {
const settings = { retry_policy: { "rate-limited": 2 } };
render(<ReliabilityRetriesSection routerSettings={settings} routerFieldsMetadata={{}} />);
// HTML input type=text strips newlines, so check that the key/value appears
const input = screen.getByRole("textbox", { name: /retry_policy/i }) as HTMLInputElement;
expect(input.value).toContain('"rate-limited"');
expect(input.value).toContain('2');
});
it("should render no inputs when routerSettings is empty", () => {
render(<ReliabilityRetriesSection routerSettings={{}} routerFieldsMetadata={{}} />);
expect(screen.queryAllByRole("textbox")).toHaveLength(0);
});
});

View file

@ -0,0 +1,133 @@
import { describe, it, expect, vi } from "vitest";
import { render, screen } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import RouterSettingsForm from "./RouterSettingsForm";
import type { RouterSettingsFormValue } from "./RouterSettingsForm";
// Use the same antd mock as RoutingStrategySelector to keep things consistent
vi.mock("antd", () => ({
Select: Object.assign(
({ value, onChange, children }: any) => (
<select
data-testid="strategy-select"
value={value ?? ""}
onChange={(e) => onChange(e.target.value)}
>
{children}
</select>
),
{
Option: ({ value, children }: any) => (
<option value={value}>{children}</option>
),
}
),
}));
vi.mock("@tremor/react", async (importOriginal) => {
const actual = await importOriginal<typeof import("@tremor/react")>();
return {
...actual,
Switch: ({ checked, onChange }: any) => (
<input
type="checkbox"
role="switch"
checked={checked}
onChange={(e) => onChange(e.target.checked)}
/>
),
};
});
const defaultValue: RouterSettingsFormValue = {
routerSettings: {},
selectedStrategy: null,
enableTagFiltering: false,
};
const baseProps = {
value: defaultValue,
onChange: vi.fn(),
routerFieldsMetadata: {},
availableRoutingStrategies: [],
routingStrategyDescriptions: {},
};
describe("RouterSettingsForm", () => {
it("should render", () => {
render(<RouterSettingsForm {...baseProps} />);
expect(screen.getByText("Routing Settings")).toBeInTheDocument();
});
it("should not show the strategy selector when no strategies are provided", () => {
render(<RouterSettingsForm {...baseProps} />);
expect(screen.queryByTestId("strategy-select")).not.toBeInTheDocument();
});
it("should show the strategy selector when strategies are available", () => {
const props = {
...baseProps,
availableRoutingStrategies: ["simple-shuffle", "latency-based-routing"],
};
render(<RouterSettingsForm {...props} />);
expect(screen.getByTestId("strategy-select")).toBeInTheDocument();
});
it("should not render LatencyBasedConfiguration for non-latency strategies", () => {
const props = {
...baseProps,
value: { ...defaultValue, selectedStrategy: "simple-shuffle" },
availableRoutingStrategies: ["simple-shuffle"],
};
render(<RouterSettingsForm {...props} />);
expect(screen.queryByText("Latency-Based Configuration")).not.toBeInTheDocument();
});
it("should render LatencyBasedConfiguration when strategy is latency-based-routing", () => {
const props = {
...baseProps,
value: {
...defaultValue,
selectedStrategy: "latency-based-routing",
routerSettings: { routing_strategy_args: { ttl: 3600, lowest_latency_buffer: 0 } },
},
availableRoutingStrategies: ["latency-based-routing"],
};
render(<RouterSettingsForm {...props} />);
expect(screen.getByText("Latency-Based Configuration")).toBeInTheDocument();
});
it("should call onChange with the updated strategy when the selector changes", async () => {
const onChange = vi.fn();
const user = userEvent.setup();
const props = {
...baseProps,
onChange,
availableRoutingStrategies: ["simple-shuffle", "latency-based-routing"],
};
render(<RouterSettingsForm {...props} />);
await user.selectOptions(screen.getByTestId("strategy-select"), "latency-based-routing");
expect(onChange).toHaveBeenCalledWith(
expect.objectContaining({ selectedStrategy: "latency-based-routing" })
);
});
it("should call onChange with the updated enableTagFiltering when the toggle changes", async () => {
const onChange = vi.fn();
const user = userEvent.setup();
render(<RouterSettingsForm {...baseProps} onChange={onChange} />);
await user.click(screen.getByRole("switch"));
expect(onChange).toHaveBeenCalledWith(
expect.objectContaining({ enableTagFiltering: true })
);
});
it("should show the Reliability & Retries section", () => {
render(<RouterSettingsForm {...baseProps} />);
expect(screen.getByText("Reliability & Retries")).toBeInTheDocument();
});
});

View file

@ -0,0 +1,94 @@
import { describe, it, expect, vi } from "vitest";
import { render, screen } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import RoutingStrategySelector from "./RoutingStrategySelector";
// Ant Design's Select is complex to drive in JSDOM; swap it for a plain
// <select> so we can assert options and fire change events normally.
vi.mock("antd", () => ({
Select: Object.assign(
({ value, onChange, children }: any) => (
<div data-testid="ant-select">
<select
data-testid="strategy-select"
value={value ?? ""}
onChange={(e) => onChange(e.target.value)}
>
{children}
</select>
</div>
),
{
Option: ({ value, children }: any) => (
<option value={value}>{children}</option>
),
}
),
}));
const baseProps = {
selectedStrategy: null,
availableStrategies: ["simple-shuffle", "latency-based-routing", "least-busy"],
routingStrategyDescriptions: {
"simple-shuffle": "Randomly pick a deployment",
"latency-based-routing": "Pick the lowest-latency deployment",
},
routerFieldsMetadata: {},
onStrategyChange: vi.fn(),
};
describe("RoutingStrategySelector", () => {
it("should render", () => {
render(<RoutingStrategySelector {...baseProps} />);
expect(screen.getByTestId("ant-select")).toBeInTheDocument();
});
it("should display default label when no metadata is provided", () => {
render(<RoutingStrategySelector {...baseProps} />);
expect(screen.getByText("Routing Strategy")).toBeInTheDocument();
});
it("should display ui_field_name from metadata when provided", () => {
const props = {
...baseProps,
routerFieldsMetadata: {
routing_strategy: {
ui_field_name: "Strategy",
field_description: "How to pick a deployment",
},
},
};
render(<RoutingStrategySelector {...props} />);
expect(screen.getByText("Strategy")).toBeInTheDocument();
expect(screen.getByText("How to pick a deployment")).toBeInTheDocument();
});
it("should render all available strategies as options", () => {
render(<RoutingStrategySelector {...baseProps} />);
expect(screen.getByText("simple-shuffle")).toBeInTheDocument();
expect(screen.getByText("latency-based-routing")).toBeInTheDocument();
expect(screen.getByText("least-busy")).toBeInTheDocument();
});
it("should display strategy descriptions alongside option labels", () => {
render(<RoutingStrategySelector {...baseProps} />);
expect(screen.getByText("Randomly pick a deployment")).toBeInTheDocument();
expect(screen.getByText("Pick the lowest-latency deployment")).toBeInTheDocument();
});
it("should not render a description for a strategy that has none", () => {
render(<RoutingStrategySelector {...baseProps} />);
// "least-busy" has no entry in routingStrategyDescriptions — it still renders without crashing
expect(screen.getByText("least-busy")).toBeInTheDocument();
});
it("should call onStrategyChange with the selected strategy value", async () => {
const onStrategyChange = vi.fn();
const user = userEvent.setup();
render(<RoutingStrategySelector {...baseProps} onStrategyChange={onStrategyChange} />);
await user.selectOptions(screen.getByTestId("strategy-select"), "latency-based-routing");
expect(onStrategyChange).toHaveBeenCalledWith("latency-based-routing");
});
});

View file

@ -0,0 +1,113 @@
import { describe, it, expect, vi } from "vitest";
import { render, screen } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import TagFilteringToggle from "./TagFilteringToggle";
// setupTests.ts mocks @tremor/react but leaves Switch as the real implementation.
// Re-mock Switch as a plain checkbox so toggle interactions are trivially testable.
vi.mock("@tremor/react", async (importOriginal) => {
const actual = await importOriginal<typeof import("@tremor/react")>();
return {
...actual,
Switch: ({ checked, onChange, className }: any) => (
<input
type="checkbox"
role="switch"
checked={checked}
onChange={(e) => onChange(e.target.checked)}
className={className}
/>
),
};
});
const baseMetadata = {
enable_tag_filtering: {
ui_field_name: "Tag Filtering",
field_description: "Route requests based on tags",
link: null,
},
};
describe("TagFilteringToggle", () => {
it("should render", () => {
render(
<TagFilteringToggle enabled={false} routerFieldsMetadata={{}} onToggle={vi.fn()} />
);
expect(screen.getByRole("switch")).toBeInTheDocument();
});
it("should display default label when no metadata is provided", () => {
render(
<TagFilteringToggle enabled={false} routerFieldsMetadata={{}} onToggle={vi.fn()} />
);
expect(screen.getByText("Enable Tag Filtering")).toBeInTheDocument();
});
it("should display the label from metadata when provided", () => {
render(
<TagFilteringToggle
enabled={false}
routerFieldsMetadata={baseMetadata}
onToggle={vi.fn()}
/>
);
expect(screen.getByText("Tag Filtering")).toBeInTheDocument();
});
it("should display the description from metadata", () => {
render(
<TagFilteringToggle
enabled={false}
routerFieldsMetadata={baseMetadata}
onToggle={vi.fn()}
/>
);
expect(screen.getByText("Route requests based on tags")).toBeInTheDocument();
});
it("should render a Learn more link when metadata provides one", () => {
const metadata = {
enable_tag_filtering: {
...baseMetadata.enable_tag_filtering,
link: "https://docs.example.com/tag-filtering",
},
};
render(
<TagFilteringToggle enabled={false} routerFieldsMetadata={metadata} onToggle={vi.fn()} />
);
const link = screen.getByRole("link", { name: /learn more/i });
expect(link).toBeInTheDocument();
expect(link).toHaveAttribute("href", "https://docs.example.com/tag-filtering");
});
it("should not render a Learn more link when metadata has no link", () => {
render(
<TagFilteringToggle
enabled={false}
routerFieldsMetadata={baseMetadata}
onToggle={vi.fn()}
/>
);
expect(screen.queryByRole("link", { name: /learn more/i })).not.toBeInTheDocument();
});
it("should reflect the enabled=true state on the switch", () => {
render(
<TagFilteringToggle enabled={true} routerFieldsMetadata={{}} onToggle={vi.fn()} />
);
expect(screen.getByRole("switch")).toBeChecked();
});
it("should call onToggle with the new value when the switch is toggled", async () => {
const onToggle = vi.fn();
const user = userEvent.setup();
render(
<TagFilteringToggle enabled={false} routerFieldsMetadata={{}} onToggle={onToggle} />
);
await user.click(screen.getByRole("switch"));
expect(onToggle).toHaveBeenCalledWith(true);
});
});

View file

@ -0,0 +1,176 @@
import { describe, it, expect, vi, beforeEach } from "vitest";
import { renderWithProviders, screen, waitFor } from "../../../tests/test-utils";
import userEvent from "@testing-library/user-event";
import RouterSettings from "./index";
vi.mock("antd", () => ({
Select: Object.assign(
({ value, onChange, children }: any) => (
<select
data-testid="strategy-select"
value={value ?? ""}
onChange={(e) => onChange(e.target.value)}
>
{children}
</select>
),
{
Option: ({ value, children }: any) => (
<option value={value}>{children}</option>
),
}
),
}));
vi.mock("@tremor/react", async (importOriginal) => {
const actual = await importOriginal<typeof import("@tremor/react")>();
return {
...actual,
Switch: ({ checked, onChange }: any) => (
<input
type="checkbox"
role="switch"
checked={checked}
onChange={(e) => onChange(e.target.checked)}
/>
),
};
});
vi.mock("@/components/networking", () => ({
getCallbacksCall: vi.fn(),
getRouterSettingsCall: vi.fn(),
setCallbacksCall: vi.fn(),
}));
import {
getCallbacksCall,
getRouterSettingsCall,
setCallbacksCall,
} from "@/components/networking";
import NotificationsManager from "@/components/molecules/notifications_manager";
const mockCallbacksResponse = {
router_settings: {
routing_strategy: "simple-shuffle",
num_retries: 3,
timeout: 30,
},
};
const mockRouterSettingsResponse = {
fields: [
{
field_name: "routing_strategy",
ui_field_name: "Routing Strategy",
field_description: "How requests are distributed",
options: ["simple-shuffle", "latency-based-routing"],
link: null,
},
{
field_name: "enable_tag_filtering",
ui_field_name: "Tag Filtering",
field_description: "Route by tag",
field_value: false,
link: null,
},
],
routing_strategy_descriptions: {
"simple-shuffle": "Randomly pick a deployment",
"latency-based-routing": "Pick the lowest-latency deployment",
},
};
const defaultProps = {
accessToken: "test-token",
userRole: "Admin",
userID: "user-1",
modelData: null,
};
describe("RouterSettings", () => {
beforeEach(() => {
vi.clearAllMocks();
vi.mocked(getCallbacksCall).mockResolvedValue(mockCallbacksResponse);
vi.mocked(getRouterSettingsCall).mockResolvedValue(mockRouterSettingsResponse);
vi.mocked(setCallbacksCall).mockResolvedValue({});
});
it("should render nothing when accessToken is null", () => {
const { container } = renderWithProviders(
<RouterSettings {...defaultProps} accessToken={null} />
);
expect(container).toBeEmptyDOMElement();
});
it("should render the Save Changes and Reset buttons when authenticated", () => {
renderWithProviders(<RouterSettings {...defaultProps} />);
expect(screen.getByRole("button", { name: /save changes/i })).toBeInTheDocument();
expect(screen.getByRole("button", { name: /reset/i })).toBeInTheDocument();
});
it("should fetch callbacks and router settings on mount", async () => {
renderWithProviders(<RouterSettings {...defaultProps} />);
await waitFor(() => {
expect(getCallbacksCall).toHaveBeenCalledWith("test-token", "user-1", "Admin");
});
expect(getRouterSettingsCall).toHaveBeenCalledWith("test-token");
});
it("should not fetch data when any required prop is missing", () => {
renderWithProviders(
<RouterSettings {...defaultProps} userRole={null} />
);
expect(getCallbacksCall).not.toHaveBeenCalled();
});
it("should render routing strategies loaded from the API", async () => {
renderWithProviders(<RouterSettings {...defaultProps} />);
await waitFor(() => {
expect(screen.getByTestId("strategy-select")).toBeInTheDocument();
});
const select = screen.getByTestId("strategy-select") as HTMLSelectElement;
const optionValues = Array.from(select.options).map((o) => o.value);
expect(optionValues).toContain("simple-shuffle");
expect(optionValues).toContain("latency-based-routing");
});
it("should call setCallbacksCall with updated settings on Save Changes", async () => {
const user = userEvent.setup();
renderWithProviders(<RouterSettings {...defaultProps} />);
// Wait for the strategy select to appear — it only renders after getRouterSettingsCall resolves
await waitFor(() => {
expect(screen.getByTestId("strategy-select")).toBeInTheDocument();
});
await user.click(screen.getByRole("button", { name: /save changes/i }));
expect(setCallbacksCall).toHaveBeenCalledWith(
"test-token",
expect.objectContaining({
router_settings: expect.objectContaining({
routing_strategy: "simple-shuffle",
}),
})
);
});
it("should show a success notification after saving", async () => {
const user = userEvent.setup();
renderWithProviders(<RouterSettings {...defaultProps} />);
// Wait for data to load before interacting
await waitFor(() => {
expect(screen.getByTestId("strategy-select")).toBeInTheDocument();
});
await user.click(screen.getByRole("button", { name: /save changes/i }));
expect(NotificationsManager.success).toHaveBeenCalledWith(
"router settings updated successfully"
);
});
});

View file

@ -363,31 +363,101 @@ describe("KeyInfoView", () => {
});
it("should show edit button in settings tab when user has write access", async () => {
vi.mocked(useAuthorized).mockReturnValue({
...baseUseAuthorizedMock,
userRole: "Admin",
describe("'Edit Settings' button visibility in the Settings tab", () => {
const renderAndOpenSettingsTab = async (keyData = MOCK_KEY_DATA) => {
render(
<KeyInfoView
keyData={keyData}
onClose={() => {}}
keyId="test-key-id"
onKeyDataUpdate={() => {}}
teams={[]}
/>,
);
await waitFor(() => {
expect(screen.getByRole("tab", { name: /settings/i })).toBeInTheDocument();
});
await userEvent.click(screen.getByRole("tab", { name: /settings/i }));
};
it("should show the Edit Settings button when the user is a proxy admin for a key they do not own", async () => {
vi.mocked(useAuthorized).mockReturnValue({
...baseUseAuthorizedMock,
userId: "proxy-admin-user-id",
userRole: "proxy_admin",
});
await renderAndOpenSettingsTab({ ...MOCK_KEY_DATA, user_id: "someone-else-id" });
expect(screen.getByRole("button", { name: /edit settings/i })).toBeInTheDocument();
});
render(
<KeyInfoView
keyData={MOCK_KEY_DATA}
onClose={() => { }}
keyId={"test-key-id"}
onKeyDataUpdate={() => { }}
teams={[]}
/>,
);
it("should show the Edit Settings button when the user is the key owner", async () => {
vi.mocked(useAuthorized).mockReturnValue({
...baseUseAuthorizedMock,
userId: "owner-user-id",
userRole: "Internal User",
});
await waitFor(() => {
const settingsTab = screen.getByRole("tab", { name: /settings/i });
expect(settingsTab).toBeInTheDocument();
await renderAndOpenSettingsTab({ ...MOCK_KEY_DATA, user_id: "owner-user-id" });
expect(screen.getByRole("button", { name: /edit settings/i })).toBeInTheDocument();
});
const settingsTab = screen.getByRole("tab", { name: /settings/i });
await userEvent.click(settingsTab);
it("should not show the Edit Settings button when an Internal User does not own the key", async () => {
vi.mocked(useAuthorized).mockReturnValue({
...baseUseAuthorizedMock,
userId: "non-owner-user-id",
userRole: "Internal User",
});
await renderAndOpenSettingsTab({ ...MOCK_KEY_DATA, user_id: "owner-user-id" });
expect(screen.queryByRole("button", { name: /edit settings/i })).not.toBeInTheDocument();
});
it("should not show the Edit Settings button when the user is an Internal Viewer even if they own the key", async () => {
vi.mocked(useAuthorized).mockReturnValue({
...baseUseAuthorizedMock,
userId: "owner-user-id",
userRole: "Internal Viewer",
});
await renderAndOpenSettingsTab({ ...MOCK_KEY_DATA, user_id: "owner-user-id" });
expect(screen.queryByRole("button", { name: /edit settings/i })).not.toBeInTheDocument();
});
it("should show the Edit Settings button when the user is a team admin for the key's team", async () => {
const teamId = "test-team-id";
const teamAdminUserId = "team-admin-user";
vi.mocked(useTeams).mockReturnValue({
teams: [
{
team_id: teamId,
team_alias: "Test Team",
models: [],
max_budget: null,
budget_duration: null,
tpm_limit: null,
rpm_limit: null,
organization_id: "org-1",
created_at: "2025-01-01T00:00:00Z",
keys: [],
members_with_roles: [{ user_id: teamAdminUserId, role: "admin" }],
spend: 0,
},
],
setTeams: vi.fn(),
});
vi.mocked(useAuthorized).mockReturnValue({
...baseUseAuthorizedMock,
userId: teamAdminUserId,
userRole: "user",
});
await renderAndOpenSettingsTab({ ...MOCK_KEY_DATA, team_id: teamId, user_id: "other-user-id" });
await waitFor(() => {
expect(screen.getByRole("button", { name: /edit settings/i })).toBeInTheDocument();
});
});

View file

@ -595,7 +595,7 @@ export default function KeyInfoView({
<Card className="overflow-y-auto max-h-[65vh]">
<div className="flex justify-between items-center mb-4">
<Title>Key Settings</Title>
{!isEditing && userRole && rolesWithWriteAccess.includes(userRole) && (
{!isEditing && canModifyKey && (
<Button onClick={() => setIsEditing(true)}>Edit Settings</Button>
)}
</div>