diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 22776dfded4..35748718587 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -74,42 +74,70 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): self.guardrailIdentifier, self.guardrailVersion, ) + + + def _create_bedrock_input_content_request(self, messages: Optional[List[AllMessageValues]]) -> BedrockRequest: + """ + Create a bedrock request for the input content - the LLM request. + """ + bedrock_request: BedrockRequest = BedrockRequest(source="INPUT") + bedrock_request_content: List[BedrockContentItem] = [] + if messages is None: + return bedrock_request + for message in messages: + message_text_content: Optional[List[str]] = ( + self.get_content_for_message(message=message) + ) + if message_text_content is None: + continue + for text_content in message_text_content: + bedrock_content_item = BedrockContentItem( + text=BedrockTextContent(text=text_content) + ) + bedrock_request_content.append(bedrock_content_item) + + bedrock_request["content"] = bedrock_request_content + return bedrock_request + + def _create_bedrock_output_content_request(self, response: Union[Any, ModelResponse]) -> BedrockRequest: + """ + Create a bedrock request for the output content - the LLM response. + """ + bedrock_request: BedrockRequest = BedrockRequest(source="OUTPUT") + bedrock_request_content: List[BedrockContentItem] = [] + if isinstance(response, litellm.ModelResponse): + for choice in response.choices: + if isinstance(choice, litellm.Choices): + if choice.message.content and isinstance( + choice.message.content, str + ): + bedrock_content_item = BedrockContentItem( + text=BedrockTextContent(text=choice.message.content) + ) + bedrock_request_content.append(bedrock_content_item) + bedrock_request["content"] = bedrock_request_content + return bedrock_request def convert_to_bedrock_format( self, + source: Literal["INPUT", "OUTPUT"], messages: Optional[List[AllMessageValues]] = None, response: Optional[Union[Any, ModelResponse]] = None, ) -> BedrockRequest: - bedrock_request: BedrockRequest = BedrockRequest(source="INPUT") - bedrock_request_content: List[BedrockContentItem] = [] + """ + Convert the litellm messages/response to the bedrock request format. - if messages: - for message in messages: - message_text_content: Optional[List[str]] = ( - self.get_content_for_message(message=message) - ) - if message_text_content is None: - continue - for text_content in message_text_content: - bedrock_content_item = BedrockContentItem( - text=BedrockTextContent(text=text_content) - ) - bedrock_request_content.append(bedrock_content_item) + If source is "INPUT", then messages is required. + If source is "OUTPUT", then response is required. - bedrock_request["content"] = bedrock_request_content - if response: - bedrock_request["source"] = "OUTPUT" - if isinstance(response, litellm.ModelResponse): - for choice in response.choices: - if isinstance(choice, litellm.Choices): - if choice.message.content and isinstance( - choice.message.content, str - ): - bedrock_content_item = BedrockContentItem( - text=BedrockTextContent(text=choice.message.content) - ) - bedrock_request_content.append(bedrock_content_item) - bedrock_request["content"] = bedrock_request_content + Returns: + BedrockRequest: The bedrock request object. + """ + bedrock_request: BedrockRequest = BedrockRequest(source=source) + if source == "INPUT": + bedrock_request = self._create_bedrock_input_content_request(messages=messages) + elif source == "OUTPUT": + bedrock_request = self._create_bedrock_output_content_request(response=response) return bedrock_request #### CALL HOOKS - proxy only #### @@ -187,20 +215,27 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): return prepped_request async def make_bedrock_api_request( - self, kwargs: dict, response: Optional[Union[Any, litellm.ModelResponse]] = None + self, + source: Literal["INPUT", "OUTPUT"], + messages: Optional[List[AllMessageValues]] = None, + response: Optional[Union[Any, litellm.ModelResponse]] = None, + request_data: Optional[dict] = None ) -> BedrockGuardrailResponse: credentials, aws_region_name = self._load_credentials() bedrock_request_data: dict = dict( self.convert_to_bedrock_format( - messages=kwargs.get("messages"), response=response + source=source, + messages=messages, + response=response ) ) bedrock_guardrail_response: BedrockGuardrailResponse = ( BedrockGuardrailResponse() ) - bedrock_request_data.update( - self.get_guardrail_dynamic_request_body_params(request_data=kwargs) - ) + if request_data: + bedrock_request_data.update( + self.get_guardrail_dynamic_request_body_params(request_data=request_data) + ) prepared_request = self._prepare_request( credentials=credentials, data=bedrock_request_data, @@ -357,7 +392,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ######################################################### ########## 1. Make the Bedrock API request ########## ######################################################### - bedrock_guardrail_response = await self.make_bedrock_api_request(kwargs=data) + bedrock_guardrail_response = await self.make_bedrock_api_request( + source="INPUT", messages=new_messages, request_data=data + ) ######################################################### ######################################################### @@ -411,7 +448,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ######################################################### ########## 1. Make the Bedrock API request ########## ######################################################### - bedrock_guardrail_response = await self.make_bedrock_api_request(kwargs=data) + bedrock_guardrail_response = await self.make_bedrock_api_request( + source="INPUT", messages=new_messages, request_data=data + ) ######################################################### ######################################################### @@ -461,28 +500,20 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): return ######################################################### - ########## 1. Make the Bedrock API request ########## - ######################################################### - bedrock_guardrail_response = await self.make_bedrock_api_request( - kwargs=data, response=response - ) + ########## 1. Make parallel Bedrock API requests ########## ######################################################### + output_content_bedrock = await self.make_bedrock_api_request( + source="OUTPUT", + response=response, + request_data=data + ) # Only response ######################################################### - ########## 2. Update the messages with the guardrail response ########## + ########## 2. Apply masking to response with output guardrail response ########## ######################################################### - data["messages"] = ( - self._update_messages_with_updated_bedrock_guardrail_response( - messages=new_messages, - bedrock_guardrail_response=bedrock_guardrail_response, - ) - ) - - - ########## 3. Apply masking to response if enabled ########## self._apply_masking_to_response( response=response, - bedrock_guardrail_response=bedrock_guardrail_response, + bedrock_guardrail_response=output_content_bedrock, ) ######################################################### @@ -543,9 +574,11 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): """ Process streaming response chunks. - Collect content from the stream and make a bedrock api request to get the guardrail response. + Collect content from the stream and make parallel bedrock api requests to get the guardrail responses. """ # Import here to avoid circular imports + import asyncio + from litellm.llms.base_llm.base_model_iterator import MockResponseIterator from litellm.main import stream_chunk_builder from litellm.types.utils import TextCompletionResponse @@ -562,20 +595,29 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ) if isinstance(assembled_model_response, ModelResponse): #################################################################### - ########## 1. Make the Bedrock Apply Guardrail API request ########## + ########## 1. Make parallel Bedrock Apply Guardrail API requests ########## # Bedrock will raise an exception if this violates the guardrail policy ################################################################### - bedrock_guardrail_response = await self.make_bedrock_api_request( - kwargs=request_data, response=assembled_model_response + # Create tasks for parallel execution + input_task = self.make_bedrock_api_request( + source="INPUT", messages=request_data.get("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 + _, output_guardrail_response = await asyncio.gather( + input_task, output_task ) ######################################################################### - ########## 2. Apply masking to response if enabled ########## + ########## 2. Apply masking to response with output guardrail response ########## ######################################################################### self._apply_masking_to_response( response=assembled_model_response, - bedrock_guardrail_response=bedrock_guardrail_response, + bedrock_guardrail_response=output_guardrail_response, ) ######################################################################### diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index 01235ca7f75..21a3deda5e3 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -7,15 +7,11 @@ model_list: model: openai/* -litellm_settings: - callbacks: ["prometheus"] - prometheus_metrics_config: - - group: totals - metrics: - - litellm_spend_metric - - litellm_total_tokens - - litellm_input_tokens_metric - - litellm_output_tokens_metric - include_labels: - - requested_modela - - model \ No newline at end of file +guardrails: + - guardrail_name: "bedrock-post-guard" + litellm_params: + guardrail: bedrock # supported values: "aporia", "bedrock", "lakera" + mode: "post_call" + guardrailIdentifier: wf0hkdb5x07f # your guardrail ID on bedrock + guardrailVersion: "DRAFT" # your guardrail version on bedrock + default_on: true \ No newline at end of file diff --git a/tests/guardrails_tests/test_bedrock_guardrails.py b/tests/guardrails_tests/test_bedrock_guardrails.py index a3df486bfa1..1e08c7a4b9f 100644 --- a/tests/guardrails_tests/test_bedrock_guardrails.py +++ b/tests/guardrails_tests/test_bedrock_guardrails.py @@ -254,8 +254,9 @@ async def test_bedrock_guardrails_streaming_request_body_mock(): # Call the method that should make the Bedrock API request await guardrail.make_bedrock_api_request( - kwargs=request_data, - response=mock_response + source="OUTPUT", + response=mock_response, + request_data=request_data ) # Verify the API call was made @@ -278,7 +279,6 @@ async def test_bedrock_guardrails_streaming_request_body_mock(): expected_body = { 'source': 'OUTPUT', 'content': [ - {'text': {'text': "what's the capital of spain?"}}, {'text': {'text': 'The capital of Spain is Madrid.'}} ] } @@ -322,7 +322,7 @@ async def test_bedrock_guardrail_aws_param_persistence(): mock_response.status_code = 200 mock_response.json = MagicMock(return_value={"action": "NONE", "outputs": []}) mock_post.return_value = mock_response - await guardrail.make_bedrock_api_request(kwargs=request_data, response=None) + await guardrail.make_bedrock_api_request(source="INPUT", messages=request_data.get("messages"), request_data=request_data) assert mock_get_creds.call_count == 3 for call in mock_get_creds.call_args_list: @@ -777,3 +777,247 @@ async def test_bedrock_guardrail_response_pii_masking_streaming(): assert "Sure! My email is {EMAIL} and SSN is {US_SSN}" == full_content print("✓ Streaming response PII masking test passed") + +@pytest.mark.asyncio +async def test_convert_to_bedrock_format_input_source(): + """Test convert_to_bedrock_format with INPUT source and mock messages""" + from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import BedrockGuardrail + from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import BedrockRequest + from unittest.mock import patch + + # Create the guardrail instance + guardrail = BedrockGuardrail( + guardrailIdentifier="test-guardrail", + guardrailVersion="DRAFT" + ) + + # Mock messages + mock_messages = [ + {"role": "user", "content": "Hello, how are you?"}, + {"role": "assistant", "content": "I'm doing well, thank you!"}, + {"role": "user", "content": [ + {"type": "text", "text": "What's the weather like?"}, + {"type": "text", "text": "Is it sunny today?"} + ]} + ] + + # Call the method + result = guardrail.convert_to_bedrock_format( + source="INPUT", + messages=mock_messages + ) + + # Verify the result structure + assert isinstance(result, dict) + assert result.get("source") == "INPUT" + assert "content" in result + assert isinstance(result.get("content"), list) + + # Verify content items + expected_content_items = [ + {"text": {"text": "Hello, how are you?"}}, + {"text": {"text": "I'm doing well, thank you!"}}, + {"text": {"text": "What's the weather like?"}}, + {"text": {"text": "Is it sunny today?"}} + ] + + assert result.get("content") == expected_content_items + print("✅ INPUT source test passed - result:", result) + + +@pytest.mark.asyncio +async def test_convert_to_bedrock_format_output_source(): + """Test convert_to_bedrock_format with OUTPUT source and mock ModelResponse""" + from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import BedrockGuardrail + from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import BedrockRequest + import litellm + from unittest.mock import patch + + # Create the guardrail instance + guardrail = BedrockGuardrail( + guardrailIdentifier="test-guardrail", + guardrailVersion="DRAFT" + ) + + # Mock ModelResponse + mock_response = litellm.ModelResponse( + id="test-response-id", + choices=[ + litellm.Choices( + index=0, + message=litellm.Message( + role="assistant", + content="This is a test response from the model." + ), + finish_reason="stop" + ), + litellm.Choices( + index=1, + message=litellm.Message( + role="assistant", + content="This is a second choice response." + ), + finish_reason="stop" + ) + ], + created=1234567890, + model="gpt-4o", + object="chat.completion" + ) + + # Call the method + result = guardrail.convert_to_bedrock_format( + source="OUTPUT", + response=mock_response + ) + + # Verify the result structure + assert isinstance(result, dict) + assert result.get("source") == "OUTPUT" + assert "content" in result + assert isinstance(result.get("content"), list) + + # Verify content items - should contain both choice contents + expected_content_items = [ + {"text": {"text": "This is a test response from the model."}}, + {"text": {"text": "This is a second choice response."}} + ] + + assert result.get("content") == expected_content_items + print("✅ OUTPUT source test passed - result:", result) + + +@pytest.mark.asyncio +async def test_convert_to_bedrock_format_post_call_streaming_hook(): + """Test async_post_call_streaming_iterator_hook makes OUTPUT bedrock request and applies masking""" + from unittest.mock import AsyncMock, MagicMock, patch + from litellm.proxy._types import UserAPIKeyAuth + from litellm.types.utils import ModelResponseStream + import litellm + + # Create proper mock objects + mock_user_api_key_dict = UserAPIKeyAuth() + + # Create guardrail instance + guardrail = BedrockGuardrail( + guardrailIdentifier="test-guardrail", + guardrailVersion="DRAFT" + ) + + # Mock streaming chunks that contain PII + async def mock_streaming_response(): + chunks = [ + ModelResponseStream( + id="test-id", + choices=[ + litellm.utils.StreamingChoices( + index=0, + delta=litellm.utils.Delta(content="My email is "), + finish_reason=None + ) + ], + created=1234567890, + model="gpt-4o", + object="chat.completion.chunk" + ), + ModelResponseStream( + id="test-id", + choices=[ + litellm.utils.StreamingChoices( + index=0, + delta=litellm.utils.Delta(content="john@example.com"), + finish_reason="stop" + ) + ], + created=1234567890, + model="gpt-4o", + object="chat.completion.chunk" + ) + ] + for chunk in chunks: + yield chunk + + # Mock Bedrock API response with PII masking + mock_bedrock_response = MagicMock() + mock_bedrock_response.status_code = 200 + mock_bedrock_response.json.return_value = { + "action": "GUARDRAIL_INTERVENED", + "outputs": [{ + "text": "My email is {EMAIL}" + }], + "assessments": [{ + "sensitiveInformationPolicy": { + "piiEntities": [{ + "type": "EMAIL", + "match": "john@example.com", + "action": "ANONYMIZED" + }] + } + }] + } + + request_data = { + "model": "gpt-4o", + "messages": [ + {"role": "user", "content": "What's your email?"} + ], + "stream": True + } + + # Track which bedrock API calls were made + bedrock_calls = [] + + # Mock the make_bedrock_api_request method to track calls + async def mock_make_bedrock_api_request(source, messages=None, response=None, request_data=None): + bedrock_calls.append({ + "source": source, + "messages": messages, + "response": response, + "request_data": request_data + }) + # Return the mock bedrock response + from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import BedrockGuardrailResponse + return BedrockGuardrailResponse(**mock_bedrock_response.json()) + + # Patch the bedrock API request method + with patch.object(guardrail, 'make_bedrock_api_request', side_effect=mock_make_bedrock_api_request): + + # Call the streaming hook + result_generator = guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=mock_user_api_key_dict, + response=mock_streaming_response(), + request_data=request_data + ) + + # Collect all chunks from the result + result_chunks = [] + async for chunk in result_generator: + result_chunks.append(chunk) + + # Verify bedrock API calls were made + assert len(bedrock_calls) == 2, f"Expected 2 bedrock calls (INPUT and OUTPUT), got {len(bedrock_calls)}" + + # Find the OUTPUT call + output_calls = [call for call in bedrock_calls if call["source"] == "OUTPUT"] + assert len(output_calls) == 1, f"Expected 1 OUTPUT call, got {len(output_calls)}" + + output_call = output_calls[0] + assert output_call["source"] == "OUTPUT" + assert output_call["response"] is not None + assert output_call["messages"] is None # OUTPUT calls don't need messages + + # Verify that the response content was masked + # The streaming chunks should now contain the masked content + full_content = "" + for chunk in result_chunks: + if hasattr(chunk, 'choices') and chunk.choices: + if hasattr(chunk.choices[0], 'delta') and chunk.choices[0].delta.content: + full_content += chunk.choices[0].delta.content + + # The content should be masked (contains {EMAIL} instead of john@example.com) + assert "{EMAIL}" in full_content, f"Expected masked content with {{EMAIL}}, got: {full_content}" + assert "john@example.com" not in full_content, f"Original email should be masked, got: {full_content}" + + print("✅ Post-call streaming hook test passed - OUTPUT source used for masking") + print(f"✅ Bedrock calls made: {[call['source'] for call in bedrock_calls]}") + print(f"✅ Final masked content: {full_content}")