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:
Ninad Phalak 2026-09-03 18:07:32 -05:00
parent 0de4b8b9a8
commit b6e3e6decd
No known key found for this signature in database
GPG key ID: 59119ED515433744
5 changed files with 323 additions and 110 deletions

View file

@ -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

View file

@ -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

View file

@ -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.

View file

@ -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;
}

View file

@ -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,
},
};