diff --git a/tests/test_litellm/llms/bedrock/passthrough/guardrail_translation/test_handler.py b/tests/test_litellm/llms/bedrock/passthrough/guardrail_translation/test_handler.py deleted file mode 100644 index 959c348d8b5..00000000000 --- a/tests/test_litellm/llms/bedrock/passthrough/guardrail_translation/test_handler.py +++ /dev/null @@ -1,47 +0,0 @@ -from copy import deepcopy -from typing import Final -from unittest.mock import AsyncMock, MagicMock - -import pytest - -from litellm.llms.bedrock.passthrough.guardrail_translation.handler import ( - BedrockPassthroughGuardrailHandler, -) - - -@pytest.mark.asyncio -async def test_skip_assistant_preserves_converse_history_and_masks_user() -> None: - assistant: Final = { - "role": "assistant", - "content": [ - {"text": "old reply"}, - { - "toolUse": { - "toolUseId": "t1", - "name": "search", - "input": {"q": "old"}, - } - }, - ], - } - data: Final = { - "endpoint": "model/test/converse", - "data": { - "messages": [ - assistant, - {"role": "user", "content": [{"text": "private"}]}, - ] - }, - } - original_assistant: Final = deepcopy(assistant) - guardrail: Final = MagicMock() - guardrail.apply_guardrail = AsyncMock(return_value={"texts": ["[MASKED]"]}) - guardrail.skip_system_message_in_guardrail = False - guardrail.skip_tool_message_in_guardrail = False - guardrail.skip_assistant_message_in_guardrail = True - - await BedrockPassthroughGuardrailHandler().process_input_messages(data, guardrail) - - assert guardrail.apply_guardrail.call_args.kwargs["inputs"]["texts"] == ["private"] - assert data["data"]["messages"][0] == original_assistant - assert data["data"]["messages"][1]["content"][0]["text"] == "[MASKED]" diff --git a/tests/unit/llms/bedrock/passthrough/guardrail_translation/test_handler.py b/tests/unit/llms/bedrock/passthrough/guardrail_translation/test_handler.py index c4ade50dc39..62c0e3a7402 100644 --- a/tests/unit/llms/bedrock/passthrough/guardrail_translation/test_handler.py +++ b/tests/unit/llms/bedrock/passthrough/guardrail_translation/test_handler.py @@ -10,6 +10,7 @@ Validates that: """ import copy +from typing import Final import pytest from unittest.mock import AsyncMock, MagicMock @@ -1213,3 +1214,41 @@ class TestDeAnonymizeConverseStream: hook_spy.assert_not_called() assert result is stream_bytes + + +@pytest.mark.asyncio +async def test_skip_assistant_preserves_converse_history_and_masks_user() -> None: + assistant: Final = { + "role": "assistant", + "content": [ + {"text": "old reply"}, + { + "toolUse": { + "toolUseId": "t1", + "name": "search", + "input": {"q": "old"}, + } + }, + ], + } + data: Final = { + "endpoint": "model/test/converse", + "data": { + "messages": [ + assistant, + {"role": "user", "content": [{"text": "private"}]}, + ] + }, + } + original_assistant: Final = copy.deepcopy(assistant) + guardrail: Final = MagicMock() + guardrail.apply_guardrail = AsyncMock(return_value={"texts": ["[MASKED]"]}) + guardrail.skip_system_message_in_guardrail = False + guardrail.skip_tool_message_in_guardrail = False + guardrail.skip_assistant_message_in_guardrail = True + + await BedrockPassthroughGuardrailHandler().process_input_messages(data, guardrail) + + assert guardrail.apply_guardrail.call_args.kwargs["inputs"]["texts"] == ["private"] + assert data["data"]["messages"][0] == original_assistant + assert data["data"]["messages"][1]["content"][0]["text"] == "[MASKED]"