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

View file

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

View file

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

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