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..788aefaf723 --- /dev/null +++ b/litellm/llms/anthropic/passthrough/guardrail_translation/handler.py @@ -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, + ) 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/unit/llms/anthropic/passthrough/__init__.py b/tests/unit/llms/anthropic/passthrough/__init__.py new file mode 100644 index 00000000000..e69de29bb2d 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..0594caedb04 --- /dev/null +++ b/tests/unit/llms/anthropic/passthrough/guardrail_translation/test_handler.py @@ -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": "", "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":""}}', + 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.""" + 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"] == "{}" diff --git a/tests/unit/proxy/test_common_request_processing.py b/tests/unit/proxy/test_common_request_processing.py index 8485c286a30..fc0c8db64b7 100644 --- a/tests/unit/proxy/test_common_request_processing.py +++ b/tests/unit/proxy/test_common_request_processing.py @@ -113,9 +113,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"}], } ] @@ -187,9 +185,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 @@ -275,16 +276,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", @@ -324,14 +321,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", @@ -373,15 +366,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", @@ -400,9 +389,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, @@ -637,9 +624,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, @@ -752,9 +737,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, @@ -897,9 +880,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.""" @@ -2800,16 +2781,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 () @@ -3556,9 +3531,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" @@ -3577,9 +3550,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, ) @@ -3600,9 +3571,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" @@ -3643,9 +3612,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" @@ -4097,9 +4064,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() @@ -4130,9 +4095,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() @@ -4203,9 +4166,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() @@ -4437,9 +4398,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, ) @@ -4455,9 +4414,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() @@ -4473,9 +4430,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( @@ -4542,9 +4497,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" @@ -4773,9 +4726,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 @@ -4804,9 +4755,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( @@ -4824,9 +4773,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 @@ -4851,9 +4798,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", @@ -4914,9 +4859,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( @@ -4967,9 +4910,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 @@ -4998,9 +4939,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) @@ -5084,19 +5023,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" ) @@ -5110,9 +5043,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"] @@ -5137,22 +5068,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): @@ -5168,15 +5089,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): @@ -5187,9 +5104,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, ) @@ -5221,9 +5136,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, ) @@ -5253,9 +5166,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, ) @@ -5270,9 +5181,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) @@ -5283,9 +5192,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( @@ -5304,6 +5211,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: @@ -5330,23 +5239,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() @@ -5365,9 +5268,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() @@ -5382,9 +5283,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() @@ -5393,9 +5292,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) @@ -5403,9 +5300,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() @@ -5484,9 +5379,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, ) @@ -5552,7 +5445,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, @@ -5602,7 +5497,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, @@ -5640,7 +5537,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, @@ -5681,7 +5580,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, @@ -5757,11 +5658,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, @@ -5770,6 +5671,118 @@ 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_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 ( + 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 @@ -5793,7 +5806,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, @@ -5824,9 +5839,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 @@ -5913,14 +5926,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) @@ -5929,27 +5945,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) @@ -5965,19 +5981,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) @@ -5999,14 +6019,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) @@ -6017,22 +6040,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) @@ -6040,6 +6066,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: """ @@ -6480,9 +6532,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, @@ -6553,9 +6603,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() @@ -6582,7 +6630,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 @@ -6734,9 +6781,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, @@ -6767,9 +6812,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( @@ -6895,9 +6938,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}}, @@ -6905,10 +6946,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}, @@ -6939,9 +6977,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", @@ -6971,9 +7007,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, @@ -7679,16 +7713,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 @@ -7979,12 +8009,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, @@ -7993,14 +8026,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( @@ -8497,9 +8536,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 @@ -8528,9 +8565,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 @@ -8825,9 +8860,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.""" @@ -8973,9 +9006,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"] @@ -8992,9 +9023,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()) @@ -9048,9 +9077,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.""" @@ -9096,9 +9123,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()) @@ -9137,9 +9162,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()) @@ -9156,9 +9179,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()) @@ -9425,9 +9446,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(): @@ -9435,9 +9454,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": []}),