feat: working key tool blocking

This commit is contained in:
Krrish Dholakia 2026-02-25 23:52:54 -08:00
parent 41d24b0c91
commit 3284df3bfa
4 changed files with 50 additions and 82 deletions

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

@ -1,60 +1,40 @@
import enum
import json
from datetime import datetime
from typing import TYPE_CHECKING, Any, Callable, Dict, List, Literal, Optional, Union
from typing import (TYPE_CHECKING, Any, Callable, Dict, List, Literal,
Optional, Union)
import httpx
from pydantic import (
BaseModel,
ConfigDict,
Field,
Json,
field_validator,
model_validator,
)
from pydantic import (BaseModel, ConfigDict, Field, Json, field_validator,
model_validator)
from typing_extensions import Required, TypedDict
from litellm._uuid import uuid
from litellm.types.integrations.slack_alerting import AlertType
from litellm.types.llms.openai import (
AllMessageValues,
OpenAIFileObject,
ResponsesAPIResponse,
)
from litellm.types.mcp import (
MCPAuth,
MCPAuthType,
MCPCredentials,
MCPTransport,
MCPTransportType,
)
from litellm.types.llms.openai import (AllMessageValues, OpenAIFileObject,
ResponsesAPIResponse)
from litellm.types.mcp import (MCPAuthType, MCPCredentials, MCPTransport,
MCPTransportType)
from litellm.types.mcp_server.mcp_server_manager import MCPInfo
from litellm.types.router import RouterErrors, UpdateRouterConfig
from litellm.types.secret_managers.main import KeyManagementSystem
from litellm.types.utils import (
CallTypes,
CostBreakdown,
EmbeddingResponse,
GenericBudgetConfigType,
ImageResponse,
LiteLLMBatch,
LiteLLMFineTuningJob,
LiteLLMPydanticObjectBase,
ModelResponse,
ProviderField,
StandardCallbackDynamicParams,
StandardLoggingGuardrailInformation,
StandardLoggingMCPToolCall,
StandardLoggingModelInformation,
StandardLoggingPayloadErrorInformation,
StandardLoggingPayloadStatus,
StandardLoggingVectorStoreRequest,
StandardPassThroughResponseObject,
TextCompletionResponse,
)
from litellm.types.utils import (CallTypes, CostBreakdown, EmbeddingResponse,
GenericBudgetConfigType, ImageResponse,
LiteLLMBatch, LiteLLMFineTuningJob,
LiteLLMPydanticObjectBase, ModelResponse,
ProviderField, StandardCallbackDynamicParams,
StandardLoggingGuardrailInformation,
StandardLoggingMCPToolCall,
StandardLoggingModelInformation,
StandardLoggingPayloadErrorInformation,
StandardLoggingPayloadStatus,
StandardLoggingVectorStoreRequest,
StandardPassThroughResponseObject,
TextCompletionResponse)
from litellm.types.videos.main import VideoObject
from .types_utils.utils import get_instance_fn, validate_custom_validate_return_type
from .types_utils.utils import (get_instance_fn,
validate_custom_validate_return_type)
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
@ -2369,7 +2349,8 @@ class UserAPIKeyAuth(
This is used to track number of requests/spend for health check calls.
"""
from litellm.constants import LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME
from litellm.constants import \
LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME
return cls(
api_key=LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME,
@ -2401,7 +2382,8 @@ class UserAPIKeyAuth(
This is used to track actions performed by automated system jobs.
"""
from litellm.constants import LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME
from litellm.constants import \
LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME
return cls(
api_key=LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME,
@ -2792,7 +2774,8 @@ class LiteLLM_AuditLogs(LiteLLMPydanticObjectBase):
@model_validator(mode="after")
def mask_api_keys(self):
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
from litellm.litellm_core_utils.sensitive_data_masker import \
SensitiveDataMasker
masker = SensitiveDataMasker(sensitive_patterns={"key"})

View file

@ -13,11 +13,9 @@ 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
@ -236,10 +234,12 @@ def _override_row_to_model(row: Any) -> ToolPolicyOverrideRow:
"updated_at",
)
}
def _norm(s: Optional[str]) -> Optional[str]:
if s is None or s == _TOOL_OVERRIDE_ANY:
return None
return s or None
return ToolPolicyOverrideRow(
override_id=row.get("override_id", ""),
tool_name=row.get("tool_name", ""),

View file

@ -31,6 +31,8 @@ from litellm._logging import verbose_proxy_logger
from litellm.caching.dual_cache import DualCache
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
@ -41,7 +43,9 @@ if TYPE_CHECKING:
GUARDRAIL_NAME = "tool_policy"
def _get_request_team_and_key(request_data: dict) -> Tuple[Optional[str], Optional[str]]:
def _get_request_team_and_key(
request_data: dict,
) -> Tuple[Optional[str], Optional[str]]:
"""Extract team_id and key hash from request_data (litellm_metadata or metadata)."""
if not request_data:
return None, None
@ -68,13 +72,17 @@ 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]]:
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
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:
@ -122,8 +130,7 @@ class ToolPolicyGuardrail(CustomGuardrail):
if not tool_names:
route = _get_request_route_from_data(request_data)
if route:
from litellm.proxy.guardrails.tool_name_extraction import \
extract_request_tool_names
tool_names = extract_request_tool_names(route, request_data)
else: # response
tool_calls = inputs.get("tool_calls") or []