mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
496f5b9859
commit
7247c09725
2 changed files with 247 additions and 0 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue