feat: address greptile feedback

This commit is contained in:
Krrish Dholakia 2026-02-17 22:46:42 -08:00
parent cd33993f49
commit e1db4050be
5 changed files with 385 additions and 81 deletions

View file

@ -1,35 +1,19 @@
from datetime import datetime from datetime import datetime
from typing import ( from typing import (TYPE_CHECKING, Any, Dict, List, Literal, Optional, Type,
TYPE_CHECKING, Union, get_args)
Any,
Dict,
List,
Literal,
Optional,
Type,
Union,
get_args,
)
from litellm._logging import verbose_logger from litellm._logging import verbose_logger
from litellm.caching import DualCache from litellm.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.custom_logger import CustomLogger
from litellm.types.guardrails import ( from litellm.types.guardrails import (DynamicGuardrailParams,
DynamicGuardrailParams, GuardrailEventHooks, LitellmParams, Mode)
GuardrailEventHooks,
LitellmParams,
Mode,
)
from litellm.types.llms.openai import AllMessageValues from litellm.types.llms.openai import AllMessageValues
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel from litellm.types.proxy.guardrails.guardrail_hooks.base import \
from litellm.types.utils import ( GuardrailConfigModel
CallTypes, from litellm.types.utils import (CallTypes, GenericGuardrailAPIInputs,
GenericGuardrailAPIInputs, GuardrailStatus, GuardrailTracingDetail,
GuardrailStatus, LLMResponseTypes,
GuardrailTracingDetail, StandardLoggingGuardrailInformation)
LLMResponseTypes,
StandardLoggingGuardrailInformation,
)
try: try:
from fastapi.exceptions import HTTPException from fastapi.exceptions import HTTPException
@ -37,7 +21,8 @@ except ImportError:
HTTPException = None # type: ignore HTTPException = None # type: ignore
if TYPE_CHECKING: 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() dc = DualCache()
@ -398,9 +383,8 @@ class CustomGuardrail(CustomLogger):
if self._event_hook_is_event_type(event_type): if self._event_hook_is_event_type(event_type):
if isinstance(self.event_hook, Mode): if isinstance(self.event_hook, Mode):
try: try:
from litellm_enterprise.integrations.custom_guardrail import ( from litellm_enterprise.integrations.custom_guardrail import \
EnterpriseCustomGuardrailHelper, EnterpriseCustomGuardrailHelper
)
except ImportError: except ImportError:
raise ImportError( raise ImportError(
"Setting tag-based guardrails is only available in litellm-enterprise. You must be a premium user to use this feature." "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): if isinstance(self.event_hook, Mode):
try: try:
from litellm_enterprise.integrations.custom_guardrail import ( from litellm_enterprise.integrations.custom_guardrail import \
EnterpriseCustomGuardrailHelper, EnterpriseCustomGuardrailHelper
)
except ImportError: except ImportError:
raise ImportError( raise ImportError(
"Setting tag-based guardrails is only available in litellm-enterprise. You must be a premium user to use this feature." "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: else:
guardrail_mode = self.event_hook # type: ignore[assignment] guardrail_mode = self.event_hook # type: ignore[assignment]
from litellm.litellm_core_utils.core_helpers import ( from litellm.litellm_core_utils.core_helpers import \
filter_exceptions_from_params, filter_exceptions_from_params
)
# Sanitize the response to ensure it's JSON serializable and free of circular refs # Sanitize the response to ensure it's JSON serializable and free of circular refs
# This prevents RecursionErrors in downstream loggers (Langfuse, Datadog, etc.) # 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). (this was logged previously as an API failure - guardrail_failed_to_respond).
Guardrails signal intentional blocks by raising: 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) - ModifyResponseException (passthrough mode violation)
""" """
@ -679,7 +661,7 @@ class CustomGuardrail(CustomLogger):
if ( if (
HTTPException is not None HTTPException is not None
and isinstance(e, HTTPException) and isinstance(e, HTTPException)
and e.status_code == 400 and e.status_code in (400, 403)
): ):
return True return True
return False return False
@ -795,9 +777,8 @@ class CustomGuardrail(CustomLogger):
): ):
from typing import cast from typing import cast
from litellm.responses.litellm_completion_transformation.transformation import ( from litellm.responses.litellm_completion_transformation.transformation import \
LiteLLMCompletionResponsesConfig, LiteLLMCompletionResponsesConfig
)
input_data = data.get("input") input_data = data.get("input")
if input_data is None: if input_data is None:

View file

@ -2,12 +2,11 @@
Claude Code - Block Expensive API Flags Guardrail Claude Code - Block Expensive API Flags Guardrail
Blocks Anthropic API parameters that trigger feature-specific pricing surcharges Blocks Anthropic API parameters that trigger feature-specific pricing surcharges
(fast mode, inference_geo, extended thinking). Also inherits the hosted tool (speed=fast, inference_geo, thinking.type=enabled). Optionally inherits hosted
type prefixes from hosted_tool_types.yaml so hosted tools are blocked here too. 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 Blocked params are driven by expensive_api_flags.yaml.
hosted_tool_types.yaml via `inherit_from`, following the same pattern as
harmful_child_safety.yaml inherits from harm_toxic_abuse.json.
""" """
import os import os
@ -28,7 +27,6 @@ if TYPE_CHECKING:
_DIR = os.path.dirname(__file__) _DIR = os.path.dirname(__file__)
_FLAGS_YAML = os.path.join(_DIR, "expensive_api_flags.yaml") _FLAGS_YAML = os.path.join(_DIR, "expensive_api_flags.yaml")
_TOOLS_YAML = os.path.join(_DIR, "hosted_tool_types.yaml")
def _load_config() -> dict: def _load_config() -> dict:
@ -54,7 +52,9 @@ def _load_config() -> dict:
_CONFIG: dict = _load_config() _CONFIG: dict = _load_config()
_BLOCKED_PARAMS: List[dict] = _CONFIG.get("blocked_params", []) _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]: 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) return t.startswith(_INHERITED_TOOL_TYPE_PREFIXES)
def _check_param( def _check_param(request_data: dict, param_cfg: dict) -> Optional[str]:
request_data: dict, param_cfg: dict
) -> Optional[str]:
""" """
Return an error message if the param in request_data matches a blocked value. Return an error message if the param in request_data matches a blocked value.
Returns None when the param is not blocked. Returns None when the param is not blocked.
@ -99,7 +97,9 @@ def _check_param(
blocked_values: List[str] = param_cfg.get("blocked_values", []) blocked_values: List[str] = param_cfg.get("blocked_values", [])
if "*" in blocked_values or str(value) in 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 return None
@ -108,9 +108,10 @@ class ClaudeCodeBlockExpensiveFlagsGuardrail(CustomGuardrail):
""" """
Guardrail that blocks expensive Anthropic API flags. Guardrail that blocks expensive Anthropic API flags.
Checks request_data for feature-specific pricing flags (fast mode, Checks request_data for feature-specific pricing flags (speed=fast,
inference_geo, extended thinking) and Anthropic-hosted tools inherited inference_geo, thinking.type=enabled) and optionally hosted tools via
from hosted_tool_types.yaml. Raises HTTP 403 on the first violation. inherited prefixes from expensive_api_flags.yaml. Raises HTTP 403 on
the first violation.
""" """
def __init__(self, **kwargs): 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: if _INHERITED_TOOL_TYPE_PREFIXES:
tools: List[dict] = list(inputs.get("tools") or []) # type: ignore[assignment] tools: List[dict] = list(inputs.get("tools") or []) # type: ignore[assignment]

View file

@ -7,6 +7,55 @@ from litellm.proxy._types import CallTypes, UserAPIKeyAuth
from litellm.types.utils import GuardrailTracingDetail 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: class TestCustomGuardrailDeploymentHook:
@pytest.mark.asyncio @pytest.mark.asyncio
@ -394,21 +443,21 @@ class TestCustomGuardrailPassthroughSupport:
""" """
Test that async_post_call_success_deployment_hook handles raw httpx.Response objects Test that async_post_call_success_deployment_hook handles raw httpx.Response objects
from passthrough endpoints without crashing with TypeError. from passthrough endpoints without crashing with TypeError.
This tests Fix #3: TypeError: TypedDict does not support instance and class checks This tests Fix #3: TypeError: TypedDict does not support instance and class checks
""" """
import httpx import httpx
custom_guardrail = CustomGuardrail() custom_guardrail = CustomGuardrail()
# Mock the async_post_call_success_hook to return None (guardrail didn't modify response) # 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) custom_guardrail.async_post_call_success_hook = AsyncMock(return_value=None)
# Create a mock httpx.Response object (typical passthrough response) # Create a mock httpx.Response object (typical passthrough response)
mock_response = AsyncMock(spec=httpx.Response) mock_response = AsyncMock(spec=httpx.Response)
mock_response.status_code = 200 mock_response.status_code = 200
mock_response.text = "Mock response" mock_response.text = "Mock response"
request_data = { request_data = {
"guardrails": ["test_guardrail"], "guardrails": ["test_guardrail"],
"user_api_key_user_id": "test_user", "user_api_key_user_id": "test_user",
@ -417,14 +466,14 @@ class TestCustomGuardrailPassthroughSupport:
"user_api_key_hash": "test_hash", "user_api_key_hash": "test_hash",
"user_api_key_request_route": "passthrough_route", "user_api_key_request_route": "passthrough_route",
} }
# This should not raise TypeError: TypedDict does not support instance and class checks # This should not raise TypeError: TypedDict does not support instance and class checks
result = await custom_guardrail.async_post_call_success_deployment_hook( result = await custom_guardrail.async_post_call_success_deployment_hook(
request_data=request_data, request_data=request_data,
response=mock_response, response=mock_response,
call_type=CallTypes.allm_passthrough_route, call_type=CallTypes.allm_passthrough_route,
) )
# When result is None, should return the original response # When result is None, should return the original response
assert result == mock_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): 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. 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. This ensures that even if call_type is None (before fix #1), the guardrail doesn't crash.
""" """
custom_guardrail = CustomGuardrail() custom_guardrail = CustomGuardrail()
# Mock the async_post_call_success_hook to return None # Mock the async_post_call_success_hook to return None
custom_guardrail.async_post_call_success_hook = AsyncMock(return_value=None) custom_guardrail.async_post_call_success_hook = AsyncMock(return_value=None)
mock_response = AsyncMock() mock_response = AsyncMock()
request_data = { request_data = {
"guardrails": ["test_guardrail"], "guardrails": ["test_guardrail"],
"user_api_key_user_id": "test_user", "user_api_key_user_id": "test_user",
} }
# Call with None call_type - should not crash # Call with None call_type - should not crash
result = await custom_guardrail.async_post_call_success_deployment_hook( result = await custom_guardrail.async_post_call_success_deployment_hook(
request_data=request_data, request_data=request_data,
response=mock_response, response=mock_response,
call_type=None, call_type=None,
) )
# Should return the original response when result is None # Should return the original response when result is None
assert result == mock_response assert result == mock_response
def test_is_valid_response_type_with_none(self): def test_is_valid_response_type_with_none(self):
""" """
Test _is_valid_response_type helper method correctly identifies None as invalid. 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. This is part of Fix #3: Safely handling TypedDict types that don't support isinstance checks.
""" """
custom_guardrail = CustomGuardrail() custom_guardrail = CustomGuardrail()
# None should be invalid # None should be invalid
assert custom_guardrail._is_valid_response_type(None) is False assert custom_guardrail._is_valid_response_type(None) is False
def test_is_valid_response_type_with_typeddict_error(self): def test_is_valid_response_type_with_typeddict_error(self):
""" """
Test _is_valid_response_type gracefully handles TypeError from TypedDict. Test _is_valid_response_type gracefully handles TypeError from TypedDict.
This tests Fix #3: When isinstance() is called with TypedDict types, it raises TypeError. This tests Fix #3: When isinstance() is called with TypedDict types, it raises TypeError.
The method should catch this and allow the response through. The method should catch this and allow the response through.
""" """
from litellm.types.utils import ModelResponse from litellm.types.utils import ModelResponse
custom_guardrail = CustomGuardrail() custom_guardrail = CustomGuardrail()
# Create a valid LiteLLM response object # Create a valid LiteLLM response object
response = ModelResponse( response = ModelResponse(
id="test-id", id="test-id",
@ -487,13 +536,12 @@ class TestCustomGuardrailPassthroughSupport:
model="test-model", model="test-model",
object="chat.completion", object="chat.completion",
) )
# This should return True (it's a valid response type or TypeError is caught) # This should return True (it's a valid response type or TypeError is caught)
result = custom_guardrail._is_valid_response_type(response) result = custom_guardrail._is_valid_response_type(response)
assert result is True assert result is True
class TestEventTypeLogging: class TestEventTypeLogging:
"""Tests for event_type logging in guardrail information.""" """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 Test that log_guardrail_information decorator correctly infers GuardrailEventHooks.pre_call
from async_pre_call_hook function name. 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 from litellm.types.guardrails import GuardrailEventHooks
class TestGuardrail(CustomGuardrail): class TestGuardrail(CustomGuardrail):
@ -540,7 +589,8 @@ class TestEventTypeLogging:
Test that log_guardrail_information decorator correctly infers GuardrailEventHooks.post_call Test that log_guardrail_information decorator correctly infers GuardrailEventHooks.post_call
from async_post_call_success_hook function name. 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 from litellm.types.guardrails import GuardrailEventHooks
class TestGuardrail(CustomGuardrail): class TestGuardrail(CustomGuardrail):
@ -575,7 +625,8 @@ class TestEventTypeLogging:
Test that log_guardrail_information decorator correctly infers GuardrailEventHooks.during_call Test that log_guardrail_information decorator correctly infers GuardrailEventHooks.during_call
from async_moderation_hook function name. 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 from litellm.types.guardrails import GuardrailEventHooks
class TestGuardrail(CustomGuardrail): class TestGuardrail(CustomGuardrail):
@ -610,7 +661,8 @@ class TestEventTypeLogging:
Test that log_guardrail_information decorator correctly infers GuardrailEventHooks.post_call Test that log_guardrail_information decorator correctly infers GuardrailEventHooks.post_call
from async_post_call_streaming_hook function name. 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 from litellm.types.guardrails import GuardrailEventHooks
class TestGuardrail(CustomGuardrail): class TestGuardrail(CustomGuardrail):
@ -645,7 +697,8 @@ class TestEventTypeLogging:
Test that log_guardrail_information decorator returns None for event_type 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. 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 from litellm.types.guardrails import GuardrailEventHooks
class TestGuardrail(CustomGuardrail): class TestGuardrail(CustomGuardrail):
@ -787,7 +840,9 @@ class TestTracingFieldsPopulation:
guardrail_json_response="blocked", guardrail_json_response="blocked",
request_data=request_data, request_data=request_data,
guardrail_status="guardrail_intervened", 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"] slg_list = request_data["metadata"]["standard_logging_guardrail_information"]

View file

@ -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

View file

@ -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