From d94670674434cc9f9ac50b974ea2b1e3bde9be32 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Mon, 5 Oct 2026 12:33:56 -0500 Subject: [PATCH] feat(guardrails): add llm shield pii redaction and rehydration guardrail (#42645) * feat(guardrails): add llm shield pii redaction and rehydration guardrail LLM Shield is a self-hosted PII gateway. This adds it as a guardrail so a proxy operator can redact personal data out of outbound requests and have the original values restored in the model's reply. The substitution is reversible, which is the difference from a masking guardrail. Outbound text is replaced with placeholders held in a session vault inside the operator's own LLM Shield deployment, and the reply is restored before it reaches the caller, so the end user still sees real values while the provider never received them. Streaming responses are restored incrementally. LLM Shield holds back only the trailing characters that could still turn out to be part of a placeholder, so tokens are forwarded as they arrive rather than the whole response being collected first. A placeholder split across two chunks is never emitted in fragments. The integration talks to LLM Shield over HTTP and adds no dependency. Notes for reviewers: - The guardrail sets use_native_lifecycle_hooks, since redaction and restoration need the native pre-call, post-call and streaming hooks rather than the unified path. - Per-request state lives on the request dict, never on the guardrail instance, because the proxy registers a single instance process-wide. The streaming carry-over is a local of the generator for the same reason. - Every failure blocks the request. A redaction guardrail that fails open would send the exact data it exists to protect to the provider. * feat(ui): list llm shield in the guardrail garden Adds the card, preset and logo so operators can pick LLM Shield from the guardrails page the same way as the other partner guardrails. * docs(guardrails): add llm shield example config Shows both modes on one entry. Listing only pre_call redacts the request and then hands the placeholders back to the end user, so the test asserts both hooks are enabled. * feat(ui): use the llm shield brand mark for the guardrail logo * fix(guardrails): restore llm shield values in anthropic replies The /v1/messages reply is a plain dict with a content block list and no choices, so it fell through the restore path and went back to the caller still carrying placeholders. The request was redacted correctly, which is what made this easy to miss. Found by running all three endpoints against a live provider; the mocked tests all passed because they only built the OpenAI shape. Adds tests for the message shape and for leaving non-text blocks alone. * docs(guardrails): correct the llm shield start command * fix(guardrails): redact every request shape and restore every reply shape Three gaps, all of which let an enabled guardrail hand data to the provider or hand placeholders to the caller. Requests only walked `messages`. The Responses API `input` and tool call `arguments` went out untouched. Measured against a live provider: a request sent through `/v1/responses` reached the model with the real address in it while the guardrail reported as enabled. Request traversal now covers chat content (string and multimodal), tool call arguments, and `input` as a bare string or a list of items. Fixing that exposed the matching gap on the way back: the Responses API reply carries `output` items rather than `choices`, so it returned to the caller still holding placeholders. It now gets its own walk, handling text blocks as dicts or objects. The dashboard preset seeded only pre_call, so a guardrail created from the UI would redact the request and return the placeholders to the user. Presets can now seed both modes; the form already normalised either shape. Adds tests for each request shape, for both Responses API reply forms, and replaces a test that had asserted the `input` bypass as correct behaviour. * fix(guardrails): narrow the stream delta before writing to it basedpyright could not prove the delta was non-None on the write path, and reportOptionalMemberAccess has a zero budget. The guard is also clearer than relying on the text check to imply it. * fix(guardrails): mint the vault id instead of trusting the caller's The vault id was taken from caller-supplied session metadata, and every caller shares one LLM Shield key. Someone who knew or guessed another caller's session id could send a placeholder, have the model echo it back, and get that caller's plaintext restored into their own reply. Vault ids are now minted per request behind a per-process prefix, so a caller cannot name a vault this process uses. Redaction mints, restoration reads back, and a reply whose id does not match is left holding its placeholders rather than resolved against some other vault. Also covers two more request fields that were reaching the provider intact: the Responses API `instructions`, and the legacy `function_call.arguments` alongside `tool_calls`. The collectors move to module level, which drops the traversal back under the complexity limit and lets the code carry its own explanation instead of the comments that were restating it. * fix(guardrails): drop Final from a loop-assigned local basedpyright rejects a Final assigned inside a loop, and reportGeneralTypeIssues sits one over its budget ceiling. * fix(guardrails): redact completion prompts and responses tool items Two more provider-bound request shapes were reaching the model intact while the guardrail reported as enabled. /v1/completions carries its text in a top-level `prompt`, which the traversal never looked at. It is handled as a string and as the array form, where each entry is rewritten in place. Responses input items hold tool data outside `content`: a function_call item in `arguments`, a function_call_output item in `output`. Both are now collected alongside the item's content. Adds a test per shape. * fix(guardrails): redact the anthropic system prompt and string-array input Two more provider-bound shapes, found by walking the request types rather than waiting for them to be reported. /v1/messages carries its system prompt at the top level, as a string or a list of text blocks. It is one of the endpoints this guardrail claims to cover, and a system prompt is a natural place to put a customer's details. `input` as an array of bare strings, the embeddings and moderations shape, was skipped because the loop only handled item dicts. Verified against a live provider: a system prompt holding an address now reaches the model as a stand-in and is restored in the reply. * fix(guardrails): narrow prompt and input to a list before iterating Guarding with a conditional iterable left the value un-narrowed, so passing it on was an argument-type error and the element checks read as unreachable. An early return narrows it properly and reads better. * fix(guardrails): restore every streaming choice, not just the first Streaming rehydration read and rewrote choices[0] only, so with n>1 every later choice went back to the caller still holding its placeholders. Each choice is its own token stream, so the sliding window is now tracked per choice index rather than once per stream. A single shared window would have been worse than the bug: it would splice the characters held back for one choice onto the next one's delta. The final flush walks every choice the same way, and the two helpers that only ever looked at choices[0] are gone. Adds a test that both choices come back restored, and one that each choice gets its own window handed back rather than its neighbour's. * refactor(guardrails): name the guardrail llm_shield_proxy throughout The integration was called llm_shield in code, llm-shield in the example config, and LLM Shield in the dashboard, while the product and its PyPI package are both llm-shield-proxy. An operator who saw the guardrail in LiteLLM could not tell what to install. One identifier now: llm_shield_proxy for the enum value, module, directory, class, config model, logo and environment variables, with LLM Shield Proxy as the display name. That matches `pip install llm-shield-proxy`. Renames only; no behaviour change. * feat(guardrails): redact the participant name on a message `name` on a user or assistant turn identifies a person and was going to the provider intact. The proxy this integrates with already redacts it, so the integration was the weaker of the two. On a tool or function turn the same field carries the function's name, which has to arrive unchanged or the call stops routing. That case is skipped, and a test asserts the value is never even sent to the shield. * fix(guardrails): flush every held choice, and cover tool results and suffix Three review findings. The trailing flush walked the last chunk's choices, so a choice that finished earlier and stopped appearing lost whatever text was still held for it and its answer was truncated. It is now driven by the windows themselves and emits one chunk per choice, synthesising the choice when the terminal chunk omits it. That was data loss, not just under-redaction. An Anthropic tool_result carries its own content, as a string or as further blocks, and only each part's `text` was being collected. Handled recursively; image and audio parts still fall through untouched. The legacy completions `suffix` is forwarded to providers that support it and was never collected. Note the placement: it has to be gathered before the string-prompt early return, which is what the new test pins. * fix(guardrails): walk nested tool results iteratively, with a depth bound CI flagged _collect_content as recursive. It was, and worse, it was unbounded: a tool_result nests its own content, the nesting is caller controlled, and the descent had nothing to stop it. That is a JSON bomb, not a style issue. Now an explicit queue with a depth bound of 8. Real payloads nest one or two deep. The queue is walked in document order because the shield maps its replies back by position, so collection order is part of the contract. * fix(guardrails): redact Responses PromptObject variables A Responses request can send `prompt` as a PromptObject rather than a string. Its `variables` are substituted into the stored prompt on the provider side, so they are caller text, and the dict shape was falling through untouched. `id` and `version` pick which stored prompt to run and are left unchanged. * test(guardrails): assert the depth bound instead of only reaching the end The depth test asserted nothing, so it passed whether or not the bound held, and the test-quality gate counted it as a zero-assert test. It now sends a shallow value alongside a 200-deep chain and asserts the shallow one is collected while the value past the bound is not. * fix(guardrails): keep system-prompt values out of the restored reply Redaction put every span of a request into one vault, and the reply was restored against that same vault. System prompts are written by the application and the caller never sees them, so a caller who got the model to echo a placeholder back had its plaintext restored into their own reply -- a way to read a system prompt they were never shown. Server-authored spans now go into a vault of their own: system and developer turns, Anthropic's top-level `system`, and the Responses API `instructions`. Its id is deliberately never stored, so nothing restores against it. The reply is restored against the caller's vault alone, and an echoed placeholder from a system prompt comes back as the placeholder. Values the caller also wrote themselves are unaffected -- they are in the caller's vault too, and still restore. The extra round trip happens only when a request actually carries server-authored text. * style(guardrails): satisfy ruff format and annotate the new tests `ruff format` wanted the widened `_collect_responses_fields` signature on one line, and the three tests added with the split-vault fix needed return annotations to keep ANN201 level with the base. * fix(guardrails): restore tool calls in the LLM Shield guardrail The request walk redacted a tool call's `arguments` -- plus the legacy `function_call`, Anthropic `tool_use.input` leaves and the Responses API's `function_call` / `function_call_output` fields -- while the response walk restored only `message.content`. A placeholder therefore reached the caller inside a tool call, and nothing raised. This is the same change as the out-of-tree example adapter this file is copied from, kept body-identical on purpose: the response side now collects every restorable span in one positional rehydrate batch, streaming keeps a window per (choice index, tool-call index) and flushes each into the chunk carrying the finish_reason, and `apply_guardrail` restores `inputs["tool_calls"]` on the response side. The declared limit on restoring values inside a JSON string is documented in the module. * fix(guardrails): import copy, keep the vault id off the provider, drop recursion Three defects Greptile and veria-ai found on the reopened PR, all real: - `copy.deepcopy` was called in `apply_guardrail` with no `import copy`, a guaranteed NameError on every response carrying tool calls. It landed on 2026-09-13, ten days after the review that rated this branch safe, and no test reached it: every tool-call test covered the request side. Adds the import and a regression test on the response side. - The vault session id was stored in `metadata`, which is forwarded to the provider on /v1/responses. A provider holding the placeholders and the session id can call the shield's rehydrate endpoint and read back the plaintext this guardrail exists to withhold. Moves it to `litellm_metadata`, which is not forwarded, and reads it back from there only. - `_collect_json_leaves` recursed over model-controlled JSON; the repo's recursive_detector gate rejects that. Rewritten with an explicit stack, same depth bound. 52 tests pass. ruff format, ruff-strict and check_type_discipline all clean, with LIT counts identical to the merge base. * fix(guardrails): build llm_shield_proxy stream deltas without new mutable literals The lint job's LIT002 budget gate failed on this PR: the file added 11 mutable-collection constructions and the tree sits at its limit. Build the index-only tool-call continuation in one helper, keep read-only inputs as tuples, and annotate the lists the delta and texts fields require. Adds tests for the two tool-call flush paths the refactor touches, which had no coverage: held arguments landing in the finish_reason chunk next to that chunk's own fragment, and the trailing flush of a stream that ends without a finish_reason. * fix(guardrails): drop Final from loop-body locals in llm_shield_proxy basedpyright rejects Final on a name assigned inside a loop, and the eleven such locals put reportGeneralTypeIssues over its budget (112/101). The LIT010 Final rule already exempts loop-body assignments, so the annotations go. * feat(guardrails): restore llm_shield_proxy placeholders on native streams Anthropic /v1/messages and /v1/responses streams have no `choices`, so the streaming hook passed them through with placeholders still in them. Both are now restored incrementally, with the same per-stream windows as chat: - /v1/messages arrives as raw SSE. Frames are cut at event boundaries, text_delta and input_json_delta are restored per block index, and held text is emitted as one more delta ahead of content_block_stop. Signed thinking deltas, frames from other endpoints and non-SSE raw streams pass through unchanged. - /v1/responses events are restored per item and part. Held text goes out as a copy of the stream's last delta before its .done event, and the events that repeat the reply (.done, content_part.done, output_item.done, response.completed) are restored in full. The request side now also redacts Anthropic tool_use inputs and Responses reasoning summaries, and sends tool and function descriptions (including parameter schema descriptions) and the user / safety_identifier fields to the non-restorable vault, like system prompts. Tool results stay restorable: the model reads them to answer, so restoring them returns what the caller would have seen without the guardrail. * fix(guardrails): redact llm_shield_proxy predicted outputs and output schemas `prediction.content` is the caller's own draft of the reply, so it is redacted into the caller vault and restored with the reply. The descriptions in a structured-output schema (Chat response_format.json_schema, Responses text.format) are application authored like tool schemas, so they go to the non-restorable vault. * fix(guardrails): fail closed on deep llm_shield_proxy requests, widen coverage - Request walks no longer skip what lies past their depth bound. Content nested past it, and tool inputs or schemas past the new JSON bound, now block the request instead of reaching the provider unredacted. The old depth test asserted the skip; it now asserts the block. - Tool and output schemas are walked by their JSON Schema structure, and give up `title`, `examples` and `default` as well as `description`. `enum` and `const` still go out as sent. - Responses events are matched by shape: any `*.delta` with a string delta is a token stream, and any `*.done` restores every non-identifier text field plus the `part` or `item` it repeats. This covers reasoning_summary_part.done and MCP arguments, and future families. Audio deltas are left alone. - An SSE stream whose first chunk ends partway through a field name (`b"eve"`) is no longer taken for a non-SSE stream. * fix(guardrails): scan llm_shield_proxy schemas by default The schema walk collected an allowlist of keywords, so any keyword it did not list -- draft-07 `dependencies`, `$comment`, vendor `x-` extensions -- went to the provider in clear. Invert it: every string is collected except under keywords whose value must go out verbatim (types, formats, patterns, references, required lists, enum, const). Name -> subschema maps still treat their keys as property names, so a property called `type` is walked, not skipped. * fix(guardrails): redact llm_shield_proxy schema enum and const values `enum` and `const` were skipped by the schema walk, so a value holding PII went to the provider in clear. They now go to the caller's vault rather than the non-restorable one: the model emits the stand-in in its tool arguments or structured output, and restoring the reply turns it back into the value the schema allows, so the call still routes. * fix(guardrails): redact llm_shield_proxy web search user locations Web search forwards the user's approximate location, and its free-text `city` and `region` fields can hold an address. Collect them into the non-restorable vault, from Chat `web_search_options.user_location` and from the `user_location` of Responses and Anthropic web-search tools. * fix(guardrails): drop unused llm_shield_proxy suppressions Upstream added LIT013 (a *-ok marker that suppresses nothing) and LIT014 (at most one for and one if per comprehension). Remove the 34 markers that no longer suppress anything and flatten the finished streams with itertools.chain.from_iterable. * fix(guardrails): type the llm_shield_proxy request and reply walks Narrowing with isinstance(x, dict) leaves keys and values unknown, so every call that passed a narrowed value counted against the reportUnknownArgumentType budget. Parse into dict[str, object] and list[object] once, in _as_object and _as_array, type the carry keys and accumulators, and bind writers with functools.partial instead of lambdas. The shield's batch reply is now also checked to hold only strings. * fix(guardrails): keep restored llm_shield_proxy replies out of the cache, widen coverage Addresses the open veria-ai and Cursor Bugbot findings on #42645. - Restore a copy of the reply and of each stream chunk, never LiteLLM's own object. LiteLLM caches and logs that object, and placeholders are numbered per request, so two callers' redacted requests can share a cache key: restoring in place cached one caller's plaintext for the next. The deployment hook no longer restores either, since LiteLLM caches what it returns; the proxy's post-call hook restores model-level guardrails after the cache write. - Restore /v1/completions replies, streamed and not, which carry `choice.text`. - Redact Responses replay fields the reply side already restores: tool output sent as input_text parts, custom_tool_call `input`, code_interpreter_call `code`. - Redact typed Responses prompt variables (`{"type": "input_text", "text": ...}`). - Put Responses system and developer input items in the non-restorable vault, like their Chat counterparts. - Expose LLMShieldProxyGuardrailConfigModel through get_config_model, so the dashboard can collect the Shield URL and key. * fix(guardrails): redact llm_shield_proxy plain-text document blocks An Anthropic document block carries text inline, in a text source's `data` or a content source's `content`, and that text reached the provider unredacted. Collect both, plus the block's `title` and `context`; base64, URL and file sources pass untouched. * fix(guardrails): redact llm_shield_proxy extra_body overrides LiteLLM merges extra_body over the transformed request just before sending, so text placed there (input, messages, system, ...) replaced the redacted field on the wire. Walk extra_body with the same collectors as the request, keeping the caller / application split. * test(guardrails): import InMemoryCache directly in the llm_shield_proxy cache test litellm keeps a deprecated module-level `caching` bool, so `litellm.caching.caching` resolves to that bool once an earlier test in the same worker has set it, and the test failed with AttributeError depending on test order. * fix(guardrails): restore llm_shield_proxy replies for model-level use outside the proxy 71e68fd stopped the deployment post-call hook from restoring, so the response cache never holds restored plaintext. Inside the proxy that is right: the proxy's post-call hook restores after the cache write. But with model-level `guardrails` on the SDK, the deployment hooks are the only redact and restore steps, so callers got placeholders back. When the deployment pre-call hook is the one that redacts, it now records the request's vault id and marks the request no-cache / no-store; the deployment post-call hook restores only when that record matches. The cache key there is built from the redacted request and a cache hit skips the post-call hook, so a cached reply could neither be restored nor safely shared. Proxy requests carry no record and keep restoring in the proxy's post-call hook, after the cache write. * fix(guardrails): don't repeat usage in llm_shield_proxy end-of-stream flush chunks With n>=2 and stream_options.include_usage, the end-of-stream flush copies the last chunk the stream carried, which is the one holding usage, so each synthetic flush chunk repeated it and a consumer summing usage chunks counted the request twice. The copy now drops `usage`, matching a normal mid-stream chunk. Reported by @yucheng-berri. * fix(guardrails): keep restored llm_shield_proxy values out of telemetry, refuse SDK streams - The post-call restore hook no longer goes through log_guardrail_information, which recorded its whole return value, the restored reply, as guardrail_response. That field is exported to traces even with message logging turned off. - A model-level stream outside the proxy is refused once redacted. Nothing restores an SDK stream, and its cache writer reads the request from before the deployment hook, so it also got cached despite the no-store bypass. - Drop a narrating comment, and keep example_config.yaml to config only; the how-to lives in the docs PR. * refactor(guardrails): split llm_shield_proxy into payload, request walk and stream modules The module had grown past 1,700 lines. Shared payload types and helpers move to payload.py, the request walk to request_walk.py and the stream restorers to stream_restorers.py; llm_shield_proxy.py keeps the guardrail class. No behaviour change. * style(guardrails): drop routine comments from llm_shield_proxy AGENTS.md keeps source comments to tool directives and genuinely complex logic; the rationale stays in the docstrings. --- .../llm_shield_proxy/__init__.py | 33 + .../llm_shield_proxy/example_config.yaml | 14 + .../llm_shield_proxy/llm_shield_proxy.py | 744 ++++++ .../llm_shield_proxy/payload.py | 190 ++ .../llm_shield_proxy/request_walk.py | 363 +++ .../llm_shield_proxy/stream_restorers.py | 345 +++ litellm/types/guardrails.py | 1 + .../guardrail_hooks/llm_shield_proxy.py | 24 + ruff-strict.toml | 4 + .../guardrail_hooks/test_llm_shield_proxy.py | 2003 +++++++++++++++++ .../public/assets/logos/llm_shield_proxy.svg | 6 + .../_components/add_guardrail_form.tsx | 4 +- .../_components/guardrail_garden_configs.ts | 12 +- .../_components/guardrail_garden_data.test.ts | 1 + .../_components/guardrail_garden_data.ts | 10 + .../_components/guardrail_info_helpers.tsx | 3 + 16 files changed, 3755 insertions(+), 2 deletions(-) create mode 100644 litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/__init__.py create mode 100644 litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/example_config.yaml create mode 100644 litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py create mode 100644 litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/payload.py create mode 100644 litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/request_walk.py create mode 100644 litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/stream_restorers.py create mode 100644 litellm/types/proxy/guardrails/guardrail_hooks/llm_shield_proxy.py create mode 100644 tests/unit/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py create mode 100644 ui/litellm-dashboard/public/assets/logos/llm_shield_proxy.svg diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/__init__.py new file mode 100644 index 00000000000..44b19f82218 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/__init__.py @@ -0,0 +1,33 @@ +from typing import TYPE_CHECKING, Final + +from litellm.types.guardrails import SupportedGuardrailIntegrations + +from .llm_shield_proxy import LLMShieldProxyGuardrail + +if TYPE_CHECKING: + from litellm.types.guardrails import Guardrail, LitellmParams + + +def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail") -> LLMShieldProxyGuardrail: + import litellm + + _llm_shield_guardrail_callback: Final = LLMShieldProxyGuardrail( + api_key=litellm_params.api_key, + api_base=litellm_params.api_base, + guardrail_name=guardrail.get("guardrail_name", ""), + event_hook=litellm_params.mode, + default_on=litellm_params.default_on, + ) + + litellm.logging_callback_manager.add_litellm_callback(_llm_shield_guardrail_callback) + return _llm_shield_guardrail_callback + + +guardrail_initializer_registry: Final = { + SupportedGuardrailIntegrations.LLM_SHIELD_PROXY.value: initialize_guardrail, +} + + +guardrail_class_registry: Final = { + SupportedGuardrailIntegrations.LLM_SHIELD_PROXY.value: LLMShieldProxyGuardrail, +} diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/example_config.yaml b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/example_config.yaml new file mode 100644 index 00000000000..4f732773a75 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/example_config.yaml @@ -0,0 +1,14 @@ +model_list: + - model_name: gpt-4o + litellm_params: + model: openai/gpt-4o + api_key: os.environ/OPENAI_API_KEY + +guardrails: + - guardrail_name: "llm_shield_proxy" + litellm_params: + guardrail: llm_shield_proxy + mode: ["pre_call", "post_call"] + default_on: true + api_base: "http://localhost:8000" + api_key: os.environ/LLM_SHIELD_PROXY_API_KEY diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py new file mode 100644 index 00000000000..fa23a110a74 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -0,0 +1,744 @@ +import copy +import functools +import os +import uuid +from collections.abc import AsyncGenerator, Callable, Mapping, Sequence +from typing import ( + TYPE_CHECKING, + Any, # noqa: TID251 # **kwargs forwards verbatim to CustomGuardrail.__init__ + ClassVar, + Final, + Literal, + Optional, +) + +import httpx + +from litellm._logging import verbose_proxy_logger +from litellm.exceptions import GuardrailRaisedException +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + log_guardrail_information, +) +from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, + httpxSpecialProvider, +) +from litellm.proxy._types import UserAPIKeyAuth +from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.utils import GenericGuardrailAPIInputs, TextChoices + +if TYPE_CHECKING: + from litellm.caching.caching import DualCache + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.types.proxy.guardrails.guardrail_hooks.llm_shield_proxy import ( + LLMShieldProxyGuardrailConfigModel, + ) + from litellm.types.utils import CallTypes, LLMResponseTypes +from .payload import ( + JsonBody, + MutableRequest, + RequestTooDeep, + Slot, + SlotSink, + as_array, + as_object, + choice_index, + collect_json_leaves, + collect_response_item, + detached, + read_field, + read_list, + rehydrate_slots, + write_field, +) +from .request_walk import ( + locate_request_texts, +) +from .stream_restorers import ( + AnthropicSSERestorer, + CarryKey, + CarryWindows, + ResponsesStreamRestorer, + carry_sort_key, + continuation_delta, + responses_event_type, +) + +GUARDRAIL_NAME: Final = "llm_shield_proxy" + +_DEFAULT_API_BASE: Final = "http://localhost:8000" +_REDACT_PATH: Final = "/v1/guard/redact" +_REHYDRATE_PATH: Final = "/v1/guard/rehydrate" +_REHYDRATE_STREAM_PATH: Final = "/v1/guard/rehydrate/stream" + +_SESSION_METADATA_KEY: Final = "llm_shield_session_id" + +_DEPLOYMENT_RESTORE_KEY: Final = "llm_shield_restore_at_deployment" + +_VAULT_PREFIX: Final = f"litellm-{uuid.uuid4().hex}" + +_DEFAULT_TIMEOUT_SECONDS: Final = 10.0 + + +class LLMShieldProxyGuardrail(CustomGuardrail): + """Redacts PII before it leaves the proxy and restores it in the response. + + Unlike a masking guardrail, the substitution is reversible. Outbound text is + replaced with placeholders held in a session vault inside the user's own LLM + Shield deployment; the model's reply is then restored so the end user sees the + original values while the provider never received them. + + Streaming is restored incrementally rather than by buffering the response. LLM + Shield holds back only the trailing characters that could still turn out to be + part of a placeholder, so tokens are forwarded as they arrive and a placeholder + split across two chunks is never emitted in fragments. + """ + + use_native_lifecycle_hooks: ClassVar[bool] = True + + def __init__( + self, + guardrail_name: str = GUARDRAIL_NAME, + api_base: str | None = None, + api_key: str | None = None, + **kwargs: Any, # noqa: LIT008 # kwargs-ok: forwarded verbatim to CustomGuardrail.__init__ + ) -> None: + self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback) + env_base: Final = os.environ.get("LLM_SHIELD_PROXY_API_BASE") + self.api_base: Final = (api_base or env_base or _DEFAULT_API_BASE).rstrip("/") + self.api_key: Final = api_key or os.environ.get("LLM_SHIELD_PROXY_API_KEY") + super().__init__(guardrail_name=guardrail_name, **kwargs) + + @classmethod + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: # mutable-ok: parent's signature. + return [GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call] + + @staticmethod + def get_config_model() -> type["LLMShieldProxyGuardrailConfigModel"]: + from litellm.types.proxy.guardrails.guardrail_hooks.llm_shield_proxy import ( + LLMShieldProxyGuardrailConfigModel, + ) + + return LLMShieldProxyGuardrailConfigModel + + async def async_pre_call_deployment_hook( + self, + kwargs: MutableRequest, + call_type: "CallTypes | None", + ) -> MutableRequest | None: + """Redacts a model-level guardrail's request, and keeps it out of the response cache. + + Outside the proxy this hook is the only redaction step, and the deployment post-call + hook the only restoration step. LiteLLM builds the cache key after this hook, from + the redacted request, and a cache hit returns before the post-call hook runs. So a + cached reply would either reach the caller unrestored or, stored after restoration, + hand this caller's values to the next caller whose redacted request matches. The + request is therefore neither read from nor written to the cache. Inside the proxy + this hook does not redact -- the proxy's pre-call hook already ran -- and caching + is left alone, because the proxy restores after the cache write. + + A streamed request is refused once redacted. No hook restores an SDK stream, and the + stream's cache writer reads the request from before this hook, so it would also be + cached despite the bypass. + """ + before: Final = self._minted_session_id(kwargs) + _ = await super().async_pre_call_deployment_hook(kwargs, call_type) + session_id: Final = self._minted_session_id(kwargs) + if session_id is None or session_id == before: + return kwargs + if kwargs.get("stream") is True: + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=( + "LLM Shield Proxy cannot restore a streamed reply for a model-level guardrail " + "outside the LiteLLM proxy; send the request through the proxy or without stream=True." + ), + ) + metadata: Final = as_object(kwargs.get("litellm_metadata")) + if metadata is not None: + metadata[_DEPLOYMENT_RESTORE_KEY] = session_id + cache_controls: Final = as_object(kwargs.get("cache")) + kwargs["cache"] = {**(cache_controls or {}), "no-cache": True, "no-store": True} + return kwargs + + async def async_post_call_success_deployment_hook( + self, + request_data: MutableRequest, + response: "LLMResponseTypes", + call_type: "CallTypes | None", + ) -> "LLMResponseTypes | None": + """Restores the reply here only when the deployment pre-call hook redacted it. + + LiteLLM caches what this hook returns. Inside the proxy the request was redacted by + the proxy's pre-call hook and the proxy's post-call hook restores the reply after + the cache write, so restoring here as well would cache this caller's plaintext under + a key built from the redacted request. Outside the proxy nothing restores later, and + the pre-call deployment hook has already kept that request out of the cache. + """ + metadata: Final = as_object(request_data.get("litellm_metadata")) + marker: Final = metadata.get(_DEPLOYMENT_RESTORE_KEY) if metadata is not None else None + session_id: Final = self._minted_session_id(request_data) + if session_id is None or marker != session_id: + return None + return await super().async_post_call_success_deployment_hook(request_data, response, call_type) + + def _headers(self, session_id: str) -> JsonBody: + headers: Final[JsonBody] = { + "Content-Type": "application/json", + "X-Session-ID": session_id, + } + if self.api_key: + headers["Authorization"] = f"Bearer {self.api_key}" + return headers + + async def _call_shield(self, path: str, session_id: str, payload: JsonBody) -> Mapping[str, object]: + """Posts to LLM Shield Proxy, failing closed on any transport or status error. + + A redaction guardrail that fails open sends the very data it exists to + protect to a third-party provider, so an unreachable or erroring shield + blocks the request instead of passing it through. + """ + try: + response: Final = await self.async_handler.post( + f"{self.api_base}{path}", + headers=self._headers(session_id), + json=payload, + timeout=_DEFAULT_TIMEOUT_SECONDS, + ) + response.raise_for_status() + return response.json() + except httpx.HTTPStatusError as exc: + verbose_proxy_logger.exception("LLM Shield Proxy returned %s for %s", exc.response.status_code, path) + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=f"LLM Shield Proxy returned {exc.response.status_code}; blocking the request.", + ) from exc + except Exception as exc: + verbose_proxy_logger.exception("LLM Shield Proxy call to %s failed", path) + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message="LLM Shield Proxy is unreachable; blocking the request.", + ) from exc + + async def _redact(self, texts: Sequence[str], session_id: str) -> Sequence[str]: + payload: Final[JsonBody] = {"texts": list(texts)} + body: Final = await self._call_shield(_REDACT_PATH, session_id, payload) + return self._same_length_or_raise(body.get("texts"), texts, "redact") + + async def _rehydrate(self, texts: Sequence[str], session_id: str) -> Sequence[str]: + payload: Final[JsonBody] = {"texts": list(texts)} + body: Final = await self._call_shield(_REHYDRATE_PATH, session_id, payload) + return self._same_length_or_raise(body.get("texts"), texts, "rehydrate") + + def _same_length_or_raise(self, returned: object, sent: Sequence[str], operation: str) -> Sequence[str]: + """Guards the positional mapping the callers rely on to write results back.""" + entries: Final = as_array(returned) + texts: Final = tuple(entry for entry in entries or () if isinstance(entry, str)) + if entries is None or len(entries) != len(sent) or len(texts) != len(entries): + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=f"LLM Shield Proxy {operation} returned an unexpected payload; blocking the request.", + ) + return texts + + @staticmethod + def _mint_session_id(data: MutableRequest) -> str: + """Mints a vault id for this request, overwriting anything already there. + + Redaction and restoration both happen inside one request/response pair, so + a fresh id per request is all that is needed, and it is what keeps one + caller from reaching another caller's vault. + """ + session_id: Final = f"{_VAULT_PREFIX}-{uuid.uuid4().hex}" + metadata: Final = data.setdefault("litellm_metadata", {}) + if isinstance(metadata, dict): + metadata[_SESSION_METADATA_KEY] = session_id + return session_id + + @staticmethod + def _minted_session_id(data: MutableRequest) -> str | None: + """The vault id this process minted for `data`, or None if it has none. + + Read only from `litellm_metadata`, the proxy-private store `_mint_session_id` writes + to. A caller can populate `metadata`; they cannot populate this. + """ + metadata: Final = as_object(data.get("litellm_metadata")) + existing: Final = metadata.get(_SESSION_METADATA_KEY) if metadata is not None else None + return existing if isinstance(existing, str) and existing.startswith(_VAULT_PREFIX) else None + + @staticmethod + def _session_id(data: MutableRequest) -> str: + """Reads back the vault id minted while redacting this request. + + Falls back to an unused id rather than to anything the caller supplied: a + reply that cannot be restored is a visible placeholder, while trusting a + caller-supplied id would hand them someone else's plaintext. + """ + existing: Final = LLMShieldProxyGuardrail._minted_session_id(data) + return existing if existing is not None else f"{_VAULT_PREFIX}-{uuid.uuid4().hex}" + + @staticmethod + def _locate_request_texts(data: MutableRequest) -> tuple[Sequence[Slot], Sequence[Slot]]: + return locate_request_texts(data) + + @log_guardrail_information + async def async_pre_call_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + cache: "DualCache", + data: MutableRequest, + call_type: str, + ) -> MutableRequest | None: + """Replaces PII anywhere in the outbound request with vault placeholders.""" + if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_call) is not True: + return data + + try: + slots, privileged = self._locate_request_texts(data) + except RequestTooDeep as exc: + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=f"Request {exc} nests deeper than LLM Shield Proxy inspects; blocking the request.", + ) from exc + if not slots and not privileged: + return data + + session_id: Final = self._mint_session_id(data) + if privileged: + await self._redact_into(privileged, f"{_VAULT_PREFIX}-{uuid.uuid4().hex}") + if slots: + await self._redact_into(slots, session_id) + return data + + async def _redact_into(self, slots: Sequence[Slot], session_id: str) -> None: + """Redacts every span in `slots` under one vault and writes the result back.""" + redacted: Final = await self._redact(tuple(text for text, _ in slots), session_id) + for (_, write), replacement in zip(slots, redacted): + write(replacement) + + async def async_post_call_success_hook( + self, + data: MutableRequest, + user_api_key_dict: UserAPIKeyAuth, + response: Any, + ) -> Any: + """Restores the original values in a copy of a non-streaming response. + + The copy is what keeps plaintext out of the response cache. LiteLLM caches the + reply it received from the provider, and on some paths -- a native Anthropic dict, + an in-memory cache -- it stores the object itself rather than a serialised + snapshot. Restoring that object in place would cache this caller's values under a + key built from the redacted request, which another caller's identical-looking + request then hits. Left untouched, the cached reply holds placeholders, and a hit + is restored against the new caller's own vault. + """ + if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.post_call) is not True: + return response + response = detached(response) # rebind-ok: everything below restores the copy. + + if self._is_anthropic_message_response(response): + return await self._restore_anthropic_response(response, data) + + response_slots: Final = self._responses_api_slots(response) + if response_slots: + return await self._restore_responses_api_response(response, response_slots, data) + + choices: Final = getattr(response, "choices", None) + if not choices: + return response + + pending: Final[SlotSink] = [] + for choice in choices: + message = getattr(choice, "message", None) + if message is None: + text = read_field(choice, "text") + if isinstance(text, str) and text: + pending.append((text, functools.partial(write_field, choice, "text"))) + continue + content = getattr(message, "content", None) + if isinstance(content, str) and content: + pending.append((content, functools.partial(setattr, message, "content"))) + for tool_call in getattr(message, "tool_calls", None) or (): + function = getattr(tool_call, "function", None) + arguments = getattr(function, "arguments", None) if function is not None else None + if isinstance(arguments, str) and arguments: + pending.append((arguments, functools.partial(setattr, function, "arguments"))) + legacy = getattr(message, "function_call", None) + legacy_arguments = getattr(legacy, "arguments", None) if legacy is not None else None + if isinstance(legacy_arguments, str) and legacy_arguments: + pending.append((legacy_arguments, functools.partial(setattr, legacy, "arguments"))) + if not pending: + return response + + restored: Final = await self._rehydrate(tuple(text for text, _ in pending), self._session_id(data)) + for (_, write), replacement in zip(pending, restored): + write(replacement) + return response + + @staticmethod + def _is_anthropic_message_response(response: object) -> bool: + """Anthropic's native /v1/messages reply arrives as a plain dict.""" + body: Final = as_object(response) + return body is not None and body.get("type") == "message" and isinstance(body.get("content"), list) + + async def _restore_anthropic_response(self, response: MutableRequest, data: MutableRequest) -> MutableRequest: + """Restores text blocks and tool inputs in an Anthropic native message reply. + + This shape has no `choices`, so without its own branch the reply would go + back to the caller still carrying placeholders. + + A `tool_use` block's payload is `input`, an arbitrary JSON object rather than a + string, and the request path redacts its string leaves -- so the reply's leaves + have to come back or the application invokes the tool with placeholders. + """ + slots: Final[SlotSink] = [] + for entry in read_list(response, "content"): + block = as_object(entry) + if block is None: + continue + kind = block.get("type") + text = block.get("text") + if kind == "text" and isinstance(text, str) and text: + slots.append((text, functools.partial(block.__setitem__, "text"))) + elif kind == "tool_use" and as_object(block.get("input")) is not None: + collect_json_leaves(block.get("input"), slots) + if not slots: + return response + + restored: Final = await self._rehydrate(tuple(text for text, _ in slots), self._session_id(data)) + for (_, write), replacement in zip(slots, restored): + write(replacement) + return response + + @staticmethod + def _responses_api_slots(response: object) -> Sequence[Slot]: + """Restorable spans in a Responses API reply. + + That shape carries `output` items rather than `choices`, so it needs its own + walk; without one the reply goes back to the caller still holding + placeholders even though the request was redacted correctly. Items and blocks + come through as dicts or as objects depending on how far the reply has been + deserialised, so both are handled. + + The item-level fields mirror `collect_responses_fields`, which walks the same + fields on the request side -- a function_call item holds `arguments`, a + function_call_output holds `output` -- so the two directions stay symmetric. + """ + slots: Final[SlotSink] = [] + for item in getattr(response, "output", None) or (): + collect_response_item(item, slots) + return tuple(slots) + + async def _restore_responses_api_response(self, response: Any, slots: Sequence[Slot], data: MutableRequest) -> Any: + """Puts the original values back into a Responses API reply.""" + await rehydrate_slots(slots, functools.partial(self._rehydrate, session_id=self._session_id(data))) + return response + + async def async_post_call_streaming_iterator_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + response: Any, + request_data: MutableRequest, + ) -> AsyncGenerator[Any, None]: + """Restores original values incrementally, without buffering the stream. + + Each choice -- and each tool call within a choice -- is its own token stream, so + the sliding window is tracked per (choice index, tool call) pair. One shared + window would splice the characters held back for one stream onto another. The + windows are locals of this generator, so they are scoped to a single stream and + cannot leak between concurrent requests. + + The two native stream shapes have no `choices` and are restored by their own + walkers, with the same per-stream windows: Anthropic `/v1/messages` arrives as raw + SSE frames, and the Responses API as typed events. + """ + if self.should_run_guardrail(data=request_data, event_type=GuardrailEventHooks.post_call) is not True: + async for chunk in response: + yield chunk + return + + session_id: Final = self._session_id(request_data) + step: Final = functools.partial(self._stream_step, session_id=session_id) + rehydrate: Final = functools.partial(self._rehydrate, session_id=session_id) + sse: Final = AnthropicSSERestorer(step) + events: Final = ResponsesStreamRestorer(step, rehydrate) + carries: Final[CarryWindows] = {} + last_chunk = None # rebind-ok: tracks the most recent chunk for the final flush. + + async for chunk in response: + if isinstance(chunk, (bytes, str)): + for frames in await sse.feed(chunk): + yield frames + continue + if responses_event_type(chunk) is not None: + for event in await events.restore(detached(chunk)): + yield event + continue + restored_chunk = detached(chunk) + last_chunk = restored_chunk + for choice in getattr(restored_chunk, "choices", None) or (): + await self._restore_choice(choice, carries, session_id) + yield restored_chunk + + for frames in await sse.finish(): + yield frames + for event in await events.finish(): + yield event + if last_chunk is not None and any(carries.values()): + async for trailing in self._flush_trailing(last_chunk, carries, session_id): + yield trailing + + async def _restore_choice(self, choice: object, carries: CarryWindows, session_id: str) -> None: + """Restores one choice's delta, advancing that choice's own windows. + + Content and each tool call are separate token streams, so each gets its own + window: `(choice_index, None)` for content, `(choice_index, tool_call_index)` for + one tool call's accumulating `arguments`. A shared window would splice the text + held back for one stream onto another. + """ + delta: Final = getattr(choice, "delta", None) + index: Final = choice_index(choice) + is_final: Final = bool(getattr(choice, "finish_reason", None)) + if isinstance(choice, TextChoices): + await self._restore_text_window(choice, (index, None), carries, session_id, is_final) + return + if delta is None: + return + + await self._restore_content_window(delta, (index, None), carries, session_id, is_final) + + for tool_call in getattr(delta, "tool_calls", None) or (): + await self._restore_tool_call_window(tool_call, index, carries, session_id) + + if is_final: + await self._flush_finished_choice(delta, index, carries, session_id) + + async def _restore_text_window( + self, + choice: object, + key: CarryKey, + carries: CarryWindows, + session_id: str, + is_final: bool, + ) -> None: + """Restores a Completions stream choice's `text` through its window.""" + carry: Final = carries.get(key, "") + text: Final = read_field(choice, "text") + if not isinstance(text, str) or not text: + if not (is_final and carry): + return + emitted, remaining = await self._stream_step(text if isinstance(text, str) else "", carry, is_final, session_id) + carries[key] = remaining # rebind-ok: this stream's window advances. + if emitted or text: + write_field(choice, "text", emitted) + + async def _restore_content_window( + self, + delta: Any, + key: CarryKey, + carries: CarryWindows, + session_id: str, + is_final: bool, + ) -> None: + """Restores one delta's content through its own window.""" + carry: Final = carries.get(key, "") + text: Final = getattr(delta, "content", None) + + if not isinstance(text, str) or not text: + if is_final and carry: + flushed, remaining = await self._stream_step("", carry, True, session_id) + carries[key] = remaining # rebind-ok: this stream's window advances. + if flushed: + delta.content = flushed + return + + emitted, remaining = await self._stream_step(text, carry, is_final, session_id) + carries[key] = remaining # rebind-ok: this stream's window advances. + delta.content = emitted + + async def _restore_tool_call_window( + self, + tool_call: Any, + choice_index: int, + carries: CarryWindows, + session_id: str, + ) -> None: + """Restores one streamed tool call's argument fragment. + + A tool call's `arguments` is a JSON document delivered as fragments that clients + concatenate per tool-call index, so each index gets a window of its own rather + than sharing the content stream's. + """ + tool_index: Final = read_field(tool_call, "index") + if not isinstance(tool_index, int): + return + function: Final = read_field(tool_call, "function") + if function is None: + return + arguments: Final = read_field(function, "arguments") + if not isinstance(arguments, str) or not arguments: + return + + key: Final = (choice_index, tool_index) + emitted, remaining = await self._stream_step(arguments, carries.get(key, ""), False, session_id) + carries[key] = remaining # rebind-ok: this tool call's window advances. + write_field(function, "arguments", emitted) + + async def _flush_finished_choice( + self, + delta: Any, + choice_index: int, + carries: CarryWindows, + session_id: str, + ) -> None: + """Emits everything this finishing choice still holds, into this chunk. + + A client parses a tool call's `arguments` when the chunk carrying the + finish_reason arrives, so a flush delivered afterwards is too late -- the client + has already tried to parse truncated JSON. Content lands back on `content`; held + tool-call text is appended as an index-only continuation entry, which is the shape + clients concatenate by index, so no id or name is needed. Appending is correct + even when this chunk already carried a fragment for that tool call. + """ + continuations: Final[list[dict[str, object]]] = [] # mutable-ok: built into this chunk's delta. + for key in sorted((held for held in carries if held[0] == choice_index), key=carry_sort_key): + carry = carries[key] + if not carry: + continue + _, tool_index = key + text, remaining = await self._stream_step("", carry, True, session_id) + carries[key] = remaining # rebind-ok: this stream's window advances. + if not text: + continue + if tool_index is None: + delta.content = text + else: + continuations.extend(continuation_delta(tool_index, text)) + if continuations: + existing: Final = tuple(getattr(delta, "tool_calls", None) or ()) + delta.tool_calls = [*existing, *continuations] + + async def _flush_trailing( + self, last_chunk: Any, carries: CarryWindows, session_id: str + ) -> AsyncGenerator[Any, None]: + """Empties every window still holding text, one chunk per window. + + This is the net for a stream that ended with no finish_reason at all; a stream + that ended with one is flushed into its own terminal chunk by + `_flush_finished_choice`, because that is the moment a client parses tool + arguments. + + Driven by the windows rather than by the last chunk's choices. A choice that + finished earlier is not present in the terminal chunk, and flushing only what + that chunk carries would drop its held text and truncate its answer. + """ + for key in sorted(carries, key=carry_sort_key): + carry = carries[key] + if not carry: + continue + choice_index, tool_index = key + text, remaining = await self._stream_step("", carry, True, session_id) + carries[key] = remaining # rebind-ok: this stream's window advances. + if not text: + continue + chunk = self._chunk_for_choice(last_chunk, choice_index) + if chunk is None: + continue + choice: object = chunk.choices[0] + delta = read_field(choice, "delta") + if isinstance(choice, TextChoices): + choice.text = text + elif tool_index is None: + write_field(delta, "content", text) + else: + write_field(delta, "content", None) + write_field(delta, "tool_calls", continuation_delta(tool_index, text)) + yield chunk + + @staticmethod + def _chunk_for_choice(last_chunk: Any, index: int) -> Any: + """A single-choice copy of the last chunk, carrying only `index`. + + Emitting one choice per chunk keeps a flush from reading as content on a + choice it does not belong to. + """ + chunk: Final = last_chunk.model_copy(deep=True) + raw_choices: Final = getattr(chunk, "choices", None) + if not raw_choices: + return None + choices: Final[tuple[object, ...]] = tuple(raw_choices) + position: Final = next((at for at, choice in enumerate(choices) if choice_index(choice) == index), 0) + kept: Final = raw_choices[position] + if getattr(kept, "delta", None) is None and not isinstance(kept, TextChoices): + return None + kept.index = index + kept.finish_reason = None + chunk.choices = [kept] + if hasattr(chunk, "usage"): + del chunk.usage + return chunk + + async def _stream_step(self, text: str, carry: str, final: bool, session_id: str) -> tuple[str, str]: + """Returns ``(text safe to emit now, window still being held)``.""" + body: Final = await self._call_shield( + _REHYDRATE_STREAM_PATH, + session_id, + {"text": text, "carry": carry, "final": final}, + ) + emitted: Final = body.get("text") + remaining: Final = body.get("carry") + if not isinstance(emitted, str) or not isinstance(remaining, str): + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message="LLM Shield Proxy stream rehydration returned an unexpected payload.", + ) + return emitted, remaining + + @log_guardrail_information + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: MutableRequest, + input_type: Literal["request", "response"], + logging_obj: Optional["LiteLLMLoggingObj"] = None, + ) -> GenericGuardrailAPIInputs: + """Unified entry point: what the UI's Test guardrail button and the translation + handlers call. + + `tool_calls` is handled on the response side only. LiteLLM populates the field here, + and on a reply it holds the model's tool arguments -- the same text the native hook + restores, and restoring one but not the other would leave the placeholder on + whichever path ran. The request side is left to the native pre-call hook, because + redacting it here as well would redact it twice. + """ + text_list: Final = tuple(inputs.get("texts") or ()) + tool_calls: Final = tuple(inputs.get("tool_calls") or ()) if input_type == "response" else () + if not text_list and not tool_calls: + return inputs + + restored_calls: Final[list[object]] = [copy.deepcopy(call) for call in tool_calls] # mutable-ok: a new list. + spans: Final[list[str]] = list(text_list) # mutable-ok: ordered batch, frozen before the call. + writers: Final[list[Callable[[str], None]]] = [] # mutable-ok: one per span appended below. + for call in restored_calls: + function = read_field(call, "function") + arguments = read_field(function, "arguments") if function is not None else None + if isinstance(arguments, str) and arguments: + spans.append(arguments) + writers.append(functools.partial(write_field, function, "arguments")) + + replaced: Final = ( + await self._redact(tuple(spans), self._mint_session_id(request_data)) + if input_type == "request" + else await self._rehydrate(tuple(spans), self._session_id(request_data)) + ) + restored_values: Final[list[str]] = list(replaced) # mutable-ok: sliced into the texts list. + + for write, replacement in zip(writers, restored_values[len(text_list) :]): + write(replacement) + merged: Final[JsonBody] = {**inputs} + if text_list: + merged["texts"] = restored_values[: len(text_list)] + if restored_calls: + merged["tool_calls"] = restored_calls + return merged diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/payload.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/payload.py new file mode 100644 index 00000000000..644ac76efb8 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/payload.py @@ -0,0 +1,190 @@ +import copy +import functools +from collections.abc import Awaitable, Callable, Sequence +from typing import ( + Final, + TypeAlias, +) + +MutableRequest: TypeAlias = dict[str, object] + +JsonBody: TypeAlias = dict[str, object] + +MAX_JSON_DEPTH: Final = 64 + +Slot: TypeAlias = tuple[str, Callable[[str], None]] + +StreamStep: TypeAlias = Callable[[str, str, bool], Awaitable[tuple[str, str]]] + +Rehydrate: TypeAlias = Callable[[Sequence[str]], Awaitable[Sequence[str]]] + +SlotSink: TypeAlias = list[Slot] + +MutableSeq: TypeAlias = list[object] + + +def as_object(value: object) -> MutableRequest | None: + """`value` as a JSON object, or None. + + `isinstance(value, dict)` alone leaves the keys and values unknown to the type + checker. A JSON object's keys are strings, so the type is stated once, here. + """ + return value if isinstance(value, dict) else None + + +def as_array(value: object) -> MutableSeq | None: + """`value` as a JSON array, or None. See `as_object`.""" + return value if isinstance(value, list) else None + + +def detached(value: object) -> object: + """A deep copy of a reply or chunk, for restoring without touching LiteLLM's own object. + + LiteLLM keeps the object it handed the hooks to fill its response cache and its + logs, so writing restored plaintext into that object would put it there too. + """ + return copy.deepcopy(value) + + +def is_container(value: object) -> bool: + """Whether `value` is a JSON object or array, without narrowing it to unknown types.""" + return isinstance(value, (dict, list)) + + +def collect(container: MutableRequest, key: str, slots: SlotSink) -> None: + """Records the string at `key`, along with the write that replaces it.""" + value: Final = container.get(key) + if isinstance(value, str) and value: + slots.append((value, functools.partial(container.__setitem__, key))) + + +def collect_entry(entries: MutableSeq, index: int, slots: SlotSink) -> None: + """Records a string held directly in a list, rather than under a key.""" + value: Final = entries[index] + if isinstance(value, str) and value: + slots.append((value, functools.partial(entries.__setitem__, index))) + + +class RequestTooDeep(Exception): + """A request nests text past a walk's bound. + + Skipping the rest would forward it unredacted while the guardrail reports as + enabled, so the pre-call hook refuses the request instead. + """ + + +def collect_text_parts(container: MutableRequest, key: str, slots: SlotSink) -> None: + """Collects the `text` of every part in the list held at `key`.""" + for entry in as_array(container.get(key)) or (): + part = as_object(entry) + if part is not None: + collect(part, "text", slots) + + +def choice_index(choice: object) -> int: + """Streaming choices are matched across chunks by their index.""" + index: Final = getattr(choice, "index", 0) + return index if isinstance(index, int) else 0 + + +def read_field(holder: object, name: str) -> object: + """Reads one field from a dict or from an object. + + LiteLLM's replies arrive as Pydantic models on some paths and as plain dicts on + others, depending how far they have been deserialised, so every response walk here + has to handle both shapes. + """ + fields: Final = as_object(holder) + if fields is not None: + return fields.get(name) + return getattr(holder, name, None) + + +def read_list(holder: object, name: str) -> Sequence[object]: + """Reads a list field from a dict or an object; anything else reads as empty. + + The entries are the reply's own objects, so writing through them edits the reply. + """ + value: Final = read_field(holder, name) + if isinstance(value, tuple): + return value + return tuple(as_array(value) or ()) + + +def write_field(holder: object, name: str, value: object) -> None: + """Writes one string field back into a dict or an object. Pairs with read_field.""" + if isinstance(holder, dict): + holder[name] = value + else: + setattr(holder, name, value) + + +def collect_json_leaves(node: object, slots: SlotSink, *, strict: bool = False) -> None: + """Collects every string leaf of a JSON-ish structure, with a write-back per leaf. + + An Anthropic `tool_use` block carries `input`, an arbitrary JSON object rather than a + string, so a value worth restoring can sit at any depth. Bounded by `MAX_JSON_DEPTH`: + the shape is caller or model controlled, and the bound is what stops a crafted one from + becoming an unbounded descent. Walked with an explicit stack rather than recursively, + so a deeply nested value cannot spend stack frames proportional to attacker-chosen + depth. + + `strict` is for the request side, where a leaf left behind would reach the provider + unredacted: past the bound it raises `RequestTooDeep`. On the reply side a leaf past + the bound just keeps its placeholder, which leaks nothing, so it is skipped. + """ + pending: Final[list[tuple[object, int]]] = [(node, 0)] # mutable-ok: local walk stack. + while pending: + current, current_depth = pending.pop() + if current_depth > MAX_JSON_DEPTH: + if strict and is_container(current) and current: + raise RequestTooDeep("json") + continue + current_object = as_object(current) + if current_object is not None: + for key in tuple(current_object): + value = current_object[key] + if isinstance(value, str) and value: + slots.append((value, functools.partial(current_object.__setitem__, key))) + else: + pending.append((value, current_depth + 1)) + continue + entries = as_array(current) + if entries is not None: + for index, value in enumerate(entries): + if isinstance(value, str) and value: + slots.append((value, functools.partial(entries.__setitem__, index))) + else: + pending.append((value, current_depth + 1)) + + +def collect_response_item(item: object, slots: SlotSink) -> None: + """Restorable spans in one Responses API output item, dict or object. + + Mirrors `collect_responses_fields` on the request side -- a function_call or + mcp_call item holds `arguments`, their outputs `output`, a reasoning item `summary` + parts -- so the two directions stay symmetric. A custom tool call carries `input` and + a code interpreter call `code`, both model-written. + """ + for block in read_list(item, "content"): + for field in ("text", "refusal"): + text = read_field(block, field) + if isinstance(text, str) and text: + slots.append((text, lambda new, b=block, f=field: write_field(b, f, new))) + for part in read_list(item, "summary"): + text = read_field(part, "text") + if isinstance(text, str) and text: + slots.append((text, lambda new, p=part: write_field(p, "text", new))) + for field in ("arguments", "output", "input", "code"): + value = read_field(item, field) + if isinstance(value, str) and value: + slots.append((value, lambda new, i=item, f=field: write_field(i, f, new))) + + +async def rehydrate_slots(slots: Sequence[Slot], rehydrate: Rehydrate) -> None: + """Restores every span in `slots` in one batch and writes each result back.""" + if not slots: + return + restored: Final = await rehydrate(tuple(text for text, _ in slots)) + for (_, write), replacement in zip(slots, restored): + write(replacement) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/request_walk.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/request_walk.py new file mode 100644 index 00000000000..c9c0d2ccc47 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/request_walk.py @@ -0,0 +1,363 @@ +from collections.abc import Sequence +from typing import ( + Final, +) + +from .payload import ( + MAX_JSON_DEPTH, + MutableRequest, + RequestTooDeep, + Slot, + SlotSink, + as_array, + as_object, + collect, + collect_entry, + collect_json_leaves, + collect_text_parts, + is_container, + read_list, +) + +PRIVILEGED_ROLES: Final = frozenset({"system", "developer"}) + +MAX_CONTENT_DEPTH: Final = 8 + +SCHEMA_STRUCTURAL_KEYWORDS: Final = frozenset( + ( + "type", + "format", + "pattern", + "required", + "dependentRequired", + "propertyOrdering", + "discriminator", + "contentEncoding", + "contentMediaType", + "$ref", + "$id", + "$schema", + "$anchor", + "$dynamicRef", + "$dynamicAnchor", + "$recursiveRef", + "$recursiveAnchor", + "$vocabulary", + ) +) + +SCHEMA_VALUE_KEYWORDS: Final = frozenset(("examples", "default")) + +SCHEMA_LITERAL_KEYWORDS: Final = frozenset(("enum", "const")) + +SCHEMA_MAP_KEYWORDS: Final = frozenset( + ("properties", "patternProperties", "$defs", "definitions", "dependentSchemas", "dependencies") +) + + +def collect_prompt(data: MutableRequest, slots: SlotSink) -> None: + """The Completions API sends its text in `prompt`, and its tail in `suffix`.""" + collect(data, "suffix", slots) + prompt: Final = data.get("prompt") + if isinstance(prompt, str): + collect(data, "prompt", slots) + return + prompt_object: Final = as_object(prompt) + if prompt_object is not None: + variables: Final = as_object(prompt_object.get("variables")) + if variables is not None: + for name in tuple(variables): + collect(variables, name, slots) + typed = as_object(variables[name]) + if typed is not None: + collect(typed, "text", slots) + return + entries: Final = as_array(prompt) + if entries is None: + return + for index in range(len(entries)): + collect_entry(entries, index, slots) + + +def collect_content(container: MutableRequest, slots: SlotSink) -> None: + """Collects `content`, a string or a list of typed parts. + + An Anthropic tool_result nests its own content, so this has to descend. It walks + with an explicit stack and a depth bound rather than by recursion: the nesting is + caller controlled, and an unbounded descent is a JSON bomb. Content nested past the + bound raises `RequestTooDeep` rather than being skipped. + """ + pending: Final[list[tuple[MutableRequest, int]]] = [(container, 0)] # mutable-ok: local queue, never escapes. + cursor = 0 # rebind-ok: advances through the queue. + while cursor < len(pending): + node, depth = pending[cursor] + cursor += 1 + content = node.get("content") + if isinstance(content, str): + collect(node, "content", slots) + continue + if depth >= MAX_CONTENT_DEPTH and content: + raise RequestTooDeep("content") + for item in as_array(content) or (): + part = as_object(item) + if part is None: + continue + collect(part, "text", slots) + if part.get("type") == "tool_use": + collect_json_leaves(part.get("input"), slots, strict=True) + source = as_object(part.get("source")) if part.get("type") == "document" else None + if source is not None: + collect(part, "title", slots) + collect(part, "context", slots) + if source.get("type") == "text": + collect(source, "data", slots) + elif source.get("type") == "content": + pending.append((source, depth + 1)) + if "content" in part: + pending.append((part, depth + 1)) + + +def collect_participant_name(message: MutableRequest, slots: SlotSink) -> None: + """Redacts `name` where it identifies a person, never where it names a function. + + On a user or assistant turn `name` is the participant, which is personal data. + On a tool or function turn the same field carries the function's name and has + to reach the provider unchanged, or the call no longer routes. + """ + if message.get("role") in ("tool", "function"): + return + collect(message, "name", slots) + + +def collect_tool_arguments(message: MutableRequest, slots: SlotSink) -> None: + """Tool arguments carry the values a user asked the model to act on.""" + for tool_call in read_list(message, "tool_calls"): + tool_call_object = as_object(tool_call) + function = as_object(tool_call_object.get("function")) if tool_call_object is not None else None + if function is not None: + collect(function, "arguments", slots) + legacy: Final = as_object(message.get("function_call")) + if legacy is not None: + collect(legacy, "arguments", slots) + + +def collect_system(data: MutableRequest, slots: SlotSink) -> None: + """Anthropic's /v1/messages carries its system prompt at the top level.""" + system: Final = data.get("system") + if isinstance(system, str): + collect(data, "system", slots) + return + collect_text_parts(data, "system", slots) + + +def collect_responses_fields(data: MutableRequest, slots: SlotSink, privileged: SlotSink) -> None: + """The Responses API sends text outside `messages`, in `instructions` and `input`. + + `instructions` is written by the application, not by the caller, so it is + collected into the privileged sink; `input` is the caller's own text, except for + system and developer items in it, which go to the privileged sink like their Chat + counterparts. + """ + collect(data, "instructions", privileged) + request_input: Final = data.get("input") + if isinstance(request_input, str): + collect(data, "input", slots) + return + entries: Final = as_array(request_input) + if entries is None: + return + for index, entry in enumerate(entries): + if isinstance(entry, str): + collect_entry(entries, index, slots) + continue + item = as_object(entry) + if item is None: + continue + collect_content(item, privileged if item.get("role") in PRIVILEGED_ROLES else slots) + collect(item, "arguments", slots) + collect(item, "output", slots) + collect_text_parts(item, "output", slots) + collect(item, "input", slots) + collect(item, "code", slots) + collect_text_parts(item, "summary", slots) + + +def collect_tool_definitions(data: MutableRequest, slots: SlotSink, privileged: SlotSink) -> None: + """Tool definitions are application-authored free text bound for the provider. + + A tool's description and the free text in its parameter schema are where callers put + examples and customer context, so they carry PII as often as a prompt does. They are + collected into the privileged sink, like a system prompt: redacted outbound, and never + restorable from the reply. `enum` and `const` values are the exception, and go to the + caller's vault -- see `SCHEMA_LITERAL_KEYWORDS`. Names and types are left as sent. + + Covers Chat `tools[].function`, the legacy `functions[]`, and the flat tool shape the + Responses API and Anthropic share, whose schema is `parameters` or `input_schema`. + """ + for key in ("tools", "functions"): + for entry in as_array(data.get(key)) or (): + tool = as_object(entry) + if tool is None: + continue + function = as_object(tool.get("function")) + for holder in (tool, function) if function is not None else (tool,): + collect(holder, "description", privileged) + collect_schema_text(holder.get("parameters"), slots, privileged) + collect_schema_text(holder.get("input_schema"), slots, privileged) + + +def collect_schema_text(schema: object, slots: SlotSink, privileged: SlotSink) -> None: + """Collects the text in a JSON Schema, at any depth. + + Scan by default: every string is collected except under the keywords in + `SCHEMA_STRUCTURAL_KEYWORDS`, whose values must go out verbatim. A list of keywords + *to* collect would leak every one it forgot -- draft-07 `dependencies`, a `$comment`, + a vendor `x-` extension -- which is how this walk started out. Free text goes to the + privileged sink; `enum` / `const` literals go to the caller's, so the model's use of + them is restored. + + Structure matters in two places. Under `properties` and the other name -> subschema + maps, keys are property names rather than keywords, so a property called `type` is a + subschema to walk, not a keyword to skip. And `examples` / `default` hold JSON values, + so all their strings are collected whatever the keys around them are called. Nested + past `MAX_JSON_DEPTH`, the request is refused. + """ + pending: Final[list[tuple[object, int]]] = [(schema, 0)] # mutable-ok: local walk stack. + while pending: + node, depth = pending.pop() + if depth > MAX_JSON_DEPTH: + if is_container(node) and node: + raise RequestTooDeep("schema") + continue + entries = as_array(node) + if entries is not None: + for index, item in enumerate(entries): + collect_entry(entries, index, privileged) + if is_container(item): + pending.append((item, depth + 1)) + continue + schema_object = as_object(node) + if schema_object is None: + continue + for keyword, value in tuple(schema_object.items()): + if keyword in SCHEMA_STRUCTURAL_KEYWORDS: + continue + subschemas = as_object(value) if keyword in SCHEMA_MAP_KEYWORDS else None + if keyword in SCHEMA_LITERAL_KEYWORDS: + collect(schema_object, keyword, slots) + collect_json_leaves(value, slots, strict=True) + elif keyword in SCHEMA_VALUE_KEYWORDS: + collect(schema_object, keyword, privileged) + collect_json_leaves(value, privileged, strict=True) + elif subschemas is not None: + pending.extend((child, depth + 1) for child in subschemas.values()) + elif isinstance(value, str): + collect(schema_object, keyword, privileged) + elif is_container(value): + pending.append((value, depth + 1)) + + +def collect_output_contracts(data: MutableRequest, slots: SlotSink, privileged: SlotSink) -> None: + """Text the caller sends to shape the reply rather than to prompt it. + + A predicted output (`prediction.content`) is the caller's own draft of the answer, so + it goes with their text: the model largely repeats it, and it has to come back. A + structured-output schema -- Chat `response_format.json_schema`, Responses + `text.format` -- is application-authored like a tool schema, so its free text goes + to the privileged sink, and its names and types stay as sent. + """ + prediction: Final = as_object(data.get("prediction")) + if prediction is not None: + collect(prediction, "content", slots) + collect_text_parts(prediction, "content", slots) + response_format: Final = as_object(data.get("response_format")) + text_options: Final = as_object(data.get("text")) + for declared in ( + response_format.get("json_schema") if response_format is not None else None, + text_options.get("format") if text_options is not None else None, + ): + wrapper = as_object(declared) + if wrapper is not None: + collect(wrapper, "description", privileged) + collect_schema_text(wrapper.get("schema"), slots, privileged) + + +def collect_user_locations(data: MutableRequest, privileged: SlotSink) -> None: + """Web search forwards the user's approximate location, whose `city` and `region` + are free text and can hold a street address. + + Chat carries it in `web_search_options.user_location.approximate`; the Responses + and Anthropic web-search tools carry it flat on the tool's `user_location`. Nothing + restores it from a reply, hence the privileged sink. + """ + options: Final = data.get("web_search_options") + tools: Final = as_array(data.get("tools")) or () + for declared in (options, *tools): + holder = as_object(declared) + location = as_object(holder.get("user_location")) if holder is not None else None + if location is None: + continue + approximate = as_object(location.get("approximate")) + for container in (location, approximate) if approximate is not None else (location,): + collect(container, "city", privileged) + collect(container, "region", privileged) + + +def collect_end_user_ids(data: MutableRequest, privileged: SlotSink) -> None: + """`user` and `safety_identifier` are forwarded to the provider and often hold an email. + + Only detected PII is replaced, so an opaque id reaches the provider unchanged. LiteLLM's + own end-user spend tracking reads the id resolved at authentication, before this hook + runs, so rewriting the field here does not move spend. Nothing restores these from a + reply, hence the privileged sink. + """ + collect(data, "user", privileged) + collect(data, "safety_identifier", privileged) + + +def locate_request_texts( + data: MutableRequest, +) -> tuple[Sequence[Slot], Sequence[Slot]]: + """Finds every redactable span, split by whether the caller can see it. + + Anything missed here reaches the provider in the clear while the guardrail + still reports as enabled, so the walk covers every request shape that + carries text. + + The split exists because the response is restored against one vault only. + Server-authored spans -- system and developer turns, Anthropic's top-level + `system`, the Responses API `instructions`, tool and output schemas -- go into a + vault nothing is ever restored against, so a caller who gets the model to + echo one of their placeholders back receives the placeholder, not the value + behind it. End-user identifiers go there too: nothing in a reply needs them. + + Tool *results* stay on the caller's side deliberately. The model reads them in + order to answer, so it can already repeat anything in them; restoring the + placeholder gives the caller the answer they would have had without this + guardrail, and an agent that reads a file and quotes an address from it needs + that address back. + + `extra_body` is walked the same way as the request itself. LiteLLM merges it over + the transformed request just before sending, so a field there -- `input`, + `messages`, `system` -- replaces the redacted one on the wire. + """ + slots: Final[SlotSink] = [] + privileged: Final[SlotSink] = [] + for payload in (data, as_object(data.get("extra_body"))): + if payload is None: + continue + for entry in read_list(payload, "messages"): + message = as_object(entry) + if message is not None: + sink = privileged if message.get("role") in PRIVILEGED_ROLES else slots + collect_content(message, sink) + collect_participant_name(message, sink) + collect_tool_arguments(message, sink) + collect_responses_fields(payload, slots, privileged) + collect_prompt(payload, slots) + collect_system(payload, privileged) + collect_tool_definitions(payload, slots, privileged) + collect_output_contracts(payload, slots, privileged) + collect_user_locations(payload, privileged) + collect_end_user_ids(payload, privileged) + return tuple(slots), tuple(privileged) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/stream_restorers.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/stream_restorers.py new file mode 100644 index 00000000000..4ef8bbdbb3e --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/stream_restorers.py @@ -0,0 +1,345 @@ +import copy +import functools +import itertools +import json +import re +from enum import Enum +from types import MappingProxyType +from typing import ( + Final, + TypeAlias, +) + +from .payload import ( + JsonBody, + MutableRequest, + Rehydrate, + SlotSink, + StreamStep, + as_object, + collect_response_item, + read_field, + read_list, + rehydrate_slots, + write_field, +) + +ANTHROPIC_DELTA_FIELDS: Final = MappingProxyType({"text_delta": "text", "input_json_delta": "partial_json"}) + +SSE_EVENT_BOUNDARY: Final = re.compile(rb"(\r?\n\r?\n)") + +SSE_OPENINGS: Final = (b"event:", b"data:", b"id:", b"retry:", b":") + +RESPONSES_BINARY_DELTAS: Final = frozenset(("response.audio.delta",)) + +RESPONSES_STRUCTURAL_FIELDS: Final = frozenset( + ("type", "id", "item_id", "call_id", "name", "server_label", "status", "obfuscation") +) + +RESPONSES_TERMINAL_EVENTS: Final = frozenset(("response.completed", "response.incomplete")) + +CarryKey: TypeAlias = tuple[int, int | None] +CarryWindows: TypeAlias = dict[CarryKey, str] + +ResponsesStreamKey: TypeAlias = tuple[str, object, object, object] + + +def carry_sort_key(key: CarryKey) -> tuple[int, int]: + """Orders streaming windows without ever comparing None to an int. + + `sorted()` over the raw keys raises as soon as one choice holds both a content window + and a tool-call window, because `None < 0` is not orderable. Content sorts first, then + tool calls by their index. + """ + choice_index, tool_index = key + return (choice_index, -1 if tool_index is None else tool_index) + + +def continuation_delta(tool_index: int, text: str) -> list[dict[str, object]]: + """A `tool_calls` delta carrying `text` as an index-only continuation. + + Clients concatenate tool-call fragments by index, so no id or name is needed. + """ + return [{"index": tool_index, "function": {"arguments": text}}] + + +def opens_like_sse(head: bytes) -> bool | None: + """Whether a raw stream is SSE, judged by its opening bytes; None while undecidable. + + An SSE stream opens with a field name or a `:` comment. Anything else -- a JSON array + streamed in pieces, say -- has no event boundaries to wait for. A chunk that ends + partway through a field name decides nothing yet, so that case waits for more. + """ + opening: Final = head.lstrip() + if not opening: + return None + if opening.startswith(SSE_OPENINGS): + return True + if any(field.startswith(opening) for field in SSE_OPENINGS): + return None + return False + + +def responses_event_type(chunk: object) -> str | None: + """The event type of a Responses API stream event, or None for any other chunk. + + The type arrives as a plain string on dicts and as a str-valued Enum on LiteLLM's + event models. The Enum is unwrapped because it does not hash like its value, so it + would miss every lookup in the event tables above. + """ + if isinstance(chunk, (bytes, str)): + return None + kind: Final = read_field(chunk, "type") + value: Final = kind.value if isinstance(kind, Enum) else kind + return value if isinstance(value, str) and value.startswith("response.") else None + + +class AnthropicSSERestorer: + """Restores an Anthropic `/v1/messages` stream, which reaches the hook as raw SSE. + + Each content block is its own token stream with its own window, keyed by the block's + `index`: `text_delta` carries prose and `input_json_delta` a tool call's arguments. + When a block stops, whatever its window still holds is emitted as one more delta for + that block, just ahead of the `content_block_stop` frame, so the client has the whole + block before it is told the block is complete. + + Frames are processed whole. A network chunk can end in the middle of an event, so the + unfinished tail is kept until the rest arrives; that delays one partial event, never + a completed one. A frame that is not an Anthropic event -- another endpoint's SSE, or + anything that fails to parse -- is passed through byte for byte, and a raw stream that + does not open like SSE at all is passed through chunk by chunk, never buffered. + """ + + def __init__(self, step: StreamStep) -> None: + self._step: Final = step + self._carries: Final[dict[int, str]] = {} # mutable-ok: per-block windows advanced in place. + self._delta_types: Final[dict[int, str]] = {} # mutable-ok: each block's delta type, for its flush. + self._pending = b"" + self._as_text = False + self._is_sse: bool | None = None + + async def feed(self, chunk: bytes | str) -> tuple[bytes | str, ...]: + """Restores every event this chunk completes; holds back an unfinished tail.""" + if isinstance(chunk, str): + self._as_text = True + if self._is_sse is False: + return (chunk,) + raw: Final = chunk.encode("utf-8") if isinstance(chunk, str) else chunk + buffered: Final = self._pending + raw + if self._is_sse is None: + self._is_sse = opens_like_sse(buffered) + if self._is_sse is None: + self._pending = buffered + return () + if not self._is_sse: + self._pending = b"" + return self._emit(buffered) + boundaries: Final = tuple(SSE_EVENT_BOUNDARY.finditer(buffered)) + if not boundaries: + self._pending = buffered + return () + cut: Final = boundaries[-1].end() + self._pending = buffered[cut:] + parts: Final = SSE_EVENT_BOUNDARY.split(buffered[:cut]) + restored: Final = tuple( + [await self._restore_event(parts[index]) + parts[index + 1] for index in range(0, len(parts) - 1, 2)] + ) + return self._emit(b"".join(restored)) + + async def finish(self) -> tuple[bytes | str, ...]: + """Emits an unterminated final event and any window a block never closed.""" + held: Final = self._pending + self._pending = b"" + if not self._is_sse: + return self._emit(held) + tail: Final = await self._restore_event(held) if held.strip() else held + flushed: Final = await self._flush_all() + separator: Final = b"\n\n" if tail.strip() and flushed else b"" + return self._emit(tail + separator + flushed) + + def _emit(self, frames: bytes) -> tuple[bytes | str, ...]: + if not frames: + return () + return (frames.decode("utf-8") if self._as_text else frames,) + + async def _restore_event(self, block: bytes) -> bytes: + """Rewrites one SSE event, or returns it untouched if it carries nothing to restore.""" + try: + lines: Final = block.decode("utf-8").split("\n") + except UnicodeDecodeError: + return block + data_lines: Final = tuple(index for index, line in enumerate(lines) if line.startswith("data:")) + if len(data_lines) != 1: + return block + line: Final = lines[data_lines[0]] + try: + parsed: Final[object] = json.loads(line[len("data:") :]) + except ValueError: + return block + event: Final = as_object(parsed) + if event is None: + return block + kind: Final = event.get("type") + index: Final = event.get("index") + if kind == "content_block_stop" and isinstance(index, int): + return await self._flush(index) + block + if kind == "message_stop": + return await self._flush_all() + block + if kind != "content_block_delta" or not await self._restore_delta(event): + return block + ending: Final = "\r" if line.endswith("\r") else "" + rewritten: Final = ( + *lines[: data_lines[0]], + f"data: {json.dumps(event, ensure_ascii=False)}{ending}", + *lines[data_lines[0] + 1 :], + ) + return "\n".join(rewritten).encode("utf-8") + + async def _restore_delta(self, event: MutableRequest) -> bool: + """Advances one block's window through this delta. False if it holds no text.""" + index: Final = event.get("index") + delta: Final = event.get("delta") + if not isinstance(index, int) or not isinstance(delta, dict): + return False + delta_type: Final = delta.get("type") + if not isinstance(delta_type, str): + return False + field: Final = ANTHROPIC_DELTA_FIELDS.get(delta_type) + text: Final = delta.get(field) if field is not None else None + if field is None or not isinstance(text, str) or not text: + return False + emitted, remaining = await self._step(text, self._carries.get(index, ""), False) + self._carries[index] = remaining + self._delta_types[index] = delta_type + delta[field] = emitted + return True + + async def _flush(self, index: int) -> bytes: + """One synthetic delta frame carrying whatever `index`'s window still holds.""" + carry: Final = self._carries.pop(index, "") + delta_type: Final = self._delta_types.pop(index, None) + field: Final = ANTHROPIC_DELTA_FIELDS.get(delta_type) if isinstance(delta_type, str) else None + if not carry or field is None: + return b"" + text, _ = await self._step("", carry, True) + if not text: + return b"" + event: Final[JsonBody] = { + "type": "content_block_delta", + "index": index, + "delta": {"type": delta_type, field: text}, + } + return f"event: content_block_delta\ndata: {json.dumps(event, ensure_ascii=False)}\n\n".encode() + + async def _flush_all(self) -> bytes: + flushed: Final = tuple([await self._flush(index) for index in tuple(self._carries)]) + return b"".join(flushed) + + +class ResponsesStreamRestorer: + """Restores a Responses API event stream. + + The event families are matched by shape rather than listed, so a text stream the + API adds later is restored by default instead of leaking a placeholder: + + - Any `*.delta` event whose `delta` is a string is a token stream (output_text, + refusal, function-call and MCP arguments, reasoning summaries, ...). Each gets its + own window, keyed by the family, the item id and the part index. + - Any `*.done` event closes the stream of the same family. Whatever its window still + holds goes out first, as a copy of that stream's last delta event -- so it carries + the stream's own ids, and repeats that event's `sequence_number`. Then every text + field on the done event is restored in full: its string fields other than + identifiers, plus any `part` or `item` it repeats. + - `response.completed` / `response.incomplete` repeat the whole reply, and are + restored the same way the non-streaming reply is. + """ + + def __init__(self, step: StreamStep, rehydrate: Rehydrate) -> None: + self._step: Final = step + self._rehydrate: Final = rehydrate + self._carries: Final[dict[ResponsesStreamKey, str]] = {} # mutable-ok: per-stream windows advanced in place. + self._last_deltas: Final[dict[ResponsesStreamKey, object]] = {} # mutable-ok: newest delta per stream. + + async def restore(self, event: object) -> tuple[object, ...]: + """The events to emit in place of `event`: any flush, then the event itself.""" + kind: Final = responses_event_type(event) + if kind is None: + return (event,) + if kind.endswith(".delta") and kind not in RESPONSES_BINARY_DELTAS: + await self._restore_delta(event, kind) + return (event,) + slots: Final[SlotSink] = [] + flushed: Final = await self._flush(responses_stream_key(event, kind)) if kind.endswith(".done") else () + if kind.endswith(".done"): + collect_event_text(event, slots) + part: Final = read_field(event, "part") + if part is not None: + collect_response_item({"content": [part]}, slots) + collect_response_item(read_field(event, "item"), slots) + elif kind in RESPONSES_TERMINAL_EVENTS: + for item in read_list(read_field(event, "response"), "output"): + collect_response_item(item, slots) + await rehydrate_slots(slots, self._rehydrate) + return (*flushed, event) + + async def finish(self) -> tuple[object, ...]: + """Flushes every stream the provider never closed, e.g. a truncated reply.""" + flushed: Final = tuple([await self._flush(key) for key in tuple(self._carries)]) + return tuple(itertools.chain.from_iterable(flushed)) + + async def _restore_delta(self, event: object, kind: str) -> None: + text: Final = read_field(event, "delta") + if not isinstance(text, str) or not text: + return + key: Final = responses_stream_key(event, kind) + emitted, remaining = await self._step(text, self._carries.get(key, ""), False) + self._carries[key] = remaining + self._last_deltas[key] = event + write_field(event, "delta", emitted) + + async def _flush(self, key: ResponsesStreamKey) -> tuple[object, ...]: + carry: Final = self._carries.pop(key, "") + template: Final = self._last_deltas.pop(key, None) + if not carry or template is None: + return () + text, _ = await self._step("", carry, True) + if not text: + return () + flush: Final = copy.deepcopy(template) + write_field(flush, "delta", text) + return (flush,) + + +def responses_stream_key(event: object, kind: str) -> ResponsesStreamKey: + """Identifies the delta stream an event belongs to, the same for its delta and done. + + The family is the event type without its `.delta` / `.done` suffix, so an output_text + stream and a refusal stream on the same part never share a window. + """ + family: Final = kind.rsplit(".", 1)[0] + part_index: Final = read_field(event, "content_index") + summary_index: Final = read_field(event, "summary_index") + return ( + family, + read_field(event, "item_id"), + read_field(event, "output_index"), + part_index if part_index is not None else summary_index, + ) + + +def collect_event_text(event: object, slots: SlotSink) -> None: + """Collects every top-level text field of a Responses API event, dict or model. + + Scan by default, with identifiers excluded, rather than a list of known fields: the + `.done` event of each stream family names its text differently (`text`, `refusal`, + `arguments`, ...), and a family added upstream would otherwise leak a placeholder. + """ + attributes: Final[object] = getattr(event, "__dict__", None) + fields: Final = as_object(event) or as_object(attributes) + if fields is None: + return + for name, value in tuple(fields.items()): + if name in RESPONSES_STRUCTURAL_FIELDS or name.endswith("_id"): + continue + if isinstance(value, str) and value: + slots.append((value, functools.partial(write_field, event, name))) diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 46026c12d24..8f1140cef62 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -145,6 +145,7 @@ class SupportedGuardrailIntegrations(Enum): STRAIKER = "straiker" ALICE = "alice" AGENT_365 = "agent_365" + LLM_SHIELD_PROXY = "llm_shield_proxy" CONDUCT = "conduct" diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/llm_shield_proxy.py b/litellm/types/proxy/guardrails/guardrail_hooks/llm_shield_proxy.py new file mode 100644 index 00000000000..967d7ee75c2 --- /dev/null +++ b/litellm/types/proxy/guardrails/guardrail_hooks/llm_shield_proxy.py @@ -0,0 +1,24 @@ +from pydantic import Field + +from .base import GuardrailConfigModel + + +class LLMShieldProxyGuardrailConfigModel(GuardrailConfigModel): + api_key: str | None = Field( + default=None, + description=( + "The virtual key for the LLM Shield Proxy instance. If not provided, the " + "`LLM_SHIELD_PROXY_API_KEY` environment variable is checked." + ), + ) + api_base: str | None = Field( + default=None, + description=( + "The base URL of the LLM Shield Proxy instance. If not provided, the `LLM_SHIELD_PROXY_API_BASE` " + "environment variable is checked, then `http://localhost:8000`." + ), + ) + + @staticmethod + def ui_friendly_name() -> str: + return "LLM Shield Proxy" diff --git a/ruff-strict.toml b/ruff-strict.toml index 899a8ff3af5..9e0b11a1af4 100644 --- a/ruff-strict.toml +++ b/ruff-strict.toml @@ -30,6 +30,10 @@ external = [ # grows over time; typing it concretely (`object`) broke that forwarding call outright — # basedpyright turned every named param into a reportArgumentType error. Any is correct here. "litellm/proxy/guardrails/guardrail_hooks/alice/alice.py" = ["ANN401"] +# Same reason: `**kwargs` forwards verbatim to CustomGuardrail.__init__, and the lifecycle +# hook signatures inherit `Any` for `response` from CustomLogger, so narrowing them here +# would break the override rather than describe it. +"litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py" = ["ANN401"] [lint.mccabe] max-complexity = 15 diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py b/tests/unit/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py new file mode 100644 index 00000000000..ed34c07022e --- /dev/null +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -0,0 +1,2003 @@ +import asyncio +import json +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest +from httpx import Request, Response + +import litellm +from litellm.caching.caching_handler import _PENDING_CACHE_WRITES +from litellm.caching.in_memory_cache import InMemoryCache +from litellm.exceptions import GuardrailRaisedException +from litellm.proxy.guardrails.guardrail_hooks.llm_shield_proxy.llm_shield_proxy import ( + GUARDRAIL_NAME, + LLMShieldProxyGuardrail, +) +from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2 +from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.llms.openai import ( + FunctionCallArgumentsDeltaEvent, + OutputTextDeltaEvent, + OutputTextDoneEvent, + ResponsesAPIStreamEvents, +) +from litellm.types.utils import ( + Choices, + Delta, + Message, + ModelResponse, + ModelResponseStream, + StreamingChoices, + TextChoices, + TextCompletionResponse, + Usage, +) + + +def _guardrail(**overrides: object) -> LLMShieldProxyGuardrail: + params: dict[str, object] = { + "api_key": "test-key", + "api_base": "http://shield.test", + "guardrail_name": GUARDRAIL_NAME, + "event_hook": "pre_call", + "default_on": True, + } + params.update(overrides) + return LLMShieldProxyGuardrail(**params) + + +def _response(payload: dict, status_code: int = 200) -> Response: + return Response( + status_code=status_code, + json=payload, + request=Request("POST", "http://shield.test/v1/guard/redact"), + ) + + +def _mock_post(guardrail: LLMShieldProxyGuardrail, *payloads: dict) -> AsyncMock: + """Queues one shield response per expected call.""" + mock = AsyncMock(side_effect=[_response(p) for p in payloads]) + guardrail.async_handler.post = mock # type: ignore[method-assign] + return mock + + +def _chunk(content: str | None, finish_reason: str | None = None) -> ModelResponseStream: + return ModelResponseStream( + choices=[StreamingChoices(index=0, delta=Delta(content=content), finish_reason=finish_reason)] + ) + + +async def _drain(generator) -> list: + return [chunk async for chunk in generator] + + +def _tool_chunk(arguments: str, finish_reason: str | None = None) -> ModelResponseStream: + """One streamed fragment of tool call 0's arguments.""" + tool_call = {"index": 0, "id": "call_1", "type": "function", "function": {"name": "send", "arguments": arguments}} + return ModelResponseStream( + choices=[StreamingChoices(index=0, delta=Delta(tool_calls=[tool_call]), finish_reason=finish_reason)] + ) + + +def _field(holder: object, name: str) -> object: + """Reads a field from a dict or a model; the guardrail emits both shapes.""" + return holder.get(name) if isinstance(holder, dict) else getattr(holder, name) + + +class _FakeShield: + """The three guard endpoints over one fixed vault, placeholder -> original. + + The stream endpoint holds back a trailing `[` that has not closed yet, which is the + behaviour that makes a placeholder split across two chunks come out whole. + """ + + def __init__(self, vault: dict[str, str]) -> None: + self.vault = vault + self.urls: list[str] = [] + + def _restore(self, text: str) -> str: + for placeholder, original in self.vault.items(): + text = text.replace(placeholder, original) + return text + + async def post(self, url: str, headers: dict, json: dict, timeout: float) -> Response: + self.urls.append(url) + if url.endswith("/rehydrate/stream"): + text = self._restore(json["carry"] + json["text"]) + opening = text.rfind("[") + if json["final"] or opening == -1 or "]" in text[opening:]: + return _response({"text": text, "carry": ""}) + return _response({"text": text[:opening], "carry": text[opening:]}) + return _response({"texts": [self._restore(text) for text in json["texts"]]}) + + +def _shielded(vault: dict[str, str]) -> tuple[LLMShieldProxyGuardrail, _FakeShield]: + guardrail = _guardrail(event_hook="post_call") + shield = _FakeShield(vault) + guardrail.async_handler.post = shield.post # type: ignore[method-assign] + return guardrail, shield + + +def _sse(event: dict) -> bytes: + return f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() + + +def _sse_events(frames: list) -> list[dict]: + """Parses emitted SSE output, whatever its chunking, back into event payloads.""" + raw = b"".join(frame.encode() if isinstance(frame, str) else frame for frame in frames).decode() + return [ + json.loads(line[len("data:") :]) + for event in raw.split("\n\n") + for line in event.split("\n") + if line.startswith("data:") + ] + + +def _text_block_stream(*deltas: str) -> list[bytes]: + """An Anthropic /v1/messages stream with one text block made of `deltas`.""" + return [ + _sse({"type": "message_start", "message": {"id": "msg_1", "role": "assistant", "content": []}}), + _sse({"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}), + *( + _sse({"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": d}}) + for d in deltas + ), + _sse({"type": "content_block_stop", "index": 0}), + _sse({"type": "message_stop"}), + ] + + +async def _restore_stream(guardrail: LLMShieldProxyGuardrail, chunks: list) -> list: + async def stream(): + for chunk in chunks: + yield chunk + + return await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + +def test_llm_shield_guardrail_config(monkeypatch: pytest.MonkeyPatch): + """Should register through init_guardrails_v2 like any other provider.""" + monkeypatch.setattr(litellm, "guardrail_name_config_map", {}) + monkeypatch.setenv("LLM_SHIELD_PROXY_API_KEY", "test-key") + + init_guardrails_v2( + all_guardrails=[ + { + "guardrail_name": "llm_shield_proxy", + "litellm_params": {"guardrail": "llm_shield_proxy", "mode": "pre_call", "default_on": True}, + } + ], + config_file_path="", + ) + + registered = [cb for cb in litellm.callbacks if isinstance(cb, LLMShieldProxyGuardrail)] + assert len(registered) == 1 + assert registered[0].guardrail_name == "llm_shield_proxy" + + +class TestLLMShieldProxyInitialization: + def test_api_base_defaults_to_localhost(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("LLM_SHIELD_PROXY_API_BASE", raising=False) + assert _guardrail(api_base=None).api_base == "http://localhost:8000" + + def test_api_base_reads_environment(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LLM_SHIELD_PROXY_API_BASE", "http://shield.internal:9000") + assert _guardrail(api_base=None).api_base == "http://shield.internal:9000" + + def test_trailing_slash_is_stripped(self): + assert _guardrail(api_base="http://shield.test/").api_base == "http://shield.test" + + def test_both_modes_can_be_enabled_on_one_entry(self): + """Redaction and restoration are two halves of one config entry. + + A deployment that lists only pre_call would redact the request and then hand + the placeholders straight back to the end user. + """ + guardrail = _guardrail(event_hook=["pre_call", "post_call"]) + data: dict = {"messages": []} + + assert guardrail.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_call) is True + assert guardrail.should_run_guardrail(data=data, event_type=GuardrailEventHooks.post_call) is True + assert guardrail.should_run_guardrail(data=data, event_type=GuardrailEventHooks.during_call) is False + + +class TestRedaction: + @pytest.mark.asyncio + async def test_string_content_is_redacted(self): + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["Email [EMAIL_1] about it"]}) + + data = {"messages": [{"role": "user", "content": "Email a@b.com about it"}]} + result = await guardrail.async_pre_call_hook( + user_api_key_dict=None, cache=None, data=data, call_type="completion" + ) + + assert result["messages"][0]["content"] == "Email [EMAIL_1] about it" + + @pytest.mark.asyncio + async def test_multimodal_text_parts_are_redacted(self): + """The list content shape is a historical bypass; text parts must be covered.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["call [PHONE_1]"]}) + + data = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "call 555-0100"}, + {"type": "image_url", "image_url": {"url": "http://x/y.png"}}, + ], + } + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert data["messages"][0]["content"][0]["text"] == "call [PHONE_1]" + assert data["messages"][0]["content"][1]["image_url"]["url"] == "http://x/y.png" + + @pytest.mark.asyncio + async def test_request_without_text_is_untouched(self): + """No text to redact means no call to LLM Shield Proxy. + + This deliberately uses a request with no caller text at all. An earlier + version used a Responses-API `input`, which asserted the very bypass that + let `input` reach the provider unredacted. + """ + guardrail = _guardrail() + mock = _mock_post(guardrail) + data = {"model": "gpt-4o", "temperature": 0.2} + + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + mock.assert_not_called() + + @pytest.mark.asyncio + async def test_session_id_is_reused_across_hooks(self): + """Rehydration can only resolve tokens minted under the same session.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}, {"texts": ["a@b.com"]}) + + data = {"messages": [{"role": "user", "content": "a@b.com"}]} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + await guardrail._rehydrate(["[EMAIL_1]"], guardrail._session_id(data)) + + sessions = {call.kwargs["headers"]["X-Session-ID"] for call in mock.call_args_list} + assert len(sessions) == 1 + + +class TestRequestCoverage: + """Every request shape that carries caller text must be redacted. + + A shape missed here is not a cosmetic gap: the guardrail reports as enabled + while the raw value goes to the provider. + """ + + @pytest.mark.asyncio + async def test_responses_api_string_input_is_redacted(self): + """Measured against a live provider: `input` reached the model unredacted.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["Email [EMAIL_1] the invoice"]}) + + data = {"input": "Email jane.doe@example.com the invoice"} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["Email jane.doe@example.com the invoice"] + assert data["input"] == "Email [EMAIL_1] the invoice" + + @pytest.mark.asyncio + async def test_responses_api_list_input_is_redacted(self): + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]", "[PHONE_1]"]}) + + data = { + "input": [ + {"role": "user", "content": "jane.doe@example.com"}, + {"role": "user", "content": [{"type": "input_text", "text": "555-0100"}]}, + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert data["input"][0]["content"] == "[EMAIL_1]" + assert data["input"][1]["content"][0]["text"] == "[PHONE_1]" + + @pytest.mark.asyncio + async def test_tool_call_arguments_are_redacted(self): + """Tool arguments carry the values the user asked the model to act on.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ['{"email": "[EMAIL_1]"}']}) + + data = { + "messages": [ + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "send", "arguments": '{"email": "jane.doe@example.com"}'}, + } + ], + } + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert data["messages"][0]["tool_calls"][0]["function"]["arguments"] == '{"email": "[EMAIL_1]"}' + + @pytest.mark.asyncio + async def test_responses_api_instructions_are_redacted(self): + """`instructions` is provider-bound text that sits outside `messages`.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["contact [EMAIL_1]"]}) + + data = {"instructions": "contact jane.doe@example.com", "input": ""} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["contact jane.doe@example.com"] + assert data["instructions"] == "contact [EMAIL_1]" + + @pytest.mark.asyncio + async def test_legacy_function_call_arguments_are_redacted(self): + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ['{"email": "[EMAIL_1]"}']}) + + data = { + "messages": [ + { + "role": "assistant", + "function_call": {"name": "send", "arguments": '{"email": "jane.doe@example.com"}'}, + } + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert data["messages"][0]["function_call"]["arguments"] == '{"email": "[EMAIL_1]"}' + + @pytest.mark.asyncio + async def test_completions_prompt_is_redacted(self): + """/v1/completions puts its text in a top-level `prompt`, not in messages.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["Email [EMAIL_1]"]}) + + data = {"prompt": "Email jane.doe@example.com"} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="atext_completion") + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["Email jane.doe@example.com"] + assert data["prompt"] == "Email [EMAIL_1]" + + @pytest.mark.asyncio + async def test_completions_prompt_array_is_redacted(self): + """`prompt` also accepts an array, and each entry is provider-bound.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]", "[PHONE_1]"]}) + + data = {"prompt": ["jane.doe@example.com", "555-0100"]} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="atext_completion") + + assert data["prompt"] == ["[EMAIL_1]", "[PHONE_1]"] + + @pytest.mark.asyncio + async def test_responses_function_call_items_are_redacted(self): + """Responses input items hold tool data in `arguments` and `output`.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ['{"email": "[EMAIL_1]"}', "sent to [EMAIL_1]"]}) + + data = { + "input": [ + {"type": "function_call", "name": "send", "arguments": '{"email": "jane.doe@example.com"}'}, + {"type": "function_call_output", "call_id": "c1", "output": "sent to jane.doe@example.com"}, + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert data["input"][0]["arguments"] == '{"email": "[EMAIL_1]"}' + assert data["input"][1]["output"] == "sent to [EMAIL_1]" + + @pytest.mark.asyncio + async def test_responses_tool_output_parts_are_redacted(self): + """A function_call_output can carry its result as a list of input_text parts.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["sent to [EMAIL_1]"]}) + + data = { + "input": [ + { + "type": "function_call_output", + "call_id": "c1", + "output": [{"type": "input_text", "text": "sent to jane.doe@example.com"}], + }, + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["sent to jane.doe@example.com"] + assert data["input"][0]["output"][0]["text"] == "sent to [EMAIL_1]" + + @pytest.mark.asyncio + async def test_responses_custom_tool_call_input_is_redacted(self): + """A replayed custom_tool_call carries its payload in `input`, not `arguments`.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["email [EMAIL_1]"]}) + + data = { + "input": [ + {"type": "custom_tool_call", "call_id": "c1", "name": "mail", "input": "email jane.doe@example.com"}, + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert data["input"][0]["input"] == "email [EMAIL_1]" + assert data["input"][0]["name"] == "mail" + + @pytest.mark.asyncio + async def test_extra_body_overrides_are_redacted(self): + """LiteLLM merges `extra_body` over the request just before sending, so its fields win on the wire.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["safe", "Mail [EMAIL_1]", "[EMAIL_1]"]}) + + data = { + "model": "gpt-4o", + "input": "safe", + "extra_body": { + "input": "alice@example.com", + "messages": [{"role": "user", "content": "Mail alice@example.com"}], + "service_tier": "flex", + }, + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["safe", "Mail alice@example.com", "alice@example.com"] + assert data["extra_body"]["input"] == "[EMAIL_1]" + assert data["extra_body"]["messages"][0]["content"] == "Mail [EMAIL_1]" + assert data["extra_body"]["service_tier"] == "flex" + + def test_extra_body_system_text_is_privileged(self) -> None: + """An application-authored override is no more restorable than the field it replaces.""" + data = {"messages": [{"role": "user", "content": "U"}], "extra_body": {"system": "S", "instructions": "I"}} + + caller, privileged = LLMShieldProxyGuardrail._locate_request_texts(data) + + assert [text for text, _ in caller] == ["U"] + assert sorted(text for text, _ in privileged) == ["I", "S"] + + @pytest.mark.asyncio + async def test_anthropic_text_documents_are_redacted(self): + """A document block carries text inline, as `source.data` or as `source.content`.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[NAME_1] notes", "[EMAIL_1]", "Reach [EMAIL_2]"]}) + + data = { + "messages": [ + { + "role": "user", + "content": [ + { + "type": "document", + "title": "Jane Doe notes", + "source": {"type": "text", "media_type": "text/plain", "data": "alice@example.com"}, + }, + { + "type": "document", + "source": {"type": "content", "content": [{"type": "text", "text": "Reach bob@example.com"}]}, + }, + {"type": "document", "source": {"type": "base64", "media_type": "application/pdf", "data": "JVBERi0x"}}, + ], + } + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="anthropic_messages") + + sent = mock.call_args_list[0].kwargs["json"]["texts"] + assert sent == ["Jane Doe notes", "alice@example.com", "Reach bob@example.com"] + blocks = data["messages"][0]["content"] + assert blocks[0]["source"]["data"] == "[EMAIL_1]" + assert blocks[1]["source"]["content"][0]["text"] == "Reach [EMAIL_2]" + assert blocks[2]["source"]["data"] == "JVBERi0x", "binary sources are not text" + + @pytest.mark.asyncio + async def test_responses_code_interpreter_code_is_redacted(self): + """A replayed code_interpreter_call carries the code the model wrote, which the reply side restores.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["send('[EMAIL_1]')"]}) + + data = {"input": [{"type": "code_interpreter_call", "id": "ci_1", "code": "send('jane.doe@example.com')"}]} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert data["input"][0]["code"] == "send('[EMAIL_1]')" + + @pytest.mark.asyncio + async def test_anthropic_system_prompt_is_redacted(self): + """/v1/messages carries its system prompt at the top level, not in messages.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["the user is [EMAIL_1]"]}) + + data = {"system": "the user is jane.doe@example.com", "messages": []} + await guardrail.async_pre_call_hook( + user_api_key_dict=None, cache=None, data=data, call_type="anthropic_messages" + ) + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["the user is jane.doe@example.com"] + assert data["system"] == "the user is [EMAIL_1]" + + @pytest.mark.asyncio + async def test_anthropic_system_blocks_are_redacted(self): + """`system` also accepts a list of text blocks.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + + data = {"system": [{"type": "text", "text": "jane.doe@example.com"}], "messages": []} + await guardrail.async_pre_call_hook( + user_api_key_dict=None, cache=None, data=data, call_type="anthropic_messages" + ) + + assert data["system"][0]["text"] == "[EMAIL_1]" + + @pytest.mark.asyncio + async def test_string_array_input_is_redacted(self): + """Embeddings and moderations send `input` as an array of bare strings.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]", "[PHONE_1]"]}) + + data = {"input": ["jane.doe@example.com", "555-0100"]} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aembedding") + + assert data["input"] == ["[EMAIL_1]", "[PHONE_1]"] + + @pytest.mark.asyncio + async def test_participant_name_is_redacted(self): + """`name` on a user turn identifies a person.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["hi", "[PERSON_1]"]}) + + data = {"messages": [{"role": "user", "name": "Jane Doe", "content": "hi"}]} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["hi", "Jane Doe"] + assert data["messages"][0]["name"] == "[PERSON_1]" + + @pytest.mark.asyncio + async def test_tool_function_name_is_left_alone(self): + """On a tool turn the same field is the function name. + + Redacting it would stop the call routing, so this asserts it is never sent + to the shield at all. + """ + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["result"]}) + + data = {"messages": [{"role": "tool", "name": "get_weather", "content": "result"}]} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert data["messages"][0]["name"] == "get_weather" + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["result"] + + @pytest.mark.asyncio + async def test_anthropic_tool_result_content_is_redacted(self): + """A tool_result nests its own content, as a string or as more blocks.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]", "[EMAIL_2]"]}) + + data = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "tool_result", "tool_use_id": "t1", "content": "found jane.doe@example.com"}, + { + "type": "tool_result", + "tool_use_id": "t2", + "content": [{"type": "text", "text": "also bob@example.com"}], + }, + ], + } + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert data["messages"][0]["content"][0]["content"] == "[EMAIL_1]" + assert data["messages"][0]["content"][1]["content"][0]["text"] == "[EMAIL_2]" + + @pytest.mark.asyncio + async def test_nesting_past_the_bound_blocks_the_request(self): + """Nesting is caller controlled, so the descent has to stop somewhere -- and where + it stops, the request must not go out. + + This test used to assert the opposite: that text past the bound was skipped. That + sent `past-the-bound@example.com` to the provider unredacted while the guardrail + reported as enabled. + """ + guardrail = _guardrail() + mock = _mock_post(guardrail) + + deep: dict = {"type": "tool_result", "content": "past-the-bound@example.com"} + for _ in range(200): + deep = {"type": "tool_result", "content": [deep]} + data = {"messages": [{"role": "user", "content": [{"type": "text", "text": "shallow"}, deep]}]} + + with pytest.raises(GuardrailRaisedException): + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + mock.assert_not_called() + + @pytest.mark.asyncio + async def test_deep_tool_input_blocks_the_request(self): + """A tool_use input past the JSON bound must not be forwarded half-redacted.""" + guardrail = _guardrail() + mock = _mock_post(guardrail) + + deep: dict = {"email": "past-the-bound@example.com"} + for _ in range(100): + deep = {"next": deep} + block = {"type": "tool_use", "id": "t1", "name": "f", "input": deep} + data = {"messages": [{"role": "assistant", "content": [block]}]} + + with pytest.raises(GuardrailRaisedException): + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + mock.assert_not_called() + + @pytest.mark.asyncio + async def test_realistic_nesting_is_redacted_in_full(self): + """The bounds are far past real payloads: a tool input nested inside a tool result, + several JSON levels deep, is redacted whole rather than refused.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + + tool_use = { + "type": "tool_use", + "id": "t1", + "name": "f", + "input": {"a": {"b": {"c": {"d": {"to": "x@example.com"}}}}}, + } + data = {"messages": [{"role": "user", "content": [{"type": "tool_result", "content": [tool_use]}]}]} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert tool_use["input"]["a"]["b"]["c"]["d"]["to"] == "[EMAIL_1]" + + @pytest.mark.asyncio + async def test_responses_prompt_object_variables_are_redacted(self): + """A PromptObject's variables are substituted into the prompt provider side. + + The id and version pick which stored prompt to run and have to arrive + unchanged; the variables are caller text. + """ + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + + data = {"prompt": {"id": "pmpt_123", "version": "2", "variables": {"customer": "jane.doe@example.com"}}} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert data["prompt"]["variables"]["customer"] == "[EMAIL_1]" + assert data["prompt"]["id"] == "pmpt_123" + assert data["prompt"]["version"] == "2" + + @pytest.mark.asyncio + async def test_responses_typed_prompt_variables_are_redacted(self): + """A variable can be a typed input rather than a string; its `text` is caller text.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + + data = { + "prompt": { + "id": "pmpt_123", + "variables": { + "customer": {"type": "input_text", "text": "jane.doe@example.com"}, + "logo": {"type": "input_image", "image_url": "https://example.com/logo.png"}, + }, + } + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert data["prompt"]["variables"]["customer"] == {"type": "input_text", "text": "[EMAIL_1]"} + assert data["prompt"]["variables"]["logo"]["image_url"] == "https://example.com/logo.png" + + @pytest.mark.asyncio + async def test_completions_suffix_is_redacted(self): + """LiteLLM forwards the legacy `suffix` to providers that support it.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["signed [EMAIL_1]", "write to [EMAIL_1]"]}) + + data = {"prompt": "write to jane.doe@example.com", "suffix": "signed jane.doe@example.com"} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="atext_completion") + + assert data["suffix"] == "signed [EMAIL_1]" + + @pytest.mark.asyncio + async def test_every_shape_in_one_request_is_redacted(self): + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["a", "b", "c", "d"]}) + + data = { + "messages": [ + {"role": "user", "content": "one"}, + {"role": "user", "content": [{"type": "text", "text": "two"}]}, + { + "role": "assistant", + "tool_calls": [{"function": {"name": "f", "arguments": "three"}}], + }, + ], + "input": "four", + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["one", "two", "three", "four"] + assert data["messages"][0]["content"] == "a" + assert data["messages"][1]["content"][0]["text"] == "b" + assert data["messages"][2]["tool_calls"][0]["function"]["arguments"] == "c" + assert data["input"] == "d" + + @pytest.mark.asyncio + async def test_anthropic_tool_use_input_is_redacted(self): + """A replayed tool_use block carries its arguments as a JSON object, not a string.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]", "[PHONE_1]"]}) + + data = { + "messages": [ + { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": "t1", + "name": "send", + "input": {"to": "jane.doe@example.com", "meta": {"phone": "555-0100"}}, + } + ], + } + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["jane.doe@example.com", "555-0100"] + block = data["messages"][0]["content"][0] + assert block["input"] == {"to": "[EMAIL_1]", "meta": {"phone": "[PHONE_1]"}} + assert block["name"] == "send", "the tool name has to arrive unchanged for the call to route" + + @pytest.mark.asyncio + async def test_responses_reasoning_summary_is_redacted(self): + """A replayed reasoning item quotes the conversation in its summary parts.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["user asked about [EMAIL_1]"]}) + + data = { + "input": [ + { + "type": "reasoning", + "id": "rs_1", + "summary": [{"type": "summary_text", "text": "user asked about jane.doe@example.com"}], + } + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert data["input"][0]["summary"][0]["text"] == "user asked about [EMAIL_1]" + + def test_tool_schemas_give_up_their_free_text_and_nothing_else(self): + """Every string is collected except what must reach the model verbatim: names, + types, formats, patterns and required lists.""" + data = { + "tools": [ + { + "type": "function", + "function": { + "name": "lookup", + "description": "top", + "parameters": { + "type": "object", + "title": "title", + "properties": { + # A property that is itself named "description". + "description": {"type": "string", "description": "named"}, + "kind": {"type": "string", "enum": ["a", "b"], "const": "a", "description": "enum"}, + "deep": {"type": "array", "items": {"type": "object", "description": "nested"}}, + "to": { + "type": "string", + "format": "email", + "pattern": "^.+@.+$", + "examples": ["example"], + "default": "default", + }, + "choice": {"anyOf": [{"type": "object", "default": {"type": "object-default"}}]}, + # Property names that collide with keywords are subschemas all the same. + "type": {"type": "string", "description": "named-type"}, + }, + "required": ["to"], + "$defs": {"shared": {"description": "defined"}}, + # Keywords nobody listed: scanned by default. + "dependencies": {"mode": {"description": "dependent"}}, + "$comment": "comment", + "x-note": "vendor", + }, + }, + } + ] + } + caller, privileged = LLMShieldProxyGuardrail._locate_request_texts(data) + + assert sorted(text for text, _ in caller) == ["a", "a", "b"], "enum and const go to the caller vault" + assert sorted(text for text, _ in privileged) == [ + "comment", + "default", + "defined", + "dependent", + "enum", + "example", + "named", + "named-type", + "nested", + "object-default", + "title", + "top", + "vendor", + ] + + @pytest.mark.asyncio + async def test_enum_values_are_redacted_and_restored_in_the_tool_call(self): + """An enum value holding PII is redacted, and the model's use of the stand-in is + restored in its tool arguments, so the call still carries a value the schema allows.""" + guardrail = _guardrail(event_hook=["pre_call", "post_call"]) + shield = _FakeShield({"[EMAIL_1]": "ops@example.com"}) + redact_mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + data = { + "messages": [], + "tools": [ + { + "type": "function", + "function": { + "name": "notify", + "parameters": {"properties": {"to": {"type": "string", "enum": ["ops@example.com"]}}}, + }, + } + ], + } + + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + assert redact_mock.call_args_list[0].kwargs["json"]["texts"] == ["ops@example.com"] + assert data["tools"][0]["function"]["parameters"]["properties"]["to"]["enum"] == ["[EMAIL_1]"] + + guardrail.async_handler.post = shield.post # type: ignore[method-assign] + call = SimpleNamespace(function=SimpleNamespace(name="notify", arguments='{"to": "[EMAIL_1]"}')) + reply = ModelResponse(choices=[Choices(message=Message(content=None, tool_calls=None))]) + reply.choices[0].message.tool_calls = [call] + restored = await guardrail.async_post_call_success_hook(data=data, user_api_key_dict=None, response=reply) + + assert json.loads(restored.choices[0].message.tool_calls[0].function.arguments) == {"to": "ops@example.com"} + + def test_schema_nesting_past_the_bound_is_refused(self): + schema: dict = {"type": "object", "description": "past-the-bound@example.com"} + for _ in range(100): + schema = {"type": "object", "properties": {"next": schema}} + data = {"tools": [{"type": "function", "function": {"name": "f", "parameters": schema}}]} + + with pytest.raises(Exception, match="schema"): + LLMShieldProxyGuardrail._locate_request_texts(data) + + +class TestRestoration: + @pytest.mark.asyncio + async def test_openai_shape_is_restored(self): + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["a@b.com"]}) + + response = ModelResponse(choices=[Choices(index=0, message=Message(role="assistant", content="[EMAIL_1]"))]) + result = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=response + ) + + assert result.choices[0].message.content == "a@b.com" + + @pytest.mark.asyncio + async def test_responses_api_shape_is_restored(self): + """The Responses API reply carries output items, not choices. + + Measured against a live provider: once the request side was fixed the reply + came back still holding the placeholder, because this shape has no choices + to walk. + """ + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["a@b.com"]}) + + response = SimpleNamespace(output=[SimpleNamespace(content=[{"type": "output_text", "text": "[EMAIL_1]"}])]) + result = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=response + ) + + assert result.output[0].content[0]["text"] == "a@b.com" + + @pytest.mark.asyncio + async def test_responses_api_object_blocks_are_restored(self): + """Blocks arrive as objects too, depending on how far the reply is parsed.""" + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["a@b.com"]}) + + block = SimpleNamespace(text="[EMAIL_1]") + response = SimpleNamespace(output=[SimpleNamespace(content=[block])]) + result = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=response + ) + + assert result.output[0].content[0].text == "a@b.com" + + @pytest.mark.asyncio + async def test_anthropic_message_shape_is_restored(self): + """The /v1/messages reply is a plain dict with no choices. + + Measured against a live provider: without its own branch the reply went + back to the caller still carrying the placeholder, even though the + request had been redacted correctly. + """ + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["a@b.com"]}) + + response = { + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "[EMAIL_1]"}], + } + result = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=response + ) + + assert result["content"][0]["text"] == "a@b.com" + + @pytest.mark.asyncio + async def test_anthropic_non_text_blocks_are_left_alone(self): + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["a@b.com"]}) + + response = { + "type": "message", + "role": "assistant", + "content": [ + {"type": "text", "text": "[EMAIL_1]"}, + {"type": "tool_use", "id": "t1", "name": "lookup", "input": {}}, + ], + } + result = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=response + ) + + assert result["content"][0]["text"] == "a@b.com" + assert result["content"][1] == {"type": "tool_use", "id": "t1", "name": "lookup", "input": {}} + + +class TestVaultIsolation: + """The vault id must never be something a caller can choose. + + The vault holds the plaintext behind every placeholder. If a caller could name + the vault, they could send a placeholder, have the model echo it back, and get + another caller's value restored into their own reply. + """ + + @pytest.mark.asyncio + async def test_caller_supplied_session_id_is_not_used(self): + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + + data = { + "messages": [{"role": "user", "content": "a@b.com"}], + "metadata": {"llm_shield_session_id": "victim-session"}, + "litellm_session_id": "victim-session", + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + used = mock.call_args_list[0].kwargs["headers"]["X-Session-ID"] + assert used != "victim-session" + assert data["litellm_metadata"]["llm_shield_session_id"] == used + + @pytest.mark.asyncio + async def test_session_id_is_not_forwarded_to_the_provider(self): + """The vault id is a capability, so it must stay out of provider-visible metadata. + + `metadata` is forwarded upstream on /v1/responses; `litellm_metadata` is not. A + provider holding both the placeholders and the session id could call the shield's + rehydrate endpoint and read back exactly what this guardrail withholds. + """ + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + + data = {"messages": [{"role": "user", "content": "a@b.com"}], "metadata": {}} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + used = mock.call_args_list[0].kwargs["headers"]["X-Session-ID"] + assert "llm_shield_session_id" not in data["metadata"] + assert data["litellm_metadata"]["llm_shield_session_id"] == used + + @pytest.mark.asyncio + async def test_restore_ignores_a_foreign_session_id(self): + """A reply is left unrestored rather than resolved against another vault.""" + guardrail = _guardrail(event_hook="post_call") + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + + data = {"metadata": {"llm_shield_session_id": "victim-session"}} + response = ModelResponse(choices=[Choices(index=0, message=Message(role="assistant", content="[EMAIL_1]"))]) + await guardrail.async_post_call_success_hook(data=data, user_api_key_dict=None, response=response) + + assert mock.call_args_list[0].kwargs["headers"]["X-Session-ID"] != "victim-session" + + @pytest.mark.asyncio + async def test_each_request_gets_its_own_vault(self): + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}, {"texts": ["[EMAIL_1]"]}) + + for _ in range(2): + await guardrail.async_pre_call_hook( + user_api_key_dict=None, + cache=None, + data={"messages": [{"role": "user", "content": "a@b.com"}]}, + call_type="completion", + ) + + seen = {call.kwargs["headers"]["X-Session-ID"] for call in mock.call_args_list} + assert len(seen) == 2 + + + @pytest.mark.parametrize( + "data", + [ + pytest.param( + {"messages": [{"role": "system", "content": "S"}, {"role": "user", "content": "U"}]}, + id="system-turn", + ), + pytest.param( + {"messages": [{"role": "developer", "content": "S"}, {"role": "user", "content": "U"}]}, + id="developer-turn", + ), + pytest.param( + {"system": "S", "messages": [{"role": "user", "content": "U"}]}, + id="anthropic-top-level-system", + ), + pytest.param({"instructions": "S", "input": "U"}, id="responses-instructions"), + pytest.param( + { + "messages": [{"role": "user", "content": "U"}], + "tools": [{"type": "function", "function": {"name": "f", "description": "S"}}], + }, + id="chat-tool-description", + ), + pytest.param( + {"input": "U", "tools": [{"type": "function", "name": "f", "description": "S"}]}, + id="responses-tool-description", + ), + pytest.param( + {"messages": [{"role": "user", "content": "U"}], "functions": [{"name": "f", "description": "S"}]}, + id="legacy-function-description", + ), + pytest.param( + { + "messages": [{"role": "user", "content": "U"}], + "tools": [{"name": "f", "input_schema": {"properties": {"to": {"description": "S"}}}}], + }, + id="anthropic-schema-description", + ), + pytest.param( + { + "messages": [{"role": "user", "content": "U"}], + "response_format": { + "type": "json_schema", + "json_schema": {"name": "n", "description": "S", "schema": {"type": "object"}}, + }, + }, + id="chat-response-format", + ), + pytest.param( + { + "input": "U", + "text": { + "format": { + "type": "json_schema", + "name": "n", + "schema": {"properties": {"a": {"description": "S"}}}, + } + }, + }, + id="responses-text-format", + ), + pytest.param({"prediction": {"type": "content", "content": "U"}, "instructions": "S"}, id="prediction"), + pytest.param( + {"prediction": {"type": "content", "content": [{"type": "text", "text": "U"}]}, "instructions": "S"}, + id="prediction-parts", + ), + pytest.param( + { + "messages": [{"role": "user", "content": "U"}], + "web_search_options": {"user_location": {"type": "approximate", "approximate": {"city": "S"}}}, + }, + id="chat-web-search-location", + ), + pytest.param( + { + "input": "U", + "tools": [{"type": "web_search", "user_location": {"type": "approximate", "region": "S"}}], + }, + id="responses-web-search-location", + ), + pytest.param({"messages": [{"role": "user", "content": "U"}], "user": "S"}, id="end-user-id"), + pytest.param({"input": "U", "safety_identifier": "S"}, id="safety-identifier"), + ], + ) + def test_server_authored_text_is_split_from_the_callers(self, data: dict) -> None: + """Every request shape must sort its server-authored spans out of the caller's.""" + caller, privileged = LLMShieldProxyGuardrail._locate_request_texts(data) + + assert [text for text, _ in caller] == ["U"] + assert [text for text, _ in privileged] == ["S"] + + @pytest.mark.asyncio + async def test_a_system_prompt_gets_a_vault_of_its_own(self) -> None: + """The reply is restored against the caller's vault, so the two cannot be one.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}, {"texts": ["[EMAIL_2]"]}) + + data = { + "messages": [ + {"role": "system", "content": "escalate to admin@corp.internal"}, + {"role": "user", "content": "email a@b.com"}, + ] + } + await guardrail.async_pre_call_hook( + user_api_key_dict=None, cache=None, data=data, call_type="completion" + ) + + privileged_id, caller_id = ( + call.kwargs["headers"]["X-Session-ID"] for call in mock.call_args_list + ) + assert privileged_id != caller_id + assert guardrail._session_id(data) == caller_id + + @pytest.mark.asyncio + async def test_the_system_prompt_vault_id_is_never_stored(self) -> None: + """Nothing can restore against the system vault later, because its id is not kept. + + This is what stops a caller from having the model echo a placeholder out of a + system prompt they cannot see and receiving the plaintext behind it. + """ + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}, {"texts": ["[EMAIL_2]"]}) + + data = { + "messages": [ + {"role": "system", "content": "escalate to admin@corp.internal"}, + {"role": "user", "content": "email a@b.com"}, + ] + } + await guardrail.async_pre_call_hook( + user_api_key_dict=None, cache=None, data=data, call_type="completion" + ) + + privileged_id = mock.call_args_list[0].kwargs["headers"]["X-Session-ID"] + assert privileged_id not in json.dumps(data, default=str) + + def test_responses_system_and_developer_items_are_privileged(self) -> None: + """Responses `input` carries system and developer turns as items, like Chat messages. + + In the caller's vault, a caller could have the model echo a placeholder out of a + system message they cannot see and receive the plaintext behind it. + """ + data = { + "input": [ + {"role": "system", "content": "S"}, + {"type": "message", "role": "developer", "content": [{"type": "input_text", "text": "D"}]}, + {"role": "user", "content": "U"}, + ] + } + caller, privileged = LLMShieldProxyGuardrail._locate_request_texts(data) + + assert [text for text, _ in caller] == ["U"] + assert [text for text, _ in privileged] == ["S", "D"] + + +class TestFailClosed: + @pytest.mark.asyncio + async def test_unreachable_shield_blocks_the_request(self): + """Failing open would send the PII upstream, defeating the guardrail.""" + guardrail = _guardrail() + guardrail.async_handler.post = AsyncMock(side_effect=ConnectionError("refused")) + + with pytest.raises(GuardrailRaisedException): + await guardrail.async_pre_call_hook( + user_api_key_dict=None, + cache=None, + data={"messages": [{"role": "user", "content": "a@b.com"}]}, + call_type="completion", + ) + + @pytest.mark.asyncio + async def test_error_status_blocks_the_request(self): + guardrail = _guardrail() + guardrail.async_handler.post = AsyncMock(return_value=_response({"error": "nope"}, status_code=500)) + + with pytest.raises(GuardrailRaisedException): + await guardrail.async_pre_call_hook( + user_api_key_dict=None, + cache=None, + data={"messages": [{"role": "user", "content": "a@b.com"}]}, + call_type="completion", + ) + + @pytest.mark.asyncio + async def test_short_payload_blocks_the_request(self): + """A response that loses an entry would silently misalign the write-back.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": []}) + + with pytest.raises(GuardrailRaisedException): + await guardrail.async_pre_call_hook( + user_api_key_dict=None, + cache=None, + data={"messages": [{"role": "user", "content": "a@b.com"}]}, + call_type="completion", + ) + + +class TestStreamingRehydration: + @pytest.mark.asyncio + async def test_split_placeholder_is_not_emitted_in_fragments(self): + """The window holds back a partial placeholder and releases it once complete.""" + guardrail = _guardrail(event_hook="post_call") + # Shield holds "[EMAIL" back, then releases the restored value. + _mock_post( + guardrail, + {"text": "Email ", "carry": "[EMAIL"}, + {"text": "a@b.com about it", "carry": ""}, + ) + + async def stream(): + yield _chunk("Email [EMAIL") + yield _chunk("_1] about it", finish_reason="stop") + + chunks = await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + emitted = [c.choices[0].delta.content for c in chunks] + assert emitted == ["Email ", "a@b.com about it"] + # No fragment of the placeholder ever reached the client. + assert not any("[EMAIL" in (text or "") for text in emitted) + + @pytest.mark.asyncio + async def test_carry_is_returned_to_the_next_call(self): + guardrail = _guardrail(event_hook="post_call") + mock = _mock_post( + guardrail, + {"text": "", "carry": "hold"}, + {"text": "held-and-more", "carry": ""}, + ) + + async def stream(): + yield _chunk("hold") + yield _chunk("-and-more", finish_reason="stop") + + await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + assert mock.call_args_list[0].kwargs["json"]["carry"] == "" + assert mock.call_args_list[1].kwargs["json"]["carry"] == "hold" + assert mock.call_args_list[1].kwargs["json"]["final"] is True + + @pytest.mark.asyncio + async def test_every_choice_is_restored(self): + """With n>1 a later choice must not be handed back still holding a placeholder.""" + guardrail = _guardrail(event_hook="post_call") + _mock_post( + guardrail, + {"text": "first@example.com", "carry": ""}, + {"text": "second@example.com", "carry": ""}, + ) + + async def stream(): + yield ModelResponseStream( + choices=[ + StreamingChoices(index=0, delta=Delta(content="[EMAIL_1]"), finish_reason="stop"), + StreamingChoices(index=1, delta=Delta(content="[EMAIL_2]"), finish_reason="stop"), + ] + ) + + chunks = await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + restored = [choice.delta.content for choice in chunks[0].choices] + assert restored == ["first@example.com", "second@example.com"] + + @pytest.mark.asyncio + async def test_choice_windows_do_not_cross_contaminate(self): + """Each choice is its own token stream, so each carries its own window. + + One shared window would send the characters held back for choice 0 up + against choice 1's next delta and splice the two streams together. + """ + guardrail = _guardrail(event_hook="post_call") + mock = _mock_post( + guardrail, + {"text": "", "carry": "A-held"}, + {"text": "", "carry": "B-held"}, + {"text": "a-done", "carry": ""}, + {"text": "b-done", "carry": ""}, + ) + + async def stream(): + yield ModelResponseStream( + choices=[ + StreamingChoices(index=0, delta=Delta(content="a1")), + StreamingChoices(index=1, delta=Delta(content="b1")), + ] + ) + yield ModelResponseStream( + choices=[ + StreamingChoices(index=0, delta=Delta(content="a2"), finish_reason="stop"), + StreamingChoices(index=1, delta=Delta(content="b2"), finish_reason="stop"), + ] + ) + + await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + sent = [call.kwargs["json"] for call in mock.call_args_list] + assert sent[2]["carry"] == "A-held", "choice 0 must get its own window back" + assert sent[3]["carry"] == "B-held", "choice 1 must get its own window back" + + @pytest.mark.asyncio + async def test_a_choice_missing_from_the_last_chunk_still_flushes(self): + """Held text must not be dropped because its choice ended earlier. + + Choice 1 finishes and stops appearing, then the stream ends without a + finish_reason for choice 0. Flushing only the terminal chunk's choices would + discard whatever choice 1 was still holding and truncate its answer. + """ + guardrail = _guardrail(event_hook="post_call") + _mock_post( + guardrail, + {"text": "", "carry": "held-0"}, + {"text": "", "carry": "held-1"}, + {"text": "zero-done", "carry": ""}, + {"text": "one-done", "carry": ""}, + ) + + async def stream(): + yield ModelResponseStream( + choices=[ + StreamingChoices(index=0, delta=Delta(content="a")), + StreamingChoices(index=1, delta=Delta(content="b")), + ] + ) + yield ModelResponseStream(choices=[StreamingChoices(index=0, delta=Delta(content=None))]) + + chunks = await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + flushed = { + choice.index: choice.delta.content for chunk in chunks for choice in chunk.choices if choice.delta.content + } + assert flushed.get(1) == "one-done", "choice 1's held text was dropped" + assert flushed.get(0) == "zero-done" + + @pytest.mark.asyncio + async def test_held_tool_arguments_land_in_the_finishing_chunk(self): + """A client parses tool arguments on finish_reason, so the flush must ride that chunk. + + The finishing chunk also carries its own fragment for the same tool call. That + entry has to survive, with the held text appended after it as a continuation. + """ + guardrail = _guardrail(event_hook="post_call") + _mock_post( + guardrail, + {"text": '{"to": "', "carry": "[EMA"}, + {"text": "", "carry": '[EMAIL_1]"}'}, + {"text": 'a@example.com"}', "carry": ""}, + ) + + async def stream(): + yield _tool_chunk('{"to": "[EMA') + yield _tool_chunk('IL_1]"}', finish_reason="tool_calls") + + chunks = await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + assert len(chunks) == 2, "the flush must not arrive after the finish_reason chunk" + final_calls = chunks[1].choices[0].delta.tool_calls + assert len(final_calls) == 2, "the finishing chunk's own fragment was dropped" + assert _field(final_calls[1], "index") == 0 + arguments = "".join( + _field(_field(call, "function"), "arguments") or "" + for chunk in chunks + for call in chunk.choices[0].delta.tool_calls + ) + assert json.loads(arguments) == {"to": "a@example.com"} + + @pytest.mark.asyncio + async def test_held_tool_arguments_flush_when_the_stream_ends_unfinished(self): + """No finish_reason at all: a trailing chunk carries the held arguments alone.""" + guardrail = _guardrail(event_hook="post_call") + _mock_post( + guardrail, + {"text": '{"to": "', "carry": "[EMA"}, + {"text": 'a@example.com"}', "carry": ""}, + ) + + async def stream(): + yield _tool_chunk('{"to": "[EMA') + + chunks = await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + assert len(chunks) == 2 + trailing = chunks[1].choices[0].delta.tool_calls + assert trailing == [{"index": 0, "function": {"arguments": 'a@example.com"}'}}], ( + "the copied chunk's own fragment was already delivered and must not repeat" + ) + + @pytest.mark.asyncio + async def test_chunks_are_forwarded_as_they_arrive(self): + """Restoration must not buffer the stream into a single terminal chunk.""" + guardrail = _guardrail(event_hook="post_call") + _mock_post( + guardrail, + {"text": "one ", "carry": ""}, + {"text": "two ", "carry": ""}, + {"text": "three", "carry": ""}, + ) + + async def stream(): + yield _chunk("one ") + yield _chunk("two ") + yield _chunk("three", finish_reason="stop") + + chunks = await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + assert len(chunks) == 3 + assert [c.choices[0].delta.content for c in chunks] == ["one ", "two ", "three"] + + +class TestApplyGuardrailToolCalls: + """The unified entry point the UI's Test button and the translation handlers use.""" + + @pytest.mark.asyncio + async def test_response_tool_call_arguments_are_rehydrated(self): + """Regression: this path deep-copied tool calls with `copy` never imported. + + 47 tests passed with a guaranteed NameError here, because every tool-call test + covered the request side and this is the only path that reaches the copy. + """ + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["hi", '{"email": "a@b.com"}']}) + + data = {"litellm_metadata": {"llm_shield_session_id": "shield-abc"}} + inputs = { + "texts": ["hi"], + "tool_calls": [{"function": {"name": "send", "arguments": '{"email": "[EMAIL_1]"}'}}], + } + + merged = await guardrail.apply_guardrail(inputs=inputs, request_data=data, input_type="response") + + assert merged["tool_calls"][0]["function"]["arguments"] == '{"email": "a@b.com"}' + assert inputs["tool_calls"][0]["function"]["arguments"] == '{"email": "[EMAIL_1]"}' + + +class TestAnthropicStreamRestoration: + """/v1/messages streams reach the hook as raw SSE frames, with no `choices` to walk.""" + + VAULT = {"[EMAIL_1]": "a@example.com"} + + @pytest.mark.asyncio + async def test_split_placeholder_is_restored_and_never_fragmented(self): + guardrail, _ = _shielded(self.VAULT) + + out = await _restore_stream(guardrail, _text_block_stream("Mail [EMA", "IL_1] now")) + + deltas = [e["delta"]["text"] for e in _sse_events(out) if e["type"] == "content_block_delta"] + assert "".join(deltas) == "Mail a@example.com now" + assert not any("[EMA" in d for d in deltas), "a placeholder fragment reached the client" + + @pytest.mark.asyncio + async def test_held_text_lands_before_its_block_stops(self): + """A trailing `[` that never became a placeholder is still part of the answer.""" + guardrail, _ = _shielded(self.VAULT) + + out = await _restore_stream(guardrail, _text_block_stream("Mail [EMAIL_1], x = a[")) + + types = [e["type"] for e in _sse_events(out)] + deltas = [e["delta"]["text"] for e in _sse_events(out) if e["type"] == "content_block_delta"] + assert "".join(deltas) == "Mail a@example.com, x = a[" + assert types.index("content_block_stop") > max(i for i, t in enumerate(types) if t == "content_block_delta") + + @pytest.mark.asyncio + async def test_events_split_across_network_chunks_are_restored(self): + """A chunk can end mid-event; the frame is parsed once it is whole.""" + guardrail, _ = _shielded(self.VAULT) + raw = b"".join(_text_block_stream("Mail [EMA", "IL_1] now")) + + out = await _restore_stream(guardrail, [raw[i : i + 7] for i in range(0, len(raw), 7)]) + + deltas = [e["delta"]["text"] for e in _sse_events(out) if e["type"] == "content_block_delta"] + assert "".join(deltas) == "Mail a@example.com now" + + @pytest.mark.asyncio + async def test_str_frames_stay_str(self): + guardrail, _ = _shielded(self.VAULT) + + out = await _restore_stream(guardrail, [frame.decode() for frame in _text_block_stream("[EMAIL_1]")]) + + assert all(isinstance(frame, str) for frame in out) + assert "a@example.com" in "".join(out) + + @pytest.mark.asyncio + async def test_tool_input_json_is_restored(self): + guardrail, _ = _shielded(self.VAULT) + frames = [ + _sse({"type": "content_block_start", "index": 1, "content_block": {"type": "tool_use", "input": {}}}), + *( + _sse( + { + "type": "content_block_delta", + "index": 1, + "delta": {"type": "input_json_delta", "partial_json": p}, + } + ) + for p in ('{"to": "[EMAI', 'L_1]"}') + ), + _sse({"type": "content_block_stop", "index": 1}), + ] + + out = await _restore_stream(guardrail, frames) + + partial = "".join(e["delta"]["partial_json"] for e in _sse_events(out) if e["type"] == "content_block_delta") + assert json.loads(partial) == {"to": "a@example.com"} + + @pytest.mark.asyncio + async def test_signed_thinking_and_foreign_frames_pass_through_byte_for_byte(self): + """Rewriting a signed thinking block breaks it; other frames are not ours to touch.""" + guardrail, shield = _shielded(self.VAULT) + thinking = {"type": "thinking_delta", "thinking": "[EMAIL_1]"} + frames = [ + _sse({"type": "content_block_delta", "index": 0, "delta": thinking}), + b'data: {"candidates": [{"content": {"parts": [{"text": "[EMAIL_1]"}]}}]}\n\n', + b"data: not json\n\n", + ] + + out = await _restore_stream(guardrail, frames) + + assert b"".join(out) == b"".join(frames) + assert shield.urls == [] + + @pytest.mark.asyncio + async def test_a_raw_stream_that_is_not_sse_is_never_buffered(self): + """Without event boundaries to wait for, buffering would hold the whole reply.""" + guardrail, _ = _shielded(self.VAULT) + chunks = [b'[{"candidates": []}', b', {"candidates": []}]'] + + out = await _restore_stream(guardrail, chunks) + + assert out == chunks + + @pytest.mark.asyncio + @pytest.mark.parametrize("cut", [1, 3, 5, 6]) + async def test_a_field_name_split_by_the_first_chunk_still_reads_as_sse(self, cut: int): + """`b"eve"` then `b"nt: ..."` is still SSE; deciding on the first chunk alone + would pass the whole stream through with its placeholders.""" + guardrail, _ = _shielded(self.VAULT) + raw = b"".join(_text_block_stream("Mail [EMAIL_1]")) + + out = await _restore_stream(guardrail, [raw[:cut], raw[cut:]]) + + deltas = [e["delta"]["text"] for e in _sse_events(out) if e["type"] == "content_block_delta"] + assert "".join(deltas) == "Mail a@example.com" + + +class TestResponsesStreamRestoration: + """/v1/responses streams are typed events, with no `choices` to walk.""" + + VAULT = {"[EMAIL_1]": "a@example.com"} + + @staticmethod + def _text_delta(delta: str, sequence_number: int, content_index: int = 0) -> OutputTextDeltaEvent: + return OutputTextDeltaEvent( + type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, + item_id="msg_1", + output_index=0, + content_index=content_index, + delta=delta, + sequence_number=sequence_number, + ) + + @pytest.mark.asyncio + async def test_deltas_and_done_text_are_restored(self): + guardrail, _ = _shielded(self.VAULT) + done = OutputTextDoneEvent( + type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE, + item_id="msg_1", + output_index=0, + content_index=0, + text="Mail [EMAIL_1] x[", + ) + + out = await _restore_stream( + guardrail, [self._text_delta("Mail [EMA", 1), self._text_delta("IL_1] x[", 2), done] + ) + + deltas = [e.delta for e in out if isinstance(e, OutputTextDeltaEvent)] + assert not any("[EMA" in d for d in deltas), "a placeholder fragment reached the client" + assert "".join(deltas) == "Mail a@example.com x[" + assert out[-1].text == "Mail a@example.com x[", "the done event repeats the full, restored text" + assert isinstance(out[-2], OutputTextDeltaEvent), "held text must land before the done event" + + @pytest.mark.asyncio + async def test_function_call_arguments_are_restored(self): + guardrail, _ = _shielded(self.VAULT) + events = [ + FunctionCallArgumentsDeltaEvent( + type=ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA, + item_id="fc_1", + output_index=1, + delta=part, + ) + for part in ('{"to": "[EMAI', 'L_1]"}') + ] + + out = await _restore_stream(guardrail, events) + + assert json.loads("".join(e.delta for e in out)) == {"to": "a@example.com"} + + @pytest.mark.asyncio + async def test_a_truncated_stream_still_flushes(self): + """No done event at all: whatever the window holds goes out at the end.""" + guardrail, _ = _shielded(self.VAULT) + + out = await _restore_stream(guardrail, [self._text_delta("see [EMAIL_1] a[", 1)]) + + assert "".join(e.delta for e in out) == "see a@example.com a[" + + @pytest.mark.asyncio + async def test_completed_response_is_restored(self): + """The terminal event repeats the whole reply, and clients read it as the answer.""" + guardrail, _ = _shielded(self.VAULT) + block = {"type": "output_text", "text": "Mail [EMAIL_1]"} + call = SimpleNamespace(type="function_call", arguments='{"to": "[EMAIL_1]"}') + completed = SimpleNamespace( + type="response.completed", + response=SimpleNamespace(output=[SimpleNamespace(content=[block]), call]), + ) + + (out,) = await _restore_stream(guardrail, [completed]) + + restored_block, restored_call = out.response.output[0].content[0], out.response.output[1] + assert restored_block["text"] == "Mail a@example.com" + assert restored_call.arguments == '{"to": "a@example.com"}' + + @pytest.mark.asyncio + async def test_streams_on_different_parts_do_not_share_a_window(self): + guardrail, _ = _shielded(self.VAULT) + events = [ + self._text_delta("one [EMA", 1), + self._text_delta("two", 2, content_index=1), + self._text_delta("IL_1]", 3), + ] + + out = await _restore_stream(guardrail, events) + + by_part: dict[int, str] = {} + for event in out: + by_part[event.content_index] = by_part.get(event.content_index, "") + event.delta + assert by_part == {0: "one a@example.com", 1: "two"} + + @pytest.mark.asyncio + async def test_reasoning_summary_part_done_is_restored(self): + """The summary part repeats the whole summary text after its deltas.""" + guardrail, _ = _shielded(self.VAULT) + part = SimpleNamespace(type="summary_text", text="asked about [EMAIL_1]") + event = SimpleNamespace( + type="response.reasoning_summary_part.done", item_id="rs_1", output_index=0, summary_index=0, part=part + ) + + (out,) = await _restore_stream(guardrail, [event]) + + assert out.part.text == "asked about a@example.com" + + @pytest.mark.asyncio + async def test_mcp_call_arguments_are_restored(self): + """A stream family outside the chat-era set: matched by shape, not by name.""" + guardrail, _ = _shielded(self.VAULT) + deltas = [ + {"type": "response.mcp_call_arguments.delta", "item_id": "mcp_1", "output_index": 0, "delta": d} + for d in ('{"to": "[EMAI', 'L_1]"}') + ] + done = { + "type": "response.mcp_call_arguments.done", + "item_id": "mcp_1", + "output_index": 0, + "arguments": '{"to": "[EMAIL_1]"}', + } + + out = await _restore_stream(guardrail, [*deltas, done]) + + assert json.loads("".join(e["delta"] for e in out[:-1])) == {"to": "a@example.com"} + assert json.loads(out[-1]["arguments"]) == {"to": "a@example.com"} + assert out[-1]["item_id"] == "mcp_1", "identifiers are not text and stay as sent" + + @pytest.mark.asyncio + async def test_audio_deltas_are_not_sent_to_the_shield(self): + """Audio arrives base64-encoded; restoring it would cost a round trip for nothing.""" + guardrail, shield = _shielded(self.VAULT) + audio = {"type": "response.audio.delta", "item_id": "a_1", "output_index": 0, "delta": "UklGRiQAAABXQVZF"} + + out = await _restore_stream(guardrail, [audio]) + + assert out == [audio] + assert shield.urls == [] + + +class TestResponseCacheIsolation: + """The reply LiteLLM caches must keep its placeholders. + + Placeholders are numbered per request, so two callers' redacted requests can be + identical and share a cache key. LiteLLM keeps the provider's reply object -- for a + native Anthropic dict or an in-memory cache, the object itself -- so restoring it in + place would hand one caller's values to the next caller who hits that key. + """ + + @pytest.mark.asyncio + async def test_a_cache_hit_is_restored_against_the_new_callers_vault(self): + cache = InMemoryCache() + provider_reply = {"type": "message", "content": [{"type": "text", "text": "Repeat [EMAIL_1]"}]} + cache.set_cache("redacted-request", provider_reply) + + alice, _ = _shielded({"[EMAIL_1]": "alice@example.com"}) + to_alice = await alice.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=cache.get_cache("redacted-request") + ) + bob, _ = _shielded({"[EMAIL_1]": "bob@example.com"}) + to_bob = await bob.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=cache.get_cache("redacted-request") + ) + + assert to_alice["content"][0]["text"] == "Repeat alice@example.com" + assert to_bob["content"][0]["text"] == "Repeat bob@example.com" + assert cache.get_cache("redacted-request")["content"][0]["text"] == "Repeat [EMAIL_1]" + + @pytest.mark.asyncio + async def test_a_model_response_is_restored_as_a_copy(self): + """A reply as `acompletion` returns it, hidden params and all.""" + reply = await litellm.acompletion( + model="gpt-4o-mini", messages=[{"role": "user", "content": "hi"}], mock_response="Mail [EMAIL_1]" + ) + guardrail, _ = _shielded({"[EMAIL_1]": "a@example.com"}) + + restored = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=reply + ) + + assert restored.choices[0].message.content == "Mail a@example.com" + assert reply.choices[0].message.content == "Mail [EMAIL_1]" + + @pytest.mark.asyncio + async def test_responses_reply_is_restored_as_a_copy(self): + guardrail, _ = _shielded({"[EMAIL_1]": "a@example.com"}) + block = {"type": "output_text", "text": "Mail [EMAIL_1]"} + reply = SimpleNamespace(output=[SimpleNamespace(content=[block])]) + + restored = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=reply + ) + + assert restored.output[0].content[0]["text"] == "Mail a@example.com" + assert block["text"] == "Mail [EMAIL_1]" + + @pytest.mark.asyncio + async def test_stream_chunks_are_restored_as_copies(self): + """LiteLLM assembles the reply it caches from the chunks it yielded.""" + guardrail, _ = _shielded({"[EMAIL_1]": "a@example.com"}) + chunk = _chunk("Mail [EMAIL_1]", finish_reason="stop") + event = OutputTextDeltaEvent( + type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, + item_id="msg_1", + output_index=0, + content_index=0, + delta="Mail [EMAIL_1]", + ) + + out = await _restore_stream(guardrail, [chunk, event]) + + assert out[0].choices[0].delta.content == "Mail a@example.com" + assert chunk.choices[0].delta.content == "Mail [EMAIL_1]" + assert "Mail [EMAIL_1]" == event.delta + + +class TestCompletionsRestoration: + """`/v1/completions` replies carry their text on the choice, with no message or delta.""" + + VAULT = {"[EMAIL_1]": "a@example.com"} + + @pytest.mark.asyncio + async def test_completions_reply_is_restored(self): + guardrail, _ = _shielded(self.VAULT) + reply = TextCompletionResponse(choices=[TextChoices(index=0, text="Mail [EMAIL_1]", finish_reason="stop")]) + + restored = await guardrail.async_post_call_success_hook( + data={"prompt": "x"}, user_api_key_dict=None, response=reply + ) + + assert restored.choices[0].text == "Mail a@example.com" + + @pytest.mark.asyncio + async def test_completions_stream_is_restored_across_chunks(self): + guardrail, _ = _shielded(self.VAULT) + chunks = [ + TextCompletionResponse(choices=[TextChoices(index=0, text="Mail [EMAI")]), + TextCompletionResponse(choices=[TextChoices(index=0, text="L_1] now", finish_reason="stop")]), + ] + + out = await _restore_stream(guardrail, chunks) + + assert "".join(chunk.choices[0].text for chunk in out) == "Mail a@example.com now" + + @pytest.mark.asyncio + async def test_completions_stream_without_finish_reason_is_flushed(self): + guardrail, _ = _shielded(self.VAULT) + chunks = [TextCompletionResponse(choices=[TextChoices(index=0, text="Mail [EMAIL_1")])] + + out = await _restore_stream(guardrail, chunks) + + assert "".join(chunk.choices[0].text or "" for chunk in out) == "Mail [EMAIL_1" + + +class TestProxyWiring: + def test_dashboard_config_model_is_exposed(self): + """The guardrail garden reads the provider's fields from `get_config_model`.""" + model = LLMShieldProxyGuardrail.get_config_model() + + assert model is not None + assert {"api_key", "api_base"} <= set(model.model_fields) + + @pytest.mark.asyncio + async def test_the_deployment_hook_leaves_a_proxy_reply_for_the_proxy_hook(self): + """Inside the proxy the deployment hook must not restore: LiteLLM caches what it returns. + + A proxy request was redacted by the proxy's pre-call hook, so it carries no + deployment-restore marker, and the proxy's post-call hook restores it after the cache + write. A caller-sent marker that does not match the minted vault id is ignored. + """ + guardrail, shield = _shielded({"[EMAIL_1]": "a@example.com"}) + reply = ModelResponse(choices=[Choices(index=0, message=Message(role="assistant", content="[EMAIL_1]"))]) + data = {"messages": [], "guardrails": [GUARDRAIL_NAME], "litellm_metadata": {}} + LLMShieldProxyGuardrail._mint_session_id(data) + data["litellm_metadata"]["llm_shield_restore_at_deployment"] = "caller-chosen" + + result = await guardrail.async_post_call_success_deployment_hook( + request_data=data, response=reply, call_type=None + ) + + assert result is None + assert reply.choices[0].message.content == "[EMAIL_1]" + assert shield.urls == [] + + @pytest.mark.asyncio + async def test_model_level_use_outside_the_proxy_is_restored_and_never_cached(self, monkeypatch): + """SDK use with model-level `guardrails`: the deployment hooks are the only redact and + restore steps, so the reply is restored there, and the request bypasses the cache -- + its key is built from the redacted request, and a cache hit would skip restoration. + """ + vault = {"[EMAIL_1]": "alice@example.com"} + + async def shield(url: str, headers: dict, json: dict, timeout: float) -> Response: + texts = json["texts"] + if url.endswith("/redact"): + return _response({"texts": [t.replace("alice@example.com", "[EMAIL_1]") for t in texts]}) + restored = [] + for text in texts: + for placeholder, original in vault.items(): + text = text.replace(placeholder, original) + restored.append(text) + return _response({"texts": restored}) + + guardrail = _guardrail(event_hook=["pre_call", "post_call"], default_on=False) + guardrail.async_handler.post = shield # type: ignore[method-assign] + cache = InMemoryCache() + monkeypatch.setattr(litellm, "callbacks", [guardrail]) + monkeypatch.setattr(litellm, "cache", litellm.Cache(type="local")) + monkeypatch.setattr(litellm.cache, "cache", cache) + + reply = await litellm.acompletion( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "Repeat alice@example.com"}], + mock_response="Repeat [EMAIL_1]", + guardrails=[GUARDRAIL_NAME], + ) + + # LiteLLM writes the cache from background tasks; let them land before looking. + await asyncio.gather(*_PENDING_CACHE_WRITES) + + assert reply.choices[0].message.content == "Repeat alice@example.com" + assert cache.cache_dict == {}, "the redacted request's reply must not be cached" + + @pytest.mark.asyncio + async def test_model_level_streaming_outside_the_proxy_is_refused(self, monkeypatch): + """No hook restores an SDK stream, and its cache writer misses the bypass, so it fails closed.""" + guardrail = _guardrail(event_hook=["pre_call", "post_call"], default_on=False) + _mock_post(guardrail, {"texts": ["Repeat [EMAIL_1]"]}) + cache = InMemoryCache() + monkeypatch.setattr(litellm, "callbacks", [guardrail]) + monkeypatch.setattr(litellm, "cache", litellm.Cache(type="local")) + monkeypatch.setattr(litellm.cache, "cache", cache) + + with pytest.raises(GuardrailRaisedException, match="cannot restore a streamed reply"): + await litellm.acompletion( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "Repeat alice@example.com"}], + mock_response="Repeat [EMAIL_1]", + stream=True, + guardrails=[GUARDRAIL_NAME], + ) + await asyncio.gather(*_PENDING_CACHE_WRITES) + + assert cache.cache_dict == {} + + @pytest.mark.asyncio + async def test_restored_values_are_not_recorded_as_guardrail_telemetry(self): + """Guardrail logging is exported to traces even with message logging off, so the + restored reply must not land in it.""" + guardrail, _ = _shielded({"[EMAIL_1]": "alice@example.com"}) + reply = ModelResponse(choices=[Choices(index=0, message=Message(role="assistant", content="[EMAIL_1]"))]) + data = {"messages": [], "metadata": {}} + + restored = await guardrail.async_post_call_success_hook(data=data, user_api_key_dict=None, response=reply) + + assert restored.choices[0].message.content == "alice@example.com" + assert "alice@example.com" not in json.dumps(data, default=str) + + +class TestStreamUsage: + @pytest.mark.asyncio + async def test_trailing_flush_does_not_repeat_usage(self): + """With n>=2 and include_usage, the last chunk carries usage; a flush copied from it must not. + + Any consumer that sums usage chunks would otherwise count the request twice. + """ + guardrail, _ = _shielded({"[EMAIL_1]": "a@example.com"}) + both = ModelResponseStream( + choices=[ + StreamingChoices(index=0, delta=Delta(content="Mail [EMAI")), + StreamingChoices(index=1, delta=Delta(content="Call [EMAI")), + ] + ) + usage_chunk = ModelResponseStream( + choices=[StreamingChoices(index=1, delta=Delta(content=None))], + usage=Usage(prompt_tokens=5, completion_tokens=7, total_tokens=12), + ) + + out = await _restore_stream(guardrail, [both, usage_chunk]) + + with_usage = [chunk for chunk in out if getattr(chunk, "usage", None) is not None] + assert with_usage == [out[1]], "only the provider's own usage chunk carries usage" + flushed = out[2:] + assert flushed, "the held-back text is flushed at end of stream" + assert sorted(chunk.choices[0].index for chunk in flushed) == [0, 1] + assert [chunk.choices[0].delta.content for chunk in flushed] == ["[EMAI", "[EMAI"] diff --git a/ui/litellm-dashboard/public/assets/logos/llm_shield_proxy.svg b/ui/litellm-dashboard/public/assets/logos/llm_shield_proxy.svg new file mode 100644 index 00000000000..0dd78b078c9 --- /dev/null +++ b/ui/litellm-dashboard/public/assets/logos/llm_shield_proxy.svg @@ -0,0 +1,6 @@ + + + + + + diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx index 29df7c8bf3d..a02cae097a7 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx @@ -73,7 +73,9 @@ interface GuardrailPreset { provider: string; categoryName?: string; guardrailNameSuggestion: string; - mode: string; + // A guardrail that both rewrites the request and repairs the response needs two + // modes seeded, not one; the form already normalises either shape. + mode: string | string[]; defaultOn: boolean; } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts index 10ca58294b9..684049c54a9 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts @@ -2,7 +2,9 @@ export interface GuardrailPreset { provider: string; categoryName?: string; guardrailNameSuggestion: string; - mode: string; + // A guardrail that both rewrites the request and repairs the response needs two + // modes seeded, not one; the form already normalises either shape. + mode: string | string[]; defaultOn: boolean; } @@ -325,6 +327,14 @@ export const GUARDRAIL_PRESETS: Record = { // MCP-only: default_on is the only activation path on the MCP hook defaultOn: true, }, + llm_shield_proxy: { + provider: "LLM Shield Proxy", + guardrailNameSuggestion: "LLM Shield Proxy", + // Both halves are required. With only pre_call the request is redacted and the + // placeholders are handed straight back to the caller. + mode: ["pre_call", "post_call"], + defaultOn: false, + }, conduct: { provider: "Conduct", guardrailNameSuggestion: "Conduct Guard", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts index a27dd344c95..e756c8c3e58 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts @@ -29,6 +29,7 @@ const EXPECTED_PARTNER_LOGO_FILES: Record = { straiker: "straiker.svg", alice: "alice.svg", agent_365: "microsoft_azure.svg", + llm_shield_proxy: "llm_shield_proxy.svg", conduct: "conduct.png", }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts index d88a333d6f1..d1ead5f589a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts @@ -484,6 +484,16 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ tags: ["Agentic", "MCP", "Tool Misuse", "Observability"], providerKey: "Agent365", }, + { + id: "llm_shield_proxy", + name: "LLM Shield Proxy", + description: + "Self-hosted PII redaction that puts the original values back into the model's response, so the provider never receives personal data while the end user still sees it.", + category: "partner", + logo: guardrailLogoMap["LLM Shield Proxy"], + tags: ["PII", "Data Privacy", "Compliance", "Streaming"], + providerKey: "LLM Shield Proxy", + }, { id: "conduct", name: "Conduct Guard", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx index 8df7dfb1403..09e0f21c670 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx @@ -1,6 +1,7 @@ import aimSecurityLogo from "../../../../../public/assets/logos/aim_security.jpeg"; import aktoLogo from "../../../../../public/assets/logos/akto.svg"; import aliceLogo from "../../../../../public/assets/logos/alice.svg"; +import llmShieldProxyLogo from "../../../../../public/assets/logos/llm_shield_proxy.svg"; import conductLogo from "../../../../../public/assets/logos/conduct.png"; import aporiaLogo from "../../../../../public/assets/logos/aporia.png"; import bedrockLogo from "../../../../../public/assets/logos/bedrock.svg"; @@ -86,6 +87,7 @@ export const guardrail_provider_map: Record = { QostodianNexus: "qostodian_nexus", Repelloai: "repelloai", Alice: "alice", + "LLM Shield Proxy": "llm_shield_proxy", Conduct: "conduct", }; @@ -211,6 +213,7 @@ export const guardrailLogoMap = { Straiker: straikerLogo.src, Alice: aliceLogo.src, "Microsoft Agent 365": microsoftAzureLogo.src, + "LLM Shield Proxy": llmShieldProxyLogo.src, "Conduct Guard": conductLogo.src, } satisfies Record;