From e99151bb95eaf1de26009940aa780b7f6f8b8cda Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Fri, 26 Jun 2026 11:25:56 -0700 Subject: [PATCH] feat(guardrails): make the Generic Guardrail resilient to built-in tools and errors (adopted from #31286) (#31461) * fix(guardrails): stop Generic Guardrail API 500 on built-in tools Requests carrying built-in tools (code_interpreter, file_search, ...) crashed the Generic Guardrail with a 500. GenericGuardrailAPIRequest.tools validated each tool against ChatCompletionToolParam, whose base TypedDict requires a function block, so a tool like {"type": "code_interpreter"} raised a Pydantic ValidationError before the request was ever sent. Type the field with a permissive GuardrailToolParam model (type required, extra=allow) so built-in tools validate and their config is forwarded to the guardrail intact instead of being stripped. * feat(guardrails): add complete fail-open (fail_on_error) to Generic Guardrail The Generic Guardrail already honored unreachable_fallback, which fails open only on network-unreachable errors. This wires up the existing generic fail_on_error config (so far implemented only by Model Armor) so that fail_on_error=false degrades any guardrail error to a critical-log warning and lets the request proceed as if the guardrail were absent. Only a valid guardrail response can act: a parsed BLOCKED decision still raises, while endpoint errors, malformed responses, and internal serialization or validation errors all fall through when fail_on_error=false. To cover that last class, the request construction now runs inside the protected block, so the kind of validation error that previously surfaced as a 500 is caught here too. Defaults to true (fail closed), matching today's behavior; turning it off is an explicit availability-over-security choice and is logged at critical level on every bypass. * test(guardrails): cover fail_on_error on the response path The existing fail_on_error tests all drive the request path. Add response-path (input_type=response) coverage: an endpoint error proceeds unchanged under fail_on_error=false, and a valid BLOCKED decision still raises. Guards against a future regression that special-cases input_type in the error handling. * style(guardrails): black-format the fail-open guard expression CI runs black (line-length 88) over litellm/; the unreachable_fail_open assignment exceeded it. Wrap it to satisfy the formatter. * fix(guardrails): validate tools into GuardrailToolParam at the call site Changing the request field to List[GuardrailToolParam] left the construction passing List[ChatCompletionToolParam] (list is invariant), which tripped the basedpyright reportArgumentType budget gate. Validate each tool explicitly, which is what Pydantic did implicitly, so the types line up with no Any or suppression and the serialized payload is unchanged. * fix(guardrails): make fail-open log message accurate for non-network errors The fail-open path is now shared by fail_on_error, so it fires for any guardrail error, not just unreachability. The log said 'unreachable' even for an HTTP 400 or a malformed response; reword to 'error' (the status code and exception are already logged). Addresses the Greptile review's only finding. * fix(guardrails): align GenericGuardrailAPIResponse.tools with GuardrailToolParam Greptile flagged that the request side moved to GuardrailToolParam but the response side still annotated tools as List[ChatCompletionToolParam], which mandates a function block and contradicts the new built-in-tools support. Update the response annotation (and the now-unused import) so the two sides agree. Runtime is unchanged; from_dict stores the raw dicts and the only consumer assigns through to GenericGuardrailAPIInputs without inspecting the elements. --------- Co-authored-by: Itay Ovadia --- .../generic_guardrail_api/__init__.py | 1 + .../generic_guardrail_api.py | 53 +++--- litellm/types/guardrails.py | 7 +- .../guardrail_hooks/generic_guardrail_api.py | 31 +++- .../test_generic_guardrail_api.py | 168 ++++++++++++++++++ 5 files changed, 233 insertions(+), 27 deletions(-) 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", + )