mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge remote-tracking branch 'origin/main' into litellm_batch_update_database_tasks
This commit is contained in:
commit
e1e05d27f8
73 changed files with 5657 additions and 262 deletions
|
|
@ -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"
|
||||
}
|
||||
|
|
|
|||
145
docs/my-website/blog/gpt_5_3_codex/index.md
Normal file
145
docs/my-website/blog/gpt_5_3_codex/index.md
Normal 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.
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
179
litellm/proxy/db/tool_registry_writer.py
Normal file
179
litellm/proxy/db/tool_registry_writer.py
Normal 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 {}
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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=(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
149
litellm/proxy/management_endpoints/tool_management_endpoints.py
Normal file
149
litellm/proxy/management_endpoints/tool_management_endpoints.py
Normal 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))
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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.",
|
||||
}
|
||||
)
|
||||
|
|
@ -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"},
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
42
litellm/types/tool_management.py
Normal file
42
litellm/types/tool_management.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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]
|
||||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
197
tests/test_litellm/proxy/db/test_tool_registry_writer.py
Normal file
197
tests/test_litellm/proxy/db/test_tool_registry_writer.py
Normal 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()
|
||||
|
|
@ -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"
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
@ -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",
|
||||
)
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
151
ui/litellm-dashboard/package-lock.json
generated
151
ui/litellm-dashboard/package-lock.json
generated
|
|
@ -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"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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" ? (
|
||||
|
|
|
|||
415
ui/litellm-dashboard/src/components/ToolPolicies.tsx
Normal file
415
ui/litellm-dashboard/src/components/ToolPolicies.tsx
Normal 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;
|
||||
|
|
@ -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");
|
||||
});
|
||||
});
|
||||
|
|
@ -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. "Which model costs me the most?"</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;
|
||||
|
|
@ -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} />);
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
};
|
||||
|
|
|
|||
|
|
@ -134,6 +134,12 @@ const menuGroups: MenuGroup[] = [
|
|||
label: "Vector Stores",
|
||||
icon: <DatabaseOutlined />,
|
||||
},
|
||||
{
|
||||
key: "tool-policies",
|
||||
page: "tool-policies",
|
||||
label: "Tool Policies",
|
||||
icon: <SafetyOutlined />,
|
||||
},
|
||||
],
|
||||
},
|
||||
],
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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');
|
||||
});
|
||||
});
|
||||
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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");
|
||||
});
|
||||
});
|
||||
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
|
|
@ -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"
|
||||
);
|
||||
});
|
||||
});
|
||||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue