mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
Merge pull request #41132 from BerriAI/litellm_bedrock_grounding_from_plain_messages
fix(bedrock guardrails): derive contextual grounding source and query from plain messages
This commit is contained in:
commit
8481bc27f9
7 changed files with 290 additions and 7 deletions
|
|
@ -11503,6 +11503,12 @@
|
|||
"description": "Enable content moderation to check for harmful content (harassment, hate speech, etc.).",
|
||||
"title": "Content Moderation Check"
|
||||
},
|
||||
"contextual_grounding_from_messages": {
|
||||
"default": false,
|
||||
"description": "ApplyGuardrail: when True, post-call scans of a request with no grounding_source / query content parts send the system and developer messages as the grounding source and the latest user message as the query, so the guardrail's contextual grounding policy can score the response. Bedrock bills contextual grounding units for these scans and rejects queries, sources and responses over its contextual grounding length limits, so leave this off for guardrails without a contextual grounding policy. Default False: plain messages are never sent as grounding context.",
|
||||
"title": "Contextual Grounding From Messages",
|
||||
"type": "boolean"
|
||||
},
|
||||
"credentials": {
|
||||
"anyOf": [
|
||||
{
|
||||
|
|
|
|||
|
|
@ -244,6 +244,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
prompt_attack_threshold: float | None = 0.5,
|
||||
pii_confidence_threshold: float | None = 0.5,
|
||||
chunk_budget_chars: int = BEDROCK_APPLY_GUARDRAIL_CHUNK_BUDGET_CHARS,
|
||||
contextual_grounding_from_messages: bool = False,
|
||||
streaming_buffer_until_moderated: bool | None = None,
|
||||
streaming_sampling_rate: int | None = None,
|
||||
streaming_end_of_stream_only: bool | None = None,
|
||||
|
|
@ -265,6 +266,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
self.guardrailVersion = guardrailVersion
|
||||
self.guardrail_provider = "bedrock"
|
||||
self.chunk_budget_chars = chunk_budget_chars
|
||||
self.contextual_grounding_from_messages = contextual_grounding_from_messages
|
||||
self.experimental_use_latest_role_message_only = bool(kwargs.get("experimental_use_latest_role_message_only"))
|
||||
|
||||
# Resource-less, detect-only InvokeGuardrailChecks mode. Present `checks`
|
||||
|
|
@ -459,8 +461,8 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
"""
|
||||
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.
|
||||
Untagged text keeps ``qualifier=None``; the OUTPUT scan decides whether to
|
||||
derive grounding qualifiers from it.
|
||||
"""
|
||||
content: Final = message.get("content")
|
||||
if content is None:
|
||||
|
|
@ -493,6 +495,10 @@ 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).
|
||||
|
||||
With ``contextual_grounding_from_messages`` on, 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 +510,33 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
and role in _GROUNDING_SOURCE_TRUSTED_ROLES
|
||||
):
|
||||
grounding.append(block)
|
||||
return grounding
|
||||
if grounding or not self.contextual_grounding_from_messages:
|
||||
return grounding
|
||||
return self._derive_grounding_blocks_from_plain_messages(messages)
|
||||
|
||||
def _derive_grounding_blocks_from_plain_messages(
|
||||
self, messages: list[AllMessageValues] | None
|
||||
) -> list[QualifiedTextBlock]:
|
||||
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 +3242,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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ def initialize_bedrock(litellm_params: LitellmParams, guardrail: Guardrail):
|
|||
prompt_attack_threshold=litellm_params.prompt_attack_threshold,
|
||||
pii_confidence_threshold=litellm_params.pii_confidence_threshold,
|
||||
chunk_budget_chars=litellm_params.chunk_budget_chars,
|
||||
contextual_grounding_from_messages=litellm_params.contextual_grounding_from_messages,
|
||||
default_on=litellm_params.default_on,
|
||||
disable_exception_on_block=litellm_params.disable_exception_on_block,
|
||||
mask_request_content=litellm_params.mask_request_content,
|
||||
|
|
|
|||
|
|
@ -552,6 +552,16 @@ class BedrockGuardrailConfigModel(BaseModel):
|
|||
"still rejects is bisected automatically, so this value only trades round trips against "
|
||||
"batch size and cannot fail a request on its own.",
|
||||
)
|
||||
contextual_grounding_from_messages: bool = Field(
|
||||
default=False,
|
||||
description="ApplyGuardrail: when True, post-call scans of a request with no grounding_source / "
|
||||
"query content parts send the system and developer messages as the grounding source and "
|
||||
"the latest user message as the query, so the guardrail's contextual grounding policy can "
|
||||
"score the response. Bedrock bills contextual grounding units for these scans and rejects "
|
||||
"queries, sources and responses over its contextual grounding length limits, so leave this "
|
||||
"off for guardrails without a contextual grounding policy. Default False: plain messages "
|
||||
"are never sent as grounding context.",
|
||||
)
|
||||
|
||||
|
||||
class BedrockGuardrailStreamingParams(BaseModel):
|
||||
|
|
|
|||
|
|
@ -2375,8 +2375,12 @@ _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_guardrail(from_messages: bool = False) -> BedrockGuardrail:
|
||||
return BedrockGuardrail(
|
||||
guardrailIdentifier="test-guardrail",
|
||||
guardrailVersion="DRAFT",
|
||||
contextual_grounding_from_messages=from_messages,
|
||||
)
|
||||
|
||||
|
||||
def _grounding_messages() -> list:
|
||||
|
|
@ -2418,9 +2422,11 @@ def _input_request(messages: list) -> dict:
|
|||
return _grounding_guardrail().convert_to_bedrock_format(source="INPUT", messages=messages)
|
||||
|
||||
|
||||
def _output_request(messages: list, response=None) -> dict:
|
||||
def _output_request(messages: list, response=None, from_messages: bool = False) -> dict:
|
||||
"""Arrange a guardrail and act: build the Bedrock OUTPUT payload."""
|
||||
return _grounding_guardrail().convert_to_bedrock_format(source="OUTPUT", response=response, messages=messages)
|
||||
return _grounding_guardrail(from_messages).convert_to_bedrock_format(
|
||||
source="OUTPUT", response=response, messages=messages
|
||||
)
|
||||
|
||||
|
||||
def test_grounding_input_strips_grounding_and_query_qualifiers():
|
||||
|
|
@ -2474,6 +2480,131 @@ 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():
|
||||
"""Flag on: untagged system + user text is sent as grounding_source + query."""
|
||||
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), from_messages=True)
|
||||
|
||||
assert actual_request == expected_request
|
||||
|
||||
|
||||
def test_grounding_output_plain_messages_stay_legacy_when_flag_is_off():
|
||||
"""Default config: plain system + user text is never sent as grounding context."""
|
||||
messages = [
|
||||
{"role": "system", "content": _GROUNDING_SOURCE_TEXT},
|
||||
{"role": "user", "content": _GROUNDING_QUERY_TEXT},
|
||||
]
|
||||
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_derived_query_is_latest_user_turn_only():
|
||||
"""Only the latest user turn is the query; system and developer turns are the source."""
|
||||
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), from_messages=True)
|
||||
|
||||
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 source without a query and vice versa, so send neither."""
|
||||
expected_request = {"source": "OUTPUT", "content": [{"text": {"text": _GROUNDING_RESPONSE_TEXT}}]}
|
||||
|
||||
actual_request = _output_request(messages, _model_response(_GROUNDING_RESPONSE_TEXT), from_messages=True)
|
||||
|
||||
assert actual_request == expected_request
|
||||
|
||||
|
||||
def test_grounding_output_explicit_tags_take_precedence_over_plain_messages():
|
||||
"""Tagged blocks win: untagged text around them is not added as source or 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), from_messages=True)
|
||||
|
||||
assert actual_request == expected_request
|
||||
|
||||
|
||||
def test_grounding_input_ignores_plain_message_derivation():
|
||||
"""INPUT scans never derive grounding qualifiers from plain messages."""
|
||||
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 = _grounding_guardrail(from_messages=True).convert_to_bedrock_format(
|
||||
source="INPUT", messages=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 +2728,56 @@ async def test_grounding_output_blocked_raises_400():
|
|||
assert exc_info.value.status_code == 400
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"from_messages, request_messages",
|
||||
[
|
||||
(
|
||||
True,
|
||||
[
|
||||
{"role": "system", "content": _GROUNDING_SOURCE_TEXT},
|
||||
{"role": "user", "content": _GROUNDING_QUERY_TEXT},
|
||||
],
|
||||
),
|
||||
(
|
||||
False,
|
||||
[
|
||||
{"role": "system", "content": [{"type": "grounding_source", "text": _GROUNDING_SOURCE_TEXT}]},
|
||||
{"role": "user", "content": [{"type": "query", "text": _GROUNDING_QUERY_TEXT}]},
|
||||
],
|
||||
),
|
||||
],
|
||||
ids=["plain-messages-flag-on", "tagged-messages-flag-off"],
|
||||
)
|
||||
async def test_apply_guardrail_response_forwards_request_messages_for_grounding(from_messages, request_messages):
|
||||
guardrail = _grounding_guardrail(from_messages=from_messages)
|
||||
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
|
||||
#
|
||||
|
|
|
|||
|
|
@ -71,6 +71,52 @@ def test_initialize_bedrock_forwards_chunk_budget_chars():
|
|||
assert initialized[-1].chunk_budget_chars == 60_000
|
||||
|
||||
|
||||
def test_initialize_bedrock_forwards_contextual_grounding_from_messages():
|
||||
"""`contextual_grounding_from_messages: true` in config.yaml must make the post-call
|
||||
payload carry the plain system prompt and user turn as grounding_source and query."""
|
||||
import litellm
|
||||
from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import BedrockGuardrail
|
||||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
|
||||
test_guardrail = {
|
||||
"guardrail_name": "test_bedrock_grounding_from_messages",
|
||||
"litellm_params": {
|
||||
"guardrail": SupportedGuardrailIntegrations.BEDROCK.value,
|
||||
"mode": "post_call",
|
||||
"guardrailIdentifier": "test-guardrail",
|
||||
"guardrailVersion": "DRAFT",
|
||||
"contextual_grounding_from_messages": True,
|
||||
},
|
||||
}
|
||||
messages = [
|
||||
{"role": "system", "content": "Returns are accepted for 30 days."},
|
||||
{"role": "user", "content": "How long is the return window?"},
|
||||
]
|
||||
response = ModelResponse(
|
||||
choices=[Choices(index=0, message=Message(role="assistant", content="30 days."), finish_reason="stop")]
|
||||
)
|
||||
expected_request = {
|
||||
"source": "OUTPUT",
|
||||
"content": [
|
||||
{"text": {"text": "Returns are accepted for 30 days.", "qualifiers": ["grounding_source"]}},
|
||||
{"text": {"text": "How long is the return window?", "qualifiers": ["query"]}},
|
||||
{"text": {"text": "30 days.", "qualifiers": ["guard_content"]}},
|
||||
],
|
||||
}
|
||||
|
||||
guardrail_handler = InMemoryGuardrailHandler()
|
||||
guardrail_handler.initialize_guardrail(guardrail=test_guardrail)
|
||||
|
||||
initialized = [
|
||||
callback
|
||||
for callback in litellm.callbacks
|
||||
if isinstance(callback, BedrockGuardrail) and callback.guardrail_name == "test_bedrock_grounding_from_messages"
|
||||
]
|
||||
assert initialized, "bedrock guardrail was not registered as a callback"
|
||||
actual_request = initialized[-1].convert_to_bedrock_format(source="OUTPUT", response=response, messages=messages)
|
||||
assert json.loads(json.dumps(actual_request)) == expected_request
|
||||
|
||||
|
||||
def test_initialize_guardrail_preserves_guardrail_info():
|
||||
"""
|
||||
Regression (LIT-2529): initialize_guardrail must carry guardrail_info into the
|
||||
|
|
|
|||
6
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
6
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -31006,6 +31006,12 @@ export interface components {
|
|||
* @description Enable content moderation to check for harmful content (harassment, hate speech, etc.).
|
||||
*/
|
||||
content_moderation_check?: boolean | null;
|
||||
/**
|
||||
* Contextual Grounding From Messages
|
||||
* @description ApplyGuardrail: when True, post-call scans of a request with no grounding_source / query content parts send the system and developer messages as the grounding source and the latest user message as the query, so the guardrail's contextual grounding policy can score the response. Bedrock bills contextual grounding units for these scans and rejects queries, sources and responses over its contextual grounding length limits, so leave this off for guardrails without a contextual grounding policy. Default False: plain messages are never sent as grounding context.
|
||||
* @default false
|
||||
*/
|
||||
contextual_grounding_from_messages: boolean;
|
||||
/**
|
||||
* Credentials
|
||||
* @description Path to Google Cloud credentials JSON file or JSON string
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue