mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
67a9932f7d
commit
4b21583005
3 changed files with 430 additions and 202 deletions
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue