diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/__init__.py
new file mode 100644
index 00000000000..44b19f82218
--- /dev/null
+++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/__init__.py
@@ -0,0 +1,33 @@
+from typing import TYPE_CHECKING, Final
+
+from litellm.types.guardrails import SupportedGuardrailIntegrations
+
+from .llm_shield_proxy import LLMShieldProxyGuardrail
+
+if TYPE_CHECKING:
+ from litellm.types.guardrails import Guardrail, LitellmParams
+
+
+def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail") -> LLMShieldProxyGuardrail:
+ import litellm
+
+ _llm_shield_guardrail_callback: Final = LLMShieldProxyGuardrail(
+ 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 = {
+ SupportedGuardrailIntegrations.LLM_SHIELD_PROXY.value: initialize_guardrail,
+}
+
+
+guardrail_class_registry: Final = {
+ SupportedGuardrailIntegrations.LLM_SHIELD_PROXY.value: LLMShieldProxyGuardrail,
+}
diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/example_config.yaml b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/example_config.yaml
new file mode 100644
index 00000000000..4f732773a75
--- /dev/null
+++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/example_config.yaml
@@ -0,0 +1,14 @@
+model_list:
+ - model_name: gpt-4o
+ litellm_params:
+ model: openai/gpt-4o
+ api_key: os.environ/OPENAI_API_KEY
+
+guardrails:
+ - guardrail_name: "llm_shield_proxy"
+ litellm_params:
+ guardrail: llm_shield_proxy
+ mode: ["pre_call", "post_call"]
+ default_on: true
+ api_base: "http://localhost:8000"
+ api_key: os.environ/LLM_SHIELD_PROXY_API_KEY
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
new file mode 100644
index 00000000000..fa23a110a74
--- /dev/null
+++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py
@@ -0,0 +1,744 @@
+import copy
+import functools
+import os
+import uuid
+from collections.abc import AsyncGenerator, Callable, Mapping, Sequence
+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,
+ 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, TextChoices
+
+if TYPE_CHECKING:
+ from litellm.caching.caching import DualCache
+ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
+ from litellm.types.proxy.guardrails.guardrail_hooks.llm_shield_proxy import (
+ LLMShieldProxyGuardrailConfigModel,
+ )
+ from litellm.types.utils import CallTypes, LLMResponseTypes
+from .payload import (
+ JsonBody,
+ MutableRequest,
+ RequestTooDeep,
+ Slot,
+ SlotSink,
+ as_array,
+ as_object,
+ choice_index,
+ collect_json_leaves,
+ collect_response_item,
+ detached,
+ read_field,
+ read_list,
+ rehydrate_slots,
+ write_field,
+)
+from .request_walk import (
+ locate_request_texts,
+)
+from .stream_restorers import (
+ AnthropicSSERestorer,
+ CarryKey,
+ CarryWindows,
+ ResponsesStreamRestorer,
+ carry_sort_key,
+ continuation_delta,
+ responses_event_type,
+)
+
+GUARDRAIL_NAME: Final = "llm_shield_proxy"
+
+_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"
+
+_SESSION_METADATA_KEY: Final = "llm_shield_session_id"
+
+_DEPLOYMENT_RESTORE_KEY: Final = "llm_shield_restore_at_deployment"
+
+_VAULT_PREFIX: Final = f"litellm-{uuid.uuid4().hex}"
+
+_DEFAULT_TIMEOUT_SECONDS: Final = 10.0
+
+
+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
+ 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.
+ """
+
+ 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, # noqa: LIT008 # kwargs-ok: forwarded verbatim to CustomGuardrail.__init__
+ ) -> None:
+ self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
+ 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
+ def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: # mutable-ok: parent's signature.
+ return [GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call]
+
+ @staticmethod
+ def get_config_model() -> type["LLMShieldProxyGuardrailConfigModel"]:
+ from litellm.types.proxy.guardrails.guardrail_hooks.llm_shield_proxy import (
+ LLMShieldProxyGuardrailConfigModel,
+ )
+
+ return LLMShieldProxyGuardrailConfigModel
+
+ async def async_pre_call_deployment_hook(
+ self,
+ kwargs: MutableRequest,
+ call_type: "CallTypes | None",
+ ) -> MutableRequest | None:
+ """Redacts a model-level guardrail's request, and keeps it out of the response cache.
+
+ Outside the proxy this hook is the only redaction step, and the deployment post-call
+ hook the only restoration step. LiteLLM builds the cache key after this hook, from
+ the redacted request, and a cache hit returns before the post-call hook runs. So a
+ cached reply would either reach the caller unrestored or, stored after restoration,
+ hand this caller's values to the next caller whose redacted request matches. The
+ request is therefore neither read from nor written to the cache. Inside the proxy
+ this hook does not redact -- the proxy's pre-call hook already ran -- and caching
+ is left alone, because the proxy restores after the cache write.
+
+ A streamed request is refused once redacted. No hook restores an SDK stream, and the
+ stream's cache writer reads the request from before this hook, so it would also be
+ cached despite the bypass.
+ """
+ before: Final = self._minted_session_id(kwargs)
+ _ = await super().async_pre_call_deployment_hook(kwargs, call_type)
+ session_id: Final = self._minted_session_id(kwargs)
+ if session_id is None or session_id == before:
+ return kwargs
+ if kwargs.get("stream") is True:
+ raise GuardrailRaisedException(
+ guardrail_name=self.guardrail_name,
+ message=(
+ "LLM Shield Proxy cannot restore a streamed reply for a model-level guardrail "
+ "outside the LiteLLM proxy; send the request through the proxy or without stream=True."
+ ),
+ )
+ metadata: Final = as_object(kwargs.get("litellm_metadata"))
+ if metadata is not None:
+ metadata[_DEPLOYMENT_RESTORE_KEY] = session_id
+ cache_controls: Final = as_object(kwargs.get("cache"))
+ kwargs["cache"] = {**(cache_controls or {}), "no-cache": True, "no-store": True}
+ return kwargs
+
+ async def async_post_call_success_deployment_hook(
+ self,
+ request_data: MutableRequest,
+ response: "LLMResponseTypes",
+ call_type: "CallTypes | None",
+ ) -> "LLMResponseTypes | None":
+ """Restores the reply here only when the deployment pre-call hook redacted it.
+
+ LiteLLM caches what this hook returns. Inside the proxy the request was redacted by
+ the proxy's pre-call hook and the proxy's post-call hook restores the reply after
+ the cache write, so restoring here as well would cache this caller's plaintext under
+ a key built from the redacted request. Outside the proxy nothing restores later, and
+ the pre-call deployment hook has already kept that request out of the cache.
+ """
+ metadata: Final = as_object(request_data.get("litellm_metadata"))
+ marker: Final = metadata.get(_DEPLOYMENT_RESTORE_KEY) if metadata is not None else None
+ session_id: Final = self._minted_session_id(request_data)
+ if session_id is None or marker != session_id:
+ return None
+ return await super().async_post_call_success_deployment_hook(request_data, response, call_type)
+
+ def _headers(self, session_id: str) -> JsonBody:
+ headers: Final[JsonBody] = {
+ "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: JsonBody) -> Mapping[str, object]:
+ """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
+ blocks the request instead of passing it through.
+ """
+ try:
+ response: Final = 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 Proxy returned %s for %s", exc.response.status_code, path)
+ raise GuardrailRaisedException(
+ guardrail_name=self.guardrail_name,
+ 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 Proxy call to %s failed", path)
+ raise GuardrailRaisedException(
+ guardrail_name=self.guardrail_name,
+ message="LLM Shield Proxy is unreachable; blocking the request.",
+ ) from exc
+
+ async def _redact(self, texts: Sequence[str], session_id: str) -> Sequence[str]:
+ 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)}
+ 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: object, sent: Sequence[str], operation: str) -> Sequence[str]:
+ """Guards the positional mapping the callers rely on to write results back."""
+ entries: Final = as_array(returned)
+ texts: Final = tuple(entry for entry in entries or () if isinstance(entry, str))
+ 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 texts
+
+ @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("litellm_metadata", {})
+ if isinstance(metadata, dict):
+ metadata[_SESSION_METADATA_KEY] = session_id
+ return session_id
+
+ @staticmethod
+ def _minted_session_id(data: MutableRequest) -> str | None:
+ """The vault id this process minted for `data`, or None if it has none.
+
+ Read only from `litellm_metadata`, the proxy-private store `_mint_session_id` writes
+ to. A caller can populate `metadata`; they cannot populate this.
+ """
+ metadata: Final = as_object(data.get("litellm_metadata"))
+ existing: Final = metadata.get(_SESSION_METADATA_KEY) if metadata is not None else None
+ return existing if isinstance(existing, str) and existing.startswith(_VAULT_PREFIX) else None
+
+ @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.
+ """
+ existing: Final = LLMShieldProxyGuardrail._minted_session_id(data)
+ return existing if existing is not None else f"{_VAULT_PREFIX}-{uuid.uuid4().hex}"
+
+ @staticmethod
+ def _locate_request_texts(data: MutableRequest) -> tuple[Sequence[Slot], Sequence[Slot]]:
+ return locate_request_texts(data)
+
+ @log_guardrail_information
+ async def async_pre_call_hook(
+ self,
+ user_api_key_dict: UserAPIKeyAuth,
+ cache: "DualCache",
+ data: MutableRequest,
+ call_type: str,
+ ) -> 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
+
+ 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
+
+ session_id: Final = self._mint_session_id(data)
+ if privileged:
+ 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)
+
+ async def async_post_call_success_hook(
+ self,
+ data: MutableRequest,
+ user_api_key_dict: UserAPIKeyAuth,
+ response: Any,
+ ) -> Any:
+ """Restores the original values in a copy of a non-streaming response.
+
+ The copy is what keeps plaintext out of the response cache. LiteLLM caches the
+ reply it received from the provider, and on some paths -- a native Anthropic dict,
+ an in-memory cache -- it stores the object itself rather than a serialised
+ snapshot. Restoring that object in place would cache this caller's values under a
+ key built from the redacted request, which another caller's identical-looking
+ request then hits. Left untouched, the cached reply holds placeholders, and a hit
+ is restored against the new caller's own vault.
+ """
+ if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.post_call) is not True:
+ return response
+ response = detached(response) # rebind-ok: everything below restores the copy.
+
+ if self._is_anthropic_message_response(response):
+ return await self._restore_anthropic_response(response, 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[SlotSink] = []
+ for choice in choices:
+ message = getattr(choice, "message", None)
+ if message is None:
+ text = read_field(choice, "text")
+ if isinstance(text, str) and text:
+ pending.append((text, functools.partial(write_field, choice, "text")))
+ continue
+ content = getattr(message, "content", None)
+ if isinstance(content, str) and content:
+ pending.append((content, functools.partial(setattr, message, "content")))
+ for tool_call in getattr(message, "tool_calls", None) or ():
+ 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, 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, functools.partial(setattr, legacy, "arguments")))
+ if not pending:
+ return response
+
+ 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
+ def _is_anthropic_message_response(response: object) -> bool:
+ """Anthropic's native /v1/messages reply arrives as a plain dict."""
+ body: Final = as_object(response)
+ return body is not None and body.get("type") == "message" and isinstance(body.get("content"), list)
+
+ async def _restore_anthropic_response(self, response: MutableRequest, data: MutableRequest) -> MutableRequest:
+ """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.
+ """
+ slots: Final[SlotSink] = []
+ for entry in read_list(response, "content"):
+ block = as_object(entry)
+ if block is None:
+ continue
+ kind = block.get("type")
+ 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
+
+ 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_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. 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.
+ """
+ slots: Final[SlotSink] = []
+ for item in getattr(response, "output", None) or ():
+ 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."""
+ await rehydrate_slots(slots, functools.partial(self._rehydrate, session_id=self._session_id(data)))
+ return response
+
+ async def async_post_call_streaming_iterator_hook(
+ self,
+ user_api_key_dict: UserAPIKeyAuth,
+ response: Any,
+ request_data: MutableRequest,
+ ) -> AsyncGenerator[Any, None]:
+ """Restores original values incrementally, without buffering the stream.
+
+ 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.
+
+ 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:
+ yield chunk
+ 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[CarryWindows] = {}
+ 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(detached(chunk)):
+ yield event
+ continue
+ restored_chunk = detached(chunk)
+ last_chunk = restored_chunk
+ for choice in getattr(restored_chunk, "choices", None) or ():
+ await self._restore_choice(choice, carries, session_id)
+ yield restored_chunk
+
+ 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
+
+ async def _restore_choice(self, choice: object, carries: CarryWindows, session_id: str) -> None:
+ """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)
+ index: Final = choice_index(choice)
+ is_final: Final = bool(getattr(choice, "finish_reason", None))
+ if isinstance(choice, TextChoices):
+ await self._restore_text_window(choice, (index, None), carries, session_id, is_final)
+ return
+ if delta is None:
+ return
+
+ 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:
+ await self._flush_finished_choice(delta, index, carries, session_id)
+
+ async def _restore_text_window(
+ self,
+ choice: object,
+ key: CarryKey,
+ carries: CarryWindows,
+ session_id: str,
+ is_final: bool,
+ ) -> None:
+ """Restores a Completions stream choice's `text` through its window."""
+ carry: Final = carries.get(key, "")
+ text: Final = read_field(choice, "text")
+ if not isinstance(text, str) or not text:
+ if not (is_final and carry):
+ return
+ emitted, remaining = await self._stream_step(text if isinstance(text, str) else "", carry, is_final, session_id)
+ carries[key] = remaining # rebind-ok: this stream's window advances.
+ if emitted or text:
+ write_field(choice, "text", emitted)
+
+ async def _restore_content_window(
+ self,
+ delta: Any,
+ key: CarryKey,
+ 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:
+ if is_final and carry:
+ 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[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[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:
+ 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.extend(continuation_delta(tool_index, text))
+ if continuations:
+ existing: Final = tuple(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 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 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[key] = remaining # rebind-ok: this stream's window advances.
+ if not text:
+ continue
+ chunk = self._chunk_for_choice(last_chunk, choice_index)
+ if chunk is None:
+ continue
+ choice: object = chunk.choices[0]
+ delta = read_field(choice, "delta")
+ if isinstance(choice, TextChoices):
+ choice.text = text
+ elif tool_index is None:
+ write_field(delta, "content", text)
+ else:
+ write_field(delta, "content", None)
+ write_field(delta, "tool_calls", continuation_delta(tool_index, 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[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 and not isinstance(kept, TextChoices):
+ return None
+ kept.index = index
+ kept.finish_reason = None
+ chunk.choices = [kept]
+ if hasattr(chunk, "usage"):
+ del chunk.usage
+ 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)``."""
+ body: Final = await self._call_shield(
+ _REHYDRATE_STREAM_PATH,
+ session_id,
+ {"text": text, "carry": carry, "final": final},
+ )
+ 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 Proxy stream rehydration returned an unexpected payload.",
+ )
+ return emitted, remaining
+
+ @log_guardrail_information
+ async def apply_guardrail(
+ self,
+ inputs: GenericGuardrailAPIInputs,
+ request_data: MutableRequest,
+ input_type: Literal["request", "response"],
+ logging_obj: Optional["LiteLLMLoggingObj"] = None,
+ ) -> GenericGuardrailAPIInputs:
+ """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 = 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
+
+ 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(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[str]] = list(replaced) # mutable-ok: sliced into the texts list.
+
+ for write, replacement in zip(writers, restored_values[len(text_list) :]):
+ write(replacement)
+ merged: Final[JsonBody] = {**inputs}
+ if text_list:
+ merged["texts"] = restored_values[: len(text_list)]
+ if restored_calls:
+ merged["tool_calls"] = restored_calls
+ return merged
diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/payload.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/payload.py
new file mode 100644
index 00000000000..644ac76efb8
--- /dev/null
+++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/payload.py
@@ -0,0 +1,190 @@
+import copy
+import functools
+from collections.abc import Awaitable, Callable, Sequence
+from typing import (
+ Final,
+ TypeAlias,
+)
+
+MutableRequest: TypeAlias = dict[str, object]
+
+JsonBody: TypeAlias = dict[str, object]
+
+MAX_JSON_DEPTH: Final = 64
+
+Slot: TypeAlias = tuple[str, Callable[[str], None]]
+
+StreamStep: TypeAlias = Callable[[str, str, bool], Awaitable[tuple[str, str]]]
+
+Rehydrate: TypeAlias = Callable[[Sequence[str]], Awaitable[Sequence[str]]]
+
+SlotSink: TypeAlias = list[Slot]
+
+MutableSeq: TypeAlias = list[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 detached(value: object) -> object:
+ """A deep copy of a reply or chunk, for restoring without touching LiteLLM's own object.
+
+ LiteLLM keeps the object it handed the hooks to fill its response cache and its
+ logs, so writing restored plaintext into that object would put it there too.
+ """
+ return copy.deepcopy(value)
+
+
+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, 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, functools.partial(entries.__setitem__, index)))
+
+
+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_text_parts(container: MutableRequest, key: str, slots: SlotSink) -> None:
+ """Collects the `text` of every part in the list held at `key`."""
+ for entry in as_array(container.get(key)) or ():
+ part = as_object(entry)
+ if part is not None:
+ collect(part, "text", 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
+
+
+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.
+ """
+ fields: Final = as_object(holder)
+ if fields is not None:
+ return fields.get(name)
+ 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)
+ if isinstance(value, tuple):
+ return value
+ return tuple(as_array(value) or ())
+
+
+def write_field(holder: object, name: str, value: object) -> 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, *, 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_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[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 is_container(current) and current:
+ raise RequestTooDeep("json")
+ continue
+ 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, functools.partial(current_object.__setitem__, key)))
+ else:
+ pending.append((value, current_depth + 1))
+ continue
+ entries = as_array(current)
+ if entries is not None:
+ for index, value in enumerate(entries):
+ if isinstance(value, str) and value:
+ slots.append((value, functools.partial(entries.__setitem__, index)))
+ else:
+ pending.append((value, current_depth + 1))
+
+
+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 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"):
+ 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", "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)))
+
+
+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)
diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/request_walk.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/request_walk.py
new file mode 100644
index 00000000000..c9c0d2ccc47
--- /dev/null
+++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/request_walk.py
@@ -0,0 +1,363 @@
+from collections.abc import Sequence
+from typing import (
+ Final,
+)
+
+from .payload import (
+ MAX_JSON_DEPTH,
+ MutableRequest,
+ RequestTooDeep,
+ Slot,
+ SlotSink,
+ as_array,
+ as_object,
+ collect,
+ collect_entry,
+ collect_json_leaves,
+ collect_text_parts,
+ is_container,
+ read_list,
+)
+
+PRIVILEGED_ROLES: Final = frozenset({"system", "developer"})
+
+MAX_CONTENT_DEPTH: Final = 8
+
+SCHEMA_STRUCTURAL_KEYWORDS: Final = frozenset(
+ (
+ "type",
+ "format",
+ "pattern",
+ "required",
+ "dependentRequired",
+ "propertyOrdering",
+ "discriminator",
+ "contentEncoding",
+ "contentMediaType",
+ "$ref",
+ "$id",
+ "$schema",
+ "$anchor",
+ "$dynamicRef",
+ "$dynamicAnchor",
+ "$recursiveRef",
+ "$recursiveAnchor",
+ "$vocabulary",
+ )
+)
+
+SCHEMA_VALUE_KEYWORDS: Final = frozenset(("examples", "default"))
+
+SCHEMA_LITERAL_KEYWORDS: Final = frozenset(("enum", "const"))
+
+SCHEMA_MAP_KEYWORDS: Final = frozenset(
+ ("properties", "patternProperties", "$defs", "definitions", "dependentSchemas", "dependencies")
+)
+
+
+def collect_prompt(data: MutableRequest, slots: SlotSink) -> None:
+ """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)
+ return
+ prompt_object: Final = as_object(prompt)
+ if prompt_object is not None:
+ variables: Final = as_object(prompt_object.get("variables"))
+ if variables is not None:
+ for name in tuple(variables):
+ collect(variables, name, slots)
+ typed = as_object(variables[name])
+ if typed is not None:
+ collect(typed, "text", slots)
+ return
+ entries: Final = as_array(prompt)
+ if entries is None:
+ return
+ for index in range(len(entries)):
+ collect_entry(entries, index, slots)
+
+
+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. Content nested past the
+ bound raises `RequestTooDeep` rather than being skipped.
+ """
+ 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]
+ cursor += 1
+ content = node.get("content")
+ if isinstance(content, str):
+ collect(node, "content", slots)
+ continue
+ if depth >= MAX_CONTENT_DEPTH and content:
+ raise RequestTooDeep("content")
+ for item in as_array(content) or ():
+ part = as_object(item)
+ if part is None:
+ continue
+ collect(part, "text", slots)
+ if part.get("type") == "tool_use":
+ collect_json_leaves(part.get("input"), slots, strict=True)
+ source = as_object(part.get("source")) if part.get("type") == "document" else None
+ if source is not None:
+ collect(part, "title", slots)
+ collect(part, "context", slots)
+ if source.get("type") == "text":
+ collect(source, "data", slots)
+ elif source.get("type") == "content":
+ pending.append((source, depth + 1))
+ if "content" in part:
+ pending.append((part, depth + 1))
+
+
+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 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 = as_object(message.get("function_call"))
+ if legacy is not 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
+ collect_text_parts(data, "system", 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, except for
+ system and developer items in it, which go to the privileged sink like their Chat
+ counterparts.
+ """
+ collect(data, "instructions", privileged)
+ request_input: Final = data.get("input")
+ if isinstance(request_input, str):
+ collect(data, "input", slots)
+ return
+ entries: Final = as_array(request_input)
+ if entries is None:
+ return
+ for index, entry in enumerate(entries):
+ if isinstance(entry, str):
+ collect_entry(entries, index, slots)
+ continue
+ item = as_object(entry)
+ if item is None:
+ continue
+ collect_content(item, privileged if item.get("role") in PRIVILEGED_ROLES else slots)
+ collect(item, "arguments", slots)
+ collect(item, "output", slots)
+ collect_text_parts(item, "output", slots)
+ collect(item, "input", slots)
+ collect(item, "code", slots)
+ collect_text_parts(item, "summary", slots)
+
+
+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. `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`.
+ """
+ for key in ("tools", "functions"):
+ for entry in as_array(data.get(key)) or ():
+ tool = as_object(entry)
+ if tool is None:
+ continue
+ 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)
+
+
+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. 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
+ 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[tuple[object, int]]] = [(schema, 0)] # mutable-ok: local walk stack.
+ while pending:
+ node, depth = pending.pop()
+ if depth > MAX_JSON_DEPTH:
+ if is_container(node) and node:
+ raise RequestTooDeep("schema")
+ continue
+ 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
+ schema_object = as_object(node)
+ if schema_object is None:
+ continue
+ 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(schema_object, keyword, slots)
+ collect_json_leaves(value, slots, strict=True)
+ elif keyword in SCHEMA_VALUE_KEYWORDS:
+ collect(schema_object, keyword, privileged)
+ collect_json_leaves(value, privileged, strict=True)
+ elif subschemas is not None:
+ pending.extend((child, depth + 1) for child in subschemas.values())
+ elif isinstance(value, str):
+ collect(schema_object, keyword, privileged)
+ elif is_container(value):
+ pending.append((value, depth + 1))
+
+
+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 free text goes
+ to the privileged sink, and its names and types stay as sent.
+ """
+ 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 = 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,
+ ):
+ wrapper = as_object(declared)
+ if wrapper is not None:
+ collect(wrapper, "description", 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 = 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 = 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)
+
+
+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 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 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`, 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.
+
+ 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.
+
+ `extra_body` is walked the same way as the request itself. LiteLLM merges it over
+ the transformed request just before sending, so a field there -- `input`,
+ `messages`, `system` -- replaces the redacted one on the wire.
+ """
+ slots: Final[SlotSink] = []
+ privileged: Final[SlotSink] = []
+ for payload in (data, as_object(data.get("extra_body"))):
+ if payload is None:
+ continue
+ for entry in read_list(payload, "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)
+ collect_tool_arguments(message, sink)
+ collect_responses_fields(payload, slots, privileged)
+ collect_prompt(payload, slots)
+ collect_system(payload, privileged)
+ collect_tool_definitions(payload, slots, privileged)
+ collect_output_contracts(payload, slots, privileged)
+ collect_user_locations(payload, privileged)
+ collect_end_user_ids(payload, privileged)
+ return tuple(slots), tuple(privileged)
diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/stream_restorers.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/stream_restorers.py
new file mode 100644
index 00000000000..4ef8bbdbb3e
--- /dev/null
+++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/stream_restorers.py
@@ -0,0 +1,345 @@
+import copy
+import functools
+import itertools
+import json
+import re
+from enum import Enum
+from types import MappingProxyType
+from typing import (
+ Final,
+ TypeAlias,
+)
+
+from .payload import (
+ JsonBody,
+ MutableRequest,
+ Rehydrate,
+ SlotSink,
+ StreamStep,
+ as_object,
+ collect_response_item,
+ read_field,
+ read_list,
+ rehydrate_slots,
+ write_field,
+)
+
+ANTHROPIC_DELTA_FIELDS: Final = MappingProxyType({"text_delta": "text", "input_json_delta": "partial_json"})
+
+SSE_EVENT_BOUNDARY: Final = re.compile(rb"(\r?\n\r?\n)")
+
+SSE_OPENINGS: Final = (b"event:", b"data:", b"id:", b"retry:", b":")
+
+RESPONSES_BINARY_DELTAS: Final = frozenset(("response.audio.delta",))
+
+RESPONSES_STRUCTURAL_FIELDS: Final = frozenset(
+ ("type", "id", "item_id", "call_id", "name", "server_label", "status", "obfuscation")
+)
+
+RESPONSES_TERMINAL_EVENTS: Final = frozenset(("response.completed", "response.incomplete"))
+
+CarryKey: TypeAlias = tuple[int, int | None]
+CarryWindows: TypeAlias = dict[CarryKey, str]
+
+ResponsesStreamKey: TypeAlias = tuple[str, object, object, object]
+
+
+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
+ 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)
+
+
+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}}]
+
+
+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.
+
+ 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""
+ 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."""
+ if isinstance(chunk, str):
+ self._as_text = True
+ 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:
+ 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
+ return ()
+ cut: Final = boundaries[-1].end()
+ self._pending = buffered[cut:]
+ parts: Final = SSE_EVENT_BOUNDARY.split(buffered[:cut])
+ restored: Final = tuple(
+ [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."""
+ held: Final = self._pending
+ self._pending = b""
+ if not self._is_sse:
+ return self._emit(held)
+ tail: Final = await self._restore_event(held) if held.strip() else held
+ flushed: Final = await self._flush_all()
+ 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:
+ parsed: Final[object] = json.loads(line[len("data:") :])
+ except ValueError:
+ return block
+ event: Final = as_object(parsed)
+ if event is None:
+ 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] = {
+ "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)])
+ return b"".join(flushed)
+
+
+class ResponsesStreamRestorer:
+ """Restores a Responses API event stream.
+
+ 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:
+
+ - 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:
+ self._step: Final = step
+ self._rehydrate: Final = rehydrate
+ 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."""
+ kind: Final = responses_event_type(event)
+ if kind is None:
+ return (event,)
+ if kind.endswith(".delta") and kind not in RESPONSES_BINARY_DELTAS:
+ await self._restore_delta(event, kind)
+ return (event,)
+ 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)
+ 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 (*flushed, 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)])
+ return tuple(itertools.chain.from_iterable(flushed))
+
+ 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: ResponsesStreamKey) -> 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) -> 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
+ 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 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.
+ """
+ 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 name in RESPONSES_STRUCTURAL_FIELDS or name.endswith("_id"):
+ continue
+ if isinstance(value, str) and value:
+ slots.append((value, functools.partial(write_field, event, name)))
diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py
index 46026c12d24..8f1140cef62 100644
--- a/litellm/types/guardrails.py
+++ b/litellm/types/guardrails.py
@@ -145,6 +145,7 @@ class SupportedGuardrailIntegrations(Enum):
STRAIKER = "straiker"
ALICE = "alice"
AGENT_365 = "agent_365"
+ LLM_SHIELD_PROXY = "llm_shield_proxy"
CONDUCT = "conduct"
diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/llm_shield_proxy.py b/litellm/types/proxy/guardrails/guardrail_hooks/llm_shield_proxy.py
new file mode 100644
index 00000000000..967d7ee75c2
--- /dev/null
+++ b/litellm/types/proxy/guardrails/guardrail_hooks/llm_shield_proxy.py
@@ -0,0 +1,24 @@
+from pydantic import Field
+
+from .base import GuardrailConfigModel
+
+
+class LLMShieldProxyGuardrailConfigModel(GuardrailConfigModel):
+ api_key: str | None = Field(
+ default=None,
+ description=(
+ "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 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 Proxy"
diff --git a/ruff-strict.toml b/ruff-strict.toml
index 899a8ff3af5..9e0b11a1af4 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_proxy/llm_shield_proxy.py" = ["ANN401"]
[lint.mccabe]
max-complexity = 15
diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py b/tests/unit/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py
new file mode 100644
index 00000000000..ed34c07022e
--- /dev/null
+++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py
@@ -0,0 +1,2003 @@
+import asyncio
+import json
+from types import SimpleNamespace
+from unittest.mock import AsyncMock
+
+import pytest
+from httpx import Request, Response
+
+import litellm
+from litellm.caching.caching_handler import _PENDING_CACHE_WRITES
+from litellm.caching.in_memory_cache import InMemoryCache
+from litellm.exceptions import GuardrailRaisedException
+from litellm.proxy.guardrails.guardrail_hooks.llm_shield_proxy.llm_shield_proxy import (
+ GUARDRAIL_NAME,
+ LLMShieldProxyGuardrail,
+)
+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,
+ TextChoices,
+ TextCompletionResponse,
+ Usage,
+)
+
+
+def _guardrail(**overrides: object) -> LLMShieldProxyGuardrail:
+ 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 LLMShieldProxyGuardrail(**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: 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]
+ 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 _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)
+
+
+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", {})
+ monkeypatch.setenv("LLM_SHIELD_PROXY_API_KEY", "test-key")
+
+ init_guardrails_v2(
+ all_guardrails=[
+ {
+ "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, LLMShieldProxyGuardrail)]
+ assert len(registered) == 1
+ assert registered[0].guardrail_name == "llm_shield_proxy"
+
+
+class TestLLMShieldProxyInitialization:
+ def test_api_base_defaults_to_localhost(self, monkeypatch: pytest.MonkeyPatch):
+ 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_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):
+ 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
+ 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_text_is_untouched(self):
+ """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
+ let `input` reach the provider unredacted.
+ """
+ guardrail = _guardrail()
+ mock = _mock_post(guardrail)
+ 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")
+
+ 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 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_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_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_responses_tool_output_parts_are_redacted(self):
+ """A function_call_output can carry its result as a list of input_text parts."""
+ guardrail = _guardrail()
+ mock = _mock_post(guardrail, {"texts": ["sent to [EMAIL_1]"]})
+
+ data = {
+ "input": [
+ {
+ "type": "function_call_output",
+ "call_id": "c1",
+ "output": [{"type": "input_text", "text": "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 mock.call_args_list[0].kwargs["json"]["texts"] == ["sent to jane.doe@example.com"]
+ assert data["input"][0]["output"][0]["text"] == "sent to [EMAIL_1]"
+
+ @pytest.mark.asyncio
+ async def test_responses_custom_tool_call_input_is_redacted(self):
+ """A replayed custom_tool_call carries its payload in `input`, not `arguments`."""
+ guardrail = _guardrail()
+ _mock_post(guardrail, {"texts": ["email [EMAIL_1]"]})
+
+ data = {
+ "input": [
+ {"type": "custom_tool_call", "call_id": "c1", "name": "mail", "input": "email 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]["input"] == "email [EMAIL_1]"
+ assert data["input"][0]["name"] == "mail"
+
+ @pytest.mark.asyncio
+ async def test_extra_body_overrides_are_redacted(self):
+ """LiteLLM merges `extra_body` over the request just before sending, so its fields win on the wire."""
+ guardrail = _guardrail()
+ mock = _mock_post(guardrail, {"texts": ["safe", "Mail [EMAIL_1]", "[EMAIL_1]"]})
+
+ data = {
+ "model": "gpt-4o",
+ "input": "safe",
+ "extra_body": {
+ "input": "alice@example.com",
+ "messages": [{"role": "user", "content": "Mail alice@example.com"}],
+ "service_tier": "flex",
+ },
+ }
+ 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"] == ["safe", "Mail alice@example.com", "alice@example.com"]
+ assert data["extra_body"]["input"] == "[EMAIL_1]"
+ assert data["extra_body"]["messages"][0]["content"] == "Mail [EMAIL_1]"
+ assert data["extra_body"]["service_tier"] == "flex"
+
+ def test_extra_body_system_text_is_privileged(self) -> None:
+ """An application-authored override is no more restorable than the field it replaces."""
+ data = {"messages": [{"role": "user", "content": "U"}], "extra_body": {"system": "S", "instructions": "I"}}
+
+ caller, privileged = LLMShieldProxyGuardrail._locate_request_texts(data)
+
+ assert [text for text, _ in caller] == ["U"]
+ assert sorted(text for text, _ in privileged) == ["I", "S"]
+
+ @pytest.mark.asyncio
+ async def test_anthropic_text_documents_are_redacted(self):
+ """A document block carries text inline, as `source.data` or as `source.content`."""
+ guardrail = _guardrail()
+ mock = _mock_post(guardrail, {"texts": ["[NAME_1] notes", "[EMAIL_1]", "Reach [EMAIL_2]"]})
+
+ data = {
+ "messages": [
+ {
+ "role": "user",
+ "content": [
+ {
+ "type": "document",
+ "title": "Jane Doe notes",
+ "source": {"type": "text", "media_type": "text/plain", "data": "alice@example.com"},
+ },
+ {
+ "type": "document",
+ "source": {"type": "content", "content": [{"type": "text", "text": "Reach bob@example.com"}]},
+ },
+ {"type": "document", "source": {"type": "base64", "media_type": "application/pdf", "data": "JVBERi0x"}},
+ ],
+ }
+ ]
+ }
+ await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="anthropic_messages")
+
+ sent = mock.call_args_list[0].kwargs["json"]["texts"]
+ assert sent == ["Jane Doe notes", "alice@example.com", "Reach bob@example.com"]
+ blocks = data["messages"][0]["content"]
+ assert blocks[0]["source"]["data"] == "[EMAIL_1]"
+ assert blocks[1]["source"]["content"][0]["text"] == "Reach [EMAIL_2]"
+ assert blocks[2]["source"]["data"] == "JVBERi0x", "binary sources are not text"
+
+ @pytest.mark.asyncio
+ async def test_responses_code_interpreter_code_is_redacted(self):
+ """A replayed code_interpreter_call carries the code the model wrote, which the reply side restores."""
+ guardrail = _guardrail()
+ _mock_post(guardrail, {"texts": ["send('[EMAIL_1]')"]})
+
+ data = {"input": [{"type": "code_interpreter_call", "id": "ci_1", "code": "send('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]["code"] == "send('[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_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_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_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.
+
+ 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()
+ 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")
+
+ assert tool_use["input"]["a"]["b"]["c"]["d"]["to"] == "[EMAIL_1]"
+
+ @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_responses_typed_prompt_variables_are_redacted(self):
+ """A variable can be a typed input rather than a string; its `text` is caller text."""
+ guardrail = _guardrail()
+ _mock_post(guardrail, {"texts": ["[EMAIL_1]"]})
+
+ data = {
+ "prompt": {
+ "id": "pmpt_123",
+ "variables": {
+ "customer": {"type": "input_text", "text": "jane.doe@example.com"},
+ "logo": {"type": "input_image", "image_url": "https://example.com/logo.png"},
+ },
+ }
+ }
+ await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses")
+
+ assert data["prompt"]["variables"]["customer"] == {"type": "input_text", "text": "[EMAIL_1]"}
+ assert data["prompt"]["variables"]["logo"]["image_url"] == "https://example.com/logo.png"
+
+ @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()
+ 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"
+
+ @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_their_free_text_and_nothing_else(self):
+ """Every string is collected except what must reach the model verbatim: names,
+ types, formats, patterns and required lists."""
+ data = {
+ "tools": [
+ {
+ "type": "function",
+ "function": {
+ "name": "lookup",
+ "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"], "const": "a", "description": "enum"},
+ "deep": {"type": "array", "items": {"type": "object", "description": "nested"}},
+ "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",
+ },
+ },
+ }
+ ]
+ }
+ 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",
+ "defined",
+ "dependent",
+ "enum",
+ "example",
+ "named",
+ "named-type",
+ "nested",
+ "object-default",
+ "title",
+ "top",
+ "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]
+ restored = await guardrail.async_post_call_success_hook(data=data, user_api_key_dict=None, response=reply)
+
+ assert json.loads(restored.choices[0].message.tool_calls[0].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):
+ 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:
+ @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_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])])
+ 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_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 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["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):
+ """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
+
+
+ @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"),
+ 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"}],
+ "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"}],
+ "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"),
+ ],
+ )
+ 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)
+
+ 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) -> 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]"]})
+
+ 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) -> 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
+ 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)
+
+ def test_responses_system_and_developer_items_are_privileged(self) -> None:
+ """Responses `input` carries system and developer turns as items, like Chat messages.
+
+ In the caller's vault, a caller could have the model echo a placeholder out of a
+ system message they cannot see and receive the plaintext behind it.
+ """
+ data = {
+ "input": [
+ {"role": "system", "content": "S"},
+ {"type": "message", "role": "developer", "content": [{"type": "input_text", "text": "D"}]},
+ {"role": "user", "content": "U"},
+ ]
+ }
+ caller, privileged = LLMShieldProxyGuardrail._locate_request_texts(data)
+
+ assert [text for text, _ in caller] == ["U"]
+ assert [text for text, _ in privileged] == ["S", "D"]
+
+
+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_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_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_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."""
+ 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"]
+
+
+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]"}'
+
+
+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
+
+ @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."""
+
+ 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]),
+ )
+
+ (out,) = await _restore_stream(guardrail, [completed])
+
+ restored_block, restored_call = out.response.output[0].content[0], out.response.output[1]
+ assert restored_block["text"] == "Mail a@example.com"
+ assert restored_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"}
+
+ @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
+ )
+
+ (out,) = await _restore_stream(guardrail, [event])
+
+ assert out.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 == []
+
+
+class TestResponseCacheIsolation:
+ """The reply LiteLLM caches must keep its placeholders.
+
+ Placeholders are numbered per request, so two callers' redacted requests can be
+ identical and share a cache key. LiteLLM keeps the provider's reply object -- for a
+ native Anthropic dict or an in-memory cache, the object itself -- so restoring it in
+ place would hand one caller's values to the next caller who hits that key.
+ """
+
+ @pytest.mark.asyncio
+ async def test_a_cache_hit_is_restored_against_the_new_callers_vault(self):
+ cache = InMemoryCache()
+ provider_reply = {"type": "message", "content": [{"type": "text", "text": "Repeat [EMAIL_1]"}]}
+ cache.set_cache("redacted-request", provider_reply)
+
+ alice, _ = _shielded({"[EMAIL_1]": "alice@example.com"})
+ to_alice = await alice.async_post_call_success_hook(
+ data={"messages": []}, user_api_key_dict=None, response=cache.get_cache("redacted-request")
+ )
+ bob, _ = _shielded({"[EMAIL_1]": "bob@example.com"})
+ to_bob = await bob.async_post_call_success_hook(
+ data={"messages": []}, user_api_key_dict=None, response=cache.get_cache("redacted-request")
+ )
+
+ assert to_alice["content"][0]["text"] == "Repeat alice@example.com"
+ assert to_bob["content"][0]["text"] == "Repeat bob@example.com"
+ assert cache.get_cache("redacted-request")["content"][0]["text"] == "Repeat [EMAIL_1]"
+
+ @pytest.mark.asyncio
+ async def test_a_model_response_is_restored_as_a_copy(self):
+ """A reply as `acompletion` returns it, hidden params and all."""
+ reply = await litellm.acompletion(
+ model="gpt-4o-mini", messages=[{"role": "user", "content": "hi"}], mock_response="Mail [EMAIL_1]"
+ )
+ guardrail, _ = _shielded({"[EMAIL_1]": "a@example.com"})
+
+ restored = await guardrail.async_post_call_success_hook(
+ data={"messages": []}, user_api_key_dict=None, response=reply
+ )
+
+ assert restored.choices[0].message.content == "Mail a@example.com"
+ assert reply.choices[0].message.content == "Mail [EMAIL_1]"
+
+ @pytest.mark.asyncio
+ async def test_responses_reply_is_restored_as_a_copy(self):
+ guardrail, _ = _shielded({"[EMAIL_1]": "a@example.com"})
+ block = {"type": "output_text", "text": "Mail [EMAIL_1]"}
+ reply = SimpleNamespace(output=[SimpleNamespace(content=[block])])
+
+ restored = await guardrail.async_post_call_success_hook(
+ data={"messages": []}, user_api_key_dict=None, response=reply
+ )
+
+ assert restored.output[0].content[0]["text"] == "Mail a@example.com"
+ assert block["text"] == "Mail [EMAIL_1]"
+
+ @pytest.mark.asyncio
+ async def test_stream_chunks_are_restored_as_copies(self):
+ """LiteLLM assembles the reply it caches from the chunks it yielded."""
+ guardrail, _ = _shielded({"[EMAIL_1]": "a@example.com"})
+ chunk = _chunk("Mail [EMAIL_1]", finish_reason="stop")
+ event = OutputTextDeltaEvent(
+ type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA,
+ item_id="msg_1",
+ output_index=0,
+ content_index=0,
+ delta="Mail [EMAIL_1]",
+ )
+
+ out = await _restore_stream(guardrail, [chunk, event])
+
+ assert out[0].choices[0].delta.content == "Mail a@example.com"
+ assert chunk.choices[0].delta.content == "Mail [EMAIL_1]"
+ assert "Mail [EMAIL_1]" == event.delta
+
+
+class TestCompletionsRestoration:
+ """`/v1/completions` replies carry their text on the choice, with no message or delta."""
+
+ VAULT = {"[EMAIL_1]": "a@example.com"}
+
+ @pytest.mark.asyncio
+ async def test_completions_reply_is_restored(self):
+ guardrail, _ = _shielded(self.VAULT)
+ reply = TextCompletionResponse(choices=[TextChoices(index=0, text="Mail [EMAIL_1]", finish_reason="stop")])
+
+ restored = await guardrail.async_post_call_success_hook(
+ data={"prompt": "x"}, user_api_key_dict=None, response=reply
+ )
+
+ assert restored.choices[0].text == "Mail a@example.com"
+
+ @pytest.mark.asyncio
+ async def test_completions_stream_is_restored_across_chunks(self):
+ guardrail, _ = _shielded(self.VAULT)
+ chunks = [
+ TextCompletionResponse(choices=[TextChoices(index=0, text="Mail [EMAI")]),
+ TextCompletionResponse(choices=[TextChoices(index=0, text="L_1] now", finish_reason="stop")]),
+ ]
+
+ out = await _restore_stream(guardrail, chunks)
+
+ assert "".join(chunk.choices[0].text for chunk in out) == "Mail a@example.com now"
+
+ @pytest.mark.asyncio
+ async def test_completions_stream_without_finish_reason_is_flushed(self):
+ guardrail, _ = _shielded(self.VAULT)
+ chunks = [TextCompletionResponse(choices=[TextChoices(index=0, text="Mail [EMAIL_1")])]
+
+ out = await _restore_stream(guardrail, chunks)
+
+ assert "".join(chunk.choices[0].text or "" for chunk in out) == "Mail [EMAIL_1"
+
+
+class TestProxyWiring:
+ def test_dashboard_config_model_is_exposed(self):
+ """The guardrail garden reads the provider's fields from `get_config_model`."""
+ model = LLMShieldProxyGuardrail.get_config_model()
+
+ assert model is not None
+ assert {"api_key", "api_base"} <= set(model.model_fields)
+
+ @pytest.mark.asyncio
+ async def test_the_deployment_hook_leaves_a_proxy_reply_for_the_proxy_hook(self):
+ """Inside the proxy the deployment hook must not restore: LiteLLM caches what it returns.
+
+ A proxy request was redacted by the proxy's pre-call hook, so it carries no
+ deployment-restore marker, and the proxy's post-call hook restores it after the cache
+ write. A caller-sent marker that does not match the minted vault id is ignored.
+ """
+ guardrail, shield = _shielded({"[EMAIL_1]": "a@example.com"})
+ reply = ModelResponse(choices=[Choices(index=0, message=Message(role="assistant", content="[EMAIL_1]"))])
+ data = {"messages": [], "guardrails": [GUARDRAIL_NAME], "litellm_metadata": {}}
+ LLMShieldProxyGuardrail._mint_session_id(data)
+ data["litellm_metadata"]["llm_shield_restore_at_deployment"] = "caller-chosen"
+
+ result = await guardrail.async_post_call_success_deployment_hook(
+ request_data=data, response=reply, call_type=None
+ )
+
+ assert result is None
+ assert reply.choices[0].message.content == "[EMAIL_1]"
+ assert shield.urls == []
+
+ @pytest.mark.asyncio
+ async def test_model_level_use_outside_the_proxy_is_restored_and_never_cached(self, monkeypatch):
+ """SDK use with model-level `guardrails`: the deployment hooks are the only redact and
+ restore steps, so the reply is restored there, and the request bypasses the cache --
+ its key is built from the redacted request, and a cache hit would skip restoration.
+ """
+ vault = {"[EMAIL_1]": "alice@example.com"}
+
+ async def shield(url: str, headers: dict, json: dict, timeout: float) -> Response:
+ texts = json["texts"]
+ if url.endswith("/redact"):
+ return _response({"texts": [t.replace("alice@example.com", "[EMAIL_1]") for t in texts]})
+ restored = []
+ for text in texts:
+ for placeholder, original in vault.items():
+ text = text.replace(placeholder, original)
+ restored.append(text)
+ return _response({"texts": restored})
+
+ guardrail = _guardrail(event_hook=["pre_call", "post_call"], default_on=False)
+ guardrail.async_handler.post = shield # type: ignore[method-assign]
+ cache = InMemoryCache()
+ monkeypatch.setattr(litellm, "callbacks", [guardrail])
+ monkeypatch.setattr(litellm, "cache", litellm.Cache(type="local"))
+ monkeypatch.setattr(litellm.cache, "cache", cache)
+
+ reply = await litellm.acompletion(
+ model="gpt-4o-mini",
+ messages=[{"role": "user", "content": "Repeat alice@example.com"}],
+ mock_response="Repeat [EMAIL_1]",
+ guardrails=[GUARDRAIL_NAME],
+ )
+
+ # LiteLLM writes the cache from background tasks; let them land before looking.
+ await asyncio.gather(*_PENDING_CACHE_WRITES)
+
+ assert reply.choices[0].message.content == "Repeat alice@example.com"
+ assert cache.cache_dict == {}, "the redacted request's reply must not be cached"
+
+ @pytest.mark.asyncio
+ async def test_model_level_streaming_outside_the_proxy_is_refused(self, monkeypatch):
+ """No hook restores an SDK stream, and its cache writer misses the bypass, so it fails closed."""
+ guardrail = _guardrail(event_hook=["pre_call", "post_call"], default_on=False)
+ _mock_post(guardrail, {"texts": ["Repeat [EMAIL_1]"]})
+ cache = InMemoryCache()
+ monkeypatch.setattr(litellm, "callbacks", [guardrail])
+ monkeypatch.setattr(litellm, "cache", litellm.Cache(type="local"))
+ monkeypatch.setattr(litellm.cache, "cache", cache)
+
+ with pytest.raises(GuardrailRaisedException, match="cannot restore a streamed reply"):
+ await litellm.acompletion(
+ model="gpt-4o-mini",
+ messages=[{"role": "user", "content": "Repeat alice@example.com"}],
+ mock_response="Repeat [EMAIL_1]",
+ stream=True,
+ guardrails=[GUARDRAIL_NAME],
+ )
+ await asyncio.gather(*_PENDING_CACHE_WRITES)
+
+ assert cache.cache_dict == {}
+
+ @pytest.mark.asyncio
+ async def test_restored_values_are_not_recorded_as_guardrail_telemetry(self):
+ """Guardrail logging is exported to traces even with message logging off, so the
+ restored reply must not land in it."""
+ guardrail, _ = _shielded({"[EMAIL_1]": "alice@example.com"})
+ reply = ModelResponse(choices=[Choices(index=0, message=Message(role="assistant", content="[EMAIL_1]"))])
+ data = {"messages": [], "metadata": {}}
+
+ restored = await guardrail.async_post_call_success_hook(data=data, user_api_key_dict=None, response=reply)
+
+ assert restored.choices[0].message.content == "alice@example.com"
+ assert "alice@example.com" not in json.dumps(data, default=str)
+
+
+class TestStreamUsage:
+ @pytest.mark.asyncio
+ async def test_trailing_flush_does_not_repeat_usage(self):
+ """With n>=2 and include_usage, the last chunk carries usage; a flush copied from it must not.
+
+ Any consumer that sums usage chunks would otherwise count the request twice.
+ """
+ guardrail, _ = _shielded({"[EMAIL_1]": "a@example.com"})
+ both = ModelResponseStream(
+ choices=[
+ StreamingChoices(index=0, delta=Delta(content="Mail [EMAI")),
+ StreamingChoices(index=1, delta=Delta(content="Call [EMAI")),
+ ]
+ )
+ usage_chunk = ModelResponseStream(
+ choices=[StreamingChoices(index=1, delta=Delta(content=None))],
+ usage=Usage(prompt_tokens=5, completion_tokens=7, total_tokens=12),
+ )
+
+ out = await _restore_stream(guardrail, [both, usage_chunk])
+
+ with_usage = [chunk for chunk in out if getattr(chunk, "usage", None) is not None]
+ assert with_usage == [out[1]], "only the provider's own usage chunk carries usage"
+ flushed = out[2:]
+ assert flushed, "the held-back text is flushed at end of stream"
+ assert sorted(chunk.choices[0].index for chunk in flushed) == [0, 1]
+ assert [chunk.choices[0].delta.content for chunk in flushed] == ["[EMAI", "[EMAI"]
diff --git a/ui/litellm-dashboard/public/assets/logos/llm_shield_proxy.svg b/ui/litellm-dashboard/public/assets/logos/llm_shield_proxy.svg
new file mode 100644
index 00000000000..0dd78b078c9
--- /dev/null
+++ b/ui/litellm-dashboard/public/assets/logos/llm_shield_proxy.svg
@@ -0,0 +1,6 @@
+
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 10ca58294b9..684049c54a9 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;
}
@@ -325,6 +327,14 @@ export const GUARDRAIL_PRESETS: Record = {
// MCP-only: default_on is the only activation path on the MCP hook
defaultOn: true,
},
+ 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"],
+ defaultOn: false,
+ },
conduct: {
provider: "Conduct",
guardrailNameSuggestion: "Conduct Guard",
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 a27dd344c95..e756c8c3e58 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
@@ -29,6 +29,7 @@ const EXPECTED_PARTNER_LOGO_FILES: Record = {
straiker: "straiker.svg",
alice: "alice.svg",
agent_365: "microsoft_azure.svg",
+ llm_shield_proxy: "llm_shield_proxy.svg",
conduct: "conduct.png",
};
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 d88a333d6f1..d1ead5f589a 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
@@ -484,6 +484,16 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [
tags: ["Agentic", "MCP", "Tool Misuse", "Observability"],
providerKey: "Agent365",
},
+ {
+ 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 Proxy"],
+ tags: ["PII", "Data Privacy", "Compliance", "Streaming"],
+ providerKey: "LLM Shield Proxy",
+ },
{
id: "conduct",
name: "Conduct Guard",
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 8df7dfb1403..09e0f21c670 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 llmShieldProxyLogo from "../../../../../public/assets/logos/llm_shield_proxy.svg";
import conductLogo from "../../../../../public/assets/logos/conduct.png";
import aporiaLogo from "../../../../../public/assets/logos/aporia.png";
import bedrockLogo from "../../../../../public/assets/logos/bedrock.svg";
@@ -86,6 +87,7 @@ export const guardrail_provider_map: Record = {
QostodianNexus: "qostodian_nexus",
Repelloai: "repelloai",
Alice: "alice",
+ "LLM Shield Proxy": "llm_shield_proxy",
Conduct: "conduct",
};
@@ -211,6 +213,7 @@ export const guardrailLogoMap = {
Straiker: straikerLogo.src,
Alice: aliceLogo.src,
"Microsoft Agent 365": microsoftAzureLogo.src,
+ "LLM Shield Proxy": llmShieldProxyLogo.src,
"Conduct Guard": conductLogo.src,
} satisfies Record;