mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge branch 'litellm_dev_02_25_2026_p2' into litellm_03_03_2026_tool_policies_demo
Resolve conflicts in schema files (formatting), _types.py (import style), tool_registry_writer.py (keep raw SQL approach with agent_id enrichment), proxy_server.py (keep InFlightRequestsMiddleware import), and ToolPolicies.tsx (use extracted PolicySelect component). Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
commit
f81152d5e1
34 changed files with 2648 additions and 455 deletions
|
|
@ -109,6 +109,8 @@ Key files:
|
|||
- `litellm/proxy/auth/` - Authentication logic
|
||||
- `litellm/proxy/management_endpoints/` - Admin API endpoints
|
||||
|
||||
**Database (proxy)**: Use Prisma model methods (`prisma_client.db.<model>.upsert`, `.find_many`, `.find_unique`, etc.), not raw SQL (`execute_raw`/`query_raw`). See COMMON PITFALLS for details.
|
||||
|
||||
## MCP (MODEL CONTEXT PROTOCOL) SUPPORT
|
||||
|
||||
LiteLLM supports MCP for agent workflows:
|
||||
|
|
@ -176,6 +178,7 @@ When opening issues or pull requests, follow these templates:
|
|||
5. **Dependencies**: Keep dependencies minimal and well-justified
|
||||
6. **UI/Backend Contract Mismatch**: When adding a new entity type to the UI, always check whether the backend endpoint accepts a single value or an array. Match the UI control accordingly (single-select vs. multi-select) to avoid silently dropping user selections
|
||||
7. **Missing Tests for New Entity Types**: When adding a new entity type (e.g., in `EntityUsage`, `UsageViewSelect`), always add corresponding tests in the existing test files and update any icon/component mocks
|
||||
8. **Raw SQL in proxy DB code**: Do not use `execute_raw` or `query_raw` for proxy database access. Use Prisma model methods (e.g. `prisma_client.db.litellm_tooltable.upsert()`, `.find_many()`, `.find_unique()`) so behavior stays consistent with the schema, the client stays mockable in tests, and you avoid the pitfalls of hand-written SQL (parameter ordering, type casting, schema drift)
|
||||
|
||||
8. **Do not hardcode model-specific flags**: Put model-specific capability flags in `model_prices_and_context_window.json` and read them via `get_model_info` (or existing helpers like `supports_reasoning`). This prevents users from needing to upgrade LiteLLM each time a new model supports a feature.
|
||||
|
||||
|
|
|
|||
|
|
@ -107,6 +107,10 @@ LiteLLM is a unified interface for 100+ LLM providers with two main components:
|
|||
- Migration files auto-generated with `prisma migrate dev`
|
||||
- Always test migrations against both PostgreSQL and SQLite
|
||||
|
||||
### Proxy database access
|
||||
- **Do not write raw SQL** for proxy DB operations. Use Prisma model methods instead of `execute_raw` / `query_raw`.
|
||||
- Use the generated client: `prisma_client.db.<model>` (e.g. `litellm_tooltable`, `litellm_usertable`) with `.upsert()`, `.find_many()`, `.find_unique()`, `.update()`, `.update_many()` as appropriate. This avoids schema/client drift, keeps code testable with simple mocks, and matches patterns used in spend logs and other proxy code.
|
||||
|
||||
### Enterprise Features
|
||||
- Enterprise-specific code in `enterprise/` directory
|
||||
- Optional features enabled via environment variables
|
||||
|
|
|
|||
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_ObjectPermissionTable" ADD COLUMN "blocked_tools" TEXT[] DEFAULT ARRAY[]::TEXT[];
|
||||
|
|
@ -0,0 +1,11 @@
|
|||
-- CreateTable
|
||||
CREATE TABLE "LiteLLM_SpendLogToolIndex" (
|
||||
"request_id" TEXT NOT NULL,
|
||||
"tool_name" TEXT NOT NULL,
|
||||
"start_time" TIMESTAMP(3) NOT NULL,
|
||||
|
||||
CONSTRAINT "LiteLLM_SpendLogToolIndex_pkey" PRIMARY KEY ("request_id","tool_name")
|
||||
);
|
||||
|
||||
-- CreateIndex
|
||||
CREATE INDEX "LiteLLM_SpendLogToolIndex_tool_name_start_time_idx" ON "LiteLLM_SpendLogToolIndex"("tool_name", "start_time");
|
||||
|
|
@ -260,6 +260,7 @@ model LiteLLM_ObjectPermissionTable {
|
|||
vector_stores String[] @default([])
|
||||
agents String[] @default([])
|
||||
agent_access_groups String[] @default([])
|
||||
blocked_tools String[] @default([]) // Tool names blocked for any key/team/user with this permission
|
||||
teams LiteLLM_TeamTable[]
|
||||
projects LiteLLM_ProjectTable[]
|
||||
verification_tokens LiteLLM_VerificationToken[]
|
||||
|
|
@ -390,7 +391,7 @@ model LiteLLM_DeletedVerificationToken {
|
|||
config Json @default("{}")
|
||||
user_id String?
|
||||
team_id String?
|
||||
agent_id String?
|
||||
agent_id String?
|
||||
project_id String?
|
||||
permissions Json @default("{}")
|
||||
max_parallel_requests Int?
|
||||
|
|
@ -921,6 +922,16 @@ model LiteLLM_SpendLogGuardrailIndex {
|
|||
@@index([policy_id, start_time])
|
||||
}
|
||||
|
||||
// Index for fast "last N logs for tool" from SpendLogs – see how a tool is called in production
|
||||
model LiteLLM_SpendLogToolIndex {
|
||||
request_id String
|
||||
tool_name String // matches LiteLLM_ToolTable.tool_name; join for call_policy etc.
|
||||
start_time DateTime
|
||||
|
||||
@@id([request_id, tool_name])
|
||||
@@index([tool_name, start_time])
|
||||
}
|
||||
|
||||
// Prompt table for storing prompt configurations
|
||||
model LiteLLM_PromptTable {
|
||||
id String @id @default(uuid())
|
||||
|
|
@ -1078,6 +1089,7 @@ model LiteLLM_ToolTable {
|
|||
@@index([team_id])
|
||||
}
|
||||
|
||||
// Per-(tool, team/key) policy overrides. When present, override replaces global tool policy for that scope.
|
||||
//Unified Access Groups table for storing unified access groups
|
||||
model LiteLLM_AccessGroupTable {
|
||||
access_group_id String @id @default(uuid())
|
||||
|
|
|
|||
|
|
@ -75,7 +75,7 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
if messages is None:
|
||||
return data
|
||||
|
||||
chat_completion_compatible_request, tool_name_mapping = (
|
||||
chat_completion_compatible_request, _tool_name_mapping = (
|
||||
LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai(
|
||||
# Use a shallow copy to avoid mutating request data (pop on litellm_metadata).
|
||||
anthropic_message_request=cast(AnthropicMessagesRequest, data.copy())
|
||||
|
|
@ -141,6 +141,14 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
|
||||
return data
|
||||
|
||||
def extract_request_tool_names(self, data: dict) -> List[str]:
|
||||
"""Extract tool names from Anthropic messages request (tools[].name)."""
|
||||
names: List[str] = []
|
||||
for tool in data.get("tools") or []:
|
||||
if isinstance(tool, dict) and tool.get("name"):
|
||||
names.append(str(tool["name"]))
|
||||
return names
|
||||
|
||||
def _extract_input_text_and_images(
|
||||
self,
|
||||
message: Dict[str, Any],
|
||||
|
|
|
|||
|
|
@ -98,3 +98,10 @@ class BaseTranslation(ABC):
|
|||
Optional to override in subclasses.
|
||||
"""
|
||||
return responses_so_far
|
||||
|
||||
def extract_request_tool_names(self, data: dict) -> List[str]:
|
||||
"""
|
||||
Extract tool names from the request body for allowlist/policy checks.
|
||||
Override in tool-capable handlers; default returns [].
|
||||
"""
|
||||
return []
|
||||
|
|
|
|||
|
|
@ -135,6 +135,19 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
|
||||
return data
|
||||
|
||||
def extract_request_tool_names(self, data: dict) -> List[str]:
|
||||
"""Extract tool names from OpenAI chat completions request (tools[].function.name, functions[].name)."""
|
||||
names: List[str] = []
|
||||
for tool in data.get("tools") or []:
|
||||
if isinstance(tool, dict) and tool.get("type") == "function":
|
||||
fn = tool.get("function")
|
||||
if isinstance(fn, dict) and fn.get("name"):
|
||||
names.append(str(fn["name"]))
|
||||
for fn in data.get("functions") or []:
|
||||
if isinstance(fn, dict) and fn.get("name"):
|
||||
names.append(str(fn["name"]))
|
||||
return names
|
||||
|
||||
def _extract_inputs(
|
||||
self,
|
||||
message: Dict[str, Any],
|
||||
|
|
|
|||
|
|
@ -30,27 +30,22 @@ Output: response.output is List[GenericResponseOutputItem] where each has:
|
|||
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast
|
||||
|
||||
from openai.types.responses.response_function_tool_call import ResponseFunctionToolCall
|
||||
from openai.types.responses.response_function_tool_call import \
|
||||
ResponseFunctionToolCall
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.completion_extras.litellm_responses_transformation.transformation import (
|
||||
LiteLLMResponsesTransformationHandler,
|
||||
OpenAiResponsesToChatCompletionStreamIterator,
|
||||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
|
||||
from litellm.responses.litellm_completion_transformation.transformation import (
|
||||
LiteLLMCompletionResponsesConfig,
|
||||
)
|
||||
from litellm.types.llms.openai import (
|
||||
ChatCompletionToolCallChunk,
|
||||
ChatCompletionToolParam,
|
||||
)
|
||||
from litellm.types.responses.main import (
|
||||
GenericResponseOutputItem,
|
||||
OutputFunctionToolCall,
|
||||
OutputText,
|
||||
)
|
||||
OpenAiResponsesToChatCompletionStreamIterator)
|
||||
from litellm.llms.base_llm.guardrail_translation.base_translation import \
|
||||
BaseTranslation
|
||||
from litellm.responses.litellm_completion_transformation.transformation import \
|
||||
LiteLLMCompletionResponsesConfig
|
||||
from litellm.types.llms.openai import (ChatCompletionToolCallChunk,
|
||||
ChatCompletionToolParam)
|
||||
from litellm.types.responses.main import (GenericResponseOutputItem,
|
||||
OutputFunctionToolCall, OutputText)
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -188,6 +183,18 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
|
||||
return data
|
||||
|
||||
def extract_request_tool_names(self, data: dict) -> List[str]:
|
||||
"""Extract tool names from Responses API request (tools[].name for function, tools[].server_label for mcp)."""
|
||||
names: List[str] = []
|
||||
for tool in data.get("tools") or []:
|
||||
if not isinstance(tool, dict):
|
||||
continue
|
||||
if tool.get("type") == "function" and tool.get("name"):
|
||||
names.append(str(tool["name"]))
|
||||
elif tool.get("type") == "mcp" and tool.get("server_label"):
|
||||
names.append(str(tool["server_label"]))
|
||||
return names
|
||||
|
||||
def _extract_and_transform_tools(
|
||||
self,
|
||||
tools: List[Dict[str, Any]],
|
||||
|
|
|
|||
|
|
@ -23,33 +23,11 @@ model_list:
|
|||
|
||||
|
||||
guardrails:
|
||||
- guardrail_name: "airline-competitor-intent"
|
||||
guardrail_id: "airline-competitor-intent"
|
||||
- guardrail_name: "tool_policy"
|
||||
litellm_params:
|
||||
guardrail: litellm_content_filter
|
||||
mode: pre_call
|
||||
default_on: false
|
||||
competitor_intent_config:
|
||||
brand_self:
|
||||
- emirates
|
||||
- ek
|
||||
competitors:
|
||||
- qatar airways
|
||||
- qatar
|
||||
- etihad
|
||||
locations:
|
||||
- qatar
|
||||
- doha
|
||||
- doh
|
||||
competitor_aliases:
|
||||
qatar airways: [qr, doha airline]
|
||||
qatar: [qr]
|
||||
policy:
|
||||
competitor_comparison: refuse
|
||||
possible_competitor_comparison: reframe
|
||||
threshold_high: 0.70
|
||||
threshold_medium: 0.45
|
||||
threshold_low: 0.30
|
||||
guardrail: tool_policy
|
||||
mode: [pre_call, post_call]
|
||||
default_on: true
|
||||
|
||||
mcp_servers:
|
||||
my_http_server:
|
||||
|
|
|
|||
|
|
@ -77,6 +77,7 @@ class SupportedDBObjectType(str, enum.Enum):
|
|||
PASS_THROUGH_ENDPOINTS = "pass_through_endpoints"
|
||||
PROMPTS = "prompts"
|
||||
MODEL_COST_MAP = "model_cost_map"
|
||||
TOOLS = "tools"
|
||||
|
||||
def __str__(self):
|
||||
return str(self.value)
|
||||
|
|
@ -2126,7 +2127,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
|||
user_header_mappings: Optional[List[UserHeaderMapping]] = None
|
||||
supported_db_objects: Optional[List[SupportedDBObjectType]] = Field(
|
||||
None,
|
||||
description="Fine-grained control over which object types to load from the database when store_model_in_db is True. Available types: 'models', 'mcp', 'guardrails', 'vector_stores', 'pass_through_endpoints', 'prompts', 'model_cost_map'. If not set, all objects are loaded (default behavior).",
|
||||
description="Fine-grained control over which object types to load from the database when store_model_in_db is True. Available types: 'models', 'mcp', 'guardrails', 'vector_stores', 'pass_through_endpoints', 'prompts', 'model_cost_map', 'tools'. If not set, all objects are loaded (default behavior).",
|
||||
)
|
||||
user_mcp_management_mode: Optional[UserMCPManagementMode] = Field(
|
||||
None,
|
||||
|
|
@ -3370,6 +3371,11 @@ class ProxyErrorTypes(str, enum.Enum):
|
|||
Team member is already in team
|
||||
"""
|
||||
|
||||
tool_access_denied = "tool_access_denied"
|
||||
"""
|
||||
Tool is not in the allowed tools list for this key/team
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def get_model_access_error_type_for_object(
|
||||
cls, object_type: Literal["key", "user", "team", "org", "project"]
|
||||
|
|
|
|||
|
|
@ -58,6 +58,10 @@ from litellm.proxy._types import (
|
|||
)
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
from litellm.proxy.guardrails.tool_name_extraction import (
|
||||
TOOL_CAPABLE_CALL_TYPES,
|
||||
extract_request_tool_names,
|
||||
)
|
||||
from litellm.proxy.route_llm_request import route_request
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging, log_db_metrics
|
||||
from litellm.router import Router
|
||||
|
|
@ -220,7 +224,48 @@ async def _run_project_checks(
|
|||
)
|
||||
|
||||
|
||||
async def common_checks(
|
||||
async def check_tools_allowlist(
|
||||
request_body: dict,
|
||||
valid_token: Optional[UserAPIKeyAuth],
|
||||
team_object: Optional[LiteLLM_TeamTable],
|
||||
route: str,
|
||||
) -> None:
|
||||
"""
|
||||
Enforce key/team tool allowlist (metadata.allowed_tools). No DB in hot path —
|
||||
effective allowlist is read from valid_token.metadata and valid_token.team_metadata.
|
||||
Raises ProxyException with tool_access_denied if a tool is not allowed.
|
||||
"""
|
||||
from litellm.litellm_core_utils.api_route_to_call_types import (
|
||||
get_call_types_for_route,
|
||||
)
|
||||
|
||||
if valid_token is None:
|
||||
return
|
||||
call_types = get_call_types_for_route(route)
|
||||
if not call_types or not any(ct.value in TOOL_CAPABLE_CALL_TYPES for ct in call_types):
|
||||
return
|
||||
tool_names = extract_request_tool_names(route, request_body)
|
||||
if not tool_names:
|
||||
return
|
||||
key_meta = (valid_token.metadata or {}) if isinstance(valid_token.metadata, dict) else {}
|
||||
team_meta = (valid_token.team_metadata or {}) if isinstance(valid_token.team_metadata, dict) else {}
|
||||
key_allowed = key_meta.get("allowed_tools")
|
||||
team_allowed = team_meta.get("allowed_tools")
|
||||
effective = key_allowed if (isinstance(key_allowed, list) and len(key_allowed) > 0) else team_allowed
|
||||
if not isinstance(effective, list) or len(effective) == 0:
|
||||
return
|
||||
allowed_set = {str(t) for t in effective}
|
||||
disallowed = [n for n in tool_names if n not in allowed_set]
|
||||
if disallowed:
|
||||
raise ProxyException(
|
||||
message=f"Tool(s) {disallowed} are not in the allowed tools list for this key/team.",
|
||||
type=ProxyErrorTypes.tool_access_denied,
|
||||
param="tools",
|
||||
code=status.HTTP_403_FORBIDDEN,
|
||||
)
|
||||
|
||||
|
||||
async def common_checks( # noqa: PLR0915
|
||||
request_body: dict,
|
||||
team_object: Optional[LiteLLM_TeamTable],
|
||||
user_object: Optional[LiteLLM_UserTable],
|
||||
|
|
@ -473,6 +518,14 @@ async def common_checks(
|
|||
valid_token=valid_token,
|
||||
)
|
||||
|
||||
# 12. [OPTIONAL] Tool allowlist - key/team allowed_tools (no DB in hot path)
|
||||
await check_tools_allowlist(
|
||||
request_body=request_body,
|
||||
valid_token=valid_token,
|
||||
team_object=team_object,
|
||||
route=route,
|
||||
)
|
||||
|
||||
return True
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -13,49 +13,36 @@ 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
|
||||
from litellm.caching import DualCache, RedisCache
|
||||
from litellm.constants import DB_SPEND_UPDATE_JOB_NAME
|
||||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
from litellm.proxy._types import (
|
||||
DB_CONNECTION_ERROR_TYPES,
|
||||
BaseDailySpendTransaction,
|
||||
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,
|
||||
)
|
||||
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._types import (DB_CONNECTION_ERROR_TYPES,
|
||||
BaseDailySpendTransaction,
|
||||
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
|
||||
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:
|
||||
|
|
@ -104,12 +91,10 @@ class DBSpendUpdateWriter:
|
|||
end_time: Optional[datetime],
|
||||
response_cost: Optional[float],
|
||||
):
|
||||
from litellm.proxy.proxy_server import (
|
||||
disable_spend_logs,
|
||||
litellm_proxy_budget_name,
|
||||
prisma_client,
|
||||
user_api_key_cache,
|
||||
)
|
||||
from litellm.proxy.proxy_server import (disable_spend_logs,
|
||||
litellm_proxy_budget_name,
|
||||
prisma_client,
|
||||
user_api_key_cache)
|
||||
from litellm.proxy.utils import ProxyUpdateSpend, hash_token
|
||||
|
||||
try:
|
||||
|
|
@ -124,9 +109,8 @@ class DBSpendUpdateWriter:
|
|||
hashed_token = token
|
||||
|
||||
## CREATE SPEND LOG PAYLOAD ##
|
||||
from litellm.proxy.spend_tracking.spend_tracking_utils import (
|
||||
get_logging_payload,
|
||||
)
|
||||
from litellm.proxy.spend_tracking.spend_tracking_utils import \
|
||||
get_logging_payload
|
||||
|
||||
payload = get_logging_payload(
|
||||
kwargs=kwargs,
|
||||
|
|
@ -245,11 +229,13 @@ class DBSpendUpdateWriter:
|
|||
# --- 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")
|
||||
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")
|
||||
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")
|
||||
|
|
@ -280,7 +266,9 @@ class DBSpendUpdateWriter:
|
|||
_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"):
|
||||
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:
|
||||
|
|
@ -768,19 +756,46 @@ class DBSpendUpdateWriter:
|
|||
daily_end_user_spend_update_transactions,
|
||||
daily_agent_spend_update_transactions,
|
||||
daily_tag_spend_update_transactions,
|
||||
) = await self.redis_update_buffer.get_all_transactions_from_redis_buffer_pipeline()
|
||||
) = (
|
||||
await self.redis_update_buffer.get_all_transactions_from_redis_buffer_pipeline()
|
||||
)
|
||||
|
||||
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 {}),
|
||||
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,
|
||||
|
|
@ -985,10 +1000,8 @@ class DBSpendUpdateWriter:
|
|||
Commits all the spend `UPDATE` transactions to the Database
|
||||
|
||||
"""
|
||||
from litellm.proxy.utils import (
|
||||
ProxyUpdateSpend,
|
||||
_raise_failed_update_spend_exception,
|
||||
)
|
||||
from litellm.proxy.utils import (ProxyUpdateSpend,
|
||||
_raise_failed_update_spend_exception)
|
||||
|
||||
### UPDATE USER TABLE ###
|
||||
user_list_transactions = db_spend_update_transactions["user_list_transactions"]
|
||||
|
|
@ -1523,14 +1536,14 @@ class DBSpendUpdateWriter:
|
|||
|
||||
# Add cache-related fields if they exist
|
||||
if "cache_read_input_tokens" in transaction:
|
||||
common_data[
|
||||
"cache_read_input_tokens"
|
||||
] = transaction.get("cache_read_input_tokens", 0)
|
||||
common_data["cache_read_input_tokens"] = (
|
||||
transaction.get("cache_read_input_tokens", 0)
|
||||
)
|
||||
if "cache_creation_input_tokens" in transaction:
|
||||
common_data[
|
||||
"cache_creation_input_tokens"
|
||||
] = transaction.get(
|
||||
"cache_creation_input_tokens", 0
|
||||
common_data["cache_creation_input_tokens"] = (
|
||||
transaction.get(
|
||||
"cache_creation_input_tokens", 0
|
||||
)
|
||||
)
|
||||
|
||||
if entity_type == "tag" and "request_id" in transaction:
|
||||
|
|
|
|||
147
litellm/proxy/db/spend_log_tool_index.py
Normal file
147
litellm/proxy/db/spend_log_tool_index.py
Normal file
|
|
@ -0,0 +1,147 @@
|
|||
"""
|
||||
Track tool usage for the dashboard: insert into SpendLogToolIndex when spend logs
|
||||
are written, so "last N requests for tool X" and "how is this tool called in production"
|
||||
queries are fast.
|
||||
"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Dict, List, Set
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
|
||||
def _add_tool_calls_to_set(tool_calls: Any, out: Set[str]) -> None:
|
||||
"""Extract tool names from OpenAI-style tool_calls list into out."""
|
||||
if not isinstance(tool_calls, list):
|
||||
return
|
||||
for tc in tool_calls:
|
||||
if not isinstance(tc, dict):
|
||||
continue
|
||||
fn = tc.get("function")
|
||||
if isinstance(fn, dict):
|
||||
name = fn.get("name")
|
||||
if name and isinstance(name, str) and name.strip():
|
||||
out.add(name.strip())
|
||||
|
||||
|
||||
def _parse_tool_names_from_payload(payload: Dict[str, Any]) -> Set[str]:
|
||||
"""
|
||||
Extract deduplicated tool names from a spend log payload.
|
||||
Sources: mcp_namespaced_tool_name, response (tool_calls), proxy_server_request (tools).
|
||||
"""
|
||||
tool_names: Set[str] = set()
|
||||
|
||||
# Top-level MCP tool name (single tool per request for that flow)
|
||||
mcp_name = payload.get("mcp_namespaced_tool_name")
|
||||
if mcp_name and isinstance(mcp_name, str) and mcp_name.strip():
|
||||
tool_names.add(mcp_name.strip())
|
||||
|
||||
# Response: OpenAI-style tool_calls[].function.name or choices[0].message.tool_calls
|
||||
response_raw = payload.get("response")
|
||||
if response_raw:
|
||||
response_obj = (
|
||||
safe_json_loads(response_raw, default=None)
|
||||
if isinstance(response_raw, str)
|
||||
else response_raw
|
||||
)
|
||||
if isinstance(response_obj, dict):
|
||||
_add_tool_calls_to_set(response_obj.get("tool_calls"), tool_names)
|
||||
choices = response_obj.get("choices")
|
||||
if isinstance(choices, list) and choices:
|
||||
msg = choices[0].get("message") if isinstance(choices[0], dict) else None
|
||||
if isinstance(msg, dict):
|
||||
_add_tool_calls_to_set(msg.get("tool_calls"), tool_names)
|
||||
|
||||
# Request body: tools[].function.name
|
||||
request_raw = payload.get("proxy_server_request")
|
||||
if request_raw:
|
||||
request_obj = (
|
||||
safe_json_loads(request_raw, default=None)
|
||||
if isinstance(request_raw, str)
|
||||
else request_raw
|
||||
)
|
||||
if isinstance(request_obj, dict):
|
||||
body = request_obj.get("body", request_obj)
|
||||
if isinstance(body, dict):
|
||||
request_obj = body
|
||||
if isinstance(request_obj, dict):
|
||||
tools = request_obj.get("tools")
|
||||
if isinstance(tools, list):
|
||||
for t in tools:
|
||||
if isinstance(t, dict):
|
||||
fn = t.get("function")
|
||||
if isinstance(fn, dict):
|
||||
name = fn.get("name")
|
||||
if name and isinstance(name, str) and name.strip():
|
||||
tool_names.add(name.strip())
|
||||
|
||||
return tool_names
|
||||
|
||||
|
||||
async def process_spend_logs_tool_usage(
|
||||
prisma_client: PrismaClient,
|
||||
logs_to_process: List[Dict[str, Any]],
|
||||
) -> None:
|
||||
"""
|
||||
After spend logs are written: insert SpendLogToolIndex rows from each payload.
|
||||
Extracts tool names from mcp_namespaced_tool_name, response tool_calls, and
|
||||
proxy_server_request tools.
|
||||
"""
|
||||
if not logs_to_process:
|
||||
return
|
||||
|
||||
index_rows: List[Dict[str, Any]] = []
|
||||
|
||||
for payload in logs_to_process:
|
||||
request_id = payload.get("request_id")
|
||||
start_time = payload.get("startTime")
|
||||
if not request_id or not start_time:
|
||||
continue
|
||||
if isinstance(start_time, str):
|
||||
try:
|
||||
start_time = datetime.fromisoformat(
|
||||
start_time.replace("Z", "+00:00")
|
||||
)
|
||||
except (ValueError, TypeError):
|
||||
continue
|
||||
if start_time.tzinfo is None:
|
||||
start_time = start_time.replace(tzinfo=timezone.utc)
|
||||
|
||||
tool_names = _parse_tool_names_from_payload(payload)
|
||||
for tool_name in tool_names:
|
||||
index_rows.append({
|
||||
"request_id": request_id,
|
||||
"tool_name": tool_name,
|
||||
"start_time": start_time,
|
||||
})
|
||||
|
||||
if not index_rows:
|
||||
return
|
||||
|
||||
try:
|
||||
index_data = []
|
||||
for r in index_rows:
|
||||
st = r["start_time"]
|
||||
if isinstance(st, str):
|
||||
try:
|
||||
st = datetime.fromisoformat(st.replace("Z", "+00:00"))
|
||||
except (ValueError, TypeError):
|
||||
continue
|
||||
if st.tzinfo is None:
|
||||
st = st.replace(tzinfo=timezone.utc)
|
||||
index_data.append({
|
||||
"request_id": r["request_id"],
|
||||
"tool_name": r["tool_name"],
|
||||
"start_time": st,
|
||||
})
|
||||
if index_data:
|
||||
await prisma_client.db.litellm_spendlogtoolindex.create_many(
|
||||
data=index_data,
|
||||
skip_duplicates=True,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning(
|
||||
"Tool usage tracking (SpendLogToolIndex) failed (non-fatal): %s", e
|
||||
)
|
||||
|
|
@ -3,25 +3,48 @@ 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 typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import ToolDiscoveryQueueItem
|
||||
from litellm.types.tool_management import LiteLLM_ToolTableRow, ToolCallPolicy
|
||||
from litellm.types.tool_management import (
|
||||
LiteLLM_ToolTableRow,
|
||||
ToolCallPolicy,
|
||||
ToolPolicyOverrideRow,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
|
||||
def _row_to_model(row: dict) -> LiteLLM_ToolTableRow:
|
||||
def _row_to_model(row: Union[dict, Any]) -> LiteLLM_ToolTableRow:
|
||||
"""Convert a Prisma model instance or dict to LiteLLM_ToolTableRow."""
|
||||
model_dump = getattr(row, "model_dump", None)
|
||||
if callable(model_dump):
|
||||
row = model_dump()
|
||||
elif not isinstance(row, dict):
|
||||
row = {
|
||||
k: getattr(row, k, None)
|
||||
for k in (
|
||||
"tool_id",
|
||||
"tool_name",
|
||||
"origin",
|
||||
"call_policy",
|
||||
"call_count",
|
||||
"assignments",
|
||||
"key_hash",
|
||||
"team_id",
|
||||
"key_alias",
|
||||
"created_at",
|
||||
"updated_at",
|
||||
"created_by",
|
||||
"updated_by",
|
||||
)
|
||||
}
|
||||
return LiteLLM_ToolTableRow(
|
||||
tool_id=row.get("tool_id", ""),
|
||||
tool_name=row.get("tool_name", ""),
|
||||
|
|
@ -44,7 +67,7 @@ async def batch_upsert_tools(
|
|||
items: List[ToolDiscoveryQueueItem],
|
||||
) -> None:
|
||||
"""
|
||||
Batch-upsert tool registry rows via raw SQL.
|
||||
Batch-upsert tool registry rows via Prisma.
|
||||
|
||||
On first insert: sets call_policy = "untrusted" (schema default), call_count = 1.
|
||||
On conflict: increments call_count; preserves existing call_policy.
|
||||
|
|
@ -83,7 +106,9 @@ async def batch_upsert_tools(
|
|||
"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)
|
||||
verbose_proxy_logger.error(
|
||||
"tool_registry_writer batch_upsert_tools error: %s", e
|
||||
)
|
||||
|
||||
|
||||
async def _get_agent_ids_for_key_hashes(
|
||||
|
|
@ -146,15 +171,12 @@ async def get_tool(
|
|||
) -> Optional[LiteLLM_ToolTableRow]:
|
||||
"""Return a single tool row by tool_name. Enriches with agent_id from key table if key_hash is set."""
|
||||
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,
|
||||
row = await prisma_client.db.litellm_tooltable.find_unique(
|
||||
where={"tool_name": tool_name},
|
||||
)
|
||||
if not rows:
|
||||
if row is None:
|
||||
return None
|
||||
tool = _row_to_model(rows[0])
|
||||
tool = _row_to_model(row)
|
||||
if tool.key_hash:
|
||||
key_to_agent = await _get_agent_ids_for_key_hashes(
|
||||
prisma_client, [tool.key_hash]
|
||||
|
|
@ -188,7 +210,9 @@ async def update_tool_policy(
|
|||
)
|
||||
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)
|
||||
verbose_proxy_logger.error(
|
||||
"tool_registry_writer update_tool_policy error: %s", e
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
|
|
@ -203,12 +227,207 @@ async def get_tools_by_names(
|
|||
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,
|
||||
rows = await prisma_client.db.litellm_tooltable.find_many(
|
||||
where={"tool_name": {"in": tool_names}},
|
||||
)
|
||||
return {row["tool_name"]: row["call_policy"] for row in rows}
|
||||
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)
|
||||
verbose_proxy_logger.error(
|
||||
"tool_registry_writer get_tools_by_names error: %s", e
|
||||
)
|
||||
return {}
|
||||
|
||||
|
||||
async def list_overrides_for_tool(
|
||||
prisma_client: "PrismaClient",
|
||||
tool_name: str,
|
||||
) -> List[ToolPolicyOverrideRow]:
|
||||
"""
|
||||
Return override-like rows for a tool by finding object permissions that have
|
||||
this tool in blocked_tools, then resolving each permission to key/team scope for display.
|
||||
"""
|
||||
out: List[ToolPolicyOverrideRow] = []
|
||||
try:
|
||||
perms = await prisma_client.db.litellm_objectpermissiontable.find_many(
|
||||
where={"blocked_tools": {"has": tool_name}},
|
||||
include={
|
||||
"verification_tokens": True,
|
||||
"teams": True,
|
||||
},
|
||||
)
|
||||
for perm in perms:
|
||||
op_id = getattr(perm, "object_permission_id", None) or ""
|
||||
tokens = getattr(perm, "verification_tokens", []) or []
|
||||
teams = getattr(perm, "teams", []) or []
|
||||
for t in tokens:
|
||||
out.append(
|
||||
ToolPolicyOverrideRow(
|
||||
override_id=op_id,
|
||||
tool_name=tool_name,
|
||||
team_id=None,
|
||||
key_hash=getattr(t, "token", None),
|
||||
call_policy="blocked",
|
||||
key_alias=getattr(t, "key_alias", None),
|
||||
created_at=None,
|
||||
updated_at=None,
|
||||
)
|
||||
)
|
||||
for team in teams:
|
||||
out.append(
|
||||
ToolPolicyOverrideRow(
|
||||
override_id=op_id,
|
||||
tool_name=tool_name,
|
||||
team_id=getattr(team, "team_id", None),
|
||||
key_hash=None,
|
||||
call_policy="blocked",
|
||||
key_alias=getattr(team, "team_alias", None),
|
||||
created_at=None,
|
||||
updated_at=None,
|
||||
)
|
||||
)
|
||||
return out
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(
|
||||
"tool_registry_writer list_overrides_for_tool error: %s", e
|
||||
)
|
||||
return []
|
||||
|
||||
|
||||
class ToolPolicyRegistry:
|
||||
"""
|
||||
In-memory registry of tool policies synced from DB.
|
||||
Synced in _init_tool_policy_in_db (from add_deployment / _init_non_llm_objects_in_db).
|
||||
Hot path uses get_effective_policies only — no DB, no cache.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._global_tool_policies: Dict[str, str] = {}
|
||||
self._blocked_tools_by_op_id: Dict[str, List[str]] = {}
|
||||
self._initialized: bool = False
|
||||
|
||||
def is_initialized(self) -> bool:
|
||||
return self._initialized
|
||||
|
||||
async def sync_tool_policy_from_db(self, prisma_client: "PrismaClient") -> None:
|
||||
"""Load all tool policies and object-permission blocked_tools from DB; replace in-memory state."""
|
||||
try:
|
||||
tools = await prisma_client.db.litellm_tooltable.find_many()
|
||||
self._global_tool_policies = {row.tool_name: row.call_policy for row in tools}
|
||||
|
||||
perms = await prisma_client.db.litellm_objectpermissiontable.find_many()
|
||||
self._blocked_tools_by_op_id = {}
|
||||
for row in perms:
|
||||
op_id = getattr(row, "object_permission_id", None)
|
||||
blocked = getattr(row, "blocked_tools", None) or []
|
||||
if op_id:
|
||||
self._blocked_tools_by_op_id[op_id] = list(blocked)
|
||||
|
||||
self._initialized = True
|
||||
verbose_proxy_logger.info(
|
||||
"ToolPolicyRegistry: synced %d global tool policies and %d object permissions from DB",
|
||||
len(self._global_tool_policies),
|
||||
len(self._blocked_tools_by_op_id),
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"ToolPolicyRegistry sync_tool_policy_from_db error: %s", e
|
||||
)
|
||||
raise
|
||||
|
||||
def get_effective_policies(
|
||||
self,
|
||||
tool_names: List[str],
|
||||
object_permission_id: Optional[str] = None,
|
||||
team_object_permission_id: Optional[str] = None,
|
||||
) -> Dict[str, str]:
|
||||
"""
|
||||
Return effective call_policy per tool from in-memory state.
|
||||
If tool is in key or team blocked_tools -> "blocked", else global policy or "untrusted".
|
||||
"""
|
||||
if not tool_names:
|
||||
return {}
|
||||
blocked: set = set()
|
||||
for op_id in (object_permission_id, team_object_permission_id):
|
||||
if op_id and op_id.strip():
|
||||
blocked.update(
|
||||
self._blocked_tools_by_op_id.get(op_id.strip(), [])
|
||||
)
|
||||
result: Dict[str, str] = {}
|
||||
for name in tool_names:
|
||||
if name in blocked:
|
||||
result[name] = "blocked"
|
||||
else:
|
||||
result[name] = self._global_tool_policies.get(name, "untrusted")
|
||||
return result
|
||||
|
||||
|
||||
_tool_policy_registry: Optional[ToolPolicyRegistry] = None
|
||||
|
||||
|
||||
def get_tool_policy_registry() -> ToolPolicyRegistry:
|
||||
"""Return the global ToolPolicyRegistry singleton."""
|
||||
global _tool_policy_registry
|
||||
if _tool_policy_registry is None:
|
||||
_tool_policy_registry = ToolPolicyRegistry()
|
||||
return _tool_policy_registry
|
||||
|
||||
|
||||
async def add_tool_to_object_permission_blocked(
|
||||
prisma_client: "PrismaClient",
|
||||
object_permission_id: str,
|
||||
tool_name: str,
|
||||
) -> bool:
|
||||
"""Add tool_name to the permission's blocked_tools if not already present."""
|
||||
if not object_permission_id or not tool_name:
|
||||
return False
|
||||
try:
|
||||
row = await prisma_client.db.litellm_objectpermissiontable.find_unique(
|
||||
where={"object_permission_id": object_permission_id},
|
||||
)
|
||||
if row is None:
|
||||
return False
|
||||
current = list(getattr(row, "blocked_tools", []) or [])
|
||||
if tool_name in current:
|
||||
return True
|
||||
current.append(tool_name)
|
||||
await prisma_client.db.litellm_objectpermissiontable.update(
|
||||
where={"object_permission_id": object_permission_id},
|
||||
data={"blocked_tools": current},
|
||||
)
|
||||
return True
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(
|
||||
"tool_registry_writer add_tool_to_object_permission_blocked error: %s", e
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
async def remove_tool_from_object_permission_blocked(
|
||||
prisma_client: "PrismaClient",
|
||||
object_permission_id: str,
|
||||
tool_name: str,
|
||||
) -> bool:
|
||||
"""Remove tool_name from the permission's blocked_tools. Returns False if tool was not in list."""
|
||||
if not object_permission_id or not tool_name:
|
||||
return False
|
||||
try:
|
||||
row = await prisma_client.db.litellm_objectpermissiontable.find_unique(
|
||||
where={"object_permission_id": object_permission_id},
|
||||
)
|
||||
if row is None:
|
||||
return False
|
||||
current = list(getattr(row, "blocked_tools", []) or [])
|
||||
if tool_name not in current:
|
||||
return False
|
||||
current = [t for t in current if t != tool_name]
|
||||
await prisma_client.db.litellm_objectpermissiontable.update(
|
||||
where={"object_permission_id": object_permission_id},
|
||||
data={"blocked_tools": current},
|
||||
)
|
||||
return True
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(
|
||||
"tool_registry_writer remove_tool_from_object_permission_blocked error: %s",
|
||||
e,
|
||||
)
|
||||
return False
|
||||
|
|
|
|||
|
|
@ -23,17 +23,16 @@ or both pre and post call:
|
|||
mode: during_call # runs before LLM and on response
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional
|
||||
from typing import TYPE_CHECKING, Any, Literal, Optional, Tuple
|
||||
|
||||
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.proxy.guardrails.tool_name_extraction import extract_request_tool_names
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
|
|
@ -43,12 +42,50 @@ if TYPE_CHECKING:
|
|||
GUARDRAIL_NAME = "tool_policy"
|
||||
|
||||
|
||||
def _get_request_object_permission_ids(
|
||||
request_data: dict,
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""Extract object_permission_id and team_object_permission_id from request_data."""
|
||||
if not request_data:
|
||||
return None, None
|
||||
for key in ("litellm_metadata", "metadata"):
|
||||
meta = request_data.get(key)
|
||||
if not isinstance(meta, dict):
|
||||
continue
|
||||
auth = meta.get("user_api_key_auth")
|
||||
if auth is not None and hasattr(auth, "object_permission_id"):
|
||||
key_op = getattr(auth, "object_permission_id", None)
|
||||
team_op = getattr(auth, "team_object_permission_id", None)
|
||||
if key_op is not None or team_op is not None:
|
||||
return (
|
||||
str(key_op).strip() if key_op else None,
|
||||
str(team_op).strip() if team_op else None,
|
||||
)
|
||||
key_op = meta.get("user_api_key_object_permission_id")
|
||||
team_op = meta.get("user_api_key_team_object_permission_id")
|
||||
if key_op is not None or team_op is not None:
|
||||
return (
|
||||
str(key_op).strip() if key_op else None,
|
||||
str(team_op).strip() if team_op else None,
|
||||
)
|
||||
return None, None
|
||||
|
||||
|
||||
def _get_request_route_from_data(request_data: dict) -> Optional[str]:
|
||||
"""Get request route from request_data (metadata or top-level)."""
|
||||
route = request_data.get("user_api_key_request_route")
|
||||
if route:
|
||||
return route
|
||||
meta = request_data.get("metadata") or request_data.get("litellm_metadata") or {}
|
||||
return meta.get("user_api_key_request_route")
|
||||
|
||||
|
||||
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.
|
||||
Guardrail that enforces per-tool call policies from the in-memory
|
||||
ToolPolicyRegistry (synced from DB). Key/team allowed_tools (allowlist) is
|
||||
enforced in the auth layer (check_tools_allowlist). No DB or cache in hot
|
||||
path — registry lookups only.
|
||||
"""
|
||||
|
||||
def __init__(self, **kwargs: Any) -> None:
|
||||
|
|
@ -59,7 +96,6 @@ class ToolPolicyGuardrail(CustomGuardrail):
|
|||
GuardrailEventHooks.during_call,
|
||||
]
|
||||
super().__init__(**kwargs)
|
||||
self._policy_cache: DualCache = DualCache()
|
||||
|
||||
@log_guardrail_information
|
||||
async def apply_guardrail(
|
||||
|
|
@ -70,12 +106,7 @@ class ToolPolicyGuardrail(CustomGuardrail):
|
|||
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".
|
||||
Enforce DB call_policy on request tools / response tool_calls.
|
||||
"""
|
||||
if input_type == "request":
|
||||
tools = inputs.get("tools") or []
|
||||
|
|
@ -86,6 +117,11 @@ class ToolPolicyGuardrail(CustomGuardrail):
|
|||
and isinstance(t.get("function"), dict)
|
||||
and t["function"].get("name")
|
||||
]
|
||||
if not tool_names:
|
||||
route = _get_request_route_from_data(request_data)
|
||||
if route:
|
||||
|
||||
tool_names = extract_request_tool_names(route, request_data)
|
||||
else: # response
|
||||
tool_calls = inputs.get("tool_calls") or []
|
||||
tool_names = []
|
||||
|
|
@ -101,8 +137,20 @@ class ToolPolicyGuardrail(CustomGuardrail):
|
|||
if not tool_names:
|
||||
return inputs
|
||||
|
||||
policy_map = await self._get_policies_cached(tool_names)
|
||||
object_permission_id, team_object_permission_id = (
|
||||
_get_request_object_permission_ids(request_data)
|
||||
)
|
||||
from litellm.proxy.db.tool_registry_writer import get_tool_policy_registry
|
||||
|
||||
registry = get_tool_policy_registry()
|
||||
if not registry.is_initialized():
|
||||
policy_map = {}
|
||||
else:
|
||||
policy_map = registry.get_effective_policies(
|
||||
tool_names,
|
||||
object_permission_id=object_permission_id,
|
||||
team_object_permission_id=team_object_permission_id,
|
||||
)
|
||||
blocked = [name for name in tool_names if policy_map.get(name) == "blocked"]
|
||||
if blocked:
|
||||
verbose_proxy_logger.warning(
|
||||
|
|
@ -118,46 +166,3 @@ class ToolPolicyGuardrail(CustomGuardrail):
|
|||
)
|
||||
|
||||
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
|
||||
|
|
|
|||
85
litellm/proxy/guardrails/tool_name_extraction.py
Normal file
85
litellm/proxy/guardrails/tool_name_extraction.py
Normal file
|
|
@ -0,0 +1,85 @@
|
|||
"""
|
||||
Extract tool names from request body by route/call type.
|
||||
|
||||
Used by auth (check_tools_allowlist) and ToolPolicyGuardrail so tool-format
|
||||
knowledge lives in one place. Uses guardrail translation handlers where available,
|
||||
with standalone extractors for generate_content and MCP.
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, List
|
||||
|
||||
from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route
|
||||
from litellm.llms import load_guardrail_translation_mappings
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
# Call types that have no guardrail translation handler; we use standalone extractors
|
||||
STANDALONE_EXTRACTORS: Dict[str, Any] = {}
|
||||
|
||||
|
||||
def _extract_generate_content_tool_names(data: dict) -> List[str]:
|
||||
"""Google generateContent: tools[].functionDeclarations[].name"""
|
||||
names: List[str] = []
|
||||
for tool in data.get("tools") or []:
|
||||
if not isinstance(tool, dict):
|
||||
continue
|
||||
for decl in tool.get("functionDeclarations") or []:
|
||||
if isinstance(decl, dict) and decl.get("name"):
|
||||
names.append(str(decl["name"]))
|
||||
return names
|
||||
|
||||
|
||||
def _extract_mcp_tool_names(data: dict) -> List[str]:
|
||||
"""MCP call_tool: name or mcp_tool_name in body"""
|
||||
names: List[str] = []
|
||||
name = data.get("name") or data.get("mcp_tool_name")
|
||||
if name:
|
||||
names.append(str(name))
|
||||
return names
|
||||
|
||||
|
||||
def _register_standalone_extractors() -> None:
|
||||
if STANDALONE_EXTRACTORS:
|
||||
return
|
||||
STANDALONE_EXTRACTORS[CallTypes.generate_content.value] = _extract_generate_content_tool_names
|
||||
STANDALONE_EXTRACTORS[CallTypes.agenerate_content.value] = _extract_generate_content_tool_names
|
||||
STANDALONE_EXTRACTORS[CallTypes.call_mcp_tool.value] = _extract_mcp_tool_names
|
||||
|
||||
|
||||
# Tool-capable call types (routes that can send tools in the request)
|
||||
TOOL_CAPABLE_CALL_TYPES = frozenset({
|
||||
CallTypes.completion.value,
|
||||
CallTypes.acompletion.value,
|
||||
CallTypes.responses.value,
|
||||
CallTypes.aresponses.value,
|
||||
CallTypes.anthropic_messages.value,
|
||||
CallTypes.generate_content.value,
|
||||
CallTypes.agenerate_content.value,
|
||||
CallTypes.call_mcp_tool.value,
|
||||
})
|
||||
|
||||
|
||||
def extract_request_tool_names(route: str, data: dict) -> List[str]:
|
||||
"""
|
||||
Extract tool names from the request body for the given route.
|
||||
Uses guardrail translation handlers when available, else standalone extractors
|
||||
for generate_content and MCP. Returns [] for non-tool-capable routes or when
|
||||
no tools are present.
|
||||
"""
|
||||
call_types = get_call_types_for_route(route)
|
||||
if not call_types:
|
||||
return []
|
||||
_register_standalone_extractors()
|
||||
mappings = load_guardrail_translation_mappings()
|
||||
for call_type in call_types:
|
||||
if not isinstance(call_type, CallTypes):
|
||||
continue
|
||||
if call_type.value not in TOOL_CAPABLE_CALL_TYPES:
|
||||
continue
|
||||
if call_type.value in STANDALONE_EXTRACTORS:
|
||||
return STANDALONE_EXTRACTORS[call_type.value](data)
|
||||
handler_cls = mappings.get(call_type)
|
||||
if handler_cls is not None:
|
||||
names = handler_cls().extract_request_tool_names(data)
|
||||
if names:
|
||||
return names
|
||||
return []
|
||||
|
|
@ -10,12 +10,16 @@ import litellm
|
|||
from litellm._logging import verbose_logger, verbose_proxy_logger
|
||||
from litellm._service_logger import ServiceLogging
|
||||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
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
|
||||
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(
|
||||
|
|
@ -24,9 +28,12 @@ _SPECIAL_HEADERS_CACHE = frozenset(
|
|||
from litellm.router import Router
|
||||
from litellm.types.llms.anthropic import ANTHROPIC_API_HEADERS
|
||||
from litellm.types.services import ServiceTypes
|
||||
from litellm.types.utils import (LlmProviders, ProviderSpecificHeader,
|
||||
StandardLoggingUserAPIKeyMetadata,
|
||||
SupportedCacheControls)
|
||||
from litellm.types.utils import (
|
||||
LlmProviders,
|
||||
ProviderSpecificHeader,
|
||||
StandardLoggingUserAPIKeyMetadata,
|
||||
SupportedCacheControls,
|
||||
)
|
||||
|
||||
service_logger_obj = ServiceLogging() # used for tracking latency on OTEL
|
||||
|
||||
|
|
@ -682,7 +689,8 @@ class LiteLLMProxyRequestSetup:
|
|||
return data
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_ManagementEndpoint_MetadataFields,
|
||||
LiteLLM_ManagementEndpoint_MetadataFields_Premium)
|
||||
LiteLLM_ManagementEndpoint_MetadataFields_Premium,
|
||||
)
|
||||
|
||||
# ignore any special fields
|
||||
added_metadata = {}
|
||||
|
|
@ -851,8 +859,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 (RedactedDict,
|
||||
SecretFields)
|
||||
from litellm.types.proxy.litellm_pre_call_utils import RedactedDict, SecretFields
|
||||
|
||||
_raw_headers: Dict[str, str] = RedactedDict(_safe_get_request_headers(request))
|
||||
|
||||
|
|
@ -1084,6 +1091,15 @@ async def add_litellm_data_to_request( # noqa: PLR0915
|
|||
] = user_api_key_dict.user_max_budget
|
||||
|
||||
data[_metadata_variable_name]["user_api_key_metadata"] = user_api_key_dict.metadata
|
||||
data[_metadata_variable_name]["user_api_key_team_metadata"] = (
|
||||
user_api_key_dict.team_metadata
|
||||
)
|
||||
data[_metadata_variable_name]["user_api_key_object_permission_id"] = (
|
||||
getattr(user_api_key_dict, "object_permission_id", None)
|
||||
)
|
||||
data[_metadata_variable_name]["user_api_key_team_object_permission_id"] = (
|
||||
getattr(user_api_key_dict, "team_object_permission_id", None)
|
||||
)
|
||||
data[_metadata_variable_name]["headers"] = _headers
|
||||
data[_metadata_variable_name]["endpoint"] = str(request.url)
|
||||
|
||||
|
|
@ -1527,8 +1543,7 @@ async def move_guardrails_to_metadata(
|
|||
|
||||
# Only check policy engine if no local config (avoid import + registry lookup)
|
||||
if not (has_key_config or has_team_config or has_request_config):
|
||||
from litellm.proxy.policy_engine.policy_registry import \
|
||||
get_policy_registry
|
||||
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
|
||||
|
||||
if not get_policy_registry().is_initialized():
|
||||
# Nothing configured anywhere - clean up request body fields and return
|
||||
|
|
@ -1592,16 +1607,14 @@ async def move_guardrails_to_metadata(
|
|||
|
||||
def _is_policy_version_id(s: str) -> bool:
|
||||
"""Return True if string is a policy version ID (starts with policy_<uuid> prefix)."""
|
||||
from litellm.proxy.policy_engine.policy_registry import \
|
||||
POLICY_VERSION_ID_PREFIX
|
||||
from litellm.proxy.policy_engine.policy_registry import POLICY_VERSION_ID_PREFIX
|
||||
|
||||
return isinstance(s, str) and s.startswith(POLICY_VERSION_ID_PREFIX)
|
||||
|
||||
|
||||
def _extract_policy_id(s: str) -> Optional[str]:
|
||||
"""Extract raw UUID from policy_<uuid> string, or None if not a valid version ID."""
|
||||
from litellm.proxy.policy_engine.policy_registry import \
|
||||
POLICY_VERSION_ID_PREFIX
|
||||
from litellm.proxy.policy_engine.policy_registry import POLICY_VERSION_ID_PREFIX
|
||||
|
||||
if not _is_policy_version_id(s):
|
||||
return None
|
||||
|
|
@ -1622,9 +1635,10 @@ def _match_and_track_policies(
|
|||
"""
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
add_policy_sources_to_metadata, add_policy_to_applied_policies_header)
|
||||
from litellm.proxy.policy_engine.attachment_registry import \
|
||||
get_attachment_registry
|
||||
add_policy_sources_to_metadata,
|
||||
add_policy_to_applied_policies_header,
|
||||
)
|
||||
from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry
|
||||
from litellm.proxy.policy_engine.policy_matcher import PolicyMatcher
|
||||
|
||||
# Get matching policies via attachments (with match reasons for attribution)
|
||||
|
|
@ -1769,8 +1783,7 @@ async def add_guardrails_from_policy_engine(
|
|||
user_api_key_dict: The user's API key authentication info
|
||||
"""
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy.common_utils.http_parsing_utils import \
|
||||
get_tags_from_request_body
|
||||
from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body
|
||||
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
|
||||
from litellm.types.proxy.policy_engine import PolicyMatchContext
|
||||
|
||||
|
|
|
|||
|
|
@ -8,9 +8,14 @@ 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
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING, Any, List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth
|
||||
|
|
@ -18,9 +23,12 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
|||
from litellm.types.tool_management import (
|
||||
LiteLLM_ToolTableRow,
|
||||
ToolCallPolicy,
|
||||
ToolDetailResponse,
|
||||
ToolListResponse,
|
||||
ToolPolicyUpdateRequest,
|
||||
ToolPolicyUpdateResponse,
|
||||
ToolUsageLogEntry,
|
||||
ToolUsageLogsResponse,
|
||||
)
|
||||
|
||||
router = APIRouter()
|
||||
|
|
@ -51,13 +59,204 @@ async def list_tools(
|
|||
)
|
||||
|
||||
try:
|
||||
tools = await db_list_tools(prisma_client=prisma_client, call_policy=call_policy)
|
||||
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}/detail",
|
||||
tags=["tool management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=ToolDetailResponse,
|
||||
)
|
||||
async def get_tool_detail(
|
||||
tool_name: str,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Get a single tool with its policy overrides (for UI detail view).
|
||||
|
||||
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.db.tool_registry_writer import list_overrides_for_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")
|
||||
overrides = await list_overrides_for_tool(
|
||||
prisma_client=prisma_client, tool_name=tool_name
|
||||
)
|
||||
return ToolDetailResponse(tool=tool, overrides=overrides)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error getting tool detail: %s", e)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
def _input_snippet_for_tool_log(sl: Any, max_len: int = 200) -> Optional[str]:
|
||||
"""Short snippet from messages or proxy_server_request for tool usage log row."""
|
||||
if sl is None:
|
||||
return None
|
||||
messages = getattr(sl, "messages", None)
|
||||
if messages is not None:
|
||||
s = _snippet_str(messages, max_len)
|
||||
if s:
|
||||
return s
|
||||
psr = getattr(sl, "proxy_server_request", None)
|
||||
if not psr:
|
||||
return None
|
||||
if isinstance(psr, str):
|
||||
import json
|
||||
|
||||
try:
|
||||
psr = json.loads(psr)
|
||||
except Exception:
|
||||
return _snippet_str(psr, max_len)
|
||||
if isinstance(psr, dict):
|
||||
msgs = psr.get("messages")
|
||||
if msgs is None and isinstance(psr.get("body"), dict):
|
||||
msgs = psr["body"].get("messages")
|
||||
s = _snippet_str(msgs, max_len)
|
||||
if s:
|
||||
return s
|
||||
return _snippet_str(psr, max_len)
|
||||
|
||||
|
||||
def _snippet_str(text: Any, max_len: int = 200) -> Optional[str]:
|
||||
if text is None:
|
||||
return None
|
||||
if isinstance(text, str):
|
||||
s = text
|
||||
elif isinstance(text, list):
|
||||
parts = []
|
||||
for item in text:
|
||||
if isinstance(item, dict) and "content" in item:
|
||||
c = item["content"]
|
||||
parts.append(c if isinstance(c, str) else str(c))
|
||||
else:
|
||||
parts.append(str(item))
|
||||
s = " ".join(parts)
|
||||
else:
|
||||
s = str(text)
|
||||
if not s or s == "{}":
|
||||
return None
|
||||
return (s[:max_len] + "...") if len(s) > max_len else s
|
||||
|
||||
|
||||
@router.get(
|
||||
"/v1/tool/{tool_name:path}/logs",
|
||||
tags=["tool management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=ToolUsageLogsResponse,
|
||||
)
|
||||
async def get_tool_usage_logs(
|
||||
tool_name: str,
|
||||
page: int = Query(1, ge=1),
|
||||
page_size: int = Query(50, ge=1, le=100),
|
||||
start_date: Optional[str] = Query(None, description="YYYY-MM-DD"),
|
||||
end_date: Optional[str] = Query(None, description="YYYY-MM-DD"),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Return paginated spend logs for requests that used this tool (from SpendLogToolIndex).
|
||||
"""
|
||||
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:
|
||||
where: dict = {"tool_name": tool_name}
|
||||
if start_date or end_date:
|
||||
start_time_filter: Optional[datetime] = None
|
||||
end_time_filter: Optional[datetime] = None
|
||||
if start_date:
|
||||
try:
|
||||
start_time_filter = datetime.strptime(
|
||||
start_date + "T00:00:00", "%Y-%m-%dT%H:%M:%S"
|
||||
).replace(tzinfo=timezone.utc)
|
||||
except ValueError:
|
||||
pass
|
||||
if end_date:
|
||||
try:
|
||||
end_time_filter = datetime.strptime(
|
||||
end_date + "T23:59:59", "%Y-%m-%dT%H:%M:%S"
|
||||
).replace(tzinfo=timezone.utc)
|
||||
except ValueError:
|
||||
pass
|
||||
if start_time_filter is not None or end_time_filter is not None:
|
||||
where["start_time"] = {}
|
||||
if start_time_filter is not None:
|
||||
where["start_time"]["gte"] = start_time_filter
|
||||
if end_time_filter is not None:
|
||||
where["start_time"]["lte"] = end_time_filter
|
||||
|
||||
total = await prisma_client.db.litellm_spendlogtoolindex.count(where=where)
|
||||
index_rows = await prisma_client.db.litellm_spendlogtoolindex.find_many(
|
||||
where=where,
|
||||
order={"start_time": "desc"},
|
||||
skip=(page - 1) * page_size,
|
||||
take=page_size,
|
||||
)
|
||||
request_ids = [r.request_id for r in index_rows]
|
||||
if not request_ids:
|
||||
return ToolUsageLogsResponse(
|
||||
logs=[], total=total, page=page, page_size=page_size
|
||||
)
|
||||
|
||||
spend_logs = await prisma_client.db.litellm_spendlogs.find_many(
|
||||
where={"request_id": {"in": request_ids}}
|
||||
)
|
||||
log_by_id = {s.request_id: s for s in spend_logs}
|
||||
|
||||
logs_out: List[ToolUsageLogEntry] = []
|
||||
for r in index_rows:
|
||||
sl = log_by_id.get(r.request_id)
|
||||
if not sl:
|
||||
continue
|
||||
ts = (
|
||||
sl.startTime.isoformat()
|
||||
if hasattr(sl.startTime, "isoformat")
|
||||
else str(sl.startTime)
|
||||
)
|
||||
logs_out.append(
|
||||
ToolUsageLogEntry(
|
||||
id=sl.request_id,
|
||||
timestamp=ts,
|
||||
model=getattr(sl, "model", None) or None,
|
||||
spend=getattr(sl, "spend", None),
|
||||
total_tokens=getattr(sl, "total_tokens", None),
|
||||
input_snippet=_input_snippet_for_tool_log(sl),
|
||||
)
|
||||
)
|
||||
|
||||
return ToolUsageLogsResponse(
|
||||
logs=logs_out, total=total, page=page, page_size=page_size
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error getting tool usage logs: %s", e)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.get(
|
||||
"/v1/tool/{tool_name:path}",
|
||||
tags=["tool management"],
|
||||
|
|
@ -85,9 +284,7 @@ async def get_tool(
|
|||
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"
|
||||
)
|
||||
raise HTTPException(status_code=404, detail=f"Tool '{tool_name}' not found")
|
||||
return tool
|
||||
except HTTPException:
|
||||
raise
|
||||
|
|
@ -96,6 +293,85 @@ async def get_tool(
|
|||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
async def _resolve_key_hash_to_object_permission_id(
|
||||
prisma_client: "PrismaClient",
|
||||
key_hash: str,
|
||||
) -> Optional[str]:
|
||||
"""Resolve key (hash or raw) to object_permission_id; create permission if key has none."""
|
||||
from litellm.proxy.proxy_server import hash_token
|
||||
|
||||
hashed = key_hash if "sk-" not in (key_hash or "") else hash_token(key_hash)
|
||||
if not hashed:
|
||||
return None
|
||||
row = await prisma_client.db.litellm_verificationtoken.find_unique(
|
||||
where={"token": hashed}
|
||||
)
|
||||
if row is None:
|
||||
return None
|
||||
op_id = getattr(row, "object_permission_id", None)
|
||||
if op_id:
|
||||
return op_id
|
||||
# Create new object permission and atomically assign to key.
|
||||
# Uses update_many with object_permission_id=None to prevent race conditions:
|
||||
# only one concurrent request wins; the loser cleans up its orphaned row.
|
||||
new_id = str(uuid.uuid4())
|
||||
await prisma_client.db.litellm_objectpermissiontable.create(
|
||||
data={"object_permission_id": new_id, "blocked_tools": []}
|
||||
)
|
||||
updated_count = await prisma_client.db.litellm_verificationtoken.update_many(
|
||||
where={"token": hashed, "object_permission_id": None},
|
||||
data={"object_permission_id": new_id},
|
||||
)
|
||||
if updated_count == 0:
|
||||
# Another request already assigned a permission; clean up orphan
|
||||
await prisma_client.db.litellm_objectpermissiontable.delete(
|
||||
where={"object_permission_id": new_id}
|
||||
)
|
||||
row = await prisma_client.db.litellm_verificationtoken.find_unique(
|
||||
where={"token": hashed}
|
||||
)
|
||||
return getattr(row, "object_permission_id", None) if row else None
|
||||
return new_id
|
||||
|
||||
|
||||
async def _resolve_team_id_to_object_permission_id(
|
||||
prisma_client: "PrismaClient",
|
||||
team_id: str,
|
||||
) -> Optional[str]:
|
||||
"""Resolve team_id to object_permission_id; create permission if team has none."""
|
||||
if not team_id or not team_id.strip():
|
||||
return None
|
||||
team_id_clean = team_id.strip()
|
||||
row = await prisma_client.db.litellm_teamtable.find_unique(
|
||||
where={"team_id": team_id_clean},
|
||||
select={"object_permission_id": True},
|
||||
)
|
||||
if row is None:
|
||||
return None
|
||||
op_id = getattr(row, "object_permission_id", None)
|
||||
if op_id:
|
||||
return op_id
|
||||
# Same atomic pattern as _resolve_key_hash_to_object_permission_id
|
||||
new_id = str(uuid.uuid4())
|
||||
await prisma_client.db.litellm_objectpermissiontable.create(
|
||||
data={"object_permission_id": new_id, "blocked_tools": []}
|
||||
)
|
||||
updated_count = await prisma_client.db.litellm_teamtable.update_many(
|
||||
where={"team_id": team_id_clean, "object_permission_id": None},
|
||||
data={"object_permission_id": new_id},
|
||||
)
|
||||
if updated_count == 0:
|
||||
await prisma_client.db.litellm_objectpermissiontable.delete(
|
||||
where={"object_permission_id": new_id}
|
||||
)
|
||||
row = await prisma_client.db.litellm_teamtable.find_unique(
|
||||
where={"team_id": team_id_clean},
|
||||
select={"object_permission_id": True},
|
||||
)
|
||||
return getattr(row, "object_permission_id", None) if row else None
|
||||
return new_id
|
||||
|
||||
|
||||
@router.post(
|
||||
"/v1/tool/policy",
|
||||
tags=["tool management"],
|
||||
|
|
@ -107,15 +383,24 @@ async def update_tool_policy(
|
|||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Set the call policy for a tool.
|
||||
Set the call policy for a tool (global) or for a specific team/key (override).
|
||||
|
||||
Parameters:
|
||||
- tool_name: str - The tool to update
|
||||
- call_policy: "trusted" | "untrusted" | "dual_llm" | "blocked"
|
||||
- team_id: optional - if set, create/update override for this team only
|
||||
- key_hash: optional - if set, create/update override for this key only
|
||||
- key_alias: optional - human-readable key alias for UI
|
||||
|
||||
Setting a tool to "blocked" will cause the ToolPolicyGuardrail to remove
|
||||
that tool_call from LLM responses before returning them to the client.
|
||||
If both team_id and key_hash are omitted, updates the global tool policy.
|
||||
Setting a tool to "blocked" will cause the ToolPolicyGuardrail to reject
|
||||
that tool_call for the relevant scope.
|
||||
"""
|
||||
from litellm.proxy.db.tool_registry_writer import (
|
||||
add_tool_to_object_permission_blocked,
|
||||
get_tool_policy_registry,
|
||||
remove_tool_from_object_permission_blocked,
|
||||
)
|
||||
from litellm.proxy.db.tool_registry_writer import (
|
||||
update_tool_policy as db_update_tool_policy,
|
||||
)
|
||||
|
|
@ -127,6 +412,52 @@ async def update_tool_policy(
|
|||
)
|
||||
|
||||
try:
|
||||
if data.team_id is not None or data.key_hash is not None:
|
||||
if data.team_id is not None and data.key_hash is not None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Provide either team_id or key_hash, not both",
|
||||
)
|
||||
if data.key_hash is not None:
|
||||
op_id = await _resolve_key_hash_to_object_permission_id(
|
||||
prisma_client, data.key_hash
|
||||
)
|
||||
else:
|
||||
op_id = await _resolve_team_id_to_object_permission_id(
|
||||
prisma_client, data.team_id or ""
|
||||
)
|
||||
if op_id is None:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail="Key or team not found for the given identifier",
|
||||
)
|
||||
if data.call_policy == "blocked":
|
||||
ok = await add_tool_to_object_permission_blocked(
|
||||
prisma_client=prisma_client,
|
||||
object_permission_id=op_id,
|
||||
tool_name=data.tool_name,
|
||||
)
|
||||
else:
|
||||
ok = await remove_tool_from_object_permission_blocked(
|
||||
prisma_client=prisma_client,
|
||||
object_permission_id=op_id,
|
||||
tool_name=data.tool_name,
|
||||
)
|
||||
if not ok:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail=f"Failed to update policy override for tool '{data.tool_name}'",
|
||||
)
|
||||
registry = get_tool_policy_registry()
|
||||
if registry.is_initialized():
|
||||
await registry.sync_tool_policy_from_db(prisma_client)
|
||||
return ToolPolicyUpdateResponse(
|
||||
tool_name=data.tool_name,
|
||||
call_policy=data.call_policy,
|
||||
updated=True,
|
||||
team_id=data.team_id,
|
||||
key_hash=data.key_hash,
|
||||
)
|
||||
updated = await db_update_tool_policy(
|
||||
prisma_client=prisma_client,
|
||||
tool_name=data.tool_name,
|
||||
|
|
@ -135,8 +466,12 @@ async def update_tool_policy(
|
|||
)
|
||||
if updated is None:
|
||||
raise HTTPException(
|
||||
status_code=500, detail=f"Failed to update policy for tool '{data.tool_name}'"
|
||||
status_code=500,
|
||||
detail=f"Failed to update policy for tool '{data.tool_name}'",
|
||||
)
|
||||
registry = get_tool_policy_registry()
|
||||
if registry.is_initialized():
|
||||
await registry.sync_tool_policy_from_db(prisma_client)
|
||||
return ToolPolicyUpdateResponse(
|
||||
tool_name=updated.tool_name,
|
||||
call_policy=updated.call_policy,
|
||||
|
|
@ -147,3 +482,77 @@ async def update_tool_policy(
|
|||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error updating tool policy: %s", e)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/v1/tool/{tool_name:path}/overrides",
|
||||
tags=["tool management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def delete_tool_policy_override(
|
||||
tool_name: str,
|
||||
team_id: Optional[str] = Query(
|
||||
None, description="Team ID of the override to remove"
|
||||
),
|
||||
key_hash: Optional[str] = Query(
|
||||
None, description="Key hash of the override to remove"
|
||||
),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Remove a policy override for a tool. Specify the override by team_id or key_hash
|
||||
(exactly one required).
|
||||
"""
|
||||
from litellm.proxy.db.tool_registry_writer import (
|
||||
get_tool_policy_registry,
|
||||
remove_tool_from_object_permission_blocked,
|
||||
)
|
||||
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
|
||||
)
|
||||
if team_id is None and key_hash is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="At least one of team_id or key_hash is required to identify the override",
|
||||
)
|
||||
if team_id is not None and key_hash is not None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Provide either team_id or key_hash, not both",
|
||||
)
|
||||
try:
|
||||
if key_hash is not None:
|
||||
op_id = await _resolve_key_hash_to_object_permission_id(
|
||||
prisma_client, key_hash
|
||||
)
|
||||
else:
|
||||
op_id = await _resolve_team_id_to_object_permission_id(
|
||||
prisma_client, team_id or ""
|
||||
)
|
||||
if op_id is None:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail="Key or team not found for the given identifier",
|
||||
)
|
||||
deleted = await remove_tool_from_object_permission_blocked(
|
||||
prisma_client=prisma_client,
|
||||
object_permission_id=op_id,
|
||||
tool_name=tool_name,
|
||||
)
|
||||
if not deleted:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail=f"No override found for tool '{tool_name}' with the given scope",
|
||||
)
|
||||
registry = get_tool_policy_registry()
|
||||
if registry.is_initialized():
|
||||
await registry.sync_tool_policy_from_db(prisma_client)
|
||||
return {"deleted": True, "tool_name": tool_name}
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error deleting tool policy override: %s", e)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
|
|
|||
|
|
@ -366,9 +366,7 @@ from litellm.proxy.management_endpoints.fallback_management_endpoints import (
|
|||
from litellm.proxy.management_endpoints.internal_user_endpoints import (
|
||||
router as internal_user_router,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.internal_user_endpoints import (
|
||||
user_update,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.internal_user_endpoints import user_update
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
delete_verification_tokens,
|
||||
duration_in_seconds,
|
||||
|
|
@ -434,9 +432,7 @@ from litellm.proxy.openai_evals_endpoints.endpoints import router as evals_route
|
|||
from litellm.proxy.openai_files_endpoints.files_endpoints import (
|
||||
router as openai_files_router,
|
||||
)
|
||||
from litellm.proxy.openai_files_endpoints.files_endpoints import (
|
||||
set_files_config,
|
||||
)
|
||||
from litellm.proxy.openai_files_endpoints.files_endpoints import set_files_config
|
||||
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
|
||||
passthrough_endpoint_router,
|
||||
)
|
||||
|
|
@ -535,9 +531,7 @@ from litellm.types.proxy.management_endpoints.ui_sso import (
|
|||
LiteLLM_UpperboundKeyGenerateParams,
|
||||
)
|
||||
from litellm.types.realtime import RealtimeQueryParams
|
||||
from litellm.types.router import (
|
||||
DeploymentTypedDict,
|
||||
)
|
||||
from litellm.types.router import DeploymentTypedDict
|
||||
from litellm.types.router import ModelInfo as RouterModelInfo
|
||||
from litellm.types.router import (
|
||||
RouterGeneralSettings,
|
||||
|
|
@ -4411,6 +4405,9 @@ class ProxyConfig:
|
|||
if self._should_load_db_object(object_type="search_tools"):
|
||||
await self._init_search_tools_in_db(prisma_client=prisma_client)
|
||||
|
||||
if self._should_load_db_object(object_type="tools"):
|
||||
await self._init_tool_policy_in_db(prisma_client=prisma_client)
|
||||
|
||||
if self._should_load_db_object(object_type="model_cost_map"):
|
||||
await self._check_and_reload_model_cost_map(prisma_client=prisma_client)
|
||||
|
||||
|
|
@ -4847,6 +4844,24 @@ class ProxyConfig:
|
|||
)
|
||||
)
|
||||
|
||||
async def _init_tool_policy_in_db(self, prisma_client: PrismaClient):
|
||||
"""
|
||||
Initialize tool policy from database into the in-memory registry.
|
||||
Synced periodically by add_deployment -> _init_non_llm_objects_in_db.
|
||||
"""
|
||||
from litellm.proxy.db.tool_registry_writer import get_tool_policy_registry
|
||||
|
||||
try:
|
||||
registry = get_tool_policy_registry()
|
||||
await registry.sync_tool_policy_from_db(prisma_client=prisma_client)
|
||||
verbose_proxy_logger.debug("Successfully synced tool policy from DB")
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"litellm.proxy.proxy_server.py::ProxyConfig:_init_tool_policy_in_db - {}".format(
|
||||
str(e)
|
||||
)
|
||||
)
|
||||
|
||||
async def _init_vector_stores_in_db(self, prisma_client: PrismaClient):
|
||||
from litellm.vector_stores.vector_store_registry import VectorStoreRegistry
|
||||
|
||||
|
|
@ -10577,6 +10592,12 @@ async def async_queue_request(
|
|||
data["metadata"]["user_api_key_team_id"] = getattr(
|
||||
user_api_key_dict, "team_id", None
|
||||
)
|
||||
data["metadata"]["user_api_key_object_permission_id"] = getattr(
|
||||
user_api_key_dict, "object_permission_id", None
|
||||
)
|
||||
data["metadata"]["user_api_key_team_object_permission_id"] = getattr(
|
||||
user_api_key_dict, "team_object_permission_id", None
|
||||
)
|
||||
data["metadata"]["endpoint"] = str(request.url)
|
||||
|
||||
global user_temperature, user_request_timeout, user_max_tokens, user_api_base
|
||||
|
|
@ -11093,9 +11114,7 @@ async def get_favicon():
|
|||
|
||||
if favicon_url.startswith(("http://", "https://")):
|
||||
try:
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
|
||||
async_client = get_async_httpx_client(
|
||||
|
|
|
|||
|
|
@ -260,6 +260,7 @@ model LiteLLM_ObjectPermissionTable {
|
|||
vector_stores String[] @default([])
|
||||
agents String[] @default([])
|
||||
agent_access_groups String[] @default([])
|
||||
blocked_tools String[] @default([]) // Tool names blocked for any key/team/user with this permission
|
||||
teams LiteLLM_TeamTable[]
|
||||
projects LiteLLM_ProjectTable[]
|
||||
verification_tokens LiteLLM_VerificationToken[]
|
||||
|
|
@ -390,7 +391,7 @@ model LiteLLM_DeletedVerificationToken {
|
|||
config Json @default("{}")
|
||||
user_id String?
|
||||
team_id String?
|
||||
agent_id String?
|
||||
agent_id String?
|
||||
project_id String?
|
||||
permissions Json @default("{}")
|
||||
max_parallel_requests Int?
|
||||
|
|
@ -921,6 +922,16 @@ model LiteLLM_SpendLogGuardrailIndex {
|
|||
@@index([policy_id, start_time])
|
||||
}
|
||||
|
||||
// Index for fast "last N logs for tool" from SpendLogs – see how a tool is called in production
|
||||
model LiteLLM_SpendLogToolIndex {
|
||||
request_id String
|
||||
tool_name String // matches LiteLLM_ToolTable.tool_name; join for call_policy etc.
|
||||
start_time DateTime
|
||||
|
||||
@@id([request_id, tool_name])
|
||||
@@index([tool_name, start_time])
|
||||
}
|
||||
|
||||
// Prompt table for storing prompt configurations
|
||||
model LiteLLM_PromptTable {
|
||||
id String @id @default(uuid())
|
||||
|
|
|
|||
|
|
@ -317,7 +317,7 @@ class ProxyLogging:
|
|||
self.slack_alerting_instance: SlackAlerting = SlackAlerting(
|
||||
alerting_threshold=self.alerting_threshold,
|
||||
alerting=self.alerting,
|
||||
internal_usage_cache=self.internal_usage_cache.dual_cache,
|
||||
internal_usage_cache=self.internal_usage_cache.dual_cache, # type: ignore[call-arg]
|
||||
)
|
||||
self.email_logging_instance: Optional[Any] = None
|
||||
if BaseEmailLogger is not None:
|
||||
|
|
@ -325,7 +325,7 @@ class ProxyLogging:
|
|||
if email_logger_class is not None:
|
||||
# All email logger classes now accept internal_usage_cache
|
||||
self.email_logging_instance = email_logger_class(
|
||||
internal_usage_cache=self.internal_usage_cache.dual_cache,
|
||||
internal_usage_cache=self.internal_usage_cache.dual_cache, # type: ignore[call-arg]
|
||||
)
|
||||
self.premium_user = premium_user
|
||||
self.service_logging_obj = ServiceLogging()
|
||||
|
|
@ -3583,8 +3583,9 @@ class PrismaClient:
|
|||
def _get_engine_pid(self) -> int:
|
||||
try:
|
||||
engine = self.db._original_prisma._engine # type: ignore[attr-defined]
|
||||
if engine is not None and engine.process is not None:
|
||||
return engine.process.pid
|
||||
process = getattr(engine, "process", None) if engine is not None else None
|
||||
if process is not None:
|
||||
return process.pid
|
||||
except (AttributeError, TypeError):
|
||||
pass
|
||||
return 0
|
||||
|
|
@ -4688,6 +4689,19 @@ async def update_spend_logs_job(
|
|||
guardrail_tracking_err,
|
||||
)
|
||||
|
||||
# Tool usage tracking (same batch): SpendLogToolIndex for "last N requests for tool X"
|
||||
try:
|
||||
from litellm.proxy.db.spend_log_tool_index import process_spend_logs_tool_usage
|
||||
await process_spend_logs_tool_usage(
|
||||
prisma_client=prisma_client,
|
||||
logs_to_process=logs_to_process,
|
||||
)
|
||||
except Exception as tool_tracking_err:
|
||||
verbose_proxy_logger.warning(
|
||||
"Spend tracking - tool usage tracking failed (non-fatal): %s",
|
||||
tool_tracking_err,
|
||||
)
|
||||
|
||||
|
||||
async def _monitor_spend_logs_queue(
|
||||
prisma_client: PrismaClient,
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ Pydantic models for Tool Policy management endpoints.
|
|||
from datetime import datetime
|
||||
from typing import Dict, List, Literal, Optional
|
||||
|
||||
from pydantic import BaseModel
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
ToolCallPolicy = Literal["trusted", "untrusted", "dual_llm", "blocked"]
|
||||
|
||||
|
|
@ -35,9 +35,48 @@ class ToolListResponse(BaseModel):
|
|||
class ToolPolicyUpdateRequest(BaseModel):
|
||||
tool_name: str
|
||||
call_policy: ToolCallPolicy
|
||||
team_id: Optional[str] = None # if set, create/update override for this team
|
||||
key_hash: Optional[str] = None # if set, create/update override for this key
|
||||
key_alias: Optional[str] = None # human-readable key alias for UI
|
||||
|
||||
|
||||
class ToolPolicyUpdateResponse(BaseModel):
|
||||
tool_name: str
|
||||
call_policy: ToolCallPolicy
|
||||
updated: bool
|
||||
team_id: Optional[str] = None
|
||||
key_hash: Optional[str] = None
|
||||
|
||||
|
||||
class ToolPolicyOverrideRow(BaseModel):
|
||||
override_id: str
|
||||
tool_name: str
|
||||
team_id: Optional[str] = None
|
||||
key_hash: Optional[str] = None
|
||||
call_policy: ToolCallPolicy
|
||||
key_alias: Optional[str] = None
|
||||
created_at: Optional[datetime] = None
|
||||
updated_at: Optional[datetime] = None
|
||||
|
||||
|
||||
class ToolDetailResponse(BaseModel):
|
||||
tool: LiteLLM_ToolTableRow
|
||||
overrides: List[ToolPolicyOverrideRow] = Field(default_factory=list)
|
||||
|
||||
|
||||
class ToolUsageLogEntry(BaseModel):
|
||||
"""One spend log row for a tool call (for UI "recent logs" table)."""
|
||||
|
||||
id: str # request_id
|
||||
timestamp: str
|
||||
model: Optional[str] = None
|
||||
spend: Optional[float] = None
|
||||
total_tokens: Optional[int] = None
|
||||
input_snippet: Optional[str] = None
|
||||
|
||||
|
||||
class ToolUsageLogsResponse(BaseModel):
|
||||
logs: List[ToolUsageLogEntry]
|
||||
total: int
|
||||
page: int
|
||||
page_size: int
|
||||
|
|
|
|||
|
|
@ -260,6 +260,7 @@ model LiteLLM_ObjectPermissionTable {
|
|||
vector_stores String[] @default([])
|
||||
agents String[] @default([])
|
||||
agent_access_groups String[] @default([])
|
||||
blocked_tools String[] @default([]) // Tool names blocked for any key/team/user with this permission
|
||||
teams LiteLLM_TeamTable[]
|
||||
projects LiteLLM_ProjectTable[]
|
||||
verification_tokens LiteLLM_VerificationToken[]
|
||||
|
|
@ -390,7 +391,7 @@ model LiteLLM_DeletedVerificationToken {
|
|||
config Json @default("{}")
|
||||
user_id String?
|
||||
team_id String?
|
||||
agent_id String?
|
||||
agent_id String?
|
||||
project_id String?
|
||||
permissions Json @default("{}")
|
||||
max_parallel_requests Int?
|
||||
|
|
@ -921,6 +922,16 @@ model LiteLLM_SpendLogGuardrailIndex {
|
|||
@@index([policy_id, start_time])
|
||||
}
|
||||
|
||||
// Index for fast "last N logs for tool" from SpendLogs – see how a tool is called in production
|
||||
model LiteLLM_SpendLogToolIndex {
|
||||
request_id String
|
||||
tool_name String // matches LiteLLM_ToolTable.tool_name; join for call_policy etc.
|
||||
start_time DateTime
|
||||
|
||||
@@id([request_id, tool_name])
|
||||
@@index([tool_name, start_time])
|
||||
}
|
||||
|
||||
// Prompt table for storing prompt configurations
|
||||
model LiteLLM_PromptTable {
|
||||
id String @id @default(uuid())
|
||||
|
|
|
|||
116
scripts/test_tool_allowlist_script.py
Normal file
116
scripts/test_tool_allowlist_script.py
Normal file
|
|
@ -0,0 +1,116 @@
|
|||
#!/usr/bin/env python3
|
||||
"""
|
||||
Standalone script to test tool allowlist enforcement and tool name extraction.
|
||||
|
||||
Run from repo root:
|
||||
poetry run python scripts/test_tool_allowlist_script.py
|
||||
|
||||
Or run the unit tests:
|
||||
poetry run pytest tests/test_litellm/proxy/test_tools_allowlist_enforcement.py -v
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
# Ensure repo root is on path
|
||||
repo_root = Path(__file__).resolve().parent.parent
|
||||
if str(repo_root) not in sys.path:
|
||||
sys.path.insert(0, str(repo_root))
|
||||
|
||||
|
||||
def test_extraction():
|
||||
"""Test extract_request_tool_names for each API shape."""
|
||||
from litellm.proxy.guardrails.tool_name_extraction import extract_request_tool_names
|
||||
|
||||
cases = [
|
||||
("OpenAI chat tools", "/v1/chat/completions", {"tools": [{"type": "function", "function": {"name": "get_weather"}}]}),
|
||||
("OpenAI chat functions", "/v1/chat/completions", {"functions": [{"name": "run_sql"}]}),
|
||||
("OpenAI responses function", "/v1/responses", {"tools": [{"type": "function", "name": "get_current_weather"}]}),
|
||||
("OpenAI responses MCP", "/v1/responses", {"tools": [{"type": "mcp", "server_label": "dmcp"}]}),
|
||||
("Anthropic", "/v1/messages", {"tools": [{"name": "get_weather"}, {"name": "run_sql"}]}),
|
||||
("Google generateContent", "/generate_content", {"tools": [{"functionDeclarations": [{"name": "schedule_meeting"}]}]}),
|
||||
("MCP call_tool", "/mcp/call_tool", {"name": "my_tool", "arguments": {}}),
|
||||
("Non-tool route", "/v1/embeddings", {"tools": [{"type": "function", "function": {"name": "x"}}]}),
|
||||
]
|
||||
print("=== extract_request_tool_names(route, data) ===\n")
|
||||
for label, route, data in cases:
|
||||
names = extract_request_tool_names(route, data)
|
||||
print(f" {label}: {names}")
|
||||
print()
|
||||
|
||||
|
||||
async def test_check_tools_allowlist():
|
||||
"""Test check_tools_allowlist with mock tokens."""
|
||||
from litellm.proxy._types import ProxyErrorTypes, ProxyException, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.auth_checks import check_tools_allowlist
|
||||
|
||||
def token(metadata=None, team_metadata=None):
|
||||
return UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="user",
|
||||
team_id="team",
|
||||
org_id=None,
|
||||
models=["*"],
|
||||
metadata=metadata or {},
|
||||
team_metadata=team_metadata or {},
|
||||
)
|
||||
|
||||
print("=== check_tools_allowlist (auth) ===\n")
|
||||
|
||||
# No allowlist -> pass
|
||||
await check_tools_allowlist(
|
||||
request_body={"tools": [{"type": "function", "function": {"name": "get_weather"}}]},
|
||||
valid_token=token(),
|
||||
team_object=None,
|
||||
route="/v1/chat/completions",
|
||||
)
|
||||
print(" No allowlist, body has tools: PASS")
|
||||
|
||||
# Allowed tool -> pass
|
||||
await check_tools_allowlist(
|
||||
request_body={"tools": [{"type": "function", "function": {"name": "get_weather"}}]},
|
||||
valid_token=token(metadata={"allowed_tools": ["get_weather"]}),
|
||||
team_object=None,
|
||||
route="/v1/chat/completions",
|
||||
)
|
||||
print(" allowed_tools=['get_weather'], body has get_weather: PASS")
|
||||
|
||||
# Disallowed tool -> raise
|
||||
try:
|
||||
await check_tools_allowlist(
|
||||
request_body={"tools": [{"type": "function", "function": {"name": "get_weather"}}]},
|
||||
valid_token=token(metadata={"allowed_tools": ["other_tool"]}),
|
||||
team_object=None,
|
||||
route="/v1/chat/completions",
|
||||
)
|
||||
print(" DISALLOWED: expected ProxyException")
|
||||
except ProxyException as e:
|
||||
if e.type == ProxyErrorTypes.tool_access_denied:
|
||||
print(" allowed_tools=['other_tool'], body has get_weather: PASS (raised tool_access_denied)")
|
||||
else:
|
||||
print(f" Unexpected ProxyException type: {e.type}")
|
||||
except Exception as e:
|
||||
print(f" Unexpected: {e}")
|
||||
|
||||
# Team allowlist when key empty
|
||||
await check_tools_allowlist(
|
||||
request_body={"tools": [{"type": "function", "function": {"name": "get_weather"}}]},
|
||||
valid_token=token(team_metadata={"allowed_tools": ["get_weather"]}),
|
||||
team_object=None,
|
||||
route="/v1/chat/completions",
|
||||
)
|
||||
print(" team_metadata.allowed_tools=['get_weather']: PASS")
|
||||
print()
|
||||
|
||||
|
||||
def main():
|
||||
print("Tool allowlist / tool name extraction – script checks\n")
|
||||
test_extraction()
|
||||
asyncio.run(test_check_tools_allowlist())
|
||||
print("Done. For full unit tests run:")
|
||||
print(" poetry run pytest tests/test_litellm/proxy/test_tools_allowlist_enforcement.py -v")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
|
@ -1,6 +1,6 @@
|
|||
"""
|
||||
Unit tests for tool_registry_writer.py — uses a mock prisma client
|
||||
that exposes execute_raw / query_raw (matching the actual raw-SQL implementation).
|
||||
that exposes litellm_tooltable.upsert / find_many / find_unique.
|
||||
"""
|
||||
|
||||
import os
|
||||
|
|
@ -12,18 +12,22 @@ 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,
|
||||
)
|
||||
from litellm.proxy.db.tool_registry_writer import (ToolPolicyRegistry,
|
||||
batch_upsert_tools,
|
||||
get_tool,
|
||||
get_tool_policy_registry,
|
||||
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 = {
|
||||
def _mock_row(**kwargs):
|
||||
"""Build a row-like object with real attributes (no MagicMock) for _row_to_model."""
|
||||
|
||||
class Row:
|
||||
pass
|
||||
|
||||
default = {
|
||||
"tool_id": "uuid-1",
|
||||
"tool_name": "my_tool",
|
||||
"origin": "user_defined",
|
||||
|
|
@ -38,31 +42,53 @@ def _make_prisma(query_rows=None):
|
|||
"created_by": None,
|
||||
"updated_by": None,
|
||||
}
|
||||
rows = query_rows if query_rows is not None else [default_row]
|
||||
default.update(kwargs)
|
||||
row = Row()
|
||||
for k, v in default.items():
|
||||
setattr(row, k, v)
|
||||
return row
|
||||
|
||||
|
||||
def _make_prisma(
|
||||
*,
|
||||
upsert_return=None,
|
||||
find_many_rows=None,
|
||||
find_unique_row=None,
|
||||
):
|
||||
"""Return a mock prisma_client with litellm_tooltable.upsert, find_many, find_unique."""
|
||||
prisma = MagicMock()
|
||||
prisma.db.execute_raw = AsyncMock(return_value=None)
|
||||
prisma.db.query_raw = AsyncMock(return_value=rows)
|
||||
prisma.db.litellm_tooltable = MagicMock()
|
||||
prisma.db.litellm_tooltable.upsert = AsyncMock(return_value=upsert_return)
|
||||
prisma.db.litellm_tooltable.find_many = AsyncMock(
|
||||
return_value=find_many_rows if find_many_rows is not None else []
|
||||
)
|
||||
prisma.db.litellm_tooltable.find_unique = AsyncMock(
|
||||
return_value=find_unique_row
|
||||
)
|
||||
return prisma
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_batch_upsert_tools_calls_execute_raw():
|
||||
async def test_batch_upsert_tools_calls_upsert():
|
||||
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
|
||||
prisma.db.litellm_tooltable.upsert.assert_awaited_once()
|
||||
call_kw = prisma.db.litellm_tooltable.upsert.call_args.kwargs
|
||||
assert call_kw["where"] == {"tool_name": "tool_a"}
|
||||
assert call_kw["data"]["create"]["tool_name"] == "tool_a"
|
||||
assert call_kw["data"]["create"]["origin"] == "mcp_server"
|
||||
assert call_kw["data"]["create"]["call_policy"] == "untrusted"
|
||||
assert call_kw["data"]["create"]["call_count"] == 1
|
||||
assert call_kw["data"]["update"]["call_count"] == {"increment": 1}
|
||||
assert "updated_at" in call_kw["data"]["update"]
|
||||
|
||||
|
||||
@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()
|
||||
prisma.db.litellm_tooltable.upsert.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -70,123 +96,112 @@ 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()
|
||||
prisma.db.litellm_tooltable.upsert.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_batch_upsert_multiple_tools_calls_execute_raw_per_tool():
|
||||
async def test_batch_upsert_multiple_tools_calls_upsert_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
|
||||
assert prisma.db.litellm_tooltable.upsert.await_count == 2
|
||||
calls = prisma.db.litellm_tooltable.upsert.call_args_list
|
||||
assert calls[0].kwargs["where"]["tool_name"] == "tool_a"
|
||||
assert calls[1].kwargs["where"]["tool_name"] == "tool_b"
|
||||
|
||||
|
||||
@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])
|
||||
row = _mock_row(
|
||||
tool_id="id1",
|
||||
tool_name="tool_a",
|
||||
origin="mcp",
|
||||
call_policy="untrusted",
|
||||
call_count=5,
|
||||
)
|
||||
prisma = _make_prisma(find_many_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()
|
||||
prisma.db.litellm_tooltable.find_many.assert_awaited_once()
|
||||
call_kw = prisma.db.litellm_tooltable.find_many.call_args.kwargs
|
||||
assert call_kw["where"] == {}
|
||||
assert call_kw["order"] == {"created_at": "desc"}
|
||||
|
||||
|
||||
@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])
|
||||
row = _mock_row(
|
||||
tool_id="id1",
|
||||
tool_name="blocked_tool",
|
||||
origin=None,
|
||||
call_policy="blocked",
|
||||
call_count=2,
|
||||
assignments=None,
|
||||
)
|
||||
prisma = _make_prisma(find_many_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
|
||||
call_kw = prisma.db.litellm_tooltable.find_many.call_args.kwargs
|
||||
assert call_kw["where"] == {"call_policy": "blocked"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_tool_found():
|
||||
prisma = _make_prisma()
|
||||
row = _mock_row(tool_name="my_tool")
|
||||
prisma = _make_prisma(find_unique_row=row)
|
||||
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()
|
||||
prisma.db.litellm_tooltable.find_unique.assert_awaited_once_with(
|
||||
where={"tool_name": "my_tool"}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_tool_not_found():
|
||||
prisma = _make_prisma(query_rows=[])
|
||||
prisma = _make_prisma(find_unique_row=None)
|
||||
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])
|
||||
async def test_update_tool_policy_calls_upsert_then_get_tool():
|
||||
row = _mock_row(
|
||||
tool_name="my_tool",
|
||||
call_policy="blocked",
|
||||
updated_by="admin",
|
||||
)
|
||||
prisma = _make_prisma(find_unique_row=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
|
||||
prisma.db.litellm_tooltable.upsert.assert_awaited_once()
|
||||
call_kw = prisma.db.litellm_tooltable.upsert.call_args.kwargs
|
||||
assert call_kw["where"] == {"tool_name": "my_tool"}
|
||||
assert call_kw["data"]["update"]["call_policy"] == "blocked"
|
||||
assert call_kw["data"]["update"]["updated_by"] == "admin"
|
||||
prisma.db.litellm_tooltable.find_unique.assert_awaited_with(
|
||||
where={"tool_name": "my_tool"}
|
||||
)
|
||||
|
||||
|
||||
@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"},
|
||||
_mock_row(tool_name="tool_a", call_policy="trusted"),
|
||||
_mock_row(tool_name="tool_b", call_policy="blocked"),
|
||||
]
|
||||
prisma = _make_prisma(query_rows=rows)
|
||||
prisma = _make_prisma(find_many_rows=rows)
|
||||
result = await get_tools_by_names(prisma, ["tool_a", "tool_b"])
|
||||
assert result == {"tool_a": "trusted", "tool_b": "blocked"}
|
||||
prisma.db.litellm_tooltable.find_many.assert_awaited_once_with(
|
||||
where={"tool_name": {"in": ["tool_a", "tool_b"]}}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -194,4 +209,66 @@ 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()
|
||||
prisma.db.litellm_tooltable.find_many.assert_not_awaited()
|
||||
|
||||
|
||||
# --- ToolPolicyRegistry ---
|
||||
|
||||
|
||||
def _mock_tool_row(tool_name: str, call_policy: str = "untrusted"):
|
||||
row = MagicMock()
|
||||
row.tool_name = tool_name
|
||||
row.call_policy = call_policy
|
||||
return row
|
||||
|
||||
|
||||
def _mock_perm_row(object_permission_id: str, blocked_tools: list):
|
||||
row = MagicMock()
|
||||
row.object_permission_id = object_permission_id
|
||||
row.blocked_tools = blocked_tools
|
||||
return row
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_policy_registry_sync_and_get_effective_policies():
|
||||
"""Registry syncs from DB; get_effective_policies returns merged blocked + global."""
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_tooltable.find_many = AsyncMock(
|
||||
return_value=[
|
||||
_mock_tool_row("tool_a", "trusted"),
|
||||
_mock_tool_row("tool_b", "blocked"),
|
||||
_mock_tool_row("tool_c", "untrusted"),
|
||||
]
|
||||
)
|
||||
prisma.db.litellm_objectpermissiontable.find_many = AsyncMock(
|
||||
return_value=[
|
||||
_mock_perm_row("op-key-1", ["tool_a"]),
|
||||
_mock_perm_row("op-team-1", ["tool_c"]),
|
||||
]
|
||||
)
|
||||
registry = get_tool_policy_registry()
|
||||
await registry.sync_tool_policy_from_db(prisma)
|
||||
assert registry.is_initialized()
|
||||
# Key blocked: tool_a. Team blocked: tool_c. Global: tool_b blocked.
|
||||
result = registry.get_effective_policies(
|
||||
["tool_a", "tool_b", "tool_c"],
|
||||
object_permission_id="op-key-1",
|
||||
team_object_permission_id="op-team-1",
|
||||
)
|
||||
assert result["tool_a"] == "blocked"
|
||||
assert result["tool_b"] == "blocked"
|
||||
assert result["tool_c"] == "blocked"
|
||||
# No op ids: only global
|
||||
result_global = registry.get_effective_policies(["tool_a", "tool_b", "tool_c"])
|
||||
assert result_global["tool_a"] == "trusted"
|
||||
assert result_global["tool_b"] == "blocked"
|
||||
assert result_global["tool_c"] == "untrusted"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_policy_registry_not_initialized_returns_untrusted():
|
||||
"""When not synced, get_effective_policies still returns untrusted for unknown tools."""
|
||||
registry = ToolPolicyRegistry()
|
||||
assert not registry.is_initialized()
|
||||
result = registry.get_effective_policies(["unknown_tool"])
|
||||
assert result == {"unknown_tool": "untrusted"}
|
||||
|
|
|
|||
|
|
@ -12,9 +12,8 @@ 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.proxy.guardrails.guardrail_hooks.tool_policy.tool_policy_guardrail import \
|
||||
ToolPolicyGuardrail
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
|
||||
|
|
@ -70,10 +69,21 @@ async def test_no_tool_calls_in_response_passes_through(guardrail):
|
|||
assert result is inputs
|
||||
|
||||
|
||||
def _registry_mock(policy_map: dict):
|
||||
"""Return a mock registry with is_initialized=True and get_effective_policies returning policy_map."""
|
||||
reg = MagicMock()
|
||||
reg.is_initialized.return_value = True
|
||||
reg.get_effective_policies.return_value = policy_map
|
||||
return reg
|
||||
|
||||
|
||||
@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)):
|
||||
with patch(
|
||||
"litellm.proxy.db.tool_registry_writer.get_tool_policy_registry",
|
||||
return_value=_registry_mock(policy_map),
|
||||
):
|
||||
inputs: Any = _tool_request_inputs(["search", "read_file"])
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs, request_data={}, input_type="request"
|
||||
|
|
@ -84,7 +94,10 @@ async def test_untrusted_tools_pass_through(guardrail):
|
|||
@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)):
|
||||
with patch(
|
||||
"litellm.proxy.db.tool_registry_writer.get_tool_policy_registry",
|
||||
return_value=_registry_mock(policy_map),
|
||||
):
|
||||
inputs: Any = _tool_request_inputs(["dangerous_tool"])
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail.apply_guardrail(
|
||||
|
|
@ -97,7 +110,10 @@ async def test_blocked_tool_in_request_raises_http_exception(guardrail):
|
|||
@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)):
|
||||
with patch(
|
||||
"litellm.proxy.db.tool_registry_writer.get_tool_policy_registry",
|
||||
return_value=_registry_mock(policy_map),
|
||||
):
|
||||
inputs: Any = _tool_response_inputs(["exfil_tool"])
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail.apply_guardrail(
|
||||
|
|
@ -110,7 +126,10 @@ async def test_blocked_tool_in_response_raises_http_exception(guardrail):
|
|||
@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)):
|
||||
with patch(
|
||||
"litellm.proxy.db.tool_registry_writer.get_tool_policy_registry",
|
||||
return_value=_registry_mock(policy_map),
|
||||
):
|
||||
inputs: Any = _tool_request_inputs(["safe_tool", "bad_tool"])
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail.apply_guardrail(
|
||||
|
|
@ -123,8 +142,11 @@ async def test_mixed_blocked_and_allowed_raises_for_blocked(guardrail):
|
|||
|
||||
@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={})):
|
||||
"""When registry returns no policy (or empty), tools are not blocked."""
|
||||
with patch(
|
||||
"litellm.proxy.db.tool_registry_writer.get_tool_policy_registry",
|
||||
return_value=_registry_mock({}),
|
||||
):
|
||||
inputs: Any = _tool_request_inputs(["unknown_tool"])
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs, request_data={}, input_type="request"
|
||||
|
|
@ -133,43 +155,30 @@ async def test_tool_not_in_db_passes_through(guardrail):
|
|||
|
||||
|
||||
@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"}
|
||||
async def test_registry_not_initialized_passes_through(guardrail):
|
||||
"""When registry is not initialized, no tools are blocked (empty policy map)."""
|
||||
reg = MagicMock()
|
||||
reg.is_initialized.return_value = False
|
||||
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(),
|
||||
"litellm.proxy.db.tool_registry_writer.get_tool_policy_registry",
|
||||
return_value=reg,
|
||||
):
|
||||
# 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 == {}
|
||||
inputs: Any = _tool_request_inputs(["any_tool"])
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs, request_data={}, input_type="request"
|
||||
)
|
||||
assert result is inputs
|
||||
reg.get_effective_policies.assert_not_called()
|
||||
|
||||
|
||||
@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)):
|
||||
with patch(
|
||||
"litellm.proxy.db.tool_registry_writer.get_tool_policy_registry",
|
||||
return_value=_registry_mock(policy_map),
|
||||
):
|
||||
fn = MagicMock()
|
||||
fn.name = "obj_tool"
|
||||
tc = MagicMock()
|
||||
|
|
|
|||
200
tests/test_litellm/proxy/test_tools_allowlist_enforcement.py
Normal file
200
tests/test_litellm/proxy/test_tools_allowlist_enforcement.py
Normal file
|
|
@ -0,0 +1,200 @@
|
|||
"""
|
||||
Tests for tool allowlist enforcement (key/team metadata.allowed_tools).
|
||||
|
||||
Covers:
|
||||
- check_tools_allowlist: allowed, disallowed, no allowlist, non-tool routes
|
||||
- extract_request_tool_names: OpenAI chat, responses, Anthropic, generate_content, MCP
|
||||
"""
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.proxy._types import (ProxyErrorTypes, ProxyException,
|
||||
UserAPIKeyAuth)
|
||||
from litellm.proxy.auth.auth_checks import check_tools_allowlist
|
||||
from litellm.proxy.guardrails.tool_name_extraction import (
|
||||
TOOL_CAPABLE_CALL_TYPES, extract_request_tool_names)
|
||||
|
||||
|
||||
def _token(metadata=None, team_metadata=None):
|
||||
return UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="user",
|
||||
team_id="team",
|
||||
org_id=None,
|
||||
models=["*"],
|
||||
metadata=metadata or {},
|
||||
team_metadata=team_metadata or {},
|
||||
)
|
||||
|
||||
|
||||
class TestExtractRequestToolNames:
|
||||
"""Test tool name extraction per API format."""
|
||||
|
||||
def test_openai_chat_tools(self):
|
||||
data = {
|
||||
"tools": [
|
||||
{"type": "function", "function": {"name": "get_weather"}},
|
||||
{"type": "function", "function": {"name": "run_sql"}},
|
||||
]
|
||||
}
|
||||
assert extract_request_tool_names("/v1/chat/completions", data) == [
|
||||
"get_weather",
|
||||
"run_sql",
|
||||
]
|
||||
|
||||
def test_openai_chat_functions_legacy(self):
|
||||
data = {"functions": [{"name": "get_weather"}, {"name": "run_sql"}]}
|
||||
assert extract_request_tool_names("/v1/chat/completions", data) == [
|
||||
"get_weather",
|
||||
"run_sql",
|
||||
]
|
||||
|
||||
def test_openai_responses_function_tools(self):
|
||||
data = {
|
||||
"tools": [
|
||||
{"type": "function", "name": "get_current_weather", "description": "x"},
|
||||
]
|
||||
}
|
||||
assert extract_request_tool_names("/v1/responses", data) == [
|
||||
"get_current_weather"
|
||||
]
|
||||
|
||||
def test_openai_responses_mcp_tools(self):
|
||||
data = {
|
||||
"tools": [
|
||||
{"type": "mcp", "server_label": "dmcp", "server_url": "http://x"},
|
||||
]
|
||||
}
|
||||
assert extract_request_tool_names("/v1/responses", data) == ["dmcp"]
|
||||
|
||||
def test_anthropic_tools(self):
|
||||
data = {"tools": [{"name": "get_weather"}, {"name": "run_sql"}]}
|
||||
assert extract_request_tool_names("/v1/messages", data) == [
|
||||
"get_weather",
|
||||
"run_sql",
|
||||
]
|
||||
|
||||
def test_generate_content_tools(self):
|
||||
data = {
|
||||
"tools": [
|
||||
{
|
||||
"functionDeclarations": [
|
||||
{"name": "schedule_meeting", "description": "x"},
|
||||
]
|
||||
},
|
||||
]
|
||||
}
|
||||
assert extract_request_tool_names("/generate_content", data) == [
|
||||
"schedule_meeting"
|
||||
]
|
||||
|
||||
def test_mcp_call_tool_name(self):
|
||||
data = {"name": "my_tool", "arguments": {}}
|
||||
assert extract_request_tool_names("/mcp/call_tool", data) == ["my_tool"]
|
||||
|
||||
def test_mcp_call_tool_mcp_tool_name(self):
|
||||
data = {"mcp_tool_name": "other_tool"}
|
||||
assert extract_request_tool_names("/mcp/call_tool", data) == ["other_tool"]
|
||||
|
||||
def test_non_tool_route_returns_empty(self):
|
||||
data = {"tools": [{"type": "function", "function": {"name": "x"}}]}
|
||||
assert extract_request_tool_names("/v1/embeddings", data) == []
|
||||
|
||||
|
||||
class TestCheckToolsAllowlist:
|
||||
"""Test allowlist enforcement in auth (no DB in hot path)."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_allowlist_passes(self):
|
||||
token = _token(metadata={}, team_metadata={})
|
||||
body = {
|
||||
"tools": [{"type": "function", "function": {"name": "get_weather"}}]
|
||||
}
|
||||
await check_tools_allowlist(
|
||||
request_body=body,
|
||||
valid_token=token,
|
||||
team_object=None,
|
||||
route="/v1/chat/completions",
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_allowed_tool_passes(self):
|
||||
token = _token(metadata={"allowed_tools": ["get_weather"]})
|
||||
body = {
|
||||
"tools": [{"type": "function", "function": {"name": "get_weather"}}]
|
||||
}
|
||||
await check_tools_allowlist(
|
||||
request_body=body,
|
||||
valid_token=token,
|
||||
team_object=None,
|
||||
route="/v1/chat/completions",
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disallowed_tool_raises(self):
|
||||
token = _token(metadata={"allowed_tools": ["other_tool"]})
|
||||
body = {
|
||||
"tools": [{"type": "function", "function": {"name": "get_weather"}}]
|
||||
}
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await check_tools_allowlist(
|
||||
request_body=body,
|
||||
valid_token=token,
|
||||
team_object=None,
|
||||
route="/v1/chat/completions",
|
||||
)
|
||||
assert exc_info.value.type == ProxyErrorTypes.tool_access_denied
|
||||
assert "get_weather" in str(exc_info.value.message)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_allowlist_used_when_key_empty(self):
|
||||
token = _token(
|
||||
metadata={},
|
||||
team_metadata={"allowed_tools": ["get_weather"]},
|
||||
)
|
||||
body = {
|
||||
"tools": [{"type": "function", "function": {"name": "get_weather"}}]
|
||||
}
|
||||
await check_tools_allowlist(
|
||||
request_body=body,
|
||||
valid_token=token,
|
||||
team_object=None,
|
||||
route="/v1/chat/completions",
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_allowlist_overrides_team(self):
|
||||
token = _token(
|
||||
metadata={"allowed_tools": ["get_weather"]},
|
||||
team_metadata={"allowed_tools": ["other_tool"]},
|
||||
)
|
||||
body = {
|
||||
"tools": [{"type": "function", "function": {"name": "get_weather"}}]
|
||||
}
|
||||
await check_tools_allowlist(
|
||||
request_body=body,
|
||||
valid_token=token,
|
||||
team_object=None,
|
||||
route="/v1/chat/completions",
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_valid_token_none_skips(self):
|
||||
await check_tools_allowlist(
|
||||
request_body={"tools": [{"type": "function", "function": {"name": "x"}}]},
|
||||
valid_token=None,
|
||||
team_object=None,
|
||||
route="/v1/chat/completions",
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_tools_in_body_passes(self):
|
||||
token = _token(metadata={"allowed_tools": ["get_weather"]})
|
||||
await check_tools_allowlist(
|
||||
request_body={"messages": []},
|
||||
valid_token=token,
|
||||
team_object=None,
|
||||
route="/v1/chat/completions",
|
||||
)
|
||||
|
|
@ -39,7 +39,7 @@ import UserDashboard from "@/components/user_dashboard";
|
|||
import { AccessGroupsPage } from "@/components/AccessGroups/AccessGroupsPage";
|
||||
import { ProjectsPage } from "@/components/Projects/ProjectsPage";
|
||||
import VectorStoreManagement from "@/components/vector_store_management";
|
||||
import ToolPolicies from "@/components/ToolPolicies";
|
||||
import ToolPoliciesView from "@/components/ToolPoliciesView";
|
||||
import SpendLogsTable from "@/components/view_logs";
|
||||
import ViewUserDashboard from "@/components/view_users";
|
||||
import { ThemeProvider } from "@/contexts/ThemeContext";
|
||||
|
|
@ -549,7 +549,7 @@ function CreateKeyPageContent() {
|
|||
) : page == "vector-stores" ? (
|
||||
<VectorStoreManagement accessToken={accessToken} userRole={userRole} userID={userID} />
|
||||
) : page == "tool-policies" ? (
|
||||
<ToolPolicies accessToken={accessToken} userRole={userRole} />
|
||||
<ToolPoliciesView accessToken={accessToken} userRole={userRole} />
|
||||
) : page == "guardrails-monitor" ? (
|
||||
<GuardrailsMonitorView accessToken={accessToken} />
|
||||
) : page == "new_usage" ? (
|
||||
|
|
|
|||
373
ui/litellm-dashboard/src/components/ToolDetail.tsx
Normal file
373
ui/litellm-dashboard/src/components/ToolDetail.tsx
Normal file
|
|
@ -0,0 +1,373 @@
|
|||
"use client";
|
||||
|
||||
import { ArrowLeftOutlined, HistoryOutlined, ToolOutlined } from "@ant-design/icons";
|
||||
import { useQuery, useQueryClient } from "@tanstack/react-query";
|
||||
import { Button, Select, Spin } from "antd";
|
||||
import React, { useCallback, useMemo, useState } from "react";
|
||||
import TeamDropdown from "@/components/common_components/team_dropdown";
|
||||
import { LogViewer } from "@/components/GuardrailsMonitor/LogViewer";
|
||||
import type { LogEntry } from "@/components/GuardrailsMonitor/mockData";
|
||||
import { PolicySelect } from "@/components/ToolPolicies/PolicySelect";
|
||||
import {
|
||||
deleteToolPolicyOverride,
|
||||
fetchToolDetail,
|
||||
getToolUsageLogs,
|
||||
keyListCall,
|
||||
teamListCall,
|
||||
updateToolPolicy,
|
||||
type ToolPolicyOverrideRow,
|
||||
} from "@/components/networking";
|
||||
import type { Team } from "@/components/key_team_helpers/key_list";
|
||||
|
||||
interface ToolDetailProps {
|
||||
toolName: string;
|
||||
onBack: () => void;
|
||||
accessToken: string | null;
|
||||
}
|
||||
|
||||
interface TeamOption {
|
||||
team_id: string;
|
||||
team_alias?: string;
|
||||
}
|
||||
|
||||
interface KeyOption {
|
||||
token: string;
|
||||
key_alias?: string;
|
||||
}
|
||||
|
||||
const TOOL_DETAIL_QUERY_KEY = "tool-detail";
|
||||
|
||||
const LOGS_PAGE_SIZE = 50;
|
||||
|
||||
function getDefaultLogsDateRange(): { start: string; end: string } {
|
||||
const end = new Date();
|
||||
const start = new Date();
|
||||
start.setDate(start.getDate() - 90);
|
||||
const fmt = (d: Date) =>
|
||||
d.toISOString().slice(0, 19).replace("T", " ");
|
||||
return { start: fmt(start), end: fmt(end) };
|
||||
}
|
||||
|
||||
export function ToolDetail({ toolName, onBack, accessToken }: ToolDetailProps) {
|
||||
const queryClient = useQueryClient();
|
||||
const [overrideSaving, setOverrideSaving] = useState(false);
|
||||
const [policySaving, setPolicySaving] = useState(false);
|
||||
const [blockScope, setBlockScope] = useState<"team" | "key">("team");
|
||||
const [blockTeamId, setBlockTeamId] = useState<string | null>(null);
|
||||
const [blockKey, setBlockKey] = useState<KeyOption | null>(null);
|
||||
|
||||
const logsDateRange = useMemo(() => getDefaultLogsDateRange(), []);
|
||||
|
||||
const { data: detail, isLoading: detailLoading, error: detailError } = useQuery({
|
||||
queryKey: [TOOL_DETAIL_QUERY_KEY, toolName],
|
||||
queryFn: () => fetchToolDetail(accessToken!, toolName),
|
||||
enabled: !!accessToken && !!toolName,
|
||||
});
|
||||
|
||||
const { data: teamsData } = useQuery({
|
||||
queryKey: ["teams-list-tool-detail"],
|
||||
queryFn: () => teamListCall(accessToken!, null, null),
|
||||
enabled: !!accessToken,
|
||||
});
|
||||
|
||||
const { data: keysData } = useQuery({
|
||||
queryKey: ["keys-list-tool-detail"],
|
||||
queryFn: () => keyListCall(accessToken!, null, null, null, null, null, 1, 100),
|
||||
enabled: !!accessToken,
|
||||
});
|
||||
|
||||
const { data: logsData, isLoading: logsLoading } = useQuery({
|
||||
queryKey: ["tool-usage-logs", toolName, logsDateRange.start, logsDateRange.end],
|
||||
queryFn: () =>
|
||||
getToolUsageLogs(accessToken!, toolName, {
|
||||
page: 1,
|
||||
pageSize: LOGS_PAGE_SIZE,
|
||||
startDate: logsDateRange.start,
|
||||
endDate: logsDateRange.end,
|
||||
}),
|
||||
enabled: !!accessToken && !!toolName,
|
||||
});
|
||||
|
||||
const logs: LogEntry[] = useMemo(() => {
|
||||
const list = logsData?.logs ?? [];
|
||||
return list.map((l) => ({
|
||||
id: l.id,
|
||||
timestamp: l.timestamp,
|
||||
action: "passed" as const,
|
||||
model: l.model ?? undefined,
|
||||
input_snippet: l.input_snippet ?? undefined,
|
||||
}));
|
||||
}, [logsData?.logs]);
|
||||
|
||||
const teams: Team[] = useMemo(() => {
|
||||
const arr = Array.isArray(teamsData) ? teamsData : teamsData?.data ?? [];
|
||||
return arr.map((t: { team_id?: string; id?: string; team_alias?: string }) => ({
|
||||
team_id: t.team_id ?? t.id ?? "",
|
||||
team_alias: t.team_alias ?? t.team_id ?? "",
|
||||
models: [],
|
||||
max_budget: null,
|
||||
budget_duration: null,
|
||||
tpm_limit: null,
|
||||
rpm_limit: null,
|
||||
organization_id: "",
|
||||
created_at: "",
|
||||
keys: [],
|
||||
members_with_roles: [],
|
||||
spend: 0,
|
||||
}));
|
||||
}, [teamsData]);
|
||||
|
||||
const keys: KeyOption[] = useMemo(() => {
|
||||
const keysRes = keysData?.keys ?? keysData?.data ?? [];
|
||||
return keysRes.map((k: { token?: string; api_key?: string; key_hash?: string; key_alias?: string }) => ({
|
||||
token: k.token ?? k.api_key ?? k.key_hash ?? "",
|
||||
key_alias: k.key_alias ?? (k.token ?? k.api_key ?? k.key_hash)?.toString?.()?.substring?.(0, 8),
|
||||
}));
|
||||
}, [keysData]);
|
||||
|
||||
const invalidateDetail = useCallback(() => {
|
||||
queryClient.invalidateQueries({ queryKey: [TOOL_DETAIL_QUERY_KEY, toolName] });
|
||||
}, [queryClient, toolName]);
|
||||
|
||||
const handlePolicyChange = useCallback(
|
||||
async (name: string, newPolicy: string) => {
|
||||
if (!accessToken) return;
|
||||
setPolicySaving(true);
|
||||
try {
|
||||
await updateToolPolicy(accessToken, name, newPolicy);
|
||||
invalidateDetail();
|
||||
} catch (e: unknown) {
|
||||
alert(`Failed to update policy: ${e instanceof Error ? e.message : String(e)}`);
|
||||
} finally {
|
||||
setPolicySaving(false);
|
||||
}
|
||||
},
|
||||
[accessToken, invalidateDetail]
|
||||
);
|
||||
|
||||
const handleAddOverride = useCallback(async () => {
|
||||
if (!accessToken || !toolName) return;
|
||||
const isTeam = blockScope === "team";
|
||||
if (isTeam && !blockTeamId) return;
|
||||
if (!isTeam && !blockKey?.token) return;
|
||||
setOverrideSaving(true);
|
||||
try {
|
||||
await updateToolPolicy(accessToken, toolName, "blocked", {
|
||||
team_id: isTeam ? blockTeamId : undefined,
|
||||
key_hash: !isTeam ? blockKey!.token : undefined,
|
||||
key_alias: !isTeam ? blockKey!.key_alias : undefined,
|
||||
});
|
||||
invalidateDetail();
|
||||
setBlockTeamId(null);
|
||||
setBlockKey(null);
|
||||
} catch (e: unknown) {
|
||||
alert(`Failed to add override: ${e instanceof Error ? e.message : String(e)}`);
|
||||
} finally {
|
||||
setOverrideSaving(false);
|
||||
}
|
||||
}, [accessToken, toolName, blockScope, blockTeamId, blockKey, invalidateDetail]);
|
||||
|
||||
const handleRemoveOverride = useCallback(
|
||||
async (override: ToolPolicyOverrideRow) => {
|
||||
if (!accessToken || !toolName) return;
|
||||
setOverrideSaving(true);
|
||||
try {
|
||||
await deleteToolPolicyOverride(accessToken, toolName, {
|
||||
team_id: override.team_id ?? undefined,
|
||||
key_hash: override.key_hash ?? undefined,
|
||||
});
|
||||
invalidateDetail();
|
||||
} catch (e: unknown) {
|
||||
alert(`Failed to remove override: ${e instanceof Error ? e.message : String(e)}`);
|
||||
} finally {
|
||||
setOverrideSaving(false);
|
||||
}
|
||||
},
|
||||
[accessToken, toolName, invalidateDetail]
|
||||
);
|
||||
|
||||
if (detailLoading && !detail) {
|
||||
return (
|
||||
<div className="flex items-center justify-center py-12">
|
||||
<Spin size="large" />
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
if (detailError && !detail) {
|
||||
return (
|
||||
<div>
|
||||
<Button type="link" icon={<ArrowLeftOutlined />} onClick={onBack} className="pl-0 mb-4">
|
||||
Back to Tool Policies
|
||||
</Button>
|
||||
<p className="text-red-600">Failed to load tool details.</p>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
if (!detail) {
|
||||
return null;
|
||||
}
|
||||
|
||||
const { tool, overrides } = detail;
|
||||
|
||||
return (
|
||||
<div>
|
||||
<div className="mb-6">
|
||||
<Button
|
||||
type="link"
|
||||
icon={<ArrowLeftOutlined />}
|
||||
onClick={onBack}
|
||||
className="pl-0 mb-4"
|
||||
>
|
||||
Back to Tool Policies
|
||||
</Button>
|
||||
|
||||
<div className="flex items-start justify-between">
|
||||
<div>
|
||||
<div className="flex items-center gap-3 mb-1 flex-wrap">
|
||||
<ToolOutlined className="text-xl text-gray-400" />
|
||||
<h1 className="text-xl font-semibold text-gray-900 font-mono">{tool.tool_name}</h1>
|
||||
<span className="inline-flex items-center px-2.5 py-1 text-xs font-medium rounded-md bg-gray-100 text-gray-700 border border-gray-200">
|
||||
{tool.origin ?? "—"}
|
||||
</span>
|
||||
<span className="inline-flex items-center px-2.5 py-1 text-xs font-medium rounded-md bg-indigo-50 text-indigo-700 border border-indigo-200">
|
||||
{(tool.call_count ?? 0).toLocaleString()} calls
|
||||
</span>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div className="space-y-6">
|
||||
<section className="bg-white rounded-lg border border-gray-200 p-5 shadow-sm">
|
||||
<h2 className="text-sm font-semibold text-gray-700 mb-3">Global policy</h2>
|
||||
<PolicySelect
|
||||
value={tool.call_policy}
|
||||
toolName={tool.tool_name}
|
||||
saving={policySaving}
|
||||
onChange={handlePolicyChange}
|
||||
size="middle"
|
||||
minWidth={140}
|
||||
stopPropagation={false}
|
||||
/>
|
||||
</section>
|
||||
|
||||
{overrides.length > 0 && (
|
||||
<section className="bg-white rounded-lg border border-gray-200 p-5 shadow-sm">
|
||||
<h2 className="text-sm font-semibold text-gray-700 mb-3">Blocked for team or key</h2>
|
||||
<ul className="border rounded-md divide-y divide-gray-100 bg-red-50/30">
|
||||
{overrides.map((ov) => (
|
||||
<li
|
||||
key={ov.override_id}
|
||||
className="flex items-center justify-between px-3 py-2.5 text-sm"
|
||||
>
|
||||
<span className="text-gray-700">
|
||||
{ov.team_id ? `Team: ${ov.team_id}` : ""}
|
||||
{ov.team_id && ov.key_hash ? " · " : ""}
|
||||
{ov.key_hash ? `Key: ${ov.key_alias || ov.key_hash.substring(0, 8)}` : ""}
|
||||
{!ov.team_id && !ov.key_hash ? "—" : ""}
|
||||
</span>
|
||||
<Button
|
||||
type="link"
|
||||
danger
|
||||
size="small"
|
||||
disabled={overrideSaving}
|
||||
onClick={() => handleRemoveOverride(ov)}
|
||||
>
|
||||
Remove
|
||||
</Button>
|
||||
</li>
|
||||
))}
|
||||
</ul>
|
||||
</section>
|
||||
)}
|
||||
|
||||
<section className="bg-white rounded-lg border border-gray-200 p-5 shadow-sm">
|
||||
<h2 className="text-sm font-semibold text-gray-700 mb-3">Block for team or key</h2>
|
||||
<div className="flex flex-col gap-4 max-w-md">
|
||||
<div>
|
||||
<span className="text-sm font-medium text-gray-700 block mb-2">Scope</span>
|
||||
<div className="flex items-center gap-6">
|
||||
<label className="flex items-center gap-2 cursor-pointer text-sm text-gray-700">
|
||||
<input
|
||||
type="radio"
|
||||
checked={blockScope === "team"}
|
||||
onChange={() => setBlockScope("team")}
|
||||
className="align-middle"
|
||||
/>
|
||||
Team
|
||||
</label>
|
||||
<label className="flex items-center gap-2 cursor-pointer text-sm text-gray-700">
|
||||
<input
|
||||
type="radio"
|
||||
checked={blockScope === "key"}
|
||||
onChange={() => setBlockScope("key")}
|
||||
className="align-middle"
|
||||
/>
|
||||
Key
|
||||
</label>
|
||||
</div>
|
||||
</div>
|
||||
<div>
|
||||
<span className="text-sm font-medium text-gray-700 block mb-2">
|
||||
{blockScope === "team" ? "Team" : "Key"}
|
||||
</span>
|
||||
{blockScope === "team" ? (
|
||||
<TeamDropdown
|
||||
teams={teams}
|
||||
value={blockTeamId ?? undefined}
|
||||
onChange={(id) => setBlockTeamId(id || null)}
|
||||
/>
|
||||
) : (
|
||||
<Select
|
||||
placeholder="Select key"
|
||||
allowClear
|
||||
showSearch
|
||||
optionFilterProp="label"
|
||||
value={blockKey ? blockKey.token : undefined}
|
||||
onChange={(token) => {
|
||||
const k = keys.find((x) => x.token === token);
|
||||
setBlockKey(k ?? null);
|
||||
}}
|
||||
options={keys.map((k) => ({
|
||||
value: k.token,
|
||||
label: k.key_alias || k.token?.substring?.(0, 12) || k.token,
|
||||
}))}
|
||||
className="w-full"
|
||||
style={{ minWidth: 200 }}
|
||||
/>
|
||||
)}
|
||||
</div>
|
||||
<Button
|
||||
type="primary"
|
||||
danger
|
||||
disabled={overrideSaving || (blockScope === "team" ? !blockTeamId : !blockKey?.token)}
|
||||
loading={overrideSaving}
|
||||
onClick={handleAddOverride}
|
||||
>
|
||||
Block for {blockScope}
|
||||
</Button>
|
||||
</div>
|
||||
</section>
|
||||
|
||||
<section className="bg-white rounded-lg border border-gray-200 p-5 shadow-sm">
|
||||
<h2 className="text-sm font-semibold text-gray-700 mb-3 flex items-center gap-2">
|
||||
<HistoryOutlined />
|
||||
Recent logs
|
||||
</h2>
|
||||
<LogViewer
|
||||
guardrailName={tool.tool_name}
|
||||
filterAction="passed"
|
||||
logs={logs}
|
||||
logsLoading={logsLoading}
|
||||
totalLogs={logsData?.total ?? 0}
|
||||
accessToken={accessToken}
|
||||
startDate={logsDateRange.start}
|
||||
endDate={logsDateRange.end}
|
||||
/>
|
||||
</section>
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
|
@ -1,22 +1,41 @@
|
|||
"use client";
|
||||
|
||||
import React, { useCallback, useDeferredValue, useEffect, useMemo, useState } from "react";
|
||||
import { Button, Switch, Tooltip } from "antd";
|
||||
import { Table, TableBody, TableCell, TableHead, TableHeaderCell, TableRow } from "@tremor/react";
|
||||
import { Select, Switch, Tooltip } from "antd";
|
||||
import React, { useCallback, useDeferredValue, useEffect, useState } from "react";
|
||||
import type { SortState } from "./common_components/TableHeaderSortDropdown/TableHeaderSortDropdown";
|
||||
import { TableHeaderSortDropdown } from "./common_components/TableHeaderSortDropdown/TableHeaderSortDropdown";
|
||||
import FilterComponent, { FilterOption } from "./molecules/filter";
|
||||
import { fetchToolsList, ToolRow, updateToolPolicy } from "./networking";
|
||||
import { MetricCard } from "./GuardrailsMonitor/MetricCard";
|
||||
import { PolicySelect, POLICY_OPTIONS } from "./ToolPolicies/PolicySelect";
|
||||
import { fetchToolsList, updateToolPolicy, ToolRow } from "./networking";
|
||||
import { TimeCell } from "./view_logs/time_cell";
|
||||
|
||||
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;
|
||||
// --- Date helpers (UTC) for "new tools" counts ---
|
||||
function getUTCDateKey(date: Date): string {
|
||||
return `${date.getUTCFullYear()}-${String(date.getUTCMonth() + 1).padStart(2, "0")}-${String(date.getUTCDate()).padStart(2, "0")}`;
|
||||
}
|
||||
|
||||
type PolicyValue = "trusted" | "blocked";
|
||||
function isCreatedInUTCDay(createdAt: string | undefined, utcDateKey: string): boolean {
|
||||
if (!createdAt) return false;
|
||||
try {
|
||||
const d = new Date(createdAt);
|
||||
return getUTCDateKey(d) === utcDateKey;
|
||||
} catch {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
const policyStyle = (p: string) => POLICY_OPTIONS.find((o) => o.value === p) ?? POLICY_OPTIONS[1];
|
||||
function countToolsInUTCDay(tools: ToolRow[], utcDateKey: string): number {
|
||||
return tools.filter((t) => isCreatedInUTCDay(t.created_at, utcDateKey)).length;
|
||||
}
|
||||
|
||||
function getTrendSubtitle(newToday: number, newYesterday: number): string | undefined {
|
||||
const diff = newToday - newYesterday;
|
||||
if (diff === 0) return undefined;
|
||||
if (diff > 0) return `+${diff} since yesterday`;
|
||||
return `${diff} since yesterday`;
|
||||
}
|
||||
|
||||
type SortField = "tool_name" | "call_policy" | "team_id" | "key_alias" | "created_at" | "call_count";
|
||||
|
||||
|
|
@ -27,60 +46,10 @@ interface FilterValues {
|
|||
interface ToolPoliciesProps {
|
||||
accessToken: string | null;
|
||||
userRole?: string;
|
||||
onSelectTool?: (toolName: string) => void;
|
||||
}
|
||||
|
||||
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,
|
||||
}}
|
||||
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 }) => {
|
||||
export const ToolPolicies: React.FC<ToolPoliciesProps> = ({ accessToken, onSelectTool }) => {
|
||||
const [tools, setTools] = useState<ToolRow[]>([]);
|
||||
const [loading, setLoading] = useState(true);
|
||||
const [isFetching, setIsFetching] = useState(false);
|
||||
|
|
@ -185,6 +154,41 @@ export const ToolPolicies: React.FC<ToolPoliciesProps> = ({ accessToken }) => {
|
|||
},
|
||||
];
|
||||
|
||||
// Derived counts for summary cards and "Needs Review" (UTC today/yesterday)
|
||||
const { newToday, newYesterday, trendSubtitle, totalTools, blockedCount, activeTeamsCount, needsReviewTools } =
|
||||
useMemo(() => {
|
||||
const now = new Date();
|
||||
const todayKey = getUTCDateKey(now);
|
||||
const yesterday = new Date(now);
|
||||
yesterday.setUTCDate(yesterday.getUTCDate() - 1);
|
||||
const yesterdayKey = getUTCDateKey(yesterday);
|
||||
|
||||
const newToday = countToolsInUTCDay(tools, todayKey);
|
||||
const newYesterday = countToolsInUTCDay(tools, yesterdayKey);
|
||||
const trendSubtitle = getTrendSubtitle(newToday, newYesterday);
|
||||
|
||||
const totalTools = tools.length;
|
||||
const blockedCount = tools.filter((t) => t.call_policy === "blocked").length;
|
||||
const activeTeamsCount = new Set(tools.map((t) => t.team_id).filter(Boolean)).size;
|
||||
|
||||
// New in period (today) and not yet decided — untrusted or dual_llm
|
||||
const needsReviewTools = tools.filter(
|
||||
(t) =>
|
||||
isCreatedInUTCDay(t.created_at, todayKey) &&
|
||||
(t.call_policy === "untrusted" || t.call_policy === "dual_llm")
|
||||
);
|
||||
|
||||
return {
|
||||
newToday,
|
||||
newYesterday,
|
||||
trendSubtitle,
|
||||
totalTools,
|
||||
blockedCount,
|
||||
activeTeamsCount,
|
||||
needsReviewTools,
|
||||
};
|
||||
}, [tools]);
|
||||
|
||||
const SortHeader = ({ label, field }: { label: string; field: SortField }) => (
|
||||
<div className="flex items-center gap-1">
|
||||
<span>{label}</span>
|
||||
|
|
@ -223,9 +227,76 @@ export const ToolPolicies: React.FC<ToolPoliciesProps> = ({ accessToken }) => {
|
|||
const totalPages = Math.max(1, Math.ceil(sorted.length / pageSize));
|
||||
const paginated = sorted.slice((currentPage - 1) * pageSize, currentPage * pageSize);
|
||||
|
||||
const scrollToToolRow = (toolId: string) => {
|
||||
const idx = sorted.findIndex((t) => t.tool_id === toolId);
|
||||
if (idx >= 0) {
|
||||
const page = Math.floor(idx / pageSize) + 1;
|
||||
if (page !== currentPage) setCurrentPage(page);
|
||||
// Scroll after a short delay so the table has re-rendered with the new page
|
||||
requestAnimationFrame(() => {
|
||||
setTimeout(() => {
|
||||
document.getElementById(`tool-row-${toolId}`)?.scrollIntoView({ behavior: "smooth", block: "center" });
|
||||
}, 100);
|
||||
});
|
||||
}
|
||||
};
|
||||
|
||||
return (
|
||||
<div className="p-6 w-full">
|
||||
<div className="w-full">
|
||||
<h1 className="text-2xl font-semibold text-gray-900 mb-6">Tool Policies</h1>
|
||||
|
||||
{/* Summary cards */}
|
||||
<div className="grid grid-cols-2 lg:grid-cols-4 gap-4 mb-6">
|
||||
<MetricCard
|
||||
label="New Today"
|
||||
value={newToday}
|
||||
valueColor="text-green-600"
|
||||
subtitle={trendSubtitle}
|
||||
icon={
|
||||
<svg className="w-4 h-4 text-green-500" fill="none" stroke="currentColor" viewBox="0 0 24 24">
|
||||
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M13 7h8m0 0v8m0-8l-8 8-4-4-6 6" />
|
||||
</svg>
|
||||
}
|
||||
/>
|
||||
<MetricCard label="Total Tools Discovered" value={totalTools} />
|
||||
<MetricCard
|
||||
label="Blocked Tools"
|
||||
value={blockedCount}
|
||||
valueColor={blockedCount > 0 ? "text-red-600" : undefined}
|
||||
/>
|
||||
<MetricCard label="Active Teams" value={activeTeamsCount > 0 ? activeTeamsCount : "—"} />
|
||||
</div>
|
||||
|
||||
{/* Needs Review */}
|
||||
{needsReviewTools.length > 0 && (
|
||||
<div className="bg-amber-50 border border-amber-200 rounded-lg p-4 mb-6">
|
||||
<h2 className="text-sm font-semibold text-amber-900 mb-1">Needs Review</h2>
|
||||
<p className="text-sm text-amber-800 mb-3">
|
||||
{needsReviewTools.length} new tool{needsReviewTools.length !== 1 ? "s" : ""} discovered that require
|
||||
policy decisions.
|
||||
</p>
|
||||
<div className="flex flex-wrap gap-2">
|
||||
{needsReviewTools.map((t) => (
|
||||
<span
|
||||
key={t.tool_id}
|
||||
className="inline-flex items-center gap-2 px-3 py-1.5 bg-white border border-amber-200 rounded-md text-sm"
|
||||
>
|
||||
<span className="font-mono text-amber-900 truncate max-w-[200px]" title={t.tool_name}>
|
||||
{t.tool_name}
|
||||
</span>
|
||||
<button
|
||||
type="button"
|
||||
onClick={() => scrollToToolRow(t.tool_id)}
|
||||
className="text-amber-700 hover:text-amber-900 font-medium text-xs whitespace-nowrap"
|
||||
>
|
||||
Review
|
||||
</button>
|
||||
</span>
|
||||
))}
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
<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">
|
||||
|
|
@ -377,16 +448,20 @@ export const ToolPolicies: React.FC<ToolPoliciesProps> = ({ accessToken }) => {
|
|||
</TableRow>
|
||||
) : (
|
||||
paginated.map((tool) => (
|
||||
<TableRow key={tool.tool_id} className="h-8 hover:bg-gray-50">
|
||||
<TableRow key={tool.tool_id} id={`tool-row-${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>
|
||||
<button
|
||||
type="button"
|
||||
onClick={() => onSelectTool?.(tool.tool_name)}
|
||||
className="text-left w-full font-mono text-xs max-w-[20ch] truncate block font-medium text-blue-600 hover:text-blue-800 hover:underline focus:outline-none focus:ring-0"
|
||||
>
|
||||
<Tooltip title={onSelectTool ? "Click to view details and block for team/key" : tool.tool_name}>
|
||||
<span>{tool.tool_name}</span>
|
||||
</Tooltip>
|
||||
</button>
|
||||
</TableCell>
|
||||
<TableCell className="py-0.5 max-h-8">
|
||||
<PolicySelect
|
||||
|
|
@ -396,8 +471,10 @@ export const ToolPolicies: React.FC<ToolPoliciesProps> = ({ accessToken }) => {
|
|||
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 className="py-0.5 max-h-8">
|
||||
<div className="flex items-center justify-end h-8 tabular-nums text-sm font-mono text-gray-700">
|
||||
{(tool.call_count ?? 0).toLocaleString()}
|
||||
</div>
|
||||
</TableCell>
|
||||
<TableCell className="py-0.5 max-h-8 overflow-hidden whitespace-nowrap">
|
||||
<Tooltip title={tool.team_id ?? "-"}>
|
||||
|
|
@ -453,6 +530,7 @@ export const ToolPolicies: React.FC<ToolPoliciesProps> = ({ accessToken }) => {
|
|||
</div>
|
||||
)}
|
||||
</div>
|
||||
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
|
|
|||
|
|
@ -0,0 +1,90 @@
|
|||
"use client";
|
||||
|
||||
import React from "react";
|
||||
import { Select } from "antd";
|
||||
|
||||
// DB policy values: "trusted" | "untrusted" | "dual_llm" | "blocked" — we expose all except dual_llm in the Policy dropdown
|
||||
export const POLICY_OPTIONS = [
|
||||
{ value: "trusted", label: "trusted", color: "#065f46", bg: "#d1fae5", border: "#6ee7b7" },
|
||||
{ value: "untrusted", label: "untrusted", color: "#92400e", bg: "#fef3c7", border: "#fcd34d" },
|
||||
{ value: "blocked", label: "blocked", color: "#991b1b", bg: "#fee2e2", border: "#fca5a5" },
|
||||
] as const;
|
||||
|
||||
export const policyStyle = (p: string) =>
|
||||
POLICY_OPTIONS.find((o) => o.value === p) ?? POLICY_OPTIONS[1];
|
||||
|
||||
export interface PolicySelectProps {
|
||||
value: string;
|
||||
toolName: string;
|
||||
saving: boolean;
|
||||
onChange: (toolName: string, policy: string) => void;
|
||||
size?: "small" | "middle";
|
||||
minWidth?: number;
|
||||
stopPropagation?: boolean;
|
||||
}
|
||||
|
||||
export const PolicySelect: React.FC<PolicySelectProps> = ({
|
||||
value,
|
||||
toolName,
|
||||
saving,
|
||||
onChange,
|
||||
size = "small",
|
||||
minWidth = 110,
|
||||
stopPropagation = true,
|
||||
}) => {
|
||||
const style = policyStyle(value);
|
||||
return (
|
||||
<Select
|
||||
size={size}
|
||||
value={value}
|
||||
disabled={saving}
|
||||
loading={saving}
|
||||
onChange={(v) => onChange(toolName, v)}
|
||||
onClick={(e) => stopPropagation && e.stopPropagation()}
|
||||
style={{
|
||||
minWidth,
|
||||
fontWeight: 500,
|
||||
}}
|
||||
styles={{
|
||||
selector: {
|
||||
backgroundColor: style.bg,
|
||||
borderColor: style.border,
|
||||
color: style.color,
|
||||
borderRadius: 999,
|
||||
fontSize: size === "small" ? 11 : 12,
|
||||
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>
|
||||
),
|
||||
}))}
|
||||
/>
|
||||
);
|
||||
};
|
||||
44
ui/litellm-dashboard/src/components/ToolPoliciesView.tsx
Normal file
44
ui/litellm-dashboard/src/components/ToolPoliciesView.tsx
Normal file
|
|
@ -0,0 +1,44 @@
|
|||
"use client";
|
||||
|
||||
import React, { useState } from "react";
|
||||
import { ToolDetail } from "@/components/ToolDetail";
|
||||
import { ToolPolicies } from "@/components/ToolPolicies";
|
||||
|
||||
type View =
|
||||
| { type: "overview" }
|
||||
| { type: "detail"; toolName: string };
|
||||
|
||||
interface ToolPoliciesViewProps {
|
||||
accessToken: string | null;
|
||||
userRole?: string;
|
||||
}
|
||||
|
||||
export default function ToolPoliciesView({ accessToken, userRole }: ToolPoliciesViewProps) {
|
||||
const [view, setView] = useState<View>({ type: "overview" });
|
||||
|
||||
const handleSelectTool = (toolName: string) => {
|
||||
setView({ type: "detail", toolName });
|
||||
};
|
||||
|
||||
const handleBack = () => {
|
||||
setView({ type: "overview" });
|
||||
};
|
||||
|
||||
return (
|
||||
<div className="p-6 w-full min-w-0 flex-1">
|
||||
{view.type === "detail" ? (
|
||||
<ToolDetail
|
||||
toolName={view.toolName}
|
||||
onBack={handleBack}
|
||||
accessToken={accessToken}
|
||||
/>
|
||||
) : (
|
||||
<ToolPolicies
|
||||
accessToken={accessToken}
|
||||
userRole={userRole}
|
||||
onSelectTool={handleSelectTool}
|
||||
/>
|
||||
)}
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
|
@ -10019,19 +10019,136 @@ export const fetchToolsList = async (accessToken: string): Promise<ToolRow[]> =>
|
|||
return data.tools ?? [];
|
||||
};
|
||||
|
||||
export const updateToolPolicy = async (
|
||||
export interface ToolPolicyOverrideRow {
|
||||
override_id: string;
|
||||
tool_name: string;
|
||||
team_id?: string | null;
|
||||
key_hash?: string | null;
|
||||
call_policy: string;
|
||||
key_alias?: string | null;
|
||||
created_at?: string;
|
||||
updated_at?: string;
|
||||
}
|
||||
|
||||
export interface ToolDetailResponse {
|
||||
tool: ToolRow;
|
||||
overrides: ToolPolicyOverrideRow[];
|
||||
}
|
||||
|
||||
export interface ToolUsageLogEntry {
|
||||
id: string;
|
||||
timestamp: string;
|
||||
model?: string | null;
|
||||
spend?: number | null;
|
||||
total_tokens?: number | null;
|
||||
input_snippet?: string | null;
|
||||
}
|
||||
|
||||
export interface ToolUsageLogsResponse {
|
||||
logs: ToolUsageLogEntry[];
|
||||
total: number;
|
||||
page: number;
|
||||
page_size: number;
|
||||
}
|
||||
|
||||
export const getToolUsageLogs = 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",
|
||||
options: { page?: number; pageSize?: number; startDate?: string; endDate?: string }
|
||||
): Promise<ToolUsageLogsResponse> => {
|
||||
const encoded = encodeURIComponent(toolName);
|
||||
const url = proxyBaseUrl
|
||||
? `${proxyBaseUrl}/v1/tool/${encoded}/logs`
|
||||
: `/v1/tool/${encoded}/logs`;
|
||||
const params = new URLSearchParams();
|
||||
if (options.page != null) params.append("page", String(options.page));
|
||||
if (options.pageSize != null) params.append("page_size", String(options.pageSize));
|
||||
if (options.startDate) params.append("start_date", options.startDate);
|
||||
if (options.endDate) params.append("end_date", options.endDate);
|
||||
const fullUrl = params.toString() ? `${url}?${params.toString()}` : url;
|
||||
const response = await fetch(fullUrl, {
|
||||
method: "GET",
|
||||
headers: {
|
||||
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
});
|
||||
if (!response.ok) {
|
||||
const errorData = await response.json().catch(() => ({}));
|
||||
throw new Error(deriveErrorMessage(errorData));
|
||||
}
|
||||
return response.json();
|
||||
};
|
||||
|
||||
export const fetchToolDetail = async (
|
||||
accessToken: string,
|
||||
toolName: string
|
||||
): Promise<ToolDetailResponse> => {
|
||||
const encoded = encodeURIComponent(toolName);
|
||||
const url = proxyBaseUrl
|
||||
? `${proxyBaseUrl}/v1/tool/${encoded}/detail`
|
||||
: `/v1/tool/${encoded}/detail`;
|
||||
const response = await fetch(url, {
|
||||
method: "GET",
|
||||
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();
|
||||
};
|
||||
|
||||
export const updateToolPolicy = async (
|
||||
accessToken: string,
|
||||
toolName: string,
|
||||
callPolicy: string,
|
||||
options?: { team_id?: string | null; key_hash?: string | null; key_alias?: string | null }
|
||||
): Promise<ToolRow & { team_id?: string; key_hash?: string }> => {
|
||||
const url = proxyBaseUrl ? `${proxyBaseUrl}/v1/tool/policy` : `/v1/tool/policy`;
|
||||
const body: Record<string, string | undefined | null> = {
|
||||
tool_name: toolName,
|
||||
call_policy: callPolicy,
|
||||
};
|
||||
if (options?.team_id != null) body.team_id = options.team_id || undefined;
|
||||
if (options?.key_hash != null) body.key_hash = options.key_hash || undefined;
|
||||
if (options?.key_alias != null) body.key_alias = options.key_alias || undefined;
|
||||
const response = await fetch(url, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
body: JSON.stringify(body),
|
||||
});
|
||||
if (!response.ok) {
|
||||
const errorData = await response.text();
|
||||
throw new Error(errorData);
|
||||
}
|
||||
return response.json();
|
||||
};
|
||||
|
||||
export const deleteToolPolicyOverride = async (
|
||||
accessToken: string,
|
||||
toolName: string,
|
||||
params: { team_id?: string | null; key_hash?: string | null }
|
||||
): Promise<{ deleted: boolean; tool_name: string }> => {
|
||||
const encoded = encodeURIComponent(toolName);
|
||||
const q = new URLSearchParams();
|
||||
if (params.team_id != null && params.team_id !== "") q.set("team_id", params.team_id);
|
||||
if (params.key_hash != null && params.key_hash !== "") q.set("key_hash", params.key_hash);
|
||||
const query = q.toString();
|
||||
const url = proxyBaseUrl
|
||||
? `${proxyBaseUrl}/v1/tool/${encoded}/overrides${query ? `?${query}` : ""}`
|
||||
: `/v1/tool/${encoded}/overrides${query ? `?${query}` : ""}`;
|
||||
const response = await fetch(url, {
|
||||
method: "DELETE",
|
||||
headers: {
|
||||
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
|
||||
},
|
||||
});
|
||||
if (!response.ok) {
|
||||
const errorData = await response.text();
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue