fix(guardrails): skip INPUT validation when no user messages with experimental_use_latest_role_message_only

Fixes #23476
This commit is contained in:
michelligabriele 2026-04-08 18:14:15 +02:00
parent 62757ff48f
commit 1ed38789ee
No known key found for this signature in database
2 changed files with 197 additions and 36 deletions

View file

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

View file

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