mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(policy_engine): deliver guardrail text rewrites on multi-choice, unfinished, and envelope-less streams
Post-call pipeline rewrites on buffered streams failed open on three shapes: chat streams with n > 1 (the rebuilt response collapsed every choice into index 0), streams that ended without a finish marker, and Responses streams whose final event carried no response envelope. The chat handler now rebuilds the ended stream one choice index at a time and writes each choice's rewrite back to that choice's buffered deltas. The Anthropic handler writes an unended stream's rewrite across its text deltas. The Responses handler spreads an envelope-less rewrite over the buffered output_text events, still failing open when a scanned event cannot be placed. Tool-call rewrites on n > 1 chat streams keep failing open.
This commit is contained in:
parent
12ddb35aad
commit
b3d9ba9e7b
7 changed files with 311 additions and 92 deletions
|
|
@ -1253,10 +1253,9 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
Process output streaming response by applying guardrails to text content.
|
||||
|
||||
Get the string so far, check the apply guardrail to the string so far, and return the list of responses so far.
|
||||
With ``deliver_ended_stream_rewrites``, an ended stream whose guardrail rewrote the text gets the rewrite
|
||||
written back across the buffered chunks (full rewritten text in the first ``text_delta``, the rest blanked);
|
||||
a rewrite on a stream that never reported a ``stop_reason`` has no write-back and is reported as
|
||||
undeliverable, so the pipeline executor discards it and releases the original chunks.
|
||||
With ``deliver_ended_stream_rewrites``, a stream whose guardrail rewrote the text gets the rewrite
|
||||
written back across the buffered chunks (full rewritten text in the first ``text_delta``, the rest blanked),
|
||||
whether or not the stream ever reported a ``stop_reason``.
|
||||
"""
|
||||
from litellm.integrations.custom_guardrail import ModifyResponseException
|
||||
|
||||
|
|
@ -1354,9 +1353,7 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
raise
|
||||
unended_texts: Final = _guardrailed_inputs.get("texts")
|
||||
if deliver_ended_stream_rewrites and unended_texts and tuple(unended_texts) != (string_so_far,):
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
||||
|
||||
raise UndeliverableStreamRewrite(guardrail_to_apply.guardrail_name or "unknown")
|
||||
self._write_ended_stream_text_rewrite(responses_so_far, unended_texts[0])
|
||||
return responses_so_far
|
||||
|
||||
def _prepare_request_data(
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ import json
|
|||
import time
|
||||
import uuid
|
||||
from collections.abc import Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Union, cast
|
||||
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
|
@ -651,10 +652,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
"""Ended-stream path: rebuild the full response, run the non-streaming
|
||||
output guardrail against it, and (when opted in) write any text or
|
||||
tool-call rewrite back across the buffered chunks."""
|
||||
model_response: Final = cast(
|
||||
ModelResponse,
|
||||
stream_chunk_builder(chunks=responses_so_far, logging_obj=litellm_logging_obj),
|
||||
)
|
||||
model_response: Final = self._rebuild_ended_stream_per_choice(responses_so_far, litellm_logging_obj)
|
||||
pre_guardrail_texts: Final = self._string_choice_contents(model_response)
|
||||
pre_guardrail_tool_calls: Final = self._function_tool_call_shapes(model_response)
|
||||
await self.process_output_response(
|
||||
|
|
@ -666,18 +664,61 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
)
|
||||
if not deliver_ended_stream_rewrites:
|
||||
return
|
||||
guardrail_name: Final = guardrail_to_apply.guardrail_name or "unknown"
|
||||
await self._write_ended_stream_text_rewrites(
|
||||
responses_so_far=responses_so_far,
|
||||
guardrailed_response=model_response,
|
||||
pre_guardrail_texts=pre_guardrail_texts,
|
||||
guardrail_name=guardrail_name,
|
||||
)
|
||||
self._write_ended_stream_tool_call_rewrites(
|
||||
responses_so_far=responses_so_far,
|
||||
guardrailed_response=model_response,
|
||||
pre_guardrail_tool_calls=pre_guardrail_tool_calls,
|
||||
guardrail_name=guardrail_name,
|
||||
guardrail_name=guardrail_to_apply.guardrail_name or "unknown",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _rebuild_ended_stream_per_choice(
|
||||
responses_so_far: Sequence["ModelResponseStream"],
|
||||
litellm_logging_obj: "LiteLLMLoggingObj | None",
|
||||
) -> "ModelResponse":
|
||||
"""``stream_chunk_builder`` folds every choice of a stream into one index-0
|
||||
choice, so the stream is rebuilt one choice index at a time (every chunk
|
||||
kept, its choices narrowed to that index, so usage-only chunks still
|
||||
count) and the rebuilt choices are stitched into one response, each
|
||||
carrying the index the stream gave it."""
|
||||
choice_indices: Final = tuple(
|
||||
sorted(frozenset(choice.index for response in responses_so_far for choice in response.choices))
|
||||
)
|
||||
rebuilt_by_index: Final = tuple(
|
||||
(
|
||||
index,
|
||||
cast(
|
||||
ModelResponse,
|
||||
stream_chunk_builder(
|
||||
chunks=[ # mutable-ok: callee takes a list
|
||||
response.model_copy(
|
||||
update=MappingProxyType(
|
||||
{"choices": tuple(choice for choice in response.choices if choice.index == index)}
|
||||
)
|
||||
)
|
||||
for response in responses_so_far
|
||||
],
|
||||
logging_obj=litellm_logging_obj,
|
||||
),
|
||||
),
|
||||
)
|
||||
for index in choice_indices
|
||||
)
|
||||
(_, base_response), *_ = rebuilt_by_index
|
||||
return base_response.model_copy(
|
||||
update=MappingProxyType(
|
||||
{
|
||||
"choices": tuple(
|
||||
rebuilt.choices[0].model_copy(update=MappingProxyType({"index": index}))
|
||||
for index, rebuilt in rebuilt_by_index
|
||||
)
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
def build_stream_error_items(
|
||||
|
|
@ -1058,39 +1099,28 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
responses_so_far: list["ModelResponseStream"], # mutable-ok: rewrites the caller's buffered chunks in place
|
||||
guardrailed_response: "ModelResponse",
|
||||
pre_guardrail_texts: tuple[str | None, ...],
|
||||
guardrail_name: str,
|
||||
) -> None:
|
||||
"""Write ended-stream guardrail text rewrites back across the buffered
|
||||
chunks: the full rewritten text lands in the choice's first
|
||||
content-carrying chunk and the rest are blanked, the same shape the
|
||||
in-flight write-back uses. Chunks carrying only finish_reason or usage
|
||||
stay untouched. A rewrite on a stream carrying more than one distinct
|
||||
choice index is reported as undeliverable, so the pipeline executor
|
||||
discards it and releases the original chunks."""
|
||||
chunks, one rewrite per rebuilt choice index: the full rewritten text
|
||||
lands in that choice's first content-carrying chunk and the rest are
|
||||
blanked, the same shape the in-flight write-back uses. Chunks carrying
|
||||
only finish_reason or usage stay untouched."""
|
||||
post_guardrail_texts: Final = self._string_choice_contents(guardrailed_response)
|
||||
changed: Final = tuple(
|
||||
after
|
||||
for before, after in zip(pre_guardrail_texts, post_guardrail_texts)
|
||||
if before is not None and after is not None and after != before
|
||||
rewrites_by_choice: Final = MappingProxyType(
|
||||
{
|
||||
choice.index: after
|
||||
for choice, before, after in zip(
|
||||
guardrailed_response.choices, pre_guardrail_texts, post_guardrail_texts
|
||||
)
|
||||
if before is not None and after is not None and after != before
|
||||
}
|
||||
)
|
||||
if not changed:
|
||||
if not rewrites_by_choice:
|
||||
return
|
||||
stream_choice_indices: Final = frozenset(
|
||||
choice.index for response in responses_so_far for choice in response.choices
|
||||
)
|
||||
if len(stream_choice_indices) != 1:
|
||||
# stream_chunk_builder collapses every choice into one index-0
|
||||
# choice, so a rewrite of the rebuilt response cannot be attributed
|
||||
# back to a single choice on an n>1 stream: report it undeliverable
|
||||
# rather than deliver the rewrite on the wrong choice
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
||||
|
||||
raise UndeliverableStreamRewrite(guardrail_name)
|
||||
target_choice_index: Final = next(iter(stream_choice_indices))
|
||||
await self._apply_guardrail_responses_to_output_streaming(
|
||||
responses=responses_so_far,
|
||||
guardrailed_texts=list(changed), # mutable-ok: callee takes lists
|
||||
task_mappings=[(target_choice_index, None) for _ in changed], # mutable-ok: callee takes lists
|
||||
guardrailed_texts=list(rewrites_by_choice.values()), # mutable-ok: callee takes lists
|
||||
task_mappings=[(index, None) for index in rewrites_by_choice], # mutable-ok: callee takes lists
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -209,6 +209,7 @@ _TOOL_CALL_PAYLOAD_EVENT_TYPES: Final = _TOOL_CALL_PAYLOAD_DELTA_EVENT_TYPES | f
|
|||
_TOOL_CALL_PAYLOAD_DONE_EVENT_FIELDS
|
||||
)
|
||||
_OUTPUT_ITEM_EVENT_TYPES: Final = frozenset({"response.output_item.added", "response.output_item.done"})
|
||||
_OUTPUT_TEXT_EVENT_TYPES: Final = frozenset({"response.output_text.delta", "response.output_text.done"})
|
||||
_PATCHABLE_ITEM_FIELDS: Final[Mapping[str, str]] = MappingProxyType(
|
||||
{"function_call_output": "output", "message": "content"}
|
||||
)
|
||||
|
|
@ -832,9 +833,10 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
(``response.output_text.delta`` / ``.done``,
|
||||
``response.content_part.done``, ``response.output_item.done``) are synced
|
||||
to the rewritten envelope too, so a client reading deltas sees the
|
||||
rewrite instead of the raw model output; a rewrite observed where no
|
||||
write-back is possible is reported as undeliverable, so the pipeline
|
||||
executor discards it and releases the original events.
|
||||
rewrite instead of the raw model output; a stream with no envelope
|
||||
gets its rewrite spread over the buffered text events, and a rewrite
|
||||
observed where no write-back is possible is reported as undeliverable,
|
||||
so the pipeline executor discards it and releases the original events.
|
||||
"""
|
||||
if not responses_so_far:
|
||||
return responses_so_far
|
||||
|
|
@ -958,10 +960,9 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
return responses_so_far
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# Fallback: apply guardrail to the accumulated text string. #
|
||||
# No structured write-back is possible here; guardrails that only #
|
||||
# need to block/flag (not rewrite) still work correctly, and a #
|
||||
# rewrite a caller expects delivered is reported undeliverable. #
|
||||
# Fallback: apply guardrail to the accumulated text string. With no #
|
||||
# envelope to rewrite, a rewrite a caller expects delivered is spread #
|
||||
# over the buffered text events instead. #
|
||||
# ------------------------------------------------------------------ #
|
||||
string_so_far: Final = self.get_streaming_string_so_far(responses_so_far)
|
||||
if string_so_far:
|
||||
|
|
@ -979,11 +980,54 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
)
|
||||
fallback_texts: Final = fallback_outputs.get("texts")
|
||||
if deliver_ended_stream_rewrites and fallback_texts and tuple(fallback_texts) != (string_so_far,):
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
||||
|
||||
raise UndeliverableStreamRewrite(guardrail_to_apply.guardrail_name or "unknown")
|
||||
self._spread_text_rewrite_over_stream_events(
|
||||
stream_events=responses_so_far,
|
||||
rewritten_text=fallback_texts[0],
|
||||
guardrail_name=guardrail_to_apply.guardrail_name or "unknown",
|
||||
)
|
||||
return responses_so_far
|
||||
|
||||
def _spread_text_rewrite_over_stream_events(
|
||||
self,
|
||||
stream_events: Sequence[Any],
|
||||
rewritten_text: str,
|
||||
guardrail_name: str,
|
||||
) -> None:
|
||||
"""Deliver a text rewrite on a stream with no completed envelope by
|
||||
spreading it over the text parts the guardrail scanned, in stream
|
||||
order: the whole rewrite on the first part and every later part
|
||||
blanked, through the same sync the envelope path uses. A scanned
|
||||
event the sync cannot place (one that is not an ``output_text`` delta
|
||||
or done, or lacks integer ``output_index`` / ``content_index``) makes
|
||||
the rewrite undeliverable, so the pipeline executor discards it and
|
||||
releases the original events."""
|
||||
scanned_events: Final = tuple(
|
||||
event
|
||||
for event in stream_events
|
||||
if isinstance(stream_item_field(event, "text"), str) or isinstance(stream_item_field(event, "delta"), str)
|
||||
)
|
||||
scanned_positions: Final = tuple(
|
||||
dict.fromkeys(
|
||||
(stream_item_field(event, "output_index"), stream_item_field(event, "content_index"))
|
||||
for event in scanned_events
|
||||
)
|
||||
)
|
||||
placeable_positions: Final = tuple(
|
||||
(output_index, content_index)
|
||||
for output_index, content_index in scanned_positions
|
||||
if isinstance(output_index, int) and isinstance(content_index, int)
|
||||
)
|
||||
if len(placeable_positions) != len(scanned_positions) or any(
|
||||
stream_item_field(event, "type") not in _OUTPUT_TEXT_EVENT_TYPES for event in scanned_events
|
||||
):
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
||||
|
||||
raise UndeliverableStreamRewrite(guardrail_name)
|
||||
self._sync_stream_events_with_rewrites(
|
||||
stream_events=stream_events,
|
||||
rewrites_by_position=MappingProxyType(dict(zip(placeable_positions, chain((rewritten_text,), repeat(""))))),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _write_event_field(event: object, field: str, value: str) -> None:
|
||||
if isinstance(event, dict):
|
||||
|
|
|
|||
|
|
@ -445,19 +445,22 @@ class TestAnthropicMessagesHandlerStreamingOutputProcessing:
|
|||
assert chunks == original
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unended_stream_rewrite_with_delivery_expected_fails_closed(self):
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
||||
|
||||
async def test_unended_stream_rewrite_with_delivery_expected_lands_in_the_buffered_deltas(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
chunks = self._ended_sse_chunks()[:-2]
|
||||
|
||||
with pytest.raises(UndeliverableStreamRewrite):
|
||||
await handler.process_output_streaming_response(
|
||||
responses_so_far=chunks,
|
||||
guardrail_to_apply=self._masking_guardrail(),
|
||||
litellm_logging_obj=MagicMock(),
|
||||
deliver_ended_stream_rewrites=True,
|
||||
)
|
||||
result = await handler.process_output_streaming_response(
|
||||
responses_so_far=chunks,
|
||||
guardrail_to_apply=self._masking_guardrail(),
|
||||
litellm_logging_obj=MagicMock(),
|
||||
deliver_ended_stream_rewrites=True,
|
||||
)
|
||||
|
||||
assert result is chunks
|
||||
assert self._delta_texts(chunks) == ["hello [MASKED]", ""]
|
||||
raw = b"".join(chunks).decode()
|
||||
assert "event: message_start" in raw and "event: content_block_stop" in raw
|
||||
assert "event: message_stop" not in raw
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unended_stream_without_rewrite_is_released_with_delivery_expected(self):
|
||||
|
|
|
|||
|
|
@ -1262,20 +1262,75 @@ class TestOpenAIChatCompletionsHandlerStreamingOutput:
|
|||
return MaskWorld(guardrail_name="test-mask")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deliver_ended_stream_rewrite_on_multi_choice_stream_fails_closed(self):
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
||||
|
||||
async def test_deliver_ended_stream_rewrite_lands_on_the_rewritten_choice_only(self):
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
chunks = self._two_choice_stream_chunks()
|
||||
|
||||
with pytest.raises(UndeliverableStreamRewrite):
|
||||
await handler.process_output_streaming_response(
|
||||
responses_so_far=chunks,
|
||||
guardrail_to_apply=self._world_masking_guardrail(),
|
||||
litellm_logging_obj=None,
|
||||
deliver_ended_stream_rewrites=True,
|
||||
result = await handler.process_output_streaming_response(
|
||||
responses_so_far=chunks,
|
||||
guardrail_to_apply=self._world_masking_guardrail(),
|
||||
litellm_logging_obj=None,
|
||||
deliver_ended_stream_rewrites=True,
|
||||
)
|
||||
|
||||
assert result is chunks
|
||||
assert [(c.choices[0].index, c.choices[0].delta.content) for c in chunks] == [
|
||||
(0, "safe "),
|
||||
(1, "hello [MASKED]"),
|
||||
(0, "text"),
|
||||
(1, ""),
|
||||
]
|
||||
assert [c.choices[0].finish_reason for c in chunks] == [None, None, "stop", "stop"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deliver_ended_stream_rewrites_each_choice_with_its_own_text(self):
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
guardrail = MockGuardrail(guardrail_name="test")
|
||||
chunks = self._two_choice_stream_chunks()
|
||||
|
||||
await handler.process_output_streaming_response(
|
||||
responses_so_far=chunks,
|
||||
guardrail_to_apply=guardrail,
|
||||
litellm_logging_obj=None,
|
||||
deliver_ended_stream_rewrites=True,
|
||||
)
|
||||
|
||||
assert guardrail.last_inputs["texts"] == ["safe text", "hello world"]
|
||||
assert [(c.choices[0].index, c.choices[0].delta.content) for c in chunks] == [
|
||||
(0, "SAFE TEXT"),
|
||||
(1, "HELLO WORLD"),
|
||||
(0, ""),
|
||||
(1, ""),
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deliver_rewrite_on_unfinished_stream_lands_in_the_buffered_deltas(self):
|
||||
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
|
||||
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
|
||||
def chunk(content: str) -> ModelResponseStream:
|
||||
return ModelResponseStream(
|
||||
id="chatcmpl-123",
|
||||
created=1234567890,
|
||||
model="gpt-4",
|
||||
object="chat.completion.chunk",
|
||||
choices=[StreamingChoices(index=0, delta=Delta(content=content), finish_reason=None)],
|
||||
)
|
||||
|
||||
chunks = [chunk("hello "), chunk("world")]
|
||||
|
||||
result = await handler.process_output_streaming_response(
|
||||
responses_so_far=chunks,
|
||||
guardrail_to_apply=self._world_masking_guardrail(),
|
||||
litellm_logging_obj=None,
|
||||
deliver_ended_stream_rewrites=True,
|
||||
)
|
||||
|
||||
assert result is chunks
|
||||
assert [c.choices[0].delta.content for c in chunks] == ["hello [MASKED]", ""]
|
||||
assert [c.choices[0].finish_reason for c in chunks] == [None, None]
|
||||
|
||||
@staticmethod
|
||||
def _two_choice_tool_call_stream_chunks() -> list:
|
||||
from litellm.types.utils import (
|
||||
|
|
|
|||
|
|
@ -1747,33 +1747,82 @@ class TestOpenAIResponsesHandlerStreamingOutputProcessing:
|
|||
assert events[5]["response"]["output"][0]["content"][0]["text"] == "hello [MASKED]"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fallback_rewrite_with_delivery_expected_fails_closed(self):
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
||||
|
||||
async def test_fallback_rewrite_with_delivery_expected_lands_in_the_delta_and_done_events(self):
|
||||
handler = OpenAIResponsesHandler()
|
||||
events = [
|
||||
{"type": "response.output_text.delta", "output_index": 0, "content_index": 0, "delta": "hello "},
|
||||
{"type": "response.output_text.done", "output_index": 0, "content_index": 0, "text": "hello world"},
|
||||
]
|
||||
|
||||
with pytest.raises(UndeliverableStreamRewrite):
|
||||
await handler.process_output_streaming_response(
|
||||
responses_so_far=events,
|
||||
guardrail_to_apply=self._masking_guardrail(),
|
||||
litellm_logging_obj=None,
|
||||
deliver_ended_stream_rewrites=True,
|
||||
)
|
||||
result = await handler.process_output_streaming_response(
|
||||
responses_so_far=events,
|
||||
guardrail_to_apply=self._masking_guardrail(),
|
||||
litellm_logging_obj=None,
|
||||
deliver_ended_stream_rewrites=True,
|
||||
)
|
||||
|
||||
assert result is events
|
||||
assert events[0]["delta"] == "hello [MASKED]"
|
||||
assert events[1]["text"] == "hello [MASKED]"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fallback_delta_only_rewrite_with_delivery_expected_fails_closed(self):
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
||||
|
||||
async def test_fallback_delta_only_rewrite_with_delivery_expected_spreads_over_the_deltas(self):
|
||||
handler = OpenAIResponsesHandler()
|
||||
events = [
|
||||
{"type": "response.output_text.delta", "output_index": 0, "content_index": 0, "delta": "hello "},
|
||||
{"type": "response.output_text.delta", "output_index": 0, "content_index": 0, "delta": "world"},
|
||||
]
|
||||
|
||||
result = await handler.process_output_streaming_response(
|
||||
responses_so_far=events,
|
||||
guardrail_to_apply=self._masking_guardrail(),
|
||||
litellm_logging_obj=None,
|
||||
deliver_ended_stream_rewrites=True,
|
||||
)
|
||||
|
||||
assert result is events
|
||||
assert [event["delta"] for event in events] == ["hello [MASKED]", ""]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fallback_rewrite_across_parts_lands_whole_on_the_first_part(self):
|
||||
handler = OpenAIResponsesHandler()
|
||||
events = [
|
||||
{"type": "response.output_text.delta", "output_index": 0, "content_index": 0, "delta": "hello "},
|
||||
{"type": "response.output_text.done", "output_index": 0, "content_index": 0, "text": "hello "},
|
||||
{"type": "response.output_text.delta", "output_index": 1, "content_index": 0, "delta": "wor"},
|
||||
{"type": "response.output_text.delta", "output_index": 1, "content_index": 0, "delta": "ld"},
|
||||
]
|
||||
guardrail = MockRecordingGuardrail(guardrail_name="test")
|
||||
|
||||
await handler.process_output_streaming_response(
|
||||
responses_so_far=events,
|
||||
guardrail_to_apply=guardrail,
|
||||
litellm_logging_obj=None,
|
||||
deliver_ended_stream_rewrites=True,
|
||||
)
|
||||
assert [inputs.get("texts") for inputs in guardrail.seen_inputs] == [["hello world"]]
|
||||
|
||||
await handler.process_output_streaming_response(
|
||||
responses_so_far=events,
|
||||
guardrail_to_apply=self._masking_guardrail(),
|
||||
litellm_logging_obj=None,
|
||||
deliver_ended_stream_rewrites=True,
|
||||
)
|
||||
|
||||
assert events[0]["delta"] == "hello [MASKED]"
|
||||
assert events[1]["text"] == "hello [MASKED]"
|
||||
assert [event["delta"] for event in events[2:]] == ["", ""]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fallback_rewrite_over_an_unplaceable_scanned_event_fails_open(self):
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
||||
|
||||
handler = OpenAIResponsesHandler()
|
||||
events = [
|
||||
{"type": "response.reasoning_summary_text.delta", "output_index": 0, "summary_index": 0, "delta": "hello "},
|
||||
{"type": "response.output_text.delta", "output_index": 1, "content_index": 0, "delta": "world"},
|
||||
]
|
||||
|
||||
with pytest.raises(UndeliverableStreamRewrite):
|
||||
await handler.process_output_streaming_response(
|
||||
responses_so_far=events,
|
||||
|
|
@ -1781,21 +1830,26 @@ class TestOpenAIResponsesHandlerStreamingOutputProcessing:
|
|||
litellm_logging_obj=None,
|
||||
deliver_ended_stream_rewrites=True,
|
||||
)
|
||||
assert [event["delta"] for event in events] == ["hello ", "world"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_output_item_done_last_rewrite_with_delivery_expected_fails_closed(self):
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
||||
|
||||
async def test_output_item_done_last_rewrite_with_delivery_expected_syncs_every_text_event(self):
|
||||
handler = OpenAIResponsesHandler()
|
||||
events = self._ended_stream_events()[:-1]
|
||||
|
||||
with pytest.raises(UndeliverableStreamRewrite):
|
||||
await handler.process_output_streaming_response(
|
||||
responses_so_far=events,
|
||||
guardrail_to_apply=self._masking_guardrail(),
|
||||
litellm_logging_obj=None,
|
||||
deliver_ended_stream_rewrites=True,
|
||||
)
|
||||
result = await handler.process_output_streaming_response(
|
||||
responses_so_far=events,
|
||||
guardrail_to_apply=self._masking_guardrail(),
|
||||
litellm_logging_obj=None,
|
||||
deliver_ended_stream_rewrites=True,
|
||||
)
|
||||
|
||||
assert result is events
|
||||
assert events[0]["delta"] == "hello [MASKED]"
|
||||
assert events[1]["delta"] == ""
|
||||
assert events[2]["text"] == "hello [MASKED]"
|
||||
assert events[3]["part"]["text"] == "hello [MASKED]"
|
||||
assert events[4]["item"]["content"][0]["text"] == "hello [MASKED]"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_output_item_done_last_scans_text_with_delivery_expected(self):
|
||||
|
|
|
|||
|
|
@ -1318,6 +1318,42 @@ async def test_streaming_step_records_guardrail_information_once_on_block(monkey
|
|||
assert _recorded_guardrail_statuses(result) == ["guardrail_intervened"]
|
||||
|
||||
|
||||
def _two_choice_chat_chunks():
|
||||
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
|
||||
|
||||
def chunk(index, content, finish_reason=None):
|
||||
return ModelResponseStream(
|
||||
id="chatcmpl-123",
|
||||
created=1234567890,
|
||||
model="gpt-4",
|
||||
object="chat.completion.chunk",
|
||||
choices=[StreamingChoices(index=index, delta=Delta(content=content), finish_reason=finish_reason)],
|
||||
)
|
||||
|
||||
return [chunk(0, "pers"), chunk(1, "pers"), chunk(0, "immon", "stop"), chunk(1, "immon", "stop")]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_step_delivers_text_rewrites_on_every_choice_of_a_chat_stream(monkeypatch, caplog):
|
||||
from litellm.llms.openai.chat.guardrail_translation.handler import OpenAIChatCompletionsHandler
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", [_TextReturningGuardrail(["[MASKED]", "[MASKED]"])])
|
||||
chunks = _two_choice_chat_chunks()
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
result = await _run_streaming_step(OpenAIChatCompletionsHandler(), chunks)
|
||||
|
||||
assert result.terminal_action == "allow"
|
||||
assert not any("discarded" in record.getMessage() for record in caplog.records)
|
||||
assert [(c.choices[0].index, c.choices[0].delta.content) for c in chunks] == [
|
||||
(0, "[MASKED]"),
|
||||
(1, "[MASKED]"),
|
||||
(0, ""),
|
||||
(1, ""),
|
||||
]
|
||||
assert result.modified_data["metadata"]["applied_guardrails"] == ["masker"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_step_restores_chunks_when_translation_refuses_the_rewrite(monkeypatch, caplog):
|
||||
monkeypatch.setattr(litellm, "callbacks", [_TextReturningGuardrail(["hello [MASKED]"])])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue