diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py index e486be12fe2..0f72729bc70 100644 --- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py @@ -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=( 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 bf40f781fa3..b4978b7a501 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 @@ -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."""