From ec865b7d3a2d9b4a44d22d1f7e9ee0f043b6f117 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Micha=C5=82=20Furga=C5=82a?= <83299832+00200200@users.noreply.github.com> Date: Thu, 24 Sep 2026 13:43:37 +0200 Subject: [PATCH] fix(proxy): guard CRLF-framed Anthropic streams and keep text per content block Events were only split on "\n\n", so a CRLF- or CR-framed stream parsed as one block and its text deltas skipped the post-call guardrail. Split on any SSE blank line and keep each block's own separator when rewriting it. All text deltas were also merged into one synthetic block and the rewrite put on the first delta, moving text from later content blocks across any tool or thinking blocks in between. Build one synthetic text block per content-block index and write each rewrite back to its own block. Co-Authored-By: Claude Opus 5.5 --- .../guardrail_translation/handler.py | 88 ++++++++++++------- .../proxy/test_common_request_processing.py | 73 +++++++++++++++ 2 files changed, 129 insertions(+), 32 deletions(-) diff --git a/litellm/llms/anthropic/passthrough/guardrail_translation/handler.py b/litellm/llms/anthropic/passthrough/guardrail_translation/handler.py index 2a93bed9c5c..c3da6cc1a3a 100644 --- a/litellm/llms/anthropic/passthrough/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/passthrough/guardrail_translation/handler.py @@ -3,7 +3,9 @@ from __future__ import annotations import json +import re from collections.abc import Mapping +from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final from litellm._logging import verbose_proxy_logger @@ -17,6 +19,9 @@ if TYPE_CHECKING: _EVENT_STREAM_MEDIA_TYPE: Final = "text/event-stream" _MESSAGES_SUFFIXES: Final = frozenset({"messages", "v1/messages"}) +# SSE allows CRLF, LF or CR line endings, so an event ends at a blank line in any of them. +_SSE_EVENT_END: Final = re.compile(rb"\r\n\r\n|\n\n|\r\r") +_SSE_TRAILING_END: Final = re.compile(rb"(?:\r\n\r\n|\n\n|\r\r)\Z") def _is_messages_endpoint(endpoint: str) -> bool: @@ -25,13 +30,13 @@ def _is_messages_endpoint(endpoint: str) -> bool: def _parse_sse_blocks(body_bytes: bytes) -> tuple[bytes, ...]: - """Split an SSE body into event blocks (including trailing separators).""" + """Split an SSE body into event blocks, each keeping its own trailing separator.""" if not body_bytes: return () - # Keep separators so we can rebuild the stream byte-for-byte aside from rewrites. - parts: Final = body_bytes.split(b"\n\n") - last: Final = len(parts) - 1 - return tuple(part + b"\n\n" if i < last else part for i, part in enumerate(parts) if i < last or part) + ends: Final = tuple(match.end() for match in _SSE_EVENT_END.finditer(body_bytes)) + starts: Final = (0, *ends) + stops: Final = (*ends, len(body_bytes)) + return tuple(body_bytes[start:stop] for start, stop in zip(starts, stops) if stop > start) def _event_payload(block: bytes) -> tuple[str | None, dict[str, Any] | None]: @@ -57,20 +62,21 @@ def _event_payload(block: bytes) -> tuple[str | None, dict[str, Any] | None]: return event_type, payload -def _text_delta(block: bytes) -> str | None: - """The text of a content_block_delta/text_delta event, or None for any other block.""" +def _text_delta(block: bytes) -> tuple[int, str] | None: + """The content block index and text of a text_delta event, or None for any other block.""" event_type, payload = _event_payload(block) if event_type != "content_block_delta" or not payload: return None + index = payload.get("index") delta = payload.get("delta") - if not isinstance(delta, dict) or delta.get("type") != "text_delta": + if not isinstance(index, int) or not isinstance(delta, dict) or delta.get("type") != "text_delta": return None text = delta.get("text") - return text if isinstance(text, str) else None + return (index, text) if isinstance(text, str) else None def _with_text(block: bytes, new_text: str) -> bytes: - """Rebuild a content_block_delta block carrying ``new_text``; other blocks pass through.""" + """Rebuild a content_block_delta block carrying ``new_text``, keeping its framing.""" event_type, payload = _event_payload(block) if event_type != "content_block_delta" or not payload: return block @@ -79,19 +85,25 @@ def _with_text(block: bytes, new_text: str) -> bytes: return block # ``payload`` was just parsed from ``block`` and is not shared, so editing it is local. delta["text"] = new_text - return f"event: content_block_delta\ndata: {json.dumps(payload, separators=(',', ':'))}\n\n".encode() + trailing: Final = _SSE_TRAILING_END.search(block) + separator: Final = trailing.group(0) if trailing else b"" + line_end: Final = separator[: len(separator) // 2].decode() or "\n" + data: Final = json.dumps(payload, separators=(",", ":")) + return f"event: content_block_delta{line_end}data: {data}".encode() + separator -def _first_text(processed: Mapping[str, object]) -> str | None: - """The text of the first content block in a guardrail-processed Messages response.""" +def _processed_texts(processed: Mapping[str, object], count: int) -> tuple[str, ...] | None: + """The text of the first ``count`` content blocks in a guardrail-processed response.""" content: Final = processed.get("content") - if not isinstance(content, list) or not content: + if not isinstance(content, list): return None - first: Final[object] = content[0] - if not isinstance(first, Mapping): - return None - text: Final = first.get("text") - return text if isinstance(text, str) else None + first: Final[list[object]] = content[:count] + texts: Final = tuple( + text + for text in (block.get("text") if isinstance(block, Mapping) else None for block in first) + if isinstance(text, str) + ) + return texts if len(texts) == count else None class AnthropicPassthroughGuardrailHandler(BaseTranslation): @@ -120,21 +132,28 @@ class AnthropicPassthroughGuardrailHandler(BaseTranslation): Placeholders from output_parse_pii are routinely split across multiple text_delta events, so per-frame replacement cannot work; we concatenate - first, then redistribute the de-anonymized text across the original - frames (full rewrite on the first text_delta, empty on the rest). + each content block's text first, then redistribute the de-anonymized + text across that block's frames (full rewrite on its first text_delta, + empty on the rest). """ blocks: Final = _parse_sse_blocks(body_bytes) deltas: Final = tuple( - (idx, text) for idx, text in ((i, _text_delta(block)) for i, block in enumerate(blocks)) if text is not None + (position, found) + for position, found in ((position, _text_delta(block)) for position, block in enumerate(blocks)) + if found is not None ) if not deltas: return body_bytes - combined: Final = "".join(text for _, text in deltas) + # One synthetic text block per Anthropic content block, so text from separate + # blocks is never merged or moved across the tool/thinking blocks between them. + indices: Final = tuple(sorted(frozenset(index for _, (index, _) in deltas))) synthetic_response: Final[dict] = { "type": "message", "role": "assistant", - "content": [{"type": "text", "text": combined}], + "content": [ + {"type": "text", "text": "".join(text for _, (i, text) in deltas if i == index)} for index in indices + ], "stop_reason": "end_turn", } @@ -151,17 +170,22 @@ class AnthropicPassthroughGuardrailHandler(BaseTranslation): ) return body_bytes - de_anonymized: Final = _first_text(processed) - if de_anonymized is None: + rewritten: Final = _processed_texts(processed, len(indices)) + if rewritten is None: return body_bytes - # Put the full rewrite on the first text_delta; blank the rest so - # split placeholders cannot survive across frames. - delta_indices: Final = frozenset(idx for idx, _ in deltas) - first_idx: Final = deltas[0][0] + # Put each block's full rewrite on its first text_delta and blank the rest, so + # placeholders split across frames cannot survive. + text_for_index: Final = MappingProxyType({index: text for index, text in zip(indices, rewritten)}) + index_at: Final = MappingProxyType({position: index for position, (index, _) in deltas}) + first_positions: Final = frozenset( + min(position for position, (i, _) in deltas if i == index) for index in indices + ) return b"".join( - _with_text(block, de_anonymized if idx == first_idx else "") if idx in delta_indices else block - for idx, block in enumerate(blocks) + (_with_text(block, text_for_index[index_at[position]] if position in first_positions else "")) + if position in index_at + else block + for position, block in enumerate(blocks) ) async def process_input_messages( diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 0bea655aa85..bfbbc52fcbe 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -5607,6 +5607,79 @@ class TestEventStreamAllmPassthroughRoute: assert b"Alice" in result assert b"" not in result + @pytest.mark.asyncio + async def test_anthropic_crlf_framed_stream_is_still_guarded(self): + sse = ( + b"event: content_block_delta\r\n" + b'data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":""}}\r\n\r\n' + b"event: message_stop\r\n" + b'data: {"type":"message_stop"}\r\n\r\n' + ) + + async def mock_hook(data, user_api_key_dict, response): + response = dict(response) + response["content"] = [{"type": "text", "text": "Alice"}] + return response + + proxy_logging_obj = MagicMock() + proxy_logging_obj.post_call_success_hook = mock_hook + + processing_obj = ProxyBaseLLMRequestProcessing( + data={"custom_llm_provider": "anthropic", "endpoint": "/v1/messages"} + ) + result = await processing_obj._handle_event_stream_allm_passthrough_route( + body_bytes=sse, + proxy_logging_obj=proxy_logging_obj, + user_api_key_dict=MagicMock(spec=UserAPIKeyAuth), + ) + + assert b"" not in result + assert b"Alice" in result + assert result.endswith(b'data: {"type":"message_stop"}\r\n\r\n') + assert b"\r\n\r\n" in result.split(b"event: message_stop")[0] + + @pytest.mark.asyncio + async def test_anthropic_text_blocks_keep_their_own_rewrites(self): + def text_delta(index, text): + payload = {"type": "content_block_delta", "index": index, "delta": {"type": "text_delta", "text": text}} + return b"event: content_block_delta\ndata: " + json.dumps(payload, separators=(",", ":")).encode() + b"\n\n" + + tool_delta = ( + b"event: content_block_delta\n" + b'data: {"type":"content_block_delta","index":1,"delta":{"type":"input_json_delta","partial_json":"{}"}}\n\n' + ) + sse = text_delta(0, " said") + tool_delta + text_delta(2, "bye ") + + seen = {} + + async def mock_hook(data, user_api_key_dict, response): + seen["texts"] = [block["text"] for block in response["content"]] + response = dict(response) + response["content"] = [{"type": "text", "text": "Alice said"}, {"type": "text", "text": "bye Bob"}] + return response + + proxy_logging_obj = MagicMock() + proxy_logging_obj.post_call_success_hook = mock_hook + + processing_obj = ProxyBaseLLMRequestProcessing( + data={"custom_llm_provider": "anthropic", "endpoint": "/v1/messages"} + ) + result = await processing_obj._handle_event_stream_allm_passthrough_route( + body_bytes=sse, + proxy_logging_obj=proxy_logging_obj, + user_api_key_dict=MagicMock(spec=UserAPIKeyAuth), + ) + + assert seen["texts"] == [" said", "bye "] + frames = [frame for frame in result.split(b"\n\n") if frame] + payloads = [json.loads(frame.split(b"data: ", 1)[1]) for frame in frames] + assert [(p["index"], p["delta"].get("text")) for p in payloads] == [ + (0, "Alice said"), + (0, ""), + (1, None), + (2, "bye Bob"), + ] + @pytest.mark.asyncio async def test_anthropic_supports_event_stream_de_anonymization_for_messages(self): from litellm.llms.pass_through.guardrail_translation.handler import (