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.
This commit is contained in:
Ninad Phalak 2026-10-04 10:07:14 +00:00
parent 1b45935b80
commit 71e68fd15e
No known key found for this signature in database
2 changed files with 355 additions and 25 deletions

View file

@ -39,11 +39,15 @@ from litellm.llms.custom_httpx.http_handler import (
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import GenericGuardrailAPIInputs
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
GUARDRAIL_NAME: Final = "llm_shield_proxy"
@ -202,6 +206,15 @@ def _as_array(value: object) -> MutableSeq | None:
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))
@ -232,11 +245,15 @@ def _collect_prompt(data: MutableRequest, slots: _SlotSink) -> None:
if prompt_object is not None:
# A Responses API PromptObject. `variables` are substituted into the stored
# prompt on the provider side, so they are caller text. `id` and `version`
# identify which prompt to use and must arrive unchanged.
# identify which prompt to use and must arrive unchanged. A variable is a string
# or a typed input such as `{"type": "input_text", "text": ...}`.
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:
@ -327,7 +344,9 @@ def _collect_responses_fields(data: MutableRequest, slots: _SlotSink, privileged
"""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.
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")
@ -345,10 +364,15 @@ def _collect_responses_fields(data: MutableRequest, slots: _SlotSink, privileged
item = _as_object(entry)
if item is None:
continue
_collect_content(item, slots)
# A function_call item holds `arguments`; a function_call_output holds `output`.
_collect_content(item, privileged if item.get("role") in _PRIVILEGED_ROLES else slots)
# A function_call item holds `arguments`; a function_call_output holds `output`,
# as a string or as a list of input_text parts. A custom_tool_call holds `input`
# and a code_interpreter_call `code` -- the fields the reply side restores.
_collect(item, "arguments", slots)
_collect(item, "output", slots)
_collect_text_parts(item, "output", slots)
_collect(item, "input", slots)
_collect(item, "code", slots)
# A replayed reasoning item carries the model's summary of its own reasoning,
# which quotes whatever the conversation contained.
_collect_text_parts(item, "summary", slots)
@ -946,6 +970,30 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
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_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.
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.
"""
return None
# --- transport ---------------------------------------------------------------
def _headers(self, session_id: str) -> JsonBody:
@ -1143,9 +1191,19 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
user_api_key_dict: UserAPIKeyAuth,
response: Any,
) -> Any:
"""Restores the original values in a non-streaming response."""
"""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)
@ -1168,6 +1226,10 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
for choice in choices:
message = getattr(choice, "message", None)
if message is None:
# A Completions reply carries its text on the choice itself.
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:
@ -1290,14 +1352,18 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
for frames in await sse.feed(chunk):
yield frames
continue
# Chunks are restored as copies, for the reason the non-streaming hook copies:
# LiteLLM keeps the chunks it yielded to assemble the reply it caches and logs,
# so restoring them in place would cache this caller's plaintext.
if _responses_event_type(chunk) is not None:
for event in await events.restore(chunk):
for event in await events.restore(_detached(chunk)):
yield event
continue
last_chunk = chunk
for choice in getattr(chunk, "choices", None) or ():
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 chunk
yield restored_chunk
# A stream that ended early can still leave text held back, in any shape.
for frames in await sse.finish():
@ -1308,7 +1374,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
async for trailing in self._flush_trailing(last_chunk, carries, session_id):
yield trailing
async def _restore_choice(self, choice: Any, carries: _CarryWindows, session_id: str) -> None:
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
@ -1317,10 +1383,15 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
held back for one stream onto another.
"""
delta: Final = getattr(choice, "delta", None)
if delta is None:
return
index: Final = _choice_index(choice)
is_final: Final = bool(getattr(choice, "finish_reason", None))
if isinstance(choice, TextChoices):
# A Completions stream carries its text on the choice itself, with no delta
# and no tool calls: one window, the content one.
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)
@ -1333,6 +1404,25 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
# after it produces argument JSON the client has already stopped waiting for.
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,
@ -1446,7 +1536,9 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
chunk = self._chunk_for_choice(last_chunk, choice_index)
if chunk is None:
continue
if tool_index is None:
if isinstance(chunk.choices[0], TextChoices):
chunk.choices[0].text = text
elif tool_index is None:
chunk.choices[0].delta.content = text
else:
# The copy carried this chunk's own content and tool calls, both already
@ -1470,7 +1562,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
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:
if getattr(kept, "delta", None) is None and not isinstance(kept, TextChoices):
return None
kept.index = index
# The terminal signal, if there was one, already went out with the real chunk.

View file

@ -19,7 +19,16 @@ from litellm.types.llms.openai import (
OutputTextDoneEvent,
ResponsesAPIStreamEvents,
)
from litellm.types.utils import Choices, Delta, Message, ModelResponse, ModelResponseStream, StreamingChoices
from litellm.types.utils import (
Choices,
Delta,
Message,
ModelResponse,
ModelResponseStream,
StreamingChoices,
TextChoices,
TextCompletionResponse,
)
def _guardrail(**overrides: object) -> LLMShieldProxyGuardrail:
@ -387,6 +396,53 @@ class TestRequestCoverage:
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_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."""
@ -551,6 +607,26 @@ class TestRequestCoverage:
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."""
@ -720,9 +796,9 @@ class TestRequestCoverage:
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]
await guardrail.async_post_call_success_hook(data=data, user_api_key_dict=None, response=reply)
restored = await guardrail.async_post_call_success_hook(data=data, user_api_key_dict=None, response=reply)
assert json.loads(call.function.arguments) == {"to": "ops@example.com"}
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"}
@ -773,9 +849,11 @@ class TestRestoration:
block = SimpleNamespace(text="[EMAIL_1]")
response = SimpleNamespace(output=[SimpleNamespace(content=[block])])
await guardrail.async_post_call_success_hook(data={"messages": []}, user_api_key_dict=None, response=response)
result = await guardrail.async_post_call_success_hook(
data={"messages": []}, user_api_key_dict=None, response=response
)
assert block.text == "a@b.com"
assert result.output[0].content[0].text == "a@b.com"
@pytest.mark.asyncio
async def test_anthropic_message_shape_is_restored(self):
@ -1027,6 +1105,24 @@ class TestVaultIsolation:
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
@ -1522,10 +1618,11 @@ class TestResponsesStreamRestoration:
response=SimpleNamespace(output=[SimpleNamespace(content=[block]), call]),
)
await _restore_stream(guardrail, [completed])
(out,) = await _restore_stream(guardrail, [completed])
assert block["text"] == "Mail a@example.com"
assert call.arguments == '{"to": "a@example.com"}'
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):
@ -1552,9 +1649,9 @@ class TestResponsesStreamRestoration:
type="response.reasoning_summary_part.done", item_id="rs_1", output_index=0, summary_index=0, part=part
)
await _restore_stream(guardrail, [event])
(out,) = await _restore_stream(guardrail, [event])
assert part.text == "asked about a@example.com"
assert out.part.text == "asked about a@example.com"
@pytest.mark.asyncio
async def test_mcp_call_arguments_are_restored(self):
@ -1587,3 +1684,144 @@ class TestResponsesStreamRestoration:
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 = litellm.caching.caching.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_the_reply_for_the_cache_untouched(self):
"""LiteLLM caches what the deployment hook returns, so restoring there caches plaintext.
The proxy's post-call hook, which runs after the cache write, restores model-level
guardrails instead.
"""
guardrail, shield = _shielded({"[EMAIL_1]": "a@example.com"})
reply = ModelResponse(choices=[Choices(index=0, message=Message(role="assistant", content="[EMAIL_1]"))])
result = await guardrail.async_post_call_success_deployment_hook(
request_data={"messages": [], "guardrails": [GUARDRAIL_NAME]}, response=reply, call_type=None
)
assert result is None
assert reply.choices[0].message.content == "[EMAIL_1]"
assert shield.urls == []