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..b5c9f0f8b69 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/example_config.yaml @@ -0,0 +1,57 @@ +# Example LiteLLM Proxy configuration for LLM Shield Proxy +# LLM Shield Proxy is a self-hosted PII gateway: https://github.com/ninadphalak/LLM-Shield-Proxy +# +# Unlike a masking guardrail, LLM Shield Proxy's substitution is reversible. Personal data is +# replaced with placeholders before the request goes to the provider, and the original +# values are put back into the model's reply, so the end user still sees real data while +# the provider never received it. + +model_list: + - model_name: gpt-4o + litellm_params: + model: openai/gpt-4o + api_key: os.environ/OPENAI_API_KEY + +guardrails: + # Both modes belong on ONE entry. pre_call redacts the outbound request and post_call + # restores the reply; listing only pre_call would send placeholders back to the user. + - guardrail_name: "llm_shield_proxy" + litellm_params: + guardrail: llm_shield_proxy + mode: ["pre_call", "post_call"] + default_on: true + # Your own LLM Shield Proxy deployment. Defaults to http://localhost:8000, and also reads + # LLM_SHIELD_PROXY_API_BASE from the environment. + api_base: "http://localhost:8000" + # A virtual key configured on that deployment. Also reads LLM_SHIELD_PROXY_API_KEY. + api_key: os.environ/LLM_SHIELD_PROXY_API_KEY + +# Usage: +# +# 1. Run LLM Shield Proxy somewhere the proxy can reach: +# pip install llm-shield-proxy +# llm-shield-proxy --port 8000 +# +# 2. Point this config at it and start the proxy: +# export LLM_SHIELD_PROXY_API_KEY="your-virtual-key" +# litellm --config example_config.yaml +# +# 3. Send a request containing personal data: +# curl http://localhost:4000/v1/chat/completions \ +# -H "Authorization: Bearer sk-1234" \ +# -H "Content-Type: application/json" \ +# -d '{"model":"gpt-4o","messages":[{"role":"user","content":"Email jane.doe@example.com the invoice"}]}' +# +# The provider receives a stand-in value in place of the address. The reply you get +# back carries the real address again. +# +# Notes: +# +# - Requests are refused if LLM Shield Proxy is unreachable or returns an error, rather than +# being forwarded. Sending them on would hand the provider exactly the data this +# guardrail exists to withhold. +# - Restoring a value requires the request and the reply to share a session. LiteLLM's +# session id is used when present; otherwise one is generated per request. +# - Streaming replies are restored as chunks arrive. A placeholder split across two +# chunks is held back until it is complete, so partial values are never emitted. +# - Only text is redacted; images and audio pass through untouched. diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py new file mode 100644 index 00000000000..49bcb768600 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -0,0 +1,1549 @@ +# +-------------------------------------------------------------+ +# +# Use LLM Shield Proxy for reversible PII redaction +# https://github.com/ninadphalak/LLM-Shield-Proxy +# +# +-------------------------------------------------------------+ + +import copy +import functools +import itertools +import json +import os +import re +import uuid +from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequence +from enum import Enum +from types import MappingProxyType +from typing import ( + TYPE_CHECKING, + Any, # noqa: TID251 # **kwargs forwards verbatim to CustomGuardrail.__init__ + ClassVar, + Final, + Literal, + Optional, + TypeAlias, +) + +import httpx + +from litellm._logging import verbose_proxy_logger +from litellm.exceptions import GuardrailRaisedException +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + log_guardrail_information, +) +from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, + httpxSpecialProvider, +) +from litellm.proxy._types import UserAPIKeyAuth +from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.utils import GenericGuardrailAPIInputs + +if TYPE_CHECKING: + from litellm.caching.caching import DualCache + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + +GUARDRAIL_NAME: Final = "llm_shield_proxy" + +_DEFAULT_API_BASE: Final = "http://localhost:8000" +_REDACT_PATH: Final = "/v1/guard/redact" +_REHYDRATE_PATH: Final = "/v1/guard/rehydrate" +_REHYDRATE_STREAM_PATH: Final = "/v1/guard/rehydrate/stream" + +# The session id ties a redact call to the rehydrate calls that undo it. It is +# stored on the request dict rather than on the guardrail instance: the proxy +# registers one instance process-wide, so instance attributes would be shared +# across concurrent requests. +_SESSION_METADATA_KEY: Final = "llm_shield_session_id" + +# Roles whose text the application author wrote and the caller never sees. Their +# PII is still redacted outbound, but it is not restorable from the reply. +_PRIVILEGED_ROLES: Final = frozenset({"system", "developer"}) + +# Vault ids are minted here and never derived from anything the caller sends. The +# vault holds the plaintext behind every placeholder, so an id a caller could +# supply or guess would let one user rehydrate another user's values by getting a +# placeholder echoed back. The per-process prefix means a caller cannot even name +# a vault this process uses. +_VAULT_PREFIX: Final = f"litellm-{uuid.uuid4().hex}" + +_DEFAULT_TIMEOUT_SECONDS: Final = 10.0 + +# The proxy's own request dict. Mutable by design: a pre-call guardrail rewrites +# the caller's payload in place, which is the entire point of the hook. +MutableRequest: TypeAlias = dict[str, object] + +# A JSON body on its way to httpx, which requires a real dict rather than a view. +JsonBody: TypeAlias = dict[str, object] + +# One redactable span: the text as it stands, and the write that puts the +# replacement back where it came from. +# How far a tool_result chain is followed. Real payloads nest one or two deep; the +# bound is what stops a crafted one from becoming an unbounded walk. A request that +# nests deeper is refused rather than forwarded, because text past the bound would +# otherwise reach the provider unredacted. +_MAX_CONTENT_DEPTH: Final = 8 + +# How far a JSON value -- a tool input, a parameter schema -- is followed on the request +# side. Legitimate JSON nests far deeper than content blocks do, so the bound is +# generous; past it the request is refused, for the same reason as above. +_MAX_JSON_DEPTH: Final = 64 + +_Slot: TypeAlias = tuple[str, Callable[[str], None]] + +# One incremental rehydration step for a stream the caller has already bound to its +# vault: (new text, carried window, final) -> (text safe to emit, window still held). +_StreamStep: TypeAlias = Callable[[str, str, bool], Awaitable[tuple[str, str]]] + +# A batch rehydration already bound to the request's vault. +_Rehydrate: TypeAlias = Callable[[Sequence[str]], Awaitable[Sequence[str]]] + +# Anthropic /v1/messages delta types that carry restorable text, and the field holding +# it. `thinking_delta` is left out on purpose: a thinking block is signed, and one +# rewritten here fails verification when the client sends it back on the next turn. +_ANTHROPIC_DELTA_FIELDS: Final = MappingProxyType({"text_delta": "text", "input_json_delta": "partial_json"}) + +# A blank line ends an SSE event. Frames are cut there, never inside an event, so a +# `data:` line split across two network chunks is parsed only once it is whole. +_SSE_EVENT_BOUNDARY: Final = re.compile(rb"(\r?\n\r?\n)") + +# What an SSE stream can open with: one of its fields, or a `:` comment. +_SSE_OPENINGS: Final = (b"event:", b"data:", b"id:", b"retry:", b":") + +# Responses API delta events whose `delta` is not text. Audio arrives base64-encoded; +# sending it through the shield would cost a round trip per chunk to restore nothing. +_RESPONSES_BINARY_DELTAS: Final = frozenset(("response.audio.delta",)) + +# Fields on a Responses API event that identify something rather than say something. +# Every other string field on a `.done` event is model text and is restored, so an event +# type added upstream is covered by default instead of leaking a placeholder. +_RESPONSES_STRUCTURAL_FIELDS: Final = frozenset( + ("type", "id", "item_id", "call_id", "name", "server_label", "status", "obfuscation") +) + +# Terminal Responses API events that repeat the whole reply under `response`. +_RESPONSES_TERMINAL_EVENTS: Final = frozenset(("response.completed", "response.incomplete")) + +# JSON Schema keywords whose value has to reach the model or a validator verbatim, so the +# schema walk leaves them alone: types, formats, patterns, references and +# required-property lists. Everything else is scanned. +_SCHEMA_STRUCTURAL_KEYWORDS: Final = frozenset( + ( + "type", + "format", + "pattern", + "required", + "dependentRequired", + "propertyOrdering", + "discriminator", + "contentEncoding", + "contentMediaType", + "$ref", + "$id", + "$schema", + "$anchor", + "$dynamicRef", + "$dynamicAnchor", + "$recursiveRef", + "$recursiveAnchor", + "$vocabulary", + ) +) + +# Keywords holding JSON values rather than schemas: every string in them is collected, +# whatever the keys around it are called. +_SCHEMA_VALUE_KEYWORDS: Final = frozenset(("examples", "default")) + +# Keywords holding the literal values the model must reproduce. These go to the CALLER's +# vault, not the privileged one: the model emits the stand-in in its tool arguments or +# structured output, and restoring the reply turns it back into the value the schema +# allows, so the call still routes. In the non-restorable vault it would come back as a +# stand-in no validator accepts. +_SCHEMA_LITERAL_KEYWORDS: Final = frozenset(("enum", "const")) + +# Keywords whose value maps names to subschemas. Their keys are property names, not +# keywords, so a property called `type` or `enum` is walked like any other subschema. +_SCHEMA_MAP_KEYWORDS: Final = frozenset( + ("properties", "patternProperties", "$defs", "definitions", "dependentSchemas", "dependencies") +) + +# The accumulator the collectors below append into. It never escapes +# _locate_request_texts, which freezes it into a tuple before returning. +_SlotSink: TypeAlias = list[_Slot] + +# Sliding windows keyed by (choice index, tool-call index | None), threaded through one +# stream. `None` is the content channel; an int is one tool call's accumulating +# `arguments`. Content and each tool call are separate token streams, so each needs its +# own window -- one shared window would splice one stream's held-back tail onto another. +_CarryKey: TypeAlias = tuple[int, int | None] +_CarryWindows: TypeAlias = dict[_CarryKey, str] + +# A caller-owned list whose entries are rewritten in place, such as a Completions +# `prompt` sent as an array of strings. +MutableSeq: TypeAlias = list[object] + +# A Responses API delta stream: (event family, item id, output index, part index). +_ResponsesStreamKey: TypeAlias = tuple[str, object, object, object] + + +def _as_object(value: object) -> MutableRequest | None: + """`value` as a JSON object, or None. + + `isinstance(value, dict)` alone leaves the keys and values unknown to the type + checker. A JSON object's keys are strings, so the type is stated once, here. + """ + return value if isinstance(value, dict) else None + + +def _as_array(value: object) -> MutableSeq | None: + """`value` as a JSON array, or None. See `_as_object`.""" + return value if isinstance(value, list) else None + + +def _is_container(value: object) -> bool: + """Whether `value` is a JSON object or array, without narrowing it to unknown types.""" + return isinstance(value, (dict, list)) + + +def _collect(container: MutableRequest, key: str, slots: _SlotSink) -> None: + """Records the string at `key`, along with the write that replaces it.""" + value: Final = container.get(key) + if isinstance(value, str) and value: + slots.append((value, 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))) + + +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: + # A Responses API PromptObject. `variables` are substituted into the stored + # prompt on the provider side, so they are caller text. `id` and `version` + # identify which prompt to use and must arrive unchanged. + variables: Final = _as_object(prompt_object.get("variables")) + if variables is not None: + for name in tuple(variables): + _collect(variables, name, slots) + return + entries: Final = _as_array(prompt) + if entries is None: + return + for index in range(len(entries)): + _collect_entry(entries, index, slots) + + +class _RequestTooDeep(Exception): + """A request nests text past a walk's bound. + + Skipping the rest would forward it unredacted while the guardrail reports as + enabled, so the pre-call hook refuses the request instead. + """ + + +def _collect_content(container: MutableRequest, slots: _SlotSink) -> None: + """Collects `content`, a string or a list of typed parts. + + An Anthropic tool_result nests its own content, so this has to descend. It walks + with an explicit stack and a depth bound rather than by recursion: the nesting is + caller controlled, and an unbounded descent is a JSON bomb. Content nested past the + bound raises `_RequestTooDeep` rather than being skipped. + """ + # Walked in document order: the shield maps its replies back by position, so the + # order spans are collected in is part of the contract. + 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 + # Image and audio parts have no text and fall through untouched. + _collect(part, "text", slots) + if part.get("type") == "tool_use": + # A replayed Anthropic tool call. Its `input` is a JSON object rather than + # a string, so a value can sit at any depth -- the reply side walks the + # same leaves when it restores one. Its own JSON bound applies, not the + # content one, and past it the request is refused. + _collect_json_leaves(part.get("input"), slots, strict=True) + if "content" in part: + pending.append((part, depth + 1)) + + +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. + """ + _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): + # The embeddings and moderations shape: `input` as an array of strings. + _collect_entry(entries, index, slots) + continue + item = _as_object(entry) + if item is None: + continue + _collect_content(item, slots) + # A function_call item holds `arguments`; a function_call_output holds `output`. + _collect(item, "arguments", slots) + _collect(item, "output", slots) + # A replayed reasoning item carries the model's summary of its own reasoning, + # which quotes whatever the conversation contained. + _collect_text_parts(item, "summary", slots) + + +def _collect_text_parts(container: MutableRequest, key: str, slots: _SlotSink) -> None: + """Collects the `text` of every part in the list held at `key`.""" + for entry in _as_array(container.get(key)) or (): + part = _as_object(entry) + if part is not None: + _collect(part, "text", 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 _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. + """ + if isinstance(holder, dict): + return holder.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: str) -> None: + """Writes one string field back into a dict or an object. Pairs with _read_field.""" + if isinstance(holder, dict): + holder[name] = value + else: + setattr(holder, name, value) + + +def _collect_json_leaves(node: object, slots: _SlotSink, *, 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 _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 _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) + + +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: + # Too little has arrived to tell -- `b"eve"` could still become `event:`. + self._pending = buffered + return () + if not self._is_sse: + self._pending = b"" + return self._emit(buffered) + boundaries: Final = tuple(_SSE_EVENT_BOUNDARY.finditer(buffered)) + if not boundaries: + self._pending = buffered + return () + cut: Final = boundaries[-1].end() + self._pending = buffered[cut:] + # With a capturing group, split alternates event, separator, ..., and ends in the + # empty remainder after the last separator. + parts: Final = _SSE_EVENT_BOUNDARY.split(buffered[:cut]) + restored: Final = tuple( + [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: + # The stream ended before it could be told apart from SSE: hand it back as is. + return self._emit(held) + tail: Final = await self._restore_event(held) if held.strip() else held + flushed: Final = await self._flush_all() + # The tail had no blank line after it; one is needed before another frame follows. + separator: Final = b"\n\n" if tail.strip() and flushed else b"" + 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. + """ + # A model's fields live in its `__dict__`; an empty dict has none either way. + attributes: Final[object] = getattr(event, "__dict__", None) + fields: Final = _as_object(event) or _as_object(attributes) + if fields is None: + return + for name, value in tuple(fields.items()): + if 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))) + + +class LLMShieldProxyGuardrail(CustomGuardrail): + """Redacts PII before it leaves the proxy and restores it in the response. + + Unlike a masking guardrail, the substitution is reversible. Outbound text is + replaced with placeholders held in a session vault inside the user's own LLM + Shield deployment; the model's reply is then restored so the end user sees the + original values while the provider never received them. + + Streaming is restored incrementally rather than by buffering the response. LLM + Shield holds back only the trailing characters that could still turn out to be + part of a placeholder, so tokens are forwarded as they arrive and a placeholder + split across two chunks is never emitted in fragments. + """ + + # Our redaction and restoration run in the native lifecycle hooks below. Without + # this the proxy would route every event through the unified apply_guardrail path + # and the streaming hook would never fire. + use_native_lifecycle_hooks: ClassVar[bool] = True + + def __init__( + self, + guardrail_name: str = GUARDRAIL_NAME, + api_base: str | None = None, + api_key: str | None = None, + **kwargs: Any, # noqa: LIT008 # kwargs-ok: forwarded verbatim to CustomGuardrail.__init__ + ) -> None: + self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback) + env_base: Final = os.environ.get("LLM_SHIELD_PROXY_API_BASE") + self.api_base: Final = (api_base or env_base or _DEFAULT_API_BASE).rstrip("/") + self.api_key: Final = api_key or os.environ.get("LLM_SHIELD_PROXY_API_KEY") + super().__init__(guardrail_name=guardrail_name, **kwargs) + + @classmethod + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: # mutable-ok: parent's signature. + return [GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call] + + # --- transport --------------------------------------------------------------- + + 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)) + # A non-string entry would be written into the request or reply as is. + if entries is None or len(entries) != len(sent) or len(texts) != len(entries): + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=f"LLM Shield Proxy {operation} returned an unexpected payload; blocking the request.", + ) + return texts + + # --- session ------------------------------------------------------------------ + + @staticmethod + def _mint_session_id(data: MutableRequest) -> str: + """Mints a vault id for this request, overwriting anything already there. + + Redaction and restoration both happen inside one request/response pair, so + a fresh id per request is all that is needed, and it is what keeps one + caller from reaching another caller's vault. + """ + session_id: Final = f"{_VAULT_PREFIX}-{uuid.uuid4().hex}" + # `litellm_metadata` is proxy-private; `metadata` is forwarded to the provider on + # /v1/responses. The session id is a capability against the vault's rehydrate + # endpoint, so handing it to the provider alongside the placeholders would let the + # provider read back exactly what this guardrail exists to withhold. + metadata: Final = data.setdefault("litellm_metadata", {}) + if isinstance(metadata, dict): + metadata[_SESSION_METADATA_KEY] = session_id + return session_id + + @staticmethod + def _session_id(data: MutableRequest) -> str: + """Reads back the vault id minted while redacting this request. + + Falls back to an unused id rather than to anything the caller supplied: a + reply that cannot be restored is a visible placeholder, while trusting a + caller-supplied id would hand them someone else's plaintext. + """ + # Read only from `litellm_metadata`, the same proxy-private store `_mint_session_id` + # writes to. A caller can populate `metadata`; they cannot populate this. + metadata: Final = data.get("litellm_metadata") + existing: Final = metadata.get(_SESSION_METADATA_KEY) if isinstance(metadata, dict) else None + if isinstance(existing, str) and existing.startswith(_VAULT_PREFIX): + return existing + return f"{_VAULT_PREFIX}-{uuid.uuid4().hex}" + + # --- request traversal -------------------------------------------------------- + + @staticmethod + def _locate_request_texts( + data: MutableRequest, + ) -> tuple[Sequence[_Slot], Sequence[_Slot]]: + """Finds every redactable span, split by whether the caller can see it. + + Anything missed here reaches the provider in the clear while the guardrail + still reports as enabled, so the walk covers every request shape that + carries text. + + The split exists because the response is restored against one vault only. + Server-authored spans -- system and developer turns, Anthropic's top-level + `system`, the Responses API `instructions`, 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. + """ + slots: Final[_SlotSink] = [] + privileged: Final[_SlotSink] = [] + for entry in _read_list(data, "messages"): + message = _as_object(entry) + if message is not None: + sink = privileged if message.get("role") in _PRIVILEGED_ROLES else slots + _collect_content(message, sink) + _collect_participant_name(message, sink) + _collect_tool_arguments(message, sink) + _collect_responses_fields(data, slots, privileged) + _collect_prompt(data, slots) + _collect_system(data, privileged) + _collect_tool_definitions(data, slots, privileged) + _collect_output_contracts(data, slots, privileged) + _collect_user_locations(data, privileged) + _collect_end_user_ids(data, privileged) + return tuple(slots), tuple(privileged) + + # --- hooks -------------------------------------------------------------------- + + @log_guardrail_information + async def async_pre_call_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + cache: "DualCache", + data: MutableRequest, + call_type: str, + ) -> MutableRequest | None: + """Replaces PII anywhere in the outbound request with vault placeholders.""" + if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_call) is not True: + return data + + 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: + # A vault of its own, whose id is deliberately never stored: the + # response is restored against `session_id` alone, so nothing the + # model emits can turn one of these placeholders back into plaintext. + await self._redact_into(privileged, f"{_VAULT_PREFIX}-{uuid.uuid4().hex}") + if slots: + await self._redact_into(slots, session_id) + return data + + async def _redact_into(self, slots: Sequence[_Slot], session_id: str) -> None: + """Redacts every span in `slots` under one vault and writes the result back.""" + redacted: Final = await self._redact(tuple(text for text, _ in slots), session_id) + for (_, write), replacement in zip(slots, redacted): + write(replacement) + + # KNOWN LIMIT: a tool call's `arguments` is a JSON *string*, and a restored value is + # spliced into it as raw text. If the original value contained a double quote, a + # backslash or a newline, the reassembled document is no longer valid JSON for a + # strict parser. The proxy's own tool-argument rehydration has the same property + # (`_rehydrate_json_response` in api/main.py), so this is a pre-existing limit of the + # product rather than one introduced here. Escaping is deliberately NOT applied as a + # fix: a fragment is an arbitrary slice of a JSON document, so the code cannot tell + # whether the position it writes is inside a string literal, and escaping + # unconditionally would corrupt the values that are not. + @log_guardrail_information + async def async_post_call_success_hook( + self, + data: MutableRequest, + user_api_key_dict: UserAPIKeyAuth, + response: Any, + ) -> Any: + """Restores the original values in a non-streaming response.""" + if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.post_call) is not True: + return response + + if self._is_anthropic_message_response(response): + return await self._restore_anthropic_response(response, data) + + 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 + + # One batch for every restorable span in the reply, collected in document order: + # the shield maps its answers back by position. A second round trip is not an + # option here -- /v1/guard/rehydrate caps a batch at 256 texts and 1,000,000 + # characters, and `_same_length_or_raise` is what guarantees the positional + # mapping -- so a reply carrying more spans than that fails closed, which is this + # guardrail's posture everywhere else. + pending: Final[_SlotSink] = [] + for choice in choices: + message = getattr(choice, "message", None) + if message is None: + continue + content = getattr(message, "content", None) + if isinstance(content, str) and content: + pending.append((content, functools.partial(setattr, message, "content"))) + # A tool call's `arguments` is model-generated text and the request path + # redacts it, so leaving it unrestored hands the application a placeholder to + # invoke a tool with. These are Pydantic objects on this path, not dicts. + 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.""" + return ( + isinstance(response, dict) + and response.get("type") == "message" + and isinstance(response.get("content"), list) + ) + + async def _restore_anthropic_response(self, response: MutableRequest, data: MutableRequest) -> MutableRequest: + """Restores text blocks 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(chunk): + yield event + continue + last_chunk = chunk + for choice in getattr(chunk, "choices", None) or (): + await self._restore_choice(choice, carries, session_id) + yield chunk + + # A stream that ended early can still leave text held back, in any shape. + for frames in await sse.finish(): + yield frames + for event in await events.finish(): + yield event + if last_chunk is not None and any(carries.values()): + async for trailing in self._flush_trailing(last_chunk, carries, session_id): + yield trailing + + async def _restore_choice(self, choice: Any, 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) + if delta is None: + return + index: Final = _choice_index(choice) + is_final: Final = bool(getattr(choice, "finish_reason", None)) + + await self._restore_content_window(delta, (index, None), carries, session_id, is_final) + + for tool_call in getattr(delta, "tool_calls", None) or (): + await self._restore_tool_call_window(tool_call, index, carries, session_id) + + if is_final: + # A client parses a tool call's arguments when it sees the finish_reason, so + # every window this choice still holds has to land in *this* chunk. Flushing + # after it produces argument JSON the client has already stopped waiting for. + await self._flush_finished_choice(delta, index, carries, session_id) + + async def _restore_content_window( + self, + delta: Any, + key: _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: + # Nothing to restore here, but a final chunk still has to flush the window. + 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 + if tool_index is None: + chunk.choices[0].delta.content = text + else: + # The copy carried this chunk's own content and tool calls, both already + # delivered. Replace rather than append, and drop the content, or the + # client sees them twice. + chunk.choices[0].delta.content = None + chunk.choices[0].delta.tool_calls = _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: + return None + kept.index = index + # The terminal signal, if there was one, already went out with the real chunk. + kept.finish_reason = None + chunk.choices = [kept] + 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 + + # --- unified API (powers the UI "Test guardrail" button) ----------------------- + + @log_guardrail_information + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: MutableRequest, + input_type: Literal["request", "response"], + logging_obj: Optional["LiteLLMLoggingObj"] = None, + ) -> GenericGuardrailAPIInputs: + """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 + + # Copied rather than mutated: the caller's tool calls are theirs to own, and this + # method's contract is to hand back a new mapping. + restored_calls: Final[list[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) + # Return a new mapping rather than rewriting the caller's, so this stays a + # pure transform of the inputs it was handed. + merged: Final[JsonBody] = {**inputs} + if text_list: + merged["texts"] = restored_values[: len(text_list)] + if restored_calls: + merged["tool_calls"] = restored_calls + return merged 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..b9314590f38 --- /dev/null +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -0,0 +1,1589 @@ +import json +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest +from httpx import Request, Response + +import litellm +from litellm.exceptions import GuardrailRaisedException +from litellm.proxy.guardrails.guardrail_hooks.llm_shield_proxy.llm_shield_proxy import ( + GUARDRAIL_NAME, + LLMShieldProxyGuardrail, +) +from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2 +from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.llms.openai import ( + FunctionCallArgumentsDeltaEvent, + OutputTextDeltaEvent, + OutputTextDoneEvent, + ResponsesAPIStreamEvents, +) +from litellm.types.utils import Choices, Delta, Message, ModelResponse, ModelResponseStream, StreamingChoices + + +def _guardrail(**overrides: object) -> LLMShieldProxyGuardrail: + params: dict[str, object] = { + "api_key": "test-key", + "api_base": "http://shield.test", + "guardrail_name": GUARDRAIL_NAME, + "event_hook": "pre_call", + "default_on": True, + } + params.update(overrides) + return LLMShieldProxyGuardrail(**params) + + +def _response(payload: dict, status_code: int = 200) -> Response: + return Response( + status_code=status_code, + json=payload, + request=Request("POST", "http://shield.test/v1/guard/redact"), + ) + + +def _mock_post(guardrail: LLMShieldProxyGuardrail, *payloads: dict) -> AsyncMock: + """Queues one shield response per expected call.""" + mock = AsyncMock(side_effect=[_response(p) for p in payloads]) + guardrail.async_handler.post = mock # type: ignore[method-assign] + return mock + + +def _chunk(content: str | None, finish_reason: str | None = None) -> ModelResponseStream: + return ModelResponseStream( + choices=[StreamingChoices(index=0, delta=Delta(content=content), finish_reason=finish_reason)] + ) + + +async def _drain(generator) -> list: + return [chunk async for chunk in generator] + + +def _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_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_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] + await guardrail.async_post_call_success_hook(data=data, user_api_key_dict=None, response=reply) + + assert json.loads(call.function.arguments) == {"to": "ops@example.com"} + + def test_schema_nesting_past_the_bound_is_refused(self): + schema: dict = {"type": "object", "description": "past-the-bound@example.com"} + for _ in range(100): + 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])]) + await guardrail.async_post_call_success_hook(data={"messages": []}, user_api_key_dict=None, response=response) + + assert block.text == "a@b.com" + + @pytest.mark.asyncio + async def test_anthropic_message_shape_is_restored(self): + """The /v1/messages reply is a plain dict with no choices. + + Measured against a live provider: without its own branch the reply went + back to the caller still carrying the placeholder, even though the + request had been redacted correctly. + """ + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["a@b.com"]}) + + response = { + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "[EMAIL_1]"}], + } + result = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=response + ) + + assert result["content"][0]["text"] == "a@b.com" + + @pytest.mark.asyncio + async def test_anthropic_non_text_blocks_are_left_alone(self): + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["a@b.com"]}) + + response = { + "type": "message", + "role": "assistant", + "content": [ + {"type": "text", "text": "[EMAIL_1]"}, + {"type": "tool_use", "id": "t1", "name": "lookup", "input": {}}, + ], + } + result = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=response + ) + + assert result["content"][0]["text"] == "a@b.com" + assert result["content"][1] == {"type": "tool_use", "id": "t1", "name": "lookup", "input": {}} + + +class TestVaultIsolation: + """The vault id must never be something a caller can choose. + + The vault holds the plaintext behind every placeholder. If a caller could name + the vault, they could send a placeholder, have the model echo it back, and get + another caller's value restored into their own reply. + """ + + @pytest.mark.asyncio + async def test_caller_supplied_session_id_is_not_used(self): + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + + data = { + "messages": [{"role": "user", "content": "a@b.com"}], + "metadata": {"llm_shield_session_id": "victim-session"}, + "litellm_session_id": "victim-session", + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + used = mock.call_args_list[0].kwargs["headers"]["X-Session-ID"] + assert used != "victim-session" + assert data["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) + + +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]), + ) + + await _restore_stream(guardrail, [completed]) + + assert block["text"] == "Mail a@example.com" + assert call.arguments == '{"to": "a@example.com"}' + + @pytest.mark.asyncio + async def test_streams_on_different_parts_do_not_share_a_window(self): + guardrail, _ = _shielded(self.VAULT) + events = [ + self._text_delta("one [EMA", 1), + self._text_delta("two", 2, content_index=1), + self._text_delta("IL_1]", 3), + ] + + out = await _restore_stream(guardrail, events) + + by_part: dict[int, str] = {} + for event in out: + by_part[event.content_index] = by_part.get(event.content_index, "") + event.delta + assert by_part == {0: "one a@example.com", 1: "two"} + + @pytest.mark.asyncio + async def test_reasoning_summary_part_done_is_restored(self): + """The summary part repeats the whole summary text after its deltas.""" + guardrail, _ = _shielded(self.VAULT) + part = SimpleNamespace(type="summary_text", text="asked about [EMAIL_1]") + event = SimpleNamespace( + type="response.reasoning_summary_part.done", item_id="rs_1", output_index=0, summary_index=0, part=part + ) + + await _restore_stream(guardrail, [event]) + + assert part.text == "asked about a@example.com" + + @pytest.mark.asyncio + async def test_mcp_call_arguments_are_restored(self): + """A stream family outside the chat-era set: matched by shape, not by name.""" + guardrail, _ = _shielded(self.VAULT) + deltas = [ + {"type": "response.mcp_call_arguments.delta", "item_id": "mcp_1", "output_index": 0, "delta": d} + for d in ('{"to": "[EMAI', 'L_1]"}') + ] + done = { + "type": "response.mcp_call_arguments.done", + "item_id": "mcp_1", + "output_index": 0, + "arguments": '{"to": "[EMAIL_1]"}', + } + + out = await _restore_stream(guardrail, [*deltas, done]) + + assert json.loads("".join(e["delta"] for e in out[:-1])) == {"to": "a@example.com"} + assert json.loads(out[-1]["arguments"]) == {"to": "a@example.com"} + assert out[-1]["item_id"] == "mcp_1", "identifiers are not text and stay as sent" + + @pytest.mark.asyncio + async def test_audio_deltas_are_not_sent_to_the_shield(self): + """Audio arrives base64-encoded; restoring it would cost a round trip for nothing.""" + guardrail, shield = _shielded(self.VAULT) + audio = {"type": "response.audio.delta", "item_id": "a_1", "output_index": 0, "delta": "UklGRiQAAABXQVZF"} + + out = await _restore_stream(guardrail, [audio]) + + assert out == [audio] + assert shield.urls == [] 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;