fix(guardrails): scan generateContent systemInstruction text and drop fastapi import from handler tests

This commit is contained in:
mateo-berri 2026-08-29 22:11:47 -07:00
parent a099be02fd
commit 05e4d2f946
2 changed files with 65 additions and 8 deletions

View file

@ -1,8 +1,9 @@
"""
Google GenAI generateContent handler for Unified Guardrails.
Extracts text from generateContent requests (contents[].parts[].text) and
responses (candidates[].content.parts[].text), applies the guardrail, and
Extracts text from generateContent requests (systemInstruction.parts[].text
and contents[].parts[].text) and responses (candidates[].content.parts[].text),
applies the guardrail, and
writes the guardrailed text back in place. Requests and responses may be
dicts (wire format) or google-genai SDK objects; streaming chunks may
additionally be raw SSE frames, which are scanned for detection (a blocking
@ -56,12 +57,29 @@ def _content_text_parts(content: object) -> tuple[object, ...]:
return tuple(part for part in parts if _part_text(part) is not None)
def _system_instruction(data: Mapping[str, object]) -> object | None:
return next(
(
value
for container in (data, data.get("config"))
if container is not None
for key in ("systemInstruction", "system_instruction")
for value in (_field(container, key),)
if value is not None
),
None,
)
def _request_text_parts(data: Mapping[str, object]) -> tuple[object, ...]:
contents: Final = data.get("contents")
content_list: Final = (
(contents,) if isinstance(contents, dict) else tuple(contents) if isinstance(contents, list) else ()
)
return tuple(part for content in content_list for part in _content_text_parts(content))
return (
*_content_text_parts(_system_instruction(data)),
*(part for content in content_list for part in _content_text_parts(content)),
)
def _response_text_parts(response: object) -> tuple[object, ...]:

View file

@ -7,7 +7,6 @@ from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
import pytest
from fastapi import HTTPException
from litellm.llms.gemini.google_genai.guardrail_translation.handler import (
GoogleGenAIGenerateContentHandler,
@ -15,6 +14,10 @@ from litellm.llms.gemini.google_genai.guardrail_translation.handler import (
from litellm.types.utils import CallTypes
class GuardrailBlockedError(Exception):
pass
def _mock_guardrail(returned_texts):
guardrail = MagicMock()
guardrail.apply_guardrail = AsyncMock(return_value={"texts": returned_texts})
@ -39,6 +42,42 @@ async def test_input_contents_text_is_guardrailed_and_written_back():
assert result["contents"][0]["parts"][0]["text"] == "masked question"
@pytest.mark.asyncio
async def test_input_system_instruction_text_is_scanned_and_written_back():
handler = GoogleGenAIGenerateContentHandler()
guardrail = _mock_guardrail(["masked instruction", "masked question"])
data = {
"model": "gemini-2.5-flash",
"systemInstruction": {"role": "system", "parts": [{"text": "prohibited instruction"}]},
"contents": [{"role": "user", "parts": [{"text": "benign question"}]}],
}
result = await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
assert guardrail.apply_guardrail.call_args.kwargs["inputs"]["texts"] == [
"prohibited instruction",
"benign question",
]
assert result["systemInstruction"]["parts"][0]["text"] == "masked instruction"
assert result["contents"][0]["parts"][0]["text"] == "masked question"
@pytest.mark.asyncio
async def test_input_config_nested_snake_case_system_instruction_is_scanned():
handler = GoogleGenAIGenerateContentHandler()
guardrail = _mock_guardrail(["clean"])
instruction_part = SimpleNamespace(text="prohibited instruction")
data = {
"contents": [],
"config": SimpleNamespace(system_instruction=SimpleNamespace(parts=[instruction_part])),
}
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
assert guardrail.apply_guardrail.call_args.kwargs["inputs"]["texts"] == ["prohibited instruction"]
assert instruction_part.text == "clean"
@pytest.mark.asyncio
async def test_input_without_text_skips_guardrail():
handler = GoogleGenAIGenerateContentHandler()
@ -107,10 +146,10 @@ async def test_output_without_text_skips_guardrail():
async def test_output_blocking_guardrail_exception_propagates():
handler = GoogleGenAIGenerateContentHandler()
guardrail = MagicMock()
guardrail.apply_guardrail = AsyncMock(side_effect=HTTPException(status_code=400, detail="blocked"))
guardrail.apply_guardrail = AsyncMock(side_effect=GuardrailBlockedError("blocked"))
response = {"candidates": [{"content": {"parts": [{"text": "harmful answer"}]}}]}
with pytest.raises(HTTPException):
with pytest.raises(GuardrailBlockedError):
await handler.process_output_response(response=response, guardrail_to_apply=guardrail)
@ -155,10 +194,10 @@ async def test_streaming_raw_sse_chunks_accumulate_text_across_split_frames():
async def test_streaming_blocking_guardrail_exception_propagates():
handler = GoogleGenAIGenerateContentHandler()
guardrail = MagicMock()
guardrail.apply_guardrail = AsyncMock(side_effect=HTTPException(status_code=400, detail="blocked"))
guardrail.apply_guardrail = AsyncMock(side_effect=GuardrailBlockedError("blocked"))
chunks = [{"candidates": [{"content": {"parts": [{"text": "harmful"}]}}]}]
with pytest.raises(HTTPException):
with pytest.raises(GuardrailBlockedError):
await handler.process_output_streaming_response(responses_so_far=chunks, guardrail_to_apply=guardrail)