feat(bedrock guardrails): support contextual grounding qualifiers

Bedrock contextual grounding scores a model response against a reference
source and the user query, expressed via a per-content-block `qualifiers`
array on ApplyGuardrail. The guardrail hook previously sent plain text only,
so grounding could not be driven through it even though the response-side
contextualGroundingPolicy parsing already existed.

Callers now tag message content blocks `{"type":"grounding_source"}` /
`{"type":"query"}` (mirroring the existing `guarded_text` marker). On the
generate path the bedrock converse transform renders them as plain text; at
post_call the hook harvests them from the request and assembles one
ApplyGuardrail(OUTPUT) call carrying grounding_source + query + the response
(as guard_content). Requests without these tags produce a byte-identical
payload, so existing behaviour is unchanged.
This commit is contained in:
João Costa 2026-06-10 01:55:08 +01:00
parent 7247c09725
commit 0cbd7f4a14
6 changed files with 293 additions and 112 deletions

View file

@ -4741,6 +4741,12 @@ class BedrockConverseMessagesProcessor:
guardContent={"text": {"text": element["text"]}}
)
_parts.append(_part)
elif element["type"] in ("grounding_source", "query"):
# Contextual grounding tags are guardrail metadata; the
# model only needs the underlying text, so render them
# as plain text on the generate path.
_part = BedrockContentBlock(text=element["text"])
_parts.append(_part)
elif element["type"] == "image_url":
format: Optional[str] = None
if isinstance(element["image_url"], dict):
@ -5173,6 +5179,12 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915
guardContent={"text": {"text": element["text"]}}
)
_parts.append(_part)
elif element["type"] in ("grounding_source", "query"):
# Contextual grounding tags are guardrail metadata; the
# model only needs the underlying text, so render them as
# plain text on the generate path.
_part = BedrockContentBlock(text=element["text"])
_parts.append(_part)
elif element["type"] == "image_url":
format: Optional[str] = None
if isinstance(element["image_url"], dict):

View file

@ -49,6 +49,7 @@ from litellm.types.llms.openai import AllMessageValues, ChatCompletionUserMessag
from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
BedrockContentItem,
BedrockGuardrailOutput,
BedrockGuardrailQualifier,
BedrockGuardrailResponse,
BedrockRequest,
BedrockTextContent,
@ -74,6 +75,29 @@ from litellm.types.utils import (
GUARDRAIL_NAME = "bedrock"
_BEDROCK_DYNAMIC_BODY_DENYLIST = frozenset({"content", "source"})
# Maps an OpenAI message content-block ``type`` to the Bedrock guardrail qualifier
# it represents, so callers can drive contextual grounding by tagging their content.
# The model response is qualified as ``guard_content`` directly by the OUTPUT builder;
# the existing ``guarded_text`` marker is intentionally left unmapped here so its
# guardrail-hook payload is unchanged by this feature.
_CONTENT_TYPE_TO_QUALIFIER: Dict[str, BedrockGuardrailQualifier] = {
"grounding_source": "grounding_source",
"query": "query",
}
# Roles whose ``grounding_source`` blocks are trusted as reference material for the
# contextual-grounding check. Only app-authored roles qualify: ``tool``/``function``
# results and ``user`` content can carry caller- or externally-influenced text, which
# must not be graded against as if it were the application's own source material.
_GROUNDING_SOURCE_TRUSTED_ROLES = frozenset({"system", "developer"})
class QualifiedTextBlock(NamedTuple):
"""A piece of message text paired with its Bedrock grounding qualifier (if any)."""
text: str
qualifier: Optional[BedrockGuardrailQualifier]
class GuardrailMessageFilterResult(NamedTuple):
payload_messages: Optional[List[AllMessageValues]]
@ -164,41 +188,71 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
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:
blocks = self.get_content_items_for_message(message=message)
if blocks is None:
continue
for text_content in message_text_content:
bedrock_content_item = BedrockContentItem(
text=BedrockTextContent(text=text_content)
for block in blocks:
# INPUT scans send plain text only. Grounding qualifiers are attached
# exclusively when assembling the OUTPUT request, so a caller cannot use
# a grounding_source/query tag to change how input-safety policies treat
# their content (which would be an input-guardrail bypass).
bedrock_request_content.append(
BedrockContentItem(text=BedrockTextContent(text=block.text))
)
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]
self,
response: Union[Any, ModelResponse],
messages: Optional[List[AllMessageValues]] = None,
) -> BedrockRequest:
"""
Create a bedrock request for the output content - the LLM response.
Contextual grounding grades the response against the reference source and
the user query from the request. When the request tagged any
``grounding_source``/``query`` blocks, they are emitted first and the
response is qualified as ``guard_content`` so Bedrock can score grounding.
Without such tags the payload is the legacy single response block.
"""
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
grounding_blocks = self._collect_grounding_blocks(messages)
bedrock_request_content: List[BedrockContentItem] = [
self._build_content_item(block) for block in grounding_blocks
]
has_grounding = len(bedrock_request_content) > 0
# Append the response (the content to guard) after any grounding blocks; assign
# unconditionally so harvested grounding blocks survive a non-ModelResponse input.
bedrock_request_content.extend(
self._build_response_content_items(response, has_grounding=has_grounding)
)
bedrock_request["content"] = bedrock_request_content
return bedrock_request
def _build_response_content_items(
self, response: Union[Any, ModelResponse], has_grounding: bool
) -> List[BedrockContentItem]:
"""Build content item(s) from the model response. When the request supplied
grounding, the response is qualified ``guard_content`` so Bedrock can score it.
"""
items: List[BedrockContentItem] = []
if not isinstance(response, litellm.ModelResponse):
return items
for choice in response.choices:
if (
isinstance(choice, litellm.Choices)
and isinstance(choice.message.content, str)
and choice.message.content
):
block = QualifiedTextBlock(
text=choice.message.content,
qualifier="guard_content" if has_grounding else None,
)
items.append(self._build_content_item(block))
return items
def convert_to_bedrock_format(
self,
source: Literal["INPUT", "OUTPUT"],
@ -221,10 +275,68 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
)
elif source == "OUTPUT":
bedrock_request = self._create_bedrock_output_content_request(
response=response
response=response, messages=messages
)
return bedrock_request
def get_content_items_for_message(
self, message: AllMessageValues
) -> Optional[List[QualifiedTextBlock]]:
"""
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.
"""
content = message.get("content")
if content is None:
return None
blocks: List[QualifiedTextBlock] = []
if isinstance(content, str):
blocks.append(QualifiedTextBlock(text=content, qualifier=None))
elif isinstance(content, list):
for item in content:
if isinstance(item, dict) and "text" in item:
qualifier = _CONTENT_TYPE_TO_QUALIFIER.get(item.get("type", ""))
blocks.append(
QualifiedTextBlock(text=item["text"], qualifier=qualifier)
)
elif isinstance(item, str):
blocks.append(QualifiedTextBlock(text=item, qualifier=None))
return blocks
def _build_content_item(self, block: QualifiedTextBlock) -> BedrockContentItem:
"""Build a Bedrock content item, attaching qualifiers only when present."""
text_content = BedrockTextContent(text=block.text)
if block.qualifier is not None:
text_content["qualifiers"] = [block.qualifier]
return BedrockContentItem(text=text_content)
def _collect_grounding_blocks(
self, messages: Optional[List[AllMessageValues]]
) -> List[QualifiedTextBlock]:
"""Harvest grounding_source/query blocks from the request for an OUTPUT scan.
``grounding_source`` is honored only from app-authored roles (system /
developer). A grounding_source tag on a ``user``, ``tool`` or ``function``
message is ignored, so neither a forwarded end-user message nor a tool/function
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).
"""
grounding: List[QualifiedTextBlock] = []
for message in messages or []:
role = message.get("role")
for block in self.get_content_items_for_message(message=message) or []:
if block.qualifier == "query":
grounding.append(block)
elif (
block.qualifier == "grounding_source"
and role in _GROUNDING_SOURCE_TRUSTED_ROLES
):
grounding.append(block)
return grounding
def _prepare_guardrail_messages_for_role(
self,
messages: Optional[List[AllMessageValues]],
@ -1169,6 +1281,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
output_content_bedrock = await self.make_bedrock_api_request(
source="OUTPUT",
response=response,
messages=new_messages,
request_data=data,
logging_event_type=GuardrailEventHooks.post_call,
)
@ -1281,6 +1394,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
output_guardrail_response = await self.make_bedrock_api_request(
source="OUTPUT",
response=assembled_model_response,
messages=request_data.get("messages"),
request_data=request_data,
logging_event_type=GuardrailEventHooks.post_call,
)
@ -1414,28 +1528,6 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
return new_content, masking_index
def get_content_for_message(self, message: AllMessageValues) -> Optional[List[str]]:
"""
Get the content for a message.
For bedrock guardrails we create a list of all the text content in the message.
If a message has a list of content items, we flatten the list and return a list of text content.
"""
message_text_content = []
content = message.get("content")
if content is None:
return None
if isinstance(content, str):
message_text_content.append(content)
elif isinstance(content, list):
for item in content:
if isinstance(item, dict) and "text" in item:
message_text_content.append(item["text"])
elif isinstance(item, str):
message_text_content.append(item)
return message_text_content
def _apply_masking_to_response(
self,
response: Union[ModelResponse, Any],

View file

@ -802,6 +802,8 @@ ValidUserMessageContentTypes = [
"audio_url",
"document",
"guarded_text",
"grounding_source",
"query",
"video_url",
"file",
] # used for validating user messages. Prevent users from accidentally sending anthropic messages.
@ -813,6 +815,8 @@ ValidUserMessageContentTypesLiteral = Literal[
"audio_url",
"document",
"guarded_text",
"grounding_source",
"query",
"video_url",
"file",
]
@ -824,6 +828,8 @@ ValidUserMessageContentTypes = [
"audio_url",
"document",
"guarded_text",
"grounding_source",
"query",
"video_url",
"file",
] # used for validating user messages. Prevent users from accidentally sending anthropic messages.
@ -851,6 +857,8 @@ ValidChatCompletionMessageContentTypesLiteral = Literal[
"audio_url",
"document",
"guarded_text",
"grounding_source",
"query",
"video_url",
"file",
"thinking",
@ -864,6 +872,8 @@ ValidChatCompletionMessageContentTypes = [
"audio_url",
"document",
"guarded_text",
"grounding_source",
"query",
"video_url",
"file",
"thinking",

View file

@ -2,9 +2,14 @@ from typing import Any, Dict, List, Literal, Optional, Union
from typing_extensions import TypedDict
# Bedrock contextual grounding tags each content block so the guardrail knows
# which text is the reference source, the user question, and the content to grade.
BedrockGuardrailQualifier = Literal["grounding_source", "query", "guard_content"]
class BedrockTextContent(TypedDict, total=False):
text: str
qualifiers: List[BedrockGuardrailQualifier]
class BedrockContentItem(TypedDict, total=False):

View file

@ -1160,7 +1160,10 @@ async def test_convert_to_bedrock_format_post_call_streaming_hook():
output_call = bedrock_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
# OUTPUT forwards the request messages so contextual grounding can pull
# grounding_source/query blocks from them even on streamed responses. A
# plain-text (non-grounding) request still yields the single-block payload.
assert output_call["messages"] == request_data["messages"]
# Verify that the response content was masked
# The streaming chunks should now contain the masked content

View file

@ -2513,8 +2513,8 @@ async def test_post_call_success_hook_only_runs_output_scan():
# `qualifiers` array (grounding_source / query / guard_content). A caller marks
# message content blocks `{"type": "grounding_source", ...}` / `{"type": "query", ...}`;
# at post_call the hook assembles one source="OUTPUT" call carrying the source +
# query + the model response (as guard_content). A request without these tags must
# produce a byte-identical payload to before the feature existed.
# query + the model response (as guard_content). A request without these tags
# produces the plain-text payload with no qualifiers.
_GROUNDING_SOURCE_TEXT = "Tokyo is the capital of Japan."
_GROUNDING_QUERY_TEXT = "What is the capital of Japan?"
@ -2554,107 +2554,166 @@ def _model_response(content: str) -> ModelResponse:
)
def test_grounding_input_default_unchanged():
"""A plain-text request must produce the legacy INPUT payload (no qualifiers)."""
guardrail = _grounding_guardrail()
# Expected OUTPUT content blocks, keyed by their grounding qualifier, so the
# per-test assertions read as the block sequence they expect.
_GROUNDING_SOURCE_BLOCK = {
"text": {"text": _GROUNDING_SOURCE_TEXT, "qualifiers": ["grounding_source"]}
}
_QUERY_BLOCK = {"text": {"text": _GROUNDING_QUERY_TEXT, "qualifiers": ["query"]}}
_GUARD_BLOCK = {
"text": {"text": _GROUNDING_RESPONSE_TEXT, "qualifiers": ["guard_content"]}
}
request = guardrail.convert_to_bedrock_format(
source="INPUT", messages=[{"role": "user", "content": "hello"}]
def _input_request(messages: list) -> dict:
"""Arrange a guardrail and act: build the Bedrock INPUT payload."""
return _grounding_guardrail().convert_to_bedrock_format(
source="INPUT", messages=messages
)
assert request == {"source": "INPUT", "content": [{"text": {"text": "hello"}}]}
def test_grounding_input_propagates_qualifiers():
"""Tagged content blocks attach the matching Bedrock qualifiers, in order."""
guardrail = _grounding_guardrail()
request = guardrail.convert_to_bedrock_format(
source="INPUT", messages=_grounding_messages()
def _output_request(messages: list, response=None) -> dict:
"""Arrange a guardrail and act: build the Bedrock OUTPUT payload."""
return _grounding_guardrail().convert_to_bedrock_format(
source="OUTPUT", response=response, messages=messages
)
assert request["content"] == [
{"text": {"text": _GROUNDING_SOURCE_TEXT, "qualifiers": ["grounding_source"]}},
{"text": {"text": _GROUNDING_QUERY_TEXT, "qualifiers": ["query"]}},
]
def test_grounding_input_strips_grounding_and_query_qualifiers():
"""Grounding is OUTPUT-only: tagged source/query reach Bedrock as plain text on an
INPUT scan, so a tag cannot change how input-safety policies scan content (no bypass).
"""
expected_request = {
"source": "INPUT",
"content": [
{"text": {"text": _GROUNDING_SOURCE_TEXT}},
{"text": {"text": _GROUNDING_QUERY_TEXT}},
],
}
actual_request = _input_request(_grounding_messages())
assert actual_request == expected_request
def test_grounding_output_assembles_three_blocks():
"""OUTPUT assembles grounding_source + query (from request) + response (guard_content)."""
guardrail = _grounding_guardrail()
def test_grounding_input_leaves_existing_guarded_text_unqualified():
"""An existing guarded_text input block keeps its legacy unqualified payload."""
expected_request = {"source": "INPUT", "content": [{"text": {"text": "policy"}}]}
request = guardrail.convert_to_bedrock_format(
source="OUTPUT",
response=_model_response(_GROUNDING_RESPONSE_TEXT),
messages=_grounding_messages(),
actual_request = _input_request(
[{"role": "user", "content": [{"type": "guarded_text", "text": "policy"}]}]
)
assert request["source"] == "OUTPUT"
assert request["content"] == [
{"text": {"text": _GROUNDING_SOURCE_TEXT, "qualifiers": ["grounding_source"]}},
{"text": {"text": _GROUNDING_QUERY_TEXT, "qualifiers": ["query"]}},
{"text": {"text": _GROUNDING_RESPONSE_TEXT, "qualifiers": ["guard_content"]}},
]
assert actual_request == expected_request
def test_grounding_output_default_unchanged():
"""Without grounding tags, OUTPUT is the legacy single response block (no qualifiers)."""
guardrail = _grounding_guardrail()
def test_grounding_output_assembles_source_query_and_response():
"""OUTPUT emits grounding_source + query (from the request) then the response as
guard_content, so Bedrock can grade the response against the source and query."""
expected_request = {
"source": "OUTPUT",
"content": [_GROUNDING_SOURCE_BLOCK, _QUERY_BLOCK, _GUARD_BLOCK],
}
request = guardrail.convert_to_bedrock_format(
source="OUTPUT",
response=_model_response("Hi there."),
messages=[{"role": "user", "content": "hello"}],
actual_request = _output_request(
_grounding_messages(), _model_response(_GROUNDING_RESPONSE_TEXT)
)
assert request["content"] == [{"text": {"text": "Hi there."}}]
assert actual_request == expected_request
def test_grounding_multiple_sources_all_emitted():
"""Multiple grounding_source blocks are all emitted (Bedrock combines them)."""
guardrail = _grounding_guardrail()
def test_grounding_output_keeps_legacy_payload_without_tags():
"""Without grounding tags the OUTPUT payload is the legacy single response block."""
expected_request = {
"source": "OUTPUT",
"content": [{"text": {"text": "Hi there."}}],
}
actual_request = _output_request(
[{"role": "user", "content": "hello"}], _model_response("Hi there.")
)
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."
uk_source_block = {
"text": {"text": uk_source_text, "qualifiers": ["grounding_source"]}
}
messages = [
{
"role": "system",
"content": [
{"type": "grounding_source", "text": "London is the capital of UK."},
{"type": "grounding_source", "text": uk_source_text},
{"type": "grounding_source", "text": _GROUNDING_SOURCE_TEXT},
],
},
{"role": "user", "content": [{"type": "query", "text": _GROUNDING_QUERY_TEXT}]},
]
expected_request = {
"source": "OUTPUT",
"content": [
uk_source_block,
_GROUNDING_SOURCE_BLOCK,
_QUERY_BLOCK,
_GUARD_BLOCK,
],
}
request = guardrail.convert_to_bedrock_format(
source="OUTPUT",
response=_model_response(_GROUNDING_RESPONSE_TEXT),
messages=messages,
actual_request = _output_request(
messages, _model_response(_GROUNDING_RESPONSE_TEXT)
)
qualifiers = [item["text"].get("qualifiers") for item in request["content"]]
assert qualifiers == [
["grounding_source"],
["grounding_source"],
["query"],
["guard_content"],
]
assert actual_request == expected_request
def test_get_content_items_for_message_mixed():
"""The extractor preserves the grounding qualifier per block; untagged -> None."""
guardrail = _grounding_guardrail()
def test_grounding_output_keeps_grounding_for_non_model_response():
"""Harvested grounding blocks survive a non-ModelResponse output instead of being
silently dropped (regression guard for the unconditional content assignment)."""
expected_request = {
"source": "OUTPUT",
"content": [_GROUNDING_SOURCE_BLOCK, _QUERY_BLOCK],
}
blocks = guardrail.get_content_items_for_message(
actual_request = _output_request(_grounding_messages(), response=None)
assert actual_request == expected_request
@pytest.mark.parametrize(
"role, is_trusted",
[
("system", True),
("developer", True),
("tool", False),
("function", False),
("user", False),
("assistant", False),
],
)
def test_grounding_source_trusted_only_from_app_roles(role, is_trusted):
"""grounding_source is honored only from app-authored roles (system/developer). A
tag on a user, tool, function or assistant message is ignored, so neither a forwarded
end user nor an externally-influenced tool result can supply fake evidence for the
grounding check to grade the response against; query is always collected."""
messages = [
{
"role": "user",
"content": [
{"type": "grounding_source", "text": "src"},
{"type": "text", "text": "plain"},
{"type": "query", "text": "q"},
],
}
"role": role,
"content": [{"type": "grounding_source", "text": _GROUNDING_SOURCE_TEXT}],
},
{"role": "user", "content": [{"type": "query", "text": _GROUNDING_QUERY_TEXT}]},
]
expected_content = [_QUERY_BLOCK, _GUARD_BLOCK]
if is_trusted:
expected_content = [_GROUNDING_SOURCE_BLOCK, *expected_content]
actual_request = _output_request(
messages, _model_response(_GROUNDING_RESPONSE_TEXT)
)
assert blocks == [("src", "grounding_source"), ("plain", None), ("q", "query")]
assert actual_request == {"source": "OUTPUT", "content": expected_content}
@pytest.mark.asyncio