From 4509c874701f4b8243479a76b0144995120e380a Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Thu, 3 Sep 2026 11:47:57 -0500 Subject: [PATCH] feat(guardrails): add llm shield pii redaction and rehydration guardrail LLM Shield is a self-hosted PII gateway. This adds it as a guardrail so a proxy operator can redact personal data out of outbound requests and have the original values restored in the model's reply. The substitution is reversible, which is the difference from a masking guardrail. Outbound text is replaced with placeholders held in a session vault inside the operator's own LLM Shield deployment, and the reply is restored before it reaches the caller, so the end user still sees real values while the provider never received them. Streaming responses are restored incrementally. LLM Shield holds back only the trailing characters that could still turn out to be part of a placeholder, so tokens are forwarded as they arrive rather than the whole response being collected first. A placeholder split across two chunks is never emitted in fragments. The integration talks to LLM Shield over HTTP and adds no dependency. Notes for reviewers: - The guardrail sets use_native_lifecycle_hooks, since redaction and restoration need the native pre-call, post-call and streaming hooks rather than the unified path. - Per-request state lives on the request dict, never on the guardrail instance, because the proxy registers a single instance process-wide. The streaming carry-over is a local of the generator for the same reason. - Every failure blocks the request. A redaction guardrail that fails open would send the exact data it exists to protect to the provider. --- .../guardrail_hooks/llm_shield/__init__.py | 33 ++ .../guardrail_hooks/llm_shield/llm_shield.py | 351 ++++++++++++++++++ litellm/types/guardrails.py | 1 + .../guardrails/guardrail_hooks/llm_shield.py | 24 ++ .../guardrail_hooks/test_llm_shield.py | 261 +++++++++++++ 5 files changed, 670 insertions(+) create mode 100644 litellm/proxy/guardrails/guardrail_hooks/llm_shield/__init__.py create mode 100644 litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py create mode 100644 litellm/types/proxy/guardrails/guardrail_hooks/llm_shield.py create mode 100644 tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/__init__.py new file mode 100644 index 00000000000..a6cc54d5408 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/__init__.py @@ -0,0 +1,33 @@ +from typing import TYPE_CHECKING, Final + +from litellm.types.guardrails import SupportedGuardrailIntegrations + +from .llm_shield import LLMShieldGuardrail + +if TYPE_CHECKING: + from litellm.types.guardrails import Guardrail, LitellmParams + + +def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"): + import litellm + + _llm_shield_guardrail_callback: Final = LLMShieldGuardrail( + api_key=litellm_params.api_key, + api_base=litellm_params.api_base, + guardrail_name=guardrail.get("guardrail_name", ""), + event_hook=litellm_params.mode, + default_on=litellm_params.default_on, + ) + + litellm.logging_callback_manager.add_litellm_callback(_llm_shield_guardrail_callback) + return _llm_shield_guardrail_callback + + +guardrail_initializer_registry: Final = { # mutable-ok: module-level registry, built once and never mutated + SupportedGuardrailIntegrations.LLM_SHIELD.value: initialize_guardrail, +} + + +guardrail_class_registry: Final = { # mutable-ok: module-level registry, built once and never mutated + SupportedGuardrailIntegrations.LLM_SHIELD.value: LLMShieldGuardrail, +} diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py new file mode 100644 index 00000000000..4199ad5ca65 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py @@ -0,0 +1,351 @@ +# +-------------------------------------------------------------+ +# +# Use LLM Shield for reversible PII redaction +# https://github.com/ninadphalak/LLM-Shield-Proxy +# +# +-------------------------------------------------------------+ + +import os +import uuid +from collections.abc import AsyncGenerator +from typing import ( + TYPE_CHECKING, + Any, # noqa: TID251 # **kwargs forwards verbatim to CustomGuardrail.__init__ + ClassVar, + Final, + Literal, + Optional, +) + +import httpx + +from litellm._logging import verbose_proxy_logger +from litellm.exceptions import GuardrailRaisedException +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + get_session_id_from_request_data, + log_guardrail_information, +) +from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, + httpxSpecialProvider, +) +from litellm.proxy._types import UserAPIKeyAuth +from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.utils import GenericGuardrailAPIInputs + +if TYPE_CHECKING: + from litellm.caching.caching import DualCache + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + +GUARDRAIL_NAME: Final = "llm_shield" + +_DEFAULT_API_BASE: Final = "http://localhost:8000" +_REDACT_PATH: Final = "/v1/guard/redact" +_REHYDRATE_PATH: Final = "/v1/guard/rehydrate" +_REHYDRATE_STREAM_PATH: Final = "/v1/guard/rehydrate/stream" + +# The session id ties a redact call to the rehydrate calls that undo it. It is +# stored on the request dict rather than on the guardrail instance: the proxy +# registers one instance process-wide, so instance attributes would be shared +# across concurrent requests. +_SESSION_METADATA_KEY: Final = "llm_shield_session_id" + +_DEFAULT_TIMEOUT_SECONDS: Final = 10.0 + + +class LLMShieldGuardrail(CustomGuardrail): + """Redacts PII before it leaves the proxy and restores it in the response. + + Unlike a masking guardrail, the substitution is reversible. Outbound text is + replaced with placeholders held in a session vault inside the user's own LLM + Shield deployment; the model's reply is then restored so the end user sees the + original values while the provider never received them. + + Streaming is restored incrementally rather than by buffering the response. LLM + Shield holds back only the trailing characters that could still turn out to be + part of a placeholder, so tokens are forwarded as they arrive and a placeholder + split across two chunks is never emitted in fragments. + """ + + # Our redaction and restoration run in the native lifecycle hooks below. Without + # this the proxy would route every event through the unified apply_guardrail path + # and the streaming hook would never fire. + use_native_lifecycle_hooks: ClassVar[bool] = True + + def __init__( + self, + guardrail_name: str = GUARDRAIL_NAME, + api_base: str | None = None, + api_key: str | None = None, + **kwargs: Any, + ) -> None: + self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback) + self.api_base: Final = (api_base or os.environ.get("LLM_SHIELD_API_BASE") or _DEFAULT_API_BASE).rstrip("/") + self.api_key: Final = api_key or os.environ.get("LLM_SHIELD_API_KEY") + super().__init__(guardrail_name=guardrail_name, **kwargs) + + @classmethod + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: + return [GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call] + + # --- transport --------------------------------------------------------------- + + def _headers(self, session_id: str) -> dict: + headers = {"Content-Type": "application/json", "X-Session-ID": session_id} + if self.api_key: + headers["Authorization"] = f"Bearer {self.api_key}" + return headers + + async def _call_shield(self, path: str, session_id: str, payload: dict) -> dict: + """Posts to LLM Shield, failing closed on any transport or status error. + + A redaction guardrail that fails open sends the very data it exists to + protect to a third-party provider, so an unreachable or erroring shield + blocks the request instead of passing it through. + """ + try: + response = await self.async_handler.post( + f"{self.api_base}{path}", + headers=self._headers(session_id), + json=payload, + timeout=_DEFAULT_TIMEOUT_SECONDS, + ) + response.raise_for_status() + return response.json() + except httpx.HTTPStatusError as exc: + verbose_proxy_logger.exception("LLM Shield returned %s for %s", exc.response.status_code, path) + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=f"LLM Shield returned {exc.response.status_code}; blocking the request.", + ) from exc + except Exception as exc: + verbose_proxy_logger.exception("LLM Shield call to %s failed", path) + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message="LLM Shield is unreachable; blocking the request.", + ) from exc + + async def _redact(self, texts: list, session_id: str) -> list: + body = await self._call_shield(_REDACT_PATH, session_id, {"texts": texts}) + return self._same_length_or_raise(body.get("texts"), texts, "redact") + + async def _rehydrate(self, texts: list, session_id: str) -> list: + body = await self._call_shield(_REHYDRATE_PATH, session_id, {"texts": texts}) + return self._same_length_or_raise(body.get("texts"), texts, "rehydrate") + + def _same_length_or_raise(self, returned: Any, sent: list, operation: str) -> list: + """Guards the positional mapping the callers rely on to write results back.""" + if not isinstance(returned, list) or len(returned) != len(sent): + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=f"LLM Shield {operation} returned an unexpected payload; blocking the request.", + ) + return returned + + # --- session ------------------------------------------------------------------ + + def _session_id(self, data: dict) -> str: + """Returns a session id stable across this request's hooks.""" + metadata = data.setdefault("metadata", {}) + if not isinstance(metadata, dict): + return f"litellm-{uuid.uuid4().hex}" + existing = metadata.get(_SESSION_METADATA_KEY) + if isinstance(existing, str) and existing: + return existing + session_id = get_session_id_from_request_data(data) or f"litellm-{uuid.uuid4().hex}" + metadata[_SESSION_METADATA_KEY] = session_id + return session_id + + # --- message traversal -------------------------------------------------------- + + @staticmethod + def _locate_texts(messages: list) -> list: + """Finds every text span in a message list. + + Returns ``(message_index, part_index_or_None, text)``. The list form is the + multimodal shape, where only ``text`` parts carry redactable content. + """ + located = [] + for message_index, message in enumerate(messages): + if not isinstance(message, dict): + continue + content = message.get("content") + if isinstance(content, str) and content: + located.append((message_index, None, content)) + elif isinstance(content, list): + for part_index, part in enumerate(content): + if not isinstance(part, dict) or part.get("type") != "text": + continue + text = part.get("text") + if isinstance(text, str) and text: + located.append((message_index, part_index, text)) + return located + + @staticmethod + def _write_back(messages: list, located: list, replacements: list) -> None: + for (message_index, part_index, _), replacement in zip(located, replacements): + if part_index is None: + messages[message_index]["content"] = replacement + else: + messages[message_index]["content"][part_index]["text"] = replacement + + # --- hooks -------------------------------------------------------------------- + + @log_guardrail_information + async def async_pre_call_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + cache: "DualCache", + data: dict, + call_type: str, + ) -> dict | None: + """Replaces PII in the outbound messages with vault placeholders.""" + if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_call) is not True: + return data + + messages = data.get("messages") + if not isinstance(messages, list): + return data + + located = self._locate_texts(messages) + if not located: + return data + + redacted = await self._redact([text for _, _, text in located], self._session_id(data)) + self._write_back(messages, located, redacted) + return data + + @log_guardrail_information + async def async_post_call_success_hook( + self, + data: dict, + user_api_key_dict: UserAPIKeyAuth, + response: Any, + ) -> Any: + """Restores the original values in a non-streaming response.""" + if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.post_call) is not True: + return response + + choices = getattr(response, "choices", None) + if not choices: + return response + + pending = [] + for choice in choices: + message = getattr(choice, "message", None) + content = getattr(message, "content", None) + if isinstance(content, str) and content: + pending.append((message, content)) + + if not pending: + return response + + restored = await self._rehydrate([text for _, text in pending], self._session_id(data)) + for (message, _), replacement in zip(pending, restored): + message.content = replacement + return response + + async def async_post_call_streaming_iterator_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + response: Any, + request_data: dict, + ) -> AsyncGenerator[Any, None]: + """Restores original values incrementally, without buffering the stream. + + The carry-over window is a local of this generator, so it is scoped to one + stream and cannot leak between concurrent requests. LLM Shield returns the + text that is safe to emit now plus the trailing characters it is still + holding, which are sent back with the next delta. + """ + if self.should_run_guardrail(data=request_data, event_type=GuardrailEventHooks.post_call) is not True: + async for chunk in response: + yield chunk + return + + session_id = self._session_id(request_data) + carry = "" + last_chunk = None + + async for chunk in response: + last_chunk = chunk + delta = self._stream_delta(chunk) + text = getattr(delta, "content", None) if delta is not None else None + is_final = self._is_final_chunk(chunk) + + if not isinstance(text, str) or not text: + # Nothing to restore in this chunk, but a final chunk still has to + # flush whatever the window is holding. + if is_final and carry: + body = await self._stream_step("", carry, True, session_id) + carry = body["carry"] + if body["text"] and delta is not None: + delta.content = body["text"] + yield chunk + continue + + body = await self._stream_step(text, carry, is_final, session_id) + carry = body["carry"] + delta.content = body["text"] + yield chunk + + # A stream that ended without a finish_reason can still leave text held back. + if carry and last_chunk is not None: + body = await self._stream_step("", carry, True, session_id) + if body["text"]: + trailing = last_chunk.model_copy(deep=True) + trailing_delta = self._stream_delta(trailing) + if trailing_delta is not None: + trailing_delta.content = body["text"] + yield trailing + + async def _stream_step(self, text: str, carry: str, final: bool, session_id: str) -> dict: + body = await self._call_shield( + _REHYDRATE_STREAM_PATH, + session_id, + {"text": text, "carry": carry, "final": final}, + ) + emitted = body.get("text") + remaining = body.get("carry") + if not isinstance(emitted, str) or not isinstance(remaining, str): + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message="LLM Shield stream rehydration returned an unexpected payload.", + ) + return {"text": emitted, "carry": remaining} + + @staticmethod + def _stream_delta(chunk: Any) -> Any: + choices = getattr(chunk, "choices", None) + if not choices: + return None + return getattr(choices[0], "delta", None) + + @staticmethod + def _is_final_chunk(chunk: Any) -> bool: + choices = getattr(chunk, "choices", None) + if not choices: + return False + return bool(getattr(choices[0], "finish_reason", None)) + + # --- unified API (powers the UI "Test guardrail" button) ----------------------- + + @log_guardrail_information + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional["LiteLLMLoggingObj"] = None, + ) -> GenericGuardrailAPIInputs: + texts = inputs.get("texts") + if not texts: + return inputs + + session_id = self._session_id(request_data) + if input_type == "request": + inputs["texts"] = await self._redact(list(texts), session_id) + else: + inputs["texts"] = await self._rehydrate(list(texts), session_id) + return inputs diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index c17103da890..42ab6034069 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -137,6 +137,7 @@ class SupportedGuardrailIntegrations(Enum): COMPRESR = "compresr" STRAIKER = "straiker" ALICE = "alice" + LLM_SHIELD = "llm_shield" class Role(Enum): diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/llm_shield.py b/litellm/types/proxy/guardrails/guardrail_hooks/llm_shield.py new file mode 100644 index 00000000000..8d7afd907b0 --- /dev/null +++ b/litellm/types/proxy/guardrails/guardrail_hooks/llm_shield.py @@ -0,0 +1,24 @@ +from pydantic import Field + +from .base import GuardrailConfigModel + + +class LLMShieldGuardrailConfigModel(GuardrailConfigModel): + api_key: str | None = Field( + default=None, + description=( + "The virtual key for the LLM Shield instance. If not provided, the " + "`LLM_SHIELD_API_KEY` environment variable is checked." + ), + ) + api_base: str | None = Field( + default=None, + description=( + "The base URL of the LLM Shield instance. If not provided, the `LLM_SHIELD_API_BASE` " + "environment variable is checked, then `http://localhost:8000`." + ), + ) + + @staticmethod + def ui_friendly_name() -> str: + return "LLM Shield" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py new file mode 100644 index 00000000000..4b39c8bf517 --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py @@ -0,0 +1,261 @@ +from unittest.mock import AsyncMock + +import pytest +from httpx import Request, Response + +import litellm +from litellm.exceptions import GuardrailRaisedException +from litellm.proxy.guardrails.guardrail_hooks.llm_shield.llm_shield import ( + GUARDRAIL_NAME, + LLMShieldGuardrail, +) +from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2 +from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices + + +def _guardrail(**overrides: object) -> LLMShieldGuardrail: + params: dict[str, object] = { + "api_key": "test-key", + "api_base": "http://shield.test", + "guardrail_name": GUARDRAIL_NAME, + "event_hook": "pre_call", + "default_on": True, + } + params.update(overrides) + return LLMShieldGuardrail(**params) + + +def _response(payload: dict, status_code: int = 200) -> Response: + return Response( + status_code=status_code, + json=payload, + request=Request("POST", "http://shield.test/v1/guard/redact"), + ) + + +def _mock_post(guardrail: LLMShieldGuardrail, *payloads: dict) -> AsyncMock: + """Queues one shield response per expected call.""" + mock = AsyncMock(side_effect=[_response(p) for p in payloads]) + guardrail.async_handler.post = mock # type: ignore[method-assign] + return mock + + +def _chunk(content: str | None, finish_reason: str | None = None) -> ModelResponseStream: + return ModelResponseStream( + choices=[StreamingChoices(index=0, delta=Delta(content=content), finish_reason=finish_reason)] + ) + + +async def _drain(generator) -> list: + return [chunk async for chunk in generator] + + +def test_llm_shield_guardrail_config(monkeypatch: pytest.MonkeyPatch): + """Should register through init_guardrails_v2 like any other provider.""" + monkeypatch.setattr(litellm, "guardrail_name_config_map", {}) + monkeypatch.setenv("LLM_SHIELD_API_KEY", "test-key") + + init_guardrails_v2( + all_guardrails=[ + { + "guardrail_name": "llm_shield", + "litellm_params": {"guardrail": "llm_shield", "mode": "pre_call", "default_on": True}, + } + ], + config_file_path="", + ) + + registered = [cb for cb in litellm.callbacks if isinstance(cb, LLMShieldGuardrail)] + assert len(registered) == 1 + assert registered[0].guardrail_name == "llm_shield" + + +class TestLLMShieldInitialization: + def test_api_base_defaults_to_localhost(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("LLM_SHIELD_API_BASE", raising=False) + assert _guardrail(api_base=None).api_base == "http://localhost:8000" + + def test_api_base_reads_environment(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LLM_SHIELD_API_BASE", "http://shield.internal:9000") + assert _guardrail(api_base=None).api_base == "http://shield.internal:9000" + + def test_trailing_slash_is_stripped(self): + assert _guardrail(api_base="http://shield.test/").api_base == "http://shield.test" + + +class TestRedaction: + @pytest.mark.asyncio + async def test_string_content_is_redacted(self): + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["Email [EMAIL_1] about it"]}) + + data = {"messages": [{"role": "user", "content": "Email a@b.com about it"}]} + result = await guardrail.async_pre_call_hook( + user_api_key_dict=None, cache=None, data=data, call_type="completion" + ) + + assert result["messages"][0]["content"] == "Email [EMAIL_1] about it" + + @pytest.mark.asyncio + async def test_multimodal_text_parts_are_redacted(self): + """The list content shape is a historical bypass; text parts must be covered.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["call [PHONE_1]"]}) + + data = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "call 555-0100"}, + {"type": "image_url", "image_url": {"url": "http://x/y.png"}}, + ], + } + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert data["messages"][0]["content"][0]["text"] == "call [PHONE_1]" + assert data["messages"][0]["content"][1]["image_url"]["url"] == "http://x/y.png" + + @pytest.mark.asyncio + async def test_request_without_messages_is_untouched(self): + guardrail = _guardrail() + mock = _mock_post(guardrail) + data = {"input": "no messages here"} + + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + mock.assert_not_called() + + @pytest.mark.asyncio + async def test_session_id_is_reused_across_hooks(self): + """Rehydration can only resolve tokens minted under the same session.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}, {"texts": ["a@b.com"]}) + + data = {"messages": [{"role": "user", "content": "a@b.com"}]} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + await guardrail._rehydrate(["[EMAIL_1]"], guardrail._session_id(data)) + + sessions = {call.kwargs["headers"]["X-Session-ID"] for call in mock.call_args_list} + assert len(sessions) == 1 + + +class TestFailClosed: + @pytest.mark.asyncio + async def test_unreachable_shield_blocks_the_request(self): + """Failing open would send the PII upstream, defeating the guardrail.""" + guardrail = _guardrail() + guardrail.async_handler.post = AsyncMock(side_effect=ConnectionError("refused")) + + with pytest.raises(GuardrailRaisedException): + await guardrail.async_pre_call_hook( + user_api_key_dict=None, + cache=None, + data={"messages": [{"role": "user", "content": "a@b.com"}]}, + call_type="completion", + ) + + @pytest.mark.asyncio + async def test_error_status_blocks_the_request(self): + guardrail = _guardrail() + guardrail.async_handler.post = AsyncMock(return_value=_response({"error": "nope"}, status_code=500)) + + with pytest.raises(GuardrailRaisedException): + await guardrail.async_pre_call_hook( + user_api_key_dict=None, + cache=None, + data={"messages": [{"role": "user", "content": "a@b.com"}]}, + call_type="completion", + ) + + @pytest.mark.asyncio + async def test_short_payload_blocks_the_request(self): + """A response that loses an entry would silently misalign the write-back.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": []}) + + with pytest.raises(GuardrailRaisedException): + await guardrail.async_pre_call_hook( + user_api_key_dict=None, + cache=None, + data={"messages": [{"role": "user", "content": "a@b.com"}]}, + call_type="completion", + ) + + +class TestStreamingRehydration: + @pytest.mark.asyncio + async def test_split_placeholder_is_not_emitted_in_fragments(self): + """The window holds back a partial placeholder and releases it once complete.""" + guardrail = _guardrail(event_hook="post_call") + # Shield holds "[EMAIL" back, then releases the restored value. + _mock_post( + guardrail, + {"text": "Email ", "carry": "[EMAIL"}, + {"text": "a@b.com about it", "carry": ""}, + ) + + async def stream(): + yield _chunk("Email [EMAIL") + yield _chunk("_1] about it", finish_reason="stop") + + chunks = await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + emitted = [c.choices[0].delta.content for c in chunks] + assert emitted == ["Email ", "a@b.com about it"] + # No fragment of the placeholder ever reached the client. + assert not any("[EMAIL" in (text or "") for text in emitted) + + @pytest.mark.asyncio + async def test_carry_is_returned_to_the_next_call(self): + guardrail = _guardrail(event_hook="post_call") + mock = _mock_post( + guardrail, + {"text": "", "carry": "hold"}, + {"text": "held-and-more", "carry": ""}, + ) + + async def stream(): + yield _chunk("hold") + yield _chunk("-and-more", finish_reason="stop") + + await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + assert mock.call_args_list[0].kwargs["json"]["carry"] == "" + assert mock.call_args_list[1].kwargs["json"]["carry"] == "hold" + assert mock.call_args_list[1].kwargs["json"]["final"] is True + + @pytest.mark.asyncio + async def test_chunks_are_forwarded_as_they_arrive(self): + """Restoration must not buffer the stream into a single terminal chunk.""" + guardrail = _guardrail(event_hook="post_call") + _mock_post( + guardrail, + {"text": "one ", "carry": ""}, + {"text": "two ", "carry": ""}, + {"text": "three", "carry": ""}, + ) + + async def stream(): + yield _chunk("one ") + yield _chunk("two ") + yield _chunk("three", finish_reason="stop") + + chunks = await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + assert len(chunks) == 3 + assert [c.choices[0].delta.content for c in chunks] == ["one ", "two ", "three"]