diff --git a/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py index b6cb4bba138..2d91febd5e4 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py @@ -23,21 +23,21 @@ or both pre and post call: mode: during_call # runs before LLM and on response """ -from typing import TYPE_CHECKING, Any, List, Literal, Optional, Tuple +from typing import TYPE_CHECKING, Any, Literal, Optional, Tuple from fastapi import HTTPException from litellm._logging import verbose_proxy_logger -from litellm.integrations.custom_guardrail import (CustomGuardrail, - log_guardrail_information) -from litellm.proxy.guardrails.tool_name_extraction import \ - extract_request_tool_names +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 if TYPE_CHECKING: - from litellm.litellm_core_utils.litellm_logging import \ - Logging as LiteLLMLoggingObj + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj GUARDRAIL_NAME = "tool_policy" @@ -80,29 +80,12 @@ def _get_request_route_from_data(request_data: dict) -> Optional[str]: return meta.get("user_api_key_request_route") -def _get_effective_allowed_tools_from_request( - request_data: dict, -) -> Optional[List[str]]: - """Key allowed_tools overrides team; empty/missing means no restriction.""" - meta = request_data.get("metadata") or request_data.get("litellm_metadata") or {} - key_meta = meta.get("user_api_key_metadata") or {} - team_meta = meta.get("user_api_key_team_metadata") or {} - key_allowed = key_meta.get("allowed_tools") if isinstance(key_meta, dict) else None - team_allowed = ( - team_meta.get("allowed_tools") if isinstance(team_meta, dict) else None - ) - if isinstance(key_allowed, list) and len(key_allowed) > 0: - return key_allowed - if isinstance(team_allowed, list) and len(team_allowed) > 0: - return team_allowed - return None - - class ToolPolicyGuardrail(CustomGuardrail): """ Guardrail that enforces per-tool call policies from the in-memory - ToolPolicyRegistry (synced from DB). Key/team allowed_tools (allowlist) still - enforced. No DB or cache in hot path — registry lookups only. + 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: @@ -123,7 +106,7 @@ class ToolPolicyGuardrail(CustomGuardrail): logging_obj: Optional["LiteLLMLoggingObj"] = None, ) -> GenericGuardrailAPIInputs: """ - Enforce key/team allowlist then DB call_policy on request tools / response tool_calls. + Enforce DB call_policy on request tools / response tool_calls. """ if input_type == "request": tools = inputs.get("tools") or [] @@ -154,29 +137,10 @@ class ToolPolicyGuardrail(CustomGuardrail): if not tool_names: return inputs - allowed_tools = _get_effective_allowed_tools_from_request(request_data) - if isinstance(allowed_tools, list) and len(allowed_tools) > 0: - allowed_set = {str(t) for t in allowed_tools} - disallowed = [n for n in tool_names if n not in allowed_set] - if disallowed: - verbose_proxy_logger.warning( - "ToolPolicyGuardrail: tool(s) %s not in key/team allowed_tools", - disallowed, - ) - raise HTTPException( - status_code=400, - detail={ - "error": "Violated tool allowlist", - "disallowed_tools": disallowed, - "message": f"Tool(s) {disallowed} are not in the allowed tools list for this key/team.", - }, - ) - 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 + from litellm.proxy.db.tool_registry_writer import get_tool_policy_registry registry = get_tool_policy_registry() if not registry.is_initialized(): diff --git a/litellm/proxy/management_endpoints/tool_management_endpoints.py b/litellm/proxy/management_endpoints/tool_management_endpoints.py index 90c4fff5729..304a973535f 100644 --- a/litellm/proxy/management_endpoints/tool_management_endpoints.py +++ b/litellm/proxy/management_endpoints/tool_management_endpoints.py @@ -8,6 +8,7 @@ GET /v1/tool/{tool_name} - Get a single tool's details POST /v1/tool/policy - Update the call_policy for a tool """ +import uuid from datetime import datetime, timezone from typing import TYPE_CHECKING, Any, List, Optional @@ -310,17 +311,26 @@ async def _resolve_key_hash_to_object_permission_id( op_id = getattr(row, "object_permission_id", None) if op_id: return op_id - # Create new object permission and assign to key - import uuid as _uuid - - new_id = str(_uuid.uuid4()) + # 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": []} ) - await prisma_client.db.litellm_verificationtoken.update( - where={"token": hashed}, + 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 @@ -331,8 +341,9 @@ async def _resolve_team_id_to_object_permission_id( """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.strip()}, + where={"team_id": team_id_clean}, select={"object_permission_id": True}, ) if row is None: @@ -340,16 +351,24 @@ async def _resolve_team_id_to_object_permission_id( op_id = getattr(row, "object_permission_id", None) if op_id: return op_id - import uuid as _uuid - - new_id = str(_uuid.uuid4()) + # 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": []} ) - await prisma_client.db.litellm_teamtable.update( - where={"team_id": team_id.strip()}, + 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