mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
feat: address greptile feedback
This commit is contained in:
parent
cd33993f49
commit
e1db4050be
5 changed files with 385 additions and 81 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue