fix(bedrock guardrails): derive contextual grounding source and query from plain messages

Bedrock only runs a contextualGroundingPolicy when the ApplyGuardrail payload
carries grounding_source and query qualifiers. Callers sending ordinary system
and user messages never got those, so a configured grounding threshold was
silently skipped on /v1/chat/completions and /guardrails/apply_guardrail.

When no explicit grounding_source or query tags are present, system and
developer text is sent as grounding_source and the latest user message as
query. The apply_guardrail response branch now forwards the request messages,
which it previously dropped.

Resolves LIT-4224

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-14 21:18:08 +00:00
parent 4123b4bc2b
commit 26e2e208ad
2 changed files with 185 additions and 1 deletions

View file

@ -493,6 +493,9 @@ 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).
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 +507,33 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
and role in _GROUNDING_SOURCE_TRUSTED_ROLES
):
grounding.append(block)
return grounding
return grounding or self._derive_grounding_blocks_from_plain_messages(messages)
def _derive_grounding_blocks_from_plain_messages(
self, messages: list[AllMessageValues] | None
) -> list[QualifiedTextBlock]:
"""Bedrock scores grounding only when source, query and response are all present,
and rejects a source without a query, so return nothing unless both exist."""
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 +3239,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

@ -2474,6 +2474,122 @@ 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():
"""Untagged string system + user messages become grounding_source + query, so a
guardrail with a grounding threshold actually grades the response."""
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))
assert actual_request == expected_request
def test_grounding_output_derived_query_is_latest_user_turn_only():
"""In a multi-turn chat only the latest user message is the query; earlier user
turns and assistant turns are not sent as query or source. Developer messages
count as source alongside system."""
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))
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 grounding_source without a query (and vice versa), so a
request that cannot supply both from trusted roles keeps the untagged payload."""
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_explicit_tags_take_precedence_over_plain_messages():
"""A caller that tags blocks keeps full control: untagged system text is not
added as a second source and the untagged user text is not a second 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))
assert actual_request == expected_request
def test_grounding_input_ignores_plain_message_derivation():
"""Derivation is OUTPUT-only: an INPUT scan of plain system + user text stays an
untagged payload, so input policies keep scanning every block."""
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 = _input_request(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 +2713,44 @@ async def test_grounding_output_blocked_raises_400():
assert exc_info.value.status_code == 400
@pytest.mark.asyncio
async def test_apply_guardrail_response_forwards_request_messages_for_grounding():
"""/guardrails/apply_guardrail with input_type=response: the request messages
stored in request_data must reach the OUTPUT payload as grounding_source + query
around the guarded text. Before LIT-4224 the response branch dropped them, so
Bedrock never ran its contextual-grounding policy on this route."""
guardrail = _grounding_guardrail()
request_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],
}
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
#