diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py index 9f3c85b555a..b222548f4ec 100644 --- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py @@ -234,19 +234,20 @@ def _write_back_message_text(message: _WritableMessage, target: MessageTextTarge _TOOL_USE_INPUT_ADAPTER: Final = TypeAdapter(dict[str, object]) -def _write_back_tool_use(message: _WritableMessage, target: ToolUseInputTarget, shape: _ToolCallShape) -> None: +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 - try: - rewritten_input: Final = _TOOL_USE_INPUT_ADAPTER.validate_json(shape.arguments) - except ValidationError: - verbose_proxy_logger.warning( - "Anthropic Messages: guardrail returned arguments that are not a JSON object for tool_use %s; keeping its input", - block.get("id"), - ) - 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 @@ -688,6 +689,7 @@ class AnthropicMessagesHandler(BaseTranslation): 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, ) verbose_proxy_logger.debug("Anthropic Messages: Processed input messages: %s", messages) @@ -1116,15 +1118,25 @@ class AnthropicMessagesHandler(BaseTranslation): 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) ) - for item, before, after in zip(scanned_tool_calls, pre_guardrail_tool_calls, post_guardrail_tool_calls): - if before != after: - _write_back_tool_use(messages[item.target.msg_idx], item.target, after) + 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, diff --git a/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py b/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py index 287674afaa7..b73ef6453fa 100644 --- a/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py +++ b/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py @@ -2328,16 +2328,20 @@ class TestAnthropicMessagesTopLevelSystemAndToolUseInputs: assert data["messages"][2]["content"][0]["tool_use_id"] == "toolu_01" @pytest.mark.asyncio - async def test_non_json_rewritten_arguments_keep_the_tool_use_input(self): + 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="You are a careful agent harness.") + original = json.loads(json.dumps(data)) - await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + with pytest.raises(UnappliableRequestRewrite) as excinfo: + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) - assert data["messages"][1]["content"][0]["input"] == { - "cmd": "AWS_ACCESS_KEY_ID=POISON aws sts get-caller-identity" - } + 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):