mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(guardrails): scan generateContent systemInstruction text and drop fastapi import from handler tests
This commit is contained in:
parent
a099be02fd
commit
05e4d2f946
2 changed files with 65 additions and 8 deletions
|
|
@ -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, ...]:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue