diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py index c126bf0de37..26c6774f48d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py @@ -24,7 +24,6 @@ from litellm._logging import verbose_proxy_logger from litellm.exceptions import GuardrailRaisedException from litellm.integrations.custom_guardrail import ( CustomGuardrail, - get_session_id_from_request_data, log_guardrail_information, ) from litellm.llms.custom_httpx.http_handler import ( @@ -52,6 +51,13 @@ _REHYDRATE_STREAM_PATH: Final = "/v1/guard/rehydrate/stream" # across concurrent requests. _SESSION_METADATA_KEY: Final = "llm_shield_session_id" +# Vault ids are minted here and never derived from anything the caller sends. The +# vault holds the plaintext behind every placeholder, so an id a caller could +# supply or guess would let one user rehydrate another user's values by getting a +# placeholder echoed back. The per-process prefix means a caller cannot even name +# a vault this process uses. +_VAULT_PREFIX: Final = f"litellm-{uuid.uuid4().hex}" + _DEFAULT_TIMEOUT_SECONDS: Final = 10.0 # The proxy's own request dict. Mutable by design: a pre-call guardrail rewrites @@ -67,6 +73,51 @@ JsonBody: TypeAlias = dict # replacement back where it came from. _Slot: TypeAlias = tuple[str, Callable[[str], None]] # mutable-ok: Callable's param list. +# The accumulator the collectors below append into. It never escapes +# _locate_request_texts, which freezes it into a tuple before returning. +_SlotSink: TypeAlias = list[_Slot] # mutable-ok: accumulator passed between collectors. + + +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, lambda new, c=container, k=key: c.__setitem__(k, new))) + + +def _collect_content(container: MutableRequest, slots: _SlotSink) -> None: + """`content` is either a string or the multimodal list of typed parts.""" + content: Final = container.get("content") + if isinstance(content, str): + _collect(container, "content", slots) + return + for part in content if isinstance(content, list) else (): + if isinstance(part, dict): + _collect(part, "text", 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 message.get("tool_calls") or (): + function: Final = tool_call.get("function") if isinstance(tool_call, dict) else None + if isinstance(function, dict): + _collect(function, "arguments", slots) + legacy: Final = message.get("function_call") + if isinstance(legacy, dict): + _collect(legacy, "arguments", slots) + + +def _collect_responses_fields(data: MutableRequest, slots: _SlotSink) -> None: + """The Responses API sends text outside `messages`, in `instructions` and `input`.""" + _collect(data, "instructions", slots) + request_input: Final = data.get("input") + if isinstance(request_input, str): + _collect(data, "input", slots) + return + for item in request_input if isinstance(request_input, list) else (): + if isinstance(item, dict): + _collect_content(item, slots) + class LLMShieldGuardrail(CustomGuardrail): """Redacts PII before it leaves the proxy and restores it in the response. @@ -164,67 +215,50 @@ class LLMShieldGuardrail(CustomGuardrail): # --- session ------------------------------------------------------------------ - def _session_id(self, data: MutableRequest) -> str: - """Returns a session id stable across this request's hooks.""" + @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("metadata", {}) # mutable-ok: per-request store. - if not isinstance(metadata, dict): - return f"litellm-{uuid.uuid4().hex}" - existing: Final = metadata.get(_SESSION_METADATA_KEY) - if isinstance(existing, str) and existing: - return existing - session_id: Final = get_session_id_from_request_data(data) or f"litellm-{uuid.uuid4().hex}" - metadata[_SESSION_METADATA_KEY] = session_id + if isinstance(metadata, dict): + metadata[_SESSION_METADATA_KEY] = session_id return session_id + @staticmethod + def _session_id(data: MutableRequest) -> str: + """Reads back the vault id minted while redacting this request. + + Falls back to an unused id rather than to anything the caller supplied: a + reply that cannot be restored is a visible placeholder, while trusting a + caller-supplied id would hand them someone else's plaintext. + """ + metadata: Final = data.get("metadata") + existing: Final = metadata.get(_SESSION_METADATA_KEY) if isinstance(metadata, dict) else None + if isinstance(existing, str) and existing.startswith(_VAULT_PREFIX): + return existing + return f"{_VAULT_PREFIX}-{uuid.uuid4().hex}" + # --- request traversal -------------------------------------------------------- @staticmethod def _locate_request_texts(data: MutableRequest) -> Sequence[_Slot]: """Finds every redactable span in an outbound request. - Returns ``(text, write)`` pairs. Any shape missed here reaches the provider - in the clear, so this walks all of the request shapes that carry caller text: - - - chat ``messages``, both string and multimodal list ``content`` - - tool call ``arguments``, which routinely carry the values a user asked - the model to look up - - the Responses API ``input``, as a bare string or a list of items + 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 caller text. """ - slots: Final[list[_Slot]] = [] # mutable-ok: accumulator, frozen on return. - - def add(container: MutableRequest, key: str, value: object) -> None: - if isinstance(value, str) and value: - slots.append((value, lambda new, c=container, k=key: c.__setitem__(k, new))) - - def add_content(container: MutableRequest) -> None: - """Adds `content`, which is either a string or a list of typed parts.""" - content: Final = container.get("content") - if isinstance(content, str): - add(container, "content", content) - return - for part in content if isinstance(content, list) else (): - if isinstance(part, dict): - add(part, "text", part.get("text")) - - def add_tool_calls(message: MutableRequest) -> None: - for tool_call in message.get("tool_calls") or (): - function = tool_call.get("function") if isinstance(tool_call, dict) else None - if isinstance(function, dict): - add(function, "arguments", function.get("arguments")) - + slots: Final[_SlotSink] = [] # mutable-ok: accumulator, frozen on return. for message in data.get("messages") or (): if isinstance(message, dict): - add_content(message) - add_tool_calls(message) - - request_input: Final = data.get("input") - if isinstance(request_input, str): - add(data, "input", request_input) - else: - for item in request_input if isinstance(request_input, list) else (): - if isinstance(item, dict): - add_content(item) - + _collect_content(message, slots) + _collect_tool_arguments(message, slots) + _collect_responses_fields(data, slots) return tuple(slots) # --- hooks -------------------------------------------------------------------- @@ -245,7 +279,7 @@ class LLMShieldGuardrail(CustomGuardrail): if not slots: return data - redacted: Final = await self._redact(tuple(text for text, _ in slots), self._session_id(data)) + redacted: Final = await self._redact(tuple(text for text, _ in slots), self._mint_session_id(data)) for (_, write), replacement in zip(slots, redacted): write(replacement) return data @@ -451,11 +485,10 @@ class LLMShieldGuardrail(CustomGuardrail): if not texts: return inputs - session_id: Final = self._session_id(request_data) replaced: Final = ( - await self._redact(tuple(texts), session_id) + await self._redact(tuple(texts), self._mint_session_id(request_data)) if input_type == "request" - else await self._rehydrate(tuple(texts), session_id) + else await self._rehydrate(tuple(texts), self._session_id(request_data)) ) # Return a new mapping rather than rewriting the caller's, so this stays a # pure transform of the inputs it was handed. diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py index 7a07c9761a8..45d45301858 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py @@ -223,6 +223,35 @@ class TestRequestCoverage: 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_every_shape_in_one_request_is_redacted(self): guardrail = _guardrail() @@ -334,6 +363,59 @@ class TestRestoration: 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["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 + + class TestFailClosed: @pytest.mark.asyncio async def test_unreachable_shield_blocks_the_request(self):