mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
4906f2a761
commit
0f20cdd10d
3 changed files with 36 additions and 145 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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[]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue