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:
Krrish Dholakia 2026-03-03 13:30:18 -08:00
commit f81152d5e1
34 changed files with 2648 additions and 455 deletions

View file

@ -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.

View file

@ -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

View file

@ -0,0 +1,2 @@
-- AlterTable
ALTER TABLE "LiteLLM_ObjectPermissionTable" ADD COLUMN "blocked_tools" TEXT[] DEFAULT ARRAY[]::TEXT[];

View file

@ -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");

View file

@ -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())

View file

@ -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],

View file

@ -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 []

View file

@ -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],

View file

@ -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]],

View file

@ -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:

View file

@ -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"]

View file

@ -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

View file

@ -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:

View 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
)

View file

@ -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

View file

@ -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

View 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 []

View file

@ -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

View file

@ -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))

View file

@ -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(

View file

@ -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())

View file

@ -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,

View file

@ -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

View file

@ -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())

View 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()

View file

@ -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"}

View file

@ -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()

View 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",
)

View file

@ -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" ? (

View 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>
);
}

View file

@ -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>
);
};

View file

@ -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>
),
}))}
/>
);
};

View 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>
);
}

View file

@ -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();