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
@ -493,7 +542,6 @@ class TestCustomGuardrailPassthroughSupport:
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