diff --git a/litellm/proxy/db/tool_registry_writer.py b/litellm/proxy/db/tool_registry_writer.py index 1c106df675b..5a6512f1374 100644 --- a/litellm/proxy/db/tool_registry_writer.py +++ b/litellm/proxy/db/tool_registry_writer.py @@ -10,18 +10,16 @@ from datetime import datetime, timezone from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union 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.proxy._types import ToolDiscoveryQueueItem -from litellm.types.tool_management import (LiteLLM_ToolTableRow, - ToolCallPolicy, - ToolPolicyOverrideRow) +from litellm.types.tool_management import ( + LiteLLM_ToolTableRow, + ToolCallPolicy, + ToolPolicyOverrideRow, +) if TYPE_CHECKING: from litellm.proxy.utils import PrismaClient -TOOL_POLICY_CACHE_KEY_PREFIX = "tool_policy:" - def _row_to_model(row: Union[dict, Any]) -> LiteLLM_ToolTableRow: """Convert a Prisma model instance or dict to LiteLLM_ToolTableRow.""" @@ -267,67 +265,6 @@ async def list_overrides_for_tool( return [] -async def _get_merged_blocked_tools( - prisma_client: "PrismaClient", - object_permission_id: Optional[str], - team_object_permission_id: Optional[str], -) -> set: - """Return union of blocked_tools from key and team object permissions.""" - blocked: set = set() - for op_id in (object_permission_id, team_object_permission_id): - if not op_id or not op_id.strip(): - continue - try: - row = await prisma_client.db.litellm_objectpermissiontable.find_unique( - where={"object_permission_id": op_id.strip()}, - ) - if row is not None and getattr(row, "blocked_tools", None): - blocked.update(row.blocked_tools) - except Exception as e: - verbose_proxy_logger.debug( - "tool_registry_writer _get_merged_blocked_tools error for %s: %s", - op_id, - e, - ) - return blocked - - -async def get_effective_policies( - prisma_client: "PrismaClient", - 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: if tool is in key/team object permission - blocked_tools then "blocked", otherwise global policy from LiteLLM_ToolTable. - """ - if not tool_names: - return {} - try: - blocked = await _get_merged_blocked_tools( - prisma_client=prisma_client, - object_permission_id=object_permission_id, - team_object_permission_id=team_object_permission_id, - ) - result: Dict[str, str] = {} - for name in tool_names: - if name in blocked: - result[name] = "blocked" - missing = [n for n in tool_names if n not in result] - if missing: - global_map = await get_tools_by_names( - prisma_client=prisma_client, tool_names=missing - ) - result.update(global_map) - return result - except Exception as e: - verbose_proxy_logger.error( - "tool_registry_writer get_effective_policies error: %s", e - ) - return {} - - class ToolPolicyRegistry: """ In-memory registry of tool policies synced from DB. @@ -407,69 +344,6 @@ def get_tool_policy_registry() -> ToolPolicyRegistry: return _tool_policy_registry -def _effective_cache_suffix( - object_permission_id: Optional[str], - team_object_permission_id: Optional[str], -) -> str: - """Cache key suffix so different request contexts get correct policies.""" - return f":{object_permission_id or ''}:{team_object_permission_id or ''}" - - -async def get_tool_policies_cached( - tool_names: List[str], - cache: DualCache, - prisma_client: Optional["PrismaClient"], - object_permission_id: Optional[str] = None, - team_object_permission_id: Optional[str] = None, -) -> Dict[str, str]: - """ - Return effective call_policy per tool (blocked if in object permission blocked_tools, - else global). Cache-first; cache key includes object_permission_id(s). - """ - if not tool_names: - return {} - suffix = _effective_cache_suffix(object_permission_id, team_object_permission_id) - result: Dict[str, str] = {} - cache_misses: List[str] = [] - for name in tool_names: - key = f"{TOOL_POLICY_CACHE_KEY_PREFIX}{name}{suffix}" - cached = await cache.async_get_cache(key=key) - if cached is not None and isinstance(cached, str): - result[name] = cached - else: - cache_misses.append(name) - if cache_misses and prisma_client is not None: - try: - if object_permission_id or team_object_permission_id: - fetched = await get_effective_policies( - prisma_client=prisma_client, - tool_names=cache_misses, - object_permission_id=object_permission_id, - team_object_permission_id=team_object_permission_id, - ) - else: - fetched = await get_tools_by_names( - prisma_client=prisma_client, tool_names=cache_misses - ) - for name, policy in fetched.items(): - result[name] = policy - await cache.async_set_cache( - key=f"{TOOL_POLICY_CACHE_KEY_PREFIX}{name}{suffix}", - value=policy, - ttl=TOOL_POLICY_CACHE_TTL_SECONDS, - ) - verbose_proxy_logger.debug( - "get_tool_policies_cached: fetched %d from DB (hits: %d)", - len(cache_misses), - len(tool_names) - len(cache_misses), - ) - except Exception as e: - verbose_proxy_logger.error( - "tool_registry_writer get_tool_policies_cached error: %s", e - ) - return result - - async def add_tool_to_object_permission_blocked( prisma_client: "PrismaClient", object_permission_id: str, diff --git a/litellm/proxy/management_endpoints/tool_management_endpoints.py b/litellm/proxy/management_endpoints/tool_management_endpoints.py index 90115b13ab3..90c4fff5729 100644 --- a/litellm/proxy/management_endpoints/tool_management_endpoints.py +++ b/litellm/proxy/management_endpoints/tool_management_endpoints.py @@ -19,13 +19,16 @@ if TYPE_CHECKING: from litellm._logging import verbose_proxy_logger from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.types.tool_management import (LiteLLM_ToolTableRow, - ToolCallPolicy, ToolDetailResponse, - ToolListResponse, - ToolPolicyUpdateRequest, - ToolPolicyUpdateResponse, - ToolUsageLogEntry, - ToolUsageLogsResponse) +from litellm.types.tool_management import ( + LiteLLM_ToolTableRow, + ToolCallPolicy, + ToolDetailResponse, + ToolListResponse, + ToolPolicyUpdateRequest, + ToolPolicyUpdateResponse, + ToolUsageLogEntry, + ToolUsageLogsResponse, +) router = APIRouter() @@ -46,8 +49,7 @@ async def list_tools( Parameters: - call_policy: Optional filter — one of "trusted", "untrusted", "dual_llm", "blocked" """ - from litellm.proxy.db.tool_registry_writer import \ - list_tools as db_list_tools + from litellm.proxy.db.tool_registry_writer import list_tools as db_list_tools from litellm.proxy.proxy_server import prisma_client if prisma_client is None: @@ -377,9 +379,12 @@ async def update_tool_policy( """ from litellm.proxy.db.tool_registry_writer import ( add_tool_to_object_permission_blocked, - remove_tool_from_object_permission_blocked) - from litellm.proxy.db.tool_registry_writer import \ - update_tool_policy as db_update_tool_policy + 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, + ) from litellm.proxy.proxy_server import prisma_client if prisma_client is None: @@ -424,6 +429,9 @@ async def update_tool_policy( 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, @@ -442,6 +450,9 @@ async def update_tool_policy( 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, @@ -473,8 +484,10 @@ async def delete_tool_policy_override( 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 \ - remove_tool_from_object_permission_blocked + 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: @@ -515,6 +528,9 @@ async def delete_tool_policy_override( 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 diff --git a/schema.prisma b/schema.prisma index 691883ef446..cd4f9a4d247 100644 --- a/schema.prisma +++ b/schema.prisma @@ -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[]