fix(model_armor): handle Anthropic Messages and Responses streams in post_call (#39181)

* fix(model_armor): handle Anthropic Messages and Responses streams in post_call

The post_call streaming hook buffered every chunk and fed it to
stream_chunk_builder, which only understands chat-completion deltas.
/v1/messages streams raw Anthropic SSE bytes and /v1/responses streams
typed Responses events, so both raised litellm.APIError and surfaced to
the client as a 500 on every streamed request.

Assemble each surface with its own reader, frame guardrail failures as
terminal items in that surface's wire format, and pass the stream
through unscanned when it cannot be assembled instead of raising.

* 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.

* test(model_armor): pin the error-only stream predicate against content-carrying streams

is_sse_error_stream decides whether a buffered stream is forwarded to the client
untouched, so a stream that still carries content must not qualify: the frames-only
join drops typed chunks, an empty stream is not a refusal, and a content event may
carry an empty error field.

* fix(model_armor): let a streamed de-identify match mask instead of blocking

A de-identify template reports MATCH_FOUND for every redaction it makes. The
streaming block check omitted allow_sanitization, so with mask_response_content
enabled that match read as a refusal and the client got a 400 where the
non-streaming sibling returned the redacted text. Pass the flag through, as the
non-streaming hook already does, and stamp the logged status from the same
decision so the spend row agrees with what the client received.

Also drop Any from the chat-completion assembler's parameter; stream_chunk_builder
takes a bare list, so list[object] carries the mutability requirement without
erasing the element type.

* fix(model_armor): fail closed when a streamed de-identify match cannot be applied

Allowing sanitization past the streaming block check is a promise to apply the
redaction Model Armor asked for. Two paths broke that promise and released the
buffered original instead: a match that comes back with no sanitized text, and a
surface with no assembled body to rewrite.

The outcome is now resolved once, before it is recorded, so the status stamped on
request metadata agrees with what the client receives rather than reporting the
success the block check alone would have implied.

* fix: scan the deltas when a Responses stream ends without a body

response.failed and response.incomplete are terminal events like
response.completed, but a turn that broke mid-generation reports an empty
output while the deltas ahead of it already spelled the answer out to the
client. Reading only the terminal body found nothing to scan there, and the
empty-content shortcut then forwarded every buffered delta past the guardrail.

Fall back to the text the delta events carry whenever a Responses stream
assembles to nothing.

* fix: read the Responses delta event types off the event enum

The hand-listed set left out response.mcp_call_arguments.delta, so a turn that
streamed only MCP tool arguments and then reported an empty body still took the
no-content shortcut and forwarded those chunks unscanned.

Deriving the set from ResponsesAPIStreamEvents keeps it complete as the enum
grows, and the str guard in the reader already covers any event whose delta is
not text.

* fix(model_armor): scan responses deltas alongside the terminal body

A /v1/responses stream spells out reasoning summaries and tool-call arguments in
delta events that its terminal body never repeats, so scanning the body alone
handed every summary delta to the client unscanned whenever the body carried text.

* fix(model_armor): scan responses delta fields apart from each other

A Responses turn spells out its reasoning summary, its visible answer and its tool-call
arguments in separate delta events. Joining every delta into one string let a finding form
across the boundary between two fields that each carry nothing to find, so a safe stream
could be blocked. Group the deltas by the field they belong to, join a field's own deltas
as they streamed, and keep the fields apart.

* fix(model_armor): scan each responses field once, not twice

Separating delta fields stopped the terminal body from matching the delta text, so a turn
with two visible fields sent Model Armor both copies. Only the delta fields the body does not
already carry are appended now.

---------

Co-authored-by: yassin <yassin@berri.ai>
This commit is contained in:
yucheng-berri 2026-09-02 19:15:14 -07:00 • committed by GitHub
parent 8065ede40b
commit 7a81ae98e6
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 1589 additions and 105 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:
@ -111,6 +144,27 @@ def anthropic_sse_error_frames(message: str) -> tuple[bytes, ...]:
)
def is_sse_error_stream(all_chunks: Sequence[object]) -> bool:
"""Whether the buffered stream carries nothing but error frames.
post_call guardrails run in a chain, so a hook can be handed the terminal error frames an
earlier guardrail emitted when it blocked. Those carry no message to assemble, and replacing
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
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
)
def anthropic_sse_chunks_from_response(assembled: ModelResponse) -> tuple[bytes, ...]:
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (
LiteLLMAnthropicMessagesAdapter,

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
@ -27,12 +28,25 @@ from litellm.llms.custom_httpx.http_handler import (
)
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,
)
from litellm.proxy.guardrails.guardrail_hooks.model_armor.file_scanning import (
MODEL_ARMOR_MAX_FILE_SIZE_BYTES,
plan_file_scans,
)
from litellm.types.guardrails import GuardrailEventHooks, LitellmParams
from litellm.types.llms.openai import AllMessageValues
from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionToolCallChunk,
ResponsesAPIResponse,
ResponsesAPIStreamEvents,
)
from litellm.types.utils import (
CallTypes,
CallTypesLiteral,
@ -41,10 +55,33 @@ from litellm.types.utils import (
ModelResponse,
ModelResponseStream,
StandardLoggingGuardrailInformation,
TextCompletionResponse,
)
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"})
# Every event whose ``delta`` is model output already on its way to the client. Read off the event
# enum rather than listed, so an event added there cannot quietly fall out of the scan
_RESPONSES_DELTA_EVENT_TYPES: Final = frozenset(
event.value for event in ResponsesAPIStreamEvents if event.value.endswith(".delta")
)
# What makes two delta events part of the same field of the turn, rather than two fields that merely
# streamed next to each other
_RESPONSES_DELTA_FIELD_ATTRS: Final = ("type", "item_id", "output_index", "content_index", "summary_index")
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
@ -322,19 +359,9 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
else:
return {"modelResponseData": {"byteItem": {"byteDataType": file_type, "byteData": base64_data}}}
def _should_block_content(self, armor_response: dict, allow_sanitization: bool = False) -> bool:
def _should_block_content(self, armor_response: Mapping[str, Any], allow_sanitization: bool = False) -> bool:
"""Check if Model Armor response indicates content should be blocked, including both inspectResult and deidentifyResult."""
sanitization_result: Final = armor_response.get("sanitizationResult", {})
filter_results: Final = sanitization_result.get("filterResults", {})
# filterResults can be a dict (named keys) or a list (array of filter result dicts)
filter_result_items = []
if isinstance(filter_results, dict):
filter_result_items = list(filter_results.values())
elif isinstance(filter_results, list):
filter_result_items = filter_results
for filt in filter_result_items:
for filt in self._filter_result_items(armor_response):
# Check RAI, PI/Jailbreak, Malicious URI, CSAM, Virus scan as before
if filt.get("raiFilterResult", {}).get("matchState") == "MATCH_FOUND":
return True
@ -358,22 +385,12 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
# Fallback dict code removed; all cases handled above
return False
def _get_sanitized_content(self, armor_response: dict) -> str | None:
def _get_sanitized_content(self, armor_response: Mapping[str, Any]) -> str | None:
"""
Get the sanitized content from a Model Armor response, if available.
Looks for sanitized text in deidentifyResult, and falls back to root-level fields if not found.
"""
result: Final = armor_response.get("sanitizationResult", {})
filter_results: Final = result.get("filterResults", {})
# filterResults can be a dict (single filter) or a list (multiple filters)
filters: Final = (
list(filter_results.values())
if isinstance(filter_results, dict)
else filter_results
if isinstance(filter_results, list)
else []
)
filters: Final = self._filter_result_items(armor_response)
# Prefer sanitized text from deidentifyResult if present
for filter_entry in filters:
@ -397,6 +414,61 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
# Fallback: if Model Armor put sanitized text at the root, use it
return armor_response.get("sanitizedText") or armor_response.get("text")
@staticmethod
def _filter_result_items(armor_response: Mapping[str, Any]) -> Sequence[Any]:
"""Every filter result in a scan response.
filterResults is a dict of named filters on most templates and a list on some, so both
shapes are flattened to the same list of filter entries.
"""
filter_results: Final = armor_response.get("sanitizationResult", {}).get("filterResults", {})
if isinstance(filter_results, dict):
return list(filter_results.values())
if isinstance(filter_results, list):
return filter_results
return []
def _has_deidentify_match(self, armor_response: Mapping[str, Any]) -> bool:
"""Whether an SDP de-identify filter matched, i.e. Model Armor owes this response a redaction."""
for filter_entry in self._filter_result_items(armor_response):
sdp = filter_entry.get("sdpFilterResult")
if sdp and sdp.get("deidentifyResult", {}).get("matchState") == "MATCH_FOUND":
return True
return False
def _resolve_streaming_outcome(
self,
armor_response: Mapping[str, Any],
assembled_response: object,
content: str,
) -> tuple[bool, str | None]:
"""Whether to block the buffered stream, and the rewrite to emit when it is not blocked.
A de-identify match only reaches here unblocked because masking is on, so the redaction it
stands for has to be both resolvable and emittable. Where it is neither, the buffered
original still carries what Model Armor matched on, so this fails closed instead of
releasing it.
"""
if self._should_block_content(armor_response, allow_sanitization=self.mask_response_content):
return True, None
if not self.mask_response_content:
return False, None
sanitized_content: Final = self._get_sanitized_content(armor_response)
if not sanitized_content:
# No rewrite to apply. Harmless unless a match is outstanding, in which case applying
# nothing would hand back the very content that matched
return self._has_deidentify_match(armor_response), None
if sanitized_content == content:
return False, None
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"
)
return True, None
return False, sanitized_content
@staticmethod
def _append_armor_response(existing: object, armor_response: Mapping[str, object]) -> object:
"""Accumulate scan responses so a later text scan does not drop an earlier file scan.
@ -831,6 +903,185 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
return response
@staticmethod
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.
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.
"""
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
):
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:
"""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_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, *(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."""
if isinstance(assembled_response, ResponsesAPIResponse):
return self._responses_api_response_text(assembled_response)
return self._extract_content_from_response(assembled_response)
@staticmethod
def _responses_delta_field(chunk: object) -> tuple[str, ...]:
"""Which field of the turn a delta event belongs to."""
return tuple(str(getattr(chunk, attr, None)) for attr in _RESPONSES_DELTA_FIELD_ATTRS)
@staticmethod
def _responses_delta_field_texts(all_chunks: Sequence[object]) -> tuple[str, ...]:
"""Text each field of a ``/v1/responses`` turn has already spelled out in its delta events.
One field's deltas are joined as they streamed, since a finding can be split across them,
and separate fields stay apart, so a reasoning summary running into the visible answer
cannot spell out a finding that neither of them carries.
"""
deltas: Final = tuple(
(ModelArmorGuardrail._responses_delta_field(chunk), delta)
for chunk in all_chunks
if getattr(chunk, "type", None) in _RESPONSES_DELTA_EVENT_TYPES
and isinstance(delta := getattr(chunk, "delta", None), str)
)
return tuple(
"".join(delta for field, delta in deltas if field == streamed_field)
for streamed_field in dict.fromkeys(field for field, _ in deltas)
)
def _streaming_content_to_scan(
self,
assembled_response: object,
all_chunks: Sequence[object],
surface: _StreamSurface,
) -> str:
"""Text to scan for a buffered stream, which is everything the client is about to receive.
A ``/v1/responses`` stream also spells out reasoning summaries and tool-call arguments in
delta events that its terminal body never repeats, so every delta field the body does not
already carry is scanned after it.
"""
content: Final = self._extract_streaming_content(assembled_response)
if surface is not _StreamSurface.RESPONSES:
return content
unscanned: Final = tuple(text for text in self._responses_delta_field_texts(all_chunks) if text not in content)
return "\n".join(part for part in (content, *unscanned) if part)
@staticmethod
def _apply_sanitized_content(assembled_response: ModelResponse, sanitized_content: str) -> None:
"""Replace every non-empty choice message with the Model Armor sanitized text."""
for choice in assembled_response.choices:
if isinstance(choice, Choices) and choice.message.content:
choice.message.content = sanitized_content
@staticmethod
def _assemble_chat_completion_stream(
all_chunks: list[object], # mutable-ok: stream_chunk_builder only accepts a mutable list
) -> ModelResponse | TextCompletionResponse | None:
"""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: 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 _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."""
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,
user_api_key_dict: UserAPIKeyAuth,
@ -840,97 +1091,125 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
"""Process streaming response chunks."""
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
from litellm.main import stream_chunk_builder
from litellm.proxy.common_utils.callback_utils import (
add_guardrail_to_applied_guardrails_header,
)
# Collect all chunks
all_chunks: Final[list[ModelResponseStream]] = []
all_chunks: Final[list[Any]] = []
async for chunk in response:
all_chunks.append(chunk)
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 = stream_chunk_builder(chunks=all_chunks)
assembled_response: Final = self._assemble_stream(all_chunks, surface)
if isinstance(assembled_response, ModelResponse):
# Extract content
content: Final = self._extract_content_from_response(assembled_response)
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 content:
try:
# Check with Model Armor
armor_response: Final = await self.make_model_armor_request(
content=content,
source="model_response",
request_data=request_data,
)
# Forwarding an unscannable stream would silently disable the guardrail, so fail closed
add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name)
for error_item in self._stream_error_items(
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
# 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"
)
# Extract content
content: Final = self._streaming_content_to_scan(
assembled_response=assembled_response, all_chunks=all_chunks, surface=surface
)
# 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,
)
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
add_guardrail_to_applied_guardrails_header(
request_data=request_data, guardrail_name=self.guardrail_name
)
try:
# Check with Model Armor
armor_response: Final = await self.make_model_armor_request(
content=content,
source="model_response",
request_data=request_data,
)
# 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,
),
)
# Decide the outcome before recording it. Mirrors the non-streaming sibling: with
# masking on, a de-identify match is a redaction to apply rather than a refusal, but
# that only holds while the redaction can actually be delivered
blocked, sanitized_content = self._resolve_streaming_outcome(
armor_response=armor_response,
assembled_response=assembled_response,
content=content,
)
# 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:
# Update assembled response
for choice in assembled_response.choices:
if isinstance(choice, Choices):
if choice.message.content:
choice.message.content = sanitized_content
# 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 blocked else "success"
# Return sanitized stream
mock_response: Final = MockResponseIterator(model_response=assembled_response)
async for chunk in mock_response:
yield chunk
return
# 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)
except ModelArmorAPIError as e:
if self.optional_params.get("fail_on_error", True):
error_obj = {"message": e.detail, "code": "500"}
yield f"data: {json.dumps({'error': error_obj})}\n\n"
return
except HTTPException as e:
# Yield error as SSE event 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)
yield f"data: {json.dumps({'error': error_obj})}\n\n"
if blocked:
raise HTTPException(
status_code=400,
detail=self._build_block_error_detail(
"Streaming response blocked by Model Armor",
armor_response,
),
)
if sanitized_content is not None and isinstance(assembled_response, ModelResponse):
self._apply_sanitized_content(assembled_response, sanitized_content)
# 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 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")
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):
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: