diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 8ef188bb23c..0395a2608fe 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -177,6 +177,11 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): if messages is None: return bedrock_request for message in messages: + if ( + self.experimental_use_latest_role_message_only + and message.get("role") != "user" + ): + continue message_text_content: Optional[List[str]] = self.get_content_for_message( message=message ) @@ -970,25 +975,39 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): input_filter = self._prepare_guardrail_messages_for_role( messages=new_messages ) - input_messages = input_filter.payload_messages or new_messages + input_messages = input_filter.payload_messages + if input_messages is None: + if self.experimental_use_latest_role_message_only: + input_messages = None # no user messages → skip INPUT validation + else: + input_messages = new_messages - # Create tasks for parallel execution of both INPUT and OUTPUT validation - input_task = self.make_bedrock_api_request( - source="INPUT", - messages=input_messages, - request_data=data, - ) - output_task = self.make_bedrock_api_request( - source="OUTPUT", response=response, request_data=data - ) - - # Execute both requests in parallel - try: - _, output_content_bedrock = await asyncio.gather( - input_task, output_task + if input_messages is not None: + # Create tasks for parallel execution of both INPUT and OUTPUT validation + input_task = self.make_bedrock_api_request( + source="INPUT", + messages=input_messages, + request_data=data, ) - except GuardrailInterventionNormalStringError as e: - output_content_bedrock = e.message + output_task = self.make_bedrock_api_request( + source="OUTPUT", response=response, request_data=data + ) + + # Execute both requests in parallel + try: + _, output_content_bedrock = await asyncio.gather( + input_task, output_task + ) + except GuardrailInterventionNormalStringError as e: + output_content_bedrock = e.message + else: + # No user messages to validate INPUT — only run OUTPUT validation + try: + output_content_bedrock = await self.make_bedrock_api_request( + source="OUTPUT", response=response, request_data=data + ) + except GuardrailInterventionNormalStringError as e: + output_content_bedrock = e.message else: # Only run OUTPUT validation (INPUT was already validated in pre_call or during_call) try: @@ -1113,25 +1132,41 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): input_filter = self._prepare_guardrail_messages_for_role( messages=request_data.get("messages") ) - input_messages = input_filter.payload_messages or request_data.get( - "messages" - ) - input_task = self.make_bedrock_api_request( - source="INPUT", - messages=input_messages, - request_data=request_data, - ) # Only input messages - output_task = self.make_bedrock_api_request( - source="OUTPUT", response=assembled_model_response - ) # Only response + input_messages = input_filter.payload_messages + if input_messages is None: + if self.experimental_use_latest_role_message_only: + input_messages = None # no user messages → skip INPUT validation + else: + input_messages = request_data.get("messages") - # Execute both requests in parallel - try: - _, output_guardrail_response = await asyncio.gather( - input_task, output_task - ) - except GuardrailInterventionNormalStringError as e: - output_guardrail_response = e.message + if input_messages is not None: + input_task = self.make_bedrock_api_request( + source="INPUT", + messages=input_messages, + request_data=request_data, + ) # Only input messages + output_task = self.make_bedrock_api_request( + source="OUTPUT", response=assembled_model_response + ) # Only response + + # Execute both requests in parallel + try: + _, output_guardrail_response = await asyncio.gather( + input_task, output_task + ) + except GuardrailInterventionNormalStringError as e: + output_guardrail_response = e.message + else: + # No user messages to validate INPUT — only run OUTPUT validation + try: + output_guardrail_response = ( + await self.make_bedrock_api_request( + source="OUTPUT", + response=assembled_model_response, + ) + ) + except GuardrailInterventionNormalStringError as e: + output_guardrail_response = e.message else: # Only run OUTPUT validation (INPUT was already validated in pre_call or during_call) try: @@ -1407,7 +1442,12 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): filter_result = self._prepare_guardrail_messages_for_role( messages=request_messages ) - filtered_messages = filter_result.payload_messages or mock_messages + filtered_messages = filter_result.payload_messages + if filtered_messages is None: + if self.experimental_use_latest_role_message_only: + filtered_messages = None + else: + filtered_messages = mock_messages # Bedrock will throw an error if there is no text to process if filtered_messages: diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py index 84d320a0a27..17e69fd717b 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py @@ -11,6 +11,7 @@ from fastapi import HTTPException sys.path.insert(0, os.path.abspath("../../../../../..")) +import litellm from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import ( BedrockGuardrail, @@ -1189,3 +1190,123 @@ async def test_bedrock_guardrail_blocked_content_with_masking_enabled(): print("✅ BLOCKED content with masking enabled raises exception correctly") + +def test_create_bedrock_input_content_request_skips_non_user_when_flag_enabled(): + """When experimental_use_latest_role_message_only is True, + _create_bedrock_input_content_request should skip non-user messages.""" + guardrail = BedrockGuardrail( + guardrailIdentifier="test-id", + guardrailVersion="DRAFT", + experimental_use_latest_role_message_only=True, + ) + + messages = [ + {"role": "user", "content": "hello"}, + {"role": "assistant", "content": "hi there"}, + {"role": "tool", "content": "tool result"}, + ] + + result = guardrail._create_bedrock_input_content_request(messages=messages) + content_items = result.get("content", []) + + # Only the user message content should be included + assert len(content_items) == 1 + assert content_items[0]["text"]["text"] == "hello" + + +def test_create_bedrock_input_content_request_includes_all_when_flag_disabled(): + """When experimental_use_latest_role_message_only is False, + _create_bedrock_input_content_request should include all messages.""" + guardrail = BedrockGuardrail( + guardrailIdentifier="test-id", + guardrailVersion="DRAFT", + ) + + messages = [ + {"role": "user", "content": "hello"}, + {"role": "assistant", "content": "hi there"}, + ] + + result = guardrail._create_bedrock_input_content_request(messages=messages) + content_items = result.get("content", []) + + # Both messages should be included + assert len(content_items) == 2 + + +def test_prepare_guardrail_messages_no_user_messages_returns_none(): + """When experimental_use_latest_role_message_only is True and there are no + user messages, payload_messages should be None.""" + guardrail = BedrockGuardrail( + guardrailIdentifier="test-id", + guardrailVersion="DRAFT", + experimental_use_latest_role_message_only=True, + ) + + messages = [ + {"role": "assistant", "content": "response"}, + {"role": "tool", "content": "tool result"}, + ] + + result = guardrail._prepare_guardrail_messages_for_role(messages=messages) + assert result.payload_messages is None + + +@pytest.mark.asyncio +async def test_post_call_success_hook_skips_input_when_no_user_messages_and_flag_enabled(): + """When experimental_use_latest_role_message_only is True and there are no + user messages, async_post_call_success_hook should skip INPUT validation + and only run OUTPUT validation.""" + guardrail = BedrockGuardrail( + guardrailIdentifier="test-id", + guardrailVersion="DRAFT", + experimental_use_latest_role_message_only=True, + ) + + mock_user_api_key_dict = UserAPIKeyAuth() + + mock_response = litellm.ModelResponse( + id="test-id", + choices=[ + litellm.Choices( + index=0, + message=litellm.Message(role="assistant", content="safe response"), + finish_reason="stop", + ) + ], + created=1234567890, + model="gpt-4o", + object="chat.completion", + ) + + request_data = { + "model": "gpt-4o", + "messages": [ + {"role": "assistant", "content": "previous response"}, + {"role": "tool", "content": "tool result"}, + ], + } + + call_sources = [] + + async def mock_make_bedrock_api_request(source=None, **kwargs): + call_sources.append(source) + mock_bedrock = MagicMock() + mock_bedrock.get.return_value = None + return mock_bedrock + + with patch.object( + guardrail, + "make_bedrock_api_request", + side_effect=mock_make_bedrock_api_request, + ): + await guardrail.async_post_call_success_hook( + data=request_data, + user_api_key_dict=mock_user_api_key_dict, + response=mock_response, + ) + + # Only OUTPUT should have been validated, not INPUT + assert "OUTPUT" in call_sources + assert "INPUT" not in call_sources +