mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +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 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:
|
||||||
|
|
|
||||||
|
|
@ -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]
|
||||||
|
|
|
||||||
|
|
@ -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"]
|
||||||
|
|
|
||||||
|
|
@ -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