mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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 <itay@sun.security>
This commit is contained in:
parent
5a1c7839be
commit
e99151bb95
5 changed files with 233 additions and 27 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue