[Bug Fix] Fixes for bedrock guardrails post_call - applying to streaming responses (#12252)

* fixes for using bedrock guard

* fixes for output_content_bedrock guard

* test_convert_to_bedrock_format_input_source

* test_convert_to_bedrock_format_post_call_streaming_hook

* test bedrock guard
This commit is contained in:
Ishaan Jaff 2025-07-02 14:46:27 -07:00 committed by GitHub
parent 26f6dbdd3d
commit bc368707df
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 355 additions and 73 deletions

View file

@ -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,
)
#########################################################################

View file

@ -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
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

View file

@ -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}")