mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat: working key tool blocking
This commit is contained in:
parent
41d24b0c91
commit
3284df3bfa
4 changed files with 50 additions and 82 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"})
|
||||
|
||||
|
|
|
|||
|
|
@ -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", ""),
|
||||
|
|
|
|||
|
|
@ -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 []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue