Merge pull request #40984 from BerriAI/litellm_anthropic_guardrail_system_and_tool_use

fix(guardrails): scan the Anthropic top-level system prompt and tool_use arguments
This commit is contained in:
Mateo Wang 2026-09-15 00:55:42 -07:00 committed by GitHub
commit e5cb8b7534
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 511 additions and 108 deletions

View file

@ -20,6 +20,7 @@ from itertools import chain, repeat
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Protocol, cast, overload, runtime_checkable
from pydantic import TypeAdapter, ValidationError
from typing_extensions import ReadOnly, TypedDict, assert_never
from litellm._logging import verbose_proxy_logger
@ -104,9 +105,24 @@ class ToolResultBlockTextTarget:
block_idx: int
InputWriteBackTarget = (
MessageContentTarget | ContentBlockTextTarget | ToolResultStringTarget | ToolResultBlockTextTarget
)
@dataclass(frozen=True, slots=True)
class SystemStringTarget:
pass
@dataclass(frozen=True, slots=True)
class SystemBlockTextTarget:
block_idx: int
@dataclass(frozen=True, slots=True)
class ToolUseInputTarget:
msg_idx: int
content_idx: int
MessageTextTarget = MessageContentTarget | ContentBlockTextTarget | ToolResultStringTarget | ToolResultBlockTextTarget
InputWriteBackTarget = SystemStringTarget | SystemBlockTextTarget | MessageTextTarget
def _as_str_mapping(value: Mapping[str, object]) -> Mapping[str, object]:
@ -147,10 +163,17 @@ class ScannedText:
target: InputWriteBackTarget
@dataclass(frozen=True, slots=True)
class ScannedToolCall:
tool_call: ChatCompletionToolCallChunk
target: ToolUseInputTarget
@dataclass(frozen=True, slots=True)
class ExtractedInput:
scanned: tuple[ScannedText, ...]
images: tuple[str, ...]
tool_calls: tuple[ScannedToolCall, ...] = ()
EMPTY_EXTRACTED_INPUT: Final = ExtractedInput(scanned=(), images=())
@ -162,6 +185,74 @@ class _ToolCallShape:
arguments: str
def _is_client_tool_use(block: Mapping[str, object]) -> bool:
return (
block.get("type") == "tool_use"
and isinstance(block.get("id"), str)
and isinstance(block.get("name"), str)
and isinstance(block.get("input"), dict)
)
def _write_back_system_block(system: object, block_idx: int, response: str) -> None:
if not isinstance(system, list):
return
text_blocks: Final = tuple(block for block in system if isinstance(block, dict) and block.get("type") == "text")
if block_idx < len(text_blocks):
text_blocks[block_idx]["text"] = (
response # mutable-ok: guardrails rewrite the caller's request payload in place
)
def _write_back_message_text(message: _WritableMessage, target: MessageTextTarget, response: str) -> None:
content: Final = message.get("content", None)
if content is None:
return
match target:
case MessageContentTarget():
if isinstance(content, str):
message["content"] = response # mutable-ok: guardrails rewrite the caller's request payload in place
case ContentBlockTextTarget(content_idx=content_idx):
if isinstance(content, list):
content[content_idx]["text"] = (
response # mutable-ok: guardrails rewrite the caller's request payload in place
)
case ToolResultStringTarget(content_idx=content_idx):
if isinstance(content, list):
content[content_idx]["content"] = (
response # mutable-ok: guardrails rewrite the caller's request payload in place
)
case ToolResultBlockTextTarget(content_idx=content_idx, block_idx=block_idx):
if isinstance(content, list):
content[content_idx]["content"][block_idx]["text"] = (
response # mutable-ok: guardrails rewrite the caller's request payload in place
)
case _:
assert_never(target)
_TOOL_USE_INPUT_ADAPTER: Final = TypeAdapter(dict[str, object])
def _rewritten_tool_use_input(arguments: str) -> Mapping[str, object] | None:
try:
return _TOOL_USE_INPUT_ADAPTER.validate_json(arguments)
except ValidationError:
return None
def _write_back_tool_use(
message: _WritableMessage, target: ToolUseInputTarget, shape: _ToolCallShape, rewritten_input: Mapping[str, object]
) -> None:
content: Final = message.get("content", None)
block: Final = content[target.content_idx] if isinstance(content, list) else None
if not isinstance(block, dict):
return
block["input"] = rewritten_input # mutable-ok: guardrails rewrite the caller's request payload in place
if shape.name is not None and shape.name != block.get("name"):
block["name"] = shape.name # mutable-ok: guardrails rewrite the caller's request payload in place
@dataclass(frozen=True, slots=True)
class _SSEFieldRewrite:
"""One field of one nested section of a buffered SSE event, rewritten."""
@ -453,9 +544,8 @@ class AnthropicMessagesHandler(BaseTranslation):
skip_tool: Final = effective_skip_tool_message_for_guardrail(guardrail_to_apply)
scan_only_tool_results: Final = effective_scan_only_tool_results_for_guardrail(guardrail_to_apply)
# Exclude only the trusted top-level prompt. In-sequence system entries are untrusted
# and must stay aligned with texts_to_check for positional masking. When the top-level
# prompt is included, the pre-existing count mismatch disables positional masking.
# The top-level prompt is translated on its own below so it can be hoisted in front of
# any mid-turn system entries and scanned first, aligned with that structured position.
translation_source: Final = { # mutable-ok: API message payload
key: value for key, value in data.items() if key != "system"
}
@ -491,7 +581,12 @@ class AnthropicMessagesHandler(BaseTranslation):
]
)
# Step 1: Extract all text content and images
# Step 1: Extract all text content, images, and tool calls
top_level_system_scanned: Final = (
()
if hoisted_system_message is None or scan_only_tool_results
else self._extract_top_level_system_text(hoisted_system_message)
)
extracted: Final = tuple(
self._extract_input_text_and_images(
message=message,
@ -502,17 +597,27 @@ class AnthropicMessagesHandler(BaseTranslation):
)
for msg_idx, message in enumerate(messages)
)
scanned: Final = tuple(item for one_message in extracted for item in one_message.scanned)
scanned: Final = (
*top_level_system_scanned,
*(item for one_message in extracted for item in one_message.scanned),
)
texts_to_check: Final = [item.text for item in scanned] # mutable-ok: GenericGuardrailAPIInputs takes list[str]
images_to_check: Final = [
image for one_message in extracted for image in one_message.images
] # mutable-ok: GenericGuardrailAPIInputs takes list[str]
scanned_tool_calls: Final = tuple(item for one_message in extracted for item in one_message.tool_calls)
tool_calls_to_check: Final = [
item.tool_call for item in scanned_tool_calls
] # mutable-ok: GenericGuardrailAPIInputs takes list[ChatCompletionToolCallChunk]
pre_guardrail_tool_calls: Final = _tool_call_shapes(tool_calls_to_check)
# Step 2: Apply guardrail to all texts in batch
if texts_to_check:
# Step 2: Apply guardrail to all texts and tool calls in batch
if texts_to_check or tool_calls_to_check:
inputs: Final = GenericGuardrailAPIInputs(texts=texts_to_check)
if images_to_check:
inputs["images"] = images_to_check
if tool_calls_to_check:
inputs["tool_calls"] = tool_calls_to_check
if tools_to_check:
inputs["tools"] = tools_to_check
original_structured_messages: Final = structured_messages
@ -573,9 +678,16 @@ class AnthropicMessagesHandler(BaseTranslation):
else:
if guardrailed_texts and len(guardrailed_texts) != len(scanned):
raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name)
self._apply_guardrail_tool_calls_to_input(
messages=messages,
scanned_tool_calls=scanned_tool_calls,
pre_guardrail_tool_calls=pre_guardrail_tool_calls,
returned_tool_calls=guardrailed_inputs.get("tool_calls"),
guardrail_name=guardrail_to_apply.guardrail_name,
)
# Step 3: Map guardrail responses back to original message structure
await self._apply_guardrail_responses_to_input(
messages=messages,
data=data,
responses=guardrailed_texts,
scanned=scanned,
)
@ -601,6 +713,19 @@ class AnthropicMessagesHandler(BaseTranslation):
hoisted: Final = probe.get("messages") or [] # mutable-ok: API message payload
return hoisted[0] if hoisted else None
@staticmethod
def _extract_top_level_system_text(hoisted_system_message: AllMessageValues) -> tuple[ScannedText, ...]:
content: Final = hoisted_system_message.get("content")
if isinstance(content, str):
return (ScannedText(content, SystemStringTarget()),)
if not isinstance(content, list):
return ()
return tuple(
ScannedText(text_str, SystemBlockTextTarget(block_idx))
for block_idx, block in enumerate(content)
if isinstance(block, dict) and isinstance(text_str := block.get("text"), str)
)
@staticmethod
def _openai_system_message_to_anthropic(
message: Mapping[str, object],
@ -855,9 +980,25 @@ class AnthropicMessagesHandler(BaseTranslation):
for content_idx, content_item in enumerate(content)
if isinstance(content_item, dict)
)
tool_use_blocks: Final = (
()
if scan_only_tool_results
else tuple(
(content_idx, content_item)
for content_idx, content_item in enumerate(content)
if isinstance(content_item, dict) and _is_client_tool_use(content_item)
)
)
return ExtractedInput(
scanned=tuple(item for block in blocks for item in block.scanned),
images=tuple(image for block in blocks for image in block.images),
tool_calls=tuple(
ScannedToolCall(
tool_call=AnthropicConfig.convert_tool_use_to_openai_format(content_item, tool_call_idx),
target=ToolUseInputTarget(msg_idx, content_idx),
)
for tool_call_idx, (content_idx, content_item) in enumerate(tool_use_blocks)
),
)
@classmethod
@ -943,43 +1084,59 @@ class AnthropicMessagesHandler(BaseTranslation):
async def _apply_guardrail_responses_to_input(
self,
messages: Sequence[_WritableMessage],
responses: list[str],
data: dict[str, object], # mutable-ok: API message payload
responses: Sequence[str],
scanned: tuple[ScannedText, ...],
) -> None:
"""
Apply guardrail responses back to input messages.
Apply guardrail responses back to the top-level system prompt and the input messages.
"""
raw_messages: Final = data.get("messages")
messages: Final[Sequence[_WritableMessage]] = raw_messages if isinstance(raw_messages, list) else ()
for item, guardrail_response in zip(scanned, responses):
target = item.target
message = messages[target.msg_idx]
content = message.get("content", None)
if content is None:
continue
match target:
case MessageContentTarget():
if isinstance(content, str):
message["content"] = (
guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place
)
case ContentBlockTextTarget(content_idx=content_idx):
if isinstance(content, list):
content[content_idx]["text"] = (
guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place
)
case ToolResultStringTarget(content_idx=content_idx):
if isinstance(content, list):
content[content_idx]["content"] = (
guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place
)
case ToolResultBlockTextTarget(content_idx=content_idx, block_idx=block_idx):
if isinstance(content, list):
content[content_idx]["content"][block_idx]["text"] = (
match item.target:
case SystemStringTarget():
if isinstance(data.get("system"), str):
data["system"] = (
guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place
)
case SystemBlockTextTarget(block_idx=block_idx):
_write_back_system_block(data.get("system"), block_idx, guardrail_response)
case (
MessageContentTarget()
| ContentBlockTextTarget()
| ToolResultStringTarget()
| ToolResultBlockTextTarget() as message_target
):
_write_back_message_text(messages[message_target.msg_idx], message_target, guardrail_response)
case _:
assert_never(target)
assert_never(item.target)
@staticmethod
def _apply_guardrail_tool_calls_to_input(
messages: Sequence[_WritableMessage],
scanned_tool_calls: tuple[ScannedToolCall, ...],
pre_guardrail_tool_calls: tuple[_ToolCallShape, ...],
returned_tool_calls: Sequence[object] | None,
guardrail_name: str | None,
) -> None:
post_guardrail_tool_calls: Final = _tool_call_shapes(
returned_tool_calls
if returned_tool_calls is not None and len(returned_tool_calls) == len(pre_guardrail_tool_calls)
else tuple(item.tool_call for item in scanned_tool_calls)
)
rewritten: Final = tuple(
(item, after, _rewritten_tool_use_input(after.arguments))
for item, before, after in zip(scanned_tool_calls, pre_guardrail_tool_calls, post_guardrail_tool_calls)
if before != after
)
applicable: Final = tuple(
(item, after, rewritten_input) for item, after, rewritten_input in rewritten if rewritten_input is not None
)
if len(applicable) != len(rewritten):
raise unappliable_request_rewrite(guardrail_name)
for item, after, rewritten_input in applicable:
_write_back_tool_use(messages[item.target.msg_idx], item.target, after, rewritten_input)
async def process_output_response(
self,

View file

@ -1600,8 +1600,8 @@ class PanwPrismaAirsHandler(CustomGuardrail):
Args:
texts: Flattened text entries from the framework.
messages: Original request messages (request_data["messages"]),
NOT structured_messages (which may have injected system content).
messages: The structured messages the framework flattened into ``texts``,
hoisted top-level system prompt included, so positions line up.
Returns a set of scannable indices, or None on count mismatch or no user/developer
message (safety fallback to existing role-filter behavior).
@ -1788,15 +1788,10 @@ class PanwPrismaAirsHandler(CustomGuardrail):
structured_messages: Final = inputs.get("structured_messages")
if structured_messages:
# For Anthropic /v1/messages: default to latest-user-only scanning.
# Uses request_data["messages"] (original format), NOT structured_messages
# (which has injected system content from adapter translation).
if self._use_latest_user_only(request_data, logging_obj):
original_messages: Final = request_data.get("messages")
if original_messages:
scannable_indices = self._get_latest_user_text_indices(texts, original_messages)
scannable_indices = self._get_latest_user_text_indices(texts, structured_messages)
# Fall through to existing role filtering if:
# - not Anthropic, OR flag explicitly False, OR
# - no original messages, OR
# - latest-user extraction returned None (no user / count mismatch)
if scannable_indices is None:
scannable_indices = self._get_scannable_text_indices(texts, structured_messages)

View file

@ -13,6 +13,7 @@ import pytest
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.guardrail_translation.base_translation import StreamingScanKey
from litellm.llms.anthropic.chat.guardrail_translation.handler import (
AnthropicMessagesHandler,
@ -635,14 +636,19 @@ class TestAnthropicMessagesHandlerInputProcessing:
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
assert guardrail.inputs is not None
assert guardrail.inputs["texts"] == ["safe text", "prohibited correction"]
assert guardrail.inputs["texts"] == [
"trusted top-level system prompt",
"safe text",
"prohibited correction",
]
structured = guardrail.inputs["structured_messages"]
assert [m["role"] for m in structured] == ["system", "user", "system"]
assert structured[0]["content"] == "trusted top-level system prompt"
assert data["system"] == "trusted top-level system prompt"
assert data["messages"][1]["content"] == "[MASKED]"
@pytest.mark.asyncio
async def test_bedrock_masking_slice_is_unavailable_when_top_level_system_is_included(
async def test_bedrock_masking_slice_lines_up_when_top_level_system_is_included(
self,
):
from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
@ -668,25 +674,25 @@ class TestAnthropicMessagesHandlerInputProcessing:
structured = guardrail.inputs["structured_messages"]
bedrock = BedrockGuardrail(guardrailIdentifier="gi", guardrailVersion="1")
assert sum(bedrock._count_message_texts(m) for m in structured) == len(texts) + 1
assert sum(bedrock._count_message_texts(m) for m in structured) == len(texts)
latest_user_index = bedrock._find_latest_message_index(structured, target_role="user")
assert (
bedrock._locate_message_texts_slice(
structured_messages=structured,
target_index=latest_user_index,
texts=texts,
)
is None
)
assert (
bedrock._merge_masked_texts(
masked_texts=["{MASKED}"],
texts=texts,
scanned_slice=None,
scanned_role_subset=True,
)
== texts
scanned_slice = bedrock._locate_message_texts_slice(
structured_messages=structured,
target_index=latest_user_index,
texts=texts,
)
assert scanned_slice == (3, 1)
assert bedrock._merge_masked_texts(
masked_texts=["{MASKED}"],
texts=texts,
scanned_slice=scanned_slice,
scanned_role_subset=True,
) == [
"trusted top-level system prompt",
"safe text",
"prohibited correction",
"{MASKED}",
]
@pytest.mark.asyncio
@pytest.mark.parametrize("skip_system_message_in_guardrail", [True, None])
@ -1611,7 +1617,8 @@ class TestAnthropicMessagesIncrementalScan:
)
assert mock_api.call_count == 1
assert [m["content"] for m in mock_api.call_args.kwargs["messages"]] == [
"What is the capital of France?"
"You are a helpful geography assistant.",
"What is the capital of France?",
]
mock_api.reset_mock()
await handler.process_input_messages(
@ -2150,6 +2157,213 @@ class TestAnthropicMessagesScanOnlyToolResults:
assert guardrail.captured_inputs.get("images") == ["TOOL_IMG"]
class ToolCallArgumentsMaskingGuardrail(InputsRecordingGuardrail):
"""Masks the canary inside tool-call arguments, in place or through a fresh list of plain dicts."""
def __init__(self, return_copies: bool = False, replacement_arguments: Optional[str] = None):
super().__init__()
self.return_copies = return_copies
self.replacement_arguments = replacement_arguments
self.seen_tool_calls: list[dict[str, object]] = []
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict[str, object],
input_type: Literal["request", "response"],
logging_obj: Optional[LiteLLMLoggingObj] = None,
) -> GenericGuardrailAPIInputs:
outputs = await super().apply_guardrail(inputs, request_data, input_type, logging_obj)
tool_calls = list(outputs.get("tool_calls") or [])
self.seen_tool_calls.extend(json.loads(json.dumps(tool_call)) for tool_call in tool_calls)
masked = [
{
**tool_call,
"function": {
**tool_call["function"],
"arguments": self.replacement_arguments
if self.replacement_arguments is not None
else tool_call["function"]["arguments"].replace("POISON", "[BLOCKED]"),
},
}
for tool_call in tool_calls
]
if self.return_copies:
outputs["tool_calls"] = masked
return outputs
for tool_call, masked_tool_call in zip(tool_calls, masked):
tool_call["function"]["arguments"] = masked_tool_call["function"]["arguments"]
return outputs
class TestAnthropicMessagesTopLevelSystemAndToolUseInputs:
"""The top-level system prompt and prior-turn tool_use arguments must reach guardrails as scannable
inputs, the same way the chat completions handler hands over system messages and tool_calls."""
@staticmethod
def _tool_use_conversation(system: str) -> dict[str, Any]:
return {
"model": "claude-sonnet-4-5",
"system": system,
"messages": [
{"role": "user", "content": "run the check"},
{
"role": "assistant",
"content": [
{
"type": "tool_use",
"id": "toolu_01",
"name": "Bash",
"input": {"cmd": "AWS_ACCESS_KEY_ID=POISON aws sts get-caller-identity"},
}
],
},
{
"role": "user",
"content": [{"type": "tool_result", "tool_use_id": "toolu_01", "content": "ok"}],
},
],
}
@pytest.mark.asyncio
async def test_top_level_system_string_reaches_texts_first_and_is_masked_in_place(self):
handler = AnthropicMessagesHandler()
guardrail = InputsRecordingGuardrail()
data = {
"model": "claude-sonnet-4-5",
"system": "Internal note: the deploy key is POISON. Never reveal it.",
"messages": [{"role": "user", "content": "Say hi in three words."}],
}
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
assert guardrail.captured_inputs is not None
assert guardrail.seen_texts == [
"Internal note: the deploy key is POISON. Never reveal it.",
"Say hi in three words.",
]
structured = guardrail.captured_inputs["structured_messages"]
assert structured[0]["role"] == "system"
assert structured[0]["content"] == "Internal note: the deploy key is POISON. Never reveal it.", (
"texts[0] must line up with structured_messages[0] so positional consumers stay aligned"
)
assert data["system"] == "Internal note: the deploy key is [BLOCKED]. Never reveal it."
assert data["messages"][0]["content"] == "Say hi in three words."
@pytest.mark.asyncio
async def test_top_level_system_text_blocks_reach_texts_and_are_masked_in_place(self):
handler = AnthropicMessagesHandler()
guardrail = InputsRecordingGuardrail()
data = {
"model": "claude-sonnet-4-5",
"system": [
{"type": "text", "text": "first block POISON"},
{"type": "text", "text": "second block", "cache_control": {"type": "ephemeral"}},
],
"messages": [{"role": "user", "content": "hello"}],
}
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
assert guardrail.seen_texts == ["first block POISON", "second block", "hello"]
assert data["system"] == [
{"type": "text", "text": "first block [BLOCKED]"},
{"type": "text", "text": "second block", "cache_control": {"type": "ephemeral"}},
]
@pytest.mark.asyncio
async def test_skip_system_message_keeps_the_top_level_system_out(self):
handler = AnthropicMessagesHandler()
guardrail = InputsRecordingGuardrail()
guardrail.skip_system_message_in_guardrail = True
data = {
"model": "claude-sonnet-4-5",
"system": "trusted POISON prompt",
"messages": [{"role": "user", "content": "hello"}],
}
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
assert guardrail.seen_texts == ["hello"]
assert data["system"] == "trusted POISON prompt"
@pytest.mark.asyncio
async def test_prior_turn_tool_use_input_reaches_tool_calls_in_openai_shape(self):
handler = AnthropicMessagesHandler()
guardrail = InputsRecordingGuardrail()
data = self._tool_use_conversation(system="You are a careful agent harness.")
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
assert guardrail.captured_inputs is not None
tool_calls = guardrail.captured_inputs.get("tool_calls")
assert tool_calls is not None and len(tool_calls) == 1
assert tool_calls[0]["id"] == "toolu_01"
assert tool_calls[0]["type"] == "function"
assert tool_calls[0]["function"]["name"] == "Bash"
assert json.loads(tool_calls[0]["function"]["arguments"]) == {
"cmd": "AWS_ACCESS_KEY_ID=POISON aws sts get-caller-identity"
}
assert data["messages"][1]["content"][0]["input"] == {
"cmd": "AWS_ACCESS_KEY_ID=POISON aws sts get-caller-identity"
}, "a guardrail that leaves tool_calls alone must leave the tool_use input alone"
@pytest.mark.asyncio
@pytest.mark.parametrize("return_copies", [False, True])
async def test_masked_tool_call_arguments_write_back_into_the_tool_use_input(self, return_copies: bool):
handler = AnthropicMessagesHandler()
guardrail = ToolCallArgumentsMaskingGuardrail(return_copies=return_copies)
data = self._tool_use_conversation(system="You are a careful agent harness.")
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
assert [tool_call["function"]["name"] for tool_call in guardrail.seen_tool_calls] == ["Bash"]
tool_use = data["messages"][1]["content"][0]
assert tool_use == {
"type": "tool_use",
"id": "toolu_01",
"name": "Bash",
"input": {"cmd": "AWS_ACCESS_KEY_ID=[BLOCKED] aws sts get-caller-identity"},
}
assert data["messages"][2]["content"][0]["tool_use_id"] == "toolu_01"
@pytest.mark.asyncio
async def test_non_json_rewritten_arguments_are_rejected_by_name(self):
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
handler = AnthropicMessagesHandler()
guardrail = ToolCallArgumentsMaskingGuardrail(replacement_arguments="[REDACTED]")
data = self._tool_use_conversation(system="Internal note: the deploy key is POISON. Never reveal it.")
data["messages"][2]["content"][0]["content"] = "fetched POISON page"
original = json.loads(json.dumps(data))
with pytest.raises(UnappliableRequestRewrite) as excinfo:
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
assert excinfo.value.guardrail_name == "scan-only-capture"
assert data["system"] == original["system"], "a rejected rewrite must leave the request untouched"
assert data["messages"] == original["messages"], "a rejected rewrite must leave the request untouched"
@pytest.mark.asyncio
async def test_scan_only_tool_results_keeps_system_and_tool_use_out(self):
handler = AnthropicMessagesHandler()
guardrail = InputsRecordingGuardrail()
guardrail.scan_only_tool_results = True
data = self._tool_use_conversation(system="trusted POISON prompt")
data["messages"][2]["content"][0]["content"] = "fetched POISON page"
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
assert guardrail.seen_texts == ["fetched POISON page"]
assert guardrail.captured_inputs is not None
assert guardrail.captured_inputs.get("tool_calls") is None
assert data["system"] == "trusted POISON prompt"
assert data["messages"][1]["content"][0]["input"] == {
"cmd": "AWS_ACCESS_KEY_ID=POISON aws sts get-caller-identity"
}
assert data["messages"][2]["content"][0]["content"] == "fetched [BLOCKED] page"
class TestStructuredWriteBackKeepsToolResults:
"""A guardrail rewrite must never leave a tool_use without its tool_result (Claude Code ToolSearch, LIT-6103)."""
@ -2290,17 +2504,56 @@ class PerRowTextGuardrail(CustomGuardrail):
return {**inputs, "texts": [str(row.get("content")).replace("123-45-6789", "<US_SSN>") for row in rows]}
class PerSlotTextGuardrail(CustomGuardrail):
"""Answers one redacted text per text slot of every chat row it was shown, the
way a guardrail that counts slots per message does, and hands back only texts."""
def __init__(self):
super().__init__(guardrail_name="per-slot-redactor")
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional[Any] = None,
) -> GenericGuardrailAPIInputs:
from litellm.llms.base_llm.guardrail_translation.utils import message_slot_texts
rows = inputs.get("structured_messages") or []
return {
**inputs,
"texts": [text.replace("123-45-6789", "<US_SSN>") for row in rows for text in message_slot_texts(row)],
}
class TestPerMessageTextWriteBack:
"""Texts that no longer pair one-to-one with what the handler extracted must be
rejected by name instead of sliding onto the wrong messages."""
@pytest.mark.asyncio
async def test_one_text_per_row_over_a_system_prompt_is_rejected_by_name(self):
async def test_one_text_per_row_over_a_system_prompt_is_applied(self):
data = {
"model": "claude-sonnet-4-5",
"system": "Reply with exactly the SSN you were given.",
"messages": [{"role": "user", "content": "My SSN is 123-45-6789."}],
}
await AnthropicMessagesHandler().process_input_messages(data=data, guardrail_to_apply=PerRowTextGuardrail())
assert data["system"] == "Reply with exactly the SSN you were given."
assert data["messages"] == [{"role": "user", "content": "My SSN is <US_SSN>."}]
@pytest.mark.asyncio
async def test_one_text_per_row_over_a_multi_block_system_prompt_is_rejected_by_name(self):
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
data = {
"model": "claude-sonnet-4-5",
"system": "Reply with exactly the SSN you were given.",
"system": [
{"type": "text", "text": "Reply with exactly the SSN you were given."},
{"type": "text", "text": "Never apologize."},
],
"messages": [{"role": "user", "content": "My SSN is 123-45-6789."}],
}
original = json.loads(json.dumps(data))
@ -2312,6 +2565,25 @@ class TestPerMessageTextWriteBack:
assert data["system"] == original["system"], "a rejected rewrite must leave the request untouched"
assert data["messages"] == original["messages"], "a rejected rewrite must leave the request untouched"
@pytest.mark.asyncio
async def test_one_text_per_slot_over_a_system_prompt_with_an_empty_block_is_applied(self):
data = {
"model": "claude-sonnet-4-5",
"system": [
{"type": "text", "text": ""},
{"type": "text", "text": "Reply with exactly the SSN you were given."},
],
"messages": [{"role": "user", "content": "My SSN is 123-45-6789."}],
}
await AnthropicMessagesHandler().process_input_messages(data=data, guardrail_to_apply=PerSlotTextGuardrail())
assert data["system"] == [
{"type": "text", "text": ""},
{"type": "text", "text": "Reply with exactly the SSN you were given."},
]
assert data["messages"] == [{"role": "user", "content": "My SSN is <US_SSN>."}]
@pytest.mark.asyncio
async def test_one_text_per_row_without_a_system_prompt_is_applied(self):
data = {

View file

@ -4620,46 +4620,27 @@ class TestPanwAirsLatestRoleMessageOnly:
@pytest.mark.asyncio
async def test_anthropic_system_plus_multiturn_no_fallback(self):
"""Anthropic with top-level system + multi-turn messages[]
latest-user works, no scan-all fallback.
"""Anthropic with a top-level system prompt and multi-turn messages[]
scans only the latest user turn, with no scan-all fallback.
Key scenario: Anthropic top-level `system` field causes
structured_messages to have an injected system entry, but
request_data["messages"] does NOT include it.
The Anthropic handler hoists the top-level `system` field into both
`texts` and `structured_messages`, so the latest-user walk has to
count the same entries the framework flattened.
"""
handler = PanwPrismaAirsHandler(
guardrail_name="test_panw_airs",
api_key="test_api_key",
profile_name="test_profile",
default_on=True,
from litellm.llms.anthropic.chat.guardrail_translation.handler import (
AnthropicMessagesHandler,
)
# Original Anthropic messages (no system in messages array)
original_messages = [
{"role": "user", "content": "First user turn"},
{"role": "assistant", "content": "First assistant turn"},
{"role": "user", "content": "Latest user turn"},
]
# texts extracted from original_messages (3 text entries)
texts = ["First user turn", "First assistant turn", "Latest user turn"]
# structured_messages has an INJECTED system message from translation
structured_messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "First user turn"},
{"role": "assistant", "content": "First assistant turn"},
{"role": "user", "content": "Latest user turn"},
]
inputs: GenericGuardrailAPIInputs = {
"texts": texts,
"structured_messages": structured_messages,
}
handler = make_handler()
request_data = {
"litellm_call_id": "test-call-id",
"model": "anthropic/claude-sonnet-4-20250514",
"messages": original_messages,
"system": "You are a helpful assistant.",
"messages": [
{"role": "user", "content": "First user turn"},
{"role": "assistant", "content": "First assistant turn"},
{"role": "user", "content": "Latest user turn"},
],
"proxy_server_request": {
"url": "http://localhost:4000/v1/messages",
},
@ -4670,13 +4651,11 @@ class TestPanwAirsLatestRoleMessageOnly:
) as mock_api:
mock_api.return_value = {"action": "allow", "category": "benign"}
await handler.apply_guardrail(
inputs=inputs,
request_data=request_data,
input_type="request",
await AnthropicMessagesHandler().process_input_messages(
data=request_data,
guardrail_to_apply=handler,
)
# Should scan ONLY the latest user message, not fall back to scan-all
assert mock_api.call_count == 1
assert mock_api.call_args.kwargs["content"] == "Latest user turn"