mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(guardrails): redact every request shape and restore every reply shape
Three gaps, all of which let an enabled guardrail hand data to the provider or hand placeholders to the caller. Requests only walked `messages`. The Responses API `input` and tool call `arguments` went out untouched. Measured against a live provider: a request sent through `/v1/responses` reached the model with the real address in it while the guardrail reported as enabled. Request traversal now covers chat content (string and multimodal), tool call arguments, and `input` as a bare string or a list of items. Fixing that exposed the matching gap on the way back: the Responses API reply carries `output` items rather than `choices`, so it returned to the caller still holding placeholders. It now gets its own walk, handling text blocks as dicts or objects. The dashboard preset seeded only pre_call, so a guardrail created from the UI would redact the request and return the placeholders to the user. Presets can now seed both modes; the form already normalised either shape. Adds tests for each request shape, for both Responses API reply forms, and replaces a test that had asserted the `input` bypass as correct behaviour.
This commit is contained in:
parent
0de4b8b9a8
commit
b6e3e6decd
5 changed files with 323 additions and 110 deletions
|
|
@ -7,7 +7,7 @@
|
|||
|
||||
import os
|
||||
import uuid
|
||||
from collections.abc import AsyncGenerator
|
||||
from collections.abc import AsyncGenerator, Callable, Mapping, Sequence
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any, # noqa: TID251 # **kwargs forwards verbatim to CustomGuardrail.__init__
|
||||
|
|
@ -15,6 +15,7 @@ from typing import (
|
|||
Final,
|
||||
Literal,
|
||||
Optional,
|
||||
TypeAlias,
|
||||
)
|
||||
|
||||
import httpx
|
||||
|
|
@ -53,6 +54,19 @@ _SESSION_METADATA_KEY: Final = "llm_shield_session_id"
|
|||
|
||||
_DEFAULT_TIMEOUT_SECONDS: Final = 10.0
|
||||
|
||||
# The proxy's own request dict. Mutable by design: a pre-call guardrail rewrites
|
||||
# the caller's payload in place, which is the entire point of the hook.
|
||||
# mutable-ok: the shape is fixed by CustomLogger's hook signatures.
|
||||
MutableRequest: TypeAlias = dict
|
||||
|
||||
# A JSON body on its way to httpx, which requires a real dict rather than a view.
|
||||
# mutable-ok: handed straight to the HTTP client.
|
||||
JsonBody: TypeAlias = dict
|
||||
|
||||
# One redactable span: the text as it stands, and the write that puts the
|
||||
# replacement back where it came from.
|
||||
_Slot: TypeAlias = tuple[str, Callable[[str], None]] # mutable-ok: Callable's param list.
|
||||
|
||||
|
||||
class LLMShieldGuardrail(CustomGuardrail):
|
||||
"""Redacts PII before it leaves the proxy and restores it in the response.
|
||||
|
|
@ -78,7 +92,7 @@ class LLMShieldGuardrail(CustomGuardrail):
|
|||
guardrail_name: str = GUARDRAIL_NAME,
|
||||
api_base: str | None = None,
|
||||
api_key: str | None = None,
|
||||
**kwargs: Any,
|
||||
**kwargs: Any, # noqa: LIT008 # kwargs-ok: forwarded verbatim to CustomGuardrail.__init__
|
||||
) -> None:
|
||||
self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
|
||||
self.api_base: Final = (api_base or os.environ.get("LLM_SHIELD_API_BASE") or _DEFAULT_API_BASE).rstrip("/")
|
||||
|
|
@ -86,18 +100,21 @@ class LLMShieldGuardrail(CustomGuardrail):
|
|||
super().__init__(guardrail_name=guardrail_name, **kwargs)
|
||||
|
||||
@classmethod
|
||||
def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
|
||||
return [GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call]
|
||||
def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: # mutable-ok: parent's signature.
|
||||
return [GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call] # mutable-ok: parent's signature.
|
||||
|
||||
# --- transport ---------------------------------------------------------------
|
||||
|
||||
def _headers(self, session_id: str) -> dict:
|
||||
headers = {"Content-Type": "application/json", "X-Session-ID": session_id}
|
||||
def _headers(self, session_id: str) -> JsonBody:
|
||||
headers: Final[JsonBody] = { # mutable-ok: httpx requires a real dict.
|
||||
"Content-Type": "application/json",
|
||||
"X-Session-ID": session_id,
|
||||
}
|
||||
if self.api_key:
|
||||
headers["Authorization"] = f"Bearer {self.api_key}"
|
||||
return headers
|
||||
|
||||
async def _call_shield(self, path: str, session_id: str, payload: dict) -> dict:
|
||||
async def _call_shield(self, path: str, session_id: str, payload: JsonBody) -> Mapping[str, object]:
|
||||
"""Posts to LLM Shield, failing closed on any transport or status error.
|
||||
|
||||
A redaction guardrail that fails open sends the very data it exists to
|
||||
|
|
@ -105,7 +122,7 @@ class LLMShieldGuardrail(CustomGuardrail):
|
|||
blocks the request instead of passing it through.
|
||||
"""
|
||||
try:
|
||||
response = await self.async_handler.post(
|
||||
response: Final = await self.async_handler.post(
|
||||
f"{self.api_base}{path}",
|
||||
headers=self._headers(session_id),
|
||||
json=payload,
|
||||
|
|
@ -126,69 +143,89 @@ class LLMShieldGuardrail(CustomGuardrail):
|
|||
message="LLM Shield is unreachable; blocking the request.",
|
||||
) from exc
|
||||
|
||||
async def _redact(self, texts: list, session_id: str) -> list:
|
||||
body = await self._call_shield(_REDACT_PATH, session_id, {"texts": texts})
|
||||
async def _redact(self, texts: Sequence[str], session_id: str) -> Sequence[str]:
|
||||
payload: Final[JsonBody] = {"texts": list(texts)} # mutable-ok: JSON body for httpx.
|
||||
body: Final = await self._call_shield(_REDACT_PATH, session_id, payload)
|
||||
return self._same_length_or_raise(body.get("texts"), texts, "redact")
|
||||
|
||||
async def _rehydrate(self, texts: list, session_id: str) -> list:
|
||||
body = await self._call_shield(_REHYDRATE_PATH, session_id, {"texts": texts})
|
||||
async def _rehydrate(self, texts: Sequence[str], session_id: str) -> Sequence[str]:
|
||||
payload: Final[JsonBody] = {"texts": list(texts)} # mutable-ok: JSON body for httpx.
|
||||
body: Final = await self._call_shield(_REHYDRATE_PATH, session_id, payload)
|
||||
return self._same_length_or_raise(body.get("texts"), texts, "rehydrate")
|
||||
|
||||
def _same_length_or_raise(self, returned: Any, sent: list, operation: str) -> list:
|
||||
def _same_length_or_raise(self, returned: object, sent: Sequence[str], operation: str) -> Sequence[str]:
|
||||
"""Guards the positional mapping the callers rely on to write results back."""
|
||||
if not isinstance(returned, list) or len(returned) != len(sent):
|
||||
raise GuardrailRaisedException(
|
||||
guardrail_name=self.guardrail_name,
|
||||
message=f"LLM Shield {operation} returned an unexpected payload; blocking the request.",
|
||||
)
|
||||
return returned
|
||||
return tuple(returned)
|
||||
|
||||
# --- session ------------------------------------------------------------------
|
||||
|
||||
def _session_id(self, data: dict) -> str:
|
||||
def _session_id(self, data: MutableRequest) -> str:
|
||||
"""Returns a session id stable across this request's hooks."""
|
||||
metadata = data.setdefault("metadata", {})
|
||||
metadata: Final = data.setdefault("metadata", {}) # mutable-ok: per-request store.
|
||||
if not isinstance(metadata, dict):
|
||||
return f"litellm-{uuid.uuid4().hex}"
|
||||
existing = metadata.get(_SESSION_METADATA_KEY)
|
||||
existing: Final = metadata.get(_SESSION_METADATA_KEY)
|
||||
if isinstance(existing, str) and existing:
|
||||
return existing
|
||||
session_id = get_session_id_from_request_data(data) or f"litellm-{uuid.uuid4().hex}"
|
||||
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
|
||||
|
||||
# --- message traversal --------------------------------------------------------
|
||||
# --- request traversal --------------------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
def _locate_texts(messages: list) -> list:
|
||||
"""Finds every text span in a message list.
|
||||
def _locate_request_texts(data: MutableRequest) -> Sequence[_Slot]:
|
||||
"""Finds every redactable span in an outbound request.
|
||||
|
||||
Returns ``(message_index, part_index_or_None, text)``. The list form is the
|
||||
multimodal shape, where only ``text`` parts carry redactable content.
|
||||
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
|
||||
"""
|
||||
located = []
|
||||
for message_index, message in enumerate(messages):
|
||||
if not isinstance(message, dict):
|
||||
continue
|
||||
content = message.get("content")
|
||||
if isinstance(content, str) and content:
|
||||
located.append((message_index, None, content))
|
||||
elif isinstance(content, list):
|
||||
for part_index, part in enumerate(content):
|
||||
if not isinstance(part, dict) or part.get("type") != "text":
|
||||
continue
|
||||
text = part.get("text")
|
||||
if isinstance(text, str) and text:
|
||||
located.append((message_index, part_index, text))
|
||||
return located
|
||||
slots: Final[list[_Slot]] = [] # mutable-ok: accumulator, frozen on return.
|
||||
|
||||
@staticmethod
|
||||
def _write_back(messages: list, located: list, replacements: list) -> None:
|
||||
for (message_index, part_index, _), replacement in zip(located, replacements):
|
||||
if part_index is None:
|
||||
messages[message_index]["content"] = replacement
|
||||
else:
|
||||
messages[message_index]["content"][part_index]["text"] = replacement
|
||||
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 ():
|
||||
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)
|
||||
|
||||
return tuple(slots)
|
||||
|
||||
# --- hooks --------------------------------------------------------------------
|
||||
|
||||
|
|
@ -197,29 +234,26 @@ class LLMShieldGuardrail(CustomGuardrail):
|
|||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
cache: "DualCache",
|
||||
data: dict,
|
||||
data: MutableRequest,
|
||||
call_type: str,
|
||||
) -> dict | None:
|
||||
"""Replaces PII in the outbound messages with vault placeholders."""
|
||||
) -> MutableRequest | None:
|
||||
"""Replaces PII anywhere in the outbound request with vault placeholders."""
|
||||
if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_call) is not True:
|
||||
return data
|
||||
|
||||
messages = data.get("messages")
|
||||
if not isinstance(messages, list):
|
||||
slots: Final = self._locate_request_texts(data)
|
||||
if not slots:
|
||||
return data
|
||||
|
||||
located = self._locate_texts(messages)
|
||||
if not located:
|
||||
return data
|
||||
|
||||
redacted = await self._redact([text for _, _, text in located], self._session_id(data))
|
||||
self._write_back(messages, located, redacted)
|
||||
redacted: Final = await self._redact(tuple(text for text, _ in slots), self._session_id(data))
|
||||
for (_, write), replacement in zip(slots, redacted):
|
||||
write(replacement)
|
||||
return data
|
||||
|
||||
@log_guardrail_information
|
||||
async def async_post_call_success_hook(
|
||||
self,
|
||||
data: dict,
|
||||
data: MutableRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
response: Any,
|
||||
) -> Any:
|
||||
|
|
@ -230,27 +264,31 @@ class LLMShieldGuardrail(CustomGuardrail):
|
|||
if self._is_anthropic_message_response(response):
|
||||
return await self._restore_anthropic_response(response, data)
|
||||
|
||||
choices = getattr(response, "choices", None)
|
||||
text_blocks: Final = self._responses_api_text_blocks(response)
|
||||
if text_blocks:
|
||||
return await self._restore_responses_api_response(response, text_blocks, data)
|
||||
|
||||
choices: Final = getattr(response, "choices", None)
|
||||
if not choices:
|
||||
return response
|
||||
|
||||
pending = []
|
||||
for choice in choices:
|
||||
message = getattr(choice, "message", None)
|
||||
content = getattr(message, "content", None)
|
||||
if isinstance(content, str) and content:
|
||||
pending.append((message, content))
|
||||
|
||||
pending: Final = tuple(
|
||||
(choice.message, choice.message.content)
|
||||
for choice in choices
|
||||
if getattr(choice, "message", None) is not None
|
||||
and isinstance(getattr(choice.message, "content", None), str)
|
||||
and choice.message.content
|
||||
)
|
||||
if not pending:
|
||||
return response
|
||||
|
||||
restored = await self._rehydrate([text for _, text in pending], self._session_id(data))
|
||||
restored: Final = await self._rehydrate(tuple(text for _, text in pending), self._session_id(data))
|
||||
for (message, _), replacement in zip(pending, restored):
|
||||
message.content = replacement
|
||||
return response
|
||||
|
||||
@staticmethod
|
||||
def _is_anthropic_message_response(response: Any) -> bool:
|
||||
def _is_anthropic_message_response(response: object) -> bool:
|
||||
"""Anthropic's native /v1/messages reply arrives as a plain dict."""
|
||||
return (
|
||||
isinstance(response, dict)
|
||||
|
|
@ -258,30 +296,68 @@ class LLMShieldGuardrail(CustomGuardrail):
|
|||
and isinstance(response.get("content"), list)
|
||||
)
|
||||
|
||||
async def _restore_anthropic_response(self, response: dict, data: dict) -> dict:
|
||||
async def _restore_anthropic_response(self, response: MutableRequest, data: MutableRequest) -> MutableRequest:
|
||||
"""Restores text blocks in an Anthropic native message reply.
|
||||
|
||||
This shape has no `choices`, so without its own branch the reply would go
|
||||
back to the caller still carrying placeholders.
|
||||
"""
|
||||
blocks = [
|
||||
blocks: Final = tuple(
|
||||
block
|
||||
for block in response["content"]
|
||||
if isinstance(block, dict) and block.get("type") == "text" and isinstance(block.get("text"), str)
|
||||
]
|
||||
)
|
||||
if not blocks:
|
||||
return response
|
||||
|
||||
restored = await self._rehydrate([block["text"] for block in blocks], self._session_id(data))
|
||||
restored: Final = await self._rehydrate(tuple(block["text"] for block in blocks), self._session_id(data))
|
||||
for block, replacement in zip(blocks, restored):
|
||||
block["text"] = replacement
|
||||
return response
|
||||
|
||||
@staticmethod
|
||||
def _responses_api_text_blocks(response: object) -> Sequence[object]:
|
||||
"""Text blocks in a Responses API reply.
|
||||
|
||||
That shape carries `output` items rather than `choices`, so it needs its own
|
||||
walk; without one the reply goes back to the caller still holding
|
||||
placeholders even though the request was redacted correctly. Blocks come
|
||||
through as dicts or as objects depending on how far the reply has been
|
||||
deserialised, so both are handled.
|
||||
"""
|
||||
blocks: Final[list[object]] = [] # mutable-ok: accumulator, frozen on return.
|
||||
for item in getattr(response, "output", None) or ():
|
||||
for block in getattr(item, "content", None) or ():
|
||||
if isinstance(block, dict):
|
||||
if isinstance(block.get("text"), str) and block["text"]:
|
||||
blocks.append(block)
|
||||
elif isinstance(getattr(block, "text", None), str) and block.text:
|
||||
blocks.append(block)
|
||||
return tuple(blocks)
|
||||
|
||||
@staticmethod
|
||||
def _block_text(block: object) -> str:
|
||||
return block["text"] if isinstance(block, dict) else block.text
|
||||
|
||||
async def _restore_responses_api_response(
|
||||
self, response: Any, blocks: Sequence[object], data: MutableRequest
|
||||
) -> Any:
|
||||
"""Puts the original values back into a Responses API reply."""
|
||||
restored: Final = await self._rehydrate(
|
||||
tuple(self._block_text(block) for block in blocks), self._session_id(data)
|
||||
)
|
||||
for block, replacement in zip(blocks, restored):
|
||||
if isinstance(block, dict):
|
||||
block["text"] = replacement
|
||||
else:
|
||||
block.text = replacement
|
||||
return response
|
||||
|
||||
async def async_post_call_streaming_iterator_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
response: Any,
|
||||
request_data: dict,
|
||||
request_data: MutableRequest,
|
||||
) -> AsyncGenerator[Any, None]:
|
||||
"""Restores original values incrementally, without buffering the stream.
|
||||
|
||||
|
|
@ -295,9 +371,9 @@ class LLMShieldGuardrail(CustomGuardrail):
|
|||
yield chunk
|
||||
return
|
||||
|
||||
session_id = self._session_id(request_data)
|
||||
carry = ""
|
||||
last_chunk = None
|
||||
session_id: Final = self._session_id(request_data)
|
||||
carry = "" # rebind-ok: the sliding window advances with every delta.
|
||||
last_chunk = None # rebind-ok: tracks the most recent chunk for the final flush.
|
||||
|
||||
async for chunk in response:
|
||||
last_chunk = chunk
|
||||
|
|
@ -309,53 +385,54 @@ class LLMShieldGuardrail(CustomGuardrail):
|
|||
# Nothing to restore in this chunk, but a final chunk still has to
|
||||
# flush whatever the window is holding.
|
||||
if is_final and carry:
|
||||
body = await self._stream_step("", carry, True, session_id)
|
||||
carry = body["carry"]
|
||||
if body["text"] and delta is not None:
|
||||
delta.content = body["text"]
|
||||
emitted, carry = await self._stream_step("", carry, True, session_id)
|
||||
if emitted and delta is not None:
|
||||
delta.content = emitted
|
||||
yield chunk
|
||||
continue
|
||||
|
||||
body = await self._stream_step(text, carry, is_final, session_id)
|
||||
carry = body["carry"]
|
||||
delta.content = body["text"]
|
||||
emitted, carry = await self._stream_step(text, carry, is_final, session_id)
|
||||
delta.content = emitted
|
||||
yield chunk
|
||||
|
||||
# A stream that ended without a finish_reason can still leave text held back.
|
||||
if carry and last_chunk is not None:
|
||||
body = await self._stream_step("", carry, True, session_id)
|
||||
if body["text"]:
|
||||
trailing = last_chunk.model_copy(deep=True)
|
||||
trailing_delta = self._stream_delta(trailing)
|
||||
flushed: Final = await self._stream_step("", carry, True, session_id)
|
||||
trailing_text, carry = flushed # rebind-ok: window advances.
|
||||
if trailing_text:
|
||||
trailing: Final = last_chunk.model_copy(deep=True)
|
||||
trailing_delta: Final = self._stream_delta(trailing)
|
||||
if trailing_delta is not None:
|
||||
trailing_delta.content = body["text"]
|
||||
trailing_delta.content = trailing_text
|
||||
yield trailing
|
||||
|
||||
async def _stream_step(self, text: str, carry: str, final: bool, session_id: str) -> dict:
|
||||
body = await self._call_shield(
|
||||
async def _stream_step(self, text: str, carry: str, final: bool, session_id: str) -> tuple[str, str]:
|
||||
"""Returns ``(text safe to emit now, window still being held)``."""
|
||||
body: Final = await self._call_shield(
|
||||
_REHYDRATE_STREAM_PATH,
|
||||
session_id,
|
||||
{"text": text, "carry": carry, "final": final},
|
||||
# mutable-ok: JSON request body for httpx.
|
||||
{"text": text, "carry": carry, "final": final}, # mutable-ok: JSON request body for httpx.
|
||||
)
|
||||
emitted = body.get("text")
|
||||
remaining = body.get("carry")
|
||||
emitted: Final = body.get("text")
|
||||
remaining: Final = body.get("carry")
|
||||
if not isinstance(emitted, str) or not isinstance(remaining, str):
|
||||
raise GuardrailRaisedException(
|
||||
guardrail_name=self.guardrail_name,
|
||||
message="LLM Shield stream rehydration returned an unexpected payload.",
|
||||
)
|
||||
return {"text": emitted, "carry": remaining}
|
||||
return emitted, remaining
|
||||
|
||||
@staticmethod
|
||||
def _stream_delta(chunk: Any) -> Any:
|
||||
choices = getattr(chunk, "choices", None)
|
||||
def _stream_delta(chunk: object) -> Any:
|
||||
choices: Final = getattr(chunk, "choices", None)
|
||||
if not choices:
|
||||
return None
|
||||
return getattr(choices[0], "delta", None)
|
||||
|
||||
@staticmethod
|
||||
def _is_final_chunk(chunk: Any) -> bool:
|
||||
choices = getattr(chunk, "choices", None)
|
||||
def _is_final_chunk(chunk: object) -> bool:
|
||||
choices: Final = getattr(chunk, "choices", None)
|
||||
if not choices:
|
||||
return False
|
||||
return bool(getattr(choices[0], "finish_reason", None))
|
||||
|
|
@ -366,17 +443,21 @@ class LLMShieldGuardrail(CustomGuardrail):
|
|||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
request_data: MutableRequest,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional["LiteLLMLoggingObj"] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
texts = inputs.get("texts")
|
||||
texts: Final = inputs.get("texts")
|
||||
if not texts:
|
||||
return inputs
|
||||
|
||||
session_id = self._session_id(request_data)
|
||||
if input_type == "request":
|
||||
inputs["texts"] = await self._redact(list(texts), session_id)
|
||||
else:
|
||||
inputs["texts"] = await self._rehydrate(list(texts), session_id)
|
||||
return inputs
|
||||
session_id: Final = self._session_id(request_data)
|
||||
replaced: Final = (
|
||||
await self._redact(tuple(texts), session_id)
|
||||
if input_type == "request"
|
||||
else await self._rehydrate(tuple(texts), session_id)
|
||||
)
|
||||
# Return a new mapping rather than rewriting the caller's, so this stays a
|
||||
# pure transform of the inputs it was handed.
|
||||
merged: Final[JsonBody] = {**inputs, "texts": list(replaced)} # mutable-ok: TypedDict.
|
||||
return merged
|
||||
|
|
|
|||
|
|
@ -30,6 +30,10 @@ external = [
|
|||
# grows over time; typing it concretely (`object`) broke that forwarding call outright —
|
||||
# basedpyright turned every named param into a reportArgumentType error. Any is correct here.
|
||||
"litellm/proxy/guardrails/guardrail_hooks/alice/alice.py" = ["ANN401"]
|
||||
# Same reason: `**kwargs` forwards verbatim to CustomGuardrail.__init__, and the lifecycle
|
||||
# hook signatures inherit `Any` for `response` from CustomLogger, so narrowing them here
|
||||
# would break the override rather than describe it.
|
||||
"litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py" = ["ANN401"]
|
||||
|
||||
[lint.mccabe]
|
||||
max-complexity = 15
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
|
|
@ -133,10 +134,16 @@ class TestRedaction:
|
|||
assert data["messages"][0]["content"][1]["image_url"]["url"] == "http://x/y.png"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_request_without_messages_is_untouched(self):
|
||||
async def test_request_without_text_is_untouched(self):
|
||||
"""No text to redact means no call to LLM Shield.
|
||||
|
||||
This deliberately uses a request with no caller text at all. An earlier
|
||||
version used a Responses-API `input`, which asserted the very bypass that
|
||||
let `input` reach the provider unredacted.
|
||||
"""
|
||||
guardrail = _guardrail()
|
||||
mock = _mock_post(guardrail)
|
||||
data = {"input": "no messages here"}
|
||||
data = {"model": "gpt-4o", "temperature": 0.2}
|
||||
|
||||
await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion")
|
||||
|
||||
|
|
@ -156,6 +163,91 @@ class TestRedaction:
|
|||
assert len(sessions) == 1
|
||||
|
||||
|
||||
class TestRequestCoverage:
|
||||
"""Every request shape that carries caller text must be redacted.
|
||||
|
||||
A shape missed here is not a cosmetic gap: the guardrail reports as enabled
|
||||
while the raw value goes to the provider.
|
||||
"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_api_string_input_is_redacted(self):
|
||||
"""Measured against a live provider: `input` reached the model unredacted."""
|
||||
guardrail = _guardrail()
|
||||
mock = _mock_post(guardrail, {"texts": ["Email [EMAIL_1] the invoice"]})
|
||||
|
||||
data = {"input": "Email jane.doe@example.com the invoice"}
|
||||
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"] == ["Email jane.doe@example.com the invoice"]
|
||||
assert data["input"] == "Email [EMAIL_1] the invoice"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_api_list_input_is_redacted(self):
|
||||
guardrail = _guardrail()
|
||||
_mock_post(guardrail, {"texts": ["[EMAIL_1]", "[PHONE_1]"]})
|
||||
|
||||
data = {
|
||||
"input": [
|
||||
{"role": "user", "content": "jane.doe@example.com"},
|
||||
{"role": "user", "content": [{"type": "input_text", "text": "555-0100"}]},
|
||||
]
|
||||
}
|
||||
await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses")
|
||||
|
||||
assert data["input"][0]["content"] == "[EMAIL_1]"
|
||||
assert data["input"][1]["content"][0]["text"] == "[PHONE_1]"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_call_arguments_are_redacted(self):
|
||||
"""Tool arguments carry the values the user asked the model to act on."""
|
||||
guardrail = _guardrail()
|
||||
_mock_post(guardrail, {"texts": ['{"email": "[EMAIL_1]"}']})
|
||||
|
||||
data = {
|
||||
"messages": [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": {"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]["tool_calls"][0]["function"]["arguments"] == '{"email": "[EMAIL_1]"}'
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_every_shape_in_one_request_is_redacted(self):
|
||||
guardrail = _guardrail()
|
||||
mock = _mock_post(guardrail, {"texts": ["a", "b", "c", "d"]})
|
||||
|
||||
data = {
|
||||
"messages": [
|
||||
{"role": "user", "content": "one"},
|
||||
{"role": "user", "content": [{"type": "text", "text": "two"}]},
|
||||
{
|
||||
"role": "assistant",
|
||||
"tool_calls": [{"function": {"name": "f", "arguments": "three"}}],
|
||||
},
|
||||
],
|
||||
"input": "four",
|
||||
}
|
||||
await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion")
|
||||
|
||||
assert mock.call_args_list[0].kwargs["json"]["texts"] == ["one", "two", "three", "four"]
|
||||
assert data["messages"][0]["content"] == "a"
|
||||
assert data["messages"][1]["content"][0]["text"] == "b"
|
||||
assert data["messages"][2]["tool_calls"][0]["function"]["arguments"] == "c"
|
||||
assert data["input"] == "d"
|
||||
|
||||
|
||||
class TestRestoration:
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_shape_is_restored(self):
|
||||
|
|
@ -169,6 +261,36 @@ class TestRestoration:
|
|||
|
||||
assert result.choices[0].message.content == "a@b.com"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_api_shape_is_restored(self):
|
||||
"""The Responses API reply carries output items, not choices.
|
||||
|
||||
Measured against a live provider: once the request side was fixed the reply
|
||||
came back still holding the placeholder, because this shape has no choices
|
||||
to walk.
|
||||
"""
|
||||
guardrail = _guardrail(event_hook="post_call")
|
||||
_mock_post(guardrail, {"texts": ["a@b.com"]})
|
||||
|
||||
response = SimpleNamespace(output=[SimpleNamespace(content=[{"type": "output_text", "text": "[EMAIL_1]"}])])
|
||||
result = await guardrail.async_post_call_success_hook(
|
||||
data={"messages": []}, user_api_key_dict=None, response=response
|
||||
)
|
||||
|
||||
assert result.output[0].content[0]["text"] == "a@b.com"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_api_object_blocks_are_restored(self):
|
||||
"""Blocks arrive as objects too, depending on how far the reply is parsed."""
|
||||
guardrail = _guardrail(event_hook="post_call")
|
||||
_mock_post(guardrail, {"texts": ["a@b.com"]})
|
||||
|
||||
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)
|
||||
|
||||
assert block.text == "a@b.com"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_anthropic_message_shape_is_restored(self):
|
||||
"""The /v1/messages reply is a plain dict with no choices.
|
||||
|
|
|
|||
|
|
@ -73,7 +73,9 @@ interface GuardrailPreset {
|
|||
provider: string;
|
||||
categoryName?: string;
|
||||
guardrailNameSuggestion: string;
|
||||
mode: string;
|
||||
// A guardrail that both rewrites the request and repairs the response needs two
|
||||
// modes seeded, not one; the form already normalises either shape.
|
||||
mode: string | string[];
|
||||
defaultOn: boolean;
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -2,7 +2,9 @@ export interface GuardrailPreset {
|
|||
provider: string;
|
||||
categoryName?: string;
|
||||
guardrailNameSuggestion: string;
|
||||
mode: string;
|
||||
// A guardrail that both rewrites the request and repairs the response needs two
|
||||
// modes seeded, not one; the form already normalises either shape.
|
||||
mode: string | string[];
|
||||
defaultOn: boolean;
|
||||
}
|
||||
|
||||
|
|
@ -321,7 +323,9 @@ export const GUARDRAIL_PRESETS: Record<string, GuardrailPreset> = {
|
|||
llm_shield: {
|
||||
provider: "LLM Shield",
|
||||
guardrailNameSuggestion: "LLM Shield",
|
||||
mode: "pre_call",
|
||||
// Both halves are required. With only pre_call the request is redacted and the
|
||||
// placeholders are handed straight back to the caller.
|
||||
mode: ["pre_call", "post_call"],
|
||||
defaultOn: false,
|
||||
},
|
||||
};
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue