test: add failing tests for Bedrock contextual grounding (request-side)

Drive the request-side of Bedrock contextual grounding: callers tag message
content blocks as grounding_source/query, the post_call hook assembles an
ApplyGuardrail(OUTPUT) call carrying source + query + response(guard_content),
and the bedrock converse transform must render the tags as prompt text instead
of silently dropping them. Non-grounding payloads must stay byte-identical.
This commit is contained in:
João Costa 2026-06-09 22:16:09 +01:00
parent 496f5b9859
commit 7247c09725
2 changed files with 247 additions and 0 deletions

View file

@ -5386,3 +5386,45 @@ def test_converse_top_k_zero_forwarded_on_models_that_accept_it():
)
assert result["additionalModelRequestFields"]["top_k"] == 0
@pytest.mark.asyncio
async def test_grounding_source_and_query_rendered_as_text():
"""grounding_source / query content blocks must render as plain text on the
generate path (the model needs to see the RAG context + question). The bedrock
converse dispatch silently drops unrecognised content types, so these would
otherwise vanish from the prompt."""
from litellm.litellm_core_utils.prompt_templates.factory import (
BedrockConverseMessagesProcessor,
_bedrock_converse_messages_pt,
)
messages = [
{
"role": "user",
"content": [
{"type": "grounding_source", "text": "Tokyo is the capital of Japan."},
{"type": "query", "text": "What is the capital of Japan?"},
],
}
]
result = _bedrock_converse_messages_pt(
messages=messages,
model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
llm_provider="bedrock_converse",
)
async_result = (
await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async(
messages=messages,
model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
llm_provider="bedrock_converse",
)
)
assert result == async_result
assert len(result) == 1
assert result[0]["role"] == "user"
user_content = result[0]["content"]
assert {"text": "Tokyo is the capital of Japan."} in user_content
assert {"text": "What is the capital of Japan?"} in user_content

View file

@ -2503,3 +2503,208 @@ async def test_post_call_success_hook_only_runs_output_scan():
mock_make.call_args.kwargs.get("logging_event_type")
== GuardrailEventHooks.post_call
)
# ---------------------------------------------------------------------------
# Contextual grounding: request-side qualifiers
# ---------------------------------------------------------------------------
#
# Bedrock contextual grounding tags each ApplyGuardrail content block with a
# `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.
_GROUNDING_SOURCE_TEXT = "Tokyo is the capital of Japan."
_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_messages() -> list:
return [
{
"role": "system",
"content": [{"type": "grounding_source", "text": _GROUNDING_SOURCE_TEXT}],
},
{
"role": "user",
"content": [{"type": "query", "text": _GROUNDING_QUERY_TEXT}],
},
]
def _model_response(content: str) -> ModelResponse:
from litellm.types.utils import Choices, Message, ModelResponse
return ModelResponse(
choices=[
Choices(
index=0,
message=Message(role="assistant", content=content),
finish_reason="stop",
)
]
)
def test_grounding_input_default_unchanged():
"""A plain-text request must produce the legacy INPUT payload (no qualifiers)."""
guardrail = _grounding_guardrail()
request = guardrail.convert_to_bedrock_format(
source="INPUT", messages=[{"role": "user", "content": "hello"}]
)
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()
)
assert request["content"] == [
{"text": {"text": _GROUNDING_SOURCE_TEXT, "qualifiers": ["grounding_source"]}},
{"text": {"text": _GROUNDING_QUERY_TEXT, "qualifiers": ["query"]}},
]
def test_grounding_output_assembles_three_blocks():
"""OUTPUT assembles grounding_source + query (from request) + response (guard_content)."""
guardrail = _grounding_guardrail()
request = guardrail.convert_to_bedrock_format(
source="OUTPUT",
response=_model_response(_GROUNDING_RESPONSE_TEXT),
messages=_grounding_messages(),
)
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"]}},
]
def test_grounding_output_default_unchanged():
"""Without grounding tags, OUTPUT is the legacy single response block (no qualifiers)."""
guardrail = _grounding_guardrail()
request = guardrail.convert_to_bedrock_format(
source="OUTPUT",
response=_model_response("Hi there."),
messages=[{"role": "user", "content": "hello"}],
)
assert request["content"] == [{"text": {"text": "Hi there."}}]
def test_grounding_multiple_sources_all_emitted():
"""Multiple grounding_source blocks are all emitted (Bedrock combines them)."""
guardrail = _grounding_guardrail()
messages = [
{
"role": "system",
"content": [
{"type": "grounding_source", "text": "London is the capital of UK."},
{"type": "grounding_source", "text": _GROUNDING_SOURCE_TEXT},
],
},
{"role": "user", "content": [{"type": "query", "text": _GROUNDING_QUERY_TEXT}]},
]
request = guardrail.convert_to_bedrock_format(
source="OUTPUT",
response=_model_response(_GROUNDING_RESPONSE_TEXT),
messages=messages,
)
qualifiers = [item["text"].get("qualifiers") for item in request["content"]]
assert qualifiers == [
["grounding_source"],
["grounding_source"],
["query"],
["guard_content"],
]
def test_get_content_items_for_message_mixed():
"""The extractor preserves the grounding qualifier per block; untagged -> None."""
guardrail = _grounding_guardrail()
blocks = guardrail.get_content_items_for_message(
{
"role": "user",
"content": [
{"type": "grounding_source", "text": "src"},
{"type": "text", "text": "plain"},
{"type": "query", "text": "q"},
],
}
)
assert blocks == [("src", "grounding_source"), ("plain", None), ("q", "query")]
@pytest.mark.asyncio
async def test_grounding_output_blocked_raises_400():
"""A BLOCKED contextualGroundingPolicy filter raises HTTP 400."""
guardrail = _grounding_guardrail()
mock_bedrock_response = MagicMock()
mock_bedrock_response.status_code = 200
mock_bedrock_response.json.return_value = {
"action": "GUARDRAIL_INTERVENED",
"assessments": [
{
"contextualGroundingPolicy": {
"filters": [
{
"type": "GROUNDING",
"threshold": 0.7,
"score": 0.1,
"action": "BLOCKED",
}
]
}
}
],
"outputs": [{"text": "Response blocked: not grounded in the provided source."}],
}
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()),
):
mock_post.return_value = mock_bedrock_response
with pytest.raises(HTTPException) as exc_info:
await guardrail.make_bedrock_api_request(
source="OUTPUT",
response=_model_response("The capital of Japan is Paris."),
messages=_grounding_messages(),
request_data={"messages": _grounding_messages()},
)
assert exc_info.value.status_code == 400