mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
7247c09725
commit
0cbd7f4a14
6 changed files with 293 additions and 112 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue