feat(guardrails): make the Generic Guardrail resilient to built-in tools and errors (adopted from #31286) (#31461)
Some checks failed
GitHub Actions Security Analysis / zizmor (push) Waiting to run
LiteLLM Rust / rustfmt, clippy, test (push) Has been cancelled

* 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 <itay@sun.security>
This commit is contained in:
yucheng-berri 2026-06-26 11:25:56 -07:00 • committed by GitHub
parent 5a1c7839be
commit e99151bb95
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 233 additions and 27 deletions

View file

@ -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,

View file

@ -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(

View file

@ -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(

View file

@ -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

View file

@ -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",
)