diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 4a1e3e41e96..64b87669564 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -1,35 +1,19 @@ from datetime import datetime -from typing import ( - TYPE_CHECKING, - Any, - Dict, - List, - Literal, - Optional, - Type, - Union, - get_args, -) +from typing import (TYPE_CHECKING, Any, Dict, List, Literal, Optional, Type, + Union, get_args) from litellm._logging import verbose_logger from litellm.caching import DualCache from litellm.integrations.custom_logger import CustomLogger -from litellm.types.guardrails import ( - DynamicGuardrailParams, - GuardrailEventHooks, - LitellmParams, - Mode, -) +from litellm.types.guardrails import (DynamicGuardrailParams, + GuardrailEventHooks, LitellmParams, Mode) from litellm.types.llms.openai import AllMessageValues -from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel -from litellm.types.utils import ( - CallTypes, - GenericGuardrailAPIInputs, - GuardrailStatus, - GuardrailTracingDetail, - LLMResponseTypes, - StandardLoggingGuardrailInformation, -) +from litellm.types.proxy.guardrails.guardrail_hooks.base import \ + GuardrailConfigModel +from litellm.types.utils import (CallTypes, GenericGuardrailAPIInputs, + GuardrailStatus, GuardrailTracingDetail, + LLMResponseTypes, + StandardLoggingGuardrailInformation) try: from fastapi.exceptions import HTTPException @@ -37,7 +21,8 @@ except ImportError: HTTPException = None # type: ignore 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 dc = DualCache() @@ -398,9 +383,8 @@ class CustomGuardrail(CustomLogger): if self._event_hook_is_event_type(event_type): if isinstance(self.event_hook, Mode): try: - from litellm_enterprise.integrations.custom_guardrail import ( - EnterpriseCustomGuardrailHelper, - ) + from litellm_enterprise.integrations.custom_guardrail import \ + EnterpriseCustomGuardrailHelper except ImportError: raise ImportError( "Setting tag-based guardrails is only available in litellm-enterprise. You must be a premium user to use this feature." @@ -425,9 +409,8 @@ class CustomGuardrail(CustomLogger): if isinstance(self.event_hook, Mode): try: - from litellm_enterprise.integrations.custom_guardrail import ( - EnterpriseCustomGuardrailHelper, - ) + from litellm_enterprise.integrations.custom_guardrail import \ + EnterpriseCustomGuardrailHelper except ImportError: raise ImportError( "Setting tag-based guardrails is only available in litellm-enterprise. You must be a premium user to use this feature." @@ -546,9 +529,8 @@ class CustomGuardrail(CustomLogger): else: guardrail_mode = self.event_hook # type: ignore[assignment] - from litellm.litellm_core_utils.core_helpers import ( - filter_exceptions_from_params, - ) + from litellm.litellm_core_utils.core_helpers import \ + filter_exceptions_from_params # Sanitize the response to ensure it's JSON serializable and free of circular refs # This prevents RecursionErrors in downstream loggers (Langfuse, Datadog, etc.) @@ -670,7 +652,7 @@ class CustomGuardrail(CustomLogger): (this was logged previously as an API failure - guardrail_failed_to_respond). Guardrails signal intentional blocks by raising: - - HTTPException with status 400 (content policy violation) + - HTTPException with status 400 (content policy violation) or 403 (forbidden) - ModifyResponseException (passthrough mode violation) """ @@ -679,7 +661,7 @@ class CustomGuardrail(CustomLogger): if ( HTTPException is not None and isinstance(e, HTTPException) - and e.status_code == 400 + and e.status_code in (400, 403) ): return True return False @@ -795,9 +777,8 @@ class CustomGuardrail(CustomLogger): ): from typing import cast - from litellm.responses.litellm_completion_transformation.transformation import ( - LiteLLMCompletionResponsesConfig, - ) + from litellm.responses.litellm_completion_transformation.transformation import \ + LiteLLMCompletionResponsesConfig input_data = data.get("input") if input_data is None: diff --git a/litellm/proxy/guardrails/guardrail_hooks/claude_code/block_expensive_flags.py b/litellm/proxy/guardrails/guardrail_hooks/claude_code/block_expensive_flags.py index f2945124771..449891c2c40 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/claude_code/block_expensive_flags.py +++ b/litellm/proxy/guardrails/guardrail_hooks/claude_code/block_expensive_flags.py @@ -2,12 +2,11 @@ Claude Code - Block Expensive API Flags Guardrail Blocks Anthropic API parameters that trigger feature-specific pricing surcharges -(fast mode, inference_geo, extended thinking). Also inherits the hosted tool -type prefixes from hosted_tool_types.yaml so hosted tools are blocked here too. +(speed=fast, inference_geo, thinking.type=enabled). Optionally inherits hosted +tool type prefixes from expensive_api_flags.yaml's `inherit_from` (e.g. +block_hosted_tools/anthropic.yaml) so hosted tools can be blocked as well. -Blocked params are driven by expensive_api_flags.yaml which references -hosted_tool_types.yaml via `inherit_from`, following the same pattern as -harmful_child_safety.yaml inherits from harm_toxic_abuse.json. +Blocked params are driven by expensive_api_flags.yaml. """ import os @@ -28,7 +27,6 @@ if TYPE_CHECKING: _DIR = os.path.dirname(__file__) _FLAGS_YAML = os.path.join(_DIR, "expensive_api_flags.yaml") -_TOOLS_YAML = os.path.join(_DIR, "hosted_tool_types.yaml") def _load_config() -> dict: @@ -54,7 +52,9 @@ def _load_config() -> dict: _CONFIG: dict = _load_config() _BLOCKED_PARAMS: List[dict] = _CONFIG.get("blocked_params", []) -_INHERITED_TOOL_TYPE_PREFIXES: tuple = tuple(_CONFIG.get("_inherited_tool_type_prefixes", [])) +_INHERITED_TOOL_TYPE_PREFIXES: tuple = tuple( + _CONFIG.get("_inherited_tool_type_prefixes", []) +) def _tool_type(tool: dict) -> Optional[str]: @@ -73,9 +73,7 @@ def _is_hosted_tool(tool: dict) -> bool: return t.startswith(_INHERITED_TOOL_TYPE_PREFIXES) -def _check_param( - request_data: dict, param_cfg: dict -) -> Optional[str]: +def _check_param(request_data: dict, param_cfg: dict) -> Optional[str]: """ Return an error message if the param in request_data matches a blocked value. Returns None when the param is not blocked. @@ -99,7 +97,9 @@ def _check_param( blocked_values: List[str] = param_cfg.get("blocked_values", []) if "*" in blocked_values or str(value) in blocked_values: - return param_cfg.get("error_message", f"{param} is disabled by your organization's policy") + return param_cfg.get( + "error_message", f"{param} is disabled by your organization's policy" + ) return None @@ -108,9 +108,10 @@ class ClaudeCodeBlockExpensiveFlagsGuardrail(CustomGuardrail): """ Guardrail that blocks expensive Anthropic API flags. - Checks request_data for feature-specific pricing flags (fast mode, - inference_geo, extended thinking) and Anthropic-hosted tools inherited - from hosted_tool_types.yaml. Raises HTTP 403 on the first violation. + Checks request_data for feature-specific pricing flags (speed=fast, + inference_geo, thinking.type=enabled) and optionally hosted tools via + inherited prefixes from expensive_api_flags.yaml. Raises HTTP 403 on + the first violation. """ def __init__(self, **kwargs): @@ -147,7 +148,7 @@ class ClaudeCodeBlockExpensiveFlagsGuardrail(CustomGuardrail): ) # ------------------------------------------------------------------ # - # 2. Check for inherited hosted tools (from hosted_tool_types.yaml) # + # 2. Check for inherited hosted tools (from expensive_api_flags) # # ------------------------------------------------------------------ # if _INHERITED_TOOL_TYPE_PREFIXES: tools: List[dict] = list(inputs.get("tools") or []) # type: ignore[assignment] diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/test_litellm/integrations/test_custom_guardrail.py index ae4082662f9..b5bb61c0c96 100644 --- a/tests/test_litellm/integrations/test_custom_guardrail.py +++ b/tests/test_litellm/integrations/test_custom_guardrail.py @@ -7,6 +7,55 @@ from litellm.proxy._types import CallTypes, UserAPIKeyAuth from litellm.types.utils import GuardrailTracingDetail +class TestIsGuardrailIntervention: + """Tests for CustomGuardrail._is_guardrail_intervention classification.""" + + def test_http_400_returns_true(self): + """HTTP 400 (content policy violation) should be classified as guardrail_intervened.""" + from fastapi import HTTPException + + assert ( + CustomGuardrail._is_guardrail_intervention( + HTTPException(status_code=400, detail="blocked") + ) + is True + ) + + def test_http_403_returns_true(self): + """HTTP 403 (forbidden) should be classified as guardrail_intervened.""" + from fastapi import HTTPException + + assert ( + CustomGuardrail._is_guardrail_intervention( + HTTPException(status_code=403, detail="blocked") + ) + is True + ) + + def test_modify_response_exception_returns_true(self): + """ModifyResponseException should be classified as guardrail_intervened.""" + from litellm.integrations.custom_guardrail import \ + ModifyResponseException + + exc = ModifyResponseException( + message="blocked", + model="test", + request_data={}, + ) + assert CustomGuardrail._is_guardrail_intervention(exc) is True + + def test_http_500_returns_false(self): + """HTTP 500 (server error) should be classified as guardrail_failed_to_respond.""" + from fastapi import HTTPException + + assert ( + CustomGuardrail._is_guardrail_intervention( + HTTPException(status_code=500, detail="internal error") + ) + is False + ) + + class TestCustomGuardrailDeploymentHook: @pytest.mark.asyncio @@ -394,21 +443,21 @@ class TestCustomGuardrailPassthroughSupport: """ Test that async_post_call_success_deployment_hook handles raw httpx.Response objects from passthrough endpoints without crashing with TypeError. - + This tests Fix #3: TypeError: TypedDict does not support instance and class checks """ import httpx custom_guardrail = CustomGuardrail() - + # Mock the async_post_call_success_hook to return None (guardrail didn't modify response) custom_guardrail.async_post_call_success_hook = AsyncMock(return_value=None) - + # Create a mock httpx.Response object (typical passthrough response) mock_response = AsyncMock(spec=httpx.Response) mock_response.status_code = 200 mock_response.text = "Mock response" - + request_data = { "guardrails": ["test_guardrail"], "user_api_key_user_id": "test_user", @@ -417,14 +466,14 @@ class TestCustomGuardrailPassthroughSupport: "user_api_key_hash": "test_hash", "user_api_key_request_route": "passthrough_route", } - + # This should not raise TypeError: TypedDict does not support instance and class checks result = await custom_guardrail.async_post_call_success_deployment_hook( request_data=request_data, response=mock_response, call_type=CallTypes.allm_passthrough_route, ) - + # When result is None, should return the original response assert result == mock_response @@ -432,53 +481,53 @@ class TestCustomGuardrailPassthroughSupport: async def test_async_post_call_success_deployment_hook_with_none_call_type(self): """ Test that async_post_call_success_deployment_hook handles None call_type gracefully. - + This ensures that even if call_type is None (before fix #1), the guardrail doesn't crash. """ custom_guardrail = CustomGuardrail() - + # Mock the async_post_call_success_hook to return None custom_guardrail.async_post_call_success_hook = AsyncMock(return_value=None) - + mock_response = AsyncMock() - + request_data = { "guardrails": ["test_guardrail"], "user_api_key_user_id": "test_user", } - + # Call with None call_type - should not crash result = await custom_guardrail.async_post_call_success_deployment_hook( request_data=request_data, response=mock_response, call_type=None, ) - + # Should return the original response when result is None assert result == mock_response def test_is_valid_response_type_with_none(self): """ Test _is_valid_response_type helper method correctly identifies None as invalid. - + This is part of Fix #3: Safely handling TypedDict types that don't support isinstance checks. """ custom_guardrail = CustomGuardrail() - + # None should be invalid assert custom_guardrail._is_valid_response_type(None) is False def test_is_valid_response_type_with_typeddict_error(self): """ Test _is_valid_response_type gracefully handles TypeError from TypedDict. - + This tests Fix #3: When isinstance() is called with TypedDict types, it raises TypeError. The method should catch this and allow the response through. """ from litellm.types.utils import ModelResponse - + custom_guardrail = CustomGuardrail() - + # Create a valid LiteLLM response object response = ModelResponse( id="test-id", @@ -487,13 +536,12 @@ class TestCustomGuardrailPassthroughSupport: model="test-model", object="chat.completion", ) - + # This should return True (it's a valid response type or TypeError is caught) result = custom_guardrail._is_valid_response_type(response) assert result is True - class TestEventTypeLogging: """Tests for event_type logging in guardrail information.""" @@ -505,7 +553,8 @@ class TestEventTypeLogging: Test that log_guardrail_information decorator correctly infers GuardrailEventHooks.pre_call from async_pre_call_hook function name. """ - from litellm.integrations.custom_guardrail import log_guardrail_information + from litellm.integrations.custom_guardrail import \ + log_guardrail_information from litellm.types.guardrails import GuardrailEventHooks class TestGuardrail(CustomGuardrail): @@ -540,7 +589,8 @@ class TestEventTypeLogging: Test that log_guardrail_information decorator correctly infers GuardrailEventHooks.post_call from async_post_call_success_hook function name. """ - from litellm.integrations.custom_guardrail import log_guardrail_information + from litellm.integrations.custom_guardrail import \ + log_guardrail_information from litellm.types.guardrails import GuardrailEventHooks class TestGuardrail(CustomGuardrail): @@ -575,7 +625,8 @@ class TestEventTypeLogging: Test that log_guardrail_information decorator correctly infers GuardrailEventHooks.during_call from async_moderation_hook function name. """ - from litellm.integrations.custom_guardrail import log_guardrail_information + from litellm.integrations.custom_guardrail import \ + log_guardrail_information from litellm.types.guardrails import GuardrailEventHooks class TestGuardrail(CustomGuardrail): @@ -610,7 +661,8 @@ class TestEventTypeLogging: Test that log_guardrail_information decorator correctly infers GuardrailEventHooks.post_call from async_post_call_streaming_hook function name. """ - from litellm.integrations.custom_guardrail import log_guardrail_information + from litellm.integrations.custom_guardrail import \ + log_guardrail_information from litellm.types.guardrails import GuardrailEventHooks class TestGuardrail(CustomGuardrail): @@ -645,7 +697,8 @@ class TestEventTypeLogging: Test that log_guardrail_information decorator returns None for event_type when function name doesn't match known patterns, and falls back to self.event_hook. """ - from litellm.integrations.custom_guardrail import log_guardrail_information + from litellm.integrations.custom_guardrail import \ + log_guardrail_information from litellm.types.guardrails import GuardrailEventHooks class TestGuardrail(CustomGuardrail): @@ -787,7 +840,9 @@ class TestTracingFieldsPopulation: guardrail_json_response="blocked", request_data=request_data, guardrail_status="guardrail_intervened", - tracing_detail=GuardrailTracingDetail(policy_template="EU AI Act Article 5"), + tracing_detail=GuardrailTracingDetail( + policy_template="EU AI Act Article 5" + ), ) slg_list = request_data["metadata"]["standard_logging_guardrail_information"] diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/block_hosted_tools/test_block_hosted_tools.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/block_hosted_tools/test_block_hosted_tools.py new file mode 100644 index 00000000000..e82f0f9e2e6 --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/block_hosted_tools/test_block_hosted_tools.py @@ -0,0 +1,148 @@ +""" +Tests for the Block Hosted Tools Guardrail. +""" + +import pytest +from fastapi import HTTPException + +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.proxy.guardrails.guardrail_hooks.block_hosted_tools.guardrail import \ + BlockHostedToolsGuardrail + + +class TestBlockHostedToolsGuardrail: + """Test BlockHostedToolsGuardrail.""" + + def test_initialization(self): + """Test that guardrail initializes with pre_call hook.""" + guardrail = BlockHostedToolsGuardrail( + guardrail_name="test-block-hosted-tools" + ) + assert guardrail.guardrail_name == "test-block-hosted-tools" + assert "pre_call" in str(guardrail.supported_event_hooks) + + @pytest.mark.asyncio + async def test_blocks_anthropic_bash_tool(self): + """Test that Anthropic bash_* hosted tool is blocked.""" + guardrail = BlockHostedToolsGuardrail( + guardrail_name="test-block-hosted-tools" + ) + tools = [{"type": "bash_20250124", "name": "run_bash"}] + inputs = {"tools": tools} + request_data = {} + + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + ) + + assert exc_info.value.status_code == 403 + assert "bash" in str(exc_info.value.detail).lower() + assert "disabled" in str(exc_info.value.detail).lower() + + @pytest.mark.asyncio + async def test_blocks_openai_code_interpreter(self): + """Test that OpenAI code_interpreter hosted tool is blocked.""" + guardrail = BlockHostedToolsGuardrail( + guardrail_name="test-block-hosted-tools" + ) + tools = [{"type": "code_interpreter"}] + inputs = {"tools": tools} + request_data = {} + + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + ) + + assert exc_info.value.status_code == 403 + assert "disabled" in str(exc_info.value.detail).lower() + + @pytest.mark.asyncio + async def test_blocks_gemini_google_search_top_level_key(self): + """Test that Gemini native googleSearch top-level key is blocked.""" + guardrail = BlockHostedToolsGuardrail( + guardrail_name="test-block-hosted-tools" + ) + tools = [{"googleSearch": {}}] + inputs = {"tools": tools} + request_data = {} + + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + ) + + assert exc_info.value.status_code == 403 + assert "disabled" in str(exc_info.value.detail).lower() + + @pytest.mark.asyncio + async def test_allows_function_tools(self): + """Test that user-defined function tools pass through.""" + guardrail = BlockHostedToolsGuardrail( + guardrail_name="test-block-hosted-tools" + ) + tools = [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get weather", + }, + } + ] + inputs = {"tools": tools} + request_data = {} + + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + ) + assert result == inputs + + @pytest.mark.asyncio + async def test_allows_empty_tools(self): + """Test that empty tools list passes through.""" + guardrail = BlockHostedToolsGuardrail( + guardrail_name="test-block-hosted-tools" + ) + inputs = {"tools": []} + request_data = {} + + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + ) + assert result == inputs + + @pytest.mark.asyncio + async def test_skips_response_input_type(self): + """Test that response input_type is skipped.""" + guardrail = BlockHostedToolsGuardrail( + guardrail_name="test-block-hosted-tools" + ) + tools = [{"type": "bash_20250124"}] + inputs = {"tools": tools} + request_data = {} + + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="response", + ) + assert result == inputs + + @pytest.mark.asyncio + async def test_http_403_classified_as_guardrail_intervention(self): + """Test that HTTP 403 from guardrail is classified as guardrail_intervened.""" + assert CustomGuardrail._is_guardrail_intervention( + HTTPException(status_code=403, detail="blocked") + ) is True diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/claude_code/test_block_expensive_flags.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/claude_code/test_block_expensive_flags.py new file mode 100644 index 00000000000..3bb060c398b --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/claude_code/test_block_expensive_flags.py @@ -0,0 +1,119 @@ +""" +Tests for the Claude Code Block Expensive Flags Guardrail. +""" + +import pytest +from fastapi import HTTPException + +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.proxy.guardrails.guardrail_hooks.claude_code.block_expensive_flags import ( + ClaudeCodeBlockExpensiveFlagsGuardrail, +) + + +class TestClaudeCodeBlockExpensiveFlagsGuardrail: + """Test ClaudeCodeBlockExpensiveFlagsGuardrail.""" + + def test_initialization(self): + """Test that guardrail initializes with pre_call hook.""" + guardrail = ClaudeCodeBlockExpensiveFlagsGuardrail( + guardrail_name="test-block-expensive" + ) + assert guardrail.guardrail_name == "test-block-expensive" + assert "pre_call" in str(guardrail.supported_event_hooks) + + @pytest.mark.asyncio + async def test_blocks_speed_fast(self): + """Test that speed=fast is blocked.""" + guardrail = ClaudeCodeBlockExpensiveFlagsGuardrail( + guardrail_name="test-block-expensive" + ) + request_data = {"speed": "fast"} + inputs = {"tools": None} + + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + ) + + assert exc_info.value.status_code == 403 + assert "fast" in str(exc_info.value.detail).lower() + + @pytest.mark.asyncio + async def test_blocks_inference_geo(self): + """Test that inference_geo is blocked.""" + guardrail = ClaudeCodeBlockExpensiveFlagsGuardrail( + guardrail_name="test-block-expensive" + ) + request_data = {"inference_geo": "us"} + inputs = {"tools": None} + + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + ) + + assert exc_info.value.status_code == 403 + assert "inference" in str(exc_info.value.detail).lower() + + @pytest.mark.asyncio + async def test_blocks_thinking_type_enabled(self): + """Test that thinking.type=enabled is blocked.""" + guardrail = ClaudeCodeBlockExpensiveFlagsGuardrail( + guardrail_name="test-block-expensive" + ) + request_data = {"thinking": {"type": "enabled"}} + inputs = {"tools": None} + + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + ) + + assert exc_info.value.status_code == 403 + assert "thinking" in str(exc_info.value.detail).lower() + + @pytest.mark.asyncio + async def test_allows_clean_request(self): + """Test that requests without expensive flags pass through.""" + guardrail = ClaudeCodeBlockExpensiveFlagsGuardrail( + guardrail_name="test-block-expensive" + ) + request_data = {"model": "claude-3-5-sonnet"} + inputs = {"tools": None} + + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + ) + assert result == inputs + + @pytest.mark.asyncio + async def test_skips_response_input_type(self): + """Test that response input_type is skipped.""" + guardrail = ClaudeCodeBlockExpensiveFlagsGuardrail( + guardrail_name="test-block-expensive" + ) + request_data = {"speed": "fast"} + inputs = {"tools": None} + + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="response", + ) + assert result == inputs + + @pytest.mark.asyncio + async def test_http_403_classified_as_guardrail_intervention(self): + """Test that HTTP 403 from guardrail is classified as guardrail_intervened.""" + assert CustomGuardrail._is_guardrail_intervention( + HTTPException(status_code=403, detail="blocked") + ) is True