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.
This commit is contained in:
Ninad Phalak 2026-10-05 07:52:27 +00:00
parent ac6af96883
commit ae3cebd6ed
No known key found for this signature in database
2 changed files with 116 additions and 26 deletions

View file

@ -62,6 +62,13 @@ _REHYDRATE_STREAM_PATH: Final = "/v1/guard/rehydrate/stream"
# across concurrent requests.
_SESSION_METADATA_KEY: Final = "llm_shield_session_id"
# Set when the deployment pre-call hook redacted the request -- model-level `guardrails`
# outside the proxy -- to that request's vault id. Only then is the reply restored at the
# deployment, because only then does no later hook restore it. Matching it against the
# minted id, which carries the unguessable per-process prefix, means a caller cannot opt
# a proxy request into deployment-level restoration by sending the key themselves.
_DEPLOYMENT_RESTORE_KEY: Final = "llm_shield_restore_at_deployment"
# Roles whose text the application author wrote and the caller never sees. Their
# PII is still redacted outbound, but it is not restorable from the reply.
_PRIVILEGED_ROLES: Final = frozenset({"system", "developer"})
@ -544,8 +551,9 @@ def _read_field(holder: object, name: str) -> object:
others, depending how far they have been deserialised, so every response walk here
has to handle both shapes.
"""
if isinstance(holder, dict):
return holder.get(name)
fields: Final = _as_object(holder)
if fields is not None:
return fields.get(name)
return getattr(holder, name, None)
@ -990,21 +998,55 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
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.
"""
before: Final = self._minted_session_id(kwargs)
# The parent rewrites `kwargs` in place and hands the same dict back.
_ = 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
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":
"""Leaves the reply alone at the deployment, where LiteLLM caches what this returns.
"""Restores the reply here only when the deployment pre-call hook redacted it.
The inherited hook restores a model-level guardrail's reply here, before
`litellm/utils.py` writes it to the response cache, so the cache would hold this
caller's plaintext under a key built from the redacted request. The proxy's own
post-call hook runs model-level guardrails too, after the cache write, and
restores the reply there.
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.
"""
return None
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)
# --- transport ---------------------------------------------------------------
@ -1088,6 +1130,17 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
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.
@ -1096,13 +1149,8 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
reply that cannot be restored is a visible placeholder, while trusting a
caller-supplied id would hand them someone else's plaintext.
"""
# Read only from `litellm_metadata`, the same proxy-private store `_mint_session_id`
# writes to. A caller can populate `metadata`; they cannot populate this.
metadata: Final = data.get("litellm_metadata")
existing: Final = metadata.get(_SESSION_METADATA_KEY) if isinstance(metadata, dict) else None
if isinstance(existing, str) and existing.startswith(_VAULT_PREFIX):
return existing
return f"{_VAULT_PREFIX}-{uuid.uuid4().hex}"
existing: Final = LLMShieldProxyGuardrail._minted_session_id(data)
return existing if existing is not None else f"{_VAULT_PREFIX}-{uuid.uuid4().hex}"
# --- request traversal --------------------------------------------------------
@ -1276,11 +1324,8 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
@staticmethod
def _is_anthropic_message_response(response: object) -> bool:
"""Anthropic's native /v1/messages reply arrives as a plain dict."""
return (
isinstance(response, dict)
and response.get("type") == "message"
and isinstance(response.get("content"), list)
)
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.

View file

@ -1,3 +1,4 @@
import asyncio
import json
from types import SimpleNamespace
from unittest.mock import AsyncMock
@ -6,6 +7,7 @@ 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 (
@ -1875,19 +1877,62 @@ class TestProxyWiring:
assert {"api_key", "api_base"} <= set(model.model_fields)
@pytest.mark.asyncio
async def test_the_deployment_hook_leaves_the_reply_for_the_cache_untouched(self):
"""LiteLLM caches what the deployment hook returns, so restoring there caches plaintext.
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.
The proxy's post-call hook, which runs after the cache write, restores model-level
guardrails instead.
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={"messages": [], "guardrails": [GUARDRAIL_NAME]}, response=reply, call_type=None
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"