mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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.
This commit is contained in:
parent
ee0c958ea2
commit
06f375c5d7
2 changed files with 105 additions and 1 deletions
|
|
@ -62,6 +62,17 @@ 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
|
||||
_RESPONSES_DELTA_EVENT_TYPES: Final = frozenset(
|
||||
{
|
||||
"response.output_text.delta",
|
||||
"response.refusal.delta",
|
||||
"response.function_call_arguments.delta",
|
||||
"response.custom_tool_call_input.delta",
|
||||
"response.reasoning_summary_text.delta",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class _StreamSurface(Enum):
|
||||
"""Wire format of a buffered streaming response, which decides how it is read and how it is refused."""
|
||||
|
|
@ -967,6 +978,34 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
return self._responses_api_response_text(assembled_response)
|
||||
return self._extract_content_from_response(assembled_response)
|
||||
|
||||
@staticmethod
|
||||
def _responses_delta_text(all_chunks: Sequence[object]) -> str:
|
||||
"""Text a ``/v1/responses`` stream has already spelled out in its delta events."""
|
||||
return "".join(
|
||||
delta
|
||||
for chunk in all_chunks
|
||||
if getattr(chunk, "type", None) in _RESPONSES_DELTA_EVENT_TYPES
|
||||
and isinstance(delta := getattr(chunk, "delta", None), str)
|
||||
)
|
||||
|
||||
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 whatever the client is about to receive.
|
||||
|
||||
``response.failed`` and ``response.incomplete`` are terminal like ``response.completed`` but
|
||||
report a turn that broke mid-generation, so their body can be empty while the deltas ahead of
|
||||
them already spelled the answer out. Reading the body alone finds nothing to scan there and
|
||||
releases those deltas untouched, so an empty body falls back to the deltas themselves.
|
||||
"""
|
||||
content: Final = self._extract_streaming_content(assembled_response)
|
||||
if content or surface is not _StreamSurface.RESPONSES:
|
||||
return content
|
||||
return self._responses_delta_text(all_chunks)
|
||||
|
||||
@staticmethod
|
||||
def _apply_sanitized_content(assembled_response: ModelResponse, sanitized_content: str) -> None:
|
||||
"""Replace every non-empty choice message with the Model Armor sanitized text."""
|
||||
|
|
@ -1080,7 +1119,9 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
return
|
||||
|
||||
# Extract content
|
||||
content: Final = self._extract_streaming_content(assembled_response)
|
||||
content: Final = self._streaming_content_to_scan(
|
||||
assembled_response=assembled_response, all_chunks=all_chunks, surface=surface
|
||||
)
|
||||
|
||||
if not content:
|
||||
verbose_proxy_logger.debug("Model Armor: No text content in streaming response, skipping guardrail")
|
||||
|
|
|
|||
|
|
@ -4552,3 +4552,66 @@ async def test_streaming_status_records_a_surface_that_cannot_carry_the_rewrite_
|
|||
assert b"4111-1111-1111-1111" not in body
|
||||
assert b"Streaming response blocked by Model Armor" in body
|
||||
assert request_data["metadata"]["_model_armor_status"] == "blocked"
|
||||
|
||||
|
||||
def _responses_api_events_truncated(terminal: str):
|
||||
"""A /v1/responses stream whose text went out as deltas and whose terminal event reports no body.
|
||||
|
||||
``response.failed`` and ``response.incomplete`` are terminal 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.
|
||||
"""
|
||||
from litellm.types.llms.openai import (
|
||||
OutputTextDeltaEvent,
|
||||
ResponseFailedEvent,
|
||||
ResponseIncompleteEvent,
|
||||
ResponsesAPIResponse,
|
||||
ResponsesAPIStreamEvents,
|
||||
)
|
||||
|
||||
empty_body = ResponsesAPIResponse(
|
||||
id="resp_1",
|
||||
created_at=0,
|
||||
model="gpt-4o-mini",
|
||||
object="response",
|
||||
output=[],
|
||||
parallel_tool_calls=False,
|
||||
tool_choice="auto",
|
||||
tools=[],
|
||||
)
|
||||
terminal_event = (
|
||||
ResponseFailedEvent(type=ResponsesAPIStreamEvents.RESPONSE_FAILED, response=empty_body)
|
||||
if terminal == "failed"
|
||||
else ResponseIncompleteEvent(type=ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE, response=empty_body)
|
||||
)
|
||||
return (
|
||||
OutputTextDeltaEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA,
|
||||
item_id="msg_1",
|
||||
output_index=0,
|
||||
content_index=0,
|
||||
delta="my card is 4111-1111-1111-1111",
|
||||
),
|
||||
terminal_event,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("terminal", ["failed", "incomplete"])
|
||||
async def test_streaming_responses_terminal_event_without_a_body_still_scans_the_deltas(terminal):
|
||||
"""A /v1/responses turn that broke mid-generation has still delivered its deltas.
|
||||
|
||||
Reading only the terminal body would find nothing to scan and hand every buffered delta to the
|
||||
client untouched, so the deltas themselves are what gets scanned.
|
||||
"""
|
||||
guardrail = _surface_guardrail()
|
||||
post = _armor_post_mock(_MODEL_ARMOR_BLOCK)
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", post):
|
||||
delivered = await _drain_surface_hook(guardrail, _responses_api_events_truncated(terminal))
|
||||
|
||||
post.assert_called_once()
|
||||
assert "4111-1111-1111-1111" in post.call_args.kwargs["json"]["modelResponseData"]["text"]
|
||||
rendered = "".join(str(item) for item in delivered)
|
||||
assert "4111-1111-1111-1111" not in rendered
|
||||
assert "Streaming response blocked by Model Armor" in rendered
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue