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 index 81a47bbc0aa..27af12ae930 100644 --- 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 @@ -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. 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 index b218d8113e2..d6d1f42e9bd 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -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"