Merge pull request #38869 from BerriAI/litellm_fix_guardrail_route_call_types

fix(guardrails): resolve generateContent routes and async-first passthrough call types
This commit is contained in:
Mateo Wang 2026-09-07 15:15:51 -07:00 committed by GitHub
commit d0dd3ce2d8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 666 additions and 10 deletions

View file

@ -14,21 +14,48 @@ from typing import Final
from litellm.types.utils import API_ROUTE_TO_CALL_TYPES, CallTypes
def _segment_matches(route_segment: str, pattern_segment: str) -> bool:
"""
Match one concrete path segment against one pattern segment.
A bare placeholder ({param}) matches any segment; a placeholder with a
literal suffix ({model}:generateContent) requires the segment to end with
that suffix and have a non-empty value before it.
"""
if not pattern_segment.startswith("{"):
return route_segment == pattern_segment
placeholder_end: Final = pattern_segment.find("}")
if placeholder_end == -1:
return route_segment == pattern_segment
literal_suffix: Final = pattern_segment[placeholder_end + 1 :]
if not literal_suffix:
return True
return route_segment.endswith(literal_suffix) and len(route_segment) > len(literal_suffix)
def _pattern_tail_spans_segments(pattern_tail: str) -> bool:
"""
Whether the pattern's last segment is a suffixed placeholder
({model}:generateContent) that may absorb extra route segments, mirroring
FastAPI's {model_name:path} converter for slash-containing model names.
"""
return pattern_tail.startswith("{") and "}" in pattern_tail and not pattern_tail.endswith("}")
def _route_matches_pattern(route: str, pattern: str) -> bool:
"""
Return True if the concrete route matches the pattern.
Pattern segments like {param} match any single path segment.
Pattern segments like {param} match any single path segment, and a
suffixed placeholder in the last segment may span multiple segments.
"""
route_parts: Final = route.strip("/").split("/")
pattern_parts: Final = pattern.strip("/").split("/")
if len(route_parts) != len(pattern_parts):
if len(route_parts) < len(pattern_parts):
return False
for r, p in zip(route_parts, pattern_parts):
if p.startswith("{") and p.endswith("}"):
continue
if r != p:
return False
return True
if len(route_parts) > len(pattern_parts) and not _pattern_tail_spans_segments(pattern_parts[-1]):
return False
head_count: Final = len(pattern_parts) - 1
merged_parts: Final = (*route_parts[:head_count], "/".join(route_parts[head_count:]))
return all(_segment_matches(r, p) for r, p in zip(merged_parts, pattern_parts))
def get_call_types_for_route(route: str) -> Sequence[CallTypes] | None:

View file

@ -0,0 +1,20 @@
"""Google GenAI generateContent guardrail translation handler."""
from typing import Final
from litellm.llms.gemini.google_genai.guardrail_translation.handler import (
GoogleGenAIGenerateContentHandler,
)
from litellm.types.utils import CallTypes
guardrail_translation_mappings: Final = { # mutable-ok: discover_guardrail_translation_mappings only accepts isinstance(mappings, dict)
CallTypes.generate_content: GoogleGenAIGenerateContentHandler,
CallTypes.agenerate_content: GoogleGenAIGenerateContentHandler,
CallTypes.generate_content_stream: GoogleGenAIGenerateContentHandler,
CallTypes.agenerate_content_stream: GoogleGenAIGenerateContentHandler,
}
__all__ = (
"GoogleGenAIGenerateContentHandler",
"guardrail_translation_mappings",
)

View file

