fix: address PR review feedback for tool access control

- Add missing blocked_tools column to root schema.prisma (schema drift)
- Invalidate ToolPolicyRegistry after policy mutations so changes take effect immediately
- Remove dead code: unused get_effective_policies, get_tool_policies_cached, and helpers

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
Krrish Dholakia 2026-02-26 16:41:04 -08:00
parent 4906f2a761
commit 0f20cdd10d
3 changed files with 36 additions and 145 deletions

View file

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

View file

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

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