This commit is contained in:
00200200 2026-10-03 21:29:18 +02:00 • committed by GitHub
commit 5ae21b884b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 788 additions and 315 deletions

View file

@ -0,0 +1,226 @@
"""Anthropic /v1/messages passthrough guardrail translation (SSE event stream)."""
from __future__ import annotations
import json
import re
from collections.abc import Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, ValidationError
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
from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse
_EVENT_STREAM_MEDIA_TYPE: Final = "text/event-stream"
_MESSAGES_SUFFIXES: Final = frozenset({"messages", "v1/messages"})
_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")
_SSE_PAYLOAD: Final = TypeAdapter(Mapping[str, JsonValue])
class _GuardrailText(BaseModel):
model_config = ConfigDict(frozen=True, strict=True)
text: str
_GUARDRAIL_TEXTS: Final = TypeAdapter(tuple[_GuardrailText, ...])
def _is_messages_endpoint(endpoint: str) -> bool:
normalized: Final = endpoint.rstrip("/").split("?")[0]
return any(normalized.endswith(suffix) for suffix in _MESSAGES_SUFFIXES)
def _parse_sse_blocks(body_bytes: bytes) -> tuple[bytes, ...]:
"""Split an SSE body into event blocks, each keeping its own trailing separator."""
if not body_bytes:
return ()
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, Mapping[str, JsonValue] | None]:
try:
text: Final = block.decode("utf-8")
except UnicodeDecodeError:
return None, 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:
payload: Final = _SSE_PAYLOAD.validate_json(data_line)
except ValidationError:
return event_type, None
return event_type, payload
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: Final = payload.get("index")
delta: Final = payload.get("delta")
if not isinstance(index, int) or not isinstance(delta, Mapping) or delta.get("type") != "text_delta":
return None
text: Final = delta.get("text")
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``, keeping its framing."""
event_type, payload = _event_payload(block)
if event_type != "content_block_delta" or not payload:
return block
delta: Final = payload.get("delta")
if not isinstance(delta, Mapping):
return block
rewritten: Final = {**payload, "delta": {**delta, "text": new_text}}
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(rewritten, separators=(",", ":"))
return f"event: content_block_delta{line_end}data: {data}".encode() + separator
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):
return None
try:
first: Final = _GUARDRAIL_TEXTS.validate_python(content[:count])
except ValidationError:
return None
texts: Final = tuple(block.text for block in first)
return texts if len(texts) == count else None
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[str, object], # mutable-ok: post-call hooks share and update the request dictionary
) -> 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
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(
(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
indices: Final = tuple(sorted(frozenset(index for _, (index, _) in deltas)))
synthetic_response: Final[AnthropicMessagesResponse] = {
"type": "message",
"role": "assistant",
"content": [
{"type": "text", "text": "".join(text for _, (i, text) in deltas if i == index)} for index in indices
],
"stop_reason": "end_turn",
}
processed: Final = await proxy_logging_obj.post_call_success_hook( # pyright: ignore[reportUnknownMemberType] # untyped request
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
rewritten: Final = _processed_texts(processed, len(indices))
if rewritten is None:
return body_bytes
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, 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(
self,
data: dict[str, object], # mutable-ok: the delegated guardrail updates the original request dictionary
guardrail_to_apply: CustomGuardrail,
litellm_logging_obj: LiteLLMLoggingObj | None = None,
) -> Mapping[str, object]:
from litellm.llms.pass_through.guardrail_translation.handler import (
PassThroughEndpointHandler,
)
handler: Final = PassThroughEndpointHandler()
return await handler.process_input_messages( # pyright: ignore[reportUnknownMemberType] # untyped request
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: LiteLLMLoggingObj | None = None,
user_api_key_dict: UserAPIKeyAuth | None = None,
request_data: dict[str, object] | None = None, # mutable-ok: delegated hooks update shared request state
) -> object:
from litellm.llms.pass_through.guardrail_translation.handler import (
PassThroughEndpointHandler,
)
handler: Final = PassThroughEndpointHandler()
return await handler.process_output_response( # pyright: ignore[reportUnknownMemberType] # untyped request
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,
)

View file

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

View file

@ -0,0 +1,224 @@
"""
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 collections.abc import Mapping
from unittest.mock import MagicMock
import pytest
from litellm.llms.anthropic.passthrough.guardrail_translation.handler import (
AnthropicPassthroughGuardrailHandler,
_parse_sse_blocks,
_processed_texts,
_text_delta,
_with_text,
)
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\nevent: b\r\ndata: 2\r\n\r\nevent: 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")
@pytest.mark.parametrize("separator", [b"\n\n", b"\r\n\r\n", b"\r\r", b""])
def test_text_rewrite_preserves_index_metadata_and_framing(separator: bytes) -> None:
line_end = separator[: len(separator) // 2] or b"\n"
payload = {
"type": "content_block_delta",
"index": 3,
"delta": {"type": "text_delta", "text": "<PERSON_1>", "metadata": {"source": "guardrail"}},
"extra": [True, None, {"value": 1}],
}
frame = b"event: content_block_delta" + line_end + b"data: " + json.dumps(payload).encode() + separator
rewritten = _with_text(frame, "Alice")
assert rewritten == (
b"event: content_block_delta"
+ line_end
+ b"data: "
+ json.dumps({**payload, "delta": {**payload["delta"], "text": "Alice"}}, separators=(",", ":")).encode()
+ separator
)
@pytest.mark.parametrize(
"frame",
[
b"event: content_block_delta\ndata: {invalid}\n\n",
b"event: content_block_delta\ndata: []\n\n",
b"event: content_block_delta\ndata: null\n\n",
b"event: content_block_delta\ndata: \xff\n\n",
b'event: content_block_delta\ndata: {"index":0,"delta":"text"}\n\n',
_message_stop_frame(),
],
)
def test_invalid_or_non_delta_frames_stay_unchanged(frame: bytes) -> None:
assert _text_delta(frame) is None
assert _with_text(frame, "Alice") == frame
@pytest.mark.parametrize(
"processed, expected",
[
({"content": [{"text": "Alice"}, {"text": "Bob"}]}, ("Alice", "Bob")),
({"content": [{"text": "Alice"}]}, None),
({"content": [{"text": "Alice"}, {"text": 42}]}, None),
({"content": [{"text": "Alice"}, {"type": "text"}]}, None),
({"content": None}, None),
],
)
def test_guardrail_texts_require_a_complete_string_rewrite(
processed: Mapping[str, object], expected: tuple[str, ...] | None
) -> None:
assert _processed_texts(processed, 2) == expected
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
@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":"<PERSON_1>"}}',
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": "<PERSON_1>"}]
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"<PERSON_1>" 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."""
sse = _text_delta_frame(0, "<PERSON_1>", 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"] == "<PERSON_1>"
assert b"<PERSON_1>" 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, "<PERSON_")
+ _text_delta_frame(0, "1> said")
+ tool_delta
+ _text_delta_frame(2, "bye <PERSON_2>")
)
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"] == ["<PERSON_1> said", "bye <PERSON_2>"]
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"] == "{}"

File diff suppressed because it is too large Load diff