From 802a57e5f70da4f6d8916822c89f03db61ee71d7 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Micha=C5=82=20Furga=C5=82a?= <83299832+00200200@users.noreply.github.com> Date: Tue, 22 Sep 2026 23:06:28 +0200 Subject: [PATCH 1/7] Apply post-call stream guardrails on Anthropic messages passthrough Register an Anthropic SSE passthrough handler so streaming /v1/messages buffers through post-call guardrails (e.g. output_parse_pii). Fixes #42476 --- .../llms/anthropic/passthrough/__init__.py | 0 .../guardrail_translation/__init__.py | 0 .../guardrail_translation/handler.py | 198 ++++++++++++++++++ .../guardrail_translation/handler.py | 8 +- .../proxy/test_common_request_processing.py | 49 ++++- 5 files changed, 252 insertions(+), 3 deletions(-) create mode 100644 litellm/llms/anthropic/passthrough/__init__.py create mode 100644 litellm/llms/anthropic/passthrough/guardrail_translation/__init__.py create mode 100644 litellm/llms/anthropic/passthrough/guardrail_translation/handler.py diff --git a/litellm/llms/anthropic/passthrough/__init__.py b/litellm/llms/anthropic/passthrough/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/anthropic/passthrough/guardrail_translation/__init__.py b/litellm/llms/anthropic/passthrough/guardrail_translation/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/anthropic/passthrough/guardrail_translation/handler.py b/litellm/llms/anthropic/passthrough/guardrail_translation/handler.py new file mode 100644 index 00000000000..65c2629bb7c --- /dev/null +++ b/litellm/llms/anthropic/passthrough/guardrail_translation/handler.py @@ -0,0 +1,198 @@ +"""Anthropic /v1/messages passthrough guardrail translation (SSE event stream).""" + +from __future__ import annotations + +import json +from collections.abc import Mapping +from typing import TYPE_CHECKING, Any, Final, Optional + +from litellm._logging import verbose_proxy_logger +from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation + +if TYPE_CHECKING: + from litellm.integrations.custom_guardrail import CustomGuardrail + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.utils import ProxyLogging + +_EVENT_STREAM_MEDIA_TYPE: Final = "text/event-stream" +_MESSAGES_SUFFIXES: Final = frozenset({"messages", "v1/messages"}) + + +def _is_messages_endpoint(endpoint: str) -> bool: + normalized = endpoint.rstrip("/").split("?")[0] + return any(normalized.endswith(suffix) for suffix in _MESSAGES_SUFFIXES) + + +def _parse_sse_blocks(body_bytes: bytes) -> list[bytes]: + """Split an SSE body into event blocks (including trailing separators).""" + if not body_bytes: + return [] + # Keep separators so we can rebuild the stream byte-for-byte aside from rewrites. + parts = body_bytes.split(b"\n\n") + blocks: list[bytes] = [] + for i, part in enumerate(parts): + if i < len(parts) - 1: + blocks.append(part + b"\n\n") + elif part: + blocks.append(part) + return blocks + + +def _event_payload(block: bytes) -> tuple[str | None, dict[str, Any] | None]: + try: + text = block.decode("utf-8") + except UnicodeDecodeError: + return None, None + event_type: str | None = None + data_line: str | None = None + for line in text.splitlines(): + if line.startswith("event:"): + event_type = line[6:].strip() + elif line.startswith("data:"): + data_line = line[5:].strip() + if not data_line: + return event_type, None + try: + payload = json.loads(data_line) + except json.JSONDecodeError: + return event_type, None + if not isinstance(payload, dict): + return event_type, None + return event_type, payload + + +class AnthropicPassthroughGuardrailHandler(BaseTranslation): + @staticmethod + def is_event_stream_content_type(content_type: str) -> bool: + return "text/event-stream" in content_type + + @staticmethod + def event_stream_media_type() -> str: + return _EVENT_STREAM_MEDIA_TYPE + + @staticmethod + def event_stream_endpoint_is_de_anonymizable(endpoint: str) -> bool: + return _is_messages_endpoint(endpoint) + + @staticmethod + async def de_anonymize_event_stream( + body_bytes: bytes, + proxy_logging_obj: "ProxyLogging", + user_api_key_dict: "UserAPIKeyAuth", + data: dict, + ) -> bytes: + """ + Buffer Anthropic SSE frames, run post-call guardrails on concatenated + text_delta content, and rewrite text_delta payloads in place. + + 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). + """ + blocks = _parse_sse_blocks(body_bytes) + text_block_indices: list[int] = [] + texts: list[str] = [] + + for idx, block in enumerate(blocks): + event_type, payload = _event_payload(block) + if event_type != "content_block_delta" or not payload: + continue + delta = payload.get("delta") + if not isinstance(delta, dict) or delta.get("type") != "text_delta": + continue + text = delta.get("text") + if not isinstance(text, str): + continue + text_block_indices.append(idx) + texts.append(text) + + if not texts: + return body_bytes + + combined = "".join(texts) + synthetic_response: Final[dict] = { + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": combined}], + "stop_reason": "end_turn", + } + + processed = await proxy_logging_obj.post_call_success_hook( + data=data, + user_api_key_dict=user_api_key_dict, + response=synthetic_response, + ) + if not isinstance(processed, dict): + verbose_proxy_logger.debug( + "AnthropicPassthroughGuardrailHandler: post_call_success_hook returned %s, " + "leaving event stream unmodified", + type(processed).__name__, + ) + return body_bytes + + try: + content = processed["content"] + de_anonymized = content[0]["text"] + if not isinstance(de_anonymized, str): + return body_bytes + except (KeyError, IndexError, TypeError): + return body_bytes + + # Put the full rewrite on the first text_delta; blank the rest so + # split placeholders cannot survive across frames. + replacements = [de_anonymized] + [""] * (len(text_block_indices) - 1) + out: list[bytes] = list(blocks) + for block_idx, new_text in zip(text_block_indices, replacements): + event_type, payload = _event_payload(out[block_idx]) + if event_type != "content_block_delta" or not payload: + continue + delta = payload.get("delta") + if not isinstance(delta, dict): + continue + delta["text"] = new_text + payload["delta"] = delta + # Rebuild a minimal SSE block; preserve event name. + new_block = ( + f"event: content_block_delta\ndata: {json.dumps(payload, separators=(',', ':'))}\n\n" + ).encode("utf-8") + out[block_idx] = new_block + + return b"".join(out) + + async def process_input_messages( + self, + data: dict, + guardrail_to_apply: "CustomGuardrail", + litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None, + ) -> Mapping[str, object]: + from litellm.llms.pass_through.guardrail_translation.handler import ( + PassThroughEndpointHandler, + ) + + return await PassThroughEndpointHandler().process_input_messages( + data=data, + guardrail_to_apply=guardrail_to_apply, + litellm_logging_obj=litellm_logging_obj, + ) + + async def process_output_response( + self, + response: object, + guardrail_to_apply: "CustomGuardrail", + litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None, + user_api_key_dict: Optional["UserAPIKeyAuth"] = None, + request_data: dict | None = None, + ) -> object: + from litellm.llms.pass_through.guardrail_translation.handler import ( + PassThroughEndpointHandler, + ) + + return await PassThroughEndpointHandler().process_output_response( + response=response, + guardrail_to_apply=guardrail_to_apply, + litellm_logging_obj=litellm_logging_obj, + user_api_key_dict=user_api_key_dict, + request_data=request_data, + ) diff --git a/litellm/llms/pass_through/guardrail_translation/handler.py b/litellm/llms/pass_through/guardrail_translation/handler.py index 1f295a6e656..a9a795c073c 100644 --- a/litellm/llms/pass_through/guardrail_translation/handler.py +++ b/litellm/llms/pass_through/guardrail_translation/handler.py @@ -199,11 +199,17 @@ _PROVIDER_HANDLERS: dict[str, type[BaseTranslation]] = {} def _get_provider_handlers() -> dict[str, type[BaseTranslation]]: global _PROVIDER_HANDLERS if not _PROVIDER_HANDLERS: + from litellm.llms.anthropic.passthrough.guardrail_translation.handler import ( + AnthropicPassthroughGuardrailHandler, + ) from litellm.llms.bedrock.passthrough.guardrail_translation.handler import ( BedrockPassthroughGuardrailHandler, ) - _PROVIDER_HANDLERS = {"bedrock": BedrockPassthroughGuardrailHandler} + _PROVIDER_HANDLERS = { + "anthropic": AnthropicPassthroughGuardrailHandler, + "bedrock": BedrockPassthroughGuardrailHandler, + } return _PROVIDER_HANDLERS diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 5b9cd761dda..8c5c4cc7661 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -5664,11 +5664,11 @@ class TestEventStreamAllmPassthroughRoute: assert result == expected_bytes @pytest.mark.asyncio - async def test_non_bedrock_provider_returns_original_bytes(self): + async def test_unknown_provider_returns_original_bytes(self): stream_bytes = _build_event_stream_frame("messageStart", {"role": "assistant"}) proxy_logging_obj = MagicMock() - processing_obj = ProxyBaseLLMRequestProcessing(data={"custom_llm_provider": "anthropic"}) + processing_obj = ProxyBaseLLMRequestProcessing(data={"custom_llm_provider": "unknown-provider"}) result = await processing_obj._handle_event_stream_allm_passthrough_route( body_bytes=stream_bytes, proxy_logging_obj=proxy_logging_obj, @@ -5677,6 +5677,51 @@ class TestEventStreamAllmPassthroughRoute: assert result is stream_bytes + @pytest.mark.asyncio + async def test_anthropic_provider_dispatches_to_handler(self): + sse = ( + b'event: content_block_delta\n' + b'data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":""}}\n\n' + b'event: message_stop\n' + b'data: {"type":"message_stop"}\n\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"Alice" in result + assert b"" not in result + + @pytest.mark.asyncio + async def test_anthropic_supports_event_stream_de_anonymization_for_messages(self): + from litellm.llms.pass_through.guardrail_translation.handler import ( + LlmPassthroughRouteHandler, + ) + + assert LlmPassthroughRouteHandler.supports_event_stream_de_anonymization( + "anthropic", "/v1/messages" + ) + assert LlmPassthroughRouteHandler.supports_event_stream_de_anonymization( + "anthropic", "messages" + ) + assert not LlmPassthroughRouteHandler.supports_event_stream_de_anonymization( + "anthropic", "/v1/complete" + ) + @pytest.mark.asyncio async def test_non_streaming_response_includes_custom_headers(self): import json From 82801209cdbbae903fbfbba0c69b42c636445e42 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Micha=C5=82=20Furga=C5=82a?= <83299832+00200200@users.noreply.github.com> Date: Wed, 23 Sep 2026 08:46:49 +0200 Subject: [PATCH 2/7] fix(proxy): format the Anthropic passthrough handler and update the SSE test Registering an Anthropic event-stream handler means its unbuffered /v1/messages stream now carries text/event-stream, so the test that used Anthropic as a provider without a registered media type pointed at the wrong provider. It now uses one with no handler, and a new test pins the Anthropic media type. Also applies ruff format and the pyupgrade autofixes to the new handler module. Co-Authored-By: Claude Opus 5.5 --- .../guardrail_translation/handler.py | 20 +- .../proxy/test_common_request_processing.py | 545 +++++++----------- 2 files changed, 231 insertions(+), 334 deletions(-) diff --git a/litellm/llms/anthropic/passthrough/guardrail_translation/handler.py b/litellm/llms/anthropic/passthrough/guardrail_translation/handler.py index 65c2629bb7c..cb5532c93e5 100644 --- a/litellm/llms/anthropic/passthrough/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/passthrough/guardrail_translation/handler.py @@ -4,7 +4,7 @@ from __future__ import annotations import json from collections.abc import Mapping -from typing import TYPE_CHECKING, Any, Final, Optional +from typing import TYPE_CHECKING, Any, Final from litellm._logging import verbose_proxy_logger from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation @@ -78,8 +78,8 @@ class AnthropicPassthroughGuardrailHandler(BaseTranslation): @staticmethod async def de_anonymize_event_stream( body_bytes: bytes, - proxy_logging_obj: "ProxyLogging", - user_api_key_dict: "UserAPIKeyAuth", + proxy_logging_obj: ProxyLogging, + user_api_key_dict: UserAPIKeyAuth, data: dict, ) -> bytes: """ @@ -154,9 +154,7 @@ class AnthropicPassthroughGuardrailHandler(BaseTranslation): delta["text"] = new_text payload["delta"] = delta # Rebuild a minimal SSE block; preserve event name. - new_block = ( - f"event: content_block_delta\ndata: {json.dumps(payload, separators=(',', ':'))}\n\n" - ).encode("utf-8") + new_block = (f"event: content_block_delta\ndata: {json.dumps(payload, separators=(',', ':'))}\n\n").encode() out[block_idx] = new_block return b"".join(out) @@ -164,8 +162,8 @@ class AnthropicPassthroughGuardrailHandler(BaseTranslation): async def process_input_messages( self, data: dict, - guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None, + guardrail_to_apply: CustomGuardrail, + litellm_logging_obj: LiteLLMLoggingObj | None = None, ) -> Mapping[str, object]: from litellm.llms.pass_through.guardrail_translation.handler import ( PassThroughEndpointHandler, @@ -180,9 +178,9 @@ class AnthropicPassthroughGuardrailHandler(BaseTranslation): async def process_output_response( self, response: object, - guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None, - user_api_key_dict: Optional["UserAPIKeyAuth"] = None, + guardrail_to_apply: CustomGuardrail, + litellm_logging_obj: LiteLLMLoggingObj | None = None, + user_api_key_dict: UserAPIKeyAuth | None = None, request_data: dict | None = None, ) -> object: from litellm.llms.pass_through.guardrail_translation.handler import ( diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 8c5c4cc7661..0bea655aa85 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -109,9 +109,7 @@ def test_attach_guardrail_information_redacts_matched_content(): { "guardrail_name": "cf", "guardrail_status": "success", - "guardrail_response": [ - {"type": "blocked_word", "keyword": "secret-word", "action": "MASK"} - ], + "guardrail_response": [{"type": "blocked_word", "keyword": "secret-word", "action": "MASK"}], "match_details": [{"snippet": "secret-word", "detection_method": "keyword"}], } ] @@ -183,9 +181,12 @@ def test_include_guardrail_response_requested_reads_flag_from_metadata_when_rout def test_include_guardrail_response_requested_is_false_without_exact_true(): - assert include_guardrail_response_requested( - {"metadata": {"include_guardrail_response": "true"}, "litellm_metadata": {}} - ) is False + assert ( + include_guardrail_response_requested( + {"metadata": {"include_guardrail_response": "true"}, "litellm_metadata": {}} + ) + is False + ) assert include_guardrail_response_requested({}) is False @@ -271,16 +272,12 @@ class TestProxyBaseLLMRequestProcessing: assert json.loads(result.body) == guardrailed_body @pytest.mark.asyncio - async def test_handle_non_streaming_allm_passthrough_route_forwards_upstream_headers( - self, monkeypatch - ): + async def test_handle_non_streaming_allm_passthrough_route_forwards_upstream_headers(self, monkeypatch): """The guardrail JSON path must forward upstream response headers (e.g. x-amzn-requestid) alongside the x-litellm-* headers, matching the non-guardrail passthrough path, while dropping length headers that no longer match the rewritten body.""" - processing_obj = ProxyBaseLLMRequestProcessing( - data={"custom_llm_provider": "bedrock"} - ) + processing_obj = ProxyBaseLLMRequestProcessing(data={"custom_llm_provider": "bedrock"}) monkeypatch.setattr( processing_obj, "_has_post_call_guardrails_for_passthrough", @@ -320,14 +317,10 @@ class TestProxyBaseLLMRequestProcessing: assert result.headers["content-length"] == str(len(result.body)) @pytest.mark.asyncio - async def test_handle_event_stream_allm_passthrough_route_forwards_upstream_headers( - self, monkeypatch - ): + async def test_handle_event_stream_allm_passthrough_route_forwards_upstream_headers(self, monkeypatch): """The guardrail event-stream branch must also forward upstream response headers alongside the x-litellm-* headers.""" - processing_obj = ProxyBaseLLMRequestProcessing( - data={"custom_llm_provider": "bedrock"} - ) + processing_obj = ProxyBaseLLMRequestProcessing(data={"custom_llm_provider": "bedrock"}) monkeypatch.setattr( processing_obj, "_has_post_call_guardrails_for_passthrough", @@ -369,15 +362,11 @@ class TestProxyBaseLLMRequestProcessing: assert result.headers["x-litellm-call-id"] == "test-call-id" @pytest.mark.asyncio - async def test_handle_non_streaming_allm_passthrough_route_applies_response_headers_hook( - self, monkeypatch - ): + async def test_handle_non_streaming_allm_passthrough_route_applies_response_headers_hook(self, monkeypatch): """Guardrailed non-streaming passthrough responses must include headers injected by post_call_response_headers_hook, matching the headers a non-guardrailed passthrough response would carry.""" - processing_obj = ProxyBaseLLMRequestProcessing( - data={"custom_llm_provider": "bedrock"} - ) + processing_obj = ProxyBaseLLMRequestProcessing(data={"custom_llm_provider": "bedrock"}) monkeypatch.setattr( processing_obj, "_has_post_call_guardrails_for_passthrough", @@ -396,9 +385,7 @@ class TestProxyBaseLLMRequestProcessing: return kwargs["response"] proxy_logging_obj.post_call_success_hook = fake_post_call_success_hook - proxy_logging_obj.post_call_response_headers_hook = AsyncMock( - return_value={"x-litellm-custom": "from-hook"} - ) + proxy_logging_obj.post_call_response_headers_hook = AsyncMock(return_value={"x-litellm-custom": "from-hook"}) result = await processing_obj._handle_non_streaming_allm_passthrough_route( response=upstream, @@ -601,9 +588,7 @@ class TestProxyBaseLLMRequestProcessing: ) @pytest.mark.asyncio - async def test_common_processing_pre_call_logic_enforces_tag_budget_for_guardrail_added_tags( - self, monkeypatch - ): + async def test_common_processing_pre_call_logic_enforces_tag_budget_for_guardrail_added_tags(self, monkeypatch): processing_obj, mock_request, mock_proxy_logging_obj, mock_proxy_config, tag_budget_check = ( self._guardrail_tag_budget_harness( monkeypatch, @@ -716,9 +701,7 @@ class TestProxyBaseLLMRequestProcessing: tag_budget_check.assert_not_awaited() @pytest.mark.asyncio - async def test_common_processing_pre_call_logic_rechecks_guardrail_added_tag_on_fallback_retry( - self, monkeypatch - ): + async def test_common_processing_pre_call_logic_rechecks_guardrail_added_tag_on_fallback_retry(self, monkeypatch): processing_obj, mock_request, mock_proxy_logging_obj, mock_proxy_config, tag_budget_check = ( self._guardrail_tag_budget_harness( monkeypatch, @@ -861,9 +844,7 @@ class TestProxyBaseLLMRequestProcessing: tag_budget_check.assert_awaited_once() @pytest.mark.asyncio - async def test_common_processing_pre_call_logic_arms_auto_router_compression_before_guardrails( - self, monkeypatch - ): + async def test_common_processing_pre_call_logic_arms_auto_router_compression_before_guardrails(self, monkeypatch): """arm_pre_call must run before pre_call_hook: an auto router's own compression policy has to be in `data["metadata"]` (naming the model-side guardrail so it runs even if it isn't default_on) by the time guardrails see the request.""" @@ -2707,16 +2688,10 @@ class TestCommonRequestProcessingHelpers: def _stringified_none_paths(node: object, path: str = "error") -> tuple[str, ...]: if isinstance(node, dict): - return tuple( - found - for key, value in node.items() - for found in _stringified_none_paths(value, f"{path}.{key}") - ) + return tuple(found for key, value in node.items() for found in _stringified_none_paths(value, f"{path}.{key}")) if isinstance(node, (list, tuple)): return tuple( - found - for index, value in enumerate(node) - for found in _stringified_none_paths(value, f"{path}[{index}]") + found for index, value in enumerate(node) for found in _stringified_none_paths(value, f"{path}[{index}]") ) return (path,) if node == "None" else () @@ -3463,9 +3438,7 @@ class TestStreamingOverheadHeader: user_api_key_dict=mock_user_api_key_dict, call_id="test-call-id", hidden_params={}, - litellm_logging_obj=self._timing_logging_obj( - {"_response_ms": 500.0, "litellm_overhead_time_ms": 42.5} - ), + litellm_logging_obj=self._timing_logging_obj({"_response_ms": 500.0, "litellm_overhead_time_ms": 42.5}), ) assert headers["x-litellm-response-duration-ms"] == "500.0" @@ -3484,9 +3457,7 @@ class TestStreamingOverheadHeader: user_api_key_dict=mock_user_api_key_dict, call_id="test-call-id", hidden_params={}, - litellm_logging_obj=self._timing_logging_obj( - {"_response_ms": 500.0, "litellm_overhead_time_ms": 42.5} - ), + litellm_logging_obj=self._timing_logging_obj({"_response_ms": 500.0, "litellm_overhead_time_ms": 42.5}), read_timing_from_logging_obj=False, ) @@ -3507,9 +3478,7 @@ class TestStreamingOverheadHeader: user_api_key_dict=mock_user_api_key_dict, call_id="test-call-id", hidden_params={"_response_ms": 300.0}, - litellm_logging_obj=self._timing_logging_obj( - {"_response_ms": 500.0, "litellm_overhead_time_ms": 42.5} - ), + litellm_logging_obj=self._timing_logging_obj({"_response_ms": 500.0, "litellm_overhead_time_ms": 42.5}), ) assert headers["x-litellm-response-duration-ms"] == "300.0" @@ -3550,9 +3519,7 @@ class TestStreamingOverheadHeader: user_api_key_dict=mock_user_api_key_dict, call_id="test-call-id", hidden_params={"_response_ms": 300.0, "litellm_overhead_time_ms": 7.5}, - litellm_logging_obj=self._timing_logging_obj( - {"_response_ms": 500.0, "litellm_overhead_time_ms": 42.5} - ), + litellm_logging_obj=self._timing_logging_obj({"_response_ms": 500.0, "litellm_overhead_time_ms": 42.5}), ) assert headers["x-litellm-response-duration-ms"] == "300.0" @@ -4004,9 +3971,7 @@ class TestStreamCloseOnDisconnect: finally: closed.set() - response = _UpstreamClosingStreamingResponse( - body(), media_type="text/event-stream" - ) + response = _UpstreamClosingStreamingResponse(body(), media_type="text/event-stream") async def receive(): await asyncio.Event().wait() @@ -4037,9 +4002,7 @@ class TestStreamCloseOnDisconnect: finally: closed.set() - response = _UpstreamClosingStreamingResponse( - body(), media_type="text/event-stream" - ) + response = _UpstreamClosingStreamingResponse(body(), media_type="text/event-stream") async def receive(): await disconnected.wait() @@ -4110,9 +4073,7 @@ class TestStreamCloseOnDisconnect: finally: inner_closed.set() - response = await create_response( - generator=wrapped(), media_type="text/event-stream", headers={} - ) + response = await create_response(generator=wrapped(), media_type="text/event-stream", headers={}) async def receive(): await asyncio.Event().wait() @@ -4344,9 +4305,7 @@ class TestStreamCloseOnDisconnect: with pytest.raises(_ClientDisconnectedBeforeFirstChunk): await asyncio.wait_for( - _buffer_first_chunk_honoring_disconnect( - AcloseRaises(), request=self._request_that_disconnects() - ), + _buffer_first_chunk_honoring_disconnect(AcloseRaises(), request=self._request_that_disconnects()), timeout=5, ) @@ -4362,9 +4321,7 @@ class TestStreamCloseOnDisconnect: with pytest.raises(_ClientDisconnectedBeforeFirstChunk): await asyncio.wait_for( - _buffer_first_chunk_honoring_disconnect( - blocking_gen(), request=self._request_that_disconnects() - ), + _buffer_first_chunk_honoring_disconnect(blocking_gen(), request=self._request_that_disconnects()), timeout=5, ) assert closed.is_set() @@ -4380,9 +4337,7 @@ class TestHandleLLMApiExceptionRetryAfter: user_api_key_dict = UserAPIKeyAuth(api_key="sk-test") proxy_logging_obj = MagicMock() proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) - proxy_logging_obj.post_call_response_headers_hook = AsyncMock( - return_value=callback_headers or {} - ) + proxy_logging_obj.post_call_response_headers_hook = AsyncMock(return_value=callback_headers or {}) try: await processor._handle_llm_api_exception( @@ -4449,9 +4404,7 @@ class TestHandleLLMApiExceptionRetryAfter: enable_pre_call_checks=False, cooldown_list=[], ) - proxy_exc = await self._invoke( - exc, callback_headers={"retry-after": "", "x-custom": "1"} - ) + proxy_exc = await self._invoke(exc, callback_headers={"retry-after": "", "x-custom": "1"}) assert proxy_exc.headers["retry-after"] == "43" assert proxy_exc.headers["x-custom"] == "1" @@ -4680,9 +4633,7 @@ class TestDisconnectGatherCleanup: return Request(scope={"type": "http", "headers": []}, receive=receive) @pytest.mark.asyncio - async def test_base_process_llm_request_raises_499_on_client_disconnect( - self, monkeypatch - ): + async def test_base_process_llm_request_raises_499_on_client_disconnect(self, monkeypatch): """With cancel_on_disconnect enabled, base_process_llm_request returns 499.""" import asyncio @@ -4711,9 +4662,7 @@ class TestDisconnectGatherCleanup: "common_processing_pre_call_logic", AsyncMock(return_value=({"model": "gemini-2.0-flash"}, mock_logging_obj)), ) - monkeypatch.setattr( - processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False) - ) + monkeypatch.setattr(processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False)) with pytest.raises(HTTPException) as exc_info: await processing_obj.base_process_llm_request( @@ -4731,9 +4680,7 @@ class TestDisconnectGatherCleanup: assert "disconnected" in exc_info.value.detail.lower() @pytest.mark.asyncio - async def test_base_process_llm_request_reraises_cancelled_error_without_client_disconnect( - self, monkeypatch - ): + async def test_base_process_llm_request_reraises_cancelled_error_without_client_disconnect(self, monkeypatch): import asyncio import litellm.proxy.common_request_processing as cpr @@ -4758,9 +4705,7 @@ class TestDisconnectGatherCleanup: "common_processing_pre_call_logic", AsyncMock(return_value=({"model": "gemini-2.0-flash"}, mock_logging_obj)), ) - monkeypatch.setattr( - processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False) - ) + monkeypatch.setattr(processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False)) monkeypatch.setattr( cpr, "route_request", @@ -4821,9 +4766,7 @@ class TestDisconnectGatherCleanup: "common_processing_pre_call_logic", AsyncMock(return_value=({"model": "gemini-2.0-flash"}, mock_logging_obj)), ) - monkeypatch.setattr( - processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False) - ) + monkeypatch.setattr(processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False)) with pytest.raises(HTTPException): await processing_obj.base_process_llm_request( @@ -4874,9 +4817,7 @@ class TestDisconnectGatherCleanup: assert task.done() @pytest.mark.asyncio - async def test_base_process_llm_request_preserves_llm_error_after_gather( - self, monkeypatch - ): + async def test_base_process_llm_request_preserves_llm_error_after_gather(self, monkeypatch): import litellm.proxy.common_request_processing as cpr from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing @@ -4905,9 +4846,7 @@ class TestDisconnectGatherCleanup: "common_processing_pre_call_logic", AsyncMock(return_value=({"model": "gemini-2.0-flash"}, mock_logging_obj)), ) - monkeypatch.setattr( - processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False) - ) + monkeypatch.setattr(processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False)) mock_request = MagicMock(spec=Request) mock_request.is_disconnected = AsyncMock(return_value=False) @@ -4991,19 +4930,13 @@ class TestStreamingClientDisconnectLogging: "litellm_params": {"metadata": {}}, } - recorded = await _record_streaming_client_disconnect_if_needed( - mock_request, request_data - ) + recorded = await _record_streaming_client_disconnect_if_needed(mock_request, request_data) assert recorded is True assert request_data["metadata"]["client_disconnected"] is True + assert request_data["metadata"]["error_information"]["error_code"] == "499" assert ( - request_data["metadata"]["error_information"]["error_code"] == "499" - ) - assert ( - mock_logging_obj.model_call_details["litellm_params"]["metadata"][ - "error_information" - ]["error_code"] + mock_logging_obj.model_call_details["litellm_params"]["metadata"]["error_information"]["error_code"] == "499" ) @@ -5017,9 +4950,7 @@ class TestStreamingClientDisconnectLogging: mock_request.is_disconnected = AsyncMock(return_value=False) request_data = {"metadata": {}} - recorded = await _record_streaming_client_disconnect_if_needed( - mock_request, request_data - ) + recorded = await _record_streaming_client_disconnect_if_needed(mock_request, request_data) assert recorded is False assert "client_disconnected" not in request_data["metadata"] @@ -5044,22 +4975,12 @@ class TestStreamingClientDisconnectLogging: "litellm_params": {"metadata": {}}, } - recorded = await _record_streaming_client_disconnect_if_needed( - mock_request, request_data - ) + recorded = await _record_streaming_client_disconnect_if_needed(mock_request, request_data) assert recorded is True assert request_data["metadata"]["client_disconnected"] is True - assert ( - mock_logging_obj.model_call_details["litellm_params"]["metadata"][ - "client_disconnected" - ] - is True - ) - assert ( - mock_logging_obj.model_call_details["metadata"]["client_disconnected"] - is True - ) + assert mock_logging_obj.model_call_details["litellm_params"]["metadata"]["client_disconnected"] is True + assert mock_logging_obj.model_call_details["metadata"]["client_disconnected"] is True @pytest.mark.asyncio async def test_record_streaming_client_disconnect_handles_none_request_data_metadata(self): @@ -5075,15 +4996,11 @@ class TestStreamingClientDisconnectLogging: "litellm_params": {"metadata": None}, } - recorded = await _record_streaming_client_disconnect_if_needed( - mock_request, request_data - ) + recorded = await _record_streaming_client_disconnect_if_needed(mock_request, request_data) assert recorded is True assert request_data["metadata"]["client_disconnected"] is True - assert ( - request_data["litellm_params"]["metadata"]["client_disconnected"] is True - ) + assert request_data["litellm_params"]["metadata"]["client_disconnected"] is True @pytest.mark.asyncio async def test_apply_client_disconnect_metadata_none_returns_early(self): @@ -5094,9 +5011,7 @@ class TestStreamingClientDisconnectLogging: _apply_client_disconnect_metadata(None) @pytest.mark.asyncio - async def test_finalize_streaming_generator_cleanup_fires_deferred_logging( - self, monkeypatch - ): + async def test_finalize_streaming_generator_cleanup_fires_deferred_logging(self, monkeypatch): from litellm.proxy.common_request_processing import ( ProxyBaseLLMRequestProcessing, ) @@ -5128,9 +5043,7 @@ class TestStreamingClientDisconnectLogging: assert request_data["metadata"]["error_information"]["error_code"] == "499" @pytest.mark.asyncio - async def test_finalize_streaming_generator_cleanup_skips_disconnect_after_completion( - self, monkeypatch - ): + async def test_finalize_streaming_generator_cleanup_skips_disconnect_after_completion(self, monkeypatch): from litellm.proxy.common_request_processing import ( ProxyBaseLLMRequestProcessing, ) @@ -5160,9 +5073,7 @@ class TestStreamingClientDisconnectLogging: assert "client_disconnected" not in request_data["metadata"] @pytest.mark.asyncio - async def test_async_streaming_data_generator_records_499_on_early_aclose( - self, monkeypatch - ): + async def test_async_streaming_data_generator_records_499_on_early_aclose(self, monkeypatch): from litellm.proxy.common_request_processing import ( ProxyBaseLLMRequestProcessing, ) @@ -5177,9 +5088,7 @@ class TestStreamingClientDisconnectLogging: yield {"choices": [{"delta": {"content": " there"}}]} mock_proxy_logging = MagicMock(spec=ProxyLogging) - mock_proxy_logging.async_post_call_streaming_iterator_hook = ( - mock_streaming_iterator - ) + mock_proxy_logging.async_post_call_streaming_iterator_hook = mock_streaming_iterator ProxyLogging._callback_capabilities_cache.clear() mock_request = MagicMock(spec=Request) @@ -5190,9 +5099,7 @@ class TestStreamingClientDisconnectLogging: "model": "gemini-2.0-flash", "metadata": {}, "litellm_params": {"metadata": {}}, - "litellm_logging_obj": MagicMock( - model_call_details={"metadata": {}, "litellm_params": {}} - ), + "litellm_logging_obj": MagicMock(model_call_details={"metadata": {}, "litellm_params": {}}), } gen = ProxyBaseLLMRequestProcessing.async_streaming_data_generator( @@ -5211,6 +5118,8 @@ class TestStreamingClientDisconnectLogging: assert request_data["metadata"]["error_information"]["error_code"] == "499" ProxyLogging._callback_capabilities_cache.clear() + + class TestCancelOnDisconnect: """ Coverage for the opt-in `general_settings.cancel_on_disconnect` flag: @@ -5237,23 +5146,17 @@ class TestCancelOnDisconnect: llm_call = asyncio.get_running_loop().create_future() disconnect_event = asyncio.Event() - await _cancel_llm_call_on_client_disconnect( - request, llm_call, disconnect_event - ) + await _cancel_llm_call_on_client_disconnect(request, llm_call, disconnect_event) assert llm_call.cancelled() assert disconnect_event.is_set() async def test_monitor_is_noop_while_client_stays_connected(self): - request = self._request( - [{"type": "http.request", "body": b"", "more_body": False}] - ) + request = self._request([{"type": "http.request", "body": b"", "more_body": False}]) llm_call = asyncio.get_running_loop().create_future() disconnect_event = asyncio.Event() - monitor = asyncio.create_task( - _cancel_llm_call_on_client_disconnect(request, llm_call, disconnect_event) - ) + monitor = asyncio.create_task(_cancel_llm_call_on_client_disconnect(request, llm_call, disconnect_event)) await asyncio.sleep(0.01) assert not monitor.done() @@ -5272,9 +5175,7 @@ class TestCancelOnDisconnect: llm_call = asyncio.get_running_loop().create_future() disconnect_event = asyncio.Event() - await _cancel_llm_call_on_client_disconnect( - request, llm_call, disconnect_event - ) + await _cancel_llm_call_on_client_disconnect(request, llm_call, disconnect_event) assert not llm_call.cancelled() assert not disconnect_event.is_set() @@ -5289,9 +5190,7 @@ class TestCancelOnDisconnect: with pytest.raises(asyncio.CancelledError): await _await_llm_call_cancelling_on_disconnect(request, llm_call) - async def _drive_base_process_llm_request( - self, monkeypatch, general_settings: dict, llm_call, request: Request - ): + async def _drive_base_process_llm_request(self, monkeypatch, general_settings: dict, llm_call, request: Request): from litellm.proxy._types import UserAPIKeyAuth logging_obj = MagicMock() @@ -5300,9 +5199,7 @@ class TestCancelOnDisconnect: logging_obj._on_deferred_stream_complete = None logging_obj.cost_breakdown = None - processor = ProxyBaseLLMRequestProcessing( - data={"model": "fake-model", "litellm_logging_obj": logging_obj} - ) + processor = ProxyBaseLLMRequestProcessing(data={"model": "fake-model", "litellm_logging_obj": logging_obj}) proxy_logging_obj = MagicMock(spec=ProxyLogging) proxy_logging_obj.during_call_hook = AsyncMock(return_value=None) @@ -5310,9 +5207,7 @@ class TestCancelOnDisconnect: proxy_logging_obj.post_call_success_hook = AsyncMock( side_effect=lambda data, user_api_key_dict, response: response ) - proxy_logging_obj.post_call_response_headers_hook = AsyncMock( - return_value=None - ) + proxy_logging_obj.post_call_response_headers_hook = AsyncMock(return_value=None) async def fake_route_request(**kwargs): return llm_call() @@ -5391,9 +5286,7 @@ class TestCancelOnDisconnect: with pytest.raises(ProxyException) as exc_info: await processor._handle_llm_api_exception( - e=HTTPException( - status_code=499, detail="Client disconnected the request" - ), + e=HTTPException(status_code=499, detail="Client disconnected the request"), user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"), proxy_logging_obj=proxy_logging_obj, ) @@ -5459,7 +5352,9 @@ class TestAllmPassthroughRoutePostCallGuardrails: proxy_logging_obj = ProxyLogging(user_api_key_cache=MagicMock()) monkeypatch.setattr(proxy_logging_obj, "post_call_success_hook", capture_hook) - with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True): + with patch.object( + ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True + ): processing_obj = ProxyBaseLLMRequestProcessing(data={}) result = await processing_obj._handle_non_streaming_allm_passthrough_route( response=httpx_response, @@ -5509,7 +5404,9 @@ class TestAllmPassthroughRoutePostCallGuardrails: proxy_logging_obj = ProxyLogging(user_api_key_cache=MagicMock()) monkeypatch.setattr(proxy_logging_obj, "post_call_success_hook", non_dict_hook) - with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True): + with patch.object( + ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True + ): processing_obj = ProxyBaseLLMRequestProcessing(data={}) result = await processing_obj._handle_non_streaming_allm_passthrough_route( response=httpx_response, @@ -5547,7 +5444,9 @@ class TestAllmPassthroughRoutePostCallGuardrails: hook_spy = AsyncMock() monkeypatch.setattr(proxy_logging_obj, "post_call_success_hook", hook_spy) - with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True): + with patch.object( + ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True + ): processing_obj = ProxyBaseLLMRequestProcessing(data={}) result = await processing_obj._handle_non_streaming_allm_passthrough_route( response=httpx_response, @@ -5588,7 +5487,9 @@ class TestAllmPassthroughRoutePostCallGuardrails: hook_spy = AsyncMock() monkeypatch.setattr(proxy_logging_obj, "post_call_success_hook", hook_spy) - with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=False): + with patch.object( + ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=False + ): processing_obj = ProxyBaseLLMRequestProcessing(data={}) result = await processing_obj._handle_non_streaming_allm_passthrough_route( response=httpx_response, @@ -5680,9 +5581,9 @@ class TestEventStreamAllmPassthroughRoute: @pytest.mark.asyncio async def test_anthropic_provider_dispatches_to_handler(self): sse = ( - b'event: content_block_delta\n' + b"event: content_block_delta\n" b'data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":""}}\n\n' - b'event: message_stop\n' + b"event: message_stop\n" b'data: {"type":"message_stop"}\n\n' ) @@ -5712,15 +5613,9 @@ class TestEventStreamAllmPassthroughRoute: LlmPassthroughRouteHandler, ) - assert LlmPassthroughRouteHandler.supports_event_stream_de_anonymization( - "anthropic", "/v1/messages" - ) - assert LlmPassthroughRouteHandler.supports_event_stream_de_anonymization( - "anthropic", "messages" - ) - assert not LlmPassthroughRouteHandler.supports_event_stream_de_anonymization( - "anthropic", "/v1/complete" - ) + assert LlmPassthroughRouteHandler.supports_event_stream_de_anonymization("anthropic", "/v1/messages") + assert LlmPassthroughRouteHandler.supports_event_stream_de_anonymization("anthropic", "messages") + assert not LlmPassthroughRouteHandler.supports_event_stream_de_anonymization("anthropic", "/v1/complete") @pytest.mark.asyncio async def test_non_streaming_response_includes_custom_headers(self): @@ -5745,7 +5640,9 @@ class TestEventStreamAllmPassthroughRoute: "content-length": "99", } - with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True): + with patch.object( + ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True + ): processing_obj = ProxyBaseLLMRequestProcessing(data={}) result = await processing_obj._handle_non_streaming_allm_passthrough_route( response=mock_response, @@ -5776,9 +5673,7 @@ class TestAllmPassthroughStreamingProviderGate: de-anonymized. """ - def _build_processing_obj( - self, custom_llm_provider: str, endpoint: str = "" - ) -> ProxyBaseLLMRequestProcessing: + def _build_processing_obj(self, custom_llm_provider: str, endpoint: str = "") -> ProxyBaseLLMRequestProcessing: logging_obj = MagicMock() logging_obj.litellm_call_id = "call-123" logging_obj.cost_breakdown = None @@ -5865,14 +5760,17 @@ class TestAllmPassthroughStreamingProviderGate: processing_obj = self._build_processing_obj("anthropic") chunks = [b"chunk-1", b"chunk-2"] - with patch.object( - ProxyBaseLLMRequestProcessing, - "_has_post_call_guardrails", - return_value=False, - ), patch.object( - ProxyBaseLLMRequestProcessing, - "_has_post_call_guardrails_for_passthrough", - return_value=True, + with ( + patch.object( + ProxyBaseLLMRequestProcessing, + "_has_post_call_guardrails", + return_value=False, + ), + patch.object( + ProxyBaseLLMRequestProcessing, + "_has_post_call_guardrails_for_passthrough", + return_value=True, + ), ): result = await self._run(processing_obj, monkeypatch, chunks) @@ -5881,27 +5779,27 @@ class TestAllmPassthroughStreamingProviderGate: assert streamed == chunks @pytest.mark.asyncio - async def test_bedrock_converse_stream_is_buffered_through_handler( - self, monkeypatch - ): - processing_obj = self._build_processing_obj( - "bedrock", "model/us.amazon.nova-lite-v1:0/converse-stream" - ) + async def test_bedrock_converse_stream_is_buffered_through_handler(self, monkeypatch): + processing_obj = self._build_processing_obj("bedrock", "model/us.amazon.nova-lite-v1:0/converse-stream") chunks = [b"raw-1", b"raw-2"] - with patch.object( - ProxyBaseLLMRequestProcessing, - "_has_post_call_guardrails", - return_value=False, - ), patch.object( - ProxyBaseLLMRequestProcessing, - "_has_post_call_guardrails_for_passthrough", - return_value=True, - ), patch( - "litellm.llms.bedrock.passthrough.guardrail_translation.handler." - "BedrockPassthroughGuardrailHandler.de_anonymize_event_stream", - new=AsyncMock(return_value=b"modified-body"), - ) as mock_handler: + with ( + patch.object( + ProxyBaseLLMRequestProcessing, + "_has_post_call_guardrails", + return_value=False, + ), + patch.object( + ProxyBaseLLMRequestProcessing, + "_has_post_call_guardrails_for_passthrough", + return_value=True, + ), + patch( + "litellm.llms.bedrock.passthrough.guardrail_translation.handler." + "BedrockPassthroughGuardrailHandler.de_anonymize_event_stream", + new=AsyncMock(return_value=b"modified-body"), + ) as mock_handler, + ): result = await self._run(processing_obj, monkeypatch, chunks) assert isinstance(result, Response) @@ -5917,19 +5815,23 @@ class TestAllmPassthroughStreamingProviderGate: ) chunks = [b"raw-1", b"raw-2"] - with patch.object( - ProxyBaseLLMRequestProcessing, - "_has_post_call_guardrails", - return_value=False, - ), patch.object( - ProxyBaseLLMRequestProcessing, - "_has_post_call_guardrails_for_passthrough", - return_value=True, - ), patch( - "litellm.llms.bedrock.passthrough.guardrail_translation.handler." - "BedrockPassthroughGuardrailHandler.de_anonymize_event_stream", - new=AsyncMock(return_value=b"modified-body"), - ) as mock_handler: + with ( + patch.object( + ProxyBaseLLMRequestProcessing, + "_has_post_call_guardrails", + return_value=False, + ), + patch.object( + ProxyBaseLLMRequestProcessing, + "_has_post_call_guardrails_for_passthrough", + return_value=True, + ), + patch( + "litellm.llms.bedrock.passthrough.guardrail_translation.handler." + "BedrockPassthroughGuardrailHandler.de_anonymize_event_stream", + new=AsyncMock(return_value=b"modified-body"), + ) as mock_handler, + ): result = await self._run(processing_obj, monkeypatch, chunks) assert isinstance(result, StreamingResponse) @@ -5951,14 +5853,17 @@ class TestAllmPassthroughStreamingProviderGate: ) chunks = [b"raw-1", b"raw-2"] - with patch.object( - ProxyBaseLLMRequestProcessing, - "_has_post_call_guardrails", - return_value=False, - ), patch.object( - ProxyBaseLLMRequestProcessing, - "_has_post_call_guardrails_for_passthrough", - return_value=False, + with ( + patch.object( + ProxyBaseLLMRequestProcessing, + "_has_post_call_guardrails", + return_value=False, + ), + patch.object( + ProxyBaseLLMRequestProcessing, + "_has_post_call_guardrails_for_passthrough", + return_value=False, + ), ): result = await self._run(processing_obj, monkeypatch, chunks) @@ -5969,22 +5874,25 @@ class TestAllmPassthroughStreamingProviderGate: assert streamed == chunks @pytest.mark.asyncio - async def test_non_bedrock_stream_keeps_default_content_type(self, monkeypatch): + async def test_unregistered_provider_stream_keeps_default_content_type(self, monkeypatch): """ A provider with no registered event-stream media type must not have one forced onto its unbuffered stream, so the response default is unchanged """ - processing_obj = self._build_processing_obj("anthropic") + processing_obj = self._build_processing_obj("gigachat") chunks = [b"chunk-1", b"chunk-2"] - with patch.object( - ProxyBaseLLMRequestProcessing, - "_has_post_call_guardrails", - return_value=False, - ), patch.object( - ProxyBaseLLMRequestProcessing, - "_has_post_call_guardrails_for_passthrough", - return_value=False, + with ( + patch.object( + ProxyBaseLLMRequestProcessing, + "_has_post_call_guardrails", + return_value=False, + ), + patch.object( + ProxyBaseLLMRequestProcessing, + "_has_post_call_guardrails_for_passthrough", + return_value=False, + ), ): result = await self._run(processing_obj, monkeypatch, chunks) @@ -5992,6 +5900,32 @@ class TestAllmPassthroughStreamingProviderGate: assert result.media_type is None assert "content-type" not in result.headers + @pytest.mark.asyncio + async def test_anthropic_stream_uses_event_stream_content_type(self, monkeypatch): + """ + Anthropic now registers a handler, so its unbuffered /v1/messages stream carries + the SSE media type it is actually served with + """ + processing_obj = self._build_processing_obj("anthropic") + chunks = [b"chunk-1", b"chunk-2"] + + with ( + patch.object( + ProxyBaseLLMRequestProcessing, + "_has_post_call_guardrails", + return_value=False, + ), + patch.object( + ProxyBaseLLMRequestProcessing, + "_has_post_call_guardrails_for_passthrough", + return_value=False, + ), + ): + result = await self._run(processing_obj, monkeypatch, chunks) + + assert isinstance(result, StreamingResponse) + assert result.media_type == "text/event-stream" + class TestResponseCostHeaderForTypedDictResponses: """ @@ -6432,9 +6366,7 @@ class TestCostHeadersForCallsPricedAtZero: fastapi_response = Response() processing_obj = ProxyBaseLLMRequestProcessing(data={"litellm_logging_obj": logging_obj}) - with patch.object( - ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails", return_value=False - ): + with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails", return_value=False): await processing_obj.base_process_llm_request( request=MagicMock(spec=Request, headers={}), fastapi_response=fastapi_response, @@ -6505,9 +6437,7 @@ class TestCostHeadersForCallsPricedAtZero: assert breakdown.tool_usage_cost == 0.0 def test_cost_breakdown_stays_empty_for_an_inference_call(self): - breakdown = _get_cost_breakdown_from_logging_obj( - litellm_logging_obj=self._logging_obj(call_type="acompletion") - ) + breakdown = _get_cost_breakdown_from_logging_obj(litellm_logging_obj=self._logging_obj(call_type="acompletion")) assert breakdown == CostBreakdownHeaderValues() @@ -6534,7 +6464,6 @@ class TestCostHeadersForCallsPricedAtZero: class TestPreCallWithFallbacksOnLocalRateLimit: - @pytest.mark.asyncio async def test_fallback_triggered_on_local_rate_limit(self): from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError @@ -6686,9 +6615,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit: mock_router.fallbacks = [{"gpt-4": ["gpt-3.5-turbo"]}] user_api_key_dict = MagicMock() - user_api_key_dict.router_settings = { - "fallbacks": [{"gpt-4": ["claude-3-haiku"]}] - } + user_api_key_dict.router_settings = {"fallbacks": [{"gpt-4": ["claude-3-haiku"]}]} with patch.object( processor, @@ -6719,9 +6646,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit: from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing - processor = ProxyBaseLLMRequestProcessing( - data={"model": "gpt-4", "disable_fallbacks": True} - ) + processor = ProxyBaseLLMRequestProcessing(data={"model": "gpt-4", "disable_fallbacks": True}) async def mock_pre_call_logic(**kwargs): raise ProxyRateLimitError( @@ -6847,9 +6772,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit: # Real per-key per-model TPM limiter + a key carrying the customer's # `model_tpm_limit` metadata (only the primary is capped). - limiter = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + limiter = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) user_api_key_dict = UserAPIKeyAuth( api_key="sk-lit3890", metadata={"model_tpm_limit": {primary_model: 100}}, @@ -6857,10 +6780,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit: # Pre-seed the primary's per-model token counter at the cap so the very # next request trips it. The counter key uses the *hashed* api_key. - counter_key = ( - f"{user_api_key_dict.api_key}::{primary_model}" - f"::{precise_minute}::request_count" - ) + counter_key = f"{user_api_key_dict.api_key}::{primary_model}::{precise_minute}::request_count" await limiter.internal_usage_cache.async_set_cache( key=counter_key, value={"current_requests": 0, "current_tpm": 100, "current_rpm": 0}, @@ -6891,9 +6811,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit: mock_router = MagicMock() mock_router.fallbacks = [{primary_model: [fallback_model]}] - with patch( - "litellm.proxy.hooks.parallel_request_limiter.datetime", _FrozenClock - ): + with patch("litellm.proxy.hooks.parallel_request_limiter.datetime", _FrozenClock): with patch.object( processor, "common_processing_pre_call_logic", @@ -6923,9 +6841,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit: # Sanity-check the premise: the limiter genuinely raises a # ProxyRateLimitError for the capped primary under the frozen clock. - with patch( - "litellm.proxy.hooks.parallel_request_limiter.datetime", _FrozenClock - ): + with patch("litellm.proxy.hooks.parallel_request_limiter.datetime", _FrozenClock): with pytest.raises(ProxyRateLimitError): await limiter.async_pre_call_hook( user_api_key_dict=user_api_key_dict, @@ -7570,16 +7486,12 @@ class TestStreamingClientDisconnectBilling: prompt_tokens=1000, completion_tokens=10, total_tokens=1010, - prompt_tokens_details=PromptTokensDetailsWrapper( - cached_tokens=500 - ), + prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=500), ), ) ) - event = await self._bill_and_collect_success_event( - append_openai_style_cached_usage_chunk - ) + event = await self._bill_and_collect_success_event(append_openai_style_cached_usage_chunk) usage = event["response_obj"].usage assert getattr(usage, "cache_read_input_tokens", None) == 500 @@ -7870,12 +7782,15 @@ class TestModelDeploymentsSupportStreamOptions: @pytest.mark.asyncio -@pytest.mark.parametrize("key_settings, expected", [ - (None, {"group": {"team": 100}}), - ({"weights": {"group": {"key": 100}}}, {"group": {"key": 100}}), - ({"timeout": 30}, None), - ({"weights": {"group": {"key": "legacy"}}}, None), -]) +@pytest.mark.parametrize( + "key_settings, expected", + [ + (None, {"group": {"team": 100}}), + ({"weights": {"group": {"key": 100}}}, {"group": {"key": 100}}), + ({"timeout": 30}, None), + ({"weights": {"group": {"key": "legacy"}}}, None), + ], +) async def test_saved_weights_override_caller_input_and_preserve_key_precedence( monkeypatch: pytest.MonkeyPatch, key_settings: dict[str, int | dict[str, dict[str, int | str]]] | None, @@ -7884,14 +7799,20 @@ async def test_saved_weights_override_caller_input_and_preserve_key_precedence( from litellm.proxy import proxy_server monkeypatch.setattr(proxy_server, "prisma_client", None) - monkeypatch.setattr(proxy_server, "get_team_object", AsyncMock( - return_value=SimpleNamespace(router_settings={"weights": {"group": {"team": 100}}}) - )) + monkeypatch.setattr( + proxy_server, + "get_team_object", + AsyncMock(return_value=SimpleNamespace(router_settings={"weights": {"group": {"team": 100}}})), + ) forged = {"group": {"caller": 100}} - processor = ProxyBaseLLMRequestProcessing(data={ - "model": "group", "weights": forged, "_router_weights": forged, - "router_settings_override": {"weights": forged}, - }) + processor = ProxyBaseLLMRequestProcessing( + data={ + "model": "group", + "weights": forged, + "_router_weights": forged, + "router_settings_override": {"weights": forged}, + } + ) logging = MagicMock(spec=ProxyLogging) logging.pre_call_hook = AsyncMock(side_effect=lambda **kwargs: kwargs["data"]) data, _ = await processor.common_processing_pre_call_logic( @@ -8388,9 +8309,7 @@ class TestInjectCostIntoUsageDict: logging_obj.model_call_details["custom_llm_provider"] = "anthropic" assert logging_obj.cost_breakdown is None - model_response = ModelResponse( - usage=Usage(prompt_tokens=3216, completion_tokens=8, total_tokens=3224) - ) + model_response = ModelResponse(usage=Usage(prompt_tokens=3216, completion_tokens=8, total_tokens=3224)) cost = ProxyBaseLLMRequestProcessing._logging_obj_cost_or_none(model_response, logging_obj) assert cost is not None and cost > 0 @@ -8419,9 +8338,7 @@ class TestInjectCostIntoUsageDict: ) existing = logging_obj.cost_breakdown - model_response = ModelResponse( - usage=Usage(prompt_tokens=3216, completion_tokens=8, total_tokens=3224) - ) + model_response = ModelResponse(usage=Usage(prompt_tokens=3216, completion_tokens=8, total_tokens=3224)) ProxyBaseLLMRequestProcessing._logging_obj_cost_or_none(model_response, logging_obj) assert logging_obj.cost_breakdown is existing @@ -8716,9 +8633,7 @@ def test_ttft_keepalive_interval_only_arms_for_a_streaming_request(request_data, @pytest.mark.asyncio @pytest.mark.parametrize("stream_requested, expect_ping", [(True, True), (False, False)]) -async def test_base_process_llm_request_pings_while_the_upstream_call_is_still_running( - stream_requested, expect_ping -): +async def test_base_process_llm_request_pings_while_the_upstream_call_is_still_running(stream_requested, expect_ping): """The wiring, not the helper: every route funnels through this method, and the whole time-to-first-token is spent inside the call it wraps.""" @@ -8864,9 +8779,7 @@ async def test_a_late_failure_is_reported_to_the_failure_hook(): async def record(exc): audited.append(exc) - response = await open_sse_before_first_byte( - slow_failure(), ping_interval_seconds=0.05, on_late_failure=record - ) + response = await open_sse_before_first_byte(slow_failure(), ping_interval_seconds=0.05, on_late_failure=record) collected = await _drain(response) assert [type(exc).__name__ for exc in audited] == ["HTTPException"] @@ -8883,9 +8796,7 @@ async def test_a_failing_audit_hook_never_costs_the_client_its_error_frame(): async def broken_hook(exc): raise RuntimeError("the audit backend is down") - response = await open_sse_before_first_byte( - slow_failure(), ping_interval_seconds=0.05, on_late_failure=broken_hook - ) + response = await open_sse_before_first_byte(slow_failure(), ping_interval_seconds=0.05, on_late_failure=broken_hook) collected = await _drain(response) error_frame = json.loads(collected[-2].decode().removeprefix("data: ").strip()) @@ -8939,9 +8850,7 @@ async def test_base_process_llm_request_audits_a_failure_that_lands_after_its_ke [(0, False), (None, True)], ids=["operator-hard-disabled-this-deployment", "deployment-says-nothing"], ) -async def test_base_process_llm_request_honours_a_deployment_hard_disable( - deployment_keepalive, expect_ping -): +async def test_base_process_llm_request_honours_a_deployment_hard_disable(deployment_keepalive, expect_ping): """`keepalive_seconds: 0` is documented as a disable a request cannot lift. The funnel has to hand its router to the gate for that to hold before the upstream has answered, since no deployment has served the request yet.""" @@ -8987,9 +8896,7 @@ async def test_a_hook_returning_a_replacement_decides_what_the_client_sees(): async def sanitize(exc): return HTTPException(status_code=502, detail="upstream unavailable") - response = await open_sse_before_first_byte( - slow_failure(), ping_interval_seconds=0.05, on_late_failure=sanitize - ) + response = await open_sse_before_first_byte(slow_failure(), ping_interval_seconds=0.05, on_late_failure=sanitize) collected = await _drain(response) error_frame = json.loads(collected[-2].decode().removeprefix("data: ").strip()) @@ -9028,9 +8935,7 @@ async def test_a_hook_that_returns_nothing_leaves_the_real_error_intact(): async def audit_only(exc): return None - response = await open_sse_before_first_byte( - slow_failure(), ping_interval_seconds=0.05, on_late_failure=audit_only - ) + response = await open_sse_before_first_byte(slow_failure(), ping_interval_seconds=0.05, on_late_failure=audit_only) collected = await _drain(response) error_frame = json.loads(collected[-2].decode().removeprefix("data: ").strip()) @@ -9047,9 +8952,7 @@ async def test_a_broken_hook_does_not_replace_the_real_error_with_its_own_bug(): async def broken_hook(exc): raise RuntimeError("the audit backend is down") - response = await open_sse_before_first_byte( - slow_failure(), ping_interval_seconds=0.05, on_late_failure=broken_hook - ) + response = await open_sse_before_first_byte(slow_failure(), ping_interval_seconds=0.05, on_late_failure=broken_hook) collected = await _drain(response) error_frame = json.loads(collected[-2].decode().removeprefix("data: ").strip()) @@ -9316,9 +9219,7 @@ class TestStreamingResponseHeadersFollowFallback: proxy_logging_obj.post_call_success_hook = AsyncMock( side_effect=lambda data, user_api_key_dict, response: response ) - proxy_logging_obj.post_call_response_headers_hook = AsyncMock( - return_value={"x-callback-header": "kept"} - ) + proxy_logging_obj.post_call_response_headers_hook = AsyncMock(return_value={"x-callback-header": "kept"}) async def fake_route_request(**kwargs): async def call(): @@ -9326,9 +9227,7 @@ class TestStreamingResponseHeadersFollowFallback: return call() - monkeypatch.setattr( - litellm.proxy.common_request_processing, "route_request", fake_route_request - ) + monkeypatch.setattr(litellm.proxy.common_request_processing, "route_request", fake_route_request) result = await processor.base_process_llm_request( request=Request(scope={"type": "http", "headers": []}), From e8adcabcd8ac03500973b95ccc705729bbce39c5 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Micha=C5=82=20Furga=C5=82a?= <83299832+00200200@users.noreply.github.com> Date: Wed, 23 Sep 2026 15:52:04 +0200 Subject: [PATCH 3/7] fix(proxy): build the Anthropic SSE rewrite without mutable accumulators The type-discipline gate failed because the new handler added 10 LIT002 violations, over the codebase ceiling. Parse blocks into a tuple, collect the text deltas with a generator, and rebuild the stream in one join instead of appending to lists. The synthetic response stays a dict, as in the Bedrock handler, because post-call guardrail hooks receive and may rewrite it. Behaviour is unchanged; the gate now passes against the merge base. Co-Authored-By: Claude Opus 5.5 --- .../guardrail_translation/handler.py | 96 +++++++++---------- 1 file changed, 46 insertions(+), 50 deletions(-) diff --git a/litellm/llms/anthropic/passthrough/guardrail_translation/handler.py b/litellm/llms/anthropic/passthrough/guardrail_translation/handler.py index cb5532c93e5..d9c508c959e 100644 --- a/litellm/llms/anthropic/passthrough/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/passthrough/guardrail_translation/handler.py @@ -24,19 +24,14 @@ def _is_messages_endpoint(endpoint: str) -> bool: return any(normalized.endswith(suffix) for suffix in _MESSAGES_SUFFIXES) -def _parse_sse_blocks(body_bytes: bytes) -> list[bytes]: +def _parse_sse_blocks(body_bytes: bytes) -> tuple[bytes, ...]: """Split an SSE body into event blocks (including trailing separators).""" if not body_bytes: - return [] + return () # Keep separators so we can rebuild the stream byte-for-byte aside from rewrites. - parts = body_bytes.split(b"\n\n") - blocks: list[bytes] = [] - for i, part in enumerate(parts): - if i < len(parts) - 1: - blocks.append(part + b"\n\n") - elif part: - blocks.append(part) - return blocks + 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) def _event_payload(block: bytes) -> tuple[str | None, dict[str, Any] | None]: @@ -62,6 +57,31 @@ 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.""" + event_type, payload = _event_payload(block) + if event_type != "content_block_delta" or not payload: + return None + delta = payload.get("delta") + if not isinstance(delta, dict) or delta.get("type") != "text_delta": + return None + text = delta.get("text") + return 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.""" + event_type, payload = _event_payload(block) + if event_type != "content_block_delta" or not payload: + return block + delta = payload.get("delta") + if not isinstance(delta, dict): + 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() + + class AnthropicPassthroughGuardrailHandler(BaseTranslation): @staticmethod def is_event_stream_content_type(content_type: str) -> bool: @@ -91,27 +111,14 @@ class AnthropicPassthroughGuardrailHandler(BaseTranslation): first, then redistribute the de-anonymized text across the original frames (full rewrite on the first text_delta, empty on the rest). """ - blocks = _parse_sse_blocks(body_bytes) - text_block_indices: list[int] = [] - texts: list[str] = [] - - for idx, block in enumerate(blocks): - event_type, payload = _event_payload(block) - if event_type != "content_block_delta" or not payload: - continue - delta = payload.get("delta") - if not isinstance(delta, dict) or delta.get("type") != "text_delta": - continue - text = delta.get("text") - if not isinstance(text, str): - continue - text_block_indices.append(idx) - texts.append(text) - - if not texts: + 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 + ) + if not deltas: return body_bytes - combined = "".join(texts) + combined: Final = "".join(text for _, text in deltas) synthetic_response: Final[dict] = { "type": "message", "role": "assistant", @@ -119,7 +126,7 @@ class AnthropicPassthroughGuardrailHandler(BaseTranslation): "stop_reason": "end_turn", } - processed = await proxy_logging_obj.post_call_success_hook( + processed: Final = await proxy_logging_obj.post_call_success_hook( data=data, user_api_key_dict=user_api_key_dict, response=synthetic_response, @@ -133,31 +140,20 @@ class AnthropicPassthroughGuardrailHandler(BaseTranslation): return body_bytes try: - content = processed["content"] - de_anonymized = content[0]["text"] - if not isinstance(de_anonymized, str): - return body_bytes + de_anonymized: Final = processed["content"][0]["text"] except (KeyError, IndexError, TypeError): return body_bytes + if not isinstance(de_anonymized, str): + return body_bytes # Put the full rewrite on the first text_delta; blank the rest so # split placeholders cannot survive across frames. - replacements = [de_anonymized] + [""] * (len(text_block_indices) - 1) - out: list[bytes] = list(blocks) - for block_idx, new_text in zip(text_block_indices, replacements): - event_type, payload = _event_payload(out[block_idx]) - if event_type != "content_block_delta" or not payload: - continue - delta = payload.get("delta") - if not isinstance(delta, dict): - continue - delta["text"] = new_text - payload["delta"] = delta - # Rebuild a minimal SSE block; preserve event name. - new_block = (f"event: content_block_delta\ndata: {json.dumps(payload, separators=(',', ':'))}\n\n").encode() - out[block_idx] = new_block - - return b"".join(out) + delta_indices: Final = frozenset(idx for idx, _ in deltas) + first_idx: Final = deltas[0][0] + 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) + ) async def process_input_messages( self, From f77a9f92287479e0641114c9d7677ce6191cbd2b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Micha=C5=82=20Furga=C5=82a?= <83299832+00200200@users.noreply.github.com> Date: Wed, 23 Sep 2026 21:07:44 +0200 Subject: [PATCH 4/7] fix(proxy): read the processed Anthropic text without unchecked indexing processed["content"][0]["text"] indexed a union of response TypedDicts and pydantic blocks, adding reportGeneralTypeIssues, reportIndexIssue and reportOptionalSubscript errors over the basedpyright budget. Read the first block through isinstance checks instead; behaviour is unchanged and every basedpyright and LIT rule is at or below its previous count. Co-Authored-By: Claude Opus 5.5 --- .../guardrail_translation/handler.py | 19 ++++++++++++++----- 1 file changed, 14 insertions(+), 5 deletions(-) diff --git a/litellm/llms/anthropic/passthrough/guardrail_translation/handler.py b/litellm/llms/anthropic/passthrough/guardrail_translation/handler.py index d9c508c959e..2a93bed9c5c 100644 --- a/litellm/llms/anthropic/passthrough/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/passthrough/guardrail_translation/handler.py @@ -82,6 +82,18 @@ def _with_text(block: bytes, new_text: str) -> bytes: return f"event: content_block_delta\ndata: {json.dumps(payload, separators=(',', ':'))}\n\n".encode() +def _first_text(processed: Mapping[str, object]) -> str | None: + """The text of the first content block in a guardrail-processed Messages response.""" + content: Final = processed.get("content") + if not isinstance(content, list) or not content: + 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 + + class AnthropicPassthroughGuardrailHandler(BaseTranslation): @staticmethod def is_event_stream_content_type(content_type: str) -> bool: @@ -139,11 +151,8 @@ class AnthropicPassthroughGuardrailHandler(BaseTranslation): ) return body_bytes - try: - de_anonymized: Final = processed["content"][0]["text"] - except (KeyError, IndexError, TypeError): - return body_bytes - if not isinstance(de_anonymized, str): + de_anonymized: Final = _first_text(processed) + if de_anonymized is None: return body_bytes # Put the full rewrite on the first text_delta; blank the rest so 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 5/7] 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 ( From 5ac545edd385eb9c9cd7e4281da0a6e5b10ca61a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Micha=C5=82=20Furga=C5=82a?= <83299832+00200200@users.noreply.github.com> Date: Tue, 29 Sep 2026 12:21:45 +0200 Subject: [PATCH 6/7] test(proxy): pin Anthropic SSE CRLF and multi-block guardrail rewrites Regression coverage for the Greptile P1s so CRLF-framed streams still hit post-call rewriting and each content-block index keeps its own text --- .../guardrail_translation/__init__.py | 0 .../guardrail_translation/test_handler.py | 145 ++++++++++++++++++ 2 files changed, 145 insertions(+) create mode 100644 tests/unit/llms/anthropic/passthrough/guardrail_translation/__init__.py create mode 100644 tests/unit/llms/anthropic/passthrough/guardrail_translation/test_handler.py diff --git a/tests/unit/llms/anthropic/passthrough/guardrail_translation/__init__.py b/tests/unit/llms/anthropic/passthrough/guardrail_translation/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/anthropic/passthrough/guardrail_translation/test_handler.py b/tests/unit/llms/anthropic/passthrough/guardrail_translation/test_handler.py new file mode 100644 index 00000000000..7fc896b9c4f --- /dev/null +++ b/tests/unit/llms/anthropic/passthrough/guardrail_translation/test_handler.py @@ -0,0 +1,145 @@ +""" +Regression tests for AnthropicPassthroughGuardrailHandler SSE rewriting. + +Pins the two Greptile P1s from PR #42585: +- CRLF-framed streams must still reach post-call rewriting +- Multi-index text blocks must keep their own rewrites (not merge into index 0) +""" + +from __future__ import annotations + +import json +from unittest.mock import MagicMock + +import pytest + +from litellm.llms.anthropic.passthrough.guardrail_translation.handler import ( + AnthropicPassthroughGuardrailHandler, + _parse_sse_blocks, +) + + +def _text_delta_frame(index: int, text: str, sep: bytes = b"\n\n", line_end: bytes | None = None) -> bytes: + if line_end is None: + line_end = sep[: len(sep) // 2] or b"\n" + payload = { + "type": "content_block_delta", + "index": index, + "delta": {"type": "text_delta", "text": text}, + } + return ( + b"event: content_block_delta" + + line_end + + b"data: " + + json.dumps(payload, separators=(",", ":")).encode() + + sep + ) + + +def _message_stop_frame(sep: bytes = b"\n\n", line_end: bytes | None = None) -> bytes: + if line_end is None: + line_end = sep[: len(sep) // 2] or b"\n" + return b"event: message_stop" + line_end + b'data: {"type":"message_stop"}' + sep + + +def _frame_payloads(body: bytes) -> list[dict]: + payloads: list[dict] = [] + for block in _parse_sse_blocks(body): + for line in block.decode().splitlines(): + if line.startswith("data:"): + payloads.append(json.loads(line[5:].strip())) + break + return payloads + + +class TestParseSseBlocks: + def test_splits_lf_crlf_and_cr_blank_lines(self): + body = ( + b"event: a\ndata: 1\n\n" + b"event: b\r\ndata: 2\r\n\r\n" + b"event: c\rdata: 3\r\r" + ) + blocks = _parse_sse_blocks(body) + assert len(blocks) == 3 + assert blocks[0].endswith(b"\n\n") + assert blocks[1].endswith(b"\r\n\r\n") + assert blocks[2].endswith(b"\r\r") + + +class TestDeAnonymizeEventStream: + def _proxy(self, mock_hook) -> MagicMock: + proxy_logging_obj = MagicMock() + proxy_logging_obj.post_call_success_hook = mock_hook + return proxy_logging_obj + + @pytest.mark.asyncio + async def test_crlf_framed_stream_still_invokes_guardrail(self): + """P1: CRLF frames must not merge so message_stop wins and deltas skip rewriting.""" + sse = _text_delta_frame(0, "", sep=b"\r\n\r\n") + _message_stop_frame(sep=b"\r\n\r\n") + hook_calls: list[dict] = [] + + async def mock_hook(data, user_api_key_dict, response): + hook_calls.append(response) + response = dict(response) + response["content"] = [{"type": "text", "text": "Alice"}] + return response + + # Precondition: a naive LF-only split would leave one merged block. + assert len(sse.split(b"\n\n")) == 1 + + result = await AnthropicPassthroughGuardrailHandler.de_anonymize_event_stream( + body_bytes=sse, + proxy_logging_obj=self._proxy(mock_hook), + user_api_key_dict=MagicMock(), + data={}, + ) + + assert len(hook_calls) == 1 + assert hook_calls[0]["content"][0]["text"] == "" + assert b"" not in result + assert b"Alice" in result + assert result.endswith(b'data: {"type":"message_stop"}\r\n\r\n') + + @pytest.mark.asyncio + async def test_multi_index_text_blocks_keep_their_own_rewrites(self): + """P1: rewrites must stay on their content-block index around tool/thinking events.""" + tool_delta = ( + b"event: content_block_delta\n" + b'data: {"type":"content_block_delta","index":1,' + b'"delta":{"type":"input_json_delta","partial_json":"{}"}}\n\n' + ) + sse = ( + _text_delta_frame(0, " said") + + tool_delta + + _text_delta_frame(2, "bye ") + ) + seen: dict[str, list[str]] = {} + + 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 + + result = await AnthropicPassthroughGuardrailHandler.de_anonymize_event_stream( + body_bytes=sse, + proxy_logging_obj=self._proxy(mock_hook), + user_api_key_dict=MagicMock(), + data={}, + ) + + assert seen["texts"] == [" said", "bye "] + payloads = _frame_payloads(result) + assert [(p["index"], p["delta"].get("text")) for p in payloads] == [ + (0, "Alice said"), + (0, ""), + (1, None), + (2, "bye Bob"), + ] + # Tool frame between text blocks must stay untouched and in order. + assert payloads[2]["delta"]["type"] == "input_json_delta" + assert payloads[2]["delta"]["partial_json"] == "{}" From a28fe54e834077fe5cc29d53f034229358b7c180 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Micha=C5=82=20Furga=C5=82a?= <83299832+00200200@users.noreply.github.com> Date: Fri, 2 Oct 2026 20:54:58 +0200 Subject: [PATCH 7/7] fix(anthropic): join multiline SSE data before guardrail processing --- .../guardrail_translation/handler.py | 6 ++-- .../guardrail_translation/test_handler.py | 29 +++++++++++++++++++ 2 files changed, 32 insertions(+), 3 deletions(-) diff --git a/litellm/llms/anthropic/passthrough/guardrail_translation/handler.py b/litellm/llms/anthropic/passthrough/guardrail_translation/handler.py index 282f86f1fb0..788aefaf723 100644 --- a/litellm/llms/anthropic/passthrough/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/passthrough/guardrail_translation/handler.py @@ -55,9 +55,9 @@ def _event_payload(block: bytes) -> tuple[str | None, Mapping[str, JsonValue] | text: Final = block.decode("utf-8") except UnicodeDecodeError: return None, None - lines: Final = tuple(reversed(text.splitlines())) - event_type: Final = next((line[6:].strip() for line in lines if line.startswith("event:")), None) - data_line: Final = next((line[5:].strip() for line in lines if line.startswith("data:")), None) + lines: Final = tuple(text.splitlines()) + event_type: Final = next((line[6:].strip() for line in reversed(lines) if line.startswith("event:")), None) + data_line: Final = "\n".join(line[5:].removeprefix(" ") for line in lines if line.startswith("data:")) if not data_line: return event_type, None try: diff --git a/tests/unit/llms/anthropic/passthrough/guardrail_translation/test_handler.py b/tests/unit/llms/anthropic/passthrough/guardrail_translation/test_handler.py index 78055fe535a..0594caedb04 100644 --- a/tests/unit/llms/anthropic/passthrough/guardrail_translation/test_handler.py +++ b/tests/unit/llms/anthropic/passthrough/guardrail_translation/test_handler.py @@ -122,6 +122,35 @@ class TestDeAnonymizeEventStream: proxy_logging_obj.post_call_success_hook = mock_hook return proxy_logging_obj + @pytest.mark.asyncio + @pytest.mark.parametrize("line_end", [b"\n", b"\r\n", b"\r"]) + async def test_multiline_data_reaches_guardrail(self, line_end: bytes): + frame = line_end.join( + ( + b"event: content_block_delta", + b'data: {"type":"content_block_delta","index":3,', + b'data: "delta":{"type":"text_delta","text":""}}', + b"", + b"", + ) + ) + stop = _message_stop_frame(sep=line_end * 2) + + async def hook(data, user_api_key_dict, response): + assert response["content"] == [{"type": "text", "text": ""}] + return {**response, "content": [{"type": "text", "text": "Alice"}]} + + result = await AnthropicPassthroughGuardrailHandler.de_anonymize_event_stream( + body_bytes=frame + stop, + proxy_logging_obj=self._proxy(hook), + user_api_key_dict=MagicMock(), + data={}, + ) + + assert _text_delta(_parse_sse_blocks(result)[0]) == (3, "Alice") + assert b"" not in result + assert result.endswith(stop) + @pytest.mark.asyncio async def test_crlf_framed_stream_still_invokes_guardrail(self): """P1: CRLF frames must not merge so message_stop wins and deltas skip rewriting."""