mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge a28fe54e83 into 1b98748528
This commit is contained in:
commit
5ae21b884b
8 changed files with 788 additions and 315 deletions
0
litellm/llms/anthropic/passthrough/__init__.py
Normal file
0
litellm/llms/anthropic/passthrough/__init__.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
0
tests/unit/llms/anthropic/passthrough/__init__.py
Normal file
0
tests/unit/llms/anthropic/passthrough/__init__.py
Normal 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
Loading…
Add table
Reference in a new issue