fix(anthropic): scan historical tool uses in pre-call guardrails

Co-authored-by: rpc-772 <rpc-772@users.noreply.github.com>

Co-authored-by: muyu-or <muyu-or@users.noreply.github.com>

Co-authored-by: guowei-su <guowei-su@users.noreply.github.com>
This commit is contained in:
Kangwenqiao 2026-09-09 16:10:14 +08:00
parent 47b15ffb67
commit 5cfb42da11
2 changed files with 205 additions and 1 deletions

View file

@ -101,6 +101,12 @@ class ToolResultBlockTextTarget:
block_idx: int
@dataclass(frozen=True, slots=True)
class ToolCallTarget:
msg_idx: int
content_idx: int
InputWriteBackTarget = (
MessageContentTarget | ContentBlockTextTarget | ToolResultStringTarget | ToolResultBlockTextTarget
)
@ -148,6 +154,8 @@ class ScannedText:
class ExtractedInput:
scanned: tuple[ScannedText, ...]
images: tuple[str, ...]
tool_calls: tuple[ChatCompletionToolCallChunk, ...] = ()
tool_call_targets: tuple[ToolCallTarget, ...] = ()
EMPTY_EXTRACTED_INPUT: Final = ExtractedInput(scanned=(), images=())
@ -454,16 +462,22 @@ 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)
tool_calls_to_check: Final = [ # mutable-ok: GenericGuardrailAPIInputs takes a list of tool calls
tool_call for one_message in extracted for tool_call in one_message.tool_calls
]
tool_call_targets: Final = [target for one_message in extracted for target in one_message.tool_call_targets]
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]
# Step 2: Apply guardrail to all texts in batch
if texts_to_check:
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
@ -481,6 +495,7 @@ class AnthropicMessagesHandler(BaseTranslation):
)
guardrailed_texts: Final = guardrailed_inputs.get("texts", [])
guardrailed_tool_calls: Final = guardrailed_inputs.get("tool_calls", [])
guardrailed_tools: Final = guardrailed_inputs.get("tools")
if guardrailed_tools is not None:
# Convert tools back from OpenAI format to Anthropic format
@ -528,11 +543,69 @@ class AnthropicMessagesHandler(BaseTranslation):
responses=guardrailed_texts,
scanned=scanned,
)
if guardrailed_tool_calls:
self._apply_guardrail_responses_to_input_tool_calls(
messages=messages,
tool_calls=guardrailed_tool_calls,
targets=tool_call_targets,
)
verbose_proxy_logger.debug("Anthropic Messages: Processed input messages: %s", messages)
return data
@staticmethod
def _tool_call_field(tool_call: object, field: str) -> object:
if isinstance(tool_call, Mapping):
return tool_call.get(field)
return getattr(tool_call, field, None)
@classmethod
def _tool_call_to_anthropic_block(cls, tool_call: object) -> dict[str, object] | None:
tool_id: Final = cls._tool_call_field(tool_call, "id")
function: Final = cls._tool_call_field(tool_call, "function")
name: Final = cls._tool_call_field(function, "name")
arguments: Final = cls._tool_call_field(function, "arguments")
if not isinstance(tool_id, str) or not isinstance(name, str):
return None
if isinstance(arguments, str):
try:
tool_input: Final = json.loads(arguments)
except json.JSONDecodeError:
return None
else:
tool_input = arguments
if not isinstance(tool_input, dict):
return None
block: Final[dict[str, object]] = {
"type": "tool_use",
"id": tool_id,
"name": name,
"input": tool_input,
}
caller: Final = cls._tool_call_field(tool_call, "caller")
if caller is not None:
block["caller"] = caller
return block
@classmethod
def _apply_guardrail_responses_to_input_tool_calls(
cls,
messages: Sequence[_WritableMessage],
tool_calls: Sequence[object],
targets: Sequence[ToolCallTarget],
) -> None:
"""Write guardrail-modified OpenAI tool calls back to Anthropic blocks."""
for tool_call, target in zip(tool_calls, targets):
anthropic_block = cls._tool_call_to_anthropic_block(tool_call)
if anthropic_block is None:
continue
content = messages[target.msg_idx].get("content")
if isinstance(content, list) and target.content_idx < len(content):
content[target.content_idx] = anthropic_block # mutable-ok: guardrail rewrite
def _hoisted_top_level_system_message(
self, data: dict
) -> AllMessageValues | None: # mutable-ok: API message payload
@ -807,6 +880,8 @@ class AnthropicMessagesHandler(BaseTranslation):
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(tool_call for block in blocks for tool_call in block.tool_calls),
tool_call_targets=tuple(target for block in blocks for target in block.tool_call_targets),
)
@classmethod
@ -826,6 +901,19 @@ class AnthropicMessagesHandler(BaseTranslation):
if scan_only_tool_results:
return EMPTY_EXTRACTED_INPUT
if content_item.get("type") == "tool_use":
return ExtractedInput(
scanned=(),
images=(),
tool_calls=(
AnthropicConfig.convert_tool_use_to_openai_format(
anthropic_tool_content=dict(content_item), # mutable-ok: conversion helper requires a dict
index=content_idx,
),
),
tool_call_targets=(ToolCallTarget(msg_idx=msg_idx, content_idx=content_idx),),
)
text_str: Final[str | None] = content_item.get("text")
return ExtractedInput(
scanned=(

View file

@ -1834,6 +1834,102 @@ class TestAnthropicMessagesToolResultScanning:
assert messages[0]["content"] == "keep me [BLOCKED]"
class TestAnthropicMessagesToolUseScanning:
def _data(self, messages):
return {"model": "claude-sonnet-4-5", "messages": messages}
@pytest.mark.asyncio
async def test_historical_tool_use_is_passed_to_pre_call_guardrail(self):
handler = AnthropicMessagesHandler()
guardrail = InputsRecordingGuardrail()
messages = [
{"role": "user", "content": "store my key"},
{
"role": "assistant",
"content": [
{
"type": "tool_use",
"id": "tu1",
"name": "store_credential",
"input": {"value": "secret-value"},
}
],
},
{"role": "user", "content": "what did you store?"},
]
await handler.process_input_messages(data=self._data(messages), 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
assert len(tool_calls) == 1
assert tool_calls[0]["function"]["name"] == "store_credential"
assert json.loads(tool_calls[0]["function"]["arguments"]) == {"value": "secret-value"}
@pytest.mark.asyncio
async def test_tool_use_only_request_still_invokes_pre_call_guardrail(self):
handler = AnthropicMessagesHandler()
guardrail = InputsRecordingGuardrail()
messages = [
{
"role": "assistant",
"content": [{"type": "tool_use", "id": "tu1", "name": "run", "input": {"command": "pwd"}}],
}
]
await handler.process_input_messages(data=self._data(messages), guardrail_to_apply=guardrail)
assert guardrail.captured_inputs is not None
assert guardrail.captured_inputs.get("texts") == []
assert guardrail.captured_inputs.get("tool_calls")
@pytest.mark.asyncio
async def test_scan_only_tool_results_excludes_tool_use(self):
handler = AnthropicMessagesHandler()
guardrail = InputsRecordingGuardrail()
guardrail.scan_only_tool_results = True
messages = [
{
"role": "assistant",
"content": [{"type": "tool_use", "id": "tu1", "name": "run", "input": {"command": "pwd"}}],
},
{
"role": "user",
"content": [{"type": "tool_result", "tool_use_id": "tu1", "content": "done"}],
},
]
await handler.process_input_messages(data=self._data(messages), guardrail_to_apply=guardrail)
assert guardrail.captured_inputs is not None
assert guardrail.captured_inputs.get("tool_calls") is None
@pytest.mark.asyncio
async def test_guardrail_tool_call_rewrite_is_written_back_to_tool_use(self):
handler = AnthropicMessagesHandler()
guardrail = ToolCallRewritingGuardrail()
data = self._data(
[
{
"role": "assistant",
"content": [
{
"type": "tool_use",
"id": "tu1",
"name": "store_credential",
"input": {"value": "secret-value"},
}
],
}
]
)
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
assert data["messages"][0]["content"][0]["input"] == {"value": "[MASKED]"}
class InputsRecordingGuardrail(MockCanaryMaskingGuardrail):
def __init__(self):
super().__init__(guardrail_name="scan-only-capture")
@ -1850,6 +1946,26 @@ class InputsRecordingGuardrail(MockCanaryMaskingGuardrail):
return await super().apply_guardrail(inputs, request_data, input_type, logging_obj)
class ToolCallRewritingGuardrail(CustomGuardrail):
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional[Any] = None,
) -> GenericGuardrailAPIInputs:
tool_calls = inputs.get("tool_calls")
if not tool_calls:
return inputs
rewritten = inputs.copy()
rewritten_call = dict(tool_calls[0])
rewritten_function = dict(rewritten_call["function"])
rewritten_function["arguments"] = json.dumps({"value": "[MASKED]"})
rewritten_call["function"] = rewritten_function
rewritten["tool_calls"] = [rewritten_call]
return rewritten
class StructuredMessagesRewritingGuardrail(CustomGuardrail):
"""Returns a new structured_messages list with a canary redacted, like redaction guardrails do."""