@ -0,0 +1,255 @@
"""
Google GenAI generateContent handler for Unified Guardrails.
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
guardrail raises) without rewriting the frames.
"""
import json
from collections.abc import Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Optional
from litellm._logging import verbose_proxy_logger
from litellm.llms.base_llm.guardrail_translation.base_translation import (
BaseTranslation,
StreamTransformSink,
)
from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._types import UserAPIKeyAuth
_EMPTY_REQUEST_DATA: Final[Mapping[str, object]] = MappingProxyType({})
def _field(container: object, name: str) -> object | None:
if isinstance(container, dict):
return container.get(name)
return getattr(container, name, None)
def _part_text(part: object) -> str | None:
text: Final = _field(part, "text")
if isinstance(text, str) and text:
return text
return None
def _write_part_text(part: object, text: str) -> None:
if isinstance(part, dict):
part["text"] = text # rebind-ok: guardrail write-back rewrites the caller's part in place by handler contract
return
setattr(part, "text", text) # noqa: B010 # SDK parts are typed as object here; direct assignment cannot type-check
def _content_text_parts(content: object) -> tuple[object, ...]:
parts: Final = _field(content, "parts")
if not isinstance(parts, (list, tuple)):
return ()
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 (
*_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, ...]:
candidates: Final = _field(response, "candidates")
if not isinstance(candidates, (list, tuple)):
return ()
return tuple(part for candidate in candidates for part in _content_text_parts(_field(candidate, "content")))
def _part_texts(text_parts: Sequence[object]) -> tuple[str, ...]:
return tuple(text for part in text_parts for text in (_part_text(part),) if text is not None)
def _texts_payload(
texts: Sequence[str],
) -> list[str]: # mutable-ok: GenericGuardrailAPIInputs.texts is declared list[str]
return list(texts) # mutable-ok: GenericGuardrailAPIInputs.texts is declared list[str]
def _write_back_texts(text_parts: Sequence[object], guardrailed_texts: Sequence[str] | None) -> None:
if not guardrailed_texts or len(guardrailed_texts) != len(text_parts):
return
for part, text in zip(text_parts, guardrailed_texts):
_write_part_text(part, text)
def _parse_json_dict_or_none(payload: str) -> Mapping[str, object] | None:
try:
parsed: Final = json.loads(payload)
except json.JSONDecodeError:
return None
if isinstance(parsed, dict):
return parsed
return None
def _sse_payload_texts(sse_text: str) -> tuple[str, ...]:
return tuple(
text
for line in sse_text.splitlines()
if line.startswith("data:")
for payload in (line[len("data:") :].strip(),)
if payload and payload != "[DONE]"
for parsed in (_parse_json_dict_or_none(payload),)
if parsed is not None
for text in _part_texts(_response_text_parts(parsed))
)
def _chunk_sse_text(chunk: object) -> str | None:
if isinstance(chunk, bytes):
return chunk.decode("utf-8", errors="replace")
if isinstance(chunk, str):
return chunk
return None
def _accumulated_stream_text(responses_so_far: Sequence[object]) -> str:
object_texts: Final = tuple(
text
for chunk in responses_so_far
if _chunk_sse_text(chunk) is None
for text in _part_texts(_response_text_parts(chunk))
)
sse_text: Final = "".join(sse for chunk in responses_so_far for sse in (_chunk_sse_text(chunk),) if sse is not None)
return "".join(object_texts) + "".join(_sse_payload_texts(sse_text))
class GoogleGenAIGenerateContentHandler(BaseTranslation):
"""
Guardrail translation for the google genai generateContent surface
(/models/{model}:generateContent, :streamGenerateContent, and the
litellm SDK generate_content call types).
"""
async def process_input_messages(
self,
data: dict, # mutable-ok: base handler contract passes the proxy's request dict through to apply_guardrail
guardrail_to_apply: "CustomGuardrail",
litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None,
) -> object:
text_parts: Final = _request_text_parts(data)
if not text_parts:
verbose_proxy_logger.debug("Google GenAI guardrail: no request text found, skipping")
return data
model: Final = data.get("model")
inputs: Final = (
GenericGuardrailAPIInputs(texts=_texts_payload(_part_texts(text_parts)), model=model)
if isinstance(model, str)
else GenericGuardrailAPIInputs(texts=_texts_payload(_part_texts(text_parts)))
)
guardrailed_inputs: Final = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
request_data=data,
input_type="request",
logging_obj=litellm_logging_obj,
)
_write_back_texts(text_parts, guardrailed_inputs.get("texts"))
return data
async def process_output_response(
self,
response: object,
guardrail_to_apply: "CustomGuardrail",
litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None,
user_api_key_dict: Optional["UserAPIKeyAuth"] = None,
request_data: Mapping[str, object] | None = None,
) -> object:
text_parts: Final = _response_text_parts(response)
if not text_parts:
verbose_proxy_logger.debug("Google GenAI guardrail: no response text found, skipping")
return response
guardrail_request_data: Final = self._merged_request_data(
request_data=request_data,
user_api_key_dict=user_api_key_dict,
context_key="response",
context_value=response,
)
model: Final = guardrail_request_data.get("model")
inputs: Final = (
GenericGuardrailAPIInputs(texts=_texts_payload(_part_texts(text_parts)), model=model)
if isinstance(model, str)
else GenericGuardrailAPIInputs(texts=_texts_payload(_part_texts(text_parts)))
)
guardrailed_inputs: Final = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
request_data=guardrail_request_data,
input_type="response",
logging_obj=litellm_logging_obj,
)
_write_back_texts(text_parts, guardrailed_inputs.get("texts"))
return response
async def process_output_streaming_response(
self,
responses_so_far: Sequence[object],
guardrail_to_apply: "CustomGuardrail",
litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None,
user_api_key_dict: Optional["UserAPIKeyAuth"] = None,
request_data: Mapping[str, object] | None = None,
stream_transform_sink: StreamTransformSink | None = None,
) -> object:
accumulated_text: Final = _accumulated_stream_text(responses_so_far)
if not accumulated_text:
return responses_so_far
guardrail_request_data: Final = self._merged_request_data(
request_data=request_data,
user_api_key_dict=user_api_key_dict,
context_key="responses_so_far",
context_value=responses_so_far,
)
_guardrailed_inputs: Final = await guardrail_to_apply.apply_guardrail(
inputs=GenericGuardrailAPIInputs(texts=_texts_payload((accumulated_text,))),
request_data=guardrail_request_data,
input_type="response",
logging_obj=litellm_logging_obj,
)
return responses_so_far
def _merged_request_data(
self,
request_data: Mapping[str, object] | None,
user_api_key_dict: Optional["UserAPIKeyAuth"],
context_key: str,
context_value: object,
) -> dict: # mutable-ok: CustomGuardrail.apply_guardrail requires a plain dict request payload
base: Final = request_data if request_data is not None else _EMPTY_REQUEST_DATA
user_metadata: Final = self.transform_user_api_key_dict_to_metadata(user_api_key_dict)
context_pairs: Final = ((context_key, context_value),) if context_key not in base else ()
metadata_pairs: Final = (
(("litellm_metadata", user_metadata),) if user_metadata and "litellm_metadata" not in base else ()
)
return dict((*base.items(), *context_pairs, *metadata_pairs)) # mutable-ok: apply_guardrail takes a plain dict

