mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Fail closed on rewrites for buffers that never reached their terminal event
An Anthropic buffer without a stop_reason only ran the flat text scan, so a rewrite there was dropped while the executor trusted the translation to have delivered it. A Responses buffer ending at response.output_item.done returned after the tool-call scan without ever checking the text. Both now reach the flat scan and raise UndeliverableStreamRewrite when a caller expects the rewrite delivered, matching the existing Responses no-envelope fallback.
This commit is contained in:
parent
2d7697569a
commit
c09db7c7a3
4 changed files with 104 additions and 3 deletions
|
|
@ -1020,7 +1020,8 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
|
||||
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).
|
||||
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 fails closed instead.
|
||||
"""
|
||||
from litellm.integrations.custom_guardrail import ModifyResponseException
|
||||
|
||||
|
|
@ -1098,6 +1099,11 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
if e.original_response is None:
|
||||
e.original_response = self._build_streaming_usage_response(responses_so_far, request_data)
|
||||
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")
|
||||
return responses_so_far
|
||||
|
||||
def _prepare_request_data(
|
||||
|
|
|
|||
|
|
@ -638,7 +638,9 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
return responses_so_far
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# Case 2: response.output_item.done — extract tool calls only. #
|
||||
# Case 2: response.output_item.done — extract tool calls only, then #
|
||||
# fall through to the text fallback when a caller expects rewrites #
|
||||
# delivered, so a buffer truncated here still fails closed on text. #
|
||||
# ------------------------------------------------------------------ #
|
||||
if final_chunk.get("type") == "response.output_item.done":
|
||||
model_response_stream: Final = (
|
||||
|
|
@ -656,7 +658,8 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
input_type="response",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
return responses_so_far
|
||||
if not deliver_ended_stream_rewrites:
|
||||
return responses_so_far
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# Fallback: apply guardrail to the accumulated text string. #
|
||||
|
|
|
|||
|
|
@ -328,6 +328,52 @@ 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
|
||||
|
||||
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,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unended_stream_without_rewrite_is_released_with_delivery_expected(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
chunks = self._ended_sse_chunks()[:-2]
|
||||
original = [bytes(chunk) for chunk in chunks]
|
||||
|
||||
result = await handler.process_output_streaming_response(
|
||||
responses_so_far=chunks,
|
||||
guardrail_to_apply=MockPassThroughGuardrail(guardrail_name="test"),
|
||||
litellm_logging_obj=MagicMock(),
|
||||
deliver_ended_stream_rewrites=True,
|
||||
)
|
||||
|
||||
assert result is chunks
|
||||
assert chunks == original
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unended_stream_rewrite_without_delivery_expected_does_not_raise(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
chunks = self._ended_sse_chunks()[:-2]
|
||||
original = [bytes(chunk) for chunk in chunks]
|
||||
|
||||
result = await handler.process_output_streaming_response(
|
||||
responses_so_far=chunks,
|
||||
guardrail_to_apply=self._masking_guardrail(),
|
||||
litellm_logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
assert result is chunks
|
||||
assert chunks == original
|
||||
|
||||
|
||||
class TestAnthropicMessagesHandlerInputProcessing:
|
||||
"""Test input processing preserves litellm_metadata for dynamic guardrails."""
|
||||
|
|
|
|||
|
|
@ -1248,6 +1248,52 @@ class TestOpenAIResponsesHandlerStreamingOutputProcessing:
|
|||
deliver_ended_stream_rewrites=True,
|
||||
)
|
||||
|
||||
@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
|
||||
|
||||
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,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_output_item_done_last_scans_text_with_delivery_expected(self):
|
||||
handler = OpenAIResponsesHandler()
|
||||
events = self._ended_stream_events()[:-1]
|
||||
guardrail = MockRecordingGuardrail(guardrail_name="test")
|
||||
|
||||
result = await handler.process_output_streaming_response(
|
||||
responses_so_far=events,
|
||||
guardrail_to_apply=guardrail,
|
||||
litellm_logging_obj=None,
|
||||
deliver_ended_stream_rewrites=True,
|
||||
)
|
||||
|
||||
assert result is events
|
||||
assert [inputs.get("texts") for inputs in guardrail.seen_inputs] == [["hello world"]]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_output_item_done_last_without_delivery_expected_skips_text(self):
|
||||
handler = OpenAIResponsesHandler()
|
||||
events = self._ended_stream_events()[:-1]
|
||||
guardrail = MockRecordingGuardrail(guardrail_name="test")
|
||||
|
||||
result = await handler.process_output_streaming_response(
|
||||
responses_so_far=events,
|
||||
guardrail_to_apply=guardrail,
|
||||
litellm_logging_obj=None,
|
||||
)
|
||||
|
||||
assert result is events
|
||||
assert guardrail.seen_inputs == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fallback_rewrite_without_delivery_expected_does_not_raise(self):
|
||||
handler = OpenAIResponsesHandler()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue