fix(model_armor): classify the stream surface and fail closed when it cannot be assembled

Decide the wire format explicitly instead of inferring it from a boolean pair, so an
opaque raw SSE stream (the Google :streamGenerateContent route) is never refused in
Anthropic framing, and a stream that cannot be assembled is blocked rather than
released unscanned unless fail_on_error is disabled.

Also scan Responses tool-call arguments, read the body only off a terminal Responses
event, and record the applied guardrail on the fail-closed path.
This commit is contained in:
Yucheng Zhu 2026-09-01 12:56:09 -07:00
parent 67a9932f7d
commit 4b21583005
3 changed files with 430 additions and 202 deletions

View file

@ -13,6 +13,19 @@ from typing import Final
from litellm.types.utils import Choices, ModelResponse
_ANTHROPIC_EVENT_TYPES: Final = frozenset(
{
"message_start",
"message_delta",
"message_stop",
"content_block_start",
"content_block_delta",
"content_block_stop",
"ping",
"error",
}
)
def is_raw_sse_stream(all_chunks: Sequence[object]) -> bool:
return any(isinstance(chunk, (str, bytes)) for chunk in all_chunks)
@ -30,23 +43,43 @@ def _joined_sse_stream(all_chunks: Sequence[object]) -> str | None:
return None
def _anthropic_message_start(sse_stream: str) -> Mapping[str, object] | None:
def _parsed_sse_events(sse_stream: str) -> tuple[Mapping[str, object], ...]:
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import (
AnthropicPassthroughLoggingHandler,
)
return tuple(
event_data
for event in AnthropicPassthroughLoggingHandler._split_sse_chunk_into_events(sse_stream) # pyright: ignore[reportPrivateUsage] # same parser the assembler uses
if (event_data := AnthropicPassthroughLoggingHandler._extract_sse_data(event)) is not None # pyright: ignore[reportPrivateUsage] # same parser the assembler uses; a private import beats forking SSE parsing
)
def _anthropic_message_start(sse_stream: str) -> Mapping[str, object] | None:
return next(
(
message
for event in AnthropicPassthroughLoggingHandler._split_sse_chunk_into_events(sse_stream) # pyright: ignore[reportPrivateUsage] # same parser the assembler uses
if (event_data := AnthropicPassthroughLoggingHandler._extract_sse_data(event)) is not None # pyright: ignore[reportPrivateUsage] # same parser the assembler uses; a private import beats forking SSE parsing
and event_data.get("type") == "message_start"
and isinstance(message := event_data.get("message"), dict)
for event_data in _parsed_sse_events(sse_stream)
if event_data.get("type") == "message_start" and isinstance(message := event_data.get("message"), dict)
),
None,
)
def is_anthropic_sse_stream(all_chunks: Sequence[object]) -> bool:
"""Whether raw SSE frames are Anthropic Messages events.
``is_raw_sse_stream`` only says the chunks are unparsed bytes, and ``/v1/messages`` is not the
only endpoint that streams those: the Google ``:streamGenerateContent`` route marks its own
stream raw too. Reading its frames as Anthropic ones would refuse the response in a wire format
its client cannot parse, so the surface is decided on the event types actually present.
"""
sse_stream: Final = _joined_sse_stream(all_chunks)
if sse_stream is None:
return False
return any(event.get("type") in _ANTHROPIC_EVENT_TYPES for event in _parsed_sse_events(sse_stream))
def assemble_anthropic_sse_stream(
all_chunks: Sequence[object], *, restore_identity: bool = False
) -> ModelResponse | None:
@ -119,20 +152,18 @@ def is_sse_error_stream(all_chunks: Sequence[object]) -> bool:
them would hide the refusal the client is owed. Covers both wire forms a guardrail emits: the
Anthropic ``error`` event and the chat-completions ``{"error": ...}`` payload.
"""
if not all(isinstance(chunk, (str, bytes)) for chunk in all_chunks):
# A stream mixing typed chunks with an error frame still carries content to scan, and the
# frames-only join below would drop exactly the part that has to be scanned
return False
sse_stream: Final = _joined_sse_stream(all_chunks)
if sse_stream is None:
return False
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import (
AnthropicPassthroughLoggingHandler,
events: Final = _parsed_sse_events(sse_stream)
return len(events) > 0 and all(
event.get("type") == "error" or isinstance(event.get("error"), Mapping) for event in events
)
events: Final = tuple(
event_data
for event in AnthropicPassthroughLoggingHandler._split_sse_chunk_into_events(sse_stream) # pyright: ignore[reportPrivateUsage] # same parser the assembler uses
if (event_data := AnthropicPassthroughLoggingHandler._extract_sse_data(event)) is not None # pyright: ignore[reportPrivateUsage] # same parser the assembler uses
)
return len(events) > 0 and all(event.get("type") == "error" or "error" in event for event in events)
def anthropic_sse_chunks_from_response(assembled: ModelResponse) -> tuple[bytes, ...]:
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (

View file

@ -1,4 +1,5 @@
from collections.abc import AsyncGenerator, Mapping, Sequence
from enum import Enum, auto
from typing import TYPE_CHECKING, Any, Final, Literal
import httpx
@ -25,15 +26,13 @@ from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.llms.openai.responses.guardrail_translation.handler import (
OpenAIResponsesHandler,
)
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.anthropic_sse import (
anthropic_sse_chunks_from_response,
anthropic_sse_error_frames,
assemble_anthropic_sse_stream,
is_anthropic_sse_stream,
is_raw_sse_stream,
is_sse_error_stream,
)
@ -42,7 +41,11 @@ from litellm.proxy.guardrails.guardrail_hooks.model_armor.file_scanning import (
plan_file_scans,
)
from litellm.types.guardrails import GuardrailEventHooks, LitellmParams
from litellm.types.llms.openai import AllMessageValues, ResponsesAPIResponse
from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionToolCallChunk,
ResponsesAPIResponse,
)
from litellm.types.utils import (
CallTypes,
CallTypesLiteral,
@ -56,6 +59,18 @@ from litellm.types.utils import (
GUARDRAIL_NAME: Final = "model_armor"
# Only these carry the finished output; response.created carries an empty body
_RESPONSES_TERMINAL_EVENT_TYPES: Final = frozenset({"response.completed", "response.incomplete", "response.failed"})
class _StreamSurface(Enum):
"""Wire format of a buffered streaming response, which decides how it is read and how it is refused."""
CHAT_COMPLETIONS = auto()
ANTHROPIC_MESSAGES = auto()
RESPONSES = auto()
OPAQUE_SSE = auto()
class ModelArmorAPIError(Exception):
"""Model Armor API failure (non-2xx), distinct from a content-block decision so
@ -843,42 +858,73 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
return response
@staticmethod
def _final_responses_api_response(
all_chunks: Sequence[object],
) -> tuple[bool, ResponsesAPIResponse | None]:
"""Detect a ``/v1/responses`` event stream and return its final response body.
def _is_terminal_error_stream(all_chunks: Sequence[object]) -> bool:
"""Whether the buffered stream is only the refusal an earlier guardrail in the chain emitted.
Returns ``(is_responses_stream, final_response)`` so the caller can tell a stream
that is not a Responses stream apart from one whose terminal ``response.completed``
event never arrived.
post_call guardrails are composed, so this hook can be handed the terminal error items a
preceding one produced. They carry no message to scan, and replacing them would hide the
refusal the client is owed.
"""
events: Final = tuple(
(event_type, getattr(chunk, "response", None))
if all(getattr(chunk, "type", None) == "error" for chunk in all_chunks):
return True
return is_sse_error_stream(all_chunks)
@staticmethod
def _classify_stream(all_chunks: Sequence[object]) -> _StreamSurface:
"""Wire format the buffered chunks belong to."""
if is_raw_sse_stream(all_chunks):
return (
_StreamSurface.ANTHROPIC_MESSAGES if is_anthropic_sse_stream(all_chunks) else _StreamSurface.OPAQUE_SSE
)
if any(
isinstance(event_type := getattr(chunk, "type", None), str) and event_type.startswith("response.")
for chunk in all_chunks
if isinstance(event_type := getattr(chunk, "type", None), str)
)
return (
any(event_type.startswith("response.") for event_type, _ in events),
next(
(body for _, body in reversed(events) if isinstance(body, ResponsesAPIResponse)),
None,
):
return _StreamSurface.RESPONSES
return _StreamSurface.CHAT_COMPLETIONS
@staticmethod
def _final_responses_api_response(all_chunks: Sequence[object]) -> ResponsesAPIResponse | None:
"""Response body carried by a terminal ``/v1/responses`` event.
A stream cut short before it completes has to read as unassembled rather than as a clean
empty response: ``response.created`` also carries a body, but an empty one, and scanning
that would release every buffered delta unscanned.
"""
return next(
(
body
for chunk in reversed(all_chunks)
if getattr(chunk, "type", None) in _RESPONSES_TERMINAL_EVENT_TYPES
and isinstance(body := getattr(chunk, "response", None), ResponsesAPIResponse)
),
None,
)
@staticmethod
def _responses_api_response_text(response: ResponsesAPIResponse) -> str:
"""Concatenate the output text carried by a Responses API response."""
"""Text to scan in a Responses API response, tool-call arguments included.
Tool calls are folded in because ``get_content_from_model_response`` folds them into what
the chat surface scans, and a Responses turn can carry its whole payload in them.
"""
from litellm.llms.openai.responses.guardrail_translation.handler import (
OpenAIResponsesHandler,
)
texts: Final[list[str]] = [] # mutable-ok: the shared extractor below appends into caller-owned lists
tool_calls: Final[list[ChatCompletionToolCallChunk]] = [] # mutable-ok: the same extractor's tool-call sink
handler: Final = OpenAIResponsesHandler()
for output_idx, output_item in enumerate(response.output or ()):
handler._extract_output_text_and_images( # pyright: ignore[reportPrivateUsage] # the shared Responses output extractor; forking it would duplicate per-item parsing
output_item,
output_idx,
texts,
[], # mutable-ok: the extractor's images sink, unused here
[], # mutable-ok: the extractor's task-mapping sink, unused here
output_item=output_item,
output_idx=output_idx,
texts_to_check=texts,
images_to_check=[], # mutable-ok: the extractor's images sink, unused here
task_mappings=[], # mutable-ok: the extractor's task-mapping sink, unused here
tool_calls_to_check=tool_calls,
)
return "".join(texts)
return "".join((*texts, *(json.dumps(tool_call) for tool_call in tool_calls)))
def _extract_streaming_content(self, assembled_response: object) -> str:
"""Text to scan from an assembled stream, for every endpoint shape this hook serves."""
@ -897,35 +943,56 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
def _assemble_chat_completion_stream(
all_chunks: list[Any], # mutable-ok: stream_chunk_builder only accepts a mutable list
) -> ModelResponse | TextCompletionResponse | None:
"""Assemble chat-completion chunks, returning ``None`` when they are not chat deltas."""
"""Assemble chat-completion chunks, returning ``None`` when they cannot be assembled."""
from litellm.main import stream_chunk_builder
try:
return stream_chunk_builder(chunks=all_chunks)
except Exception as exc:
verbose_proxy_logger.warning(
"Model Armor: could not assemble the streamed response for scanning (%s), forwarding it unscanned",
exc,
)
verbose_proxy_logger.warning("Model Armor: chat-completion stream assembly failed (%s)", exc)
return None
def _assemble_stream(
self, all_chunks: Sequence[object], surface: _StreamSurface
) -> ModelResponse | TextCompletionResponse | ResponsesAPIResponse | None:
"""Assemble the buffered stream into the scannable response its surface produces."""
if surface is _StreamSurface.ANTHROPIC_MESSAGES:
return assemble_anthropic_sse_stream(all_chunks, restore_identity=True)
if surface is _StreamSurface.RESPONSES:
return self._final_responses_api_response(all_chunks)
if surface is _StreamSurface.OPAQUE_SSE:
return None
return self._assemble_chat_completion_stream(list(all_chunks))
@staticmethod
def _stream_error_items(
exc: HTTPException,
error_obj: Mapping[str, object],
*,
raw_sse: bool,
responses_stream: bool,
) -> Sequence[object]:
def _error_payload(exc: HTTPException) -> Mapping[str, object]:
"""Error object for a terminal stream item, carrying the status the frame would otherwise lose."""
detail: Final = exc.detail if isinstance(exc.detail, Mapping) else {"message": str(exc.detail)}
error_value: Final = detail.get("error", detail)
return {
**(dict(error_value) if isinstance(error_value, Mapping) else {"message": str(error_value)}),
"code": str(exc.status_code),
}
@staticmethod
def _build_responses_error_items(exc: HTTPException) -> Sequence[object] | None:
"""Responses API error events for a failure discovered after the stream started."""
from litellm.llms.openai.responses.guardrail_translation.handler import (
OpenAIResponsesHandler,
)
return OpenAIResponsesHandler().build_stream_error_items(exc, responses_so_far=None)
def _stream_error_items(self, exc: HTTPException, *, surface: _StreamSurface) -> Sequence[object]:
"""Frame a guardrail failure as terminal stream items in this endpoint's wire format."""
chat_completions_form: Final = (f"data: {json.dumps({'error': error_obj})}\n\n",)
if raw_sse:
return anthropic_sse_error_frames(str(error_obj.get("message", "")))
if responses_stream:
return (
OpenAIResponsesHandler().build_stream_error_items(exc, responses_so_far=None) or chat_completions_form
)
return chat_completions_form
payload: Final = self._error_payload(exc)
if surface is _StreamSurface.ANTHROPIC_MESSAGES:
return anthropic_sse_error_frames(str(payload.get("message", "")))
if surface is _StreamSurface.RESPONSES and (responses_items := self._build_responses_error_items(exc)):
return responses_items
# Also the fallback when a surface cannot frame its own error: create_response() reads the
# status back out of this form, so the refusal keeps its code instead of arriving as a 200
return (f"data: {json.dumps({'error': payload})}\n\n",)
async def async_post_call_streaming_iterator_hook(
self,
@ -936,78 +1003,93 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
"""Process streaming response chunks."""
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
from litellm.proxy.common_utils.callback_utils import (
add_guardrail_to_applied_guardrails_header,
)
# Collect all chunks
all_chunks: Final[list[Any]] = []
async for chunk in response:
all_chunks.append(chunk)
raw_sse: Final = is_raw_sse_stream(all_chunks)
responses_stream, final_responses_api_response = (
(False, None) if raw_sse else self._final_responses_api_response(all_chunks)
)
if not all_chunks or self._is_terminal_error_stream(all_chunks):
for chunk in all_chunks:
yield chunk
return
surface: Final = self._classify_stream(all_chunks)
# Build complete response
assembled_response: Final = (
assemble_anthropic_sse_stream(all_chunks, restore_identity=True)
if raw_sse
else final_responses_api_response
if responses_stream
else self._assemble_chat_completion_stream(all_chunks)
)
assembled_response: Final = self._assemble_stream(all_chunks, surface)
if assembled_response is None:
if not self.optional_params.get("fail_on_error", True):
verbose_proxy_logger.warning(
"Model Armor: streamed response could not be assembled for scanning, "
"forwarding it unscanned because fail_on_error is disabled"
)
for chunk in all_chunks:
yield chunk
return
if (
assembled_response is None
and (raw_sse or responses_stream)
and not (raw_sse and is_sse_error_stream(all_chunks))
):
# Forwarding an unscannable stream would silently disable the guardrail, so fail closed
unscannable: Final = HTTPException(
status_code=500,
detail=f"{self.guardrail_name}: streamed response could not be assembled for scanning, blocking it",
)
add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name)
for error_item in self._stream_error_items(
unscannable,
{"message": str(unscannable.detail), "code": "500"},
raw_sse=raw_sse,
responses_stream=responses_stream,
HTTPException(
status_code=500,
detail=f"{self.guardrail_name}: streamed response could not be assembled for scanning, blocking it",
),
surface=surface,
):
yield error_item
return
if assembled_response is not None:
# Extract content
content: Final = self._extract_streaming_content(assembled_response)
# Extract content
content: Final = self._extract_streaming_content(assembled_response)
if content:
try:
# Check with Model Armor
armor_response: Final = await self.make_model_armor_request(
content=content,
source="model_response",
request_data=request_data,
)
if not content:
verbose_proxy_logger.debug("Model Armor: No text content in streaming response, skipping guardrail")
for chunk in all_chunks:
yield chunk
return
# Attach Model Armor response & status to this request's metadata to avoid race conditions
if isinstance(request_data, dict):
_, metadata = get_or_create_metadata_bucket(request_data)
metadata["_model_armor_response"] = self._build_logging_response(armor_response)
metadata["_model_armor_status"] = (
"blocked" if self._should_block_content(armor_response) else "success"
try:
# Check with Model Armor
armor_response: Final = await self.make_model_armor_request(
content=content,
source="model_response",
request_data=request_data,
)
# Attach Model Armor response & status to this request's metadata to avoid race conditions
if isinstance(request_data, dict):
_, metadata = get_or_create_metadata_bucket(request_data)
metadata["_model_armor_response"] = self._build_logging_response(armor_response)
metadata["_model_armor_status"] = "blocked" if self._should_block_content(armor_response) else "success"
# Add guardrail to applied_guardrails BEFORE potential blocking
# This ensures guardrail is recorded even when it blocks the request
add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name)
# Check if blocked
if self._should_block_content(armor_response):
raise HTTPException(
status_code=400,
detail=self._build_block_error_detail(
"Streaming response blocked by Model Armor",
armor_response,
),
)
# Apply sanitization if enabled
if self.mask_response_content:
sanitized_content: Final = self._get_sanitized_content(armor_response)
if sanitized_content and sanitized_content != content:
if not isinstance(assembled_response, ModelResponse):
verbose_proxy_logger.warning(
"Model Armor: sanitized content cannot be re-emitted on this "
"streaming endpoint, blocking the response instead"
)
# Add guardrail to applied_guardrails BEFORE potential blocking
# This ensures guardrail is recorded even when it blocks the request
from litellm.proxy.common_utils.callback_utils import (
add_guardrail_to_applied_guardrails_header,
)
add_guardrail_to_applied_guardrails_header(
request_data=request_data, guardrail_name=self.guardrail_name
)
# Check if blocked
if self._should_block_content(armor_response):
raise HTTPException(
status_code=400,
detail=self._build_block_error_detail(
@ -1016,72 +1098,37 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
),
)
# Apply sanitization if enabled
if self.mask_response_content:
sanitized_content: Final = self._get_sanitized_content(armor_response)
if sanitized_content and sanitized_content != content:
if not isinstance(assembled_response, ModelResponse):
verbose_proxy_logger.warning(
"Model Armor: sanitized content cannot be re-emitted on this "
"streaming endpoint, blocking the response instead"
)
raise HTTPException(
status_code=400,
detail=self._build_block_error_detail(
"Streaming response blocked by Model Armor",
armor_response,
),
)
# Update assembled response
self._apply_sanitized_content(assembled_response, sanitized_content)
# Update assembled response
self._apply_sanitized_content(assembled_response, sanitized_content)
# Return sanitized stream
if raw_sse:
for sse_chunk in anthropic_sse_chunks_from_response(assembled_response):
yield sse_chunk
return
mock_response: Final = MockResponseIterator(model_response=assembled_response)
async for chunk in mock_response:
yield chunk
return
except ModelArmorAPIError as e:
if self.optional_params.get("fail_on_error", True):
error_obj = {"message": e.detail, "code": "500"}
for error_item in self._stream_error_items(
HTTPException(status_code=500, detail=e.detail),
error_obj,
raw_sse=raw_sse,
responses_stream=responses_stream,
):
yield error_item
# Return sanitized stream
if surface is _StreamSurface.ANTHROPIC_MESSAGES:
for sse_chunk in anthropic_sse_chunks_from_response(assembled_response):
yield sse_chunk
return
except HTTPException as e:
# Yield the error as a terminal stream item so create_response() detects
# it and returns a proper JSON error response with the correct status code.
# (Raising from a generator hits create_response's generic except → 500.)
detail: Final = e.detail if isinstance(e.detail, dict) else {"message": str(e.detail)}
error_value: Final = detail.get("error", detail)
if isinstance(error_value, dict):
error_obj = dict(error_value)
else:
error_obj = {"message": str(error_value)}
error_obj["code"] = str(e.status_code)
for error_item in self._stream_error_items(
e,
error_obj,
raw_sse=raw_sse,
responses_stream=responses_stream,
):
yield error_item
mock_response: Final = MockResponseIterator(model_response=assembled_response)
async for chunk in mock_response:
yield chunk
return
except Exception as e:
verbose_proxy_logger.error("Model Armor streaming error: %s", str(e), exc_info=True)
if self.optional_params.get("fail_on_error", True):
raise
else:
verbose_proxy_logger.debug("Model Armor: No text content in streaming response, skipping guardrail")
except ModelArmorAPIError as e:
if self.optional_params.get("fail_on_error", True):
for error_item in self._stream_error_items(
HTTPException(status_code=500, detail=e.detail), surface=surface
):
yield error_item
return
except HTTPException as e:
# Yield the error as a terminal stream item so create_response() detects it and returns
# a proper JSON error response with the correct status code. Raising from a generator
# instead hits create_response's generic except and becomes a 500.
for error_item in self._stream_error_items(e, surface=surface):
yield error_item
return
except Exception as e:
verbose_proxy_logger.error("Model Armor streaming error: %s", str(e), exc_info=True)
if self.optional_params.get("fail_on_error", True):
raise
# Return original chunks if no sanitization needed
for chunk in all_chunks:

View file

@ -15,9 +15,7 @@ import litellm.types.utils
from litellm._logging import verbose_proxy_logger
from litellm.caching import DualCache
from litellm.llms.custom_httpx.http_handler import MaskedHTTPStatusError
from litellm.llms.openai.responses.guardrail_translation.handler import (
OpenAIResponsesHandler,
)
from litellm.proxy.guardrails.anthropic_sse import anthropic_sse_error_frames
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.model_armor import ModelArmorGuardrail
from litellm.proxy.guardrails.guardrail_hooks.model_armor.model_armor import (
@ -3892,7 +3890,7 @@ def _responses_api_events():
)
async def _drain_surface_hook(guardrail, chunks):
async def _drain_surface_hook(guardrail, chunks, request_data=None):
async def _stream():
for chunk in chunks:
yield chunk
@ -3902,7 +3900,9 @@ async def _drain_surface_hook(guardrail, chunks):
async for item in guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(),
response=_stream(),
request_data={
request_data=request_data
if request_data is not None
else {
"model": "claude-haiku",
"messages": [{"role": "user", "content": "show me a card"}],
"metadata": {"guardrails": ["model-armor-test"]},
@ -4059,6 +4059,7 @@ async def test_streaming_api_failure_frames_error_per_surface(surface):
id="anthropic-sse-without-message-start",
),
pytest.param(None, id="responses-stream-without-completed-event"),
pytest.param("created", id="responses-stream-cut-off-after-response-created"),
],
)
async def test_streaming_hook_fails_closed_when_a_surface_stream_cannot_be_assembled(chunks):
@ -4070,16 +4071,17 @@ async def test_streaming_hook_fails_closed_when_a_surface_stream_cannot_be_assem
ResponsesAPIStreamEvents,
)
if chunks is None:
chunks = (
OutputTextDeltaEvent(
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA,
item_id="msg_1",
output_index=0,
content_index=0,
delta="hi",
),
if chunks is None or chunks == "created":
delta = OutputTextDeltaEvent(
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA,
item_id="msg_1",
output_index=0,
content_index=0,
delta="my card is 4111-1111-1111-1111",
)
# response.created carries a ResponsesAPIResponse too, but an empty one: reading the body
# off it would scan "" and release every buffered delta unscanned
chunks = (delta,) if chunks is None else (_responses_created_event(), delta)
guardrail = _surface_guardrail()
post = _armor_post_mock(_MODEL_ARMOR_CLEAN)
@ -4145,8 +4147,6 @@ async def test_streaming_hook_forwards_a_preceding_guardrails_error_item():
async def test_streaming_hook_forwards_a_preceding_guardrails_error_frame(chunks):
"""Chained post_call guardrails hand each other their output. An earlier guardrail's error
frame carries no message to assemble, and replacing it would hide the real refusal."""
from litellm.proxy.guardrails.anthropic_sse import anthropic_sse_error_frames
if chunks is None:
chunks = anthropic_sse_error_frames("Streaming response blocked by the first guardrail")
guardrail = _surface_guardrail()
@ -4163,17 +4163,167 @@ async def test_streaming_hook_forwards_a_preceding_guardrails_error_frame(chunks
async def test_streaming_responses_error_falls_back_to_sse_when_the_handler_declines():
"""build_stream_error_items may return None, which must not swallow the block into a clean
200: the refusal falls back to the chat-completions SSE form that still carries the status."""
guardrail = _surface_guardrail()
from litellm.proxy.guardrails.guardrail_hooks.model_armor.model_armor import (
_StreamSurface,
)
class _DecliningGuardrail(ModelArmorGuardrail):
@staticmethod
def _build_responses_error_items(exc):
return None
guardrail = _DecliningGuardrail(
template_id="test-template",
project_id="test-project",
location="us-central1",
guardrail_name="model-armor-test",
)
exc = HTTPException(status_code=400, detail={"message": "blocked"})
with patch.object(OpenAIResponsesHandler, "build_stream_error_items", return_value=None):
items = guardrail._stream_error_items(
exc,
{"message": "blocked", "code": "400"},
raw_sse=False,
responses_stream=True,
)
items = guardrail._stream_error_items(exc, surface=_StreamSurface.RESPONSES)
assert len(items) == 1
assert '"code": "400"' in items[0]
assert "blocked" in items[0]
def _responses_created_event():
from litellm.types.llms.openai import (
ResponseCreatedEvent,
ResponsesAPIResponse,
ResponsesAPIStreamEvents,
)
return ResponseCreatedEvent(
type=ResponsesAPIStreamEvents.RESPONSE_CREATED,
response=ResponsesAPIResponse(
id="resp_1",
created_at=0,
model="gpt-4o-mini",
object="response",
output=[],
parallel_tool_calls=False,
tool_choice="auto",
tools=[],
),
)
@pytest.mark.asyncio
async def test_streaming_hook_refuses_an_opaque_sse_stream_without_anthropic_framing():
"""/v1/messages is not the only endpoint that streams raw SSE: the Google generateContent
route marks its own stream raw too. Its frames carry no Anthropic event types, so refusing
them in Anthropic's format would hand a Google client a body it cannot parse."""
guardrail = _surface_guardrail()
post = _armor_post_mock(_MODEL_ARMOR_CLEAN)
chunks = (b'data: {"candidates":[{"content":{"parts":[{"text":"my card is 4111"}]}}]}\n\n',)
with patch.object(guardrail.async_handler, "post", post):
delivered = await _drain_surface_hook(guardrail, chunks)
post.assert_not_called()
assert tuple(delivered) != chunks
body = "".join(item.decode() if isinstance(item, bytes) else item for item in delivered)
assert "could not be assembled for scanning" in body
assert "event: error" not in body
assert '"code": "500"' in body
@pytest.mark.asyncio
async def test_streaming_unassemblable_stream_is_forwarded_when_fail_on_error_is_disabled():
"""fail_on_error: false is a deliberate choice to degrade open, and it governs every other
path in this hook. The fail-closed refusal has to honour it too."""
guardrail = _surface_guardrail(fail_on_error=False)
post = _armor_post_mock(_MODEL_ARMOR_CLEAN)
chunks = (
b'event: content_block_delta\ndata: {"type":"content_block_delta","index":0,'
b'"delta":{"type":"text_delta","text":"hi"}}\n\n',
)
with patch.object(guardrail.async_handler, "post", post):
delivered = await _drain_surface_hook(guardrail, chunks)
post.assert_not_called()
assert tuple(delivered) == chunks
@pytest.mark.asyncio
async def test_streaming_fail_closed_records_the_applied_guardrail():
"""A refusal that no header or log attributes to the guardrail leaves on-call unable to tell
a guardrail block apart from a provider failure."""
guardrail = _surface_guardrail()
request_data = {
"model": "claude-haiku",
"messages": [{"role": "user", "content": "show me a card"}],
"metadata": {"guardrails": ["model-armor-test"]},
}
chunks = (
b'event: content_block_delta\ndata: {"type":"content_block_delta","index":0,'
b'"delta":{"type":"text_delta","text":"hi"}}\n\n',
)
with patch.object(guardrail.async_handler, "post", _armor_post_mock(_MODEL_ARMOR_CLEAN)):
await _drain_surface_hook(guardrail, chunks, request_data=request_data)
assert request_data["metadata"]["applied_guardrails"] == ["model-armor-test"]
@pytest.mark.asyncio
async def test_streaming_responses_tool_call_output_is_scanned():
"""An agentic /v1/responses turn can carry its whole payload in tool-call arguments, which
is what the chat surface already folds into the scanned text."""
from litellm.types.llms.openai import (
ResponseCompletedEvent,
ResponsesAPIResponse,
ResponsesAPIStreamEvents,
)
guardrail = _surface_guardrail()
post = _armor_post_mock(_MODEL_ARMOR_CLEAN)
completed = ResponseCompletedEvent(
type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
response=ResponsesAPIResponse(
id="resp_1",
created_at=0,
model="gpt-4o-mini",
object="response",
output=[
{
"type": "function_call",
"id": "fc_1",
"call_id": "call_1",
"name": "send_email",
"arguments": '{"body": "my card is 4111-1111-1111-1111"}',
}
],
parallel_tool_calls=False,
tool_choice="auto",
tools=[],
),
)
with patch.object(guardrail.async_handler, "post", post):
delivered = await _drain_surface_hook(guardrail, (completed,))
post.assert_called_once()
scanned = post.call_args.kwargs["json"]["modelResponseData"]["text"]
assert "4111-1111-1111-1111" in scanned
assert tuple(delivered) == (completed,)
@pytest.mark.asyncio
async def test_streaming_hook_refuses_a_content_stream_that_ends_with_an_error_frame():
"""The chain-aware passthrough must stay narrow. A stream carrying real content plus a
trailing error frame is not a bare refusal to forward: the assembler cannot read it, and
releasing it would ship the buffered content unscanned."""
guardrail = _surface_guardrail()
post = _armor_post_mock(_MODEL_ARMOR_CLEAN)
chunks = (*_ANTHROPIC_SSE_CHUNKS, *anthropic_sse_error_frames("upstream gave up"))
with patch.object(guardrail.async_handler, "post", post):
delivered = await _drain_surface_hook(guardrail, chunks)
post.assert_not_called()
body = b"".join(delivered)
assert b"4111-1111-1111-1111" not in body
assert b"could not be assembled for scanning" in body