diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py index bb0d0a99b31..777c45bcd83 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py @@ -21,6 +21,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" unreachable_fallback=getattr( litellm_params, "unreachable_fallback", "fail_closed" ), + fail_on_error=getattr(litellm_params, "fail_on_error", True), extra_headers=getattr(litellm_params, "extra_headers", None), guardrail_name=guardrail.get("guardrail_name", ""), event_hook=litellm_params.mode, diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py index 790ee31f2e0..8d3abe2bdc3 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py @@ -27,6 +27,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPIMetadata, GenericGuardrailAPIRequest, GenericGuardrailAPIResponse, + GuardrailToolParam, ) from litellm.types.utils import GenericGuardrailAPIInputs @@ -187,6 +188,7 @@ class GenericGuardrailAPI(CustomGuardrail): api_key: Optional[str] = None, additional_provider_specific_params: Optional[Dict[str, Any]] = None, unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed", + fail_on_error: Optional[bool] = True, extra_headers: Optional[list] = None, **kwargs, ): @@ -223,6 +225,8 @@ class GenericGuardrailAPI(CustomGuardrail): unreachable_fallback ) + self.fail_on_error: bool = True if fail_on_error is None else fail_on_error + # Set supported event hooks if "supported_event_hooks" not in kwargs: kwargs["supported_event_hooks"] = [ @@ -299,7 +303,7 @@ class GenericGuardrailAPI(CustomGuardrail): f" http_status_code={http_status_code}" if http_status_code else "" ) verbose_proxy_logger.critical( - "Generic Guardrail API unreachable (fail-open). Proceeding without guardrail.%s " + "Generic Guardrail API error (fail-open). Proceeding without guardrail.%s " "guardrail_name=%s api_base=%s input_type=%s litellm_call_id=%s litellm_trace_id=%s", status_suffix, getattr(self, "guardrail_name", None), @@ -351,7 +355,10 @@ class GenericGuardrailAPI(CustomGuardrail): logging_obj: Optional["LiteLLMLoggingObj"], is_unreachable: bool = True, ) -> GenericGuardrailAPIInputs: - if is_unreachable and self.unreachable_fallback == "fail_open": + unreachable_fail_open = ( + is_unreachable and self.unreachable_fallback == "fail_open" + ) + if unreachable_fail_open or not self.fail_on_error: http_status_code = getattr( getattr(error, "response", None), "status_code", None ) @@ -432,26 +439,30 @@ class GenericGuardrailAPI(CustomGuardrail): extra_allowlist=extra_allowlist, ) - # Create request payload - guardrail_request = GenericGuardrailAPIRequest( - litellm_call_id=logging_obj.litellm_call_id if logging_obj else None, - litellm_trace_id=logging_obj.litellm_trace_id if logging_obj else None, - texts=texts, - request_data=user_metadata, - request_headers=inbound_headers, - litellm_version=litellm_version, - images=images, - tools=tools, - structured_messages=structured_messages, - tool_calls=tool_calls, - additional_provider_specific_params=additional_params, - input_type=input_type, - model=model, - ) - - headers = self._build_request_headers() - try: + # Create request payload + guardrail_request = GenericGuardrailAPIRequest( + litellm_call_id=logging_obj.litellm_call_id if logging_obj else None, + litellm_trace_id=logging_obj.litellm_trace_id if logging_obj else None, + texts=texts, + request_data=user_metadata, + request_headers=inbound_headers, + litellm_version=litellm_version, + images=images, + tools=( + [GuardrailToolParam.model_validate(t) for t in tools] + if tools + else None + ), + structured_messages=structured_messages, + tool_calls=tool_calls, + additional_provider_specific_params=additional_params, + input_type=input_type, + model=model, + ) + + headers = self._build_request_headers() + # Make the API request # Use mode="json" to ensure all iterables are converted to lists response = await self.async_handler.post( diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index c9623d8595a..6a76dd337fa 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -750,7 +750,12 @@ class BaseLitellmParams( ) fail_on_error: Optional[bool] = Field( default=True, - description="Whether to fail the request if Model Armor encounters an error", + description=( + "Whether to fail the request if the guardrail encounters an error. " + "Implemented by guardrail='model_armor' and 'generic_guardrail_api'. " + "True (default) raises the error. False logs a critical error and lets the request proceed, " + "so only a valid guardrail response can block or modify it." + ), ) additional_provider_specific_params: Optional[Dict[str, Any]] = Field( diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py index 94f219a5fc6..2fe976463c4 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py @@ -1,17 +1,28 @@ from typing import Any, Dict, List, Literal, Optional, Union -from pydantic import BaseModel, Field +from pydantic import BaseModel, ConfigDict, Field from typing_extensions import TYPE_CHECKING, TypedDict from litellm.types.llms.openai import ( AllMessageValues, ChatCompletionToolCallChunk, - ChatCompletionToolParam, ) from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel from litellm.types.utils import ChatCompletionMessageToolCall +class GuardrailToolParam(BaseModel): + """A tool forwarded verbatim to the guardrail for inspection. + + Built-in tools (code_interpreter, file_search, ...) have no ``function`` block + and stash their config in tool-specific keys, so only ``type`` is required and + ``extra="allow"`` preserves the rest instead of stripping it. + """ + + model_config = ConfigDict(extra="allow") + type: str + + class GenericGuardrailAPIMetadata(TypedDict, total=False): user_api_key_hash: Optional[str] user_api_key_alias: Optional[str] @@ -39,6 +50,16 @@ class GenericGuardrailAPIOptionalParams(BaseModel): ), ) + fail_on_error: Optional[bool] = Field( + default=True, + description=( + "Behavior on any guardrail error, not just unreachability. " + "True (default) raises and blocks the request on error. " + "False logs a critical error and allows the request to proceed, so only a valid " + "guardrail response can block or modify it; broader than unreachable_fallback." + ), + ) + class GenericGuardrailAPIConfigModel( GuardrailConfigModel[GenericGuardrailAPIOptionalParams], @@ -65,7 +86,7 @@ class GenericGuardrailAPIRequest(BaseModel): ) structured_messages: Optional[List[AllMessageValues]] = None images: Optional[List[str]] = None - tools: Optional[List[ChatCompletionToolParam]] = None + tools: Optional[List[GuardrailToolParam]] = None texts: Optional[List[str]] = None request_data: GenericGuardrailAPIMetadata request_headers: Optional[Dict[str, str]] = Field( @@ -88,7 +109,7 @@ class GenericGuardrailAPIResponse: texts: Optional[List[str]] images: Optional[List[str]] - tools: Optional[List[ChatCompletionToolParam]] + tools: Optional[List[GuardrailToolParam]] action: str blocked_reason: Optional[str] @@ -98,7 +119,7 @@ class GenericGuardrailAPIResponse: texts: Optional[List[str]] = None, blocked_reason: Optional[str] = None, images: Optional[List[str]] = None, - tools: Optional[List[ChatCompletionToolParam]] = None, + tools: Optional[List[GuardrailToolParam]] = None, ): self.action = action self.blocked_reason = blocked_reason diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py index 6ec793a1bb0..399442a5f71 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py @@ -1021,3 +1021,171 @@ class TestMultimodalSupport: call_args = mock_post.call_args json_payload = call_args.kwargs["json"] assert isinstance(json_payload["structured_messages"], list) + + +class TestToolSupport: + """Test tool handling in guardrail requests""" + + @pytest.mark.asyncio + async def test_builtin_tools_without_function_block_do_not_crash( + self, generic_guardrail + ): + """Built-in tools (code_interpreter, file_search) have no `function` block. + + Regression for a 500 where serializing them raised a Pydantic + ValidationError because the tool schema required `function`. The full + tool, including built-in tool config, must reach the guardrail intact. + """ + tools = [ + {"type": "function", "function": {"name": "get_weather", "parameters": {}}}, + {"type": "code_interpreter"}, + { + "type": "file_search", + "vector_store_ids": ["vs_1"], + "max_num_results": 5, + }, + ] + + mock_response = MagicMock() + mock_response.json.return_value = {"action": "NONE", "texts": ["hi"]} + mock_response.raise_for_status = MagicMock() + + with patch.object( + generic_guardrail.async_handler, "post", return_value=mock_response + ) as mock_post: + await generic_guardrail.apply_guardrail( + inputs={"texts": ["hi"], "tools": tools}, + request_data={}, + input_type="request", + ) + + forwarded_tools = mock_post.call_args.kwargs["json"]["tools"] + + assert forwarded_tools == tools + + +class TestFailOnError: + """Test fail_on_error: complete fail-open on any guardrail error""" + + @pytest.fixture + def fail_open_guardrail(self): + return GenericGuardrailAPI( + api_base="https://api.test.guardrail.com", + guardrail_name="test-fail-open-guardrail", + event_hook="pre_call", + default_on=True, + fail_on_error=False, + ) + + @pytest.mark.asyncio + async def test_endpoint_error_continues_when_fail_on_error_false( + self, fail_open_guardrail + ): + """A non-unreachable endpoint error (HTTP 400) is swallowed and the request proceeds unchanged.""" + error = httpx.HTTPStatusError( + "bad request", request=MagicMock(), response=MagicMock(status_code=400) + ) + with patch.object( + fail_open_guardrail.async_handler, "post", side_effect=error + ): + result = await fail_open_guardrail.apply_guardrail( + inputs={"texts": ["hi"]}, + request_data={}, + input_type="request", + ) + + assert result == {"texts": ["hi"]} + + @pytest.mark.asyncio + async def test_internal_error_continues_without_calling_endpoint( + self, fail_open_guardrail + ): + """An error while building the request (here: invalid input_type) fails open too. + + Proves the request construction runs inside the protected block: the + endpoint is never called, yet the request still proceeds unchanged. + """ + with patch.object(fail_open_guardrail.async_handler, "post") as mock_post: + result = await fail_open_guardrail.apply_guardrail( + inputs={"texts": ["hi"]}, + request_data={}, + input_type="bogus", # type: ignore[arg-type] + ) + + mock_post.assert_not_called() + assert result == {"texts": ["hi"]} + + @pytest.mark.asyncio + async def test_valid_block_still_blocks_when_fail_on_error_false( + self, fail_open_guardrail + ): + """Only a valid response acts: a BLOCKED decision still raises even with fail_on_error=False.""" + mock_response = MagicMock() + mock_response.json.return_value = { + "action": "BLOCKED", + "blocked_reason": "policy violation", + } + mock_response.raise_for_status = MagicMock() + + with patch.object( + fail_open_guardrail.async_handler, "post", return_value=mock_response + ): + with pytest.raises(GuardrailRaisedException): + await fail_open_guardrail.apply_guardrail( + inputs={"texts": ["hi"]}, + request_data={}, + input_type="request", + ) + + @pytest.mark.asyncio + async def test_endpoint_error_raises_by_default(self, generic_guardrail): + """Default fail_on_error=True keeps blocking on a non-unreachable endpoint error.""" + error = httpx.HTTPStatusError( + "bad request", request=MagicMock(), response=MagicMock(status_code=400) + ) + with patch.object(generic_guardrail.async_handler, "post", side_effect=error): + with pytest.raises(Exception, match="Generic Guardrail API failed"): + await generic_guardrail.apply_guardrail( + inputs={"texts": ["hi"]}, + request_data={}, + input_type="request", + ) + + @pytest.mark.asyncio + async def test_response_path_continues_when_fail_on_error_false( + self, fail_open_guardrail + ): + """fail_on_error governs the response path identically to the request path.""" + error = httpx.HTTPStatusError( + "bad request", request=MagicMock(), response=MagicMock(status_code=400) + ) + with patch.object( + fail_open_guardrail.async_handler, "post", side_effect=error + ): + result = await fail_open_guardrail.apply_guardrail( + inputs={"texts": ["model output"]}, + request_data={}, + input_type="response", + ) + + assert result == {"texts": ["model output"]} + + @pytest.mark.asyncio + async def test_response_path_valid_block_still_blocks(self, fail_open_guardrail): + """On the response path too, a valid BLOCKED decision raises despite fail_on_error=False.""" + mock_response = MagicMock() + mock_response.json.return_value = { + "action": "BLOCKED", + "blocked_reason": "policy violation", + } + mock_response.raise_for_status = MagicMock() + + with patch.object( + fail_open_guardrail.async_handler, "post", return_value=mock_response + ): + with pytest.raises(GuardrailRaisedException): + await fail_open_guardrail.apply_guardrail( + inputs={"texts": ["model output"]}, + request_data={}, + input_type="response", + )