mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
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:
commit
d0dd3ce2d8
8 changed files with 666 additions and 10 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
0
tests/test_litellm/llms/gemini/google_genai/__init__.py
Normal file
0
tests/test_litellm/llms/gemini/google_genai/__init__.py
Normal 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
|
||||
Loading…
Add table
Reference in a new issue