mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
fix(guardrails): skip INPUT validation when no user messages with experimental_use_latest_role_message_only
Fixes #23476
This commit is contained in:
parent
62757ff48f
commit
1ed38789ee
2 changed files with 197 additions and 36 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue