mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(guardrails): mint the vault id instead of trusting the caller's
The vault id was taken from caller-supplied session metadata, and every caller shares one LLM Shield key. Someone who knew or guessed another caller's session id could send a placeholder, have the model echo it back, and get that caller's plaintext restored into their own reply. Vault ids are now minted per request behind a per-process prefix, so a caller cannot name a vault this process uses. Redaction mints, restoration reads back, and a reply whose id does not match is left holding its placeholders rather than resolved against some other vault. Also covers two more request fields that were reaching the provider intact: the Responses API `instructions`, and the legacy `function_call.arguments` alongside `tool_calls`. The collectors move to module level, which drops the traversal back under the complexity limit and lets the code carry its own explanation instead of the comments that were restating it.
This commit is contained in:
parent
295ad527d5
commit
47421b7541
2 changed files with 169 additions and 54 deletions
|
|
@ -24,7 +24,6 @@ from litellm._logging import verbose_proxy_logger
|
||||||
from litellm.exceptions import GuardrailRaisedException
|
from litellm.exceptions import GuardrailRaisedException
|
||||||
from litellm.integrations.custom_guardrail import (
|
from litellm.integrations.custom_guardrail import (
|
||||||
CustomGuardrail,
|
CustomGuardrail,
|
||||||
get_session_id_from_request_data,
|
|
||||||
log_guardrail_information,
|
log_guardrail_information,
|
||||||
)
|
)
|
||||||
from litellm.llms.custom_httpx.http_handler import (
|
from litellm.llms.custom_httpx.http_handler import (
|
||||||
|
|
@ -52,6 +51,13 @@ _REHYDRATE_STREAM_PATH: Final = "/v1/guard/rehydrate/stream"
|
||||||
# across concurrent requests.
|
# across concurrent requests.
|
||||||
_SESSION_METADATA_KEY: Final = "llm_shield_session_id"
|
_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
|
_DEFAULT_TIMEOUT_SECONDS: Final = 10.0
|
||||||
|
|
||||||
# The proxy's own request dict. Mutable by design: a pre-call guardrail rewrites
|
# 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.
|
# replacement back where it came from.
|
||||||
_Slot: TypeAlias = tuple[str, Callable[[str], None]] # mutable-ok: Callable's param list.
|
_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):
|
class LLMShieldGuardrail(CustomGuardrail):
|
||||||
"""Redacts PII before it leaves the proxy and restores it in the response.
|
"""Redacts PII before it leaves the proxy and restores it in the response.
|
||||||
|
|
@ -164,67 +215,50 @@ class LLMShieldGuardrail(CustomGuardrail):
|
||||||
|
|
||||||
# --- session ------------------------------------------------------------------
|
# --- session ------------------------------------------------------------------
|
||||||
|
|
||||||
def _session_id(self, data: MutableRequest) -> str:
|
@staticmethod
|
||||||
"""Returns a session id stable across this request's hooks."""
|
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.
|
metadata: Final = data.setdefault("metadata", {}) # mutable-ok: per-request store.
|
||||||
if not isinstance(metadata, dict):
|
if isinstance(metadata, dict):
|
||||||
return f"litellm-{uuid.uuid4().hex}"
|
metadata[_SESSION_METADATA_KEY] = session_id
|
||||||
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
|
|
||||||
return 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 --------------------------------------------------------
|
# --- request traversal --------------------------------------------------------
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _locate_request_texts(data: MutableRequest) -> Sequence[_Slot]:
|
def _locate_request_texts(data: MutableRequest) -> Sequence[_Slot]:
|
||||||
"""Finds every redactable span in an outbound request.
|
"""Finds every redactable span in an outbound request.
|
||||||
|
|
||||||
Returns ``(text, write)`` pairs. Any shape missed here reaches the provider
|
Anything missed here reaches the provider in the clear while the guardrail
|
||||||
in the clear, so this walks all of the request shapes that carry caller text:
|
still reports as enabled, so the walk covers every request shape that
|
||||||
|
carries 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
|
|
||||||
"""
|
"""
|
||||||
slots: Final[list[_Slot]] = [] # mutable-ok: accumulator, frozen on return.
|
slots: Final[_SlotSink] = [] # 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"))
|
|
||||||
|
|
||||||
for message in data.get("messages") or ():
|
for message in data.get("messages") or ():
|
||||||
if isinstance(message, dict):
|
if isinstance(message, dict):
|
||||||
add_content(message)
|
_collect_content(message, slots)
|
||||||
add_tool_calls(message)
|
_collect_tool_arguments(message, slots)
|
||||||
|
_collect_responses_fields(data, slots)
|
||||||
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)
|
|
||||||
|
|
||||||
return tuple(slots)
|
return tuple(slots)
|
||||||
|
|
||||||
# --- hooks --------------------------------------------------------------------
|
# --- hooks --------------------------------------------------------------------
|
||||||
|
|
@ -245,7 +279,7 @@ class LLMShieldGuardrail(CustomGuardrail):
|
||||||
if not slots:
|
if not slots:
|
||||||
return data
|
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):
|
for (_, write), replacement in zip(slots, redacted):
|
||||||
write(replacement)
|
write(replacement)
|
||||||
return data
|
return data
|
||||||
|
|
@ -451,11 +485,10 @@ class LLMShieldGuardrail(CustomGuardrail):
|
||||||
if not texts:
|
if not texts:
|
||||||
return inputs
|
return inputs
|
||||||
|
|
||||||
session_id: Final = self._session_id(request_data)
|
|
||||||
replaced: Final = (
|
replaced: Final = (
|
||||||
await self._redact(tuple(texts), session_id)
|
await self._redact(tuple(texts), self._mint_session_id(request_data))
|
||||||
if input_type == "request"
|
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
|
# Return a new mapping rather than rewriting the caller's, so this stays a
|
||||||
# pure transform of the inputs it was handed.
|
# pure transform of the inputs it was handed.
|
||||||
|
|
|
||||||
|
|
@ -223,6 +223,35 @@ class TestRequestCoverage:
|
||||||
|
|
||||||
assert data["messages"][0]["tool_calls"][0]["function"]["arguments"] == '{"email": "[EMAIL_1]"}'
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_every_shape_in_one_request_is_redacted(self):
|
async def test_every_shape_in_one_request_is_redacted(self):
|
||||||
guardrail = _guardrail()
|
guardrail = _guardrail()
|
||||||
|
|
@ -334,6 +363,59 @@ class TestRestoration:
|
||||||
assert result["content"][1] == {"type": "tool_use", "id": "t1", "name": "lookup", "input": {}}
|
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:
|
class TestFailClosed:
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_unreachable_shield_blocks_the_request(self):
|
async def test_unreachable_shield_blocks_the_request(self):
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue