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;