fix guardrail status error (#20972)

* fix guardrail status error

* fix function imports
This commit is contained in:
mubashir1osmani
2026-02-12 20:19:13 -08:00
committed by GitHub
parent da31dd19da
commit 1bc90d2db7
2 changed files with 101 additions and 1 deletions
+24 -1
View File
@@ -29,6 +29,7 @@ from litellm.types.utils import (
LLMResponseTypes,
StandardLoggingGuardrailInformation,
)
from fastapi.exceptions import HTTPException
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@@ -648,6 +649,23 @@ class CustomGuardrail(CustomLogger):
)
return response
@staticmethod
def _is_guardrail_intervention(e: Exception) -> bool:
"""
Returns True if the exception represents an intentional guardrail block
(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)
- ModifyResponseException (passthrough mode violation)
"""
if isinstance(e, ModifyResponseException):
return True
if isinstance(e, HTTPException) and e.status_code == 400:
return True
return False
def _process_error(
self,
e: Exception,
@@ -662,6 +680,11 @@ class CustomGuardrail(CustomLogger):
This gets logged on downsteam Langfuse, DataDog, etc.
"""
guardrail_status: GuardrailStatus = (
"guardrail_intervened"
if self._is_guardrail_intervention(e)
else "guardrail_failed_to_respond"
)
# For custom_code_guardrail scenario, log as "deny" instead of full exception
# Check if this is from custom_code_guardrail by checking the class name
guardrail_response: Union[Exception, str] = e
@@ -671,7 +694,7 @@ class CustomGuardrail(CustomLogger):
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_json_response=guardrail_response,
request_data=request_data,
guardrail_status="guardrail_failed_to_respond",
guardrail_status=guardrail_status,
duration=duration,
start_time=start_time,
end_time=end_time,
@@ -1122,6 +1122,83 @@ async def test_model_armor_non_model_response():
assert not guardrail.async_handler.post.called
@pytest.mark.asyncio
async def test_model_armor_guardrail_status_intervened_vs_failed():
"""
regression test for bug where _process_error always set 'guardrail_failed_to_respond'
even for intentional blocks (error 400).
"""
mock_user_api_key_dict = UserAPIKeyAuth()
mock_cache = MagicMock(spec=DualCache)
#1: Blocked content should raise exception and show guardrail status: guardrail_intervened"
guardrail = ModelArmorGuardrail(
template_id="test-template",
project_id="test-project",
location="us-central1",
guardrail_name="model-armor-test",
)
mock_response = AsyncMock()
mock_response.status_code = 200
mock_response.json = AsyncMock(return_value={
"sanitizationResult": {
"filterMatchState": "MATCH_FOUND",
"filterResults": {
"rai": {
"raiFilterResult": {
"matchState": "MATCH_FOUND",
}
}
}
}
})
guardrail._ensure_access_token_async = AsyncMock(return_value=("token", "test-project"))
with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=mock_response)):
request_data = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "bad content"}],
"metadata": {"guardrails": ["model-armor-test"]},
}
with pytest.raises(HTTPException):
await guardrail.async_pre_call_hook(
user_api_key_dict=mock_user_api_key_dict,
cache=mock_cache,
data=request_data,
call_type="completion",
)
info = request_data["metadata"]["standard_logging_guardrail_information"]
assert info[0]["guardrail_status"] == "guardrail_intervened"
#2: if an API error - guardrail status should be guardrail_failed_to_respond"
guardrail2 = ModelArmorGuardrail(
template_id="test-template",
project_id="test-project",
location="us-central1",
guardrail_name="model-armor-test2",
fail_on_error=True,
)
guardrail2._ensure_access_token_async = AsyncMock(side_effect=ConnectionError("timeout"))
request_data2 = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "hello"}],
"metadata": {"guardrails": ["model-armor-test2"]},
}
with pytest.raises(ConnectionError):
await guardrail2.async_pre_call_hook(
user_api_key_dict=mock_user_api_key_dict,
cache=mock_cache,
data=request_data2,
call_type="completion",
)
info2 = request_data2["metadata"]["standard_logging_guardrail_information"]
assert info2[0]["guardrail_status"] == "guardrail_failed_to_respond"
def mock_open(read_data=''):
"""Helper to create a mock file object"""
import io