diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index 6b84d90a327..508c1c94659 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -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: diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 65ea9cb42d8..2735b3780f5 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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"}) diff --git a/litellm/proxy/db/tool_registry_writer.py b/litellm/proxy/db/tool_registry_writer.py index a0dffefdc59..93d17685ccc 100644 --- a/litellm/proxy/db/tool_registry_writer.py +++ b/litellm/proxy/db/tool_registry_writer.py @@ -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", ""), 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 213b248929e..08b28824cc0 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 @@ -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 []