mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
[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:
parent
26f6dbdd3d
commit
bc368707df
3 changed files with 355 additions and 73 deletions
|
|
@ -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,
|
||||
)
|
||||
|
||||
#########################################################################
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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}")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue