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] 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