mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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
This commit is contained in:
parent
c6c3881d7f
commit
802a57e5f7
5 changed files with 252 additions and 3 deletions
0
litellm/llms/anthropic/passthrough/__init__.py
Normal file
0
litellm/llms/anthropic/passthrough/__init__.py
Normal file
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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":"<PERSON_1>"}}\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"<PERSON_1>" 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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue