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..c8ca68a8967 --- /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 = { # mutable-ok: module-level registry, built once and never mutated + SupportedGuardrailIntegrations.LLM_SHIELD_PROXY.value: initialize_guardrail, +} + + +guardrail_class_registry: Final = { # mutable-ok: module-level registry, built once and never mutated + 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..b5c9f0f8b69 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/example_config.yaml @@ -0,0 +1,57 @@ +# Example LiteLLM Proxy configuration for LLM Shield Proxy +# LLM Shield Proxy is a self-hosted PII gateway: https://github.com/ninadphalak/LLM-Shield-Proxy +# +# Unlike a masking guardrail, LLM Shield Proxy's substitution is reversible. Personal data is +# replaced with placeholders before the request goes to the provider, and the original +# values are put back into the model's reply, so the end user still sees real data while +# the provider never received it. + +model_list: + - model_name: gpt-4o + litellm_params: + model: openai/gpt-4o + api_key: os.environ/OPENAI_API_KEY + +guardrails: + # Both modes belong on ONE entry. pre_call redacts the outbound request and post_call + # restores the reply; listing only pre_call would send placeholders back to the user. + - guardrail_name: "llm_shield_proxy" + litellm_params: + guardrail: llm_shield_proxy + mode: ["pre_call", "post_call"] + default_on: true + # Your own LLM Shield Proxy deployment. Defaults to http://localhost:8000, and also reads + # LLM_SHIELD_PROXY_API_BASE from the environment. + api_base: "http://localhost:8000" + # A virtual key configured on that deployment. Also reads LLM_SHIELD_PROXY_API_KEY. + api_key: os.environ/LLM_SHIELD_PROXY_API_KEY + +# Usage: +# +# 1. Run LLM Shield Proxy somewhere the proxy can reach: +# pip install llm-shield-proxy +# llm-shield-proxy --port 8000 +# +# 2. Point this config at it and start the proxy: +# export LLM_SHIELD_PROXY_API_KEY="your-virtual-key" +# litellm --config example_config.yaml +# +# 3. Send a request containing personal data: +# curl http://localhost:4000/v1/chat/completions \ +# -H "Authorization: Bearer sk-1234" \ +# -H "Content-Type: application/json" \ +# -d '{"model":"gpt-4o","messages":[{"role":"user","content":"Email jane.doe@example.com the invoice"}]}' +# +# The provider receives a stand-in value in place of the address. The reply you get +# back carries the real address again. +# +# Notes: +# +# - Requests are refused if LLM Shield Proxy is unreachable or returns an error, rather than +# being forwarded. Sending them on would hand the provider exactly the data this +# guardrail exists to withhold. +# - Restoring a value requires the request and the reply to share a session. LiteLLM's +# session id is used when present; otherwise one is generated per request. +# - Streaming replies are restored as chunks arrive. A placeholder split across two +# chunks is held back until it is complete, so partial values are never emitted. +# - Only text is redacted; images and audio pass through untouched. diff --git a/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..3cab7b20cef --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -0,0 +1,660 @@ +# +-------------------------------------------------------------+ +# +# Use LLM Shield Proxy for reversible PII redaction +# https://github.com/ninadphalak/LLM-Shield-Proxy +# +# +-------------------------------------------------------------+ + +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, + TypeAlias, +) + +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 + +if TYPE_CHECKING: + from litellm.caching.caching import DualCache + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + +GUARDRAIL_NAME: Final = "llm_shield_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" + +# The session id ties a redact call to the rehydrate calls that undo it. It is +# stored on the request dict rather than on the guardrail instance: the proxy +# registers one instance process-wide, so instance attributes would be shared +# across concurrent requests. +_SESSION_METADATA_KEY: Final = "llm_shield_session_id" + +# Roles whose text the application author wrote and the caller never sees. Their +# PII is still redacted outbound, but it is not restorable from the reply. +_PRIVILEGED_ROLES: Final = frozenset({"system", "developer"}) + +# Vault ids are minted here and never derived from anything the caller sends. The +# vault holds the plaintext behind every placeholder, so an id a caller could +# supply or guess would let one user rehydrate another user's values by getting a +# placeholder echoed back. The per-process prefix means a caller cannot even name +# a vault this process uses. +_VAULT_PREFIX: Final = f"litellm-{uuid.uuid4().hex}" + +_DEFAULT_TIMEOUT_SECONDS: Final = 10.0 + +# The proxy's own request dict. Mutable by design: a pre-call guardrail rewrites +# the caller's payload in place, which is the entire point of the hook. +# mutable-ok: the shape is fixed by CustomLogger's hook signatures. +MutableRequest: TypeAlias = dict + +# A JSON body on its way to httpx, which requires a real dict rather than a view. +# mutable-ok: handed straight to the HTTP client. +JsonBody: TypeAlias = dict + +# One redactable span: the text as it stands, and the write that puts the +# replacement back where it came from. +# How far a tool_result chain is followed. Real payloads nest one or two deep; the +# bound is what stops a crafted one from becoming an unbounded walk. +_MAX_CONTENT_DEPTH: Final = 8 + +_Slot: TypeAlias = tuple[str, Callable[[str], None]] # mutable-ok: Callable's param list. + +# The accumulator the collectors below append into. It never escapes +# _locate_request_texts, which freezes it into a tuple before returning. +_SlotSink: TypeAlias = list[_Slot] # mutable-ok: accumulator passed between collectors. + +# Sliding windows keyed by streaming choice index, threaded through one stream. +_CarryWindows: TypeAlias = dict # mutable-ok: per-choice windows advanced in place. + +# A caller-owned list whose entries are rewritten in place, such as a Completions +# `prompt` sent as an array of strings. +MutableSeq: TypeAlias = list # mutable-ok: the request payload's own list. + + +def _collect(container: MutableRequest, key: str, slots: _SlotSink) -> None: + """Records the string at `key`, along with the write that replaces it.""" + value: Final = container.get(key) + if isinstance(value, str) and value: + slots.append((value, lambda new, c=container, k=key: c.__setitem__(k, new))) + + +def _collect_entry(entries: MutableSeq, index: int, slots: _SlotSink) -> None: + """Records a string held directly in a list, rather than under a key.""" + value: Final = entries[index] + if isinstance(value, str) and value: + slots.append((value, lambda new, e=entries, i=index: e.__setitem__(i, new))) + + +def _collect_prompt(data: MutableRequest, slots: _SlotSink) -> None: + """The Completions API sends its text in `prompt`, and its tail in `suffix`.""" + _collect(data, "suffix", slots) + prompt: Final = data.get("prompt") + if isinstance(prompt, str): + _collect(data, "prompt", slots) + return + if isinstance(prompt, dict): + # A Responses API PromptObject. `variables` are substituted into the stored + # prompt on the provider side, so they are caller text. `id` and `version` + # identify which prompt to use and must arrive unchanged. + variables: Final = prompt.get("variables") + if isinstance(variables, dict): + for name in tuple(variables): + _collect(variables, name, slots) + return + if not isinstance(prompt, list): + return + for index in range(len(prompt)): + _collect_entry(prompt, 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. + """ + # Walked in document order: the shield maps its replies back by position, so the + # order spans are collected in is part of the contract. + pending: Final[list] = [(container, 0)] # mutable-ok: local queue, never escapes. + cursor = 0 # rebind-ok: advances through the queue. + while cursor < len(pending): + node, depth = pending[cursor] + cursor += 1 + content = node.get("content") + if isinstance(content, str): + _collect(node, "content", slots) + continue + if depth >= _MAX_CONTENT_DEPTH: + continue + for part in content if isinstance(content, list) else (): + if not isinstance(part, dict): + continue + # Image and audio parts have no text and fall through untouched. + _collect(part, "text", slots) + if "content" in part: + pending.append((part, depth + 1)) + + +def _collect_participant_name(message: MutableRequest, slots: _SlotSink) -> None: + """Redacts `name` where it identifies a person, never where it names a function. + + On a user or assistant turn `name` is the participant, which is personal data. + On a tool or function turn the same field carries the function's name and has + to reach the provider unchanged, or the call no longer routes. + """ + if message.get("role") in ("tool", "function"): + return + _collect(message, "name", slots) + + +def _collect_tool_arguments(message: MutableRequest, slots: _SlotSink) -> None: + """Tool arguments carry the values a user asked the model to act on.""" + for tool_call in message.get("tool_calls") or (): + function = tool_call.get("function") if isinstance(tool_call, dict) else None # rebind-ok: loop variable. + if isinstance(function, dict): + _collect(function, "arguments", slots) + legacy: Final = message.get("function_call") + if isinstance(legacy, dict): + _collect(legacy, "arguments", slots) + + +def _collect_system(data: MutableRequest, slots: _SlotSink) -> None: + """Anthropic's /v1/messages carries its system prompt at the top level.""" + system: Final = data.get("system") + if isinstance(system, str): + _collect(data, "system", slots) + return + for part in system if isinstance(system, list) else (): + if isinstance(part, dict): + _collect(part, "text", slots) + + +def _collect_responses_fields(data: MutableRequest, slots: _SlotSink, privileged: _SlotSink) -> None: + """The Responses API sends text outside `messages`, in `instructions` and `input`. + + `instructions` is written by the application, not by the caller, so it is + collected into the privileged sink; `input` is the caller's own text. + """ + _collect(data, "instructions", privileged) + request_input: Final = data.get("input") + if isinstance(request_input, str): + _collect(data, "input", slots) + return + if not isinstance(request_input, list): + return + for index, item in enumerate(request_input): + if isinstance(item, str): + # The embeddings and moderations shape: `input` as an array of strings. + _collect_entry(request_input, index, slots) + continue + if not isinstance(item, dict): + continue + _collect_content(item, slots) + # A function_call item holds `arguments`; a function_call_output holds `output`. + _collect(item, "arguments", slots) + _collect(item, "output", slots) + + +def _choice_index(choice: object) -> int: + """Streaming choices are matched across chunks by their index.""" + index: Final = getattr(choice, "index", 0) + return index if isinstance(index, int) else 0 + + +class 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. + """ + + # Our redaction and restoration run in the native lifecycle hooks below. Without + # this the proxy would route every event through the unified apply_guardrail path + # and the streaming hook would never fire. + use_native_lifecycle_hooks: ClassVar[bool] = True + + def __init__( + self, + guardrail_name: str = GUARDRAIL_NAME, + api_base: str | None = None, + api_key: str | None = None, + **kwargs: Any, # 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] # mutable-ok: parent's signature. + + # --- transport --------------------------------------------------------------- + + def _headers(self, session_id: str) -> JsonBody: + headers: Final[JsonBody] = { # mutable-ok: httpx requires a real dict. + "Content-Type": "application/json", + "X-Session-ID": session_id, + } + if self.api_key: + headers["Authorization"] = f"Bearer {self.api_key}" + return headers + + async def _call_shield(self, path: str, session_id: str, payload: 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)} # mutable-ok: JSON body for httpx. + body: Final = await self._call_shield(_REDACT_PATH, session_id, payload) + return self._same_length_or_raise(body.get("texts"), texts, "redact") + + async def _rehydrate(self, texts: Sequence[str], session_id: str) -> Sequence[str]: + payload: Final[JsonBody] = {"texts": list(texts)} # mutable-ok: JSON body for httpx. + body: Final = await self._call_shield(_REHYDRATE_PATH, session_id, payload) + return self._same_length_or_raise(body.get("texts"), texts, "rehydrate") + + def _same_length_or_raise(self, returned: object, sent: Sequence[str], operation: str) -> Sequence[str]: + """Guards the positional mapping the callers rely on to write results back.""" + if not isinstance(returned, list) or len(returned) != len(sent): + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=f"LLM Shield Proxy {operation} returned an unexpected payload; blocking the request.", + ) + return tuple(returned) + + # --- session ------------------------------------------------------------------ + + @staticmethod + def _mint_session_id(data: MutableRequest) -> str: + """Mints a vault id for this request, overwriting anything already there. + + Redaction and restoration both happen inside one request/response pair, so + a fresh id per request is all that is needed, and it is what keeps one + caller from reaching another caller's vault. + """ + session_id: Final = f"{_VAULT_PREFIX}-{uuid.uuid4().hex}" + metadata: Final = data.setdefault("metadata", {}) # mutable-ok: per-request store. + if isinstance(metadata, dict): + metadata[_SESSION_METADATA_KEY] = session_id + return session_id + + @staticmethod + def _session_id(data: MutableRequest) -> str: + """Reads back the vault id minted while redacting this request. + + Falls back to an unused id rather than to anything the caller supplied: a + reply that cannot be restored is a visible placeholder, while trusting a + caller-supplied id would hand them someone else's plaintext. + """ + metadata: Final = data.get("metadata") + existing: Final = metadata.get(_SESSION_METADATA_KEY) if isinstance(metadata, dict) else None + if isinstance(existing, str) and existing.startswith(_VAULT_PREFIX): + return existing + return f"{_VAULT_PREFIX}-{uuid.uuid4().hex}" + + # --- request traversal -------------------------------------------------------- + + @staticmethod + def _locate_request_texts( + data: MutableRequest, + ) -> 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` -- go into a vault nothing is + ever restored against, so a caller who gets the model to echo one of their + placeholders back receives the placeholder, not the value behind it. + """ + slots: Final[_SlotSink] = [] # mutable-ok: accumulator, frozen on return. + privileged: Final[_SlotSink] = [] # mutable-ok: accumulator, frozen on return. + for message in data.get("messages") or (): + if isinstance(message, dict): + sink = privileged if message.get("role") in _PRIVILEGED_ROLES else slots + _collect_content(message, sink) + _collect_participant_name(message, sink) + _collect_tool_arguments(message, sink) + _collect_responses_fields(data, slots, privileged) + _collect_prompt(data, slots) + _collect_system(data, privileged) + return tuple(slots), tuple(privileged) + + # --- hooks -------------------------------------------------------------------- + + @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 + + slots, privileged = self._locate_request_texts(data) + if not slots and not privileged: + return data + + session_id: Final = self._mint_session_id(data) + if privileged: + # A vault of its own, whose id is deliberately never stored: the + # response is restored against `session_id` alone, so nothing the + # model emits can turn one of these placeholders back into plaintext. + await self._redact_into(privileged, f"{_VAULT_PREFIX}-{uuid.uuid4().hex}") + if slots: + await self._redact_into(slots, session_id) + return data + + async def _redact_into(self, slots: Sequence[_Slot], session_id: str) -> None: + """Redacts every span in `slots` under one vault and writes the result back.""" + redacted: Final = await self._redact(tuple(text for text, _ in slots), session_id) + for (_, write), replacement in zip(slots, redacted): + write(replacement) + + @log_guardrail_information + async def async_post_call_success_hook( + self, + data: MutableRequest, + user_api_key_dict: UserAPIKeyAuth, + response: Any, + ) -> Any: + """Restores the original values in a non-streaming response.""" + if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.post_call) is not True: + return response + + if self._is_anthropic_message_response(response): + return await self._restore_anthropic_response(response, data) + + text_blocks: Final = self._responses_api_text_blocks(response) + if text_blocks: + return await self._restore_responses_api_response(response, text_blocks, data) + + choices: Final = getattr(response, "choices", None) + if not choices: + return response + + pending: Final = tuple( + (choice.message, choice.message.content) + for choice in choices + if getattr(choice, "message", None) is not None + and isinstance(getattr(choice.message, "content", None), str) + and choice.message.content + ) + if not pending: + return response + + restored: Final = await self._rehydrate(tuple(text for _, text in pending), self._session_id(data)) + for (message, _), replacement in zip(pending, restored): + message.content = replacement + return response + + @staticmethod + def _is_anthropic_message_response(response: object) -> bool: + """Anthropic's native /v1/messages reply arrives as a plain dict.""" + return ( + isinstance(response, dict) + and response.get("type") == "message" + and isinstance(response.get("content"), list) + ) + + async def _restore_anthropic_response(self, response: MutableRequest, data: MutableRequest) -> MutableRequest: + """Restores text blocks in an Anthropic native message reply. + + This shape has no `choices`, so without its own branch the reply would go + back to the caller still carrying placeholders. + """ + blocks: Final = tuple( + block + for block in response["content"] + if isinstance(block, dict) and block.get("type") == "text" and isinstance(block.get("text"), str) + ) + if not blocks: + return response + + restored: Final = await self._rehydrate(tuple(block["text"] for block in blocks), self._session_id(data)) + for block, replacement in zip(blocks, restored): + block["text"] = replacement + return response + + @staticmethod + def _responses_api_text_blocks(response: object) -> Sequence[object]: + """Text blocks in a Responses API reply. + + That shape carries `output` items rather than `choices`, so it needs its own + walk; without one the reply goes back to the caller still holding + placeholders even though the request was redacted correctly. Blocks come + through as dicts or as objects depending on how far the reply has been + deserialised, so both are handled. + """ + blocks: Final[list[object]] = [] # mutable-ok: accumulator, frozen on return. + for item in getattr(response, "output", None) or (): + for block in getattr(item, "content", None) or (): + if isinstance(block, dict): + if isinstance(block.get("text"), str) and block["text"]: + blocks.append(block) + elif isinstance(getattr(block, "text", None), str) and block.text: + blocks.append(block) + return tuple(blocks) + + @staticmethod + def _block_text(block: object) -> str: + return block["text"] if isinstance(block, dict) else block.text + + async def _restore_responses_api_response( + self, response: Any, blocks: Sequence[object], data: MutableRequest + ) -> Any: + """Puts the original values back into a Responses API reply.""" + restored: Final = await self._rehydrate( + tuple(self._block_text(block) for block in blocks), self._session_id(data) + ) + for block, replacement in zip(blocks, restored): + if isinstance(block, dict): + block["text"] = replacement + else: + block.text = replacement + return response + + async def async_post_call_streaming_iterator_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + response: Any, + request_data: MutableRequest, + ) -> AsyncGenerator[Any, None]: + """Restores original values incrementally, without buffering the stream. + + Each choice is its own token stream, so the sliding window is tracked per + choice index. One shared window would splice the characters held back for + one choice onto the next. The windows are locals of this generator, so they + are scoped to a single stream and cannot leak between concurrent requests. + """ + 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) + carries: Final[dict] = {} # mutable-ok: per-choice windows, local to this stream. + last_chunk = None # rebind-ok: tracks the most recent chunk for the final flush. + + async for chunk in response: + last_chunk = chunk + for choice in getattr(chunk, "choices", None) or (): + await self._restore_choice(choice, carries, session_id) + yield chunk + + # A stream that ended without a finish_reason can still leave text held back. + if 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: Any, carries: _CarryWindows, session_id: str) -> None: + """Restores one choice's delta, advancing that choice's own window.""" + delta: Final = getattr(choice, "delta", None) + if delta is None: + return + index: Final = _choice_index(choice) + carry: Final = carries.get(index, "") + text: Final = getattr(delta, "content", None) + is_final: Final = bool(getattr(choice, "finish_reason", None)) + + if not isinstance(text, str) or not text: + # Nothing to restore here, but a final chunk still has to flush the window. + if is_final and carry: + flushed, flushed_carry = await self._stream_step("", carry, True, session_id) + carries[index] = flushed_carry # rebind-ok: this choice's window advances. + if flushed: + delta.content = flushed + return + + emitted, remaining = await self._stream_step(text, carry, is_final, session_id) + carries[index] = remaining # rebind-ok: this choice's window advances. + delta.content = emitted + + async def _flush_trailing( + self, last_chunk: Any, carries: _CarryWindows, session_id: str + ) -> AsyncGenerator[Any, None]: + """Empties every window still holding text, one chunk per choice. + + Driven by the windows rather than by the last chunk's choices. A choice that + finished earlier is not present in the terminal chunk, and flushing only what + that chunk carries would drop its held text and truncate its answer. + """ + for index in sorted(carries): + carry = carries[index] + if not carry: + continue + text, remaining = await self._stream_step("", carry, True, session_id) + carries[index] = remaining # rebind-ok: this choice's window advances. + if not text: + continue + chunk = self._chunk_for_choice(last_chunk, index) + if chunk is None: + continue + chunk.choices[0].delta.content = text + yield chunk + + @staticmethod + def _chunk_for_choice(last_chunk: Any, index: int) -> Any: + """A single-choice copy of the last chunk, carrying only `index`. + + Emitting one choice per chunk keeps a flush from reading as content on a + choice it does not belong to. + """ + chunk: Final = last_chunk.model_copy(deep=True) + raw_choices: Final = getattr(chunk, "choices", None) + if not raw_choices: + return None + choices: Final[tuple] = tuple(raw_choices) + matching: Final = tuple(choice for choice in choices if _choice_index(choice) == index) + kept: Final = matching[0] if matching else choices[0] + if getattr(kept, "delta", None) is None: + return None + kept.index = index + # The terminal signal, if there was one, already went out with the real chunk. + kept.finish_reason = None + chunk.choices = [kept] # mutable-ok: the chunk model requires a list. + return chunk + + async def _stream_step(self, text: str, carry: str, final: bool, session_id: str) -> tuple[str, str]: + """Returns ``(text safe to emit now, window still being held)``.""" + body: Final = await self._call_shield( + _REHYDRATE_STREAM_PATH, + session_id, + # mutable-ok: JSON request body for httpx. + {"text": text, "carry": carry, "final": final}, # mutable-ok: JSON request body for httpx. + ) + 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 + + # --- unified API (powers the UI "Test guardrail" button) ----------------------- + + @log_guardrail_information + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: MutableRequest, + input_type: Literal["request", "response"], + logging_obj: Optional["LiteLLMLoggingObj"] = None, + ) -> GenericGuardrailAPIInputs: + texts: Final = inputs.get("texts") + if not texts: + return inputs + + replaced: Final = ( + await self._redact(tuple(texts), self._mint_session_id(request_data)) + if input_type == "request" + else await self._rehydrate(tuple(texts), self._session_id(request_data)) + ) + # Return a new mapping rather than rewriting the caller's, so this stays a + # pure transform of the inputs it was handed. + merged: Final[JsonBody] = {**inputs, "texts": list(replaced)} # mutable-ok: TypedDict. + return merged diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index ef28181eba5..da26479e8b3 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -137,6 +137,7 @@ class SupportedGuardrailIntegrations(Enum): COMPRESR = "compresr" STRAIKER = "straiker" ALICE = "alice" + LLM_SHIELD_PROXY = "llm_shield_proxy" class Role(Enum): 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 ae092bdde7d..21595556d28 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/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py new file mode 100644 index 00000000000..700702f3bb2 --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -0,0 +1,901 @@ +import json +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest +from httpx import Request, Response + +import litellm +from litellm.exceptions import GuardrailRaisedException +from litellm.proxy.guardrails.guardrail_hooks.llm_shield_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.utils import Choices, Delta, Message, ModelResponse, ModelResponseStream, StreamingChoices + + +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 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_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_deeply_nested_tool_results_are_bounded(self): + """Nesting is caller controlled, so the descent has to stop somewhere. + + The walk must terminate on a payload built to be pathological, rather than + following it as far as it goes. + """ + guardrail = _guardrail() + + captured: list = [] + + async def echo(url, headers, json, timeout): # noqa: ARG001 + captured.append(json["texts"]) + return _response({"texts": list(json["texts"])}) + + guardrail.async_handler.post = AsyncMock(side_effect=echo) # type: ignore[method-assign] + + deep: dict = {"type": "tool_result", "content": "past-the-bound@example.com"} + for _ in range(200): + deep = {"type": "tool_result", "content": [deep]} + data = {"messages": [{"role": "user", "content": [{"type": "text", "text": "shallow"}, deep]}]} + + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + sent = captured[0] + assert "shallow" in sent + assert "past-the-bound@example.com" not in sent, "the walk followed the chain past its bound" + assert len(sent) < 200 + + @pytest.mark.asyncio + async def test_responses_prompt_object_variables_are_redacted(self): + """A PromptObject's variables are substituted into the prompt provider side. + + The id and version pick which stored prompt to run and have to arrive + unchanged; the variables are caller text. + """ + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + + data = {"prompt": {"id": "pmpt_123", "version": "2", "variables": {"customer": "jane.doe@example.com"}}} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert data["prompt"]["variables"]["customer"] == "[EMAIL_1]" + assert data["prompt"]["id"] == "pmpt_123" + assert data["prompt"]["version"] == "2" + + @pytest.mark.asyncio + async def test_completions_suffix_is_redacted(self): + """LiteLLM forwards the legacy `suffix` to providers that support it.""" + 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" + + +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])]) + await guardrail.async_post_call_success_hook(data={"messages": []}, user_api_key_dict=None, response=response) + + assert block.text == "a@b.com" + + @pytest.mark.asyncio + async def test_anthropic_message_shape_is_restored(self): + """The /v1/messages reply is a plain dict with no choices. + + 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["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"), + ], + ) + 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) + + +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_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"] 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 7785a8e44ab..38751cb1d43 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts @@ -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; } @@ -318,4 +320,12 @@ export const GUARDRAIL_PRESETS: Record = { mode: "pre_call", defaultOn: false, }, + 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, + }, }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts index 1e486639840..293c447604f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts @@ -28,6 +28,7 @@ const EXPECTED_PARTNER_LOGO_FILES: Record = { repelloai: "repelloai.png", straiker: "straiker.svg", alice: "alice.svg", + llm_shield_proxy: "llm_shield_proxy.svg", }; describe("guardrail_garden_data logos", () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts index 931b3a111d8..74a13f2f611 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts @@ -474,6 +474,16 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ tags: ["Content Moderation", "Prompt Injection", "PII", "Policy"], providerKey: "Alice", }, + { + 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", + }, ]; export const ALL_CARDS = [...LITELLM_CONTENT_FILTER_CARDS, ...PARTNER_GUARDRAIL_CARDS]; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx index f686ff5644a..1bf7056bd5c 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx @@ -1,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 aporiaLogo from "../../../../../public/assets/logos/aporia.png"; import bedrockLogo from "../../../../../public/assets/logos/bedrock.svg"; import catoNetworksLogo from "../../../../../public/assets/logos/cato_networks.svg"; @@ -85,6 +86,7 @@ export const guardrail_provider_map: Record = { QostodianNexus: "qostodian_nexus", Repelloai: "repelloai", Alice: "alice", + "LLM Shield Proxy": "llm_shield_proxy", }; // Function to populate provider map from API response - updates the original map @@ -208,6 +210,7 @@ export const guardrailLogoMap = { "RepelloAI Argus": repelloAiLogo.src, Straiker: straikerLogo.src, Alice: aliceLogo.src, + "LLM Shield Proxy": llmShieldProxyLogo.src, } satisfies Record; export const getGuardrailLogo = (displayName: string): string | undefined =>