From 4509c874701f4b8243479a76b0144995120e380a Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Thu, 3 Sep 2026 11:47:57 -0500 Subject: [PATCH 01/34] 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"] From 8ab969d56caf5a125b88a1ada46b51b21cf4eb60 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Thu, 3 Sep 2026 12:02:48 -0500 Subject: [PATCH 02/34] feat(ui): list llm shield in the guardrail garden Adds the card, preset and logo so operators can pick LLM Shield from the guardrails page the same way as the other partner guardrails. --- .../public/assets/logos/llm_shield.svg | 5 +++++ .../guardrails/_components/guardrail_garden_configs.ts | 6 ++++++ .../_components/guardrail_garden_data.test.ts | 1 + .../guardrails/_components/guardrail_garden_data.ts | 10 ++++++++++ .../guardrails/_components/guardrail_info_helpers.tsx | 3 +++ 5 files changed, 25 insertions(+) create mode 100644 ui/litellm-dashboard/public/assets/logos/llm_shield.svg diff --git a/ui/litellm-dashboard/public/assets/logos/llm_shield.svg b/ui/litellm-dashboard/public/assets/logos/llm_shield.svg new file mode 100644 index 00000000000..d61edff4473 --- /dev/null +++ b/ui/litellm-dashboard/public/assets/logos/llm_shield.svg @@ -0,0 +1,5 @@ + + + + + diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts index 7785a8e44ab..b445cc9c5ad 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts @@ -318,4 +318,10 @@ export const GUARDRAIL_PRESETS: Record = { mode: "pre_call", defaultOn: false, }, + llm_shield: { + provider: "LLM Shield", + guardrailNameSuggestion: "LLM Shield", + mode: "pre_call", + defaultOn: false, + }, }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts index 1e486639840..d3212d66737 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts @@ -28,6 +28,7 @@ const EXPECTED_PARTNER_LOGO_FILES: Record = { repelloai: "repelloai.png", straiker: "straiker.svg", alice: "alice.svg", + llm_shield: "llm_shield.svg", }; describe("guardrail_garden_data logos", () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts index 931b3a111d8..d46847a5a80 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts @@ -474,6 +474,16 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ tags: ["Content Moderation", "Prompt Injection", "PII", "Policy"], providerKey: "Alice", }, + { + id: "llm_shield", + name: "LLM Shield", + description: + "Self-hosted PII redaction that puts the original values back into the model's response, so the provider never receives personal data while the end user still sees it.", + category: "partner", + logo: guardrailLogoMap["LLM Shield"], + tags: ["PII", "Data Privacy", "Compliance", "Streaming"], + providerKey: "LLM Shield", + }, ]; export const ALL_CARDS = [...LITELLM_CONTENT_FILTER_CARDS, ...PARTNER_GUARDRAIL_CARDS]; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx index f686ff5644a..f620ea6dcd0 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx @@ -1,6 +1,7 @@ import aimSecurityLogo from "../../../../../public/assets/logos/aim_security.jpeg"; import aktoLogo from "../../../../../public/assets/logos/akto.svg"; import aliceLogo from "../../../../../public/assets/logos/alice.svg"; +import llmShieldLogo from "../../../../../public/assets/logos/llm_shield.svg"; import aporiaLogo from "../../../../../public/assets/logos/aporia.png"; import bedrockLogo from "../../../../../public/assets/logos/bedrock.svg"; import catoNetworksLogo from "../../../../../public/assets/logos/cato_networks.svg"; @@ -85,6 +86,7 @@ export const guardrail_provider_map: Record = { QostodianNexus: "qostodian_nexus", Repelloai: "repelloai", Alice: "alice", + "LLM Shield": "llm_shield", }; // Function to populate provider map from API response - updates the original map @@ -208,6 +210,7 @@ export const guardrailLogoMap = { "RepelloAI Argus": repelloAiLogo.src, Straiker: straikerLogo.src, Alice: aliceLogo.src, + "LLM Shield": llmShieldLogo.src, } satisfies Record; export const getGuardrailLogo = (displayName: string): string | undefined => From 2da630debd3e28be0ef5c2840b6941e4c9733bba Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Thu, 3 Sep 2026 12:15:38 -0500 Subject: [PATCH 03/34] docs(guardrails): add llm shield example config Shows both modes on one entry. Listing only pre_call redacts the request and then hands the placeholders back to the end user, so the test asserts both hooks are enabled. --- .../llm_shield/example_config.yaml | 57 +++++++++++++++++++ .../guardrail_hooks/test_llm_shield.py | 14 +++++ 2 files changed, 71 insertions(+) create mode 100644 litellm/proxy/guardrails/guardrail_hooks/llm_shield/example_config.yaml diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/example_config.yaml b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/example_config.yaml new file mode 100644 index 00000000000..3a4b43d5432 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/example_config.yaml @@ -0,0 +1,57 @@ +# Example LiteLLM Proxy configuration for LLM Shield +# LLM Shield is a self-hosted PII gateway: https://github.com/ninadphalak/LLM-Shield-Proxy +# +# Unlike a masking guardrail, LLM Shield's substitution is reversible. Personal data is +# replaced with placeholders before the request goes to the provider, and the original +# values are put back into the model's reply, so the end user still sees real data while +# the provider never received it. + +model_list: + - model_name: gpt-4o + litellm_params: + model: openai/gpt-4o + api_key: os.environ/OPENAI_API_KEY + +guardrails: + # Both modes belong on ONE entry. pre_call redacts the outbound request and post_call + # restores the reply; listing only pre_call would send placeholders back to the user. + - guardrail_name: "llm-shield" + litellm_params: + guardrail: llm_shield + mode: ["pre_call", "post_call"] + default_on: true + # Your own LLM Shield deployment. Defaults to http://localhost:8000, and also reads + # LLM_SHIELD_API_BASE from the environment. + api_base: "http://localhost:8000" + # A virtual key configured on that deployment. Also reads LLM_SHIELD_API_KEY. + api_key: os.environ/LLM_SHIELD_API_KEY + +# Usage: +# +# 1. Run LLM Shield somewhere the proxy can reach: +# pip install llm-shield-proxy +# llm-shield-proxy serve +# +# 2. Point this config at it and start the proxy: +# export LLM_SHIELD_API_KEY="your-virtual-key" +# litellm --config example_config.yaml +# +# 3. Send a request containing personal data: +# curl http://localhost:4000/v1/chat/completions \ +# -H "Authorization: Bearer sk-1234" \ +# -H "Content-Type: application/json" \ +# -d '{"model":"gpt-4o","messages":[{"role":"user","content":"Email jane.doe@example.com the invoice"}]}' +# +# The provider receives a stand-in value in place of the address. The reply you get +# back carries the real address again. +# +# Notes: +# +# - Requests are refused if LLM Shield is unreachable or returns an error, rather than +# being forwarded. Sending them on would hand the provider exactly the data this +# guardrail exists to withhold. +# - Restoring a value requires the request and the reply to share a session. LiteLLM's +# session id is used when present; otherwise one is generated per request. +# - Streaming replies are restored as chunks arrive. A placeholder split across two +# chunks is held back until it is complete, so partial values are never emitted. +# - Only text is redacted; images and audio pass through untouched. 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 index 4b39c8bf517..160c2300690 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py @@ -10,6 +10,7 @@ from litellm.proxy.guardrails.guardrail_hooks.llm_shield.llm_shield import ( LLMShieldGuardrail, ) from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2 +from litellm.types.guardrails import GuardrailEventHooks from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices @@ -82,6 +83,19 @@ class TestLLMShieldInitialization: def test_trailing_slash_is_stripped(self): assert _guardrail(api_base="http://shield.test/").api_base == "http://shield.test" + def test_both_modes_can_be_enabled_on_one_entry(self): + """Redaction and restoration are two halves of one config entry. + + A deployment that lists only pre_call would redact the request and then hand + the placeholders straight back to the end user. + """ + guardrail = _guardrail(event_hook=["pre_call", "post_call"]) + data: dict = {"messages": []} + + assert guardrail.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_call) is True + assert guardrail.should_run_guardrail(data=data, event_type=GuardrailEventHooks.post_call) is True + assert guardrail.should_run_guardrail(data=data, event_type=GuardrailEventHooks.during_call) is False + class TestRedaction: @pytest.mark.asyncio From 5fccbfe49f83913270116b28d18646412b9f87af Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Thu, 3 Sep 2026 12:28:18 -0500 Subject: [PATCH 04/34] feat(ui): use the llm shield brand mark for the guardrail logo --- ui/litellm-dashboard/public/assets/logos/llm_shield.svg | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/ui/litellm-dashboard/public/assets/logos/llm_shield.svg b/ui/litellm-dashboard/public/assets/logos/llm_shield.svg index d61edff4473..0dd78b078c9 100644 --- a/ui/litellm-dashboard/public/assets/logos/llm_shield.svg +++ b/ui/litellm-dashboard/public/assets/logos/llm_shield.svg @@ -1,5 +1,6 @@ - - - + + + + From ee0bac5148e56d400d699ef707ade6f8dabe55a7 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Thu, 3 Sep 2026 14:23:14 -0500 Subject: [PATCH 05/34] fix(guardrails): restore llm shield values in anthropic replies The /v1/messages reply is a plain dict with a content block list and no choices, so it fell through the restore path and went back to the caller still carrying placeholders. The request was redacted correctly, which is what made this easy to miss. Found by running all three endpoints against a live provider; the mocked tests all passed because they only built the OpenAI shape. Adds tests for the message shape and for leaving non-text blocks alone. --- .../guardrail_hooks/llm_shield/llm_shield.py | 31 ++++++++++ .../guardrail_hooks/test_llm_shield.py | 58 ++++++++++++++++++- 2 files changed, 88 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py index 4199ad5ca65..125bea5590e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py @@ -227,6 +227,9 @@ class LLMShieldGuardrail(CustomGuardrail): if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.post_call) is not True: return response + if self._is_anthropic_message_response(response): + return await self._restore_anthropic_response(response, data) + choices = getattr(response, "choices", None) if not choices: return response @@ -246,6 +249,34 @@ class LLMShieldGuardrail(CustomGuardrail): message.content = replacement return response + @staticmethod + def _is_anthropic_message_response(response: Any) -> bool: + """Anthropic's native /v1/messages reply arrives as a plain dict.""" + return ( + isinstance(response, dict) + and response.get("type") == "message" + and isinstance(response.get("content"), list) + ) + + async def _restore_anthropic_response(self, response: dict, data: dict) -> dict: + """Restores text blocks in an Anthropic native message reply. + + This shape has no `choices`, so without its own branch the reply would go + back to the caller still carrying placeholders. + """ + blocks = [ + block + for block in response["content"] + if isinstance(block, dict) and block.get("type") == "text" and isinstance(block.get("text"), str) + ] + if not blocks: + return response + + restored = await self._rehydrate([block["text"] for block in blocks], self._session_id(data)) + for block, replacement in zip(blocks, restored): + block["text"] = replacement + return response + async def async_post_call_streaming_iterator_hook( self, user_api_key_dict: UserAPIKeyAuth, 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 index 160c2300690..c2d50bf5fa7 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py @@ -11,7 +11,7 @@ from litellm.proxy.guardrails.guardrail_hooks.llm_shield.llm_shield import ( ) from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2 from litellm.types.guardrails import GuardrailEventHooks -from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices +from litellm.types.utils import Choices, Delta, Message, ModelResponse, ModelResponseStream, StreamingChoices def _guardrail(**overrides: object) -> LLMShieldGuardrail: @@ -156,6 +156,62 @@ class TestRedaction: assert len(sessions) == 1 +class TestRestoration: + @pytest.mark.asyncio + async def test_openai_shape_is_restored(self): + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["a@b.com"]}) + + response = ModelResponse(choices=[Choices(index=0, message=Message(role="assistant", content="[EMAIL_1]"))]) + result = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=response + ) + + assert result.choices[0].message.content == "a@b.com" + + @pytest.mark.asyncio + async def test_anthropic_message_shape_is_restored(self): + """The /v1/messages reply is a plain dict with no choices. + + Measured against a live provider: without its own branch the reply went + back to the caller still carrying the placeholder, even though the + request had been redacted correctly. + """ + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["a@b.com"]}) + + response = { + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "[EMAIL_1]"}], + } + result = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=response + ) + + assert result["content"][0]["text"] == "a@b.com" + + @pytest.mark.asyncio + async def test_anthropic_non_text_blocks_are_left_alone(self): + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["a@b.com"]}) + + response = { + "type": "message", + "role": "assistant", + "content": [ + {"type": "text", "text": "[EMAIL_1]"}, + {"type": "tool_use", "id": "t1", "name": "lookup", "input": {}}, + ], + } + result = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=response + ) + + assert result["content"][0]["text"] == "a@b.com" + assert result["content"][1] == {"type": "tool_use", "id": "t1", "name": "lookup", "input": {}} + + class TestFailClosed: @pytest.mark.asyncio async def test_unreachable_shield_blocks_the_request(self): From 0de4b8b9a83ade05a7b3d6b82932568ddb18f537 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Thu, 3 Sep 2026 14:46:28 -0500 Subject: [PATCH 06/34] docs(guardrails): correct the llm shield start command --- .../guardrails/guardrail_hooks/llm_shield/example_config.yaml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/example_config.yaml b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/example_config.yaml index 3a4b43d5432..aa63fa9d252 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/example_config.yaml +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/example_config.yaml @@ -30,7 +30,7 @@ guardrails: # # 1. Run LLM Shield somewhere the proxy can reach: # pip install llm-shield-proxy -# llm-shield-proxy serve +# llm-shield-proxy --port 8000 # # 2. Point this config at it and start the proxy: # export LLM_SHIELD_API_KEY="your-virtual-key" From b6e3e6decd82c249255e6dc7dbdd2d9b20992237 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Thu, 3 Sep 2026 18:07:32 -0500 Subject: [PATCH 07/34] fix(guardrails): redact every request shape and restore every reply shape Three gaps, all of which let an enabled guardrail hand data to the provider or hand placeholders to the caller. Requests only walked `messages`. The Responses API `input` and tool call `arguments` went out untouched. Measured against a live provider: a request sent through `/v1/responses` reached the model with the real address in it while the guardrail reported as enabled. Request traversal now covers chat content (string and multimodal), tool call arguments, and `input` as a bare string or a list of items. Fixing that exposed the matching gap on the way back: the Responses API reply carries `output` items rather than `choices`, so it returned to the caller still holding placeholders. It now gets its own walk, handling text blocks as dicts or objects. The dashboard preset seeded only pre_call, so a guardrail created from the UI would redact the request and return the placeholders to the user. Presets can now seed both modes; the form already normalised either shape. Adds tests for each request shape, for both Responses API reply forms, and replaces a test that had asserted the `input` bypass as correct behaviour. --- .../guardrail_hooks/llm_shield/llm_shield.py | 291 +++++++++++------- ruff-strict.toml | 4 + .../guardrail_hooks/test_llm_shield.py | 126 +++++++- .../_components/add_guardrail_form.tsx | 4 +- .../_components/guardrail_garden_configs.ts | 8 +- 5 files changed, 323 insertions(+), 110 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py index 125bea5590e..fe8b70c08b8 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py @@ -7,7 +7,7 @@ import os import uuid -from collections.abc import AsyncGenerator +from collections.abc import AsyncGenerator, Callable, Mapping, Sequence from typing import ( TYPE_CHECKING, Any, # noqa: TID251 # **kwargs forwards verbatim to CustomGuardrail.__init__ @@ -15,6 +15,7 @@ from typing import ( Final, Literal, Optional, + TypeAlias, ) import httpx @@ -53,6 +54,19 @@ _SESSION_METADATA_KEY: Final = "llm_shield_session_id" _DEFAULT_TIMEOUT_SECONDS: Final = 10.0 +# The proxy's own request dict. Mutable by design: a pre-call guardrail rewrites +# the caller's payload in place, which is the entire point of the hook. +# mutable-ok: the shape is fixed by CustomLogger's hook signatures. +MutableRequest: TypeAlias = dict + +# A JSON body on its way to httpx, which requires a real dict rather than a view. +# mutable-ok: handed straight to the HTTP client. +JsonBody: TypeAlias = dict + +# One redactable span: the text as it stands, and the write that puts the +# replacement back where it came from. +_Slot: TypeAlias = tuple[str, Callable[[str], None]] # mutable-ok: Callable's param list. + class LLMShieldGuardrail(CustomGuardrail): """Redacts PII before it leaves the proxy and restores it in the response. @@ -78,7 +92,7 @@ class LLMShieldGuardrail(CustomGuardrail): guardrail_name: str = GUARDRAIL_NAME, api_base: str | None = None, api_key: str | None = None, - **kwargs: Any, + **kwargs: Any, # noqa: LIT008 # kwargs-ok: forwarded verbatim to CustomGuardrail.__init__ ) -> 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("/") @@ -86,18 +100,21 @@ class LLMShieldGuardrail(CustomGuardrail): super().__init__(guardrail_name=guardrail_name, **kwargs) @classmethod - def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: - return [GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call] + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: # mutable-ok: parent's signature. + return [GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call] # mutable-ok: parent's signature. # --- transport --------------------------------------------------------------- - def _headers(self, session_id: str) -> dict: - headers = {"Content-Type": "application/json", "X-Session-ID": session_id} + def _headers(self, session_id: str) -> JsonBody: + headers: Final[JsonBody] = { # mutable-ok: httpx requires a real dict. + "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: + async def _call_shield(self, path: str, session_id: str, payload: JsonBody) -> Mapping[str, object]: """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 @@ -105,7 +122,7 @@ class LLMShieldGuardrail(CustomGuardrail): blocks the request instead of passing it through. """ try: - response = await self.async_handler.post( + response: Final = await self.async_handler.post( f"{self.api_base}{path}", headers=self._headers(session_id), json=payload, @@ -126,69 +143,89 @@ class LLMShieldGuardrail(CustomGuardrail): 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}) + async def _redact(self, texts: Sequence[str], session_id: str) -> Sequence[str]: + payload: Final[JsonBody] = {"texts": list(texts)} # mutable-ok: JSON body for httpx. + body: Final = await self._call_shield(_REDACT_PATH, session_id, payload) 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}) + async def _rehydrate(self, texts: Sequence[str], session_id: str) -> Sequence[str]: + payload: Final[JsonBody] = {"texts": list(texts)} # mutable-ok: JSON body for httpx. + body: Final = await self._call_shield(_REHYDRATE_PATH, session_id, payload) return self._same_length_or_raise(body.get("texts"), texts, "rehydrate") - def _same_length_or_raise(self, returned: Any, sent: list, operation: str) -> list: + def _same_length_or_raise(self, returned: object, sent: Sequence[str], operation: str) -> Sequence[str]: """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 + return tuple(returned) # --- session ------------------------------------------------------------------ - def _session_id(self, data: dict) -> str: + def _session_id(self, data: MutableRequest) -> str: """Returns a session id stable across this request's hooks.""" - metadata = data.setdefault("metadata", {}) + metadata: Final = data.setdefault("metadata", {}) # mutable-ok: per-request store. if not isinstance(metadata, dict): return f"litellm-{uuid.uuid4().hex}" - existing = metadata.get(_SESSION_METADATA_KEY) + existing: Final = 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}" + session_id: Final = get_session_id_from_request_data(data) or f"litellm-{uuid.uuid4().hex}" metadata[_SESSION_METADATA_KEY] = session_id return session_id - # --- message traversal -------------------------------------------------------- + # --- request traversal -------------------------------------------------------- @staticmethod - def _locate_texts(messages: list) -> list: - """Finds every text span in a message list. + def _locate_request_texts(data: MutableRequest) -> Sequence[_Slot]: + """Finds every redactable span in an outbound request. - Returns ``(message_index, part_index_or_None, text)``. The list form is the - multimodal shape, where only ``text`` parts carry redactable content. + Returns ``(text, write)`` pairs. Any shape missed here reaches the provider + in the clear, so this walks all of the request shapes that carry caller text: + + - chat ``messages``, both string and multimodal list ``content`` + - tool call ``arguments``, which routinely carry the values a user asked + the model to look up + - the Responses API ``input``, as a bare string or a list of items """ - 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 + slots: Final[list[_Slot]] = [] # mutable-ok: accumulator, frozen on return. - @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 + def add(container: MutableRequest, key: str, value: object) -> None: + if isinstance(value, str) and value: + slots.append((value, lambda new, c=container, k=key: c.__setitem__(k, new))) + + def add_content(container: MutableRequest) -> None: + """Adds `content`, which is either a string or a list of typed parts.""" + content: Final = container.get("content") + if isinstance(content, str): + add(container, "content", content) + return + for part in content if isinstance(content, list) else (): + if isinstance(part, dict): + add(part, "text", part.get("text")) + + def add_tool_calls(message: MutableRequest) -> None: + for tool_call in message.get("tool_calls") or (): + function = tool_call.get("function") if isinstance(tool_call, dict) else None + if isinstance(function, dict): + add(function, "arguments", function.get("arguments")) + + for message in data.get("messages") or (): + if isinstance(message, dict): + add_content(message) + add_tool_calls(message) + + request_input: Final = data.get("input") + if isinstance(request_input, str): + add(data, "input", request_input) + else: + for item in request_input if isinstance(request_input, list) else (): + if isinstance(item, dict): + add_content(item) + + return tuple(slots) # --- hooks -------------------------------------------------------------------- @@ -197,29 +234,26 @@ class LLMShieldGuardrail(CustomGuardrail): self, user_api_key_dict: UserAPIKeyAuth, cache: "DualCache", - data: dict, + data: MutableRequest, call_type: str, - ) -> dict | None: - """Replaces PII in the outbound messages with vault placeholders.""" + ) -> MutableRequest | None: + """Replaces PII anywhere in the outbound request 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): + slots: Final = self._locate_request_texts(data) + if not slots: 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) + redacted: Final = await self._redact(tuple(text for text, _ in slots), self._session_id(data)) + for (_, write), replacement in zip(slots, redacted): + write(replacement) return data @log_guardrail_information async def async_post_call_success_hook( self, - data: dict, + data: MutableRequest, user_api_key_dict: UserAPIKeyAuth, response: Any, ) -> Any: @@ -230,27 +264,31 @@ class LLMShieldGuardrail(CustomGuardrail): if self._is_anthropic_message_response(response): return await self._restore_anthropic_response(response, data) - choices = getattr(response, "choices", None) + text_blocks: Final = self._responses_api_text_blocks(response) + if text_blocks: + return await self._restore_responses_api_response(response, text_blocks, data) + + choices: Final = 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)) - + pending: Final = tuple( + (choice.message, choice.message.content) + for choice in choices + if getattr(choice, "message", None) is not None + and isinstance(getattr(choice.message, "content", None), str) + and choice.message.content + ) if not pending: return response - restored = await self._rehydrate([text for _, text in pending], self._session_id(data)) + restored: Final = await self._rehydrate(tuple(text for _, text in pending), self._session_id(data)) for (message, _), replacement in zip(pending, restored): message.content = replacement return response @staticmethod - def _is_anthropic_message_response(response: Any) -> bool: + def _is_anthropic_message_response(response: object) -> bool: """Anthropic's native /v1/messages reply arrives as a plain dict.""" return ( isinstance(response, dict) @@ -258,30 +296,68 @@ class LLMShieldGuardrail(CustomGuardrail): and isinstance(response.get("content"), list) ) - async def _restore_anthropic_response(self, response: dict, data: dict) -> dict: + async def _restore_anthropic_response(self, response: MutableRequest, data: MutableRequest) -> MutableRequest: """Restores text blocks in an Anthropic native message reply. This shape has no `choices`, so without its own branch the reply would go back to the caller still carrying placeholders. """ - blocks = [ + blocks: Final = tuple( block for block in response["content"] if isinstance(block, dict) and block.get("type") == "text" and isinstance(block.get("text"), str) - ] + ) if not blocks: return response - restored = await self._rehydrate([block["text"] for block in blocks], self._session_id(data)) + restored: Final = await self._rehydrate(tuple(block["text"] for block in blocks), self._session_id(data)) for block, replacement in zip(blocks, restored): block["text"] = replacement return response + @staticmethod + def _responses_api_text_blocks(response: object) -> Sequence[object]: + """Text blocks in a Responses API reply. + + That shape carries `output` items rather than `choices`, so it needs its own + walk; without one the reply goes back to the caller still holding + placeholders even though the request was redacted correctly. Blocks come + through as dicts or as objects depending on how far the reply has been + deserialised, so both are handled. + """ + blocks: Final[list[object]] = [] # mutable-ok: accumulator, frozen on return. + for item in getattr(response, "output", None) or (): + for block in getattr(item, "content", None) or (): + if isinstance(block, dict): + if isinstance(block.get("text"), str) and block["text"]: + blocks.append(block) + elif isinstance(getattr(block, "text", None), str) and block.text: + blocks.append(block) + return tuple(blocks) + + @staticmethod + def _block_text(block: object) -> str: + return block["text"] if isinstance(block, dict) else block.text + + async def _restore_responses_api_response( + self, response: Any, blocks: Sequence[object], data: MutableRequest + ) -> Any: + """Puts the original values back into a Responses API reply.""" + restored: Final = await self._rehydrate( + tuple(self._block_text(block) for block in blocks), self._session_id(data) + ) + for block, replacement in zip(blocks, restored): + if isinstance(block, dict): + block["text"] = replacement + else: + block.text = replacement + return response + async def async_post_call_streaming_iterator_hook( self, user_api_key_dict: UserAPIKeyAuth, response: Any, - request_data: dict, + request_data: MutableRequest, ) -> AsyncGenerator[Any, None]: """Restores original values incrementally, without buffering the stream. @@ -295,9 +371,9 @@ class LLMShieldGuardrail(CustomGuardrail): yield chunk return - session_id = self._session_id(request_data) - carry = "" - last_chunk = None + session_id: Final = self._session_id(request_data) + carry = "" # rebind-ok: the sliding window advances with every delta. + last_chunk = None # rebind-ok: tracks the most recent chunk for the final flush. async for chunk in response: last_chunk = chunk @@ -309,53 +385,54 @@ class LLMShieldGuardrail(CustomGuardrail): # 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"] + emitted, carry = await self._stream_step("", carry, True, session_id) + if emitted and delta is not None: + delta.content = emitted yield chunk continue - body = await self._stream_step(text, carry, is_final, session_id) - carry = body["carry"] - delta.content = body["text"] + emitted, carry = await self._stream_step(text, carry, is_final, session_id) + delta.content = emitted 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) + flushed: Final = await self._stream_step("", carry, True, session_id) + trailing_text, carry = flushed # rebind-ok: window advances. + if trailing_text: + trailing: Final = last_chunk.model_copy(deep=True) + trailing_delta: Final = self._stream_delta(trailing) if trailing_delta is not None: - trailing_delta.content = body["text"] + trailing_delta.content = trailing_text yield trailing - async def _stream_step(self, text: str, carry: str, final: bool, session_id: str) -> dict: - body = await self._call_shield( + async def _stream_step(self, text: str, carry: str, final: bool, session_id: str) -> tuple[str, str]: + """Returns ``(text safe to emit now, window still being held)``.""" + body: Final = await self._call_shield( _REHYDRATE_STREAM_PATH, session_id, - {"text": text, "carry": carry, "final": final}, + # mutable-ok: JSON request body for httpx. + {"text": text, "carry": carry, "final": final}, # mutable-ok: JSON request body for httpx. ) - emitted = body.get("text") - remaining = body.get("carry") + emitted: Final = body.get("text") + remaining: Final = 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} + return emitted, remaining @staticmethod - def _stream_delta(chunk: Any) -> Any: - choices = getattr(chunk, "choices", None) + def _stream_delta(chunk: object) -> Any: + choices: Final = 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) + def _is_final_chunk(chunk: object) -> bool: + choices: Final = getattr(chunk, "choices", None) if not choices: return False return bool(getattr(choices[0], "finish_reason", None)) @@ -366,17 +443,21 @@ class LLMShieldGuardrail(CustomGuardrail): async def apply_guardrail( self, inputs: GenericGuardrailAPIInputs, - request_data: dict, + request_data: MutableRequest, input_type: Literal["request", "response"], logging_obj: Optional["LiteLLMLoggingObj"] = None, ) -> GenericGuardrailAPIInputs: - texts = inputs.get("texts") + texts: Final = 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 + session_id: Final = self._session_id(request_data) + replaced: Final = ( + await self._redact(tuple(texts), session_id) + if input_type == "request" + else await self._rehydrate(tuple(texts), session_id) + ) + # Return a new mapping rather than rewriting the caller's, so this stays a + # pure transform of the inputs it was handed. + merged: Final[JsonBody] = {**inputs, "texts": list(replaced)} # mutable-ok: TypedDict. + return merged diff --git a/ruff-strict.toml b/ruff-strict.toml index ae092bdde7d..f9d026c6011 100644 --- a/ruff-strict.toml +++ b/ruff-strict.toml @@ -30,6 +30,10 @@ external = [ # grows over time; typing it concretely (`object`) broke that forwarding call outright — # basedpyright turned every named param into a reportArgumentType error. Any is correct here. "litellm/proxy/guardrails/guardrail_hooks/alice/alice.py" = ["ANN401"] +# Same reason: `**kwargs` forwards verbatim to CustomGuardrail.__init__, and the lifecycle +# hook signatures inherit `Any` for `response` from CustomLogger, so narrowing them here +# would break the override rather than describe it. +"litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py" = ["ANN401"] [lint.mccabe] max-complexity = 15 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 index c2d50bf5fa7..7a07c9761a8 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py @@ -1,3 +1,4 @@ +from types import SimpleNamespace from unittest.mock import AsyncMock import pytest @@ -133,10 +134,16 @@ class TestRedaction: 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): + async def test_request_without_text_is_untouched(self): + """No text to redact means no call to LLM Shield. + + This deliberately uses a request with no caller text at all. An earlier + version used a Responses-API `input`, which asserted the very bypass that + let `input` reach the provider unredacted. + """ guardrail = _guardrail() mock = _mock_post(guardrail) - data = {"input": "no messages here"} + data = {"model": "gpt-4o", "temperature": 0.2} await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") @@ -156,6 +163,91 @@ class TestRedaction: assert len(sessions) == 1 +class TestRequestCoverage: + """Every request shape that carries caller text must be redacted. + + A shape missed here is not a cosmetic gap: the guardrail reports as enabled + while the raw value goes to the provider. + """ + + @pytest.mark.asyncio + async def test_responses_api_string_input_is_redacted(self): + """Measured against a live provider: `input` reached the model unredacted.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["Email [EMAIL_1] the invoice"]}) + + data = {"input": "Email jane.doe@example.com the invoice"} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["Email jane.doe@example.com the invoice"] + assert data["input"] == "Email [EMAIL_1] the invoice" + + @pytest.mark.asyncio + async def test_responses_api_list_input_is_redacted(self): + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]", "[PHONE_1]"]}) + + data = { + "input": [ + {"role": "user", "content": "jane.doe@example.com"}, + {"role": "user", "content": [{"type": "input_text", "text": "555-0100"}]}, + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert data["input"][0]["content"] == "[EMAIL_1]" + assert data["input"][1]["content"][0]["text"] == "[PHONE_1]" + + @pytest.mark.asyncio + async def test_tool_call_arguments_are_redacted(self): + """Tool arguments carry the values the user asked the model to act on.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ['{"email": "[EMAIL_1]"}']}) + + data = { + "messages": [ + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "send", "arguments": '{"email": "jane.doe@example.com"}'}, + } + ], + } + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert data["messages"][0]["tool_calls"][0]["function"]["arguments"] == '{"email": "[EMAIL_1]"}' + + @pytest.mark.asyncio + async def test_every_shape_in_one_request_is_redacted(self): + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["a", "b", "c", "d"]}) + + data = { + "messages": [ + {"role": "user", "content": "one"}, + {"role": "user", "content": [{"type": "text", "text": "two"}]}, + { + "role": "assistant", + "tool_calls": [{"function": {"name": "f", "arguments": "three"}}], + }, + ], + "input": "four", + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["one", "two", "three", "four"] + assert data["messages"][0]["content"] == "a" + assert data["messages"][1]["content"][0]["text"] == "b" + assert data["messages"][2]["tool_calls"][0]["function"]["arguments"] == "c" + assert data["input"] == "d" + + class TestRestoration: @pytest.mark.asyncio async def test_openai_shape_is_restored(self): @@ -169,6 +261,36 @@ class TestRestoration: assert result.choices[0].message.content == "a@b.com" + @pytest.mark.asyncio + async def test_responses_api_shape_is_restored(self): + """The Responses API reply carries output items, not choices. + + Measured against a live provider: once the request side was fixed the reply + came back still holding the placeholder, because this shape has no choices + to walk. + """ + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["a@b.com"]}) + + response = SimpleNamespace(output=[SimpleNamespace(content=[{"type": "output_text", "text": "[EMAIL_1]"}])]) + result = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=response + ) + + assert result.output[0].content[0]["text"] == "a@b.com" + + @pytest.mark.asyncio + async def test_responses_api_object_blocks_are_restored(self): + """Blocks arrive as objects too, depending on how far the reply is parsed.""" + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["a@b.com"]}) + + block = SimpleNamespace(text="[EMAIL_1]") + response = SimpleNamespace(output=[SimpleNamespace(content=[block])]) + await guardrail.async_post_call_success_hook(data={"messages": []}, user_api_key_dict=None, response=response) + + assert block.text == "a@b.com" + @pytest.mark.asyncio async def test_anthropic_message_shape_is_restored(self): """The /v1/messages reply is a plain dict with no choices. diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx index 29df7c8bf3d..a02cae097a7 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx @@ -73,7 +73,9 @@ interface GuardrailPreset { provider: string; categoryName?: string; guardrailNameSuggestion: string; - mode: string; + // A guardrail that both rewrites the request and repairs the response needs two + // modes seeded, not one; the form already normalises either shape. + mode: string | string[]; defaultOn: boolean; } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts index b445cc9c5ad..3579457bcd5 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts @@ -2,7 +2,9 @@ export interface GuardrailPreset { provider: string; categoryName?: string; guardrailNameSuggestion: string; - mode: string; + // A guardrail that both rewrites the request and repairs the response needs two + // modes seeded, not one; the form already normalises either shape. + mode: string | string[]; defaultOn: boolean; } @@ -321,7 +323,9 @@ export const GUARDRAIL_PRESETS: Record = { llm_shield: { provider: "LLM Shield", guardrailNameSuggestion: "LLM Shield", - mode: "pre_call", + // Both halves are required. With only pre_call the request is redacted and the + // placeholders are handed straight back to the caller. + mode: ["pre_call", "post_call"], defaultOn: false, }, }; From 295ad527d50540a3eb66cb5a29933ef4cac07102 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Thu, 3 Sep 2026 19:13:31 -0500 Subject: [PATCH 08/34] fix(guardrails): narrow the stream delta before writing to it basedpyright could not prove the delta was non-None on the write path, and reportOptionalMemberAccess has a zero budget. The guard is also clearer than relying on the text check to imply it. --- .../guardrails/guardrail_hooks/llm_shield/llm_shield.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py index fe8b70c08b8..c126bf0de37 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py @@ -381,12 +381,12 @@ class LLMShieldGuardrail(CustomGuardrail): 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: + if delta is None or 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: + if is_final and carry and delta is not None: emitted, carry = await self._stream_step("", carry, True, session_id) - if emitted and delta is not None: + if emitted: delta.content = emitted yield chunk continue From 47421b7541c6684db3f91c8ba627ceed15a40468 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Thu, 3 Sep 2026 19:24:12 -0500 Subject: [PATCH 09/34] fix(guardrails): mint the vault id instead of trusting the caller's The vault id was taken from caller-supplied session metadata, and every caller shares one LLM Shield key. Someone who knew or guessed another caller's session id could send a placeholder, have the model echo it back, and get that caller's plaintext restored into their own reply. Vault ids are now minted per request behind a per-process prefix, so a caller cannot name a vault this process uses. Redaction mints, restoration reads back, and a reply whose id does not match is left holding its placeholders rather than resolved against some other vault. Also covers two more request fields that were reaching the provider intact: the Responses API `instructions`, and the legacy `function_call.arguments` alongside `tool_calls`. The collectors move to module level, which drops the traversal back under the complexity limit and lets the code carry its own explanation instead of the comments that were restating it. --- .../guardrail_hooks/llm_shield/llm_shield.py | 141 +++++++++++------- .../guardrail_hooks/test_llm_shield.py | 82 ++++++++++ 2 files changed, 169 insertions(+), 54 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py index c126bf0de37..26c6774f48d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py @@ -24,7 +24,6 @@ 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 ( @@ -52,6 +51,13 @@ _REHYDRATE_STREAM_PATH: Final = "/v1/guard/rehydrate/stream" # across concurrent requests. _SESSION_METADATA_KEY: Final = "llm_shield_session_id" +# Vault ids are minted here and never derived from anything the caller sends. The +# vault holds the plaintext behind every placeholder, so an id a caller could +# supply or guess would let one user rehydrate another user's values by getting a +# placeholder echoed back. The per-process prefix means a caller cannot even name +# a vault this process uses. +_VAULT_PREFIX: Final = f"litellm-{uuid.uuid4().hex}" + _DEFAULT_TIMEOUT_SECONDS: Final = 10.0 # The proxy's own request dict. Mutable by design: a pre-call guardrail rewrites @@ -67,6 +73,51 @@ JsonBody: TypeAlias = dict # replacement back where it came from. _Slot: TypeAlias = tuple[str, Callable[[str], None]] # mutable-ok: Callable's param list. +# The accumulator the collectors below append into. It never escapes +# _locate_request_texts, which freezes it into a tuple before returning. +_SlotSink: TypeAlias = list[_Slot] # mutable-ok: accumulator passed between collectors. + + +def _collect(container: MutableRequest, key: str, slots: _SlotSink) -> None: + """Records the string at `key`, along with the write that replaces it.""" + value: Final = container.get(key) + if isinstance(value, str) and value: + slots.append((value, lambda new, c=container, k=key: c.__setitem__(k, new))) + + +def _collect_content(container: MutableRequest, slots: _SlotSink) -> None: + """`content` is either a string or the multimodal list of typed parts.""" + content: Final = container.get("content") + if isinstance(content, str): + _collect(container, "content", slots) + return + for part in content if isinstance(content, list) else (): + if isinstance(part, dict): + _collect(part, "text", slots) + + +def _collect_tool_arguments(message: MutableRequest, slots: _SlotSink) -> None: + """Tool arguments carry the values a user asked the model to act on.""" + for tool_call in message.get("tool_calls") or (): + function: Final = tool_call.get("function") if isinstance(tool_call, dict) else None + if isinstance(function, dict): + _collect(function, "arguments", slots) + legacy: Final = message.get("function_call") + if isinstance(legacy, dict): + _collect(legacy, "arguments", slots) + + +def _collect_responses_fields(data: MutableRequest, slots: _SlotSink) -> None: + """The Responses API sends text outside `messages`, in `instructions` and `input`.""" + _collect(data, "instructions", slots) + request_input: Final = data.get("input") + if isinstance(request_input, str): + _collect(data, "input", slots) + return + for item in request_input if isinstance(request_input, list) else (): + if isinstance(item, dict): + _collect_content(item, slots) + class LLMShieldGuardrail(CustomGuardrail): """Redacts PII before it leaves the proxy and restores it in the response. @@ -164,67 +215,50 @@ class LLMShieldGuardrail(CustomGuardrail): # --- session ------------------------------------------------------------------ - def _session_id(self, data: MutableRequest) -> str: - """Returns a session id stable across this request's hooks.""" + @staticmethod + def _mint_session_id(data: MutableRequest) -> str: + """Mints a vault id for this request, overwriting anything already there. + + Redaction and restoration both happen inside one request/response pair, so + a fresh id per request is all that is needed, and it is what keeps one + caller from reaching another caller's vault. + """ + session_id: Final = f"{_VAULT_PREFIX}-{uuid.uuid4().hex}" metadata: Final = data.setdefault("metadata", {}) # mutable-ok: per-request store. - if not isinstance(metadata, dict): - return f"litellm-{uuid.uuid4().hex}" - existing: Final = metadata.get(_SESSION_METADATA_KEY) - if isinstance(existing, str) and existing: - return existing - session_id: Final = get_session_id_from_request_data(data) or f"litellm-{uuid.uuid4().hex}" - metadata[_SESSION_METADATA_KEY] = session_id + if isinstance(metadata, dict): + metadata[_SESSION_METADATA_KEY] = session_id return session_id + @staticmethod + def _session_id(data: MutableRequest) -> str: + """Reads back the vault id minted while redacting this request. + + Falls back to an unused id rather than to anything the caller supplied: a + reply that cannot be restored is a visible placeholder, while trusting a + caller-supplied id would hand them someone else's plaintext. + """ + metadata: Final = data.get("metadata") + existing: Final = metadata.get(_SESSION_METADATA_KEY) if isinstance(metadata, dict) else None + if isinstance(existing, str) and existing.startswith(_VAULT_PREFIX): + return existing + return f"{_VAULT_PREFIX}-{uuid.uuid4().hex}" + # --- request traversal -------------------------------------------------------- @staticmethod def _locate_request_texts(data: MutableRequest) -> Sequence[_Slot]: """Finds every redactable span in an outbound request. - Returns ``(text, write)`` pairs. Any shape missed here reaches the provider - in the clear, so this walks all of the request shapes that carry caller text: - - - chat ``messages``, both string and multimodal list ``content`` - - tool call ``arguments``, which routinely carry the values a user asked - the model to look up - - the Responses API ``input``, as a bare string or a list of items + Anything missed here reaches the provider in the clear while the guardrail + still reports as enabled, so the walk covers every request shape that + carries caller text. """ - slots: Final[list[_Slot]] = [] # mutable-ok: accumulator, frozen on return. - - def add(container: MutableRequest, key: str, value: object) -> None: - if isinstance(value, str) and value: - slots.append((value, lambda new, c=container, k=key: c.__setitem__(k, new))) - - def add_content(container: MutableRequest) -> None: - """Adds `content`, which is either a string or a list of typed parts.""" - content: Final = container.get("content") - if isinstance(content, str): - add(container, "content", content) - return - for part in content if isinstance(content, list) else (): - if isinstance(part, dict): - add(part, "text", part.get("text")) - - def add_tool_calls(message: MutableRequest) -> None: - for tool_call in message.get("tool_calls") or (): - function = tool_call.get("function") if isinstance(tool_call, dict) else None - if isinstance(function, dict): - add(function, "arguments", function.get("arguments")) - + slots: Final[_SlotSink] = [] # mutable-ok: accumulator, frozen on return. for message in data.get("messages") or (): if isinstance(message, dict): - add_content(message) - add_tool_calls(message) - - request_input: Final = data.get("input") - if isinstance(request_input, str): - add(data, "input", request_input) - else: - for item in request_input if isinstance(request_input, list) else (): - if isinstance(item, dict): - add_content(item) - + _collect_content(message, slots) + _collect_tool_arguments(message, slots) + _collect_responses_fields(data, slots) return tuple(slots) # --- hooks -------------------------------------------------------------------- @@ -245,7 +279,7 @@ class LLMShieldGuardrail(CustomGuardrail): if not slots: return data - redacted: Final = await self._redact(tuple(text for text, _ in slots), self._session_id(data)) + redacted: Final = await self._redact(tuple(text for text, _ in slots), self._mint_session_id(data)) for (_, write), replacement in zip(slots, redacted): write(replacement) return data @@ -451,11 +485,10 @@ class LLMShieldGuardrail(CustomGuardrail): if not texts: return inputs - session_id: Final = self._session_id(request_data) replaced: Final = ( - await self._redact(tuple(texts), session_id) + await self._redact(tuple(texts), self._mint_session_id(request_data)) if input_type == "request" - else await self._rehydrate(tuple(texts), session_id) + else await self._rehydrate(tuple(texts), self._session_id(request_data)) ) # Return a new mapping rather than rewriting the caller's, so this stays a # pure transform of the inputs it was handed. 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 index 7a07c9761a8..45d45301858 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py @@ -223,6 +223,35 @@ class TestRequestCoverage: assert data["messages"][0]["tool_calls"][0]["function"]["arguments"] == '{"email": "[EMAIL_1]"}' + @pytest.mark.asyncio + async def test_responses_api_instructions_are_redacted(self): + """`instructions` is provider-bound text that sits outside `messages`.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["contact [EMAIL_1]"]}) + + data = {"instructions": "contact jane.doe@example.com", "input": ""} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["contact jane.doe@example.com"] + assert data["instructions"] == "contact [EMAIL_1]" + + @pytest.mark.asyncio + async def test_legacy_function_call_arguments_are_redacted(self): + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ['{"email": "[EMAIL_1]"}']}) + + data = { + "messages": [ + { + "role": "assistant", + "function_call": {"name": "send", "arguments": '{"email": "jane.doe@example.com"}'}, + } + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert data["messages"][0]["function_call"]["arguments"] == '{"email": "[EMAIL_1]"}' + @pytest.mark.asyncio async def test_every_shape_in_one_request_is_redacted(self): guardrail = _guardrail() @@ -334,6 +363,59 @@ class TestRestoration: assert result["content"][1] == {"type": "tool_use", "id": "t1", "name": "lookup", "input": {}} +class TestVaultIsolation: + """The vault id must never be something a caller can choose. + + The vault holds the plaintext behind every placeholder. If a caller could name + the vault, they could send a placeholder, have the model echo it back, and get + another caller's value restored into their own reply. + """ + + @pytest.mark.asyncio + async def test_caller_supplied_session_id_is_not_used(self): + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + + data = { + "messages": [{"role": "user", "content": "a@b.com"}], + "metadata": {"llm_shield_session_id": "victim-session"}, + "litellm_session_id": "victim-session", + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + used = mock.call_args_list[0].kwargs["headers"]["X-Session-ID"] + assert used != "victim-session" + assert data["metadata"]["llm_shield_session_id"] == used + + @pytest.mark.asyncio + async def test_restore_ignores_a_foreign_session_id(self): + """A reply is left unrestored rather than resolved against another vault.""" + guardrail = _guardrail(event_hook="post_call") + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + + data = {"metadata": {"llm_shield_session_id": "victim-session"}} + response = ModelResponse(choices=[Choices(index=0, message=Message(role="assistant", content="[EMAIL_1]"))]) + await guardrail.async_post_call_success_hook(data=data, user_api_key_dict=None, response=response) + + assert mock.call_args_list[0].kwargs["headers"]["X-Session-ID"] != "victim-session" + + @pytest.mark.asyncio + async def test_each_request_gets_its_own_vault(self): + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}, {"texts": ["[EMAIL_1]"]}) + + for _ in range(2): + await guardrail.async_pre_call_hook( + user_api_key_dict=None, + cache=None, + data={"messages": [{"role": "user", "content": "a@b.com"}]}, + call_type="completion", + ) + + seen = {call.kwargs["headers"]["X-Session-ID"] for call in mock.call_args_list} + assert len(seen) == 2 + + class TestFailClosed: @pytest.mark.asyncio async def test_unreachable_shield_blocks_the_request(self): From f3eb108f86a5b02063f35a6d2f50d69ceb57cf91 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Thu, 3 Sep 2026 19:41:46 -0500 Subject: [PATCH 10/34] fix(guardrails): drop Final from a loop-assigned local basedpyright rejects a Final assigned inside a loop, and reportGeneralTypeIssues sits one over its budget ceiling. --- .../proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py index 26c6774f48d..172de67020b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py @@ -99,7 +99,7 @@ def _collect_content(container: MutableRequest, slots: _SlotSink) -> None: def _collect_tool_arguments(message: MutableRequest, slots: _SlotSink) -> None: """Tool arguments carry the values a user asked the model to act on.""" for tool_call in message.get("tool_calls") or (): - function: Final = tool_call.get("function") if isinstance(tool_call, dict) else None + function = tool_call.get("function") if isinstance(tool_call, dict) else None # rebind-ok: loop variable. if isinstance(function, dict): _collect(function, "arguments", slots) legacy: Final = message.get("function_call") From 8d0b1881d7b028e1eb6b09c1bb76f1527923768f Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Thu, 3 Sep 2026 19:59:36 -0500 Subject: [PATCH 11/34] fix(guardrails): redact completion prompts and responses tool items Two more provider-bound request shapes were reaching the model intact while the guardrail reported as enabled. /v1/completions carries its text in a top-level `prompt`, which the traversal never looked at. It is handled as a string and as the array form, where each entry is rewritten in place. Responses input items hold tool data outside `content`: a function_call item in `arguments`, a function_call_output item in `output`. Both are now collected alongside the item's content. Adds a test per shape. --- .../guardrail_hooks/llm_shield/llm_shield.py | 30 +++++++++++++- .../guardrail_hooks/test_llm_shield.py | 40 +++++++++++++++++++ 2 files changed, 68 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py index 172de67020b..8989fabc321 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py @@ -77,6 +77,10 @@ _Slot: TypeAlias = tuple[str, Callable[[str], None]] # mutable-ok: Callable's p # _locate_request_texts, which freezes it into a tuple before returning. _SlotSink: TypeAlias = list[_Slot] # mutable-ok: accumulator passed between collectors. +# A caller-owned list whose entries are rewritten in place, such as a Completions +# `prompt` sent as an array of strings. +MutableSeq: TypeAlias = list # mutable-ok: the request payload's own list. + def _collect(container: MutableRequest, key: str, slots: _SlotSink) -> None: """Records the string at `key`, along with the write that replaces it.""" @@ -85,6 +89,23 @@ def _collect(container: MutableRequest, key: str, slots: _SlotSink) -> None: slots.append((value, lambda new, c=container, k=key: c.__setitem__(k, new))) +def _collect_entry(entries: MutableSeq, index: int, slots: _SlotSink) -> None: + """Records a string held directly in a list, rather than under a key.""" + value: Final = entries[index] + if isinstance(value, str) and value: + slots.append((value, lambda new, e=entries, i=index: e.__setitem__(i, new))) + + +def _collect_prompt(data: MutableRequest, slots: _SlotSink) -> None: + """The Completions API sends its text in a top-level `prompt`.""" + prompt: Final = data.get("prompt") + if isinstance(prompt, str): + _collect(data, "prompt", slots) + return + for index in range(len(prompt)) if isinstance(prompt, list) else (): + _collect_entry(prompt, index, slots) + + def _collect_content(container: MutableRequest, slots: _SlotSink) -> None: """`content` is either a string or the multimodal list of typed parts.""" content: Final = container.get("content") @@ -115,8 +136,12 @@ def _collect_responses_fields(data: MutableRequest, slots: _SlotSink) -> None: _collect(data, "input", slots) return for item in request_input if isinstance(request_input, list) else (): - if isinstance(item, dict): - _collect_content(item, slots) + if not isinstance(item, dict): + continue + _collect_content(item, slots) + # A function_call item holds `arguments`; a function_call_output holds `output`. + _collect(item, "arguments", slots) + _collect(item, "output", slots) class LLMShieldGuardrail(CustomGuardrail): @@ -259,6 +284,7 @@ class LLMShieldGuardrail(CustomGuardrail): _collect_content(message, slots) _collect_tool_arguments(message, slots) _collect_responses_fields(data, slots) + _collect_prompt(data, slots) return tuple(slots) # --- hooks -------------------------------------------------------------------- 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 index 45d45301858..ffd6ede28b6 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py @@ -252,6 +252,46 @@ class TestRequestCoverage: assert data["messages"][0]["function_call"]["arguments"] == '{"email": "[EMAIL_1]"}' + @pytest.mark.asyncio + async def test_completions_prompt_is_redacted(self): + """/v1/completions puts its text in a top-level `prompt`, not in messages.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["Email [EMAIL_1]"]}) + + data = {"prompt": "Email jane.doe@example.com"} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="atext_completion") + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["Email jane.doe@example.com"] + assert data["prompt"] == "Email [EMAIL_1]" + + @pytest.mark.asyncio + async def test_completions_prompt_array_is_redacted(self): + """`prompt` also accepts an array, and each entry is provider-bound.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]", "[PHONE_1]"]}) + + data = {"prompt": ["jane.doe@example.com", "555-0100"]} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="atext_completion") + + assert data["prompt"] == ["[EMAIL_1]", "[PHONE_1]"] + + @pytest.mark.asyncio + async def test_responses_function_call_items_are_redacted(self): + """Responses input items hold tool data in `arguments` and `output`.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ['{"email": "[EMAIL_1]"}', "sent to [EMAIL_1]"]}) + + data = { + "input": [ + {"type": "function_call", "name": "send", "arguments": '{"email": "jane.doe@example.com"}'}, + {"type": "function_call_output", "call_id": "c1", "output": "sent to jane.doe@example.com"}, + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert data["input"][0]["arguments"] == '{"email": "[EMAIL_1]"}' + assert data["input"][1]["output"] == "sent to [EMAIL_1]" + @pytest.mark.asyncio async def test_every_shape_in_one_request_is_redacted(self): guardrail = _guardrail() From 46f13807a513801e107592f1eec79ea24c47f1a6 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Thu, 3 Sep 2026 20:13:45 -0500 Subject: [PATCH 12/34] fix(guardrails): redact the anthropic system prompt and string-array input Two more provider-bound shapes, found by walking the request types rather than waiting for them to be reported. /v1/messages carries its system prompt at the top level, as a string or a list of text blocks. It is one of the endpoints this guardrail claims to cover, and a system prompt is a natural place to put a customer's details. `input` as an array of bare strings, the embeddings and moderations shape, was skipped because the loop only handled item dicts. Verified against a live provider: a system prompt holding an address now reaches the model as a stand-in and is restored in the reply. --- .../guardrail_hooks/llm_shield/llm_shield.py | 18 ++++++++- .../guardrail_hooks/test_llm_shield.py | 38 +++++++++++++++++++ 2 files changed, 55 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py index 8989fabc321..e495f2e99d4 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py @@ -128,6 +128,17 @@ def _collect_tool_arguments(message: MutableRequest, slots: _SlotSink) -> None: _collect(legacy, "arguments", slots) +def _collect_system(data: MutableRequest, slots: _SlotSink) -> None: + """Anthropic's /v1/messages carries its system prompt at the top level.""" + system: Final = data.get("system") + if isinstance(system, str): + _collect(data, "system", slots) + return + for part in system if isinstance(system, list) else (): + if isinstance(part, dict): + _collect(part, "text", slots) + + def _collect_responses_fields(data: MutableRequest, slots: _SlotSink) -> None: """The Responses API sends text outside `messages`, in `instructions` and `input`.""" _collect(data, "instructions", slots) @@ -135,7 +146,11 @@ def _collect_responses_fields(data: MutableRequest, slots: _SlotSink) -> None: if isinstance(request_input, str): _collect(data, "input", slots) return - for item in request_input if isinstance(request_input, list) else (): + for index, item in enumerate(request_input if isinstance(request_input, list) else ()): + if isinstance(item, str): + # The embeddings and moderations shape: `input` as an array of strings. + _collect_entry(request_input, index, slots) + continue if not isinstance(item, dict): continue _collect_content(item, slots) @@ -285,6 +300,7 @@ class LLMShieldGuardrail(CustomGuardrail): _collect_tool_arguments(message, slots) _collect_responses_fields(data, slots) _collect_prompt(data, slots) + _collect_system(data, slots) return tuple(slots) # --- hooks -------------------------------------------------------------------- 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 index ffd6ede28b6..f16584fbd1d 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py @@ -292,6 +292,44 @@ class TestRequestCoverage: assert data["input"][0]["arguments"] == '{"email": "[EMAIL_1]"}' assert data["input"][1]["output"] == "sent to [EMAIL_1]" + @pytest.mark.asyncio + async def test_anthropic_system_prompt_is_redacted(self): + """/v1/messages carries its system prompt at the top level, not in messages.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["the user is [EMAIL_1]"]}) + + data = {"system": "the user is jane.doe@example.com", "messages": []} + await guardrail.async_pre_call_hook( + user_api_key_dict=None, cache=None, data=data, call_type="anthropic_messages" + ) + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["the user is jane.doe@example.com"] + assert data["system"] == "the user is [EMAIL_1]" + + @pytest.mark.asyncio + async def test_anthropic_system_blocks_are_redacted(self): + """`system` also accepts a list of text blocks.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + + data = {"system": [{"type": "text", "text": "jane.doe@example.com"}], "messages": []} + await guardrail.async_pre_call_hook( + user_api_key_dict=None, cache=None, data=data, call_type="anthropic_messages" + ) + + assert data["system"][0]["text"] == "[EMAIL_1]" + + @pytest.mark.asyncio + async def test_string_array_input_is_redacted(self): + """Embeddings and moderations send `input` as an array of bare strings.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]", "[PHONE_1]"]}) + + data = {"input": ["jane.doe@example.com", "555-0100"]} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aembedding") + + assert data["input"] == ["[EMAIL_1]", "[PHONE_1]"] + @pytest.mark.asyncio async def test_every_shape_in_one_request_is_redacted(self): guardrail = _guardrail() From a35c5028196c07880ae7000f5c99dcf40d069c59 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Thu, 3 Sep 2026 20:26:57 -0500 Subject: [PATCH 13/34] fix(guardrails): narrow prompt and input to a list before iterating Guarding with a conditional iterable left the value un-narrowed, so passing it on was an argument-type error and the element checks read as unreachable. An early return narrows it properly and reads better. --- .../guardrails/guardrail_hooks/llm_shield/llm_shield.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py index e495f2e99d4..39c0930d7cf 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py @@ -102,7 +102,9 @@ def _collect_prompt(data: MutableRequest, slots: _SlotSink) -> None: if isinstance(prompt, str): _collect(data, "prompt", slots) return - for index in range(len(prompt)) if isinstance(prompt, list) else (): + if not isinstance(prompt, list): + return + for index in range(len(prompt)): _collect_entry(prompt, index, slots) @@ -146,7 +148,9 @@ def _collect_responses_fields(data: MutableRequest, slots: _SlotSink) -> None: if isinstance(request_input, str): _collect(data, "input", slots) return - for index, item in enumerate(request_input if isinstance(request_input, list) else ()): + if not isinstance(request_input, list): + return + for index, item in enumerate(request_input): if isinstance(item, str): # The embeddings and moderations shape: `input` as an array of strings. _collect_entry(request_input, index, slots) From 46438d7cf74913f8f6694f5f1adf1f48585d2d83 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Thu, 3 Sep 2026 20:48:50 -0500 Subject: [PATCH 14/34] fix(guardrails): restore every streaming choice, not just the first Streaming rehydration read and rewrote choices[0] only, so with n>1 every later choice went back to the caller still holding its placeholders. Each choice is its own token stream, so the sliding window is now tracked per choice index rather than once per stream. A single shared window would have been worse than the bug: it would splice the characters held back for one choice onto the next one's delta. The final flush walks every choice the same way, and the two helpers that only ever looked at choices[0] are gone. Adds a test that both choices come back restored, and one that each choice gets its own window handed back rather than its neighbour's. --- .../guardrail_hooks/llm_shield/llm_shield.py | 105 ++++++++++-------- .../guardrail_hooks/test_llm_shield.py | 67 +++++++++++ 2 files changed, 128 insertions(+), 44 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py index 39c0930d7cf..f84f0f429aa 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py @@ -77,6 +77,9 @@ _Slot: TypeAlias = tuple[str, Callable[[str], None]] # mutable-ok: Callable's p # _locate_request_texts, which freezes it into a tuple before returning. _SlotSink: TypeAlias = list[_Slot] # mutable-ok: accumulator passed between collectors. +# Sliding windows keyed by streaming choice index, threaded through one stream. +_CarryWindows: TypeAlias = dict # mutable-ok: per-choice windows advanced in place. + # A caller-owned list whose entries are rewritten in place, such as a Completions # `prompt` sent as an array of strings. MutableSeq: TypeAlias = list # mutable-ok: the request payload's own list. @@ -163,6 +166,12 @@ def _collect_responses_fields(data: MutableRequest, slots: _SlotSink) -> None: _collect(item, "output", slots) +def _choice_index(choice: object) -> int: + """Streaming choices are matched across chunks by their index.""" + index: Final = getattr(choice, "index", 0) + return index if isinstance(index, int) else 0 + + class LLMShieldGuardrail(CustomGuardrail): """Redacts PII before it leaves the proxy and restores it in the response. @@ -441,10 +450,10 @@ class LLMShieldGuardrail(CustomGuardrail): ) -> 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. + Each choice is its own token stream, so the sliding window is tracked per + choice index. One shared window would splice the characters held back for + one choice onto the next. The windows are locals of this generator, so they + are scoped to a single stream and cannot leak between concurrent requests. """ if self.should_run_guardrail(data=request_data, event_type=GuardrailEventHooks.post_call) is not True: async for chunk in response: @@ -452,39 +461,61 @@ class LLMShieldGuardrail(CustomGuardrail): return session_id: Final = self._session_id(request_data) - carry = "" # rebind-ok: the sliding window advances with every delta. + carries: Final[dict] = {} # mutable-ok: per-choice windows, local to this stream. last_chunk = None # rebind-ok: tracks the most recent chunk for the final flush. 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 delta is None or 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 and delta is not None: - emitted, carry = await self._stream_step("", carry, True, session_id) - if emitted: - delta.content = emitted - yield chunk - continue - - emitted, carry = await self._stream_step(text, carry, is_final, session_id) - delta.content = emitted + for choice in getattr(chunk, "choices", None) or (): + await self._restore_choice(choice, carries, session_id) yield chunk # A stream that ended without a finish_reason can still leave text held back. - if carry and last_chunk is not None: - flushed: Final = await self._stream_step("", carry, True, session_id) - trailing_text, carry = flushed # rebind-ok: window advances. - if trailing_text: - trailing: Final = last_chunk.model_copy(deep=True) - trailing_delta: Final = self._stream_delta(trailing) - if trailing_delta is not None: - trailing_delta.content = trailing_text - yield trailing + if last_chunk is not None and any(carries.values()): + trailing: Final = last_chunk.model_copy(deep=True) + if await self._flush_trailing(trailing, carries, session_id): + yield trailing + + async def _restore_choice(self, choice: Any, carries: _CarryWindows, session_id: str) -> None: + """Restores one choice's delta, advancing that choice's own window.""" + delta: Final = getattr(choice, "delta", None) + if delta is None: + return + index: Final = _choice_index(choice) + carry: Final = carries.get(index, "") + text: Final = getattr(delta, "content", None) + is_final: Final = bool(getattr(choice, "finish_reason", None)) + + if not isinstance(text, str) or not text: + # Nothing to restore here, but a final chunk still has to flush the window. + if is_final and carry: + flushed, flushed_carry = await self._stream_step("", carry, True, session_id) + carries[index] = flushed_carry # rebind-ok: this choice's window advances. + if flushed: + delta.content = flushed + return + + emitted, remaining = await self._stream_step(text, carry, is_final, session_id) + carries[index] = remaining # rebind-ok: this choice's window advances. + delta.content = emitted + + async def _flush_trailing(self, trailing: Any, carries: _CarryWindows, session_id: str) -> bool: + """Empties every still-held window into a copy of the last chunk.""" + emitted_any = False # rebind-ok: set once any choice contributes text. + for choice in getattr(trailing, "choices", None) or (): + delta = getattr(choice, "delta", None) + if delta is None: + continue + index = _choice_index(choice) + carry = carries.get(index, "") + if not carry: + delta.content = None + continue + text, remaining = await self._stream_step("", carry, True, session_id) + carries[index] = remaining # rebind-ok: this choice's window advances. + delta.content = text or None + emitted_any = emitted_any or bool(text) + return emitted_any async def _stream_step(self, text: str, carry: str, final: bool, session_id: str) -> tuple[str, str]: """Returns ``(text safe to emit now, window still being held)``.""" @@ -503,20 +534,6 @@ class LLMShieldGuardrail(CustomGuardrail): ) return emitted, remaining - @staticmethod - def _stream_delta(chunk: object) -> Any: - choices: Final = getattr(chunk, "choices", None) - if not choices: - return None - return getattr(choices[0], "delta", None) - - @staticmethod - def _is_final_chunk(chunk: object) -> bool: - choices: Final = 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 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 index f16584fbd1d..e9cc1a36b07 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py @@ -587,6 +587,73 @@ class TestStreamingRehydration: 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_every_choice_is_restored(self): + """With n>1 a later choice must not be handed back still holding a placeholder.""" + guardrail = _guardrail(event_hook="post_call") + _mock_post( + guardrail, + {"text": "first@example.com", "carry": ""}, + {"text": "second@example.com", "carry": ""}, + ) + + async def stream(): + yield ModelResponseStream( + choices=[ + StreamingChoices(index=0, delta=Delta(content="[EMAIL_1]"), finish_reason="stop"), + StreamingChoices(index=1, delta=Delta(content="[EMAIL_2]"), finish_reason="stop"), + ] + ) + + chunks = await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + restored = [choice.delta.content for choice in chunks[0].choices] + assert restored == ["first@example.com", "second@example.com"] + + @pytest.mark.asyncio + async def test_choice_windows_do_not_cross_contaminate(self): + """Each choice is its own token stream, so each carries its own window. + + One shared window would send the characters held back for choice 0 up + against choice 1's next delta and splice the two streams together. + """ + guardrail = _guardrail(event_hook="post_call") + mock = _mock_post( + guardrail, + {"text": "", "carry": "A-held"}, + {"text": "", "carry": "B-held"}, + {"text": "a-done", "carry": ""}, + {"text": "b-done", "carry": ""}, + ) + + async def stream(): + yield ModelResponseStream( + choices=[ + StreamingChoices(index=0, delta=Delta(content="a1")), + StreamingChoices(index=1, delta=Delta(content="b1")), + ] + ) + yield ModelResponseStream( + choices=[ + StreamingChoices(index=0, delta=Delta(content="a2"), finish_reason="stop"), + StreamingChoices(index=1, delta=Delta(content="b2"), finish_reason="stop"), + ] + ) + + await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + sent = [call.kwargs["json"] for call in mock.call_args_list] + assert sent[2]["carry"] == "A-held", "choice 0 must get its own window back" + assert sent[3]["carry"] == "B-held", "choice 1 must get its own window back" + @pytest.mark.asyncio async def test_chunks_are_forwarded_as_they_arrive(self): """Restoration must not buffer the stream into a single terminal chunk.""" From a0abb9a4992489b91deea3cf69ab3b354314c9fb Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Thu, 3 Sep 2026 21:09:48 -0500 Subject: [PATCH 15/34] refactor(guardrails): name the guardrail llm_shield_proxy throughout The integration was called llm_shield in code, llm-shield in the example config, and LLM Shield in the dashboard, while the product and its PyPI package are both llm-shield-proxy. An operator who saw the guardrail in LiteLLM could not tell what to install. One identifier now: llm_shield_proxy for the enum value, module, directory, class, config model, logo and environment variables, with LLM Shield Proxy as the display name. That matches `pip install llm-shield-proxy`. Renames only; no behaviour change. --- .../__init__.py | 10 +++---- .../example_config.yaml | 24 ++++++++-------- .../llm_shield_proxy.py} | 25 +++++++++-------- litellm/types/guardrails.py | 2 +- .../{llm_shield.py => llm_shield_proxy.py} | 10 +++---- ruff-strict.toml | 2 +- ...llm_shield.py => test_llm_shield_proxy.py} | 28 +++++++++---------- .../{llm_shield.svg => llm_shield_proxy.svg} | 0 .../_components/guardrail_garden_configs.ts | 6 ++-- .../_components/guardrail_garden_data.test.ts | 2 +- .../_components/guardrail_garden_data.ts | 8 +++--- .../_components/guardrail_info_helpers.tsx | 6 ++-- 12 files changed, 62 insertions(+), 61 deletions(-) rename litellm/proxy/guardrails/guardrail_hooks/{llm_shield => llm_shield_proxy}/__init__.py (72%) rename litellm/proxy/guardrails/guardrail_hooks/{llm_shield => llm_shield_proxy}/example_config.yaml (70%) rename litellm/proxy/guardrails/guardrail_hooks/{llm_shield/llm_shield.py => llm_shield_proxy/llm_shield_proxy.py} (95%) rename litellm/types/proxy/guardrails/guardrail_hooks/{llm_shield.py => llm_shield_proxy.py} (51%) rename tests/test_litellm/proxy/guardrails/guardrail_hooks/{test_llm_shield.py => test_llm_shield_proxy.py} (96%) rename ui/litellm-dashboard/public/assets/logos/{llm_shield.svg => llm_shield_proxy.svg} (100%) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/__init__.py similarity index 72% rename from litellm/proxy/guardrails/guardrail_hooks/llm_shield/__init__.py rename to litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/__init__.py index a6cc54d5408..c8ca68a8967 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/__init__.py @@ -2,16 +2,16 @@ from typing import TYPE_CHECKING, Final from litellm.types.guardrails import SupportedGuardrailIntegrations -from .llm_shield import LLMShieldGuardrail +from .llm_shield_proxy import LLMShieldProxyGuardrail if TYPE_CHECKING: from litellm.types.guardrails import Guardrail, LitellmParams -def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"): +def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail") -> LLMShieldProxyGuardrail: import litellm - _llm_shield_guardrail_callback: Final = LLMShieldGuardrail( + _llm_shield_guardrail_callback: Final = LLMShieldProxyGuardrail( api_key=litellm_params.api_key, api_base=litellm_params.api_base, guardrail_name=guardrail.get("guardrail_name", ""), @@ -24,10 +24,10 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" guardrail_initializer_registry: Final = { # mutable-ok: module-level registry, built once and never mutated - SupportedGuardrailIntegrations.LLM_SHIELD.value: initialize_guardrail, + SupportedGuardrailIntegrations.LLM_SHIELD_PROXY.value: initialize_guardrail, } guardrail_class_registry: Final = { # mutable-ok: module-level registry, built once and never mutated - SupportedGuardrailIntegrations.LLM_SHIELD.value: LLMShieldGuardrail, + SupportedGuardrailIntegrations.LLM_SHIELD_PROXY.value: LLMShieldProxyGuardrail, } diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/example_config.yaml b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/example_config.yaml similarity index 70% rename from litellm/proxy/guardrails/guardrail_hooks/llm_shield/example_config.yaml rename to litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/example_config.yaml index aa63fa9d252..b5c9f0f8b69 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/example_config.yaml +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/example_config.yaml @@ -1,7 +1,7 @@ -# Example LiteLLM Proxy configuration for LLM Shield -# LLM Shield is a self-hosted PII gateway: https://github.com/ninadphalak/LLM-Shield-Proxy +# Example LiteLLM Proxy configuration for LLM Shield Proxy +# LLM Shield Proxy is a self-hosted PII gateway: https://github.com/ninadphalak/LLM-Shield-Proxy # -# Unlike a masking guardrail, LLM Shield's substitution is reversible. Personal data is +# Unlike a masking guardrail, LLM Shield Proxy's substitution is reversible. Personal data is # replaced with placeholders before the request goes to the provider, and the original # values are put back into the model's reply, so the end user still sees real data while # the provider never received it. @@ -15,25 +15,25 @@ model_list: guardrails: # Both modes belong on ONE entry. pre_call redacts the outbound request and post_call # restores the reply; listing only pre_call would send placeholders back to the user. - - guardrail_name: "llm-shield" + - guardrail_name: "llm_shield_proxy" litellm_params: - guardrail: llm_shield + guardrail: llm_shield_proxy mode: ["pre_call", "post_call"] default_on: true - # Your own LLM Shield deployment. Defaults to http://localhost:8000, and also reads - # LLM_SHIELD_API_BASE from the environment. + # Your own LLM Shield Proxy deployment. Defaults to http://localhost:8000, and also reads + # LLM_SHIELD_PROXY_API_BASE from the environment. api_base: "http://localhost:8000" - # A virtual key configured on that deployment. Also reads LLM_SHIELD_API_KEY. - api_key: os.environ/LLM_SHIELD_API_KEY + # A virtual key configured on that deployment. Also reads LLM_SHIELD_PROXY_API_KEY. + api_key: os.environ/LLM_SHIELD_PROXY_API_KEY # Usage: # -# 1. Run LLM Shield somewhere the proxy can reach: +# 1. Run LLM Shield Proxy somewhere the proxy can reach: # pip install llm-shield-proxy # llm-shield-proxy --port 8000 # # 2. Point this config at it and start the proxy: -# export LLM_SHIELD_API_KEY="your-virtual-key" +# export LLM_SHIELD_PROXY_API_KEY="your-virtual-key" # litellm --config example_config.yaml # # 3. Send a request containing personal data: @@ -47,7 +47,7 @@ guardrails: # # Notes: # -# - Requests are refused if LLM Shield is unreachable or returns an error, rather than +# - Requests are refused if LLM Shield Proxy is unreachable or returns an error, rather than # being forwarded. Sending them on would hand the provider exactly the data this # guardrail exists to withhold. # - Restoring a value requires the request and the reply to share a session. LiteLLM's diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py similarity index 95% rename from litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py rename to litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index f84f0f429aa..56bee894714 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -1,6 +1,6 @@ # +-------------------------------------------------------------+ # -# Use LLM Shield for reversible PII redaction +# Use LLM Shield Proxy for reversible PII redaction # https://github.com/ninadphalak/LLM-Shield-Proxy # # +-------------------------------------------------------------+ @@ -38,7 +38,7 @@ 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" +GUARDRAIL_NAME: Final = "llm_shield_proxy" _DEFAULT_API_BASE: Final = "http://localhost:8000" _REDACT_PATH: Final = "/v1/guard/redact" @@ -172,7 +172,7 @@ def _choice_index(choice: object) -> int: return index if isinstance(index, int) else 0 -class LLMShieldGuardrail(CustomGuardrail): +class LLMShieldProxyGuardrail(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 @@ -199,8 +199,9 @@ class LLMShieldGuardrail(CustomGuardrail): **kwargs: Any, # noqa: LIT008 # kwargs-ok: forwarded verbatim to CustomGuardrail.__init__ ) -> 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") + env_base: Final = os.environ.get("LLM_SHIELD_PROXY_API_BASE") + self.api_base: Final = (api_base or env_base or _DEFAULT_API_BASE).rstrip("/") + self.api_key: Final = api_key or os.environ.get("LLM_SHIELD_PROXY_API_KEY") super().__init__(guardrail_name=guardrail_name, **kwargs) @classmethod @@ -219,7 +220,7 @@ class LLMShieldGuardrail(CustomGuardrail): return headers async def _call_shield(self, path: str, session_id: str, payload: JsonBody) -> Mapping[str, object]: - """Posts to LLM Shield, failing closed on any transport or status error. + """Posts to LLM Shield Proxy, 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 @@ -235,16 +236,16 @@ class LLMShieldGuardrail(CustomGuardrail): 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) + verbose_proxy_logger.exception("LLM Shield Proxy 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.", + message=f"LLM Shield Proxy 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) + verbose_proxy_logger.exception("LLM Shield Proxy call to %s failed", path) raise GuardrailRaisedException( guardrail_name=self.guardrail_name, - message="LLM Shield is unreachable; blocking the request.", + message="LLM Shield Proxy is unreachable; blocking the request.", ) from exc async def _redact(self, texts: Sequence[str], session_id: str) -> Sequence[str]: @@ -262,7 +263,7 @@ class LLMShieldGuardrail(CustomGuardrail): 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.", + message=f"LLM Shield Proxy {operation} returned an unexpected payload; blocking the request.", ) return tuple(returned) @@ -530,7 +531,7 @@ class LLMShieldGuardrail(CustomGuardrail): 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.", + message="LLM Shield Proxy stream rehydration returned an unexpected payload.", ) return emitted, remaining diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 42ab6034069..41b33433d92 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -137,7 +137,7 @@ class SupportedGuardrailIntegrations(Enum): COMPRESR = "compresr" STRAIKER = "straiker" ALICE = "alice" - LLM_SHIELD = "llm_shield" + LLM_SHIELD_PROXY = "llm_shield_proxy" class Role(Enum): diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/llm_shield.py b/litellm/types/proxy/guardrails/guardrail_hooks/llm_shield_proxy.py similarity index 51% rename from litellm/types/proxy/guardrails/guardrail_hooks/llm_shield.py rename to litellm/types/proxy/guardrails/guardrail_hooks/llm_shield_proxy.py index 8d7afd907b0..967d7ee75c2 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/llm_shield.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/llm_shield_proxy.py @@ -3,22 +3,22 @@ from pydantic import Field from .base import GuardrailConfigModel -class LLMShieldGuardrailConfigModel(GuardrailConfigModel): +class LLMShieldProxyGuardrailConfigModel(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." + "The virtual key for the LLM Shield Proxy instance. If not provided, the " + "`LLM_SHIELD_PROXY_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` " + "The base URL of the LLM Shield Proxy instance. If not provided, the `LLM_SHIELD_PROXY_API_BASE` " "environment variable is checked, then `http://localhost:8000`." ), ) @staticmethod def ui_friendly_name() -> str: - return "LLM Shield" + return "LLM Shield Proxy" diff --git a/ruff-strict.toml b/ruff-strict.toml index f9d026c6011..21595556d28 100644 --- a/ruff-strict.toml +++ b/ruff-strict.toml @@ -33,7 +33,7 @@ external = [ # Same reason: `**kwargs` forwards verbatim to CustomGuardrail.__init__, and the lifecycle # hook signatures inherit `Any` for `response` from CustomLogger, so narrowing them here # would break the override rather than describe it. -"litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py" = ["ANN401"] +"litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py" = ["ANN401"] [lint.mccabe] max-complexity = 15 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_proxy.py similarity index 96% rename from tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py rename to tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py index e9cc1a36b07..bc74b42a0e7 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -6,16 +6,16 @@ from httpx import Request, Response import litellm from litellm.exceptions import GuardrailRaisedException -from litellm.proxy.guardrails.guardrail_hooks.llm_shield.llm_shield import ( +from litellm.proxy.guardrails.guardrail_hooks.llm_shield_proxy.llm_shield_proxy import ( GUARDRAIL_NAME, - LLMShieldGuardrail, + LLMShieldProxyGuardrail, ) from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2 from litellm.types.guardrails import GuardrailEventHooks from litellm.types.utils import Choices, Delta, Message, ModelResponse, ModelResponseStream, StreamingChoices -def _guardrail(**overrides: object) -> LLMShieldGuardrail: +def _guardrail(**overrides: object) -> LLMShieldProxyGuardrail: params: dict[str, object] = { "api_key": "test-key", "api_base": "http://shield.test", @@ -24,7 +24,7 @@ def _guardrail(**overrides: object) -> LLMShieldGuardrail: "default_on": True, } params.update(overrides) - return LLMShieldGuardrail(**params) + return LLMShieldProxyGuardrail(**params) def _response(payload: dict, status_code: int = 200) -> Response: @@ -35,7 +35,7 @@ def _response(payload: dict, status_code: int = 200) -> Response: ) -def _mock_post(guardrail: LLMShieldGuardrail, *payloads: dict) -> AsyncMock: +def _mock_post(guardrail: LLMShieldProxyGuardrail, *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] @@ -55,30 +55,30 @@ async def _drain(generator) -> list: 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") + monkeypatch.setenv("LLM_SHIELD_PROXY_API_KEY", "test-key") init_guardrails_v2( all_guardrails=[ { - "guardrail_name": "llm_shield", - "litellm_params": {"guardrail": "llm_shield", "mode": "pre_call", "default_on": True}, + "guardrail_name": "llm_shield_proxy", + "litellm_params": {"guardrail": "llm_shield_proxy", "mode": "pre_call", "default_on": True}, } ], config_file_path="", ) - registered = [cb for cb in litellm.callbacks if isinstance(cb, LLMShieldGuardrail)] + registered = [cb for cb in litellm.callbacks if isinstance(cb, LLMShieldProxyGuardrail)] assert len(registered) == 1 - assert registered[0].guardrail_name == "llm_shield" + assert registered[0].guardrail_name == "llm_shield_proxy" -class TestLLMShieldInitialization: +class TestLLMShieldProxyInitialization: def test_api_base_defaults_to_localhost(self, monkeypatch: pytest.MonkeyPatch): - monkeypatch.delenv("LLM_SHIELD_API_BASE", raising=False) + monkeypatch.delenv("LLM_SHIELD_PROXY_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") + monkeypatch.setenv("LLM_SHIELD_PROXY_API_BASE", "http://shield.internal:9000") assert _guardrail(api_base=None).api_base == "http://shield.internal:9000" def test_trailing_slash_is_stripped(self): @@ -135,7 +135,7 @@ class TestRedaction: @pytest.mark.asyncio async def test_request_without_text_is_untouched(self): - """No text to redact means no call to LLM Shield. + """No text to redact means no call to LLM Shield Proxy. This deliberately uses a request with no caller text at all. An earlier version used a Responses-API `input`, which asserted the very bypass that diff --git a/ui/litellm-dashboard/public/assets/logos/llm_shield.svg b/ui/litellm-dashboard/public/assets/logos/llm_shield_proxy.svg similarity index 100% rename from ui/litellm-dashboard/public/assets/logos/llm_shield.svg rename to ui/litellm-dashboard/public/assets/logos/llm_shield_proxy.svg diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts index 3579457bcd5..38751cb1d43 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts @@ -320,9 +320,9 @@ export const GUARDRAIL_PRESETS: Record = { mode: "pre_call", defaultOn: false, }, - llm_shield: { - provider: "LLM Shield", - guardrailNameSuggestion: "LLM Shield", + llm_shield_proxy: { + provider: "LLM Shield Proxy", + guardrailNameSuggestion: "LLM Shield Proxy", // Both halves are required. With only pre_call the request is redacted and the // placeholders are handed straight back to the caller. mode: ["pre_call", "post_call"], diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts index d3212d66737..293c447604f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts @@ -28,7 +28,7 @@ const EXPECTED_PARTNER_LOGO_FILES: Record = { repelloai: "repelloai.png", straiker: "straiker.svg", alice: "alice.svg", - llm_shield: "llm_shield.svg", + llm_shield_proxy: "llm_shield_proxy.svg", }; describe("guardrail_garden_data logos", () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts index d46847a5a80..74a13f2f611 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts @@ -475,14 +475,14 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ providerKey: "Alice", }, { - id: "llm_shield", - name: "LLM Shield", + id: "llm_shield_proxy", + name: "LLM Shield Proxy", description: "Self-hosted PII redaction that puts the original values back into the model's response, so the provider never receives personal data while the end user still sees it.", category: "partner", - logo: guardrailLogoMap["LLM Shield"], + logo: guardrailLogoMap["LLM Shield Proxy"], tags: ["PII", "Data Privacy", "Compliance", "Streaming"], - providerKey: "LLM Shield", + providerKey: "LLM Shield Proxy", }, ]; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx index f620ea6dcd0..1bf7056bd5c 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx @@ -1,7 +1,7 @@ import aimSecurityLogo from "../../../../../public/assets/logos/aim_security.jpeg"; import aktoLogo from "../../../../../public/assets/logos/akto.svg"; import aliceLogo from "../../../../../public/assets/logos/alice.svg"; -import llmShieldLogo from "../../../../../public/assets/logos/llm_shield.svg"; +import llmShieldProxyLogo from "../../../../../public/assets/logos/llm_shield_proxy.svg"; import aporiaLogo from "../../../../../public/assets/logos/aporia.png"; import bedrockLogo from "../../../../../public/assets/logos/bedrock.svg"; import catoNetworksLogo from "../../../../../public/assets/logos/cato_networks.svg"; @@ -86,7 +86,7 @@ export const guardrail_provider_map: Record = { QostodianNexus: "qostodian_nexus", Repelloai: "repelloai", Alice: "alice", - "LLM Shield": "llm_shield", + "LLM Shield Proxy": "llm_shield_proxy", }; // Function to populate provider map from API response - updates the original map @@ -210,7 +210,7 @@ export const guardrailLogoMap = { "RepelloAI Argus": repelloAiLogo.src, Straiker: straikerLogo.src, Alice: aliceLogo.src, - "LLM Shield": llmShieldLogo.src, + "LLM Shield Proxy": llmShieldProxyLogo.src, } satisfies Record; export const getGuardrailLogo = (displayName: string): string | undefined => From 6536d61adf80d668d4aebfbcce6c0b272e57de77 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Thu, 3 Sep 2026 21:22:30 -0500 Subject: [PATCH 16/34] feat(guardrails): redact the participant name on a message `name` on a user or assistant turn identifies a person and was going to the provider intact. The proxy this integrates with already redacts it, so the integration was the weaker of the two. On a tool or function turn the same field carries the function's name, which has to arrive unchanged or the call stops routing. That case is skipped, and a test asserts the value is never even sent to the shield. --- .../llm_shield_proxy/llm_shield_proxy.py | 13 +++++++++ .../guardrail_hooks/test_llm_shield_proxy.py | 28 +++++++++++++++++++ 2 files changed, 41 insertions(+) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index 56bee894714..b1ad9563eb2 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -122,6 +122,18 @@ def _collect_content(container: MutableRequest, slots: _SlotSink) -> None: _collect(part, "text", slots) +def _collect_participant_name(message: MutableRequest, slots: _SlotSink) -> None: + """Redacts `name` where it identifies a person, never where it names a function. + + On a user or assistant turn `name` is the participant, which is personal data. + On a tool or function turn the same field carries the function's name and has + to reach the provider unchanged, or the call no longer routes. + """ + if message.get("role") in ("tool", "function"): + return + _collect(message, "name", slots) + + def _collect_tool_arguments(message: MutableRequest, slots: _SlotSink) -> None: """Tool arguments carry the values a user asked the model to act on.""" for tool_call in message.get("tool_calls") or (): @@ -311,6 +323,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): for message in data.get("messages") or (): if isinstance(message, dict): _collect_content(message, slots) + _collect_participant_name(message, slots) _collect_tool_arguments(message, slots) _collect_responses_fields(data, slots) _collect_prompt(data, slots) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py index bc74b42a0e7..372bb6ed2ab 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -330,6 +330,34 @@ class TestRequestCoverage: assert data["input"] == ["[EMAIL_1]", "[PHONE_1]"] + @pytest.mark.asyncio + async def test_participant_name_is_redacted(self): + """`name` on a user turn identifies a person.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["hi", "[PERSON_1]"]}) + + data = {"messages": [{"role": "user", "name": "Jane Doe", "content": "hi"}]} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["hi", "Jane Doe"] + assert data["messages"][0]["name"] == "[PERSON_1]" + + @pytest.mark.asyncio + async def test_tool_function_name_is_left_alone(self): + """On a tool turn the same field is the function name. + + Redacting it would stop the call routing, so this asserts it is never sent + to the shield at all. + """ + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["result"]}) + + data = {"messages": [{"role": "tool", "name": "get_weather", "content": "result"}]} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert data["messages"][0]["name"] == "get_weather" + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["result"] + @pytest.mark.asyncio async def test_every_shape_in_one_request_is_redacted(self): guardrail = _guardrail() From 403b06ea76b43ae0eeef635daea3cc1b08db1730 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Thu, 3 Sep 2026 23:28:26 -0500 Subject: [PATCH 17/34] fix(guardrails): flush every held choice, and cover tool results and suffix Three review findings. The trailing flush walked the last chunk's choices, so a choice that finished earlier and stopped appearing lost whatever text was still held for it and its answer was truncated. It is now driven by the windows themselves and emits one chunk per choice, synthesising the choice when the terminal chunk omits it. That was data loss, not just under-redaction. An Anthropic tool_result carries its own content, as a string or as further blocks, and only each part's `text` was being collected. Handled recursively; image and audio parts still fall through untouched. The legacy completions `suffix` is forwarded to providers that support it and was never collected. Note the placement: it has to be gathered before the string-prompt early return, which is what the new test pins. --- .../llm_shield_proxy/llm_shield_proxy.py | 68 ++++++++++++----- .../guardrail_hooks/test_llm_shield_proxy.py | 75 +++++++++++++++++++ 2 files changed, 125 insertions(+), 18 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index b1ad9563eb2..47adc39aa9e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -100,7 +100,8 @@ def _collect_entry(entries: MutableSeq, index: int, slots: _SlotSink) -> None: def _collect_prompt(data: MutableRequest, slots: _SlotSink) -> None: - """The Completions API sends its text in a top-level `prompt`.""" + """The Completions API sends its text in `prompt`, and its tail in `suffix`.""" + _collect(data, "suffix", slots) prompt: Final = data.get("prompt") if isinstance(prompt, str): _collect(data, "prompt", slots) @@ -118,8 +119,13 @@ def _collect_content(container: MutableRequest, slots: _SlotSink) -> None: _collect(container, "content", slots) return for part in content if isinstance(content, list) else (): - if isinstance(part, dict): - _collect(part, "text", slots) + if not isinstance(part, dict): + continue + _collect(part, "text", slots) + # An Anthropic tool_result carries its own content, as a string or as more + # blocks. Image and audio parts have no text and fall through untouched. + if "content" in part: + _collect_content(part, slots) def _collect_participant_name(message: MutableRequest, slots: _SlotSink) -> None: @@ -486,8 +492,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): # A stream that ended without a finish_reason can still leave text held back. if last_chunk is not None and any(carries.values()): - trailing: Final = last_chunk.model_copy(deep=True) - if await self._flush_trailing(trailing, carries, session_id): + async for trailing in self._flush_trailing(last_chunk, carries, session_id): yield trailing async def _restore_choice(self, choice: Any, carries: _CarryWindows, session_id: str) -> None: @@ -513,23 +518,50 @@ class LLMShieldProxyGuardrail(CustomGuardrail): carries[index] = remaining # rebind-ok: this choice's window advances. delta.content = emitted - async def _flush_trailing(self, trailing: Any, carries: _CarryWindows, session_id: str) -> bool: - """Empties every still-held window into a copy of the last chunk.""" - emitted_any = False # rebind-ok: set once any choice contributes text. - for choice in getattr(trailing, "choices", None) or (): - delta = getattr(choice, "delta", None) - if delta is None: - continue - index = _choice_index(choice) - carry = carries.get(index, "") + async def _flush_trailing( + self, last_chunk: Any, carries: _CarryWindows, session_id: str + ) -> AsyncGenerator[Any, None]: + """Empties every window still holding text, one chunk per choice. + + Driven by the windows rather than by the last chunk's choices. A choice that + finished earlier is not present in the terminal chunk, and flushing only what + that chunk carries would drop its held text and truncate its answer. + """ + for index in sorted(carries): + carry = carries[index] if not carry: - delta.content = None continue text, remaining = await self._stream_step("", carry, True, session_id) carries[index] = remaining # rebind-ok: this choice's window advances. - delta.content = text or None - emitted_any = emitted_any or bool(text) - return emitted_any + if not text: + continue + chunk = self._chunk_for_choice(last_chunk, index) + if chunk is None: + continue + chunk.choices[0].delta.content = text + yield chunk + + @staticmethod + def _chunk_for_choice(last_chunk: Any, index: int) -> Any: + """A single-choice copy of the last chunk, carrying only `index`. + + Emitting one choice per chunk keeps a flush from reading as content on a + choice it does not belong to. + """ + chunk: Final = last_chunk.model_copy(deep=True) + raw_choices: Final = getattr(chunk, "choices", None) + if not raw_choices: + return None + choices: Final[tuple] = tuple(raw_choices) + matching: Final = tuple(choice for choice in choices if _choice_index(choice) == index) + kept: Final = matching[0] if matching else choices[0] + if getattr(kept, "delta", None) is None: + return None + kept.index = index + # The terminal signal, if there was one, already went out with the real chunk. + kept.finish_reason = None + chunk.choices = [kept] # mutable-ok: the chunk model requires a list. + return chunk async def _stream_step(self, text: str, carry: str, final: bool, session_id: str) -> tuple[str, str]: """Returns ``(text safe to emit now, window still being held)``.""" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py index 372bb6ed2ab..472e2606608 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -358,6 +358,43 @@ class TestRequestCoverage: assert data["messages"][0]["name"] == "get_weather" assert mock.call_args_list[0].kwargs["json"]["texts"] == ["result"] + @pytest.mark.asyncio + async def test_anthropic_tool_result_content_is_redacted(self): + """A tool_result nests its own content, as a string or as more blocks.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]", "[EMAIL_2]"]}) + + data = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "tool_result", "tool_use_id": "t1", "content": "found jane.doe@example.com"}, + { + "type": "tool_result", + "tool_use_id": "t2", + "content": [{"type": "text", "text": "also bob@example.com"}], + }, + ], + } + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert data["messages"][0]["content"][0]["content"] == "[EMAIL_1]" + assert data["messages"][0]["content"][1]["content"][0]["text"] == "[EMAIL_2]" + + @pytest.mark.asyncio + async def test_completions_suffix_is_redacted(self): + """LiteLLM forwards the legacy `suffix` to providers that support it.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["signed [EMAIL_1]", "write to [EMAIL_1]"]}) + + data = {"prompt": "write to jane.doe@example.com", "suffix": "signed jane.doe@example.com"} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="atext_completion") + + assert data["suffix"] == "signed [EMAIL_1]" + @pytest.mark.asyncio async def test_every_shape_in_one_request_is_redacted(self): guardrail = _guardrail() @@ -682,6 +719,44 @@ class TestStreamingRehydration: assert sent[2]["carry"] == "A-held", "choice 0 must get its own window back" assert sent[3]["carry"] == "B-held", "choice 1 must get its own window back" + @pytest.mark.asyncio + async def test_a_choice_missing_from_the_last_chunk_still_flushes(self): + """Held text must not be dropped because its choice ended earlier. + + Choice 1 finishes and stops appearing, then the stream ends without a + finish_reason for choice 0. Flushing only the terminal chunk's choices would + discard whatever choice 1 was still holding and truncate its answer. + """ + guardrail = _guardrail(event_hook="post_call") + _mock_post( + guardrail, + {"text": "", "carry": "held-0"}, + {"text": "", "carry": "held-1"}, + {"text": "zero-done", "carry": ""}, + {"text": "one-done", "carry": ""}, + ) + + async def stream(): + yield ModelResponseStream( + choices=[ + StreamingChoices(index=0, delta=Delta(content="a")), + StreamingChoices(index=1, delta=Delta(content="b")), + ] + ) + yield ModelResponseStream(choices=[StreamingChoices(index=0, delta=Delta(content=None))]) + + chunks = await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + flushed = { + choice.index: choice.delta.content for chunk in chunks for choice in chunk.choices if choice.delta.content + } + assert flushed.get(1) == "one-done", "choice 1's held text was dropped" + assert flushed.get(0) == "zero-done" + @pytest.mark.asyncio async def test_chunks_are_forwarded_as_they_arrive(self): """Restoration must not buffer the stream into a single terminal chunk.""" From ce921f49e47ae06946d7484ca334a0b24ac650f1 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Fri, 4 Sep 2026 02:31:20 -0500 Subject: [PATCH 18/34] fix(guardrails): walk nested tool results iteratively, with a depth bound CI flagged _collect_content as recursive. It was, and worse, it was unbounded: a tool_result nests its own content, the nesting is caller controlled, and the descent had nothing to stop it. That is a JSON bomb, not a style issue. Now an explicit queue with a depth bound of 8. Real payloads nest one or two deep. The queue is walked in document order because the shield maps its replies back by position, so collection order is part of the contract. --- .../llm_shield_proxy/llm_shield_proxy.py | 41 +++++++++++++------ .../guardrail_hooks/test_llm_shield_proxy.py | 17 ++++++++ 2 files changed, 46 insertions(+), 12 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index 47adc39aa9e..347b70ecd75 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -71,6 +71,10 @@ JsonBody: TypeAlias = dict # One redactable span: the text as it stands, and the write that puts the # replacement back where it came from. +# How far a tool_result chain is followed. Real payloads nest one or two deep; the +# bound is what stops a crafted one from becoming an unbounded walk. +_MAX_CONTENT_DEPTH: Final = 8 + _Slot: TypeAlias = tuple[str, Callable[[str], None]] # mutable-ok: Callable's param list. # The accumulator the collectors below append into. It never escapes @@ -113,19 +117,32 @@ def _collect_prompt(data: MutableRequest, slots: _SlotSink) -> None: def _collect_content(container: MutableRequest, slots: _SlotSink) -> None: - """`content` is either a string or the multimodal list of typed parts.""" - content: Final = container.get("content") - if isinstance(content, str): - _collect(container, "content", slots) - return - for part in content if isinstance(content, list) else (): - if not isinstance(part, dict): + """Collects `content`, a string or a list of typed parts. + + An Anthropic tool_result nests its own content, so this has to descend. It walks + with an explicit stack and a depth bound rather than by recursion: the nesting is + caller controlled, and an unbounded descent is a JSON bomb. + """ + # Walked in document order: the shield maps its replies back by position, so the + # order spans are collected in is part of the contract. + pending: Final[list] = [(container, 0)] # mutable-ok: local queue, never escapes. + cursor = 0 # rebind-ok: advances through the queue. + while cursor < len(pending): + node, depth = pending[cursor] + cursor += 1 + content = node.get("content") + if isinstance(content, str): + _collect(node, "content", slots) continue - _collect(part, "text", slots) - # An Anthropic tool_result carries its own content, as a string or as more - # blocks. Image and audio parts have no text and fall through untouched. - if "content" in part: - _collect_content(part, slots) + if depth >= _MAX_CONTENT_DEPTH: + continue + for part in content if isinstance(content, list) else (): + if not isinstance(part, dict): + continue + # Image and audio parts have no text and fall through untouched. + _collect(part, "text", slots) + if "content" in part: + pending.append((part, depth + 1)) def _collect_participant_name(message: MutableRequest, slots: _SlotSink) -> None: diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py index 472e2606608..a49bc5b1275 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -384,6 +384,23 @@ class TestRequestCoverage: assert data["messages"][0]["content"][0]["content"] == "[EMAIL_1]" assert data["messages"][0]["content"][1]["content"][0]["text"] == "[EMAIL_2]" + @pytest.mark.asyncio + async def test_deeply_nested_tool_results_are_bounded(self): + """Nesting is caller controlled, so the descent has to stop somewhere. + + The walk must terminate on a payload built to be pathological, rather than + following it as far as it goes. + """ + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["ok"] * 64}) + + deep: dict = {"type": "tool_result", "content": "jane.doe@example.com"} + for _ in range(200): + deep = {"type": "tool_result", "content": [deep]} + data = {"messages": [{"role": "user", "content": [deep]}]} + + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + @pytest.mark.asyncio async def test_completions_suffix_is_redacted(self): """LiteLLM forwards the legacy `suffix` to providers that support it.""" From 119ec629529a43a0e3504650d5458170a7a2dc17 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Fri, 4 Sep 2026 02:35:32 -0500 Subject: [PATCH 19/34] fix(guardrails): redact Responses PromptObject variables A Responses request can send `prompt` as a PromptObject rather than a string. Its `variables` are substituted into the stored prompt on the provider side, so they are caller text, and the dict shape was falling through untouched. `id` and `version` pick which stored prompt to run and are left unchanged. --- .../llm_shield_proxy/llm_shield_proxy.py | 9 +++++++++ .../guardrail_hooks/test_llm_shield_proxy.py | 17 +++++++++++++++++ 2 files changed, 26 insertions(+) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index 347b70ecd75..dea06cafab4 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -110,6 +110,15 @@ def _collect_prompt(data: MutableRequest, slots: _SlotSink) -> None: if isinstance(prompt, str): _collect(data, "prompt", slots) return + if isinstance(prompt, dict): + # A Responses API PromptObject. `variables` are substituted into the stored + # prompt on the provider side, so they are caller text. `id` and `version` + # identify which prompt to use and must arrive unchanged. + variables: Final = prompt.get("variables") + if isinstance(variables, dict): + for name in tuple(variables): + _collect(variables, name, slots) + return if not isinstance(prompt, list): return for index in range(len(prompt)): diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py index a49bc5b1275..b056756ac15 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -401,6 +401,23 @@ class TestRequestCoverage: await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + @pytest.mark.asyncio + async def test_responses_prompt_object_variables_are_redacted(self): + """A PromptObject's variables are substituted into the prompt provider side. + + The id and version pick which stored prompt to run and have to arrive + unchanged; the variables are caller text. + """ + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + + data = {"prompt": {"id": "pmpt_123", "version": "2", "variables": {"customer": "jane.doe@example.com"}}} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert data["prompt"]["variables"]["customer"] == "[EMAIL_1]" + assert data["prompt"]["id"] == "pmpt_123" + assert data["prompt"]["version"] == "2" + @pytest.mark.asyncio async def test_completions_suffix_is_redacted(self): """LiteLLM forwards the legacy `suffix` to providers that support it.""" From 76ac63abd4eaf9e39d5a3e1e8b921e09a0e6f918 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Fri, 4 Sep 2026 02:47:19 -0500 Subject: [PATCH 20/34] test(guardrails): assert the depth bound instead of only reaching the end The depth test asserted nothing, so it passed whether or not the bound held, and the test-quality gate counted it as a zero-assert test. It now sends a shallow value alongside a 200-deep chain and asserts the shallow one is collected while the value past the bound is not. --- .../guardrail_hooks/test_llm_shield_proxy.py | 18 +++++++++++++++--- 1 file changed, 15 insertions(+), 3 deletions(-) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py index b056756ac15..42781b441f7 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -392,15 +392,27 @@ class TestRequestCoverage: following it as far as it goes. """ guardrail = _guardrail() - _mock_post(guardrail, {"texts": ["ok"] * 64}) - deep: dict = {"type": "tool_result", "content": "jane.doe@example.com"} + captured: list = [] + + async def echo(url, headers, json, timeout): # noqa: ARG001 + captured.append(json["texts"]) + return _response({"texts": list(json["texts"])}) + + guardrail.async_handler.post = AsyncMock(side_effect=echo) # type: ignore[method-assign] + + deep: dict = {"type": "tool_result", "content": "past-the-bound@example.com"} for _ in range(200): deep = {"type": "tool_result", "content": [deep]} - data = {"messages": [{"role": "user", "content": [deep]}]} + data = {"messages": [{"role": "user", "content": [{"type": "text", "text": "shallow"}, deep]}]} await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + sent = captured[0] + assert "shallow" in sent + assert "past-the-bound@example.com" not in sent, "the walk followed the chain past its bound" + assert len(sent) < 200 + @pytest.mark.asyncio async def test_responses_prompt_object_variables_are_redacted(self): """A PromptObject's variables are substituted into the prompt provider side. From 6a4fc88d21a4449eeb9e068a0482714b9b54c417 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Sat, 5 Sep 2026 05:57:18 -0500 Subject: [PATCH 21/34] fix(guardrails): keep system-prompt values out of the restored reply Redaction put every span of a request into one vault, and the reply was restored against that same vault. System prompts are written by the application and the caller never sees them, so a caller who got the model to echo a placeholder back had its plaintext restored into their own reply -- a way to read a system prompt they were never shown. Server-authored spans now go into a vault of their own: system and developer turns, Anthropic's top-level `system`, and the Responses API `instructions`. Its id is deliberately never stored, so nothing restores against it. The reply is restored against the caller's vault alone, and an echoed placeholder from a system prompt comes back as the placeholder. Values the caller also wrote themselves are unaffected -- they are in the caller's vault too, and still restore. The extra round trip happens only when a request actually carries server-authored text. --- .../llm_shield_proxy/llm_shield_proxy.py | 63 +++++++++++----- .../guardrail_hooks/test_llm_shield_proxy.py | 72 +++++++++++++++++++ 2 files changed, 119 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index dea06cafab4..a25b077cee1 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -51,6 +51,10 @@ _REHYDRATE_STREAM_PATH: Final = "/v1/guard/rehydrate/stream" # across concurrent requests. _SESSION_METADATA_KEY: Final = "llm_shield_session_id" +# Roles whose text the application author wrote and the caller never sees. Their +# PII is still redacted outbound, but it is not restorable from the reply. +_PRIVILEGED_ROLES: Final = frozenset({"system", "developer"}) + # Vault ids are minted here and never derived from anything the caller sends. The # vault holds the plaintext behind every placeholder, so an id a caller could # supply or guess would let one user rehydrate another user's values by getting a @@ -188,9 +192,15 @@ def _collect_system(data: MutableRequest, slots: _SlotSink) -> None: _collect(part, "text", slots) -def _collect_responses_fields(data: MutableRequest, slots: _SlotSink) -> None: - """The Responses API sends text outside `messages`, in `instructions` and `input`.""" - _collect(data, "instructions", slots) +def _collect_responses_fields( + data: MutableRequest, slots: _SlotSink, privileged: _SlotSink +) -> None: + """The Responses API sends text outside `messages`, in `instructions` and `input`. + + `instructions` is written by the application, not by the caller, so it is + collected into the privileged sink; `input` is the caller's own text. + """ + _collect(data, "instructions", privileged) request_input: Final = data.get("input") if isinstance(request_input, str): _collect(data, "input", slots) @@ -344,23 +354,33 @@ class LLMShieldProxyGuardrail(CustomGuardrail): # --- request traversal -------------------------------------------------------- @staticmethod - def _locate_request_texts(data: MutableRequest) -> Sequence[_Slot]: - """Finds every redactable span in an outbound request. + def _locate_request_texts( + data: MutableRequest, + ) -> tuple[Sequence[_Slot], Sequence[_Slot]]: + """Finds every redactable span, split by whether the caller can see it. Anything missed here reaches the provider in the clear while the guardrail still reports as enabled, so the walk covers every request shape that - carries caller text. + carries text. + + The split exists because the response is restored against one vault only. + Server-authored spans -- system and developer turns, Anthropic's top-level + `system`, the Responses API `instructions` -- go into a vault nothing is + ever restored against, so a caller who gets the model to echo one of their + placeholders back receives the placeholder, not the value behind it. """ slots: Final[_SlotSink] = [] # mutable-ok: accumulator, frozen on return. + privileged: Final[_SlotSink] = [] # mutable-ok: accumulator, frozen on return. for message in data.get("messages") or (): if isinstance(message, dict): - _collect_content(message, slots) - _collect_participant_name(message, slots) - _collect_tool_arguments(message, slots) - _collect_responses_fields(data, slots) + sink = privileged if message.get("role") in _PRIVILEGED_ROLES else slots + _collect_content(message, sink) + _collect_participant_name(message, sink) + _collect_tool_arguments(message, sink) + _collect_responses_fields(data, slots, privileged) _collect_prompt(data, slots) - _collect_system(data, slots) - return tuple(slots) + _collect_system(data, privileged) + return tuple(slots), tuple(privileged) # --- hooks -------------------------------------------------------------------- @@ -376,14 +396,25 @@ class LLMShieldProxyGuardrail(CustomGuardrail): if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_call) is not True: return data - slots: Final = self._locate_request_texts(data) - if not slots: + slots, privileged = self._locate_request_texts(data) + if not slots and not privileged: return data - redacted: Final = await self._redact(tuple(text for text, _ in slots), self._mint_session_id(data)) + session_id: Final = self._mint_session_id(data) + if privileged: + # A vault of its own, whose id is deliberately never stored: the + # response is restored against `session_id` alone, so nothing the + # model emits can turn one of these placeholders back into plaintext. + await self._redact_into(privileged, f"{_VAULT_PREFIX}-{uuid.uuid4().hex}") + if slots: + await self._redact_into(slots, session_id) + return data + + async def _redact_into(self, slots: Sequence[_Slot], session_id: str) -> None: + """Redacts every span in `slots` under one vault and writes the result back.""" + redacted: Final = await self._redact(tuple(text for text, _ in slots), session_id) for (_, write), replacement in zip(slots, redacted): write(replacement) - return data @log_guardrail_information async def async_post_call_success_hook( diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py index 42781b441f7..70a29333297 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -1,3 +1,4 @@ +import json from types import SimpleNamespace from unittest.mock import AsyncMock @@ -605,6 +606,77 @@ class TestVaultIsolation: assert len(seen) == 2 + @pytest.mark.parametrize( + "data", + [ + pytest.param( + {"messages": [{"role": "system", "content": "S"}, {"role": "user", "content": "U"}]}, + id="system-turn", + ), + pytest.param( + {"messages": [{"role": "developer", "content": "S"}, {"role": "user", "content": "U"}]}, + id="developer-turn", + ), + pytest.param( + {"system": "S", "messages": [{"role": "user", "content": "U"}]}, + id="anthropic-top-level-system", + ), + pytest.param({"instructions": "S", "input": "U"}, id="responses-instructions"), + ], + ) + def test_server_authored_text_is_split_from_the_callers(self, data: dict): + """Every request shape must sort its server-authored spans out of the caller's.""" + caller, privileged = LLMShieldProxyGuardrail._locate_request_texts(data) + + assert [text for text, _ in caller] == ["U"] + assert [text for text, _ in privileged] == ["S"] + + @pytest.mark.asyncio + async def test_a_system_prompt_gets_a_vault_of_its_own(self): + """The reply is restored against the caller's vault, so the two cannot be one.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}, {"texts": ["[EMAIL_2]"]}) + + data = { + "messages": [ + {"role": "system", "content": "escalate to admin@corp.internal"}, + {"role": "user", "content": "email a@b.com"}, + ] + } + await guardrail.async_pre_call_hook( + user_api_key_dict=None, cache=None, data=data, call_type="completion" + ) + + privileged_id, caller_id = ( + call.kwargs["headers"]["X-Session-ID"] for call in mock.call_args_list + ) + assert privileged_id != caller_id + assert guardrail._session_id(data) == caller_id + + @pytest.mark.asyncio + async def test_the_system_prompt_vault_id_is_never_stored(self): + """Nothing can restore against the system vault later, because its id is not kept. + + This is what stops a caller from having the model echo a placeholder out of a + system prompt they cannot see and receiving the plaintext behind it. + """ + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}, {"texts": ["[EMAIL_2]"]}) + + data = { + "messages": [ + {"role": "system", "content": "escalate to admin@corp.internal"}, + {"role": "user", "content": "email a@b.com"}, + ] + } + await guardrail.async_pre_call_hook( + user_api_key_dict=None, cache=None, data=data, call_type="completion" + ) + + privileged_id = mock.call_args_list[0].kwargs["headers"]["X-Session-ID"] + assert privileged_id not in json.dumps(data, default=str) + + class TestFailClosed: @pytest.mark.asyncio async def test_unreachable_shield_blocks_the_request(self): From 8b2ee5ac7a96d5833483d6262f1dc89cc63027db Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Sat, 5 Sep 2026 06:25:12 -0500 Subject: [PATCH 22/34] style(guardrails): satisfy ruff format and annotate the new tests `ruff format` wanted the widened `_collect_responses_fields` signature on one line, and the three tests added with the split-vault fix needed return annotations to keep ANN201 level with the base. --- .../guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py | 4 +--- .../guardrails/guardrail_hooks/test_llm_shield_proxy.py | 6 +++--- 2 files changed, 4 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index a25b077cee1..3cab7b20cef 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -192,9 +192,7 @@ def _collect_system(data: MutableRequest, slots: _SlotSink) -> None: _collect(part, "text", slots) -def _collect_responses_fields( - data: MutableRequest, slots: _SlotSink, privileged: _SlotSink -) -> None: +def _collect_responses_fields(data: MutableRequest, slots: _SlotSink, privileged: _SlotSink) -> None: """The Responses API sends text outside `messages`, in `instructions` and `input`. `instructions` is written by the application, not by the caller, so it is diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py index 70a29333297..700702f3bb2 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -624,7 +624,7 @@ class TestVaultIsolation: pytest.param({"instructions": "S", "input": "U"}, id="responses-instructions"), ], ) - def test_server_authored_text_is_split_from_the_callers(self, data: dict): + def test_server_authored_text_is_split_from_the_callers(self, data: dict) -> None: """Every request shape must sort its server-authored spans out of the caller's.""" caller, privileged = LLMShieldProxyGuardrail._locate_request_texts(data) @@ -632,7 +632,7 @@ class TestVaultIsolation: assert [text for text, _ in privileged] == ["S"] @pytest.mark.asyncio - async def test_a_system_prompt_gets_a_vault_of_its_own(self): + async def test_a_system_prompt_gets_a_vault_of_its_own(self) -> None: """The reply is restored against the caller's vault, so the two cannot be one.""" guardrail = _guardrail() mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}, {"texts": ["[EMAIL_2]"]}) @@ -654,7 +654,7 @@ class TestVaultIsolation: assert guardrail._session_id(data) == caller_id @pytest.mark.asyncio - async def test_the_system_prompt_vault_id_is_never_stored(self): + async def test_the_system_prompt_vault_id_is_never_stored(self) -> None: """Nothing can restore against the system vault later, because its id is not kept. This is what stops a caller from having the model echo a placeholder out of a From d598d1bf96a05cd24fc3fbb6bfa7deb2ef31d596 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Sun, 13 Sep 2026 11:52:14 -0500 Subject: [PATCH 23/34] fix(guardrails): restore tool calls in the LLM Shield guardrail The request walk redacted a tool call's `arguments` -- plus the legacy `function_call`, Anthropic `tool_use.input` leaves and the Responses API's `function_call` / `function_call_output` fields -- while the response walk restored only `message.content`. A placeholder therefore reached the caller inside a tool call, and nothing raised. This is the same change as the out-of-tree example adapter this file is copied from, kept body-identical on purpose: the response side now collects every restorable span in one positional rehydrate batch, streaming keeps a window per (choice index, tool-call index) and flushes each into the chunk carrying the finish_reason, and `apply_guardrail` restores `inputs["tool_calls"]` on the response side. The declared limit on restoring values inside a JSON string is documented in the module. --- .../llm_shield_proxy/llm_shield_proxy.py | 369 ++++++++++++++---- 1 file changed, 298 insertions(+), 71 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index 3cab7b20cef..e093199578a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -85,8 +85,11 @@ _Slot: TypeAlias = tuple[str, Callable[[str], None]] # mutable-ok: Callable's p # _locate_request_texts, which freezes it into a tuple before returning. _SlotSink: TypeAlias = list[_Slot] # mutable-ok: accumulator passed between collectors. -# Sliding windows keyed by streaming choice index, threaded through one stream. -_CarryWindows: TypeAlias = dict # mutable-ok: per-choice windows advanced in place. +# Sliding windows keyed by (choice index, tool-call index | None), threaded through one +# stream. `None` is the content channel; an int is one tool call's accumulating +# `arguments`. Content and each tool call are separate token streams, so each needs its +# own window -- one shared window would splice one stream's held-back tail onto another. +_CarryWindows: TypeAlias = dict # mutable-ok: per-stream windows advanced in place. # A caller-owned list whose entries are rewritten in place, such as a Completions # `prompt` sent as an array of strings. @@ -224,6 +227,64 @@ def _choice_index(choice: object) -> int: return index if isinstance(index, int) else 0 +def _read_field(holder: object, name: str) -> object: + """Reads one field from a dict or from an object. + + LiteLLM's replies arrive as Pydantic models on some paths and as plain dicts on + others, depending how far they have been deserialised, so every response walk here + has to handle both shapes. + """ + if isinstance(holder, dict): + return holder.get(name) + return getattr(holder, name, None) + + +def _write_field(holder: object, name: str, value: str) -> None: + """Writes one string field back into a dict or an object. Pairs with _read_field.""" + if isinstance(holder, dict): + holder[name] = value + else: + setattr(holder, name, value) + + +def _collect_json_leaves(node: object, slots: _SlotSink, depth: int = 0) -> None: + """Collects every string leaf of a JSON-ish structure, with a write-back per leaf. + + An Anthropic `tool_use` block carries `input`, an arbitrary JSON object rather than a + string, so a value worth restoring can sit at any depth. Bounded by + `_MAX_CONTENT_DEPTH` for the same reason the request walk is: the shape is model + controlled, and the bound is what stops a crafted one from becoming an unbounded + descent. + """ + if depth > _MAX_CONTENT_DEPTH: + return + if isinstance(node, dict): + for key in tuple(node): + value = node[key] + if isinstance(value, str) and value: + slots.append((value, lambda new, d=node, k=key: d.__setitem__(k, new))) + else: + _collect_json_leaves(value, slots, depth + 1) + return + if isinstance(node, list): + for index, value in enumerate(node): + if isinstance(value, str) and value: + slots.append((value, lambda new, entries=node, i=index: entries.__setitem__(i, new))) + else: + _collect_json_leaves(value, slots, depth + 1) + + +def _carry_sort_key(key: tuple) -> tuple: + """Orders streaming windows without ever comparing None to an int. + + `sorted()` over the raw keys raises as soon as one choice holds both a content window + and a tool-call window, because `None < 0` is not orderable. Content sorts first, then + tool calls by their index. + """ + choice_index, tool_index = key + return (choice_index, -1 if tool_index is None else tool_index) + + class LLMShieldProxyGuardrail(CustomGuardrail): """Redacts PII before it leaves the proxy and restores it in the response. @@ -414,6 +475,15 @@ class LLMShieldProxyGuardrail(CustomGuardrail): for (_, write), replacement in zip(slots, redacted): write(replacement) + # KNOWN LIMIT: a tool call's `arguments` is a JSON *string*, and a restored value is + # spliced into it as raw text. If the original value contained a double quote, a + # backslash or a newline, the reassembled document is no longer valid JSON for a + # strict parser. The proxy's own tool-argument rehydration has the same property + # (`_rehydrate_json_response` in api/main.py), so this is a pre-existing limit of the + # product rather than one introduced here. Escaping is deliberately NOT applied as a + # fix: a fragment is an arbitrary slice of a JSON document, so the code cannot tell + # whether the position it writes is inside a string literal, and escaping + # unconditionally would corrupt the values that are not. @log_guardrail_information async def async_post_call_success_hook( self, @@ -428,27 +498,46 @@ class LLMShieldProxyGuardrail(CustomGuardrail): if self._is_anthropic_message_response(response): return await self._restore_anthropic_response(response, data) - text_blocks: Final = self._responses_api_text_blocks(response) - if text_blocks: - return await self._restore_responses_api_response(response, text_blocks, data) + response_slots: Final = self._responses_api_slots(response) + if response_slots: + return await self._restore_responses_api_response(response, response_slots, data) choices: Final = getattr(response, "choices", None) if not choices: return response - pending: Final = tuple( - (choice.message, choice.message.content) - for choice in choices - if getattr(choice, "message", None) is not None - and isinstance(getattr(choice.message, "content", None), str) - and choice.message.content - ) + # One batch for every restorable span in the reply, collected in document order: + # the shield maps its answers back by position. A second round trip is not an + # option here -- /v1/guard/rehydrate caps a batch at 256 texts and 1,000,000 + # characters, and `_same_length_or_raise` is what guarantees the positional + # mapping -- so a reply carrying more spans than that fails closed, which is this + # guardrail's posture everywhere else. + pending: Final[list] = [] # mutable-ok: local accumulator, frozen before use. + for choice in choices: + message: Final = getattr(choice, "message", None) + if message is None: + continue + content: Final = getattr(message, "content", None) + if isinstance(content, str) and content: + pending.append((content, lambda new, m=message: setattr(m, "content", new))) + # A tool call's `arguments` is model-generated text and the request path + # redacts it, so leaving it unrestored hands the application a placeholder to + # invoke a tool with. These are Pydantic objects on this path, not dicts. + for tool_call in getattr(message, "tool_calls", None) or (): + function: Final = getattr(tool_call, "function", None) + arguments: Final = getattr(function, "arguments", None) if function is not None else None + if isinstance(arguments, str) and arguments: + pending.append((arguments, lambda new, f=function: setattr(f, "arguments", new))) + legacy: Final = getattr(message, "function_call", None) + legacy_arguments: Final = getattr(legacy, "arguments", None) if legacy is not None else None + if isinstance(legacy_arguments, str) and legacy_arguments: + pending.append((legacy_arguments, lambda new, fn=legacy: setattr(fn, "arguments", new))) if not pending: return response - restored: Final = await self._rehydrate(tuple(text for _, text in pending), self._session_id(data)) - for (message, _), replacement in zip(pending, restored): - message.content = replacement + restored: Final = await self._rehydrate(tuple(text for text, _ in pending), self._session_id(data)) + for (_, write), replacement in zip(pending, restored): + write(replacement) return response @staticmethod @@ -461,60 +550,65 @@ class LLMShieldProxyGuardrail(CustomGuardrail): ) async def _restore_anthropic_response(self, response: MutableRequest, data: MutableRequest) -> MutableRequest: - """Restores text blocks in an Anthropic native message reply. + """Restores text blocks and tool inputs in an Anthropic native message reply. This shape has no `choices`, so without its own branch the reply would go back to the caller still carrying placeholders. + + A `tool_use` block's payload is `input`, an arbitrary JSON object rather than a + string, and the request path redacts its string leaves -- so the reply's leaves + have to come back or the application invokes the tool with placeholders. """ - blocks: Final = tuple( - block - for block in response["content"] - if isinstance(block, dict) and block.get("type") == "text" and isinstance(block.get("text"), str) - ) - if not blocks: + slots: Final[list] = [] # mutable-ok: accumulator, frozen before use. + for block in response["content"]: + if not isinstance(block, dict): + continue + kind: Final = block.get("type") + if kind == "text" and isinstance(block.get("text"), str) and block["text"]: + slots.append((block["text"], lambda new, b=block: b.__setitem__("text", new))) + elif kind == "tool_use" and isinstance(block.get("input"), dict): + _collect_json_leaves(block["input"], slots) + if not slots: return response - restored: Final = await self._rehydrate(tuple(block["text"] for block in blocks), self._session_id(data)) - for block, replacement in zip(blocks, restored): - block["text"] = replacement + restored: Final = await self._rehydrate(tuple(text for text, _ in slots), self._session_id(data)) + for (_, write), replacement in zip(slots, restored): + write(replacement) return response @staticmethod - def _responses_api_text_blocks(response: object) -> Sequence[object]: - """Text blocks in a Responses API reply. + def _responses_api_slots(response: object) -> Sequence[_Slot]: + """Restorable spans in a Responses API reply. That shape carries `output` items rather than `choices`, so it needs its own walk; without one the reply goes back to the caller still holding - placeholders even though the request was redacted correctly. Blocks come - through as dicts or as objects depending on how far the reply has been + placeholders even though the request was redacted correctly. Items and blocks + come through as dicts or as objects depending on how far the reply has been deserialised, so both are handled. + + The item-level fields mirror `_collect_responses_fields`, which walks the same + fields on the request side -- a function_call item holds `arguments`, a + function_call_output holds `output` -- so the two directions stay symmetric. """ - blocks: Final[list[object]] = [] # mutable-ok: accumulator, frozen on return. + slots: Final[list] = [] # mutable-ok: accumulator, frozen on return. for item in getattr(response, "output", None) or (): for block in getattr(item, "content", None) or (): - if isinstance(block, dict): - if isinstance(block.get("text"), str) and block["text"]: - blocks.append(block) - elif isinstance(getattr(block, "text", None), str) and block.text: - blocks.append(block) - return tuple(blocks) - - @staticmethod - def _block_text(block: object) -> str: - return block["text"] if isinstance(block, dict) else block.text + text: Final = _read_field(block, "text") + if isinstance(text, str) and text: + slots.append((text, lambda new, b=block: _write_field(b, "text", new))) + for field in ("arguments", "output"): + value: Final = _read_field(item, field) + if isinstance(value, str) and value: + slots.append((value, lambda new, i=item, f=field: _write_field(i, f, new))) + return tuple(slots) async def _restore_responses_api_response( - self, response: Any, blocks: Sequence[object], data: MutableRequest + self, response: Any, slots: Sequence[_Slot], data: MutableRequest ) -> Any: """Puts the original values back into a Responses API reply.""" - restored: Final = await self._rehydrate( - tuple(self._block_text(block) for block in blocks), self._session_id(data) - ) - for block, replacement in zip(blocks, restored): - if isinstance(block, dict): - block["text"] = replacement - else: - block.text = replacement + restored: Final = await self._rehydrate(tuple(text for text, _ in slots), self._session_id(data)) + for (_, write), replacement in zip(slots, restored): + write(replacement) return response async def async_post_call_streaming_iterator_hook( @@ -525,10 +619,11 @@ class LLMShieldProxyGuardrail(CustomGuardrail): ) -> AsyncGenerator[Any, None]: """Restores original values incrementally, without buffering the stream. - Each choice is its own token stream, so the sliding window is tracked per - choice index. One shared window would splice the characters held back for - one choice onto the next. The windows are locals of this generator, so they - are scoped to a single stream and cannot leak between concurrent requests. + Each choice -- and each tool call within a choice -- is its own token stream, so + the sliding window is tracked per (choice index, tool call) pair. One shared + window would splice the characters held back for one stream onto another. The + windows are locals of this generator, so they are scoped to a single stream and + cannot leak between concurrent requests. """ if self.should_run_guardrail(data=request_data, event_type=GuardrailEventHooks.post_call) is not True: async for chunk in response: @@ -536,7 +631,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): return session_id: Final = self._session_id(request_data) - carries: Final[dict] = {} # mutable-ok: per-choice windows, local to this stream. + carries: Final[dict] = {} # mutable-ok: per-stream windows, local to this generator. last_chunk = None # rebind-ok: tracks the most recent chunk for the final flush. async for chunk in response: @@ -551,49 +646,151 @@ class LLMShieldProxyGuardrail(CustomGuardrail): yield trailing async def _restore_choice(self, choice: Any, carries: _CarryWindows, session_id: str) -> None: - """Restores one choice's delta, advancing that choice's own window.""" + """Restores one choice's delta, advancing that choice's own windows. + + Content and each tool call are separate token streams, so each gets its own + window: `(choice_index, None)` for content, `(choice_index, tool_call_index)` for + one tool call's accumulating `arguments`. A shared window would splice the text + held back for one stream onto another. + """ delta: Final = getattr(choice, "delta", None) if delta is None: return index: Final = _choice_index(choice) - carry: Final = carries.get(index, "") - text: Final = getattr(delta, "content", None) is_final: Final = bool(getattr(choice, "finish_reason", None)) + await self._restore_content_window(delta, (index, None), carries, session_id, is_final) + + for tool_call in getattr(delta, "tool_calls", None) or (): + await self._restore_tool_call_window(tool_call, index, carries, session_id) + + if is_final: + # A client parses a tool call's arguments when it sees the finish_reason, so + # every window this choice still holds has to land in *this* chunk. Flushing + # after it produces argument JSON the client has already stopped waiting for. + await self._flush_finished_choice(delta, index, carries, session_id) + + async def _restore_content_window( + self, + delta: Any, + key: tuple, + carries: _CarryWindows, + session_id: str, + is_final: bool, + ) -> None: + """Restores one delta's content through its own window.""" + carry: Final = carries.get(key, "") + text: Final = getattr(delta, "content", None) + if not isinstance(text, str) or not text: # Nothing to restore here, but a final chunk still has to flush the window. if is_final and carry: - flushed, flushed_carry = await self._stream_step("", carry, True, session_id) - carries[index] = flushed_carry # rebind-ok: this choice's window advances. + flushed, remaining = await self._stream_step("", carry, True, session_id) + carries[key] = remaining # rebind-ok: this stream's window advances. if flushed: delta.content = flushed return emitted, remaining = await self._stream_step(text, carry, is_final, session_id) - carries[index] = remaining # rebind-ok: this choice's window advances. + carries[key] = remaining # rebind-ok: this stream's window advances. delta.content = emitted + async def _restore_tool_call_window( + self, + tool_call: Any, + choice_index: int, + carries: _CarryWindows, + session_id: str, + ) -> None: + """Restores one streamed tool call's argument fragment. + + A tool call's `arguments` is a JSON document delivered as fragments that clients + concatenate per tool-call index, so each index gets a window of its own rather + than sharing the content stream's. + """ + tool_index: Final = _read_field(tool_call, "index") + if not isinstance(tool_index, int): + return + function: Final = _read_field(tool_call, "function") + if function is None: + return + arguments: Final = _read_field(function, "arguments") + if not isinstance(arguments, str) or not arguments: + return + + key: Final = (choice_index, tool_index) + emitted, remaining = await self._stream_step(arguments, carries.get(key, ""), False, session_id) + carries[key] = remaining # rebind-ok: this tool call's window advances. + _write_field(function, "arguments", emitted) + + async def _flush_finished_choice( + self, + delta: Any, + choice_index: int, + carries: _CarryWindows, + session_id: str, + ) -> None: + """Emits everything this finishing choice still holds, into this chunk. + + A client parses a tool call's `arguments` when the chunk carrying the + finish_reason arrives, so a flush delivered afterwards is too late -- the client + has already tried to parse truncated JSON. Content lands back on `content`; held + tool-call text is appended as an index-only continuation entry, which is the shape + clients concatenate by index, so no id or name is needed. Appending is correct + even when this chunk already carried a fragment for that tool call. + """ + continuations: Final[list] = [] # mutable-ok: built into this chunk's delta. + for key in sorted((held for held in carries if held[0] == choice_index), key=_carry_sort_key): + carry = carries[key] + if not carry: + continue + _, tool_index = key + text, remaining = await self._stream_step("", carry, True, session_id) + carries[key] = remaining # rebind-ok: this stream's window advances. + if not text: + continue + if tool_index is None: + delta.content = text + else: + continuations.append({"index": tool_index, "function": {"arguments": text}}) + if continuations: + existing: Final[list] = list(getattr(delta, "tool_calls", None) or []) + delta.tool_calls = existing + continuations + async def _flush_trailing( self, last_chunk: Any, carries: _CarryWindows, session_id: str ) -> AsyncGenerator[Any, None]: - """Empties every window still holding text, one chunk per choice. + """Empties every window still holding text, one chunk per window. + + This is the net for a stream that ended with no finish_reason at all; a stream + that ended with one is flushed into its own terminal chunk by + `_flush_finished_choice`, because that is the moment a client parses tool + arguments. Driven by the windows rather than by the last chunk's choices. A choice that finished earlier is not present in the terminal chunk, and flushing only what that chunk carries would drop its held text and truncate its answer. """ - for index in sorted(carries): - carry = carries[index] + for key in sorted(carries, key=_carry_sort_key): + carry = carries[key] if not carry: continue + choice_index, tool_index = key text, remaining = await self._stream_step("", carry, True, session_id) - carries[index] = remaining # rebind-ok: this choice's window advances. + carries[key] = remaining # rebind-ok: this stream's window advances. if not text: continue - chunk = self._chunk_for_choice(last_chunk, index) + chunk = self._chunk_for_choice(last_chunk, choice_index) if chunk is None: continue - chunk.choices[0].delta.content = text + if tool_index is None: + chunk.choices[0].delta.content = text + else: + # The copy carried this chunk's own content and tool calls, both already + # delivered. Replace rather than append, and drop the content, or the + # client sees them twice. + chunk.choices[0].delta.content = None + chunk.choices[0].delta.tool_calls = [{"index": tool_index, "function": {"arguments": text}}] yield chunk @staticmethod @@ -645,16 +842,46 @@ class LLMShieldProxyGuardrail(CustomGuardrail): input_type: Literal["request", "response"], logging_obj: Optional["LiteLLMLoggingObj"] = None, ) -> GenericGuardrailAPIInputs: - texts: Final = inputs.get("texts") - if not texts: + """Unified entry point: what the UI's Test guardrail button and the translation + handlers call. + + `tool_calls` is handled on the response side only. LiteLLM populates the field here, + and on a reply it holds the model's tool arguments -- the same text the native hook + restores, and restoring one but not the other would leave the placeholder on + whichever path ran. The request side is left to the native pre-call hook, because + redacting it here as well would redact it twice. + """ + text_list: Final[list] = list(inputs.get("texts") or ()) + tool_calls: Final[list] = list(inputs.get("tool_calls") or ()) if input_type == "response" else [] + if not text_list and not tool_calls: return inputs + # Copied rather than mutated: the caller's tool calls are theirs to own, and this + # method's contract is to hand back a new mapping. + restored_calls: Final[list] = [copy.deepcopy(call) for call in tool_calls] # mutable-ok: a new list. + spans: Final[list] = list(text_list) # mutable-ok: ordered batch, frozen before the call. + writers: Final[list] = [] # mutable-ok: one per span appended below. + for call in restored_calls: + function: Final = _read_field(call, "function") + arguments: Final = _read_field(function, "arguments") if function is not None else None + if isinstance(arguments, str) and arguments: + spans.append(arguments) + writers.append(lambda new, f=function: _write_field(f, "arguments", new)) + replaced: Final = ( - await self._redact(tuple(texts), self._mint_session_id(request_data)) + await self._redact(tuple(spans), self._mint_session_id(request_data)) if input_type == "request" - else await self._rehydrate(tuple(texts), self._session_id(request_data)) + else await self._rehydrate(tuple(spans), self._session_id(request_data)) ) + restored_values: Final[list] = list(replaced) + + for write, replacement in zip(writers, restored_values[len(text_list) :]): + write(replacement) # Return a new mapping rather than rewriting the caller's, so this stays a # pure transform of the inputs it was handed. - merged: Final[JsonBody] = {**inputs, "texts": list(replaced)} # mutable-ok: TypedDict. + merged: Final[JsonBody] = {**inputs} # mutable-ok: TypedDict. + if text_list: + merged["texts"] = restored_values[: len(text_list)] + if restored_calls: + merged["tool_calls"] = restored_calls return merged From 28697377c3bc968e78875997f25e55dcfd007429 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Tue, 22 Sep 2026 22:02:49 -0500 Subject: [PATCH 24/34] fix(guardrails): import copy, keep the vault id off the provider, drop recursion Three defects Greptile and veria-ai found on the reopened PR, all real: - `copy.deepcopy` was called in `apply_guardrail` with no `import copy`, a guaranteed NameError on every response carrying tool calls. It landed on 2026-09-13, ten days after the review that rated this branch safe, and no test reached it: every tool-call test covered the request side. Adds the import and a regression test on the response side. - The vault session id was stored in `metadata`, which is forwarded to the provider on /v1/responses. A provider holding the placeholders and the session id can call the shield's rehydrate endpoint and read back the plaintext this guardrail exists to withhold. Moves it to `litellm_metadata`, which is not forwarded, and reads it back from there only. - `_collect_json_leaves` recursed over model-controlled JSON; the repo's recursive_detector gate rejects that. Rewritten with an explicit stack, same depth bound. 52 tests pass. ruff format, ruff-strict and check_type_discipline all clean, with LIT counts identical to the merge base. --- .../llm_shield_proxy/llm_shield_proxy.py | 53 +++++++++++-------- .../guardrail_hooks/test_llm_shield_proxy.py | 45 +++++++++++++++- 2 files changed, 75 insertions(+), 23 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index e093199578a..9877f746730 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -5,6 +5,7 @@ # # +-------------------------------------------------------------+ +import copy import os import uuid from collections.abc import AsyncGenerator, Callable, Mapping, Sequence @@ -254,24 +255,28 @@ def _collect_json_leaves(node: object, slots: _SlotSink, depth: int = 0) -> None string, so a value worth restoring can sit at any depth. Bounded by `_MAX_CONTENT_DEPTH` for the same reason the request walk is: the shape is model controlled, and the bound is what stops a crafted one from becoming an unbounded - descent. + descent. Walked with an explicit stack rather than recursively, so a deeply nested + tool input cannot spend stack frames proportional to attacker-chosen depth. """ - if depth > _MAX_CONTENT_DEPTH: - return - if isinstance(node, dict): - for key in tuple(node): - value = node[key] - if isinstance(value, str) and value: - slots.append((value, lambda new, d=node, k=key: d.__setitem__(k, new))) - else: - _collect_json_leaves(value, slots, depth + 1) - return - if isinstance(node, list): - for index, value in enumerate(node): - if isinstance(value, str) and value: - slots.append((value, lambda new, entries=node, i=index: entries.__setitem__(i, new))) - else: - _collect_json_leaves(value, slots, depth + 1) + pending: Final[list] = [(node, depth)] # mutable-ok: local walk stack. + while pending: + current, current_depth = pending.pop() + if current_depth > _MAX_CONTENT_DEPTH: + continue + if isinstance(current, dict): + for key in tuple(current): + value = current[key] + if isinstance(value, str) and value: + slots.append((value, lambda new, d=current, k=key: d.__setitem__(k, new))) + else: + pending.append((value, current_depth + 1)) + continue + if isinstance(current, list): + for index, value in enumerate(current): + if isinstance(value, str) and value: + slots.append((value, lambda new, entries=current, i=index: entries.__setitem__(i, new))) + else: + pending.append((value, current_depth + 1)) def _carry_sort_key(key: tuple) -> tuple: @@ -391,7 +396,11 @@ class LLMShieldProxyGuardrail(CustomGuardrail): caller from reaching another caller's vault. """ session_id: Final = f"{_VAULT_PREFIX}-{uuid.uuid4().hex}" - metadata: Final = data.setdefault("metadata", {}) # mutable-ok: per-request store. + # `litellm_metadata` is proxy-private; `metadata` is forwarded to the provider on + # /v1/responses. The session id is a capability against the vault's rehydrate + # endpoint, so handing it to the provider alongside the placeholders would let the + # provider read back exactly what this guardrail exists to withhold. + metadata: Final = data.setdefault("litellm_metadata", {}) # mutable-ok: per-request store. if isinstance(metadata, dict): metadata[_SESSION_METADATA_KEY] = session_id return session_id @@ -404,7 +413,9 @@ class LLMShieldProxyGuardrail(CustomGuardrail): reply that cannot be restored is a visible placeholder, while trusting a caller-supplied id would hand them someone else's plaintext. """ - metadata: Final = data.get("metadata") + # Read only from `litellm_metadata`, the same proxy-private store `_mint_session_id` + # writes to. A caller can populate `metadata`; they cannot populate this. + metadata: Final = data.get("litellm_metadata") existing: Final = metadata.get(_SESSION_METADATA_KEY) if isinstance(metadata, dict) else None if isinstance(existing, str) and existing.startswith(_VAULT_PREFIX): return existing @@ -602,9 +613,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): slots.append((value, lambda new, i=item, f=field: _write_field(i, f, new))) return tuple(slots) - async def _restore_responses_api_response( - self, response: Any, slots: Sequence[_Slot], data: MutableRequest - ) -> Any: + async def _restore_responses_api_response(self, response: Any, slots: Sequence[_Slot], data: MutableRequest) -> Any: """Puts the original values back into a Responses API reply.""" restored: Final = await self._rehydrate(tuple(text for text, _ in slots), self._session_id(data)) for (_, write), replacement in zip(slots, restored): diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py index 700702f3bb2..bfa8c04c603 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -575,7 +575,25 @@ class TestVaultIsolation: used = mock.call_args_list[0].kwargs["headers"]["X-Session-ID"] assert used != "victim-session" - assert data["metadata"]["llm_shield_session_id"] == used + assert data["litellm_metadata"]["llm_shield_session_id"] == used + + @pytest.mark.asyncio + async def test_session_id_is_not_forwarded_to_the_provider(self): + """The vault id is a capability, so it must stay out of provider-visible metadata. + + `metadata` is forwarded upstream on /v1/responses; `litellm_metadata` is not. A + provider holding both the placeholders and the session id could call the shield's + rehydrate endpoint and read back exactly what this guardrail withholds. + """ + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + + data = {"messages": [{"role": "user", "content": "a@b.com"}], "metadata": {}} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + used = mock.call_args_list[0].kwargs["headers"]["X-Session-ID"] + assert "llm_shield_session_id" not in data["metadata"] + assert data["litellm_metadata"]["llm_shield_session_id"] == used @pytest.mark.asyncio async def test_restore_ignores_a_foreign_session_id(self): @@ -899,3 +917,28 @@ class TestStreamingRehydration: assert len(chunks) == 3 assert [c.choices[0].delta.content for c in chunks] == ["one ", "two ", "three"] + + +class TestApplyGuardrailToolCalls: + """The unified entry point the UI's Test button and the translation handlers use.""" + + @pytest.mark.asyncio + async def test_response_tool_call_arguments_are_rehydrated(self): + """Regression: this path deep-copied tool calls with `copy` never imported. + + 47 tests passed with a guaranteed NameError here, because every tool-call test + covered the request side and this is the only path that reaches the copy. + """ + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["hi", '{"email": "a@b.com"}']}) + + data = {"litellm_metadata": {"llm_shield_session_id": "shield-abc"}} + inputs = { + "texts": ["hi"], + "tool_calls": [{"function": {"name": "send", "arguments": '{"email": "[EMAIL_1]"}'}}], + } + + merged = await guardrail.apply_guardrail(inputs=inputs, request_data=data, input_type="response") + + assert merged["tool_calls"][0]["function"]["arguments"] == '{"email": "a@b.com"}' + assert inputs["tool_calls"][0]["function"]["arguments"] == '{"email": "[EMAIL_1]"}' From b9656734cdf549fa7c442570a2f4c9049544c9c4 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Sat, 26 Sep 2026 01:06:05 -0500 Subject: [PATCH 25/34] fix(guardrails): build llm_shield_proxy stream deltas without new mutable literals The lint job's LIT002 budget gate failed on this PR: the file added 11 mutable-collection constructions and the tree sits at its limit. Build the index-only tool-call continuation in one helper, keep read-only inputs as tuples, and annotate the lists the delta and texts fields require. Adds tests for the two tool-call flush paths the refactor touches, which had no coverage: held arguments landing in the finish_reason chunk next to that chunk's own fragment, and the trailing flush of a stream that ends without a finish_reason. --- .../llm_shield_proxy/llm_shield_proxy.py | 22 ++++-- .../guardrail_hooks/test_llm_shield_proxy.py | 74 +++++++++++++++++++ 2 files changed, 89 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index 9877f746730..5ae8a434793 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -290,6 +290,14 @@ def _carry_sort_key(key: tuple) -> tuple: return (choice_index, -1 if tool_index is None else tool_index) +def _continuation_delta(tool_index: int, text: str) -> list[dict[str, object]]: + """A `tool_calls` delta carrying `text` as an index-only continuation. + + Clients concatenate tool-call fragments by index, so no id or name is needed. + """ + return [{"index": tool_index, "function": {"arguments": text}}] # mutable-ok: delta.tool_calls is a list. + + class LLMShieldProxyGuardrail(CustomGuardrail): """Redacts PII before it leaves the proxy and restores it in the response. @@ -761,10 +769,10 @@ class LLMShieldProxyGuardrail(CustomGuardrail): if tool_index is None: delta.content = text else: - continuations.append({"index": tool_index, "function": {"arguments": text}}) + continuations.extend(_continuation_delta(tool_index, text)) if continuations: - existing: Final[list] = list(getattr(delta, "tool_calls", None) or []) - delta.tool_calls = existing + continuations + existing: Final = tuple(getattr(delta, "tool_calls", None) or ()) + delta.tool_calls = [*existing, *continuations] # mutable-ok: delta.tool_calls is a list. async def _flush_trailing( self, last_chunk: Any, carries: _CarryWindows, session_id: str @@ -799,7 +807,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): # delivered. Replace rather than append, and drop the content, or the # client sees them twice. chunk.choices[0].delta.content = None - chunk.choices[0].delta.tool_calls = [{"index": tool_index, "function": {"arguments": text}}] + chunk.choices[0].delta.tool_calls = _continuation_delta(tool_index, text) yield chunk @staticmethod @@ -860,8 +868,8 @@ class LLMShieldProxyGuardrail(CustomGuardrail): whichever path ran. The request side is left to the native pre-call hook, because redacting it here as well would redact it twice. """ - text_list: Final[list] = list(inputs.get("texts") or ()) - tool_calls: Final[list] = list(inputs.get("tool_calls") or ()) if input_type == "response" else [] + text_list: Final = tuple(inputs.get("texts") or ()) + tool_calls: Final = tuple(inputs.get("tool_calls") or ()) if input_type == "response" else () if not text_list and not tool_calls: return inputs @@ -882,7 +890,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): if input_type == "request" else await self._rehydrate(tuple(spans), self._session_id(request_data)) ) - restored_values: Final[list] = list(replaced) + restored_values: Final[list] = list(replaced) # mutable-ok: sliced into the texts list. for write, replacement in zip(writers, restored_values[len(text_list) :]): write(replacement) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py index bfa8c04c603..0197b83c9f0 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -53,6 +53,19 @@ async def _drain(generator) -> list: return [chunk async for chunk in generator] +def _tool_chunk(arguments: str, finish_reason: str | None = None) -> ModelResponseStream: + """One streamed fragment of tool call 0's arguments.""" + tool_call = {"index": 0, "id": "call_1", "type": "function", "function": {"name": "send", "arguments": arguments}} + return ModelResponseStream( + choices=[StreamingChoices(index=0, delta=Delta(tool_calls=[tool_call]), finish_reason=finish_reason)] + ) + + +def _field(holder: object, name: str) -> object: + """Reads a field from a dict or a model; the guardrail emits both shapes.""" + return holder.get(name) if isinstance(holder, dict) else getattr(holder, name) + + 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", {}) @@ -893,6 +906,67 @@ class TestStreamingRehydration: assert flushed.get(1) == "one-done", "choice 1's held text was dropped" assert flushed.get(0) == "zero-done" + @pytest.mark.asyncio + async def test_held_tool_arguments_land_in_the_finishing_chunk(self): + """A client parses tool arguments on finish_reason, so the flush must ride that chunk. + + The finishing chunk also carries its own fragment for the same tool call. That + entry has to survive, with the held text appended after it as a continuation. + """ + guardrail = _guardrail(event_hook="post_call") + _mock_post( + guardrail, + {"text": '{"to": "', "carry": "[EMA"}, + {"text": "", "carry": '[EMAIL_1]"}'}, + {"text": 'a@example.com"}', "carry": ""}, + ) + + async def stream(): + yield _tool_chunk('{"to": "[EMA') + yield _tool_chunk('IL_1]"}', finish_reason="tool_calls") + + chunks = await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + assert len(chunks) == 2, "the flush must not arrive after the finish_reason chunk" + final_calls = chunks[1].choices[0].delta.tool_calls + assert len(final_calls) == 2, "the finishing chunk's own fragment was dropped" + assert _field(final_calls[1], "index") == 0 + arguments = "".join( + _field(_field(call, "function"), "arguments") or "" + for chunk in chunks + for call in chunk.choices[0].delta.tool_calls + ) + assert json.loads(arguments) == {"to": "a@example.com"} + + @pytest.mark.asyncio + async def test_held_tool_arguments_flush_when_the_stream_ends_unfinished(self): + """No finish_reason at all: a trailing chunk carries the held arguments alone.""" + guardrail = _guardrail(event_hook="post_call") + _mock_post( + guardrail, + {"text": '{"to": "', "carry": "[EMA"}, + {"text": 'a@example.com"}', "carry": ""}, + ) + + async def stream(): + yield _tool_chunk('{"to": "[EMA') + + chunks = await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + assert len(chunks) == 2 + trailing = chunks[1].choices[0].delta.tool_calls + assert trailing == [{"index": 0, "function": {"arguments": 'a@example.com"}'}}], ( + "the copied chunk's own fragment was already delivered and must not repeat" + ) + @pytest.mark.asyncio async def test_chunks_are_forwarded_as_they_arrive(self): """Restoration must not buffer the stream into a single terminal chunk.""" From 12e3b14a030469dd5de1752e79656352c4abf00b Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Sat, 26 Sep 2026 01:21:36 -0500 Subject: [PATCH 26/34] fix(guardrails): drop Final from loop-body locals in llm_shield_proxy basedpyright rejects Final on a name assigned inside a loop, and the eleven such locals put reportGeneralTypeIssues over its budget (112/101). The LIT010 Final rule already exempts loop-body assignments, so the annotations go. --- .../llm_shield_proxy/llm_shield_proxy.py | 22 +++++++++---------- 1 file changed, 11 insertions(+), 11 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index 5ae8a434793..8e3e6d81dbc 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -533,22 +533,22 @@ class LLMShieldProxyGuardrail(CustomGuardrail): # guardrail's posture everywhere else. pending: Final[list] = [] # mutable-ok: local accumulator, frozen before use. for choice in choices: - message: Final = getattr(choice, "message", None) + message = getattr(choice, "message", None) if message is None: continue - content: Final = getattr(message, "content", None) + content = getattr(message, "content", None) if isinstance(content, str) and content: pending.append((content, lambda new, m=message: setattr(m, "content", new))) # A tool call's `arguments` is model-generated text and the request path # redacts it, so leaving it unrestored hands the application a placeholder to # invoke a tool with. These are Pydantic objects on this path, not dicts. for tool_call in getattr(message, "tool_calls", None) or (): - function: Final = getattr(tool_call, "function", None) - arguments: Final = getattr(function, "arguments", None) if function is not None else None + function = getattr(tool_call, "function", None) + arguments = getattr(function, "arguments", None) if function is not None else None if isinstance(arguments, str) and arguments: pending.append((arguments, lambda new, f=function: setattr(f, "arguments", new))) - legacy: Final = getattr(message, "function_call", None) - legacy_arguments: Final = getattr(legacy, "arguments", None) if legacy is not None else None + legacy = getattr(message, "function_call", None) + legacy_arguments = getattr(legacy, "arguments", None) if legacy is not None else None if isinstance(legacy_arguments, str) and legacy_arguments: pending.append((legacy_arguments, lambda new, fn=legacy: setattr(fn, "arguments", new))) if not pending: @@ -582,7 +582,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): for block in response["content"]: if not isinstance(block, dict): continue - kind: Final = block.get("type") + kind = block.get("type") if kind == "text" and isinstance(block.get("text"), str) and block["text"]: slots.append((block["text"], lambda new, b=block: b.__setitem__("text", new))) elif kind == "tool_use" and isinstance(block.get("input"), dict): @@ -612,11 +612,11 @@ class LLMShieldProxyGuardrail(CustomGuardrail): slots: Final[list] = [] # mutable-ok: accumulator, frozen on return. for item in getattr(response, "output", None) or (): for block in getattr(item, "content", None) or (): - text: Final = _read_field(block, "text") + text = _read_field(block, "text") if isinstance(text, str) and text: slots.append((text, lambda new, b=block: _write_field(b, "text", new))) for field in ("arguments", "output"): - value: Final = _read_field(item, field) + value = _read_field(item, field) if isinstance(value, str) and value: slots.append((value, lambda new, i=item, f=field: _write_field(i, f, new))) return tuple(slots) @@ -879,8 +879,8 @@ class LLMShieldProxyGuardrail(CustomGuardrail): spans: Final[list] = list(text_list) # mutable-ok: ordered batch, frozen before the call. writers: Final[list] = [] # mutable-ok: one per span appended below. for call in restored_calls: - function: Final = _read_field(call, "function") - arguments: Final = _read_field(function, "arguments") if function is not None else None + function = _read_field(call, "function") + arguments = _read_field(function, "arguments") if function is not None else None if isinstance(arguments, str) and arguments: spans.append(arguments) writers.append(lambda new, f=function: _write_field(f, "arguments", new)) From f1a689059f4ce5c22a46857e5206a8bd138a0a0e Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Sat, 26 Sep 2026 09:12:40 -0500 Subject: [PATCH 27/34] feat(guardrails): restore llm_shield_proxy placeholders on native streams Anthropic /v1/messages and /v1/responses streams have no `choices`, so the streaming hook passed them through with placeholders still in them. Both are now restored incrementally, with the same per-stream windows as chat: - /v1/messages arrives as raw SSE. Frames are cut at event boundaries, text_delta and input_json_delta are restored per block index, and held text is emitted as one more delta ahead of content_block_stop. Signed thinking deltas, frames from other endpoints and non-SSE raw streams pass through unchanged. - /v1/responses events are restored per item and part. Held text goes out as a copy of the stream's last delta before its .done event, and the events that repeat the reply (.done, content_part.done, output_item.done, response.completed) are restored in full. The request side now also redacts Anthropic tool_use inputs and Responses reasoning summaries, and sends tool and function descriptions (including parameter schema descriptions) and the user / safety_identifier fields to the non-restorable vault, like system prompts. Tool results stay restorable: the model reads them to answer, so restoring them returns what the caller would have seen without the guardrail. --- .../llm_shield_proxy/llm_shield_proxy.py | 468 +++++++++++++++++- .../guardrail_hooks/test_llm_shield_proxy.py | 371 ++++++++++++++ 2 files changed, 822 insertions(+), 17 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index 8e3e6d81dbc..ebafb10831a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -6,9 +6,14 @@ # +-------------------------------------------------------------+ import copy +import functools +import json import os +import re import uuid -from collections.abc import AsyncGenerator, Callable, Mapping, Sequence +from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequence +from enum import Enum +from types import MappingProxyType from typing import ( TYPE_CHECKING, Any, # noqa: TID251 # **kwargs forwards verbatim to CustomGuardrail.__init__ @@ -80,8 +85,54 @@ JsonBody: TypeAlias = dict # bound is what stops a crafted one from becoming an unbounded walk. _MAX_CONTENT_DEPTH: Final = 8 +# How far a tool's parameter schema is followed. Deeper than content: every nested +# object costs two levels (`properties`, then the property), and a description missed +# here goes to the provider in the clear. +_MAX_SCHEMA_DEPTH: Final = 32 + _Slot: TypeAlias = tuple[str, Callable[[str], None]] # mutable-ok: Callable's param list. +# One incremental rehydration step for a stream the caller has already bound to its +# vault: (new text, carried window, final) -> (text safe to emit, window still held). +_StreamStep: TypeAlias = Callable[[str, str, bool], Awaitable[tuple[str, str]]] # mutable-ok: Callable's param list. + +# A batch rehydration already bound to the request's vault. +_Rehydrate: TypeAlias = Callable[[Sequence[str]], Awaitable[Sequence[str]]] # mutable-ok: Callable's param list. + +# Anthropic /v1/messages delta types that carry restorable text, and the field holding +# it. `thinking_delta` is left out on purpose: a thinking block is signed, and one +# rewritten here fails verification when the client sends it back on the next turn. +_ANTHROPIC_DELTA_FIELDS: Final = MappingProxyType({"text_delta": "text", "input_json_delta": "partial_json"}) + +# A blank line ends an SSE event. Frames are cut there, never inside an event, so a +# `data:` line split across two network chunks is parsed only once it is whole. +_SSE_EVENT_BOUNDARY: Final = re.compile(rb"(\r?\n\r?\n)") + +# Responses API events whose `delta` is model text. Each belongs to the stream that the +# matching `.done` event in `_RESPONSES_DONE_FIELDS` closes. +_RESPONSES_DELTA_EVENTS: Final = frozenset( + ( + "response.output_text.delta", + "response.refusal.delta", + "response.function_call_arguments.delta", + "response.reasoning_summary_text.delta", + ) +) + +# The `.done` event that closes each delta stream, and the field that repeats the +# stream's full text on it. +_RESPONSES_DONE_FIELDS: Final = MappingProxyType( + { + "response.output_text.done": "text", + "response.refusal.done": "refusal", + "response.function_call_arguments.done": "arguments", + "response.reasoning_summary_text.done": "text", + } +) + +# Terminal Responses API events that repeat the whole reply under `response`. +_RESPONSES_TERMINAL_EVENTS: Final = frozenset(("response.completed", "response.incomplete")) + # The accumulator the collectors below append into. It never escapes # _locate_request_texts, which freezes it into a tuple before returning. _SlotSink: TypeAlias = list[_Slot] # mutable-ok: accumulator passed between collectors. @@ -158,6 +209,11 @@ def _collect_content(container: MutableRequest, slots: _SlotSink) -> None: continue # Image and audio parts have no text and fall through untouched. _collect(part, "text", slots) + if part.get("type") == "tool_use": + # A replayed Anthropic tool call. Its `input` is a JSON object rather than + # a string, so a value can sit at any depth -- the reply side walks the + # same leaves when it restores one. + _collect_json_leaves(part.get("input"), slots, depth + 1) if "content" in part: pending.append((part, depth + 1)) @@ -220,6 +276,72 @@ def _collect_responses_fields(data: MutableRequest, slots: _SlotSink, privileged # A function_call item holds `arguments`; a function_call_output holds `output`. _collect(item, "arguments", slots) _collect(item, "output", slots) + # A replayed reasoning item carries the model's summary of its own reasoning, + # which quotes whatever the conversation contained. + _collect_text_parts(item, "summary", slots) + + +def _collect_text_parts(container: MutableRequest, key: str, slots: _SlotSink) -> None: + """Collects the `text` of every part in the list held at `key`.""" + parts: Final = container.get(key) + for part in parts if isinstance(parts, list) else (): + if isinstance(part, dict): + _collect(part, "text", slots) + + +def _collect_tool_definitions(data: MutableRequest, privileged: _SlotSink) -> None: + """Tool definitions are application-authored free text bound for the provider. + + A description -- on the tool, or on any property of its parameter schema -- is where + callers put examples and customer context, so it carries PII as often as a prompt + does. It is collected into the privileged sink, like a system prompt: redacted + outbound, and never restorable from the reply. Names, types and enum values are left + as sent, because the model has to reproduce them exactly for a call to route. + + Covers Chat `tools[].function`, the legacy `functions[]`, and the flat tool shape the + Responses API and Anthropic share, whose schema is `parameters` or `input_schema`. + """ + for key in ("tools", "functions"): + declared = data.get(key) + for tool in declared if isinstance(declared, list) else (): + if not isinstance(tool, dict): + continue + function = tool.get("function") + for holder in (tool, function) if isinstance(function, dict) else (tool,): + _collect(holder, "description", privileged) + _collect_schema_descriptions(holder.get("parameters"), privileged) + _collect_schema_descriptions(holder.get("input_schema"), privileged) + + +def _collect_schema_descriptions(schema: object, privileged: _SlotSink) -> None: + """Collects every string `description` in a JSON schema, at any depth. + + Only `description` is free text. A property that is itself *named* "description" + holds a schema object rather than a string, so it is descended into, not collected. + Walked with an explicit stack and a depth bound, like the other request walks. + """ + pending: Final[list] = [(schema, 0)] # mutable-ok: local walk stack. + while pending: + node, depth = pending.pop() + if depth > _MAX_SCHEMA_DEPTH: + continue + if isinstance(node, dict): + _collect(node, "description", privileged) + pending.extend((value, depth + 1) for value in node.values() if isinstance(value, (dict, list))) + elif isinstance(node, list): + pending.extend((value, depth + 1) for value in node if isinstance(value, (dict, list))) + + +def _collect_end_user_ids(data: MutableRequest, privileged: _SlotSink) -> None: + """`user` and `safety_identifier` are forwarded to the provider and often hold an email. + + Only detected PII is replaced, so an opaque id reaches the provider unchanged. LiteLLM's + own end-user spend tracking reads the id resolved at authentication, before this hook + runs, so rewriting the field here does not move spend. Nothing restores these from a + reply, hence the privileged sink. + """ + _collect(data, "user", privileged) + _collect(data, "safety_identifier", privileged) def _choice_index(choice: object) -> int: @@ -240,6 +362,15 @@ def _read_field(holder: object, name: str) -> object: return getattr(holder, name, None) +def _read_list(holder: object, name: str) -> Sequence[object]: + """Reads a list field from a dict or an object; anything else reads as empty. + + The entries are the reply's own objects, so writing through them edits the reply. + """ + value: Final = _read_field(holder, name) + return tuple(value) if isinstance(value, (list, tuple)) else () + + def _write_field(holder: object, name: str, value: str) -> None: """Writes one string field back into a dict or an object. Pairs with _read_field.""" if isinstance(holder, dict): @@ -298,6 +429,289 @@ def _continuation_delta(tool_index: int, text: str) -> list[dict[str, object]]: return [{"index": tool_index, "function": {"arguments": text}}] # mutable-ok: delta.tool_calls is a list. +def _collect_response_item(item: object, slots: _SlotSink) -> None: + """Restorable spans in one Responses API output item, dict or object. + + Mirrors `_collect_responses_fields` on the request side -- a function_call item holds + `arguments`, a function_call_output holds `output`, a reasoning item holds `summary` + parts -- so the two directions stay symmetric. + """ + for block in _read_list(item, "content"): + for field in ("text", "refusal"): + text = _read_field(block, field) + if isinstance(text, str) and text: + slots.append((text, lambda new, b=block, f=field: _write_field(b, f, new))) + for part in _read_list(item, "summary"): + text = _read_field(part, "text") + if isinstance(text, str) and text: + slots.append((text, lambda new, p=part: _write_field(p, "text", new))) + for field in ("arguments", "output"): + value = _read_field(item, field) + if isinstance(value, str) and value: + slots.append((value, lambda new, i=item, f=field: _write_field(i, f, new))) + + +async def _rehydrate_slots(slots: Sequence[_Slot], rehydrate: _Rehydrate) -> None: + """Restores every span in `slots` in one batch and writes each result back.""" + if not slots: + return + restored: Final = await rehydrate(tuple(text for text, _ in slots)) + for (_, write), replacement in zip(slots, restored): + write(replacement) + + +def _responses_event_type(chunk: object) -> str | None: + """The event type of a Responses API stream event, or None for any other chunk. + + The type arrives as a plain string on dicts and as a str-valued Enum on LiteLLM's + event models. The Enum is unwrapped because it does not hash like its value, so it + would miss every lookup in the event tables above. + """ + if isinstance(chunk, (bytes, str)): + return None + kind: Final = _read_field(chunk, "type") + value: Final = kind.value if isinstance(kind, Enum) else kind + return value if isinstance(value, str) and value.startswith("response.") else None + + +class _AnthropicSSERestorer: + """Restores an Anthropic `/v1/messages` stream, which reaches the hook as raw SSE. + + Each content block is its own token stream with its own window, keyed by the block's + `index`: `text_delta` carries prose and `input_json_delta` a tool call's arguments. + When a block stops, whatever its window still holds is emitted as one more delta for + that block, just ahead of the `content_block_stop` frame, so the client has the whole + block before it is told the block is complete. + + Frames are processed whole. A network chunk can end in the middle of an event, so the + unfinished tail is kept until the rest arrives; that delays one partial event, never + a completed one. A frame that is not an Anthropic event -- another endpoint's SSE, or + anything that fails to parse -- is passed through byte for byte, and a raw stream that + does not open like SSE at all is passed through chunk by chunk, never buffered. + """ + + def __init__(self, step: _StreamStep) -> None: + self._step: Final = step + self._carries: Final[dict[int, str]] = {} # mutable-ok: per-block windows advanced in place. + self._delta_types: Final[dict[int, str]] = {} # mutable-ok: each block's delta type, for its flush. + self._pending = b"" # rebind-ok: the unfinished tail of the stream. + self._as_text = False # rebind-ok: set once if the stream arrives as str rather than bytes. + self._is_sse: bool | None = None # rebind-ok: decided once, by the stream's first chunk. + + async def feed(self, chunk: bytes | str) -> tuple[bytes | str, ...]: + """Restores every event this chunk completes; holds back an unfinished tail.""" + if isinstance(chunk, str): + self._as_text = True + raw: Final = chunk.encode("utf-8") if isinstance(chunk, str) else chunk + if self._is_sse is None and raw.strip(): + # An SSE stream opens with a field or a comment. Anything else (a JSON array + # streamed in pieces, say) has no event boundaries to wait for. + self._is_sse = raw.lstrip().startswith((b"event:", b"data:", b":")) + if not self._is_sse: + return (chunk,) + buffered: Final = self._pending + raw + boundaries: Final = tuple(_SSE_EVENT_BOUNDARY.finditer(buffered)) + if not boundaries: + self._pending = buffered + return () + cut: Final = boundaries[-1].end() + self._pending = buffered[cut:] + # With a capturing group, split alternates event, separator, ..., and ends in the + # empty remainder after the last separator. + parts: Final = _SSE_EVENT_BOUNDARY.split(buffered[:cut]) + restored: Final = tuple( + [ # mutable-ok: an await needs a list comprehension; frozen at once. + await self._restore_event(parts[index]) + parts[index + 1] for index in range(0, len(parts) - 1, 2) + ] + ) + return self._emit(b"".join(restored)) + + async def finish(self) -> tuple[bytes | str, ...]: + """Emits an unterminated final event and any window a block never closed.""" + tail: Final = await self._restore_event(self._pending) if self._pending.strip() else self._pending + self._pending = b"" + flushed: Final = await self._flush_all() + # The tail had no blank line after it; one is needed before another frame follows. + separator: Final = b"\n\n" if tail.strip() and flushed else b"" + return self._emit(tail + separator + flushed) + + def _emit(self, frames: bytes) -> tuple[bytes | str, ...]: + if not frames: + return () + return (frames.decode("utf-8") if self._as_text else frames,) + + async def _restore_event(self, block: bytes) -> bytes: + """Rewrites one SSE event, or returns it untouched if it carries nothing to restore.""" + try: + lines: Final = block.decode("utf-8").split("\n") + except UnicodeDecodeError: + return block + data_lines: Final = tuple(index for index, line in enumerate(lines) if line.startswith("data:")) + if len(data_lines) != 1: + return block + line: Final = lines[data_lines[0]] + try: + event: Final = json.loads(line[len("data:") :]) + except ValueError: + return block + if not isinstance(event, dict): + return block + kind: Final = event.get("type") + index: Final = event.get("index") + if kind == "content_block_stop" and isinstance(index, int): + return await self._flush(index) + block + if kind == "message_stop": + return await self._flush_all() + block + if kind != "content_block_delta" or not await self._restore_delta(event): + return block + ending: Final = "\r" if line.endswith("\r") else "" + rewritten: Final = ( + *lines[: data_lines[0]], + f"data: {json.dumps(event, ensure_ascii=False)}{ending}", + *lines[data_lines[0] + 1 :], + ) + return "\n".join(rewritten).encode("utf-8") + + async def _restore_delta(self, event: MutableRequest) -> bool: + """Advances one block's window through this delta. False if it holds no text.""" + index: Final = event.get("index") + delta: Final = event.get("delta") + if not isinstance(index, int) or not isinstance(delta, dict): + return False + delta_type: Final = delta.get("type") + if not isinstance(delta_type, str): + return False + field: Final = _ANTHROPIC_DELTA_FIELDS.get(delta_type) + text: Final = delta.get(field) if field is not None else None + if field is None or not isinstance(text, str) or not text: + return False + emitted, remaining = await self._step(text, self._carries.get(index, ""), False) + self._carries[index] = remaining + self._delta_types[index] = delta_type + delta[field] = emitted + return True + + async def _flush(self, index: int) -> bytes: + """One synthetic delta frame carrying whatever `index`'s window still holds.""" + carry: Final = self._carries.pop(index, "") + delta_type: Final = self._delta_types.pop(index, None) + field: Final = _ANTHROPIC_DELTA_FIELDS.get(delta_type) if isinstance(delta_type, str) else None + if not carry or field is None: + return b"" + text, _ = await self._step("", carry, True) + if not text: + return b"" + event: Final[JsonBody] = { # mutable-ok: serialised on the next line. + "type": "content_block_delta", + "index": index, + "delta": {"type": delta_type, field: text}, + } + return f"event: content_block_delta\ndata: {json.dumps(event, ensure_ascii=False)}\n\n".encode() + + async def _flush_all(self) -> bytes: + flushed: Final = tuple([await self._flush(index) for index in tuple(self._carries)]) # mutable-ok: frozen. + return b"".join(flushed) + + +class _ResponsesStreamRestorer: + """Restores a Responses API event stream. + + Every delta stream -- one output_text content part, one refusal, one function call's + arguments, one reasoning summary part -- gets its own window, keyed by the event + family, the item id and the part index. When its `.done` event arrives, whatever the + window still holds goes out first, as a copy of that stream's last delta event -- so + it carries the stream's own ids, and repeats that event's `sequence_number` -- and + the `.done` event's full text is then restored in one call. + + The events that repeat the reply wholesale -- `content_part.done`, + `output_item.done`, and `response.completed` / `response.incomplete` -- are restored + the same way the non-streaming reply is. + """ + + def __init__(self, step: _StreamStep, rehydrate: _Rehydrate) -> None: + self._step: Final = step + self._rehydrate: Final = rehydrate + self._carries: Final[dict[tuple, str]] = {} # mutable-ok: per-stream windows advanced in place. + self._last_deltas: Final[dict[tuple, object]] = {} # mutable-ok: newest delta per stream. + + async def restore(self, event: object) -> tuple[object, ...]: + """The events to emit in place of `event`: any flush, then the event itself.""" + kind: Final = _responses_event_type(event) + if kind is None: + return (event,) + if kind in _RESPONSES_DELTA_EVENTS: + await self._restore_delta(event, kind) + return (event,) + done_field: Final = _RESPONSES_DONE_FIELDS.get(kind) + if done_field is not None: + flushed: Final = await self._flush(_responses_stream_key(event, kind)) + await _rehydrate_slots(_field_slot(event, done_field), self._rehydrate) + return (*flushed, event) + slots: Final[_SlotSink] = [] # mutable-ok: accumulator, restored in one batch. + if kind == "response.content_part.done": + part: Final = _read_field(event, "part") + _collect_response_item({"content": [part]}, slots) # mutable-ok: a one-part view. + elif kind == "response.output_item.done": + _collect_response_item(_read_field(event, "item"), slots) + elif kind in _RESPONSES_TERMINAL_EVENTS: + for item in _read_list(_read_field(event, "response"), "output"): + _collect_response_item(item, slots) + await _rehydrate_slots(slots, self._rehydrate) + return (event,) + + async def finish(self) -> tuple[object, ...]: + """Flushes every stream the provider never closed, e.g. a truncated reply.""" + flushed: Final = tuple([await self._flush(key) for key in tuple(self._carries)]) # mutable-ok: frozen. + return tuple(event for events in flushed for event in events) + + async def _restore_delta(self, event: object, kind: str) -> None: + text: Final = _read_field(event, "delta") + if not isinstance(text, str) or not text: + return + key: Final = _responses_stream_key(event, kind) + emitted, remaining = await self._step(text, self._carries.get(key, ""), False) + self._carries[key] = remaining + self._last_deltas[key] = event + _write_field(event, "delta", emitted) + + async def _flush(self, key: tuple) -> tuple[object, ...]: + carry: Final = self._carries.pop(key, "") + template: Final = self._last_deltas.pop(key, None) + if not carry or template is None: + return () + text, _ = await self._step("", carry, True) + if not text: + return () + flush: Final = copy.deepcopy(template) + _write_field(flush, "delta", text) + return (flush,) + + +def _responses_stream_key(event: object, kind: str) -> tuple: + """Identifies the delta stream an event belongs to, the same for its delta and done. + + The family is the event type without its `.delta` / `.done` suffix, so an output_text + stream and a refusal stream on the same part never share a window. + """ + family: Final = kind.rsplit(".", 1)[0] + part_index: Final = _read_field(event, "content_index") + summary_index: Final = _read_field(event, "summary_index") + return ( + family, + _read_field(event, "item_id"), + _read_field(event, "output_index"), + part_index if part_index is not None else summary_index, + ) + + +def _field_slot(holder: object, field: str) -> Sequence[_Slot]: + """The one restorable span at `field` on `holder`, if it holds text.""" + text: Final = _read_field(holder, field) + if not isinstance(text, str) or not text: + return () + return ((text, lambda new: _write_field(holder, field, new)),) + + class LLMShieldProxyGuardrail(CustomGuardrail): """Redacts PII before it leaves the proxy and restores it in the response. @@ -443,9 +857,16 @@ class LLMShieldProxyGuardrail(CustomGuardrail): The split exists because the response is restored against one vault only. Server-authored spans -- system and developer turns, Anthropic's top-level - `system`, the Responses API `instructions` -- go into a vault nothing is - ever restored against, so a caller who gets the model to echo one of their - placeholders back receives the placeholder, not the value behind it. + `system`, the Responses API `instructions`, tool definitions -- go into a + vault nothing is ever restored against, so a caller who gets the model to + echo one of their placeholders back receives the placeholder, not the value + behind it. End-user identifiers go there too: nothing in a reply needs them. + + Tool *results* stay on the caller's side deliberately. The model reads them in + order to answer, so it can already repeat anything in them; restoring the + placeholder gives the caller the answer they would have had without this + guardrail, and an agent that reads a file and quotes an address from it needs + that address back. """ slots: Final[_SlotSink] = [] # mutable-ok: accumulator, frozen on return. privileged: Final[_SlotSink] = [] # mutable-ok: accumulator, frozen on return. @@ -458,6 +879,8 @@ class LLMShieldProxyGuardrail(CustomGuardrail): _collect_responses_fields(data, slots, privileged) _collect_prompt(data, slots) _collect_system(data, privileged) + _collect_tool_definitions(data, privileged) + _collect_end_user_ids(data, privileged) return tuple(slots), tuple(privileged) # --- hooks -------------------------------------------------------------------- @@ -609,23 +1032,14 @@ class LLMShieldProxyGuardrail(CustomGuardrail): fields on the request side -- a function_call item holds `arguments`, a function_call_output holds `output` -- so the two directions stay symmetric. """ - slots: Final[list] = [] # mutable-ok: accumulator, frozen on return. + slots: Final[_SlotSink] = [] # mutable-ok: accumulator, frozen on return. for item in getattr(response, "output", None) or (): - for block in getattr(item, "content", None) or (): - text = _read_field(block, "text") - if isinstance(text, str) and text: - slots.append((text, lambda new, b=block: _write_field(b, "text", new))) - for field in ("arguments", "output"): - value = _read_field(item, field) - if isinstance(value, str) and value: - slots.append((value, lambda new, i=item, f=field: _write_field(i, f, new))) + _collect_response_item(item, slots) return tuple(slots) async def _restore_responses_api_response(self, response: Any, slots: Sequence[_Slot], data: MutableRequest) -> Any: """Puts the original values back into a Responses API reply.""" - restored: Final = await self._rehydrate(tuple(text for text, _ in slots), self._session_id(data)) - for (_, write), replacement in zip(slots, restored): - write(replacement) + await _rehydrate_slots(slots, functools.partial(self._rehydrate, session_id=self._session_id(data))) return response async def async_post_call_streaming_iterator_hook( @@ -641,6 +1055,10 @@ class LLMShieldProxyGuardrail(CustomGuardrail): window would splice the characters held back for one stream onto another. The windows are locals of this generator, so they are scoped to a single stream and cannot leak between concurrent requests. + + The two native stream shapes have no `choices` and are restored by their own + walkers, with the same per-stream windows: Anthropic `/v1/messages` arrives as raw + SSE frames, and the Responses API as typed events. """ if self.should_run_guardrail(data=request_data, event_type=GuardrailEventHooks.post_call) is not True: async for chunk in response: @@ -648,16 +1066,32 @@ class LLMShieldProxyGuardrail(CustomGuardrail): return session_id: Final = self._session_id(request_data) + step: Final = functools.partial(self._stream_step, session_id=session_id) + rehydrate: Final = functools.partial(self._rehydrate, session_id=session_id) + sse: Final = _AnthropicSSERestorer(step) + events: Final = _ResponsesStreamRestorer(step, rehydrate) carries: Final[dict] = {} # mutable-ok: per-stream windows, local to this generator. last_chunk = None # rebind-ok: tracks the most recent chunk for the final flush. async for chunk in response: + if isinstance(chunk, (bytes, str)): + for frames in await sse.feed(chunk): + yield frames + continue + if _responses_event_type(chunk) is not None: + for event in await events.restore(chunk): + yield event + continue last_chunk = chunk for choice in getattr(chunk, "choices", None) or (): await self._restore_choice(choice, carries, session_id) yield chunk - # A stream that ended without a finish_reason can still leave text held back. + # A stream that ended early can still leave text held back, in any shape. + for frames in await sse.finish(): + yield frames + for event in await events.finish(): + yield event if last_chunk is not None and any(carries.values()): async for trailing in self._flush_trailing(last_chunk, carries, session_id): yield trailing diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py index 0197b83c9f0..a24b3845f9b 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -13,6 +13,12 @@ from litellm.proxy.guardrails.guardrail_hooks.llm_shield_proxy.llm_shield_proxy ) from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2 from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.llms.openai import ( + FunctionCallArgumentsDeltaEvent, + OutputTextDeltaEvent, + OutputTextDoneEvent, + ResponsesAPIStreamEvents, +) from litellm.types.utils import Choices, Delta, Message, ModelResponse, ModelResponseStream, StreamingChoices @@ -66,6 +72,81 @@ def _field(holder: object, name: str) -> object: return holder.get(name) if isinstance(holder, dict) else getattr(holder, name) +class _FakeShield: + """The three guard endpoints over one fixed vault, placeholder -> original. + + The stream endpoint holds back a trailing `[` that has not closed yet, which is the + behaviour that makes a placeholder split across two chunks come out whole. + """ + + def __init__(self, vault: dict[str, str]) -> None: + self.vault = vault + self.urls: list[str] = [] + + def _restore(self, text: str) -> str: + for placeholder, original in self.vault.items(): + text = text.replace(placeholder, original) + return text + + async def post(self, url: str, headers: dict, json: dict, timeout: float) -> Response: + self.urls.append(url) + if url.endswith("/rehydrate/stream"): + text = self._restore(json["carry"] + json["text"]) + opening = text.rfind("[") + if json["final"] or opening == -1 or "]" in text[opening:]: + return _response({"text": text, "carry": ""}) + return _response({"text": text[:opening], "carry": text[opening:]}) + return _response({"texts": [self._restore(text) for text in json["texts"]]}) + + +def _shielded(vault: dict[str, str]) -> tuple[LLMShieldProxyGuardrail, _FakeShield]: + guardrail = _guardrail(event_hook="post_call") + shield = _FakeShield(vault) + guardrail.async_handler.post = shield.post # type: ignore[method-assign] + return guardrail, shield + + +def _sse(event: dict) -> bytes: + return f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() + + +def _sse_events(frames: list) -> list[dict]: + """Parses emitted SSE output, whatever its chunking, back into event payloads.""" + raw = b"".join(frame.encode() if isinstance(frame, str) else frame for frame in frames).decode() + return [ + json.loads(line[len("data:") :]) + for event in raw.split("\n\n") + for line in event.split("\n") + if line.startswith("data:") + ] + + +def _text_block_stream(*deltas: str) -> list[bytes]: + """An Anthropic /v1/messages stream with one text block made of `deltas`.""" + return [ + _sse({"type": "message_start", "message": {"id": "msg_1", "role": "assistant", "content": []}}), + _sse({"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}), + *( + _sse({"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": d}}) + for d in deltas + ), + _sse({"type": "content_block_stop", "index": 0}), + _sse({"type": "message_stop"}), + ] + + +async def _restore_stream(guardrail: LLMShieldProxyGuardrail, chunks: list) -> list: + async def stream(): + for chunk in chunks: + yield chunk + + return await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + 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", {}) @@ -479,6 +560,79 @@ class TestRequestCoverage: assert data["messages"][2]["tool_calls"][0]["function"]["arguments"] == "c" assert data["input"] == "d" + @pytest.mark.asyncio + async def test_anthropic_tool_use_input_is_redacted(self): + """A replayed tool_use block carries its arguments as a JSON object, not a string.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]", "[PHONE_1]"]}) + + data = { + "messages": [ + { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": "t1", + "name": "send", + "input": {"to": "jane.doe@example.com", "meta": {"phone": "555-0100"}}, + } + ], + } + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["jane.doe@example.com", "555-0100"] + block = data["messages"][0]["content"][0] + assert block["input"] == {"to": "[EMAIL_1]", "meta": {"phone": "[PHONE_1]"}} + assert block["name"] == "send", "the tool name has to arrive unchanged for the call to route" + + @pytest.mark.asyncio + async def test_responses_reasoning_summary_is_redacted(self): + """A replayed reasoning item quotes the conversation in its summary parts.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["user asked about [EMAIL_1]"]}) + + data = { + "input": [ + { + "type": "reasoning", + "id": "rs_1", + "summary": [{"type": "summary_text", "text": "user asked about jane.doe@example.com"}], + } + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert data["input"][0]["summary"][0]["text"] == "user asked about [EMAIL_1]" + + def test_tool_schemas_give_up_descriptions_and_nothing_else(self): + """Only free text is collected; names, types and enum values must reach the model.""" + data = { + "tools": [ + { + "type": "function", + "function": { + "name": "lookup", + "description": "top", + "parameters": { + "type": "object", + "properties": { + # A property that is itself named "description". + "description": {"type": "string", "description": "named"}, + "kind": {"type": "string", "enum": ["a", "b"], "description": "enum"}, + "deep": {"type": "array", "items": {"type": "object", "description": "nested"}}, + }, + }, + }, + } + ] + } + _, privileged = LLMShieldProxyGuardrail._locate_request_texts(data) + + assert sorted(text for text, _ in privileged) == ["enum", "named", "nested", "top"] + class TestRestoration: @pytest.mark.asyncio @@ -653,6 +807,30 @@ class TestVaultIsolation: id="anthropic-top-level-system", ), pytest.param({"instructions": "S", "input": "U"}, id="responses-instructions"), + pytest.param( + { + "messages": [{"role": "user", "content": "U"}], + "tools": [{"type": "function", "function": {"name": "f", "description": "S"}}], + }, + id="chat-tool-description", + ), + pytest.param( + {"input": "U", "tools": [{"type": "function", "name": "f", "description": "S"}]}, + id="responses-tool-description", + ), + pytest.param( + {"messages": [{"role": "user", "content": "U"}], "functions": [{"name": "f", "description": "S"}]}, + id="legacy-function-description", + ), + pytest.param( + { + "messages": [{"role": "user", "content": "U"}], + "tools": [{"name": "f", "input_schema": {"properties": {"to": {"description": "S"}}}}], + }, + id="anthropic-schema-description", + ), + pytest.param({"messages": [{"role": "user", "content": "U"}], "user": "S"}, id="end-user-id"), + pytest.param({"input": "U", "safety_identifier": "S"}, id="safety-identifier"), ], ) def test_server_authored_text_is_split_from_the_callers(self, data: dict) -> None: @@ -1016,3 +1194,196 @@ class TestApplyGuardrailToolCalls: assert merged["tool_calls"][0]["function"]["arguments"] == '{"email": "a@b.com"}' assert inputs["tool_calls"][0]["function"]["arguments"] == '{"email": "[EMAIL_1]"}' + + +class TestAnthropicStreamRestoration: + """/v1/messages streams reach the hook as raw SSE frames, with no `choices` to walk.""" + + VAULT = {"[EMAIL_1]": "a@example.com"} + + @pytest.mark.asyncio + async def test_split_placeholder_is_restored_and_never_fragmented(self): + guardrail, _ = _shielded(self.VAULT) + + out = await _restore_stream(guardrail, _text_block_stream("Mail [EMA", "IL_1] now")) + + deltas = [e["delta"]["text"] for e in _sse_events(out) if e["type"] == "content_block_delta"] + assert "".join(deltas) == "Mail a@example.com now" + assert not any("[EMA" in d for d in deltas), "a placeholder fragment reached the client" + + @pytest.mark.asyncio + async def test_held_text_lands_before_its_block_stops(self): + """A trailing `[` that never became a placeholder is still part of the answer.""" + guardrail, _ = _shielded(self.VAULT) + + out = await _restore_stream(guardrail, _text_block_stream("Mail [EMAIL_1], x = a[")) + + types = [e["type"] for e in _sse_events(out)] + deltas = [e["delta"]["text"] for e in _sse_events(out) if e["type"] == "content_block_delta"] + assert "".join(deltas) == "Mail a@example.com, x = a[" + assert types.index("content_block_stop") > max(i for i, t in enumerate(types) if t == "content_block_delta") + + @pytest.mark.asyncio + async def test_events_split_across_network_chunks_are_restored(self): + """A chunk can end mid-event; the frame is parsed once it is whole.""" + guardrail, _ = _shielded(self.VAULT) + raw = b"".join(_text_block_stream("Mail [EMA", "IL_1] now")) + + out = await _restore_stream(guardrail, [raw[i : i + 7] for i in range(0, len(raw), 7)]) + + deltas = [e["delta"]["text"] for e in _sse_events(out) if e["type"] == "content_block_delta"] + assert "".join(deltas) == "Mail a@example.com now" + + @pytest.mark.asyncio + async def test_str_frames_stay_str(self): + guardrail, _ = _shielded(self.VAULT) + + out = await _restore_stream(guardrail, [frame.decode() for frame in _text_block_stream("[EMAIL_1]")]) + + assert all(isinstance(frame, str) for frame in out) + assert "a@example.com" in "".join(out) + + @pytest.mark.asyncio + async def test_tool_input_json_is_restored(self): + guardrail, _ = _shielded(self.VAULT) + frames = [ + _sse({"type": "content_block_start", "index": 1, "content_block": {"type": "tool_use", "input": {}}}), + *( + _sse( + { + "type": "content_block_delta", + "index": 1, + "delta": {"type": "input_json_delta", "partial_json": p}, + } + ) + for p in ('{"to": "[EMAI', 'L_1]"}') + ), + _sse({"type": "content_block_stop", "index": 1}), + ] + + out = await _restore_stream(guardrail, frames) + + partial = "".join(e["delta"]["partial_json"] for e in _sse_events(out) if e["type"] == "content_block_delta") + assert json.loads(partial) == {"to": "a@example.com"} + + @pytest.mark.asyncio + async def test_signed_thinking_and_foreign_frames_pass_through_byte_for_byte(self): + """Rewriting a signed thinking block breaks it; other frames are not ours to touch.""" + guardrail, shield = _shielded(self.VAULT) + thinking = {"type": "thinking_delta", "thinking": "[EMAIL_1]"} + frames = [ + _sse({"type": "content_block_delta", "index": 0, "delta": thinking}), + b'data: {"candidates": [{"content": {"parts": [{"text": "[EMAIL_1]"}]}}]}\n\n', + b"data: not json\n\n", + ] + + out = await _restore_stream(guardrail, frames) + + assert b"".join(out) == b"".join(frames) + assert shield.urls == [] + + @pytest.mark.asyncio + async def test_a_raw_stream_that_is_not_sse_is_never_buffered(self): + """Without event boundaries to wait for, buffering would hold the whole reply.""" + guardrail, _ = _shielded(self.VAULT) + chunks = [b'[{"candidates": []}', b', {"candidates": []}]'] + + out = await _restore_stream(guardrail, chunks) + + assert out == chunks + + +class TestResponsesStreamRestoration: + """/v1/responses streams are typed events, with no `choices` to walk.""" + + VAULT = {"[EMAIL_1]": "a@example.com"} + + @staticmethod + def _text_delta(delta: str, sequence_number: int, content_index: int = 0) -> OutputTextDeltaEvent: + return OutputTextDeltaEvent( + type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, + item_id="msg_1", + output_index=0, + content_index=content_index, + delta=delta, + sequence_number=sequence_number, + ) + + @pytest.mark.asyncio + async def test_deltas_and_done_text_are_restored(self): + guardrail, _ = _shielded(self.VAULT) + done = OutputTextDoneEvent( + type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE, + item_id="msg_1", + output_index=0, + content_index=0, + text="Mail [EMAIL_1] x[", + ) + + out = await _restore_stream( + guardrail, [self._text_delta("Mail [EMA", 1), self._text_delta("IL_1] x[", 2), done] + ) + + deltas = [e.delta for e in out if isinstance(e, OutputTextDeltaEvent)] + assert not any("[EMA" in d for d in deltas), "a placeholder fragment reached the client" + assert "".join(deltas) == "Mail a@example.com x[" + assert out[-1].text == "Mail a@example.com x[", "the done event repeats the full, restored text" + assert isinstance(out[-2], OutputTextDeltaEvent), "held text must land before the done event" + + @pytest.mark.asyncio + async def test_function_call_arguments_are_restored(self): + guardrail, _ = _shielded(self.VAULT) + events = [ + FunctionCallArgumentsDeltaEvent( + type=ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA, + item_id="fc_1", + output_index=1, + delta=part, + ) + for part in ('{"to": "[EMAI', 'L_1]"}') + ] + + out = await _restore_stream(guardrail, events) + + assert json.loads("".join(e.delta for e in out)) == {"to": "a@example.com"} + + @pytest.mark.asyncio + async def test_a_truncated_stream_still_flushes(self): + """No done event at all: whatever the window holds goes out at the end.""" + guardrail, _ = _shielded(self.VAULT) + + out = await _restore_stream(guardrail, [self._text_delta("see [EMAIL_1] a[", 1)]) + + assert "".join(e.delta for e in out) == "see a@example.com a[" + + @pytest.mark.asyncio + async def test_completed_response_is_restored(self): + """The terminal event repeats the whole reply, and clients read it as the answer.""" + guardrail, _ = _shielded(self.VAULT) + block = {"type": "output_text", "text": "Mail [EMAIL_1]"} + call = SimpleNamespace(type="function_call", arguments='{"to": "[EMAIL_1]"}') + completed = SimpleNamespace( + type="response.completed", + response=SimpleNamespace(output=[SimpleNamespace(content=[block]), call]), + ) + + await _restore_stream(guardrail, [completed]) + + assert block["text"] == "Mail a@example.com" + assert call.arguments == '{"to": "a@example.com"}' + + @pytest.mark.asyncio + async def test_streams_on_different_parts_do_not_share_a_window(self): + guardrail, _ = _shielded(self.VAULT) + events = [ + self._text_delta("one [EMA", 1), + self._text_delta("two", 2, content_index=1), + self._text_delta("IL_1]", 3), + ] + + out = await _restore_stream(guardrail, events) + + by_part: dict[int, str] = {} + for event in out: + by_part[event.content_index] = by_part.get(event.content_index, "") + event.delta + assert by_part == {0: "one a@example.com", 1: "two"} From 40c26178fae3b8a174fa8188a671571cd353b1d6 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Sat, 26 Sep 2026 09:31:05 -0500 Subject: [PATCH 28/34] fix(guardrails): redact llm_shield_proxy predicted outputs and output schemas `prediction.content` is the caller's own draft of the reply, so it is redacted into the caller vault and restored with the reply. The descriptions in a structured-output schema (Chat response_format.json_schema, Responses text.format) are application authored like tool schemas, so they go to the non-restorable vault. --- .../llm_shield_proxy/llm_shield_proxy.py | 24 +++++++++++++++- .../guardrail_hooks/test_llm_shield_proxy.py | 28 +++++++++++++++++++ 2 files changed, 51 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index ebafb10831a..01770f83141 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -332,6 +332,27 @@ def _collect_schema_descriptions(schema: object, privileged: _SlotSink) -> None: pending.extend((value, depth + 1) for value in node if isinstance(value, (dict, list))) +def _collect_output_contracts(data: MutableRequest, slots: _SlotSink, privileged: _SlotSink) -> None: + """Text the caller sends to shape the reply rather than to prompt it. + + A predicted output (`prediction.content`) is the caller's own draft of the answer, so + it goes with their text: the model largely repeats it, and it has to come back. A + structured-output schema -- Chat `response_format.json_schema`, Responses + `text.format` -- is application-authored like a tool schema, so its descriptions go + to the privileged sink, and its names and types stay as sent. + """ + prediction: Final = data.get("prediction") + if isinstance(prediction, dict): + _collect(prediction, "content", slots) + _collect_text_parts(prediction, "content", slots) + response_format: Final = data.get("response_format") + if isinstance(response_format, dict): + _collect_schema_descriptions(response_format.get("json_schema"), privileged) + text_options: Final = data.get("text") + if isinstance(text_options, dict): + _collect_schema_descriptions(text_options.get("format"), privileged) + + def _collect_end_user_ids(data: MutableRequest, privileged: _SlotSink) -> None: """`user` and `safety_identifier` are forwarded to the provider and often hold an email. @@ -857,7 +878,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): The split exists because the response is restored against one vault only. Server-authored spans -- system and developer turns, Anthropic's top-level - `system`, the Responses API `instructions`, tool definitions -- go into a + `system`, the Responses API `instructions`, tool and output schemas -- go into a vault nothing is ever restored against, so a caller who gets the model to echo one of their placeholders back receives the placeholder, not the value behind it. End-user identifiers go there too: nothing in a reply needs them. @@ -880,6 +901,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): _collect_prompt(data, slots) _collect_system(data, privileged) _collect_tool_definitions(data, privileged) + _collect_output_contracts(data, slots, privileged) _collect_end_user_ids(data, privileged) return tuple(slots), tuple(privileged) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py index a24b3845f9b..64768b39070 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -829,6 +829,34 @@ class TestVaultIsolation: }, id="anthropic-schema-description", ), + pytest.param( + { + "messages": [{"role": "user", "content": "U"}], + "response_format": { + "type": "json_schema", + "json_schema": {"name": "n", "description": "S", "schema": {"type": "object"}}, + }, + }, + id="chat-response-format", + ), + pytest.param( + { + "input": "U", + "text": { + "format": { + "type": "json_schema", + "name": "n", + "schema": {"properties": {"a": {"description": "S"}}}, + } + }, + }, + id="responses-text-format", + ), + pytest.param({"prediction": {"type": "content", "content": "U"}, "instructions": "S"}, id="prediction"), + pytest.param( + {"prediction": {"type": "content", "content": [{"type": "text", "text": "U"}]}, "instructions": "S"}, + id="prediction-parts", + ), pytest.param({"messages": [{"role": "user", "content": "U"}], "user": "S"}, id="end-user-id"), pytest.param({"input": "U", "safety_identifier": "S"}, id="safety-identifier"), ], From d7d608fc528a1c16907f440494547ce89c0cad62 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Sat, 26 Sep 2026 10:04:13 -0500 Subject: [PATCH 29/34] fix(guardrails): fail closed on deep llm_shield_proxy requests, widen coverage - Request walks no longer skip what lies past their depth bound. Content nested past it, and tool inputs or schemas past the new JSON bound, now block the request instead of reaching the provider unredacted. The old depth test asserted the skip; it now asserts the block. - Tool and output schemas are walked by their JSON Schema structure, and give up `title`, `examples` and `default` as well as `description`. `enum` and `const` still go out as sent. - Responses events are matched by shape: any `*.delta` with a string delta is a token stream, and any `*.done` restores every non-identifier text field plus the `part` or `item` it repeats. This covers reasoning_summary_part.done and MCP arguments, and future families. Audio deltas are left alone. - An SSE stream whose first chunk ends partway through a field name (`b"eve"`) is no longer taken for a non-SSE stream. --- .../llm_shield_proxy/llm_shield_proxy.py | 287 ++++++++++++------ .../guardrail_hooks/test_llm_shield_proxy.py | 148 +++++++-- 2 files changed, 317 insertions(+), 118 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index 01770f83141..80d4446339a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -82,13 +82,15 @@ JsonBody: TypeAlias = dict # One redactable span: the text as it stands, and the write that puts the # replacement back where it came from. # How far a tool_result chain is followed. Real payloads nest one or two deep; the -# bound is what stops a crafted one from becoming an unbounded walk. +# bound is what stops a crafted one from becoming an unbounded walk. A request that +# nests deeper is refused rather than forwarded, because text past the bound would +# otherwise reach the provider unredacted. _MAX_CONTENT_DEPTH: Final = 8 -# How far a tool's parameter schema is followed. Deeper than content: every nested -# object costs two levels (`properties`, then the property), and a description missed -# here goes to the provider in the clear. -_MAX_SCHEMA_DEPTH: Final = 32 +# How far a JSON value -- a tool input, a parameter schema -- is followed on the request +# side. Legitimate JSON nests far deeper than content blocks do, so the bound is +# generous; past it the request is refused, for the same reason as above. +_MAX_JSON_DEPTH: Final = 64 _Slot: TypeAlias = tuple[str, Callable[[str], None]] # mutable-ok: Callable's param list. @@ -108,31 +110,50 @@ _ANTHROPIC_DELTA_FIELDS: Final = MappingProxyType({"text_delta": "text", "input_ # `data:` line split across two network chunks is parsed only once it is whole. _SSE_EVENT_BOUNDARY: Final = re.compile(rb"(\r?\n\r?\n)") -# Responses API events whose `delta` is model text. Each belongs to the stream that the -# matching `.done` event in `_RESPONSES_DONE_FIELDS` closes. -_RESPONSES_DELTA_EVENTS: Final = frozenset( - ( - "response.output_text.delta", - "response.refusal.delta", - "response.function_call_arguments.delta", - "response.reasoning_summary_text.delta", - ) -) +# What an SSE stream can open with: one of its fields, or a `:` comment. +_SSE_OPENINGS: Final = (b"event:", b"data:", b"id:", b"retry:", b":") -# The `.done` event that closes each delta stream, and the field that repeats the -# stream's full text on it. -_RESPONSES_DONE_FIELDS: Final = MappingProxyType( - { - "response.output_text.done": "text", - "response.refusal.done": "refusal", - "response.function_call_arguments.done": "arguments", - "response.reasoning_summary_text.done": "text", - } +# Responses API delta events whose `delta` is not text. Audio arrives base64-encoded; +# sending it through the shield would cost a round trip per chunk to restore nothing. +_RESPONSES_BINARY_DELTAS: Final = frozenset(("response.audio.delta",)) + +# Fields on a Responses API event that identify something rather than say something. +# Every other string field on a `.done` event is model text and is restored, so an event +# type added upstream is covered by default instead of leaking a placeholder. +_RESPONSES_STRUCTURAL_FIELDS: Final = frozenset( + ("type", "id", "item_id", "call_id", "name", "server_label", "status", "obfuscation") ) # Terminal Responses API events that repeat the whole reply under `response`. _RESPONSES_TERMINAL_EVENTS: Final = frozenset(("response.completed", "response.incomplete")) +# JSON Schema keywords that hold free text an application writes, and so can hold PII. +# `enum` and `const` are deliberately absent: the model has to reproduce those values +# exactly, and one redacted into the non-restorable vault would come back as a stand-in. +_SCHEMA_TEXT_KEYWORDS: Final = frozenset(("description", "title")) +_SCHEMA_VALUE_KEYWORDS: Final = frozenset(("examples", "default")) + +# JSON Schema keywords whose value is a map of name -> subschema, a single subschema, or +# a list of subschemas. Knowing which is which is what lets the walk tell a property +# *named* "description" apart from the `description` keyword. +_SCHEMA_MAP_KEYWORDS: Final = frozenset(("properties", "patternProperties", "$defs", "definitions", "dependentSchemas")) +_SCHEMA_KEYWORDS: Final = frozenset( + ( + "items", + "additionalProperties", + "additionalItems", + "unevaluatedProperties", + "unevaluatedItems", + "propertyNames", + "contains", + "not", + "if", + "then", + "else", + ) +) +_SCHEMA_LIST_KEYWORDS: Final = frozenset(("allOf", "anyOf", "oneOf", "prefixItems")) + # The accumulator the collectors below append into. It never escapes # _locate_request_texts, which freezes it into a tuple before returning. _SlotSink: TypeAlias = list[_Slot] # mutable-ok: accumulator passed between collectors. @@ -184,12 +205,21 @@ def _collect_prompt(data: MutableRequest, slots: _SlotSink) -> None: _collect_entry(prompt, index, slots) +class _RequestTooDeep(Exception): + """A request nests text past a walk's bound. + + Skipping the rest would forward it unredacted while the guardrail reports as + enabled, so the pre-call hook refuses the request instead. + """ + + def _collect_content(container: MutableRequest, slots: _SlotSink) -> None: """Collects `content`, a string or a list of typed parts. An Anthropic tool_result nests its own content, so this has to descend. It walks with an explicit stack and a depth bound rather than by recursion: the nesting is - caller controlled, and an unbounded descent is a JSON bomb. + caller controlled, and an unbounded descent is a JSON bomb. Content nested past the + bound raises `_RequestTooDeep` rather than being skipped. """ # Walked in document order: the shield maps its replies back by position, so the # order spans are collected in is part of the contract. @@ -202,8 +232,8 @@ def _collect_content(container: MutableRequest, slots: _SlotSink) -> None: if isinstance(content, str): _collect(node, "content", slots) continue - if depth >= _MAX_CONTENT_DEPTH: - continue + if depth >= _MAX_CONTENT_DEPTH and content: + raise _RequestTooDeep("content") for part in content if isinstance(content, list) else (): if not isinstance(part, dict): continue @@ -212,8 +242,9 @@ def _collect_content(container: MutableRequest, slots: _SlotSink) -> None: if part.get("type") == "tool_use": # A replayed Anthropic tool call. Its `input` is a JSON object rather than # a string, so a value can sit at any depth -- the reply side walks the - # same leaves when it restores one. - _collect_json_leaves(part.get("input"), slots, depth + 1) + # same leaves when it restores one. Its own JSON bound applies, not the + # content one, and past it the request is refused. + _collect_json_leaves(part.get("input"), slots, strict=True) if "content" in part: pending.append((part, depth + 1)) @@ -292,11 +323,11 @@ def _collect_text_parts(container: MutableRequest, key: str, slots: _SlotSink) - def _collect_tool_definitions(data: MutableRequest, privileged: _SlotSink) -> None: """Tool definitions are application-authored free text bound for the provider. - A description -- on the tool, or on any property of its parameter schema -- is where - callers put examples and customer context, so it carries PII as often as a prompt - does. It is collected into the privileged sink, like a system prompt: redacted - outbound, and never restorable from the reply. Names, types and enum values are left - as sent, because the model has to reproduce them exactly for a call to route. + A tool's description and the free text in its parameter schema are where callers put + examples and customer context, so they carry PII as often as a prompt does. They are + collected into the privileged sink, like a system prompt: redacted outbound, and never + restorable from the reply. Names, types, `enum` and `const` values are left as sent, + because the model has to reproduce them exactly for a call to route. Covers Chat `tools[].function`, the legacy `functions[]`, and the flat tool shape the Responses API and Anthropic share, whose schema is `parameters` or `input_schema`. @@ -309,27 +340,36 @@ def _collect_tool_definitions(data: MutableRequest, privileged: _SlotSink) -> No function = tool.get("function") for holder in (tool, function) if isinstance(function, dict) else (tool,): _collect(holder, "description", privileged) - _collect_schema_descriptions(holder.get("parameters"), privileged) - _collect_schema_descriptions(holder.get("input_schema"), privileged) + _collect_schema_text(holder.get("parameters"), privileged) + _collect_schema_text(holder.get("input_schema"), privileged) -def _collect_schema_descriptions(schema: object, privileged: _SlotSink) -> None: - """Collects every string `description` in a JSON schema, at any depth. +def _collect_schema_text(schema: object, privileged: _SlotSink) -> None: + """Collects the free text in a JSON Schema, at any depth. - Only `description` is free text. A property that is itself *named* "description" - holds a schema object rather than a string, so it is descended into, not collected. - Walked with an explicit stack and a depth bound, like the other request walks. + That is every `description` and `title` string, and every string inside `examples` + and `default`. The walk follows the schema's own structure -- `properties` and the + other subschema keywords -- rather than every nested dict, which is what tells a + property *named* "description" (a subschema, descended into) from the `description` + keyword (text, collected). Nested past `_MAX_JSON_DEPTH`, the request is refused. """ pending: Final[list] = [(schema, 0)] # mutable-ok: local walk stack. while pending: node, depth = pending.pop() - if depth > _MAX_SCHEMA_DEPTH: + if not isinstance(node, dict): continue - if isinstance(node, dict): - _collect(node, "description", privileged) - pending.extend((value, depth + 1) for value in node.values() if isinstance(value, (dict, list))) - elif isinstance(node, list): - pending.extend((value, depth + 1) for value in node if isinstance(value, (dict, list))) + if depth > _MAX_JSON_DEPTH: + raise _RequestTooDeep("schema") + for keyword, value in tuple(node.items()): + if keyword in _SCHEMA_TEXT_KEYWORDS: + _collect(node, keyword, privileged) + elif keyword in _SCHEMA_VALUE_KEYWORDS: + _collect(node, keyword, privileged) + _collect_json_leaves(value, privileged, strict=True) + elif keyword in _SCHEMA_MAP_KEYWORDS and isinstance(value, dict): + pending.extend((child, depth + 1) for child in value.values()) + elif keyword in _SCHEMA_KEYWORDS or keyword in _SCHEMA_LIST_KEYWORDS: + pending.extend((child, depth + 1) for child in (value if isinstance(value, list) else (value,))) def _collect_output_contracts(data: MutableRequest, slots: _SlotSink, privileged: _SlotSink) -> None: @@ -338,7 +378,7 @@ def _collect_output_contracts(data: MutableRequest, slots: _SlotSink, privileged A predicted output (`prediction.content`) is the caller's own draft of the answer, so it goes with their text: the model largely repeats it, and it has to come back. A structured-output schema -- Chat `response_format.json_schema`, Responses - `text.format` -- is application-authored like a tool schema, so its descriptions go + `text.format` -- is application-authored like a tool schema, so its free text goes to the privileged sink, and its names and types stay as sent. """ prediction: Final = data.get("prediction") @@ -346,11 +386,14 @@ def _collect_output_contracts(data: MutableRequest, slots: _SlotSink, privileged _collect(prediction, "content", slots) _collect_text_parts(prediction, "content", slots) response_format: Final = data.get("response_format") - if isinstance(response_format, dict): - _collect_schema_descriptions(response_format.get("json_schema"), privileged) text_options: Final = data.get("text") - if isinstance(text_options, dict): - _collect_schema_descriptions(text_options.get("format"), privileged) + for wrapper in ( + response_format.get("json_schema") if isinstance(response_format, dict) else None, + text_options.get("format") if isinstance(text_options, dict) else None, + ): + if isinstance(wrapper, dict): + _collect(wrapper, "description", privileged) + _collect_schema_text(wrapper.get("schema"), privileged) def _collect_end_user_ids(data: MutableRequest, privileged: _SlotSink) -> None: @@ -400,20 +443,26 @@ def _write_field(holder: object, name: str, value: str) -> None: setattr(holder, name, value) -def _collect_json_leaves(node: object, slots: _SlotSink, depth: int = 0) -> None: +def _collect_json_leaves(node: object, slots: _SlotSink, *, strict: bool = False) -> None: """Collects every string leaf of a JSON-ish structure, with a write-back per leaf. An Anthropic `tool_use` block carries `input`, an arbitrary JSON object rather than a - string, so a value worth restoring can sit at any depth. Bounded by - `_MAX_CONTENT_DEPTH` for the same reason the request walk is: the shape is model - controlled, and the bound is what stops a crafted one from becoming an unbounded - descent. Walked with an explicit stack rather than recursively, so a deeply nested - tool input cannot spend stack frames proportional to attacker-chosen depth. + string, so a value worth restoring can sit at any depth. Bounded by `_MAX_JSON_DEPTH`: + the shape is caller or model controlled, and the bound is what stops a crafted one from + becoming an unbounded descent. Walked with an explicit stack rather than recursively, + so a deeply nested value cannot spend stack frames proportional to attacker-chosen + depth. + + `strict` is for the request side, where a leaf left behind would reach the provider + unredacted: past the bound it raises `_RequestTooDeep`. On the reply side a leaf past + the bound just keeps its placeholder, which leaks nothing, so it is skipped. """ - pending: Final[list] = [(node, depth)] # mutable-ok: local walk stack. + pending: Final[list] = [(node, 0)] # mutable-ok: local walk stack. while pending: current, current_depth = pending.pop() - if current_depth > _MAX_CONTENT_DEPTH: + if current_depth > _MAX_JSON_DEPTH: + if strict and isinstance(current, (dict, list)) and current: + raise _RequestTooDeep("json") continue if isinstance(current, dict): for key in tuple(current): @@ -453,9 +502,10 @@ def _continuation_delta(tool_index: int, text: str) -> list[dict[str, object]]: def _collect_response_item(item: object, slots: _SlotSink) -> None: """Restorable spans in one Responses API output item, dict or object. - Mirrors `_collect_responses_fields` on the request side -- a function_call item holds - `arguments`, a function_call_output holds `output`, a reasoning item holds `summary` - parts -- so the two directions stay symmetric. + Mirrors `_collect_responses_fields` on the request side -- a function_call or + mcp_call item holds `arguments`, their outputs `output`, a reasoning item `summary` + parts -- so the two directions stay symmetric. A custom tool call carries `input` and + a code interpreter call `code`, both model-written. """ for block in _read_list(item, "content"): for field in ("text", "refusal"): @@ -466,7 +516,7 @@ def _collect_response_item(item: object, slots: _SlotSink) -> None: text = _read_field(part, "text") if isinstance(text, str) and text: slots.append((text, lambda new, p=part: _write_field(p, "text", new))) - for field in ("arguments", "output"): + for field in ("arguments", "output", "input", "code"): value = _read_field(item, field) if isinstance(value, str) and value: slots.append((value, lambda new, i=item, f=field: _write_field(i, f, new))) @@ -481,6 +531,23 @@ async def _rehydrate_slots(slots: Sequence[_Slot], rehydrate: _Rehydrate) -> Non write(replacement) +def _opens_like_sse(head: bytes) -> bool | None: + """Whether a raw stream is SSE, judged by its opening bytes; None while undecidable. + + An SSE stream opens with a field name or a `:` comment. Anything else -- a JSON array + streamed in pieces, say -- has no event boundaries to wait for. A chunk that ends + partway through a field name decides nothing yet, so that case waits for more. + """ + opening: Final = head.lstrip() + if not opening: + return None + if opening.startswith(_SSE_OPENINGS): + return True + if any(field.startswith(opening) for field in _SSE_OPENINGS): + return None + return False + + def _responses_event_type(chunk: object) -> str | None: """The event type of a Responses API stream event, or None for any other chunk. @@ -517,20 +584,25 @@ class _AnthropicSSERestorer: self._delta_types: Final[dict[int, str]] = {} # mutable-ok: each block's delta type, for its flush. self._pending = b"" # rebind-ok: the unfinished tail of the stream. self._as_text = False # rebind-ok: set once if the stream arrives as str rather than bytes. - self._is_sse: bool | None = None # rebind-ok: decided once, by the stream's first chunk. + self._is_sse: bool | None = None # rebind-ok: undecided until the opening bytes settle it. async def feed(self, chunk: bytes | str) -> tuple[bytes | str, ...]: """Restores every event this chunk completes; holds back an unfinished tail.""" if isinstance(chunk, str): self._as_text = True - raw: Final = chunk.encode("utf-8") if isinstance(chunk, str) else chunk - if self._is_sse is None and raw.strip(): - # An SSE stream opens with a field or a comment. Anything else (a JSON array - # streamed in pieces, say) has no event boundaries to wait for. - self._is_sse = raw.lstrip().startswith((b"event:", b"data:", b":")) - if not self._is_sse: + if self._is_sse is False: return (chunk,) + raw: Final = chunk.encode("utf-8") if isinstance(chunk, str) else chunk buffered: Final = self._pending + raw + if self._is_sse is None: + self._is_sse = _opens_like_sse(buffered) + if self._is_sse is None: + # Too little has arrived to tell -- `b"eve"` could still become `event:`. + self._pending = buffered + return () + if not self._is_sse: + self._pending = b"" + return self._emit(buffered) boundaries: Final = tuple(_SSE_EVENT_BOUNDARY.finditer(buffered)) if not boundaries: self._pending = buffered @@ -549,8 +621,12 @@ class _AnthropicSSERestorer: async def finish(self) -> tuple[bytes | str, ...]: """Emits an unterminated final event and any window a block never closed.""" - tail: Final = await self._restore_event(self._pending) if self._pending.strip() else self._pending + held: Final = self._pending self._pending = b"" + if not self._is_sse: + # The stream ended before it could be told apart from SSE: hand it back as is. + return self._emit(held) + tail: Final = await self._restore_event(held) if held.strip() else held flushed: Final = await self._flush_all() # The tail had no blank line after it; one is needed before another frame follows. separator: Final = b"\n\n" if tail.strip() and flushed else b"" @@ -637,16 +713,19 @@ class _AnthropicSSERestorer: class _ResponsesStreamRestorer: """Restores a Responses API event stream. - Every delta stream -- one output_text content part, one refusal, one function call's - arguments, one reasoning summary part -- gets its own window, keyed by the event - family, the item id and the part index. When its `.done` event arrives, whatever the - window still holds goes out first, as a copy of that stream's last delta event -- so - it carries the stream's own ids, and repeats that event's `sequence_number` -- and - the `.done` event's full text is then restored in one call. + The event families are matched by shape rather than listed, so a text stream the + API adds later is restored by default instead of leaking a placeholder: - The events that repeat the reply wholesale -- `content_part.done`, - `output_item.done`, and `response.completed` / `response.incomplete` -- are restored - the same way the non-streaming reply is. + - Any `*.delta` event whose `delta` is a string is a token stream (output_text, + refusal, function-call and MCP arguments, reasoning summaries, ...). Each gets its + own window, keyed by the family, the item id and the part index. + - Any `*.done` event closes the stream of the same family. Whatever its window still + holds goes out first, as a copy of that stream's last delta event -- so it carries + the stream's own ids, and repeats that event's `sequence_number`. Then every text + field on the done event is restored in full: its string fields other than + identifiers, plus any `part` or `item` it repeats. + - `response.completed` / `response.incomplete` repeat the whole reply, and are + restored the same way the non-streaming reply is. """ def __init__(self, step: _StreamStep, rehydrate: _Rehydrate) -> None: @@ -660,25 +739,22 @@ class _ResponsesStreamRestorer: kind: Final = _responses_event_type(event) if kind is None: return (event,) - if kind in _RESPONSES_DELTA_EVENTS: + if kind.endswith(".delta") and kind not in _RESPONSES_BINARY_DELTAS: await self._restore_delta(event, kind) return (event,) - done_field: Final = _RESPONSES_DONE_FIELDS.get(kind) - if done_field is not None: - flushed: Final = await self._flush(_responses_stream_key(event, kind)) - await _rehydrate_slots(_field_slot(event, done_field), self._rehydrate) - return (*flushed, event) slots: Final[_SlotSink] = [] # mutable-ok: accumulator, restored in one batch. - if kind == "response.content_part.done": + flushed: Final = await self._flush(_responses_stream_key(event, kind)) if kind.endswith(".done") else () + if kind.endswith(".done"): + _collect_event_text(event, slots) part: Final = _read_field(event, "part") - _collect_response_item({"content": [part]}, slots) # mutable-ok: a one-part view. - elif kind == "response.output_item.done": + if part is not None: + _collect_response_item({"content": [part]}, slots) # mutable-ok: a one-part view. _collect_response_item(_read_field(event, "item"), slots) elif kind in _RESPONSES_TERMINAL_EVENTS: for item in _read_list(_read_field(event, "response"), "output"): _collect_response_item(item, slots) await _rehydrate_slots(slots, self._rehydrate) - return (event,) + return (*flushed, event) async def finish(self) -> tuple[object, ...]: """Flushes every stream the provider never closed, e.g. a truncated reply.""" @@ -725,12 +801,21 @@ def _responses_stream_key(event: object, kind: str) -> tuple: ) -def _field_slot(holder: object, field: str) -> Sequence[_Slot]: - """The one restorable span at `field` on `holder`, if it holds text.""" - text: Final = _read_field(holder, field) - if not isinstance(text, str) or not text: - return () - return ((text, lambda new: _write_field(holder, field, new)),) +def _collect_event_text(event: object, slots: _SlotSink) -> None: + """Collects every top-level text field of a Responses API event, dict or model. + + Scan by default, with identifiers excluded, rather than a list of known fields: the + `.done` event of each stream family names its text differently (`text`, `refusal`, + `arguments`, ...), and a family added upstream would otherwise leak a placeholder. + """ + fields: Final = event if isinstance(event, dict) else getattr(event, "__dict__", None) + if not isinstance(fields, dict): + return + for name, value in tuple(fields.items()): + if not isinstance(name, str) or name in _RESPONSES_STRUCTURAL_FIELDS or name.endswith("_id"): + continue + if isinstance(value, str) and value: + slots.append((value, lambda new, n=name: _write_field(event, n, new))) class LLMShieldProxyGuardrail(CustomGuardrail): @@ -919,7 +1004,13 @@ class LLMShieldProxyGuardrail(CustomGuardrail): if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_call) is not True: return data - slots, privileged = self._locate_request_texts(data) + try: + slots, privileged = self._locate_request_texts(data) + except _RequestTooDeep as exc: + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=f"Request {exc} nests deeper than LLM Shield Proxy inspects; blocking the request.", + ) from exc if not slots and not privileged: return data diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py index 64768b39070..0929341028e 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -480,33 +480,59 @@ class TestRequestCoverage: assert data["messages"][0]["content"][1]["content"][0]["text"] == "[EMAIL_2]" @pytest.mark.asyncio - async def test_deeply_nested_tool_results_are_bounded(self): - """Nesting is caller controlled, so the descent has to stop somewhere. + async def test_nesting_past_the_bound_blocks_the_request(self): + """Nesting is caller controlled, so the descent has to stop somewhere -- and where + it stops, the request must not go out. - The walk must terminate on a payload built to be pathological, rather than - following it as far as it goes. + This test used to assert the opposite: that text past the bound was skipped. That + sent `past-the-bound@example.com` to the provider unredacted while the guardrail + reported as enabled. """ guardrail = _guardrail() - - captured: list = [] - - async def echo(url, headers, json, timeout): # noqa: ARG001 - captured.append(json["texts"]) - return _response({"texts": list(json["texts"])}) - - guardrail.async_handler.post = AsyncMock(side_effect=echo) # type: ignore[method-assign] + mock = _mock_post(guardrail) deep: dict = {"type": "tool_result", "content": "past-the-bound@example.com"} for _ in range(200): deep = {"type": "tool_result", "content": [deep]} data = {"messages": [{"role": "user", "content": [{"type": "text", "text": "shallow"}, deep]}]} + with pytest.raises(GuardrailRaisedException): + 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_deep_tool_input_blocks_the_request(self): + """A tool_use input past the JSON bound must not be forwarded half-redacted.""" + guardrail = _guardrail() + mock = _mock_post(guardrail) + + deep: dict = {"email": "past-the-bound@example.com"} + for _ in range(100): + deep = {"next": deep} + block = {"type": "tool_use", "id": "t1", "name": "f", "input": deep} + data = {"messages": [{"role": "assistant", "content": [block]}]} + + with pytest.raises(GuardrailRaisedException): + 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_realistic_nesting_is_redacted_in_full(self): + """The bounds are far past real payloads: a tool input nested inside a tool result, + several JSON levels deep, is redacted whole rather than refused.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + + tool_use = { + "type": "tool_use", + "id": "t1", + "name": "f", + "input": {"a": {"b": {"c": {"d": {"to": "x@example.com"}}}}}, + } + data = {"messages": [{"role": "user", "content": [{"type": "tool_result", "content": [tool_use]}]}]} await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") - sent = captured[0] - assert "shallow" in sent - assert "past-the-bound@example.com" not in sent, "the walk followed the chain past its bound" - assert len(sent) < 200 + assert tool_use["input"]["a"]["b"]["c"]["d"]["to"] == "[EMAIL_1]" @pytest.mark.asyncio async def test_responses_prompt_object_variables_are_redacted(self): @@ -607,8 +633,9 @@ class TestRequestCoverage: assert data["input"][0]["summary"][0]["text"] == "user asked about [EMAIL_1]" - def test_tool_schemas_give_up_descriptions_and_nothing_else(self): - """Only free text is collected; names, types and enum values must reach the model.""" + def test_tool_schemas_give_up_their_free_text_and_nothing_else(self): + """Descriptions, titles, examples and defaults are collected. Names, types, enum + and const values must reach the model exactly as sent.""" data = { "tools": [ { @@ -618,12 +645,16 @@ class TestRequestCoverage: "description": "top", "parameters": { "type": "object", + "title": "title", "properties": { # A property that is itself named "description". "description": {"type": "string", "description": "named"}, - "kind": {"type": "string", "enum": ["a", "b"], "description": "enum"}, + "kind": {"type": "string", "enum": ["a", "b"], "const": "a", "description": "enum"}, "deep": {"type": "array", "items": {"type": "object", "description": "nested"}}, + "to": {"type": "string", "examples": ["example"], "default": "default"}, + "choice": {"anyOf": [{"type": "object", "default": {"who": "object-default"}}]}, }, + "$defs": {"shared": {"description": "defined"}}, }, }, } @@ -631,7 +662,26 @@ class TestRequestCoverage: } _, privileged = LLMShieldProxyGuardrail._locate_request_texts(data) - assert sorted(text for text, _ in privileged) == ["enum", "named", "nested", "top"] + assert sorted(text for text, _ in privileged) == [ + "default", + "defined", + "enum", + "example", + "named", + "nested", + "object-default", + "title", + "top", + ] + + def test_schema_nesting_past_the_bound_is_refused(self): + schema: dict = {"type": "object", "description": "past-the-bound@example.com"} + for _ in range(100): + schema = {"type": "object", "properties": {"next": schema}} + data = {"tools": [{"type": "function", "function": {"name": "f", "parameters": schema}}]} + + with pytest.raises(Exception, match="schema"): + LLMShieldProxyGuardrail._locate_request_texts(data) class TestRestoration: @@ -1320,6 +1370,19 @@ class TestAnthropicStreamRestoration: assert out == chunks + @pytest.mark.asyncio + @pytest.mark.parametrize("cut", [1, 3, 5, 6]) + async def test_a_field_name_split_by_the_first_chunk_still_reads_as_sse(self, cut: int): + """`b"eve"` then `b"nt: ..."` is still SSE; deciding on the first chunk alone + would pass the whole stream through with its placeholders.""" + guardrail, _ = _shielded(self.VAULT) + raw = b"".join(_text_block_stream("Mail [EMAIL_1]")) + + out = await _restore_stream(guardrail, [raw[:cut], raw[cut:]]) + + deltas = [e["delta"]["text"] for e in _sse_events(out) if e["type"] == "content_block_delta"] + assert "".join(deltas) == "Mail a@example.com" + class TestResponsesStreamRestoration: """/v1/responses streams are typed events, with no `choices` to walk.""" @@ -1415,3 +1478,48 @@ class TestResponsesStreamRestoration: for event in out: by_part[event.content_index] = by_part.get(event.content_index, "") + event.delta assert by_part == {0: "one a@example.com", 1: "two"} + + @pytest.mark.asyncio + async def test_reasoning_summary_part_done_is_restored(self): + """The summary part repeats the whole summary text after its deltas.""" + guardrail, _ = _shielded(self.VAULT) + part = SimpleNamespace(type="summary_text", text="asked about [EMAIL_1]") + event = SimpleNamespace( + type="response.reasoning_summary_part.done", item_id="rs_1", output_index=0, summary_index=0, part=part + ) + + await _restore_stream(guardrail, [event]) + + assert part.text == "asked about a@example.com" + + @pytest.mark.asyncio + async def test_mcp_call_arguments_are_restored(self): + """A stream family outside the chat-era set: matched by shape, not by name.""" + guardrail, _ = _shielded(self.VAULT) + deltas = [ + {"type": "response.mcp_call_arguments.delta", "item_id": "mcp_1", "output_index": 0, "delta": d} + for d in ('{"to": "[EMAI', 'L_1]"}') + ] + done = { + "type": "response.mcp_call_arguments.done", + "item_id": "mcp_1", + "output_index": 0, + "arguments": '{"to": "[EMAIL_1]"}', + } + + out = await _restore_stream(guardrail, [*deltas, done]) + + assert json.loads("".join(e["delta"] for e in out[:-1])) == {"to": "a@example.com"} + assert json.loads(out[-1]["arguments"]) == {"to": "a@example.com"} + assert out[-1]["item_id"] == "mcp_1", "identifiers are not text and stay as sent" + + @pytest.mark.asyncio + async def test_audio_deltas_are_not_sent_to_the_shield(self): + """Audio arrives base64-encoded; restoring it would cost a round trip for nothing.""" + guardrail, shield = _shielded(self.VAULT) + audio = {"type": "response.audio.delta", "item_id": "a_1", "output_index": 0, "delta": "UklGRiQAAABXQVZF"} + + out = await _restore_stream(guardrail, [audio]) + + assert out == [audio] + assert shield.urls == [] From 1b81bb451d888ec55abd639ede75a1569b8a9b0a Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Sat, 26 Sep 2026 10:22:13 -0500 Subject: [PATCH 30/34] fix(guardrails): scan llm_shield_proxy schemas by default The schema walk collected an allowlist of keywords, so any keyword it did not list -- draft-07 `dependencies`, `$comment`, vendor `x-` extensions -- went to the provider in clear. Invert it: every string is collected except under keywords whose value must go out verbatim (types, formats, patterns, references, required lists, enum, const). Name -> subschema maps still treat their keys as property names, so a property called `type` is walked, not skipped. --- .../llm_shield_proxy/llm_shield_proxy.py | 98 ++++++++++++------- .../guardrail_hooks/test_llm_shield_proxy.py | 25 ++++- 2 files changed, 84 insertions(+), 39 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index 80d4446339a..ff5ca24c6e9 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -127,32 +127,45 @@ _RESPONSES_STRUCTURAL_FIELDS: Final = frozenset( # Terminal Responses API events that repeat the whole reply under `response`. _RESPONSES_TERMINAL_EVENTS: Final = frozenset(("response.completed", "response.incomplete")) -# JSON Schema keywords that hold free text an application writes, and so can hold PII. -# `enum` and `const` are deliberately absent: the model has to reproduce those values -# exactly, and one redacted into the non-restorable vault would come back as a stand-in. -_SCHEMA_TEXT_KEYWORDS: Final = frozenset(("description", "title")) -_SCHEMA_VALUE_KEYWORDS: Final = frozenset(("examples", "default")) - -# JSON Schema keywords whose value is a map of name -> subschema, a single subschema, or -# a list of subschemas. Knowing which is which is what lets the walk tell a property -# *named* "description" apart from the `description` keyword. -_SCHEMA_MAP_KEYWORDS: Final = frozenset(("properties", "patternProperties", "$defs", "definitions", "dependentSchemas")) -_SCHEMA_KEYWORDS: Final = frozenset( +# JSON Schema keywords whose value has to reach the model or a validator verbatim, so the +# schema walk leaves them alone: types, formats, patterns, references, required-property +# lists, and `enum` / `const`, which the model must reproduce exactly -- a value redacted +# into the non-restorable vault would come back as a stand-in and break the call. +# Everything else is scanned. +_SCHEMA_STRUCTURAL_KEYWORDS: Final = frozenset( ( - "items", - "additionalProperties", - "additionalItems", - "unevaluatedProperties", - "unevaluatedItems", - "propertyNames", - "contains", - "not", - "if", - "then", - "else", + "type", + "format", + "pattern", + "enum", + "const", + "required", + "dependentRequired", + "propertyOrdering", + "discriminator", + "contentEncoding", + "contentMediaType", + "$ref", + "$id", + "$schema", + "$anchor", + "$dynamicRef", + "$dynamicAnchor", + "$recursiveRef", + "$recursiveAnchor", + "$vocabulary", ) ) -_SCHEMA_LIST_KEYWORDS: Final = frozenset(("allOf", "anyOf", "oneOf", "prefixItems")) + +# Keywords holding JSON values rather than schemas: every string in them is collected, +# whatever the keys around it are called. +_SCHEMA_VALUE_KEYWORDS: Final = frozenset(("examples", "default")) + +# Keywords whose value maps names to subschemas. Their keys are property names, not +# keywords, so a property called `type` or `enum` is walked like any other subschema. +_SCHEMA_MAP_KEYWORDS: Final = frozenset( + ("properties", "patternProperties", "$defs", "definitions", "dependentSchemas", "dependencies") +) # The accumulator the collectors below append into. It never escapes # _locate_request_texts, which freezes it into a tuple before returning. @@ -347,29 +360,44 @@ def _collect_tool_definitions(data: MutableRequest, privileged: _SlotSink) -> No def _collect_schema_text(schema: object, privileged: _SlotSink) -> None: """Collects the free text in a JSON Schema, at any depth. - That is every `description` and `title` string, and every string inside `examples` - and `default`. The walk follows the schema's own structure -- `properties` and the - other subschema keywords -- rather than every nested dict, which is what tells a - property *named* "description" (a subschema, descended into) from the `description` - keyword (text, collected). Nested past `_MAX_JSON_DEPTH`, the request is refused. + Scan by default: every string is collected except under the keywords in + `_SCHEMA_STRUCTURAL_KEYWORDS`, whose values must go out verbatim. A list of keywords + *to* collect would leak every one it forgot -- draft-07 `dependencies`, a `$comment`, + a vendor `x-` extension -- which is how this walk started out. + + Structure matters in two places. Under `properties` and the other name -> subschema + maps, keys are property names rather than keywords, so a property called `type` is a + subschema to walk, not a keyword to skip. And `examples` / `default` hold JSON values, + so all their strings are collected whatever the keys around them are called. Nested + past `_MAX_JSON_DEPTH`, the request is refused. """ pending: Final[list] = [(schema, 0)] # mutable-ok: local walk stack. while pending: node, depth = pending.pop() + if depth > _MAX_JSON_DEPTH: + if isinstance(node, (dict, list)) and node: + raise _RequestTooDeep("schema") + continue + if isinstance(node, list): + for index, item in enumerate(node): + _collect_entry(node, index, privileged) + if isinstance(item, (dict, list)): + pending.append((item, depth + 1)) + continue if not isinstance(node, dict): continue - if depth > _MAX_JSON_DEPTH: - raise _RequestTooDeep("schema") for keyword, value in tuple(node.items()): - if keyword in _SCHEMA_TEXT_KEYWORDS: - _collect(node, keyword, privileged) - elif keyword in _SCHEMA_VALUE_KEYWORDS: + if keyword in _SCHEMA_STRUCTURAL_KEYWORDS: + continue + if keyword in _SCHEMA_VALUE_KEYWORDS: _collect(node, keyword, privileged) _collect_json_leaves(value, privileged, strict=True) elif keyword in _SCHEMA_MAP_KEYWORDS and isinstance(value, dict): pending.extend((child, depth + 1) for child in value.values()) - elif keyword in _SCHEMA_KEYWORDS or keyword in _SCHEMA_LIST_KEYWORDS: - pending.extend((child, depth + 1) for child in (value if isinstance(value, list) else (value,))) + elif isinstance(value, str): + _collect(node, keyword, privileged) + elif isinstance(value, (dict, list)): + pending.append((value, depth + 1)) def _collect_output_contracts(data: MutableRequest, slots: _SlotSink, privileged: _SlotSink) -> None: diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py index 0929341028e..475ab06994a 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -634,8 +634,8 @@ class TestRequestCoverage: assert data["input"][0]["summary"][0]["text"] == "user asked about [EMAIL_1]" def test_tool_schemas_give_up_their_free_text_and_nothing_else(self): - """Descriptions, titles, examples and defaults are collected. Names, types, enum - and const values must reach the model exactly as sent.""" + """Every string is collected except what must reach the model verbatim: names, + types, formats, patterns, required lists, enum and const values.""" data = { "tools": [ { @@ -651,10 +651,23 @@ class TestRequestCoverage: "description": {"type": "string", "description": "named"}, "kind": {"type": "string", "enum": ["a", "b"], "const": "a", "description": "enum"}, "deep": {"type": "array", "items": {"type": "object", "description": "nested"}}, - "to": {"type": "string", "examples": ["example"], "default": "default"}, - "choice": {"anyOf": [{"type": "object", "default": {"who": "object-default"}}]}, + "to": { + "type": "string", + "format": "email", + "pattern": "^.+@.+$", + "examples": ["example"], + "default": "default", + }, + "choice": {"anyOf": [{"type": "object", "default": {"type": "object-default"}}]}, + # Property names that collide with keywords are subschemas all the same. + "type": {"type": "string", "description": "named-type"}, }, + "required": ["to"], "$defs": {"shared": {"description": "defined"}}, + # Keywords nobody listed: scanned by default. + "dependencies": {"mode": {"description": "dependent"}}, + "$comment": "comment", + "x-note": "vendor", }, }, } @@ -663,15 +676,19 @@ class TestRequestCoverage: _, privileged = LLMShieldProxyGuardrail._locate_request_texts(data) assert sorted(text for text, _ in privileged) == [ + "comment", "default", "defined", + "dependent", "enum", "example", "named", + "named-type", "nested", "object-default", "title", "top", + "vendor", ] def test_schema_nesting_past_the_bound_is_refused(self): From 0f9ca3fd777c9bc1390a644f372e4f65f6f2c776 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Sat, 26 Sep 2026 10:39:16 -0500 Subject: [PATCH 31/34] fix(guardrails): redact llm_shield_proxy schema enum and const values `enum` and `const` were skipped by the schema walk, so a value holding PII went to the provider in clear. They now go to the caller's vault rather than the non-restorable one: the model emits the stand-in in its tool arguments or structured output, and restoring the reply turns it back into the value the schema allows, so the call still routes. --- .../llm_shield_proxy/llm_shield_proxy.py | 42 +++++++++++-------- .../guardrail_hooks/test_llm_shield_proxy.py | 37 +++++++++++++++- 2 files changed, 60 insertions(+), 19 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index ff5ca24c6e9..bd2796086db 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -128,17 +128,13 @@ _RESPONSES_STRUCTURAL_FIELDS: Final = frozenset( _RESPONSES_TERMINAL_EVENTS: Final = frozenset(("response.completed", "response.incomplete")) # JSON Schema keywords whose value has to reach the model or a validator verbatim, so the -# schema walk leaves them alone: types, formats, patterns, references, required-property -# lists, and `enum` / `const`, which the model must reproduce exactly -- a value redacted -# into the non-restorable vault would come back as a stand-in and break the call. -# Everything else is scanned. +# schema walk leaves them alone: types, formats, patterns, references and +# required-property lists. Everything else is scanned. _SCHEMA_STRUCTURAL_KEYWORDS: Final = frozenset( ( "type", "format", "pattern", - "enum", - "const", "required", "dependentRequired", "propertyOrdering", @@ -161,6 +157,13 @@ _SCHEMA_STRUCTURAL_KEYWORDS: Final = frozenset( # whatever the keys around it are called. _SCHEMA_VALUE_KEYWORDS: Final = frozenset(("examples", "default")) +# Keywords holding the literal values the model must reproduce. These go to the CALLER's +# vault, not the privileged one: the model emits the stand-in in its tool arguments or +# structured output, and restoring the reply turns it back into the value the schema +# allows, so the call still routes. In the non-restorable vault it would come back as a +# stand-in no validator accepts. +_SCHEMA_LITERAL_KEYWORDS: Final = frozenset(("enum", "const")) + # Keywords whose value maps names to subschemas. Their keys are property names, not # keywords, so a property called `type` or `enum` is walked like any other subschema. _SCHEMA_MAP_KEYWORDS: Final = frozenset( @@ -333,14 +336,14 @@ def _collect_text_parts(container: MutableRequest, key: str, slots: _SlotSink) - _collect(part, "text", slots) -def _collect_tool_definitions(data: MutableRequest, privileged: _SlotSink) -> None: +def _collect_tool_definitions(data: MutableRequest, slots: _SlotSink, privileged: _SlotSink) -> None: """Tool definitions are application-authored free text bound for the provider. A tool's description and the free text in its parameter schema are where callers put examples and customer context, so they carry PII as often as a prompt does. They are collected into the privileged sink, like a system prompt: redacted outbound, and never - restorable from the reply. Names, types, `enum` and `const` values are left as sent, - because the model has to reproduce them exactly for a call to route. + restorable from the reply. `enum` and `const` values are the exception, and go to the + caller's vault -- see `_SCHEMA_LITERAL_KEYWORDS`. Names and types are left as sent. Covers Chat `tools[].function`, the legacy `functions[]`, and the flat tool shape the Responses API and Anthropic share, whose schema is `parameters` or `input_schema`. @@ -353,17 +356,19 @@ def _collect_tool_definitions(data: MutableRequest, privileged: _SlotSink) -> No function = tool.get("function") for holder in (tool, function) if isinstance(function, dict) else (tool,): _collect(holder, "description", privileged) - _collect_schema_text(holder.get("parameters"), privileged) - _collect_schema_text(holder.get("input_schema"), privileged) + _collect_schema_text(holder.get("parameters"), slots, privileged) + _collect_schema_text(holder.get("input_schema"), slots, privileged) -def _collect_schema_text(schema: object, privileged: _SlotSink) -> None: - """Collects the free text in a JSON Schema, at any depth. +def _collect_schema_text(schema: object, slots: _SlotSink, privileged: _SlotSink) -> None: + """Collects the text in a JSON Schema, at any depth. Scan by default: every string is collected except under the keywords in `_SCHEMA_STRUCTURAL_KEYWORDS`, whose values must go out verbatim. A list of keywords *to* collect would leak every one it forgot -- draft-07 `dependencies`, a `$comment`, - a vendor `x-` extension -- which is how this walk started out. + a vendor `x-` extension -- which is how this walk started out. Free text goes to the + privileged sink; `enum` / `const` literals go to the caller's, so the model's use of + them is restored. Structure matters in two places. Under `properties` and the other name -> subschema maps, keys are property names rather than keywords, so a property called `type` is a @@ -389,7 +394,10 @@ def _collect_schema_text(schema: object, privileged: _SlotSink) -> None: for keyword, value in tuple(node.items()): if keyword in _SCHEMA_STRUCTURAL_KEYWORDS: continue - if keyword in _SCHEMA_VALUE_KEYWORDS: + if keyword in _SCHEMA_LITERAL_KEYWORDS: + _collect(node, keyword, slots) + _collect_json_leaves(value, slots, strict=True) + elif keyword in _SCHEMA_VALUE_KEYWORDS: _collect(node, keyword, privileged) _collect_json_leaves(value, privileged, strict=True) elif keyword in _SCHEMA_MAP_KEYWORDS and isinstance(value, dict): @@ -421,7 +429,7 @@ def _collect_output_contracts(data: MutableRequest, slots: _SlotSink, privileged ): if isinstance(wrapper, dict): _collect(wrapper, "description", privileged) - _collect_schema_text(wrapper.get("schema"), privileged) + _collect_schema_text(wrapper.get("schema"), slots, privileged) def _collect_end_user_ids(data: MutableRequest, privileged: _SlotSink) -> None: @@ -1013,7 +1021,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): _collect_responses_fields(data, slots, privileged) _collect_prompt(data, slots) _collect_system(data, privileged) - _collect_tool_definitions(data, privileged) + _collect_tool_definitions(data, slots, privileged) _collect_output_contracts(data, slots, privileged) _collect_end_user_ids(data, privileged) return tuple(slots), tuple(privileged) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py index 475ab06994a..d909181e541 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -635,7 +635,7 @@ class TestRequestCoverage: def test_tool_schemas_give_up_their_free_text_and_nothing_else(self): """Every string is collected except what must reach the model verbatim: names, - types, formats, patterns, required lists, enum and const values.""" + types, formats, patterns and required lists.""" data = { "tools": [ { @@ -673,8 +673,9 @@ class TestRequestCoverage: } ] } - _, privileged = LLMShieldProxyGuardrail._locate_request_texts(data) + caller, privileged = LLMShieldProxyGuardrail._locate_request_texts(data) + assert sorted(text for text, _ in caller) == ["a", "a", "b"], "enum and const go to the caller vault" assert sorted(text for text, _ in privileged) == [ "comment", "default", @@ -691,6 +692,38 @@ class TestRequestCoverage: "vendor", ] + @pytest.mark.asyncio + async def test_enum_values_are_redacted_and_restored_in_the_tool_call(self): + """An enum value holding PII is redacted, and the model's use of the stand-in is + restored in its tool arguments, so the call still carries a value the schema allows.""" + guardrail = _guardrail(event_hook=["pre_call", "post_call"]) + shield = _FakeShield({"[EMAIL_1]": "ops@example.com"}) + redact_mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + data = { + "messages": [], + "tools": [ + { + "type": "function", + "function": { + "name": "notify", + "parameters": {"properties": {"to": {"type": "string", "enum": ["ops@example.com"]}}}, + }, + } + ], + } + + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + assert redact_mock.call_args_list[0].kwargs["json"]["texts"] == ["ops@example.com"] + assert data["tools"][0]["function"]["parameters"]["properties"]["to"]["enum"] == ["[EMAIL_1]"] + + guardrail.async_handler.post = shield.post # type: ignore[method-assign] + call = SimpleNamespace(function=SimpleNamespace(name="notify", arguments='{"to": "[EMAIL_1]"}')) + reply = ModelResponse(choices=[Choices(message=Message(content=None, tool_calls=None))]) + reply.choices[0].message.tool_calls = [call] + await guardrail.async_post_call_success_hook(data=data, user_api_key_dict=None, response=reply) + + assert json.loads(call.function.arguments) == {"to": "ops@example.com"} + def test_schema_nesting_past_the_bound_is_refused(self): schema: dict = {"type": "object", "description": "past-the-bound@example.com"} for _ in range(100): From f1578d622ea2069a2eee8d157dd1a8d843ce67ce Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Sat, 26 Sep 2026 11:06:18 -0500 Subject: [PATCH 32/34] fix(guardrails): redact llm_shield_proxy web search user locations Web search forwards the user's approximate location, and its free-text `city` and `region` fields can hold an address. Collect them into the non-restorable vault, from Chat `web_search_options.user_location` and from the `user_location` of Responses and Anthropic web-search tools. --- .../llm_shield_proxy/llm_shield_proxy.py | 21 +++++++++++++++++++ .../guardrail_hooks/test_llm_shield_proxy.py | 14 +++++++++++++ 2 files changed, 35 insertions(+) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index bd2796086db..8c02816b844 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -432,6 +432,26 @@ def _collect_output_contracts(data: MutableRequest, slots: _SlotSink, privileged _collect_schema_text(wrapper.get("schema"), slots, privileged) +def _collect_user_locations(data: MutableRequest, privileged: _SlotSink) -> None: + """Web search forwards the user's approximate location, whose `city` and `region` + are free text and can hold a street address. + + Chat carries it in `web_search_options.user_location.approximate`; the Responses + and Anthropic web-search tools carry it flat on the tool's `user_location`. Nothing + restores it from a reply, hence the privileged sink. + """ + options: Final = data.get("web_search_options") + tools: Final = data.get("tools") + for holder in (options, *(tools if isinstance(tools, list) else ())): + location = holder.get("user_location") if isinstance(holder, dict) else None + if not isinstance(location, dict): + continue + approximate = location.get("approximate") + for container in (location, approximate) if isinstance(approximate, dict) else (location,): + _collect(container, "city", privileged) + _collect(container, "region", privileged) + + def _collect_end_user_ids(data: MutableRequest, privileged: _SlotSink) -> None: """`user` and `safety_identifier` are forwarded to the provider and often hold an email. @@ -1023,6 +1043,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): _collect_system(data, privileged) _collect_tool_definitions(data, slots, privileged) _collect_output_contracts(data, slots, privileged) + _collect_user_locations(data, privileged) _collect_end_user_ids(data, privileged) return tuple(slots), tuple(privileged) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py index d909181e541..b9314590f38 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -957,6 +957,20 @@ class TestVaultIsolation: {"prediction": {"type": "content", "content": [{"type": "text", "text": "U"}]}, "instructions": "S"}, id="prediction-parts", ), + pytest.param( + { + "messages": [{"role": "user", "content": "U"}], + "web_search_options": {"user_location": {"type": "approximate", "approximate": {"city": "S"}}}, + }, + id="chat-web-search-location", + ), + pytest.param( + { + "input": "U", + "tools": [{"type": "web_search", "user_location": {"type": "approximate", "region": "S"}}], + }, + id="responses-web-search-location", + ), pytest.param({"messages": [{"role": "user", "content": "U"}], "user": "S"}, id="end-user-id"), pytest.param({"input": "U", "safety_identifier": "S"}, id="safety-identifier"), ], From 6db026064c2251dcd4e9242384f592066e269568 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Sat, 3 Oct 2026 17:47:06 -0500 Subject: [PATCH 33/34] fix(guardrails): drop unused llm_shield_proxy suppressions Upstream added LIT013 (a *-ok marker that suppresses nothing) and LIT014 (at most one for and one if per comprehension). Remove the 34 markers that no longer suppress anything and flatten the finished streams with itertools.chain.from_iterable. --- .../llm_shield_proxy/__init__.py | 4 +- .../llm_shield_proxy/llm_shield_proxy.py | 66 +++++++++---------- 2 files changed, 33 insertions(+), 37 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/__init__.py index c8ca68a8967..44b19f82218 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/__init__.py @@ -23,11 +23,11 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" return _llm_shield_guardrail_callback -guardrail_initializer_registry: Final = { # mutable-ok: module-level registry, built once and never mutated +guardrail_initializer_registry: Final = { SupportedGuardrailIntegrations.LLM_SHIELD_PROXY.value: initialize_guardrail, } -guardrail_class_registry: Final = { # mutable-ok: module-level registry, built once and never mutated +guardrail_class_registry: Final = { SupportedGuardrailIntegrations.LLM_SHIELD_PROXY.value: LLMShieldProxyGuardrail, } diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index 8c02816b844..d02be56fc52 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -7,6 +7,7 @@ import copy import functools +import itertools import json import os import re @@ -72,11 +73,9 @@ _DEFAULT_TIMEOUT_SECONDS: Final = 10.0 # The proxy's own request dict. Mutable by design: a pre-call guardrail rewrites # the caller's payload in place, which is the entire point of the hook. -# mutable-ok: the shape is fixed by CustomLogger's hook signatures. MutableRequest: TypeAlias = dict # A JSON body on its way to httpx, which requires a real dict rather than a view. -# mutable-ok: handed straight to the HTTP client. JsonBody: TypeAlias = dict # One redactable span: the text as it stands, and the write that puts the @@ -92,14 +91,14 @@ _MAX_CONTENT_DEPTH: Final = 8 # generous; past it the request is refused, for the same reason as above. _MAX_JSON_DEPTH: Final = 64 -_Slot: TypeAlias = tuple[str, Callable[[str], None]] # mutable-ok: Callable's param list. +_Slot: TypeAlias = tuple[str, Callable[[str], None]] # One incremental rehydration step for a stream the caller has already bound to its # vault: (new text, carried window, final) -> (text safe to emit, window still held). -_StreamStep: TypeAlias = Callable[[str, str, bool], Awaitable[tuple[str, str]]] # mutable-ok: Callable's param list. +_StreamStep: TypeAlias = Callable[[str, str, bool], Awaitable[tuple[str, str]]] # A batch rehydration already bound to the request's vault. -_Rehydrate: TypeAlias = Callable[[Sequence[str]], Awaitable[Sequence[str]]] # mutable-ok: Callable's param list. +_Rehydrate: TypeAlias = Callable[[Sequence[str]], Awaitable[Sequence[str]]] # Anthropic /v1/messages delta types that carry restorable text, and the field holding # it. `thinking_delta` is left out on purpose: a thinking block is signed, and one @@ -172,17 +171,17 @@ _SCHEMA_MAP_KEYWORDS: Final = frozenset( # The accumulator the collectors below append into. It never escapes # _locate_request_texts, which freezes it into a tuple before returning. -_SlotSink: TypeAlias = list[_Slot] # mutable-ok: accumulator passed between collectors. +_SlotSink: TypeAlias = list[_Slot] # Sliding windows keyed by (choice index, tool-call index | None), threaded through one # stream. `None` is the content channel; an int is one tool call's accumulating # `arguments`. Content and each tool call are separate token streams, so each needs its # own window -- one shared window would splice one stream's held-back tail onto another. -_CarryWindows: TypeAlias = dict # mutable-ok: per-stream windows advanced in place. +_CarryWindows: TypeAlias = dict # A caller-owned list whose entries are rewritten in place, such as a Completions # `prompt` sent as an array of strings. -MutableSeq: TypeAlias = list # mutable-ok: the request payload's own list. +MutableSeq: TypeAlias = list def _collect(container: MutableRequest, key: str, slots: _SlotSink) -> None: @@ -280,7 +279,7 @@ def _collect_participant_name(message: MutableRequest, slots: _SlotSink) -> None def _collect_tool_arguments(message: MutableRequest, slots: _SlotSink) -> None: """Tool arguments carry the values a user asked the model to act on.""" for tool_call in message.get("tool_calls") or (): - function = tool_call.get("function") if isinstance(tool_call, dict) else None # rebind-ok: loop variable. + function = tool_call.get("function") if isinstance(tool_call, dict) else None if isinstance(function, dict): _collect(function, "arguments", slots) legacy: Final = message.get("function_call") @@ -552,7 +551,7 @@ def _continuation_delta(tool_index: int, text: str) -> list[dict[str, object]]: Clients concatenate tool-call fragments by index, so no id or name is needed. """ - return [{"index": tool_index, "function": {"arguments": text}}] # mutable-ok: delta.tool_calls is a list. + return [{"index": tool_index, "function": {"arguments": text}}] def _collect_response_item(item: object, slots: _SlotSink) -> None: @@ -638,9 +637,9 @@ class _AnthropicSSERestorer: self._step: Final = step self._carries: Final[dict[int, str]] = {} # mutable-ok: per-block windows advanced in place. self._delta_types: Final[dict[int, str]] = {} # mutable-ok: each block's delta type, for its flush. - self._pending = b"" # rebind-ok: the unfinished tail of the stream. - self._as_text = False # rebind-ok: set once if the stream arrives as str rather than bytes. - self._is_sse: bool | None = None # rebind-ok: undecided until the opening bytes settle it. + self._pending = b"" + self._as_text = False + self._is_sse: bool | None = None async def feed(self, chunk: bytes | str) -> tuple[bytes | str, ...]: """Restores every event this chunk completes; holds back an unfinished tail.""" @@ -669,9 +668,7 @@ class _AnthropicSSERestorer: # empty remainder after the last separator. parts: Final = _SSE_EVENT_BOUNDARY.split(buffered[:cut]) restored: Final = tuple( - [ # mutable-ok: an await needs a list comprehension; frozen at once. - await self._restore_event(parts[index]) + parts[index + 1] for index in range(0, len(parts) - 1, 2) - ] + [await self._restore_event(parts[index]) + parts[index + 1] for index in range(0, len(parts) - 1, 2)] ) return self._emit(b"".join(restored)) @@ -754,7 +751,7 @@ class _AnthropicSSERestorer: text, _ = await self._step("", carry, True) if not text: return b"" - event: Final[JsonBody] = { # mutable-ok: serialised on the next line. + event: Final[JsonBody] = { "type": "content_block_delta", "index": index, "delta": {"type": delta_type, field: text}, @@ -762,7 +759,7 @@ class _AnthropicSSERestorer: return f"event: content_block_delta\ndata: {json.dumps(event, ensure_ascii=False)}\n\n".encode() async def _flush_all(self) -> bytes: - flushed: Final = tuple([await self._flush(index) for index in tuple(self._carries)]) # mutable-ok: frozen. + flushed: Final = tuple([await self._flush(index) for index in tuple(self._carries)]) return b"".join(flushed) @@ -798,13 +795,13 @@ class _ResponsesStreamRestorer: if kind.endswith(".delta") and kind not in _RESPONSES_BINARY_DELTAS: await self._restore_delta(event, kind) return (event,) - slots: Final[_SlotSink] = [] # mutable-ok: accumulator, restored in one batch. + slots: Final[_SlotSink] = [] flushed: Final = await self._flush(_responses_stream_key(event, kind)) if kind.endswith(".done") else () if kind.endswith(".done"): _collect_event_text(event, slots) part: Final = _read_field(event, "part") if part is not None: - _collect_response_item({"content": [part]}, slots) # mutable-ok: a one-part view. + _collect_response_item({"content": [part]}, slots) _collect_response_item(_read_field(event, "item"), slots) elif kind in _RESPONSES_TERMINAL_EVENTS: for item in _read_list(_read_field(event, "response"), "output"): @@ -814,8 +811,8 @@ class _ResponsesStreamRestorer: async def finish(self) -> tuple[object, ...]: """Flushes every stream the provider never closed, e.g. a truncated reply.""" - flushed: Final = tuple([await self._flush(key) for key in tuple(self._carries)]) # mutable-ok: frozen. - return tuple(event for events in flushed for event in events) + flushed: Final = tuple([await self._flush(key) for key in tuple(self._carries)]) + return tuple(itertools.chain.from_iterable(flushed)) async def _restore_delta(self, event: object, kind: str) -> None: text: Final = _read_field(event, "delta") @@ -908,12 +905,12 @@ class LLMShieldProxyGuardrail(CustomGuardrail): @classmethod def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: # mutable-ok: parent's signature. - return [GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call] # mutable-ok: parent's signature. + return [GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call] # --- transport --------------------------------------------------------------- def _headers(self, session_id: str) -> JsonBody: - headers: Final[JsonBody] = { # mutable-ok: httpx requires a real dict. + headers: Final[JsonBody] = { "Content-Type": "application/json", "X-Session-ID": session_id, } @@ -951,12 +948,12 @@ class LLMShieldProxyGuardrail(CustomGuardrail): ) from exc async def _redact(self, texts: Sequence[str], session_id: str) -> Sequence[str]: - payload: Final[JsonBody] = {"texts": list(texts)} # mutable-ok: JSON body for httpx. + payload: Final[JsonBody] = {"texts": list(texts)} body: Final = await self._call_shield(_REDACT_PATH, session_id, payload) return self._same_length_or_raise(body.get("texts"), texts, "redact") async def _rehydrate(self, texts: Sequence[str], session_id: str) -> Sequence[str]: - payload: Final[JsonBody] = {"texts": list(texts)} # mutable-ok: JSON body for httpx. + payload: Final[JsonBody] = {"texts": list(texts)} body: Final = await self._call_shield(_REHYDRATE_PATH, session_id, payload) return self._same_length_or_raise(body.get("texts"), texts, "rehydrate") @@ -984,7 +981,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): # /v1/responses. The session id is a capability against the vault's rehydrate # endpoint, so handing it to the provider alongside the placeholders would let the # provider read back exactly what this guardrail exists to withhold. - metadata: Final = data.setdefault("litellm_metadata", {}) # mutable-ok: per-request store. + metadata: Final = data.setdefault("litellm_metadata", {}) if isinstance(metadata, dict): metadata[_SESSION_METADATA_KEY] = session_id return session_id @@ -1030,8 +1027,8 @@ class LLMShieldProxyGuardrail(CustomGuardrail): guardrail, and an agent that reads a file and quotes an address from it needs that address back. """ - slots: Final[_SlotSink] = [] # mutable-ok: accumulator, frozen on return. - privileged: Final[_SlotSink] = [] # mutable-ok: accumulator, frozen on return. + slots: Final[_SlotSink] = [] + privileged: Final[_SlotSink] = [] for message in data.get("messages") or (): if isinstance(message, dict): sink = privileged if message.get("role") in _PRIVILEGED_ROLES else slots @@ -1202,7 +1199,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): fields on the request side -- a function_call item holds `arguments`, a function_call_output holds `output` -- so the two directions stay symmetric. """ - slots: Final[_SlotSink] = [] # mutable-ok: accumulator, frozen on return. + slots: Final[_SlotSink] = [] for item in getattr(response, "output", None) or (): _collect_response_item(item, slots) return tuple(slots) @@ -1376,7 +1373,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): continuations.extend(_continuation_delta(tool_index, text)) if continuations: existing: Final = tuple(getattr(delta, "tool_calls", None) or ()) - delta.tool_calls = [*existing, *continuations] # mutable-ok: delta.tool_calls is a list. + delta.tool_calls = [*existing, *continuations] async def _flush_trailing( self, last_chunk: Any, carries: _CarryWindows, session_id: str @@ -1433,7 +1430,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): kept.index = index # The terminal signal, if there was one, already went out with the real chunk. kept.finish_reason = None - chunk.choices = [kept] # mutable-ok: the chunk model requires a list. + chunk.choices = [kept] return chunk async def _stream_step(self, text: str, carry: str, final: bool, session_id: str) -> tuple[str, str]: @@ -1441,8 +1438,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): body: Final = await self._call_shield( _REHYDRATE_STREAM_PATH, session_id, - # mutable-ok: JSON request body for httpx. - {"text": text, "carry": carry, "final": final}, # mutable-ok: JSON request body for httpx. + {"text": text, "carry": carry, "final": final}, ) emitted: Final = body.get("text") remaining: Final = body.get("carry") @@ -1500,7 +1496,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): write(replacement) # Return a new mapping rather than rewriting the caller's, so this stays a # pure transform of the inputs it was handed. - merged: Final[JsonBody] = {**inputs} # mutable-ok: TypedDict. + merged: Final[JsonBody] = {**inputs} if text_list: merged["texts"] = restored_values[: len(text_list)] if restored_calls: From 1b45935b80a43a6b9f6e4b0364ab96c31298386f Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Sat, 3 Oct 2026 18:09:45 -0500 Subject: [PATCH 34/34] fix(guardrails): type the llm_shield_proxy request and reply walks Narrowing with isinstance(x, dict) leaves keys and values unknown, so every call that passed a narrowed value counted against the reportUnknownArgumentType budget. Parse into dict[str, object] and list[object] once, in _as_object and _as_array, type the carry keys and accumulators, and bind writers with functools.partial instead of lambdas. The shield's batch reply is now also checked to hold only strings. --- .../llm_shield_proxy/llm_shield_proxy.py | 267 ++++++++++-------- 1 file changed, 156 insertions(+), 111 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index d02be56fc52..49bcb768600 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -73,10 +73,10 @@ _DEFAULT_TIMEOUT_SECONDS: Final = 10.0 # The proxy's own request dict. Mutable by design: a pre-call guardrail rewrites # the caller's payload in place, which is the entire point of the hook. -MutableRequest: TypeAlias = dict +MutableRequest: TypeAlias = dict[str, object] # A JSON body on its way to httpx, which requires a real dict rather than a view. -JsonBody: TypeAlias = dict +JsonBody: TypeAlias = dict[str, object] # One redactable span: the text as it stands, and the write that puts the # replacement back where it came from. @@ -177,25 +177,48 @@ _SlotSink: TypeAlias = list[_Slot] # stream. `None` is the content channel; an int is one tool call's accumulating # `arguments`. Content and each tool call are separate token streams, so each needs its # own window -- one shared window would splice one stream's held-back tail onto another. -_CarryWindows: TypeAlias = dict +_CarryKey: TypeAlias = tuple[int, int | None] +_CarryWindows: TypeAlias = dict[_CarryKey, str] # A caller-owned list whose entries are rewritten in place, such as a Completions # `prompt` sent as an array of strings. -MutableSeq: TypeAlias = list +MutableSeq: TypeAlias = list[object] + +# A Responses API delta stream: (event family, item id, output index, part index). +_ResponsesStreamKey: TypeAlias = tuple[str, object, object, object] + + +def _as_object(value: object) -> MutableRequest | None: + """`value` as a JSON object, or None. + + `isinstance(value, dict)` alone leaves the keys and values unknown to the type + checker. A JSON object's keys are strings, so the type is stated once, here. + """ + return value if isinstance(value, dict) else None + + +def _as_array(value: object) -> MutableSeq | None: + """`value` as a JSON array, or None. See `_as_object`.""" + return value if isinstance(value, list) else None + + +def _is_container(value: object) -> bool: + """Whether `value` is a JSON object or array, without narrowing it to unknown types.""" + return isinstance(value, (dict, list)) def _collect(container: MutableRequest, key: str, slots: _SlotSink) -> None: """Records the string at `key`, along with the write that replaces it.""" value: Final = container.get(key) if isinstance(value, str) and value: - slots.append((value, lambda new, c=container, k=key: c.__setitem__(k, new))) + slots.append((value, functools.partial(container.__setitem__, key))) def _collect_entry(entries: MutableSeq, index: int, slots: _SlotSink) -> None: """Records a string held directly in a list, rather than under a key.""" value: Final = entries[index] if isinstance(value, str) and value: - slots.append((value, lambda new, e=entries, i=index: e.__setitem__(i, new))) + slots.append((value, functools.partial(entries.__setitem__, index))) def _collect_prompt(data: MutableRequest, slots: _SlotSink) -> None: @@ -205,19 +228,21 @@ def _collect_prompt(data: MutableRequest, slots: _SlotSink) -> None: if isinstance(prompt, str): _collect(data, "prompt", slots) return - if isinstance(prompt, dict): + prompt_object: Final = _as_object(prompt) + if prompt_object is not None: # A Responses API PromptObject. `variables` are substituted into the stored # prompt on the provider side, so they are caller text. `id` and `version` # identify which prompt to use and must arrive unchanged. - variables: Final = prompt.get("variables") - if isinstance(variables, dict): + variables: Final = _as_object(prompt_object.get("variables")) + if variables is not None: for name in tuple(variables): _collect(variables, name, slots) return - if not isinstance(prompt, list): + entries: Final = _as_array(prompt) + if entries is None: return - for index in range(len(prompt)): - _collect_entry(prompt, index, slots) + for index in range(len(entries)): + _collect_entry(entries, index, slots) class _RequestTooDeep(Exception): @@ -238,7 +263,7 @@ def _collect_content(container: MutableRequest, slots: _SlotSink) -> None: """ # Walked in document order: the shield maps its replies back by position, so the # order spans are collected in is part of the contract. - pending: Final[list] = [(container, 0)] # mutable-ok: local queue, never escapes. + pending: Final[list[tuple[MutableRequest, int]]] = [(container, 0)] # mutable-ok: local queue, never escapes. cursor = 0 # rebind-ok: advances through the queue. while cursor < len(pending): node, depth = pending[cursor] @@ -249,8 +274,9 @@ def _collect_content(container: MutableRequest, slots: _SlotSink) -> None: continue if depth >= _MAX_CONTENT_DEPTH and content: raise _RequestTooDeep("content") - for part in content if isinstance(content, list) else (): - if not isinstance(part, dict): + for item in _as_array(content) or (): + part = _as_object(item) + if part is None: continue # Image and audio parts have no text and fall through untouched. _collect(part, "text", slots) @@ -278,12 +304,13 @@ def _collect_participant_name(message: MutableRequest, slots: _SlotSink) -> None def _collect_tool_arguments(message: MutableRequest, slots: _SlotSink) -> None: """Tool arguments carry the values a user asked the model to act on.""" - for tool_call in message.get("tool_calls") or (): - function = tool_call.get("function") if isinstance(tool_call, dict) else None - if isinstance(function, dict): + for tool_call in _read_list(message, "tool_calls"): + tool_call_object = _as_object(tool_call) + function = _as_object(tool_call_object.get("function")) if tool_call_object is not None else None + if function is not None: _collect(function, "arguments", slots) - legacy: Final = message.get("function_call") - if isinstance(legacy, dict): + legacy: Final = _as_object(message.get("function_call")) + if legacy is not None: _collect(legacy, "arguments", slots) @@ -293,9 +320,7 @@ def _collect_system(data: MutableRequest, slots: _SlotSink) -> None: if isinstance(system, str): _collect(data, "system", slots) return - for part in system if isinstance(system, list) else (): - if isinstance(part, dict): - _collect(part, "text", slots) + _collect_text_parts(data, "system", slots) def _collect_responses_fields(data: MutableRequest, slots: _SlotSink, privileged: _SlotSink) -> None: @@ -309,14 +334,16 @@ def _collect_responses_fields(data: MutableRequest, slots: _SlotSink, privileged if isinstance(request_input, str): _collect(data, "input", slots) return - if not isinstance(request_input, list): + entries: Final = _as_array(request_input) + if entries is None: return - for index, item in enumerate(request_input): - if isinstance(item, str): + for index, entry in enumerate(entries): + if isinstance(entry, str): # The embeddings and moderations shape: `input` as an array of strings. - _collect_entry(request_input, index, slots) + _collect_entry(entries, index, slots) continue - if not isinstance(item, dict): + item = _as_object(entry) + if item is None: continue _collect_content(item, slots) # A function_call item holds `arguments`; a function_call_output holds `output`. @@ -329,9 +356,9 @@ def _collect_responses_fields(data: MutableRequest, slots: _SlotSink, privileged def _collect_text_parts(container: MutableRequest, key: str, slots: _SlotSink) -> None: """Collects the `text` of every part in the list held at `key`.""" - parts: Final = container.get(key) - for part in parts if isinstance(parts, list) else (): - if isinstance(part, dict): + for entry in _as_array(container.get(key)) or (): + part = _as_object(entry) + if part is not None: _collect(part, "text", slots) @@ -348,12 +375,12 @@ def _collect_tool_definitions(data: MutableRequest, slots: _SlotSink, privileged Responses API and Anthropic share, whose schema is `parameters` or `input_schema`. """ for key in ("tools", "functions"): - declared = data.get(key) - for tool in declared if isinstance(declared, list) else (): - if not isinstance(tool, dict): + for entry in _as_array(data.get(key)) or (): + tool = _as_object(entry) + if tool is None: continue - function = tool.get("function") - for holder in (tool, function) if isinstance(function, dict) else (tool,): + function = _as_object(tool.get("function")) + for holder in (tool, function) if function is not None else (tool,): _collect(holder, "description", privileged) _collect_schema_text(holder.get("parameters"), slots, privileged) _collect_schema_text(holder.get("input_schema"), slots, privileged) @@ -375,35 +402,38 @@ def _collect_schema_text(schema: object, slots: _SlotSink, privileged: _SlotSink so all their strings are collected whatever the keys around them are called. Nested past `_MAX_JSON_DEPTH`, the request is refused. """ - pending: Final[list] = [(schema, 0)] # mutable-ok: local walk stack. + pending: Final[list[tuple[object, int]]] = [(schema, 0)] # mutable-ok: local walk stack. while pending: node, depth = pending.pop() if depth > _MAX_JSON_DEPTH: - if isinstance(node, (dict, list)) and node: + if _is_container(node) and node: raise _RequestTooDeep("schema") continue - if isinstance(node, list): - for index, item in enumerate(node): - _collect_entry(node, index, privileged) - if isinstance(item, (dict, list)): + entries = _as_array(node) + if entries is not None: + for index, item in enumerate(entries): + _collect_entry(entries, index, privileged) + if _is_container(item): pending.append((item, depth + 1)) continue - if not isinstance(node, dict): + schema_object = _as_object(node) + if schema_object is None: continue - for keyword, value in tuple(node.items()): + for keyword, value in tuple(schema_object.items()): if keyword in _SCHEMA_STRUCTURAL_KEYWORDS: continue + subschemas = _as_object(value) if keyword in _SCHEMA_MAP_KEYWORDS else None if keyword in _SCHEMA_LITERAL_KEYWORDS: - _collect(node, keyword, slots) + _collect(schema_object, keyword, slots) _collect_json_leaves(value, slots, strict=True) elif keyword in _SCHEMA_VALUE_KEYWORDS: - _collect(node, keyword, privileged) + _collect(schema_object, keyword, privileged) _collect_json_leaves(value, privileged, strict=True) - elif keyword in _SCHEMA_MAP_KEYWORDS and isinstance(value, dict): - pending.extend((child, depth + 1) for child in value.values()) + elif subschemas is not None: + pending.extend((child, depth + 1) for child in subschemas.values()) elif isinstance(value, str): - _collect(node, keyword, privileged) - elif isinstance(value, (dict, list)): + _collect(schema_object, keyword, privileged) + elif _is_container(value): pending.append((value, depth + 1)) @@ -416,17 +446,18 @@ def _collect_output_contracts(data: MutableRequest, slots: _SlotSink, privileged `text.format` -- is application-authored like a tool schema, so its free text goes to the privileged sink, and its names and types stay as sent. """ - prediction: Final = data.get("prediction") - if isinstance(prediction, dict): + prediction: Final = _as_object(data.get("prediction")) + if prediction is not None: _collect(prediction, "content", slots) _collect_text_parts(prediction, "content", slots) - response_format: Final = data.get("response_format") - text_options: Final = data.get("text") - for wrapper in ( - response_format.get("json_schema") if isinstance(response_format, dict) else None, - text_options.get("format") if isinstance(text_options, dict) else None, + response_format: Final = _as_object(data.get("response_format")) + text_options: Final = _as_object(data.get("text")) + for declared in ( + response_format.get("json_schema") if response_format is not None else None, + text_options.get("format") if text_options is not None else None, ): - if isinstance(wrapper, dict): + wrapper = _as_object(declared) + if wrapper is not None: _collect(wrapper, "description", privileged) _collect_schema_text(wrapper.get("schema"), slots, privileged) @@ -440,13 +471,14 @@ def _collect_user_locations(data: MutableRequest, privileged: _SlotSink) -> None restores it from a reply, hence the privileged sink. """ options: Final = data.get("web_search_options") - tools: Final = data.get("tools") - for holder in (options, *(tools if isinstance(tools, list) else ())): - location = holder.get("user_location") if isinstance(holder, dict) else None - if not isinstance(location, dict): + tools: Final = _as_array(data.get("tools")) or () + for declared in (options, *tools): + holder = _as_object(declared) + location = _as_object(holder.get("user_location")) if holder is not None else None + if location is None: continue - approximate = location.get("approximate") - for container in (location, approximate) if isinstance(approximate, dict) else (location,): + approximate = _as_object(location.get("approximate")) + for container in (location, approximate) if approximate is not None else (location,): _collect(container, "city", privileged) _collect(container, "region", privileged) @@ -487,7 +519,9 @@ def _read_list(holder: object, name: str) -> Sequence[object]: The entries are the reply's own objects, so writing through them edits the reply. """ value: Final = _read_field(holder, name) - return tuple(value) if isinstance(value, (list, tuple)) else () + if isinstance(value, tuple): + return value + return tuple(_as_array(value) or ()) def _write_field(holder: object, name: str, value: str) -> None: @@ -512,30 +546,32 @@ def _collect_json_leaves(node: object, slots: _SlotSink, *, strict: bool = False unredacted: past the bound it raises `_RequestTooDeep`. On the reply side a leaf past the bound just keeps its placeholder, which leaks nothing, so it is skipped. """ - pending: Final[list] = [(node, 0)] # mutable-ok: local walk stack. + pending: Final[list[tuple[object, int]]] = [(node, 0)] # mutable-ok: local walk stack. while pending: current, current_depth = pending.pop() if current_depth > _MAX_JSON_DEPTH: - if strict and isinstance(current, (dict, list)) and current: + if strict and _is_container(current) and current: raise _RequestTooDeep("json") continue - if isinstance(current, dict): - for key in tuple(current): - value = current[key] + current_object = _as_object(current) + if current_object is not None: + for key in tuple(current_object): + value = current_object[key] if isinstance(value, str) and value: - slots.append((value, lambda new, d=current, k=key: d.__setitem__(k, new))) + slots.append((value, functools.partial(current_object.__setitem__, key))) else: pending.append((value, current_depth + 1)) continue - if isinstance(current, list): - for index, value in enumerate(current): + entries = _as_array(current) + if entries is not None: + for index, value in enumerate(entries): if isinstance(value, str) and value: - slots.append((value, lambda new, entries=current, i=index: entries.__setitem__(i, new))) + slots.append((value, functools.partial(entries.__setitem__, index))) else: pending.append((value, current_depth + 1)) -def _carry_sort_key(key: tuple) -> tuple: +def _carry_sort_key(key: _CarryKey) -> tuple[int, int]: """Orders streaming windows without ever comparing None to an int. `sorted()` over the raw keys raises as soon as one choice holds both a content window @@ -701,10 +737,11 @@ class _AnthropicSSERestorer: return block line: Final = lines[data_lines[0]] try: - event: Final = json.loads(line[len("data:") :]) + parsed: Final[object] = json.loads(line[len("data:") :]) except ValueError: return block - if not isinstance(event, dict): + event: Final = _as_object(parsed) + if event is None: return block kind: Final = event.get("type") index: Final = event.get("index") @@ -784,8 +821,8 @@ class _ResponsesStreamRestorer: def __init__(self, step: _StreamStep, rehydrate: _Rehydrate) -> None: self._step: Final = step self._rehydrate: Final = rehydrate - self._carries: Final[dict[tuple, str]] = {} # mutable-ok: per-stream windows advanced in place. - self._last_deltas: Final[dict[tuple, object]] = {} # mutable-ok: newest delta per stream. + self._carries: Final[dict[_ResponsesStreamKey, str]] = {} # mutable-ok: per-stream windows advanced in place. + self._last_deltas: Final[dict[_ResponsesStreamKey, object]] = {} # mutable-ok: newest delta per stream. async def restore(self, event: object) -> tuple[object, ...]: """The events to emit in place of `event`: any flush, then the event itself.""" @@ -824,7 +861,7 @@ class _ResponsesStreamRestorer: self._last_deltas[key] = event _write_field(event, "delta", emitted) - async def _flush(self, key: tuple) -> tuple[object, ...]: + async def _flush(self, key: _ResponsesStreamKey) -> tuple[object, ...]: carry: Final = self._carries.pop(key, "") template: Final = self._last_deltas.pop(key, None) if not carry or template is None: @@ -837,7 +874,7 @@ class _ResponsesStreamRestorer: return (flush,) -def _responses_stream_key(event: object, kind: str) -> tuple: +def _responses_stream_key(event: object, kind: str) -> _ResponsesStreamKey: """Identifies the delta stream an event belongs to, the same for its delta and done. The family is the event type without its `.delta` / `.done` suffix, so an output_text @@ -861,14 +898,16 @@ def _collect_event_text(event: object, slots: _SlotSink) -> None: `.done` event of each stream family names its text differently (`text`, `refusal`, `arguments`, ...), and a family added upstream would otherwise leak a placeholder. """ - fields: Final = event if isinstance(event, dict) else getattr(event, "__dict__", None) - if not isinstance(fields, dict): + # A model's fields live in its `__dict__`; an empty dict has none either way. + attributes: Final[object] = getattr(event, "__dict__", None) + fields: Final = _as_object(event) or _as_object(attributes) + if fields is None: return for name, value in tuple(fields.items()): - if not isinstance(name, str) or name in _RESPONSES_STRUCTURAL_FIELDS or name.endswith("_id"): + if name in _RESPONSES_STRUCTURAL_FIELDS or name.endswith("_id"): continue if isinstance(value, str) and value: - slots.append((value, lambda new, n=name: _write_field(event, n, new))) + slots.append((value, functools.partial(_write_field, event, name))) class LLMShieldProxyGuardrail(CustomGuardrail): @@ -959,12 +998,15 @@ class LLMShieldProxyGuardrail(CustomGuardrail): def _same_length_or_raise(self, returned: object, sent: Sequence[str], operation: str) -> Sequence[str]: """Guards the positional mapping the callers rely on to write results back.""" - if not isinstance(returned, list) or len(returned) != len(sent): + entries: Final = _as_array(returned) + texts: Final = tuple(entry for entry in entries or () if isinstance(entry, str)) + # A non-string entry would be written into the request or reply as is. + if entries is None or len(entries) != len(sent) or len(texts) != len(entries): raise GuardrailRaisedException( guardrail_name=self.guardrail_name, message=f"LLM Shield Proxy {operation} returned an unexpected payload; blocking the request.", ) - return tuple(returned) + return texts # --- session ------------------------------------------------------------------ @@ -1029,8 +1071,9 @@ class LLMShieldProxyGuardrail(CustomGuardrail): """ slots: Final[_SlotSink] = [] privileged: Final[_SlotSink] = [] - for message in data.get("messages") or (): - if isinstance(message, dict): + for entry in _read_list(data, "messages"): + message = _as_object(entry) + if message is not None: sink = privileged if message.get("role") in _PRIVILEGED_ROLES else slots _collect_content(message, sink) _collect_participant_name(message, sink) @@ -1121,14 +1164,14 @@ class LLMShieldProxyGuardrail(CustomGuardrail): # characters, and `_same_length_or_raise` is what guarantees the positional # mapping -- so a reply carrying more spans than that fails closed, which is this # guardrail's posture everywhere else. - pending: Final[list] = [] # mutable-ok: local accumulator, frozen before use. + pending: Final[_SlotSink] = [] for choice in choices: message = getattr(choice, "message", None) if message is None: continue content = getattr(message, "content", None) if isinstance(content, str) and content: - pending.append((content, lambda new, m=message: setattr(m, "content", new))) + pending.append((content, functools.partial(setattr, message, "content"))) # A tool call's `arguments` is model-generated text and the request path # redacts it, so leaving it unrestored hands the application a placeholder to # invoke a tool with. These are Pydantic objects on this path, not dicts. @@ -1136,11 +1179,11 @@ class LLMShieldProxyGuardrail(CustomGuardrail): function = getattr(tool_call, "function", None) arguments = getattr(function, "arguments", None) if function is not None else None if isinstance(arguments, str) and arguments: - pending.append((arguments, lambda new, f=function: setattr(f, "arguments", new))) + pending.append((arguments, functools.partial(setattr, function, "arguments"))) legacy = getattr(message, "function_call", None) legacy_arguments = getattr(legacy, "arguments", None) if legacy is not None else None if isinstance(legacy_arguments, str) and legacy_arguments: - pending.append((legacy_arguments, lambda new, fn=legacy: setattr(fn, "arguments", new))) + pending.append((legacy_arguments, functools.partial(setattr, legacy, "arguments"))) if not pending: return response @@ -1168,15 +1211,17 @@ class LLMShieldProxyGuardrail(CustomGuardrail): string, and the request path redacts its string leaves -- so the reply's leaves have to come back or the application invokes the tool with placeholders. """ - slots: Final[list] = [] # mutable-ok: accumulator, frozen before use. - for block in response["content"]: - if not isinstance(block, dict): + slots: Final[_SlotSink] = [] + for entry in _read_list(response, "content"): + block = _as_object(entry) + if block is None: continue kind = block.get("type") - if kind == "text" and isinstance(block.get("text"), str) and block["text"]: - slots.append((block["text"], lambda new, b=block: b.__setitem__("text", new))) - elif kind == "tool_use" and isinstance(block.get("input"), dict): - _collect_json_leaves(block["input"], slots) + text = block.get("text") + if kind == "text" and isinstance(text, str) and text: + slots.append((text, functools.partial(block.__setitem__, "text"))) + elif kind == "tool_use" and _as_object(block.get("input")) is not None: + _collect_json_leaves(block.get("input"), slots) if not slots: return response @@ -1237,7 +1282,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): rehydrate: Final = functools.partial(self._rehydrate, session_id=session_id) sse: Final = _AnthropicSSERestorer(step) events: Final = _ResponsesStreamRestorer(step, rehydrate) - carries: Final[dict] = {} # mutable-ok: per-stream windows, local to this generator. + carries: Final[_CarryWindows] = {} last_chunk = None # rebind-ok: tracks the most recent chunk for the final flush. async for chunk in response: @@ -1291,7 +1336,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): async def _restore_content_window( self, delta: Any, - key: tuple, + key: _CarryKey, carries: _CarryWindows, session_id: str, is_final: bool, @@ -1357,7 +1402,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): clients concatenate by index, so no id or name is needed. Appending is correct even when this chunk already carried a fragment for that tool call. """ - continuations: Final[list] = [] # mutable-ok: built into this chunk's delta. + continuations: Final[list[dict[str, object]]] = [] # mutable-ok: built into this chunk's delta. for key in sorted((held for held in carries if held[0] == choice_index), key=_carry_sort_key): carry = carries[key] if not carry: @@ -1422,9 +1467,9 @@ class LLMShieldProxyGuardrail(CustomGuardrail): raw_choices: Final = getattr(chunk, "choices", None) if not raw_choices: return None - choices: Final[tuple] = tuple(raw_choices) - matching: Final = tuple(choice for choice in choices if _choice_index(choice) == index) - kept: Final = matching[0] if matching else choices[0] + choices: Final[tuple[object, ...]] = tuple(raw_choices) + position: Final = next((at for at, choice in enumerate(choices) if _choice_index(choice) == index), 0) + kept: Final = raw_choices[position] if getattr(kept, "delta", None) is None: return None kept.index = index @@ -1475,22 +1520,22 @@ class LLMShieldProxyGuardrail(CustomGuardrail): # Copied rather than mutated: the caller's tool calls are theirs to own, and this # method's contract is to hand back a new mapping. - restored_calls: Final[list] = [copy.deepcopy(call) for call in tool_calls] # mutable-ok: a new list. - spans: Final[list] = list(text_list) # mutable-ok: ordered batch, frozen before the call. - writers: Final[list] = [] # mutable-ok: one per span appended below. + restored_calls: Final[list[object]] = [copy.deepcopy(call) for call in tool_calls] # mutable-ok: a new list. + spans: Final[list[str]] = list(text_list) # mutable-ok: ordered batch, frozen before the call. + writers: Final[list[Callable[[str], None]]] = [] # mutable-ok: one per span appended below. for call in restored_calls: function = _read_field(call, "function") arguments = _read_field(function, "arguments") if function is not None else None if isinstance(arguments, str) and arguments: spans.append(arguments) - writers.append(lambda new, f=function: _write_field(f, "arguments", new)) + writers.append(functools.partial(_write_field, function, "arguments")) replaced: Final = ( await self._redact(tuple(spans), self._mint_session_id(request_data)) if input_type == "request" else await self._rehydrate(tuple(spans), self._session_id(request_data)) ) - restored_values: Final[list] = list(replaced) # mutable-ok: sliced into the texts list. + restored_values: Final[list[str]] = list(replaced) # mutable-ok: sliced into the texts list. for write, replacement in zip(writers, restored_values[len(text_list) :]): write(replacement)