mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
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:
commit
e5cb8b7534
4 changed files with 511 additions and 108 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue