Merge pull request #41132 from BerriAI/litellm_bedrock_grounding_from_plain_messages

fix(bedrock guardrails): derive contextual grounding source and query from plain messages
This commit is contained in:
yucheng-berri 2026-09-14 17:56:58 -07:00 committed by GitHub
commit 8481bc27f9
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 290 additions and 7 deletions

View file

@ -11503,6 +11503,12 @@
"description": "Enable content moderation to check for harmful content (harassment, hate speech, etc.).",
"title": "Content Moderation Check"
},
"contextual_grounding_from_messages": {
"default": false,
"description": "ApplyGuardrail: when True, post-call scans of a request with no grounding_source / query content parts send the system and developer messages as the grounding source and the latest user message as the query, so the guardrail's contextual grounding policy can score the response. Bedrock bills contextual grounding units for these scans and rejects queries, sources and responses over its contextual grounding length limits, so leave this off for guardrails without a contextual grounding policy. Default False: plain messages are never sent as grounding context.",
"title": "Contextual Grounding From Messages",
"type": "boolean"
},
"credentials": {
"anyOf": [
{

View file

@ -244,6 +244,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
prompt_attack_threshold: float | None = 0.5,
pii_confidence_threshold: float | None = 0.5,
chunk_budget_chars: int = BEDROCK_APPLY_GUARDRAIL_CHUNK_BUDGET_CHARS,
contextual_grounding_from_messages: bool = False,
streaming_buffer_until_moderated: bool | None = None,
streaming_sampling_rate: int | None = None,
streaming_end_of_stream_only: bool | None = None,
@ -265,6 +266,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
self.guardrailVersion = guardrailVersion
self.guardrail_provider = "bedrock"
self.chunk_budget_chars = chunk_budget_chars
self.contextual_grounding_from_messages = contextual_grounding_from_messages
self.experimental_use_latest_role_message_only = bool(kwargs.get("experimental_use_latest_role_message_only"))
# Resource-less, detect-only InvokeGuardrailChecks mode. Present `checks`
@ -459,8 +461,8 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
"""
Flatten a message into text blocks, preserving any contextual-grounding
qualifier carried by the content-block ``type`` (grounding_source / query).
Untagged text keeps ``qualifier=None`` so the payload is unchanged for
callers that do not use grounding.
Untagged text keeps ``qualifier=None``; the OUTPUT scan decides whether to
derive grounding qualifiers from it.
"""
content: Final = message.get("content")
if content is None:
@ -493,6 +495,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
result carrying externally-influenced content can supply fake evidence for the
contextual-grounding check to grade the response against. ``query`` is accepted
from any role (it is the user's question).
With ``contextual_grounding_from_messages`` on, a request with no tagged blocks
falls back to the plain messages: system / developer text is the grounding
source and the latest user message is the query.
"""
grounding: Final[list[QualifiedTextBlock]] = []
for message in messages or []:
@ -504,7 +510,33 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
and role in _GROUNDING_SOURCE_TRUSTED_ROLES
):
grounding.append(block)
return grounding
if grounding or not self.contextual_grounding_from_messages:
return grounding
return self._derive_grounding_blocks_from_plain_messages(messages)
def _derive_grounding_blocks_from_plain_messages(
self, messages: list[AllMessageValues] | None
) -> list[QualifiedTextBlock]:
if not messages:
return []
latest_user_index: Final = self._find_latest_message_index(messages, target_role="user")
if latest_user_index is None:
return []
sources: Final = tuple(
QualifiedTextBlock(text=block.text, qualifier="grounding_source")
for message in messages
if message.get("role") in _GROUNDING_SOURCE_TRUSTED_ROLES
for block in self.get_content_items_for_message(message=message) or []
if block.text
)
queries: Final = tuple(
QualifiedTextBlock(text=block.text, qualifier="query")
for block in self.get_content_items_for_message(message=messages[latest_user_index]) or []
if block.text
)
if not sources or not queries:
return []
return [*sources, *queries]
def supports_scan_only_tool_results(self) -> bool:
return self.experimental_use_latest_role_message_only is not True
@ -3210,6 +3242,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
bedrock_response = await self.make_bedrock_api_request(
source="OUTPUT",
response=synthetic_response,
messages=request_data.get("messages"),
request_data=request_data,
logging_event_type=_log_hook,
)

View file

@ -23,6 +23,7 @@ def initialize_bedrock(litellm_params: LitellmParams, guardrail: Guardrail):
prompt_attack_threshold=litellm_params.prompt_attack_threshold,
pii_confidence_threshold=litellm_params.pii_confidence_threshold,
chunk_budget_chars=litellm_params.chunk_budget_chars,
contextual_grounding_from_messages=litellm_params.contextual_grounding_from_messages,
default_on=litellm_params.default_on,
disable_exception_on_block=litellm_params.disable_exception_on_block,
mask_request_content=litellm_params.mask_request_content,

View file

@ -552,6 +552,16 @@ class BedrockGuardrailConfigModel(BaseModel):
"still rejects is bisected automatically, so this value only trades round trips against "
"batch size and cannot fail a request on its own.",
)
contextual_grounding_from_messages: bool = Field(
default=False,
description="ApplyGuardrail: when True, post-call scans of a request with no grounding_source / "
"query content parts send the system and developer messages as the grounding source and "
"the latest user message as the query, so the guardrail's contextual grounding policy can "
"score the response. Bedrock bills contextual grounding units for these scans and rejects "
"queries, sources and responses over its contextual grounding length limits, so leave this "
"off for guardrails without a contextual grounding policy. Default False: plain messages "
"are never sent as grounding context.",
)
class BedrockGuardrailStreamingParams(BaseModel):

View file

@ -2375,8 +2375,12 @@ _GROUNDING_QUERY_TEXT = "What is the capital of Japan?"
_GROUNDING_RESPONSE_TEXT = "The capital of Japan is Tokyo."
def _grounding_guardrail() -> BedrockGuardrail:
return BedrockGuardrail(guardrailIdentifier="test-guardrail", guardrailVersion="DRAFT")
def _grounding_guardrail(from_messages: bool = False) -> BedrockGuardrail:
return BedrockGuardrail(
guardrailIdentifier="test-guardrail",
guardrailVersion="DRAFT",
contextual_grounding_from_messages=from_messages,
)
def _grounding_messages() -> list:
@ -2418,9 +2422,11 @@ def _input_request(messages: list) -> dict:
return _grounding_guardrail().convert_to_bedrock_format(source="INPUT", messages=messages)
def _output_request(messages: list, response=None) -> dict:
def _output_request(messages: list, response=None, from_messages: bool = False) -> dict:
"""Arrange a guardrail and act: build the Bedrock OUTPUT payload."""
return _grounding_guardrail().convert_to_bedrock_format(source="OUTPUT", response=response, messages=messages)
return _grounding_guardrail(from_messages).convert_to_bedrock_format(
source="OUTPUT", response=response, messages=messages
)
def test_grounding_input_strips_grounding_and_query_qualifiers():
@ -2474,6 +2480,131 @@ def test_grounding_output_keeps_legacy_payload_without_tags():
assert actual_request == expected_request
def test_grounding_output_derives_source_and_query_from_plain_messages():
"""Flag on: untagged system + user text is sent as grounding_source + query."""
messages = [
{"role": "system", "content": _GROUNDING_SOURCE_TEXT},
{"role": "user", "content": _GROUNDING_QUERY_TEXT},
]
expected_request = {
"source": "OUTPUT",
"content": [_GROUNDING_SOURCE_BLOCK, _QUERY_BLOCK, _GUARD_BLOCK],
}
actual_request = _output_request(messages, _model_response(_GROUNDING_RESPONSE_TEXT), from_messages=True)
assert actual_request == expected_request
def test_grounding_output_plain_messages_stay_legacy_when_flag_is_off():
"""Default config: plain system + user text is never sent as grounding context."""
messages = [
{"role": "system", "content": _GROUNDING_SOURCE_TEXT},
{"role": "user", "content": _GROUNDING_QUERY_TEXT},
]
expected_request = {"source": "OUTPUT", "content": [{"text": {"text": _GROUNDING_RESPONSE_TEXT}}]}
actual_request = _output_request(messages, _model_response(_GROUNDING_RESPONSE_TEXT))
assert actual_request == expected_request
def test_grounding_output_derived_query_is_latest_user_turn_only():
"""Only the latest user turn is the query; system and developer turns are the source."""
developer_text = "Answer in one sentence."
messages = [
{"role": "system", "content": _GROUNDING_SOURCE_TEXT},
{"role": "developer", "content": [{"type": "text", "text": developer_text}]},
{"role": "user", "content": "Hi"},
{"role": "assistant", "content": "Hello, how can I help?"},
{"role": "user", "content": _GROUNDING_QUERY_TEXT},
]
expected_request = {
"source": "OUTPUT",
"content": [
_GROUNDING_SOURCE_BLOCK,
{"text": {"text": developer_text, "qualifiers": ["grounding_source"]}},
_QUERY_BLOCK,
_GUARD_BLOCK,
],
}
actual_request = _output_request(messages, _model_response(_GROUNDING_RESPONSE_TEXT), from_messages=True)
assert actual_request == expected_request
@pytest.mark.parametrize(
"messages",
[
pytest.param([{"role": "system", "content": _GROUNDING_SOURCE_TEXT}], id="system-without-user"),
pytest.param(
[
{"role": "tool", "content": _GROUNDING_SOURCE_TEXT, "tool_call_id": "c1"},
{"role": "user", "content": _GROUNDING_QUERY_TEXT},
],
id="tool-result-is-not-a-source",
),
pytest.param(
[
{"role": "system", "content": ""},
{"role": "user", "content": _GROUNDING_QUERY_TEXT},
],
id="empty-system-prompt",
),
pytest.param(
[
{"role": "system", "content": _GROUNDING_SOURCE_TEXT},
{"role": "user", "content": [{"type": "image_url", "image_url": {"url": "https://x.test/a.png"}}]},
],
id="image-only-user-turn",
),
],
)
def test_grounding_output_stays_legacy_when_plain_source_or_query_is_missing(messages):
"""Bedrock rejects a source without a query and vice versa, so send neither."""
expected_request = {"source": "OUTPUT", "content": [{"text": {"text": _GROUNDING_RESPONSE_TEXT}}]}
actual_request = _output_request(messages, _model_response(_GROUNDING_RESPONSE_TEXT), from_messages=True)
assert actual_request == expected_request
def test_grounding_output_explicit_tags_take_precedence_over_plain_messages():
"""Tagged blocks win: untagged text around them is not added as source or query."""
messages = [
{"role": "system", "content": "You are a helpful assistant."},
*_grounding_messages(),
{"role": "user", "content": "Please be brief."},
]
expected_request = {
"source": "OUTPUT",
"content": [_GROUNDING_SOURCE_BLOCK, _QUERY_BLOCK, _GUARD_BLOCK],
}
actual_request = _output_request(messages, _model_response(_GROUNDING_RESPONSE_TEXT), from_messages=True)
assert actual_request == expected_request
def test_grounding_input_ignores_plain_message_derivation():
"""INPUT scans never derive grounding qualifiers from plain messages."""
messages = [
{"role": "system", "content": _GROUNDING_SOURCE_TEXT},
{"role": "user", "content": _GROUNDING_QUERY_TEXT},
]
expected_request = {
"source": "INPUT",
"content": [{"text": {"text": _GROUNDING_SOURCE_TEXT}}, {"text": {"text": _GROUNDING_QUERY_TEXT}}],
}
actual_request = _grounding_guardrail(from_messages=True).convert_to_bedrock_format(
source="INPUT", messages=messages
)
assert actual_request == expected_request
def test_grounding_output_combines_multiple_sources():
"""Every grounding_source block is emitted; Bedrock combines them into one corpus."""
uk_source_text = "London is the capital of UK."
@ -2597,6 +2728,56 @@ async def test_grounding_output_blocked_raises_400():
assert exc_info.value.status_code == 400
@pytest.mark.asyncio
@pytest.mark.parametrize(
"from_messages, request_messages",
[
(
True,
[
{"role": "system", "content": _GROUNDING_SOURCE_TEXT},
{"role": "user", "content": _GROUNDING_QUERY_TEXT},
],
),
(
False,
[
{"role": "system", "content": [{"type": "grounding_source", "text": _GROUNDING_SOURCE_TEXT}]},
{"role": "user", "content": [{"type": "query", "text": _GROUNDING_QUERY_TEXT}]},
],
),
],
ids=["plain-messages-flag-on", "tagged-messages-flag-off"],
)
async def test_apply_guardrail_response_forwards_request_messages_for_grounding(from_messages, request_messages):
guardrail = _grounding_guardrail(from_messages=from_messages)
expected_request = {
"source": "OUTPUT",
"content": [_GROUNDING_SOURCE_BLOCK, _QUERY_BLOCK, _GUARD_BLOCK],
}
mock_credentials = MagicMock()
mock_credentials.access_key = "test-access-key"
mock_credentials.secret_key = "test-secret-key"
mock_credentials.token = None
with (
patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post,
patch.object(guardrail, "_load_credentials", return_value=(mock_credentials, "us-east-1")),
patch.object(guardrail, "_prepare_request", return_value=MagicMock()) as mock_prepare,
):
mock_post.return_value = _passing_bedrock_httpx_response(_GROUNDING_RESPONSE_TEXT)
await guardrail.apply_guardrail(
inputs={"texts": [_GROUNDING_RESPONSE_TEXT]},
request_data={"messages": request_messages},
input_type="response",
)
assert mock_prepare.call_count == 1
assert json.loads(json.dumps(mock_prepare.call_args.kwargs["data"])) == expected_request
###############################################################################
# LIT-4186: disable_exception_on_block regression tests
#

View file

@ -71,6 +71,52 @@ def test_initialize_bedrock_forwards_chunk_budget_chars():
assert initialized[-1].chunk_budget_chars == 60_000
def test_initialize_bedrock_forwards_contextual_grounding_from_messages():
"""`contextual_grounding_from_messages: true` in config.yaml must make the post-call
payload carry the plain system prompt and user turn as grounding_source and query."""
import litellm
from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import BedrockGuardrail
from litellm.types.utils import Choices, Message, ModelResponse
test_guardrail = {
"guardrail_name": "test_bedrock_grounding_from_messages",
"litellm_params": {
"guardrail": SupportedGuardrailIntegrations.BEDROCK.value,
"mode": "post_call",
"guardrailIdentifier": "test-guardrail",
"guardrailVersion": "DRAFT",
"contextual_grounding_from_messages": True,
},
}
messages = [
{"role": "system", "content": "Returns are accepted for 30 days."},
{"role": "user", "content": "How long is the return window?"},
]
response = ModelResponse(
choices=[Choices(index=0, message=Message(role="assistant", content="30 days."), finish_reason="stop")]
)
expected_request = {
"source": "OUTPUT",
"content": [
{"text": {"text": "Returns are accepted for 30 days.", "qualifiers": ["grounding_source"]}},
{"text": {"text": "How long is the return window?", "qualifiers": ["query"]}},
{"text": {"text": "30 days.", "qualifiers": ["guard_content"]}},
],
}
guardrail_handler = InMemoryGuardrailHandler()
guardrail_handler.initialize_guardrail(guardrail=test_guardrail)
initialized = [
callback
for callback in litellm.callbacks
if isinstance(callback, BedrockGuardrail) and callback.guardrail_name == "test_bedrock_grounding_from_messages"
]
assert initialized, "bedrock guardrail was not registered as a callback"
actual_request = initialized[-1].convert_to_bedrock_format(source="OUTPUT", response=response, messages=messages)
assert json.loads(json.dumps(actual_request)) == expected_request
def test_initialize_guardrail_preserves_guardrail_info():
"""
Regression (LIT-2529): initialize_guardrail must carry guardrail_info into the

View file

@ -31006,6 +31006,12 @@ export interface components {
* @description Enable content moderation to check for harmful content (harassment, hate speech, etc.).
*/
content_moderation_check?: boolean | null;
/**
* Contextual Grounding From Messages
* @description ApplyGuardrail: when True, post-call scans of a request with no grounding_source / query content parts send the system and developer messages as the grounding source and the latest user message as the query, so the guardrail's contextual grounding policy can score the response. Bedrock bills contextual grounding units for these scans and rejects queries, sources and responses over its contextual grounding length limits, so leave this off for guardrails without a contextual grounding policy. Default False: plain messages are never sent as grounding context.
* @default false
*/
contextual_grounding_from_messages: boolean;
/**
* Credentials
* @description Path to Google Cloud credentials JSON file or JSON string