mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
1b45935b80
commit
71e68fd15e
2 changed files with 355 additions and 25 deletions
|
|
@ -39,11 +39,15 @@ from litellm.llms.custom_httpx.http_handler import (
|
||||||
)
|
)
|
||||||
from litellm.proxy._types import UserAPIKeyAuth
|
from litellm.proxy._types import UserAPIKeyAuth
|
||||||
from litellm.types.guardrails import GuardrailEventHooks
|
from litellm.types.guardrails import GuardrailEventHooks
|
||||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
from litellm.types.utils import GenericGuardrailAPIInputs, TextChoices
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from litellm.caching.caching import DualCache
|
from litellm.caching.caching import DualCache
|
||||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
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"
|
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
|
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:
|
def _is_container(value: object) -> bool:
|
||||||
"""Whether `value` is a JSON object or array, without narrowing it to unknown types."""
|
"""Whether `value` is a JSON object or array, without narrowing it to unknown types."""
|
||||||
return isinstance(value, (dict, list))
|
return isinstance(value, (dict, list))
|
||||||
|
|
@ -232,11 +245,15 @@ def _collect_prompt(data: MutableRequest, slots: _SlotSink) -> None:
|
||||||
if prompt_object is not None:
|
if prompt_object is not None:
|
||||||
# A Responses API PromptObject. `variables` are substituted into the stored
|
# A Responses API PromptObject. `variables` are substituted into the stored
|
||||||
# prompt on the provider side, so they are caller text. `id` and `version`
|
# 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"))
|
variables: Final = _as_object(prompt_object.get("variables"))
|
||||||
if variables is not None:
|
if variables is not None:
|
||||||
for name in tuple(variables):
|
for name in tuple(variables):
|
||||||
_collect(variables, name, slots)
|
_collect(variables, name, slots)
|
||||||
|
typed = _as_object(variables[name])
|
||||||
|
if typed is not None:
|
||||||
|
_collect(typed, "text", slots)
|
||||||
return
|
return
|
||||||
entries: Final = _as_array(prompt)
|
entries: Final = _as_array(prompt)
|
||||||
if entries is None:
|
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`.
|
"""The Responses API sends text outside `messages`, in `instructions` and `input`.
|
||||||
|
|
||||||
`instructions` is written by the application, not by the caller, so it is
|
`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)
|
_collect(data, "instructions", privileged)
|
||||||
request_input: Final = data.get("input")
|
request_input: Final = data.get("input")
|
||||||
|
|
@ -345,10 +364,15 @@ def _collect_responses_fields(data: MutableRequest, slots: _SlotSink, privileged
|
||||||
item = _as_object(entry)
|
item = _as_object(entry)
|
||||||
if item is None:
|
if item is None:
|
||||||
continue
|
continue
|
||||||
_collect_content(item, slots)
|
_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`.
|
# 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, "arguments", slots)
|
||||||
_collect(item, "output", 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,
|
# A replayed reasoning item carries the model's summary of its own reasoning,
|
||||||
# which quotes whatever the conversation contained.
|
# which quotes whatever the conversation contained.
|
||||||
_collect_text_parts(item, "summary", slots)
|
_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.
|
def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: # mutable-ok: parent's signature.
|
||||||
return [GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call]
|
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 ---------------------------------------------------------------
|
# --- transport ---------------------------------------------------------------
|
||||||
|
|
||||||
def _headers(self, session_id: str) -> JsonBody:
|
def _headers(self, session_id: str) -> JsonBody:
|
||||||
|
|
@ -1143,9 +1191,19 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
|
||||||
user_api_key_dict: UserAPIKeyAuth,
|
user_api_key_dict: UserAPIKeyAuth,
|
||||||
response: Any,
|
response: Any,
|
||||||
) -> 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:
|
if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.post_call) is not True:
|
||||||
return response
|
return response
|
||||||
|
response = _detached(response) # rebind-ok: everything below restores the copy.
|
||||||
|
|
||||||
if self._is_anthropic_message_response(response):
|
if self._is_anthropic_message_response(response):
|
||||||
return await self._restore_anthropic_response(response, data)
|
return await self._restore_anthropic_response(response, data)
|
||||||
|
|
@ -1168,6 +1226,10 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
|
||||||
for choice in choices:
|
for choice in choices:
|
||||||
message = getattr(choice, "message", None)
|
message = getattr(choice, "message", None)
|
||||||
if message is 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
|
continue
|
||||||
content = getattr(message, "content", None)
|
content = getattr(message, "content", None)
|
||||||
if isinstance(content, str) and content:
|
if isinstance(content, str) and content:
|
||||||
|
|
@ -1290,14 +1352,18 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
|
||||||
for frames in await sse.feed(chunk):
|
for frames in await sse.feed(chunk):
|
||||||
yield frames
|
yield frames
|
||||||
continue
|
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:
|
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
|
yield event
|
||||||
continue
|
continue
|
||||||
last_chunk = chunk
|
restored_chunk = _detached(chunk)
|
||||||
for choice in getattr(chunk, "choices", None) or ():
|
last_chunk = restored_chunk
|
||||||
|
for choice in getattr(restored_chunk, "choices", None) or ():
|
||||||
await self._restore_choice(choice, carries, session_id)
|
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.
|
# A stream that ended early can still leave text held back, in any shape.
|
||||||
for frames in await sse.finish():
|
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):
|
async for trailing in self._flush_trailing(last_chunk, carries, session_id):
|
||||||
yield trailing
|
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.
|
"""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
|
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.
|
held back for one stream onto another.
|
||||||
"""
|
"""
|
||||||
delta: Final = getattr(choice, "delta", None)
|
delta: Final = getattr(choice, "delta", None)
|
||||||
if delta is None:
|
|
||||||
return
|
|
||||||
index: Final = _choice_index(choice)
|
index: Final = _choice_index(choice)
|
||||||
is_final: Final = bool(getattr(choice, "finish_reason", None))
|
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)
|
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.
|
# after it produces argument JSON the client has already stopped waiting for.
|
||||||
await self._flush_finished_choice(delta, index, carries, session_id)
|
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(
|
async def _restore_content_window(
|
||||||
self,
|
self,
|
||||||
delta: Any,
|
delta: Any,
|
||||||
|
|
@ -1446,7 +1536,9 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
|
||||||
chunk = self._chunk_for_choice(last_chunk, choice_index)
|
chunk = self._chunk_for_choice(last_chunk, choice_index)
|
||||||
if chunk is None:
|
if chunk is None:
|
||||||
continue
|
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
|
chunk.choices[0].delta.content = text
|
||||||
else:
|
else:
|
||||||
# The copy carried this chunk's own content and tool calls, both already
|
# 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)
|
choices: Final[tuple[object, ...]] = tuple(raw_choices)
|
||||||
position: Final = next((at for at, choice in enumerate(choices) if _choice_index(choice) == index), 0)
|
position: Final = next((at for at, choice in enumerate(choices) if _choice_index(choice) == index), 0)
|
||||||
kept: Final = raw_choices[position]
|
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
|
return None
|
||||||
kept.index = index
|
kept.index = index
|
||||||
# The terminal signal, if there was one, already went out with the real chunk.
|
# The terminal signal, if there was one, already went out with the real chunk.
|
||||||
|
|
|
||||||
|
|
@ -19,7 +19,16 @@ from litellm.types.llms.openai import (
|
||||||
OutputTextDoneEvent,
|
OutputTextDoneEvent,
|
||||||
ResponsesAPIStreamEvents,
|
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:
|
def _guardrail(**overrides: object) -> LLMShieldProxyGuardrail:
|
||||||
|
|
@ -387,6 +396,53 @@ class TestRequestCoverage:
|
||||||
assert data["input"][0]["arguments"] == '{"email": "[EMAIL_1]"}'
|
assert data["input"][0]["arguments"] == '{"email": "[EMAIL_1]"}'
|
||||||
assert data["input"][1]["output"] == "sent to [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
|
@pytest.mark.asyncio
|
||||||
async def test_anthropic_system_prompt_is_redacted(self):
|
async def test_anthropic_system_prompt_is_redacted(self):
|
||||||
"""/v1/messages carries its system prompt at the top level, not in messages."""
|
"""/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"]["id"] == "pmpt_123"
|
||||||
assert data["prompt"]["version"] == "2"
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_completions_suffix_is_redacted(self):
|
async def test_completions_suffix_is_redacted(self):
|
||||||
"""LiteLLM forwards the legacy `suffix` to providers that support it."""
|
"""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]"}'))
|
call = SimpleNamespace(function=SimpleNamespace(name="notify", arguments='{"to": "[EMAIL_1]"}'))
|
||||||
reply = ModelResponse(choices=[Choices(message=Message(content=None, tool_calls=None))])
|
reply = ModelResponse(choices=[Choices(message=Message(content=None, tool_calls=None))])
|
||||||
reply.choices[0].message.tool_calls = [call]
|
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):
|
def test_schema_nesting_past_the_bound_is_refused(self):
|
||||||
schema: dict = {"type": "object", "description": "past-the-bound@example.com"}
|
schema: dict = {"type": "object", "description": "past-the-bound@example.com"}
|
||||||
|
|
@ -773,9 +849,11 @@ class TestRestoration:
|
||||||
|
|
||||||
block = SimpleNamespace(text="[EMAIL_1]")
|
block = SimpleNamespace(text="[EMAIL_1]")
|
||||||
response = SimpleNamespace(output=[SimpleNamespace(content=[block])])
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_anthropic_message_shape_is_restored(self):
|
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"]
|
privileged_id = mock.call_args_list[0].kwargs["headers"]["X-Session-ID"]
|
||||||
assert privileged_id not in json.dumps(data, default=str)
|
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:
|
class TestFailClosed:
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|
@ -1522,10 +1618,11 @@ class TestResponsesStreamRestoration:
|
||||||
response=SimpleNamespace(output=[SimpleNamespace(content=[block]), call]),
|
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"
|
restored_block, restored_call = out.response.output[0].content[0], out.response.output[1]
|
||||||
assert call.arguments == '{"to": "a@example.com"}'
|
assert restored_block["text"] == "Mail a@example.com"
|
||||||
|
assert restored_call.arguments == '{"to": "a@example.com"}'
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_streams_on_different_parts_do_not_share_a_window(self):
|
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
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_mcp_call_arguments_are_restored(self):
|
async def test_mcp_call_arguments_are_restored(self):
|
||||||
|
|
@ -1587,3 +1684,144 @@ class TestResponsesStreamRestoration:
|
||||||
|
|
||||||
assert out == [audio]
|
assert out == [audio]
|
||||||
assert shield.urls == []
|
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 == []
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue