mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
ac6af96883
commit
ae3cebd6ed
2 changed files with 116 additions and 26 deletions
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue