diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index 70ea21320ee..fe91d6d7a28 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -11,7 +11,7 @@ import asyncio import json import threading -from collections.abc import AsyncGenerator, AsyncIterable, Awaitable, Sequence +from collections.abc import AsyncGenerator, AsyncIterable, AsyncIterator, Awaitable, Sequence from contextlib import asynccontextmanager from datetime import datetime from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypedDict, cast @@ -39,6 +39,11 @@ from litellm.integrations.custom_guardrail import ( log_guardrail_information, ) from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails.anthropic_sse import ( + anthropic_sse_chunks_from_response, + assemble_anthropic_sse_stream, + model_response_text, +) from litellm.types.guardrails import ( GuardrailEventHooks, LitellmParams, @@ -1327,30 +1332,44 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): ) return response - async def _stream_apply_output_masking( - self, - response: AsyncIterable[object], - request_data: dict, - ) -> AsyncGenerator[ModelResponseStream | bytes, None]: - """Apply Presidio masking to streaming output (apply_to_output=True path).""" + async def _mask_buffered_model_response_stream( + self, all_chunks: Sequence[ModelResponseStream], request_data: dict + ) -> tuple[object, ...]: from litellm.llms.base_llm.base_model_iterator import ( convert_model_response_to_streaming, ) from litellm.main import stream_chunk_builder from litellm.types.utils import ModelResponse + assembled: Final = stream_chunk_builder(chunks=list(all_chunks), messages=request_data.get("messages")) + if not isinstance(assembled, ModelResponse): + return tuple(all_chunks) + await self._process_response_for_pii(response=assembled, request_data=request_data, mode="mask") + return (convert_model_response_to_streaming(assembled),) + + async def _stream_apply_output_masking( + self, + response: AsyncIterable[object], + request_data: dict, + ) -> AsyncGenerator[object, None]: + """Apply Presidio masking to streaming output (apply_to_output=True path).""" all_chunks: list[ModelResponseStream] = [] passthrough_due_to_unknown_stream_shape = False try: - async for chunk in response: + stream: Final = response.__aiter__() + async for chunk in stream: if isinstance(chunk, ModelResponseStream): if passthrough_due_to_unknown_stream_shape: yield chunk else: all_chunks.append(chunk) elif isinstance(chunk, bytes): - yield chunk - continue + if passthrough_due_to_unknown_stream_shape or all_chunks: + yield chunk + continue + for masked_chunk in await self._mask_anthropic_sse_stream(chunk, stream, request_data): + yield masked_chunk + return else: if all_chunks: # Flush buffered chunks and switch to transparent passthrough for this stream shape. @@ -1375,33 +1394,39 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): if not all_chunks: verbose_proxy_logger.warning( "Presidio apply_to_output: streaming response contained no " - "ModelResponseStream chunks (e.g. raw SSE bytes or an empty " - "upstream stream). Output PII masking was skipped for this " - "response." + "ModelResponseStream chunks (an empty upstream stream). " + "Output PII masking was skipped for this response." ) return - assembled_model_response = stream_chunk_builder(chunks=all_chunks, messages=request_data.get("messages")) - - if not isinstance(assembled_model_response, ModelResponse): - for chunk in all_chunks: - yield chunk - return - - await self._process_response_for_pii( - response=assembled_model_response, - request_data=request_data, - mode="mask", - ) - - mock_response_stream: Final = convert_model_response_to_streaming(assembled_model_response) - yield mock_response_stream + for masked_chunk in await self._mask_buffered_model_response_stream(all_chunks, request_data): + yield masked_chunk except Exception as e: + if not all_chunks or isinstance(e, BlockedPiiEntityError): + raise verbose_proxy_logger.error("Error masking streaming PII output: %s", e) for chunk in all_chunks: yield chunk + async def _mask_anthropic_sse_stream( + self, first_chunk: bytes, rest: AsyncIterator[object], request_data: dict + ) -> tuple[object, ...]: + rest_chunks: Final = [chunk async for chunk in rest] # mutable-ok: tuple() cannot consume an async iterator + chunks: Final = (first_chunk, *rest_chunks) + assembled: Final = assemble_anthropic_sse_stream(chunks, restore_identity=True) + if assembled is None: + verbose_proxy_logger.warning( + "Presidio apply_to_output: raw SSE stream could not be assembled into a response. " + "Output PII masking was skipped for this response." + ) + return chunks + original_text: Final = model_response_text(assembled) + await self._process_response_for_pii(response=assembled, request_data=request_data, mode="mask") + if model_response_text(assembled) == original_text: + return chunks + return anthropic_sse_chunks_from_response(assembled) + @staticmethod def _unmask_sse_bytes_chunk(chunk: bytes, pii_tokens: dict[str, str]) -> bytes: try: @@ -1460,7 +1485,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): self, response: AsyncIterable[object], request_data: dict, - ) -> AsyncGenerator[ModelResponseStream | bytes, None]: + ) -> AsyncGenerator[object, None]: """Apply PII unmasking to streaming output (output_parse_pii=True path).""" from litellm.llms.base_llm.base_model_iterator import ( convert_model_response_to_streaming, @@ -1536,7 +1561,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): user_api_key_dict: UserAPIKeyAuth, response: AsyncIterable[object], request_data: dict, - ) -> AsyncGenerator[ModelResponseStream | bytes, None]: + ) -> AsyncGenerator[object, None]: """ Process streaming response chunks to unmask PII tokens when needed. diff --git a/litellm/proxy/guardrails/guardrail_initializers.py b/litellm/proxy/guardrails/guardrail_initializers.py index 356eb7c96c6..7bafad26569 100644 --- a/litellm/proxy/guardrails/guardrail_initializers.py +++ b/litellm/proxy/guardrails/guardrail_initializers.py @@ -165,6 +165,7 @@ def initialize_presidio(litellm_params: LitellmParams, guardrail: Guardrail) -> apply_to_output=True, event_hook=GuardrailEventHooks.post_call.value, output_parse_pii=False, + mask_response_content=True, ) if run_output else None diff --git a/tests/e2e/coverage_registry/guardrail.yaml b/tests/e2e/coverage_registry/guardrail.yaml index 86eb44f6cb1..4fe9f949fae 100644 --- a/tests/e2e/coverage_registry/guardrail.yaml +++ b/tests/e2e/coverage_registry/guardrail.yaml @@ -2,6 +2,7 @@ # Rolls up into the "Logging & Guardrails" dashboard module together with logging.* - {id: guardrail.presidio.pre_call.masks, module: guardrail, tier: P0, hook_point: pre_call, assertions: [masks], exercised_on: [chat_completions, messages], source: "guardrail_hooks/presidio.py", rationale: "PII masking pre-call; data-leak blast radius"} - {id: guardrail.presidio.post_call.masks, module: guardrail, tier: P0, hook_point: post_call, assertions: [masks], exercised_on: [chat_completions, messages], source: "guardrail_hooks/presidio.py", rationale: "Mask PII in model output"} +- {id: guardrail.presidio.post_call.masks_generated_output, module: guardrail, tier: P0, hook_point: post_call, assertions: [masks], exercised_on: [chat_completions, chat_completions_stream, anthropic_messages_stream], source: "guardrail_hooks/presidio.py", rationale: "Mask model-generated credit-card output with the UI default scope"} - {id: guardrail.presidio.logging_only.masks, module: guardrail, tier: P0, hook_point: logging_only, assertions: [masks], exercised_on: [chat_completions, messages], source: "guardrail_hooks/presidio.py", rationale: "Redact in logs without blocking"} - {id: guardrail.presidio.pre_call.logs_masked_entities, module: guardrail, tier: P0, hook_point: pre_call, assertions: [logs_masked_entities], exercised_on: [chat_completions], source: "guardrail_hooks/presidio.py", rationale: "A masking run must record itself on the spend log: the dashboard's guardrail panel renders the masked-entity counts and per-entity scores straight off metadata.guardrail_information, so a run that masks but records nothing leaves an operator unable to audit it"} - {id: guardrail.bedrock.pre_call.blocks, module: guardrail, tier: P0, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/bedrock_guardrails.py", rationale: "AWS content guardrail blocks harmful input"} diff --git a/tests/e2e/guardrails/guardrails_client.py b/tests/e2e/guardrails/guardrails_client.py index 3ceac737399..97ecac0290f 100644 --- a/tests/e2e/guardrails/guardrails_client.py +++ b/tests/e2e/guardrails/guardrails_client.py @@ -46,7 +46,7 @@ class BlockedWordBody(BaseModel): class GuardrailParamsBase(BaseModel): - mode: GuardrailMode + mode: GuardrailMode | list[GuardrailMode] default_on: bool @@ -381,6 +381,26 @@ class GuardrailsClient: ), ) + def messages_stream_raw( + self, + key: str, + model: str, + text: str, + *, + guardrails: list[str] | None = None, + max_tokens: int = 64, + ) -> StreamingResponse: + return self.proxy.messages_stream( + key, + AnthropicMessagesBody( + model=model, + messages=[ChatMessage(role="user", content=text)], + max_tokens=max_tokens, + stream=True, + guardrails=guardrails, + ), + ) + def responses( self, key: str, diff --git a/tests/e2e/guardrails/test_presidio_masking_e2e.py b/tests/e2e/guardrails/test_presidio_masking_e2e.py index 49d698938ce..d47d64be9e3 100644 --- a/tests/e2e/guardrails/test_presidio_masking_e2e.py +++ b/tests/e2e/guardrails/test_presidio_masking_e2e.py @@ -30,6 +30,7 @@ this suite deliberately requires the detected-entity details to remain visible. from __future__ import annotations import os +import re import time from collections.abc import Callable from typing import Final, Literal @@ -65,9 +66,13 @@ GUARDRAIL_PROPAGATION_POLL_INTERVAL_SECONDS = 5.0 # angle brackets; the logged payload keeps the placeholder verbatim. MASKED_EMAIL_TOKEN = "EMAIL_ADDRESS" MASKED_PHONE_TOKEN = "PHONE_NUMBER" +MASKED_CREDIT_CARD_TOKEN = "CREDIT_CARD" # Fictional NANP 555 number; a standard format Presidio's phone recognizer detects. FAKE_PHONE = "+1 415-555-0134" +FAKE_VISA_TEST_CARD = "4111 1111 1111 1111" + +_CARD_DIGIT_RUN: Final = re.compile(r"(?:\d[ -]?){13,19}") def _presidio_bases() -> tuple[str, str]: @@ -86,8 +91,8 @@ def _register_presidio( resources: ResourceManager, *, name: str, - mode: GuardrailMode = "pre_call", - filter_scope: Literal["input", "output", "both"] = "input", + mode: GuardrailMode | list[GuardrailMode] = "pre_call", + filter_scope: Literal["input", "output", "both"] | None = "input", entities: dict[PiiEntity, PiiAction] | None = None, ) -> None: analyzer, anonymizer = _presidio_bases() @@ -123,6 +128,74 @@ def _first_content(response: ChatResponse) -> str: return (message.content if message else None) or "" +class _StreamDelta(BaseModel): + content: str | None = None + + +class _StreamChoice(BaseModel): + delta: _StreamDelta + + +class _StreamChunk(BaseModel): + choices: tuple[_StreamChoice, ...] = () + + +class _AnthropicStreamDelta(BaseModel): + type: str | None = None + text: str | None = None + + +class _AnthropicStreamEvent(BaseModel): + type: str + delta: _AnthropicStreamDelta | None = None + + +def _credit_card_prompt(marker: str) -> str: + return ( + f"{marker} Reply with only the well known Visa sandbox test card number that starts with 4111, " + "the 16 digits grouped in fours separated by spaces, and nothing else." + ) + + +def _passes_luhn(digits: str) -> bool: + checksum = sum( + digit if position % 2 == 0 else (digit * 2 - 9 if digit * 2 > 9 else digit * 2) + for position, digit in enumerate(int(char) for char in reversed(digits)) + ) + return checksum % 10 == 0 + + +def _contains_card_number(text: str) -> bool: + """Presidio's CREDIT_CARD recognizer only reports Luhn-valid digit runs, so a + Luhn-invalid number the model hallucinates is not something masking can catch.""" + return any( + 13 <= len(digits) <= 19 and _passes_luhn(digits) + for digits in (re.sub(r"[ -]", "", match.group()) for match in _CARD_DIGIT_RUN.finditer(text)) + ) + + +def _stream_content(result: StreamingResponse) -> str: + return "".join( + choice.delta.content + for event in result.stream_events + if event != "[DONE]" + for choice in _StreamChunk.model_validate_json(event).choices[:1] + if choice.delta.content + ) + + +def _anthropic_stream_content(result: StreamingResponse) -> str: + return "".join( + event.delta.text + for payload in result.stream_events + for event in [_AnthropicStreamEvent.model_validate_json(payload)] + if event.type == "content_block_delta" + and event.delta is not None + and event.delta.type == "text_delta" + and event.delta.text + ) + + def _messages_text(response: AnthropicMessagesResponse) -> str: """The text of a /v1/messages answer, whichever shape the proxy produced (Anthropic-native content blocks or OpenAI-normalized choices).""" @@ -290,6 +363,124 @@ class TestPresidioPostCallMasking: time.sleep(GUARDRAIL_PROPAGATION_POLL_INTERVAL_SECONDS) +def _assert_eventually_masks_generated_card(fetch: Callable[[], str | None]) -> None: + deadline = time.monotonic() + GUARDRAIL_PROPAGATION_DEADLINE_SECONDS + last: str = "" + while True: + content = fetch() + if content is not None: + last = content + if _contains_card_number(content): + pytest.fail( + "the post_call output masking let a card number through: " + f"{content[:300]!r}" + ) + if MASKED_CREDIT_CARD_TOKEN in content: + return + if time.monotonic() >= deadline: + pytest.fail( + "presidio post_call output masking never masked the generated card within " + f"{GUARDRAIL_PROPAGATION_DEADLINE_SECONDS}s; last observation: {last[:300]!r}" + ) + time.sleep(GUARDRAIL_PROPAGATION_POLL_INTERVAL_SECONDS) + + +class TestPresidioCreditCardOutputMasking: + """Proves the UI-default Presidio scope masks model-generated card output.""" + + @pytest.mark.covers( + "guardrail.presidio.post_call.masks_generated_output", + exercised_on=["chat_completions"], + ) + def test_ui_default_scope_masks_a_card_number_the_model_generates_on_chat_completions( + self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str + ) -> None: + name = f"e2e-presidio-card-chat-{unique_marker()}" + _register_presidio( + client, + resources, + name=name, + mode=["pre_call", "post_call"], + filter_scope=None, + entities={"CREDIT_CARD": "MASK"}, + ) + prompt: Final = _credit_card_prompt(unique_marker()) + + def fetch() -> str | None: + result: Final = client.chat(scoped_key, MODEL, prompt, guardrails=[name], max_tokens=512) + match result: + case Success(data=data): + return _first_content(data) + case _: + return None + + _assert_eventually_masks_generated_card(fetch) + + @pytest.mark.covers( + "guardrail.presidio.post_call.masks_generated_output", + exercised_on=["chat_completions_stream"], + ) + def test_ui_default_scope_masks_a_card_number_the_model_generates_on_streaming_chat_completions( + self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str + ) -> None: + name = f"e2e-presidio-card-stream-{unique_marker()}" + _register_presidio( + client, + resources, + name=name, + mode=["pre_call", "post_call"], + filter_scope=None, + entities={"CREDIT_CARD": "MASK"}, + ) + prompt: Final = _credit_card_prompt(unique_marker()) + + def fetch() -> str | None: + result: Final = client.chat_stream_raw( + scoped_key, + MODEL, + prompt, + guardrails=[name], + max_tokens=512, + ) + if not result.ok or result.stream_error: + return None + return _stream_content(result) + + _assert_eventually_masks_generated_card(fetch) + + @pytest.mark.covers( + "guardrail.presidio.post_call.masks_generated_output", + exercised_on=["anthropic_messages_stream"], + ) + def test_ui_default_scope_masks_a_card_number_the_model_generates_on_streaming_anthropic_messages( + self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str + ) -> None: + name = f"e2e-presidio-card-messages-stream-{unique_marker()}" + _register_presidio( + client, + resources, + name=name, + mode=["pre_call", "post_call"], + filter_scope=None, + entities={"CREDIT_CARD": "MASK"}, + ) + prompt: Final = _credit_card_prompt(unique_marker()) + + def fetch() -> str | None: + result: Final = client.messages_stream_raw( + scoped_key, + MODEL, + prompt, + guardrails=[name], + max_tokens=512, + ) + if not result.ok or result.stream_error: + return None + return _anthropic_stream_content(result) + + _assert_eventually_masks_generated_card(fetch) + + _LOGGED_ENTITIES: dict[PiiEntity, PiiAction] = {"EMAIL_ADDRESS": "MASK", "PHONE_NUMBER": "MASK"} _ENTITY_LIST_ADAPTER: Final = TypeAdapter(list[GuardrailEntityMatch]) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py index 33614d2eeca..89e72debbf5 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py @@ -4,6 +4,7 @@ Tests PII detection and masking for different message formats """ import asyncio +import json from contextlib import asynccontextmanager from unittest.mock import MagicMock, patch @@ -18,7 +19,7 @@ from litellm.proxy.guardrails.guardrail_hooks.presidio import ( ) from litellm.exceptions import GuardrailRaisedException from litellm.types.guardrails import LitellmParams, PiiAction, PiiEntityType -from litellm.types.utils import Choices, Message, ModelResponse +from litellm.types.utils import Choices, Delta, Message, ModelResponse, StreamingChoices from litellm.exceptions import BlockedPiiEntityError @@ -2331,47 +2332,320 @@ async def test_apply_guardrail_masks_on_request(): assert "John Smith" not in result["texts"][0] +def _anthropic_sse(event_type: str, payload: dict) -> bytes: + return f"event: {event_type}\ndata: {json.dumps(payload)}\n\n".encode() + + +def _anthropic_text_deltas(chunks: list[bytes]) -> list[tuple[int, str]]: + deltas = [] + for line in b"".join(chunks).decode().split("\n"): + if not line.startswith("data: "): + continue + event = json.loads(line[6:]) + if event.get("type") == "content_block_delta" and event["delta"].get("type") == "text_delta": + deltas.append((event["index"], event["delta"]["text"])) + return deltas + + +def _chat_delta_chunk(text: str, finish_reason: str | None = None) -> ModelResponseStream: + return ModelResponseStream( + id="chatcmpl-out-mask", + choices=[StreamingChoices(index=0, delta=Delta(content=text, role="assistant"), finish_reason=finish_reason)], + created=1, + model="gpt-4", + object="chat.completion.chunk", + ) + + @pytest.mark.asyncio -async def test_apply_to_output_streaming_bytes_only_logs_warning(): +async def test_apply_to_output_streaming_chat_chunks_are_masked_as_one_response(): """ - Regression test: when apply_to_output=True and the stream contains only - bytes chunks (Anthropic native SSE), output masking is skipped. - A warning must be logged so operators are aware. + Structured chat completion chunks are buffered, assembled and masked as a + whole, so a card number split across deltas cannot reach the caller. """ guardrail = _OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, + mock_redacted_text={"text": "my card is "}, + ) + + async def mock_stream(): + yield _chat_delta_chunk("my card is 4111") + yield _chat_delta_chunk(" 1111 1111 1111") + yield _chat_delta_chunk("", finish_reason="stop") + + collected = [] + async for chunk in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), + response=mock_stream(), + request_data={"messages": [{"role": "user", "content": "what is my card"}]}, + ): + collected.append(chunk) + + assert all(isinstance(chunk, ModelResponseStream) for chunk in collected) + joined = "".join(chunk.choices[0].delta.content or "" for chunk in collected) + assert joined == "my card is " + assert collected[-1].choices[0].finish_reason == "stop" + + +@pytest.mark.asyncio +async def test_apply_to_output_streaming_bytes_after_chat_chunks_are_passed_through_in_order(): + """ + Once structured chunks have been buffered, a trailing bytes frame belongs to + the same stream and must be forwarded rather than treated as a new SSE stream. + """ + guardrail = _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + apply_to_output=True, + mock_redacted_text={"text": "hello"}, + ) + trailer = b"data: [DONE]\n\n" + + async def mock_stream(): + yield _chat_delta_chunk("hello", finish_reason="stop") + yield trailer + + collected = [] + async for chunk in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), + response=mock_stream(), + request_data={}, + ): + collected.append(chunk) + + assert collected[0] == trailer + assert len(collected) == 2 + assert isinstance(collected[1], ModelResponseStream) + assert collected[1].choices[0].delta.content == "hello" + + +@pytest.mark.asyncio +async def test_apply_to_output_streaming_anthropic_sse_bytes_masks_text_split_across_deltas(): + """ + Anthropic native /v1/messages streams reach the post_call hook as raw SSE + bytes. Output masking must run over the whole content block so a card + number split across text_delta events cannot reach the caller. + """ + guardrail = _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + apply_to_output=True, + mock_redacted_text={"text": ""}, ) byte_chunks = [ - b'data: {"type":"content_block_delta","delta":{"text":"Hello"}}\n\n', - b'data: {"type":"content_block_delta","delta":{"text":" world"}}\n\n', + _anthropic_sse( + "message_start", + {"type": "message_start", "message": {"id": "msg_1", "model": "claude", "content": [], "usage": {}}}, + ), + _anthropic_sse( + "content_block_start", + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + ), + _anthropic_sse( + "content_block_delta", + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "4111"}}, + ), + _anthropic_sse( + "content_block_delta", + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": " 1111 1111 1111"}}, + ), + _anthropic_sse("content_block_stop", {"type": "content_block_stop", "index": 0}), + _anthropic_sse("message_delta", {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {}}), + _anthropic_sse("message_stop", {"type": "message_stop"}), ] async def mock_stream(): for b in byte_chunks: yield b - mock_user_api_key = UserAPIKeyAuth(api_key="test-key") + collected = [] + async for chunk in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), + response=mock_stream(), + request_data={}, + ): + collected.append(chunk) + + assert all(isinstance(chunk, bytes) for chunk in collected) + joined = b"".join(collected).decode() + assert "4111" not in joined + assert "".join(text for _, text in _anthropic_text_deltas(collected)) == "" + assert joined.count("event: message_start") == 1 + assert joined.count("event: message_stop") == 1 + + +@pytest.mark.asyncio +async def test_apply_to_output_streaming_anthropic_sse_bytes_without_pii_are_forwarded_unchanged(): + guardrail = _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + apply_to_output=True, + mock_redacted_text={"text": "Hello world"}, + ) + + byte_chunks = [ + _anthropic_sse( + "message_start", + {"type": "message_start", "message": {"id": "msg_1", "model": "claude", "content": [], "usage": {}}}, + ), + _anthropic_sse( + "content_block_start", + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + ), + _anthropic_sse( + "content_block_delta", + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "Hello"}}, + ), + _anthropic_sse( + "content_block_delta", + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": " world"}}, + ), + _anthropic_sse("content_block_stop", {"type": "content_block_stop", "index": 0}), + _anthropic_sse("message_delta", {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {}}), + _anthropic_sse("message_stop", {"type": "message_stop"}), + ] + + async def mock_stream(): + for b in byte_chunks: + yield b collected = [] - with patch("litellm.proxy.guardrails.guardrail_hooks.presidio.verbose_proxy_logger") as mock_logger: + async for chunk in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), + response=mock_stream(), + request_data={}, + ): + collected.append(chunk) + + assert collected == byte_chunks + + +@pytest.mark.asyncio +async def test_apply_to_output_streaming_anthropic_sse_bytes_fail_closed_when_presidio_is_unreachable(): + """ + The raw SSE stream is fully drained before masking, so a Presidio outage + must surface as an error to the caller: replaying the unscanned frames + would hand over whatever PII the model generated. + """ + guardrail = _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + apply_to_output=True, + presidio_analyzer_api_base="http://127.0.0.1:9", + presidio_anonymizer_api_base="http://127.0.0.1:9", + ) + + byte_chunks = [ + _anthropic_sse( + "message_start", + {"type": "message_start", "message": {"id": "msg_1", "model": "claude", "content": [], "usage": {}}}, + ), + _anthropic_sse( + "content_block_start", + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + ), + _anthropic_sse( + "content_block_delta", + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "Hello world"}}, + ), + _anthropic_sse("content_block_stop", {"type": "content_block_stop", "index": 0}), + _anthropic_sse("message_stop", {"type": "message_stop"}), + ] + + async def mock_stream(): + for b in byte_chunks: + yield b + + collected = [] + + async def collect_masked_stream(): async for chunk in guardrail.async_post_call_streaming_iterator_hook( - user_api_key_dict=mock_user_api_key, + user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), response=mock_stream(), request_data={}, ): collected.append(chunk) - # All bytes should be yielded through - assert len(collected) == len(byte_chunks) - for original, received in zip(byte_chunks, collected): - assert original == received + with pytest.raises(Exception, match="Presidio PII analysis failed"): + await collect_masked_stream() - # Warning must be logged about skipped masking - mock_logger.warning.assert_called_once() - warning_msg = mock_logger.warning.call_args[0][0] - assert "Output PII masking was skipped" in warning_msg + assert collected == [] + + +@pytest.mark.asyncio +async def test_apply_to_output_streaming_anthropic_sse_bytes_block_action_raises_instead_of_replaying(): + """ + A BLOCK on generated PII must refuse the streaming /v1/messages response the + same way it refuses the non streaming one, not replay the raw frames. + """ + guardrail = _OPTIONAL_PresidioPIIMasking( + apply_to_output=True, + mock_testing=False, + presidio_analyzer_api_base="http://test-analyzer/", + presidio_anonymizer_api_base="http://test-anonymizer/", + pii_entities_config={PiiEntityType.CREDIT_CARD: PiiAction.BLOCK}, + ) + + byte_chunks = [ + _anthropic_sse( + "message_start", + {"type": "message_start", "message": {"id": "msg_1", "model": "claude", "content": [], "usage": {}}}, + ), + _anthropic_sse( + "content_block_start", + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + ), + _anthropic_sse( + "content_block_delta", + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "4111 1111 1111 1111"}}, + ), + _anthropic_sse("content_block_stop", {"type": "content_block_stop", "index": 0}), + _anthropic_sse("message_stop", {"type": "message_stop"}), + ] + + async def mock_stream(): + for b in byte_chunks: + yield b + + analyzer_hit = [{"entity_type": "CREDIT_CARD", "score": 0.99, "start": 0, "end": 19}] + collected = [] + + async def collect_masked_stream(): + async for chunk in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), + response=mock_stream(), + request_data={}, + ): + collected.append(chunk) + + with patch.object(guardrail, "_get_session_iterator", _make_mock_session_iterator(analyzer_hit)): + with pytest.raises(BlockedPiiEntityError): + await collect_masked_stream() + + assert collected == [] + + +@pytest.mark.asyncio +async def test_apply_to_output_streaming_propagates_upstream_error_when_nothing_was_buffered(): + """ + An upstream guardrail that rejects the stream before the first chunk must + surface as an error to the caller, not as an empty 200 stream. + """ + guardrail = _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + apply_to_output=True, + mock_redacted_text={"text": ""}, + ) + + async def failing_stream(): + raise RuntimeError("upstream guardrail rejected the stream") + yield b"" + + with pytest.raises(RuntimeError, match="upstream guardrail rejected the stream"): + async for _ in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), + response=failing_stream(), + request_data={}, + ): + pass @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/test_proxy_logging_hook_detection.py b/tests/test_litellm/proxy/test_proxy_logging_hook_detection.py index dd330d32ce6..f45715953f9 100644 --- a/tests/test_litellm/proxy/test_proxy_logging_hook_detection.py +++ b/tests/test_litellm/proxy/test_proxy_logging_hook_detection.py @@ -489,6 +489,68 @@ async def test_post_call_stream_masking_guardrail_keeps_own_iterator_on_anthropi assert delivered == chunks +@pytest.mark.asyncio +async def test_post_call_stream_presidio_output_masking_masks_anthropic_messages_stream(monkeypatch): + """Regression: the presidio output-masking callback built by initialize_presidio + was rerouted onto the unified scan-only path on /v1/messages, so a card number + the analyzer flagged still streamed to the caller unmasked.""" + import json + + from litellm.caching.caching import DualCache + from litellm.proxy.guardrails.guardrail_registry import InMemoryGuardrailHandler + from litellm.types.guardrails import SupportedGuardrailIntegrations + + handler = InMemoryGuardrailHandler() + result = handler.initialize_guardrail( + guardrail={ + "guardrail_name": "presidio-card-mask", + "litellm_params": { + "guardrail": SupportedGuardrailIntegrations.PRESIDIO.value, + "mode": ["pre_call", "post_call"], + "default_on": True, + "presidio_analyzer_api_base": "https://fakelink.com/v1/presidio/analyze", + "presidio_anonymizer_api_base": "https://fakelink.com/v1/presidio/anonymize", + "pii_entities_config": {"CREDIT_CARD": "MASK"}, + "mock_redacted_text": {"text": "", "items": []}, + }, + } + ) + guardrail_id = result["guardrail_id"] + callbacks = [ + handler.guardrail_id_to_custom_guardrail[guardrail_id], + *handler.guardrail_id_to_sibling_callbacks[guardrail_id], + ] + monkeypatch.setattr(litellm, "callbacks", callbacks) + + chunks = _anthropic_stream_chunks(["4111", " 1111 1111 1111"]) + + async def fake_stream(): + for chunk in chunks: + yield chunk + + delivered = [] + async for chunk in ProxyLogging(user_api_key_cache=DualCache()).async_post_call_streaming_iterator_hook( + response=fake_stream(), + user_api_key_dict=UserAPIKeyAuth(api_key="sk-1234", request_route="/v1/messages"), + request_data={ + "model": "claude-sonnet-5", + "litellm_logging_obj": _streaming_logging_obj(), + "metadata": {}, + }, + ): + delivered.append(chunk) + + wire = b"".join(delivered).decode() + text_deltas = [ + json.loads(line[6:])["delta"]["text"] + for line in wire.split("\n") + if line.startswith("data: ") and json.loads(line[6:]).get("delta", {}).get("type") == "text_delta" + ] + assert "4111" not in wire, wire + assert "".join(text_deltas) == "", wire + assert wire.count("event: message_stop") == 1, wire + + class _AppliesGuardrail(CustomGuardrail): """Implements the unified interface only, so the proxy routes it to unified_guardrail."""