View file

@ -957,6 +957,14 @@ API_ROUTE_TO_CALL_TYPES: Final[Mapping[str, Sequence[CallTypes]]] = {
CallTypes.agenerate_content_stream,
CallTypes.generate_content_stream,
],
"/v1beta/models/{model}:generateContent": (
CallTypes.agenerate_content,
CallTypes.generate_content,
),
"/v1beta/models/{model}:streamGenerateContent": (
CallTypes.agenerate_content_stream,
CallTypes.generate_content_stream,
),
# MCP (Model Context Protocol)
"/mcp/call_tool": [CallTypes.call_mcp_tool],
# A2A (Agent-to-Agent)
@ -964,12 +972,12 @@ API_ROUTE_TO_CALL_TYPES: Final[Mapping[str, Sequence[CallTypes]]] = {
"/a2a/{agent_id}/message/send": [CallTypes.asend_message, CallTypes.send_message],
# Passthrough endpoints
"/llm_passthrough": [
CallTypes.llm_passthrough_route,
CallTypes.allm_passthrough_route,
CallTypes.llm_passthrough_route,
],
"/v1/llm_passthrough": [
CallTypes.llm_passthrough_route,
CallTypes.allm_passthrough_route,
CallTypes.llm_passthrough_route,
],
"/v1/messages": [CallTypes.anthropic_messages],
# OCR

View file

@ -0,0 +1,112 @@
"""
Tests for route -> CallTypes resolution (api_route_to_call_types).
Regression coverage for the guardrail route table bugs:
- placeholder segments with a literal suffix ({model}:generateContent) never matched
- the /v1beta generateContent routes were missing from the table
- /llm_passthrough listed the sync call type first, resolving consumers that
take call_types[0] to a handler-less type
"""
from litellm.litellm_core_utils.api_route_to_call_types import (
get_call_types_for_route,
)
from litellm.types.utils import API_ROUTE_TO_CALL_TYPES, CallTypes
class TestGenerateContentRouteResolution:
def test_bare_generate_content_route_resolves(self):
call_types = get_call_types_for_route("/models/gemini-2.5-flash:generateContent")
assert call_types is not None
assert list(call_types) == [CallTypes.agenerate_content, CallTypes.generate_content]
def test_v1beta_generate_content_route_resolves(self):
call_types = get_call_types_for_route("/v1beta/models/gemini-2.5-flash:generateContent")
assert call_types is not None
assert list(call_types) == [CallTypes.agenerate_content, CallTypes.generate_content]
def test_bare_stream_generate_content_route_resolves(self):
call_types = get_call_types_for_route("/models/gemini-2.5-flash:streamGenerateContent")
assert call_types is not None
assert list(call_types) == [
CallTypes.agenerate_content_stream,
CallTypes.generate_content_stream,
]
def test_v1beta_stream_generate_content_route_resolves(self):
call_types = get_call_types_for_route("/v1beta/models/gemini-2.5-flash:streamGenerateContent")
assert call_types is not None
assert list(call_types) == [
CallTypes.agenerate_content_stream,
CallTypes.generate_content_stream,
]
def test_slash_containing_model_name_resolves(self):
call_types = get_call_types_for_route("/v1beta/models/gemini/gemini-2.5-flash:generateContent")
assert call_types is not None
assert list(call_types) == [CallTypes.agenerate_content, CallTypes.generate_content]
def test_empty_model_name_does_not_match(self):
assert get_call_types_for_route("/models/:generateContent") is None
def test_unrelated_model_action_does_not_match(self):
assert get_call_types_for_route("/models/gemini-2.5-flash:countTokens") is None
class TestPassthroughRouteOrdering:
def test_llm_passthrough_lists_async_call_type_first(self):
call_types = get_call_types_for_route("/llm_passthrough")
assert call_types is not None
assert list(call_types) == [
CallTypes.allm_passthrough_route,
CallTypes.llm_passthrough_route,
]
def test_v1_llm_passthrough_lists_async_call_type_first(self):
call_types = get_call_types_for_route("/v1/llm_passthrough")
assert call_types is not None
assert list(call_types) == [
CallTypes.allm_passthrough_route,
CallTypes.llm_passthrough_route,
]
class TestFirstCallTypeHasTranslationHandler:
def test_first_call_type_is_translatable_whenever_any_is(self):
"""
Consumers (unified guardrail post-call and streaming resolution) take
call_types[0]. A route whose first call type lacks a guardrail
translation handler while a later one has it silently skips guardrail
scanning, so the table must list a handler-backed call type first.
"""
from litellm.llms import load_guardrail_translation_mappings
mappings = load_guardrail_translation_mappings()
misordered = {
route: [call_type.value for call_type in call_types]
for route, call_types in API_ROUTE_TO_CALL_TYPES.items()
if call_types
and call_types[0] not in mappings
and any(call_type in mappings for call_type in call_types)
}
assert misordered == {}
class TestExistingRouteResolutionUnchanged:
def test_exact_route_still_resolves(self):
call_types = get_call_types_for_route("/chat/completions")
assert call_types is not None
assert CallTypes.acompletion in call_types
def test_single_segment_placeholder_still_resolves(self):
call_types = get_call_types_for_route("/a2a/my-agent/message/send")
assert call_types is not None
assert list(call_types) == [CallTypes.asend_message, CallTypes.send_message]
def test_longer_route_does_not_collapse_into_bare_placeholder_pattern(self):
call_types = get_call_types_for_route("/responses/resp_123/input_items")
assert call_types is not None
assert list(call_types) == [CallTypes.alist_input_items]
def test_unknown_route_returns_none(self):
assert get_call_types_for_route("/not/a/real/route") is None

View file

@ -0,0 +1,234 @@
"""
Tests for the Google GenAI generateContent guardrail translation handler.
"""
import json
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
import pytest
from litellm.llms.gemini.google_genai.guardrail_translation.handler import (
GoogleGenAIGenerateContentHandler,
)
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})
return guardrail
@pytest.mark.asyncio
async def test_input_contents_text_is_guardrailed_and_written_back():
handler = GoogleGenAIGenerateContentHandler()
guardrail = _mock_guardrail(["masked question"])
data = {
"model": "gemini-2.5-flash",
"contents": [{"role": "user", "parts": [{"text": "raw question"}]}],
}
result = await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
call_kwargs = guardrail.apply_guardrail.call_args.kwargs
assert call_kwargs["inputs"]["texts"] == ["raw question"]
assert call_kwargs["inputs"]["model"] == "gemini-2.5-flash"
assert call_kwargs["input_type"] == "request"
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()
guardrail = _mock_guardrail([])
data = {"model": "gemini-2.5-flash", "contents": [{"role": "user", "parts": [{"inlineData": {}}]}]}
result = await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
guardrail.apply_guardrail.assert_not_called()
assert result is data
@pytest.mark.asyncio
async def test_output_dict_response_text_is_guardrailed_and_written_back():
handler = GoogleGenAIGenerateContentHandler()
guardrail = _mock_guardrail(["masked answer"])
response = {
"candidates": [
{
"content": {"role": "model", "parts": [{"text": "harmful answer"}]},
"finishReason": "STOP",
}
]
}
result = await handler.process_output_response(
response=response,
guardrail_to_apply=guardrail,
request_data={"model": "gemini-2.5-flash"},
)
call_kwargs = guardrail.apply_guardrail.call_args.kwargs
assert call_kwargs["inputs"]["texts"] == ["harmful answer"]
assert call_kwargs["input_type"] == "response"
assert call_kwargs["request_data"]["response"] is response
assert result["candidates"][0]["content"]["parts"][0]["text"] == "masked answer"
@pytest.mark.asyncio
async def test_output_object_response_text_is_guardrailed_and_written_back():
handler = GoogleGenAIGenerateContentHandler()
guardrail = _mock_guardrail(["masked answer"])
part = SimpleNamespace(text="harmful answer")
response = SimpleNamespace(
candidates=[SimpleNamespace(content=SimpleNamespace(parts=[part]), finish_reason="STOP")]
)
await handler.process_output_response(response=response, guardrail_to_apply=guardrail)
assert guardrail.apply_guardrail.call_args.kwargs["inputs"]["texts"] == ["harmful answer"]
assert part.text == "masked answer"
@pytest.mark.asyncio
async def test_output_without_text_skips_guardrail():
handler = GoogleGenAIGenerateContentHandler()
guardrail = _mock_guardrail([])
result = await handler.process_output_response(response={"candidates": []}, guardrail_to_apply=guardrail)
guardrail.apply_guardrail.assert_not_called()
assert result == {"candidates": []}
@pytest.mark.asyncio
async def test_output_blocking_guardrail_exception_propagates():
handler = GoogleGenAIGenerateContentHandler()
guardrail = MagicMock()
guardrail.apply_guardrail = AsyncMock(side_effect=GuardrailBlockedError("blocked"))
response = {"candidates": [{"content": {"parts": [{"text": "harmful answer"}]}}]}
with pytest.raises(GuardrailBlockedError):
await handler.process_output_response(response=response, guardrail_to_apply=guardrail)
@pytest.mark.asyncio
async def test_streaming_dict_chunks_accumulate_text():
handler = GoogleGenAIGenerateContentHandler()
guardrail = _mock_guardrail(["clean"])
chunks = [
{"candidates": [{"content": {"parts": [{"text": "harmful "}]}}]},
{"candidates": [{"content": {"parts": [{"text": "answer"}]}}]},
]
result = await handler.process_output_streaming_response(
responses_so_far=chunks,
guardrail_to_apply=guardrail,
)
call_kwargs = guardrail.apply_guardrail.call_args.kwargs
assert call_kwargs["inputs"]["texts"] == ["harmful answer"]
assert call_kwargs["input_type"] == "response"
assert result is chunks
@pytest.mark.asyncio
async def test_streaming_raw_sse_chunks_accumulate_text_across_split_frames():
handler = GoogleGenAIGenerateContentHandler()
guardrail = _mock_guardrail(["clean"])
frame_one = "data: " + json.dumps({"candidates": [{"content": {"parts": [{"text": "harmful "}]}}]}) + "\r\n\r\n"
frame_two = "data: " + json.dumps({"candidates": [{"content": {"parts": [{"text": "answer"}]}}]}) + "\r\n\r\n"
split_at = len(frame_one) // 2
chunks = [frame_one[:split_at], frame_one[split_at:] + frame_two[:5], frame_two[5:].encode("utf-8")]
await handler.process_output_streaming_response(
responses_so_far=chunks,
guardrail_to_apply=guardrail,
)
assert guardrail.apply_guardrail.call_args.kwargs["inputs"]["texts"] == ["harmful answer"]
@pytest.mark.asyncio
async def test_streaming_blocking_guardrail_exception_propagates():
handler = GoogleGenAIGenerateContentHandler()
guardrail = MagicMock()
guardrail.apply_guardrail = AsyncMock(side_effect=GuardrailBlockedError("blocked"))
chunks = [{"candidates": [{"content": {"parts": [{"text": "harmful"}]}}]}]
with pytest.raises(GuardrailBlockedError):
await handler.process_output_streaming_response(responses_so_far=chunks, guardrail_to_apply=guardrail)
@pytest.mark.asyncio
async def test_streaming_without_text_skips_guardrail():
handler = GoogleGenAIGenerateContentHandler()
guardrail = _mock_guardrail([])
result = await handler.process_output_streaming_response(responses_so_far=[], guardrail_to_apply=guardrail)
guardrail.apply_guardrail.assert_not_called()
assert result == []
def test_generate_content_call_types_are_registered():
from litellm.llms.gemini.google_genai.guardrail_translation import (
guardrail_translation_mappings,
)
for call_type in (
CallTypes.generate_content,
CallTypes.agenerate_content,
CallTypes.generate_content_stream,
CallTypes.agenerate_content_stream,
):
assert guardrail_translation_mappings[call_type] is GoogleGenAIGenerateContentHandler
def test_discovery_finds_generate_content_handler():
from litellm.llms import load_guardrail_translation_mappings
mappings = load_guardrail_translation_mappings()
assert mappings[CallTypes.agenerate_content] is GoogleGenAIGenerateContentHandler
assert mappings[CallTypes.agenerate_content_stream] is GoogleGenAIGenerateContentHandler