mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-17 00:26:01 +00:00
fix(bedrock): scan tool definitions and additional request fields for passthrough guardrails
Converse passthrough guardrails only scanned system and message content, so a key holder could route blocked or PII text through toolConfig tool names, descriptions and input schemas or through additionalModelRequestFields, all of which are still forwarded to Bedrock. Collect strings from those fields too so key/team guardrails inspect and rewrite them, matching how the chat-completions path forwards tool definitions to guardrails.
This commit is contained in:
@@ -85,8 +85,11 @@ def _extract_converse_texts(
|
||||
top-level ``text`` blocks this scans the arbitrary-JSON fields a caller can
|
||||
hide prompt content in -- ``toolUse.input`` and
|
||||
``toolResult.content[].json`` (alongside ``toolResult.content[].text``) --
|
||||
so a blocking guardrail sees them before the request reaches Bedrock. Tool
|
||||
blocks are skipped entirely when tool messages are excluded.
|
||||
as well as the request-level fields still forwarded to Bedrock that a caller
|
||||
can route blocked content through: ``toolConfig.tools`` (tool names,
|
||||
descriptions and input schemas) and ``additionalModelRequestFields``. Tool
|
||||
message blocks are skipped when tool messages are excluded, but tool
|
||||
definitions are always scanned to match the chat-completions guardrail path.
|
||||
"""
|
||||
holders: List[_StringHolder] = []
|
||||
|
||||
@@ -114,6 +117,12 @@ def _extract_converse_texts(
|
||||
_collect_block_text(inner, holders)
|
||||
_collect_strings(inner.get("json"), holders)
|
||||
|
||||
tool_config = body.get("toolConfig")
|
||||
if isinstance(tool_config, dict):
|
||||
_collect_strings(tool_config.get("tools"), holders)
|
||||
|
||||
_collect_strings(body.get("additionalModelRequestFields"), holders)
|
||||
|
||||
texts = [container[key] for container, key in holders]
|
||||
return texts, holders
|
||||
|
||||
|
||||
@@ -170,6 +170,57 @@ class TestExtractConverseTexts:
|
||||
texts, _ = _extract_converse_texts(body, skip_system=False, skip_tool=False)
|
||||
assert texts == []
|
||||
|
||||
def test_extracts_tool_config_description_and_schema(self):
|
||||
body = {
|
||||
"messages": [{"role": "user", "content": [{"text": "hi"}]}],
|
||||
"toolConfig": {
|
||||
"tools": [
|
||||
{
|
||||
"toolSpec": {
|
||||
"name": "lookup",
|
||||
"description": "blocked tool description",
|
||||
"inputSchema": {
|
||||
"json": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"q": {
|
||||
"type": "string",
|
||||
"description": "blocked schema description",
|
||||
}
|
||||
},
|
||||
}
|
||||
},
|
||||
}
|
||||
}
|
||||
]
|
||||
},
|
||||
}
|
||||
texts, _ = _extract_converse_texts(body, skip_system=False, skip_tool=False)
|
||||
assert "blocked tool description" in texts
|
||||
assert "blocked schema description" in texts
|
||||
|
||||
def test_tool_config_scanned_even_when_tool_messages_skipped(self):
|
||||
body = {
|
||||
"messages": [{"role": "user", "content": [{"text": "hi"}]}],
|
||||
"toolConfig": {
|
||||
"tools": [
|
||||
{"toolSpec": {"name": "fn", "description": "blocked description"}}
|
||||
]
|
||||
},
|
||||
}
|
||||
texts, _ = _extract_converse_texts(body, skip_system=False, skip_tool=True)
|
||||
assert "blocked description" in texts
|
||||
|
||||
def test_extracts_additional_model_request_fields(self):
|
||||
body = {
|
||||
"messages": [{"role": "user", "content": [{"text": "hi"}]}],
|
||||
"additionalModelRequestFields": {
|
||||
"reasoning_config": {"prompt": "blocked extra field"}
|
||||
},
|
||||
}
|
||||
texts, _ = _extract_converse_texts(body, skip_system=False, skip_tool=False)
|
||||
assert "blocked extra field" in texts
|
||||
|
||||
|
||||
class TestWriteBackTexts:
|
||||
def test_writes_system_text(self):
|
||||
@@ -364,6 +415,71 @@ class TestBedrockPassthroughGuardrailHandlerInput:
|
||||
sent_texts = guardrail.apply_guardrail.call_args.kwargs["inputs"]["texts"]
|
||||
assert "blocked content" in sent_texts
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_config_description_scanned_and_masked(self):
|
||||
"""Blocked text hidden in toolConfig.tools[].toolSpec.description is still
|
||||
forwarded to Bedrock, so the guardrail must see it and mask it in place."""
|
||||
handler = BedrockPassthroughGuardrailHandler()
|
||||
data = _converse_data()
|
||||
data["data"]["toolConfig"] = {
|
||||
"tools": [
|
||||
{
|
||||
"toolSpec": {
|
||||
"name": "lookup",
|
||||
"description": "email john@example.com",
|
||||
"inputSchema": {"json": {"type": "object"}},
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
guardrail = _make_guardrail(
|
||||
{"texts": ["You are helpful.", "Hello world", "lookup", "[REDACTED]", "object"]}
|
||||
)
|
||||
|
||||
result = await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
sent_texts = guardrail.apply_guardrail.call_args.kwargs["inputs"]["texts"]
|
||||
assert "email john@example.com" in sent_texts
|
||||
tool_spec = result["data"]["toolConfig"]["tools"][0]["toolSpec"]
|
||||
assert tool_spec["description"] == "[REDACTED]"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_config_description_blocking_propagates(self):
|
||||
"""A blocking guardrail must reject content hidden in a tool description."""
|
||||
handler = BedrockPassthroughGuardrailHandler()
|
||||
data = _converse_data()
|
||||
data["data"]["toolConfig"] = {
|
||||
"tools": [{"toolSpec": {"name": "fn", "description": "blocked content"}}]
|
||||
}
|
||||
guardrail = MagicMock()
|
||||
guardrail.guardrail_name = "block-guard"
|
||||
guardrail.skip_system_message_in_guardrail = False
|
||||
guardrail.skip_tool_message_in_guardrail = False
|
||||
guardrail.apply_guardrail = AsyncMock(side_effect=GuardrailBlocked("Blocked"))
|
||||
|
||||
with pytest.raises(GuardrailBlocked):
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
sent_texts = guardrail.apply_guardrail.call_args.kwargs["inputs"]["texts"]
|
||||
assert "blocked content" in sent_texts
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_additional_model_request_fields_scanned_and_masked(self):
|
||||
"""Blocked text hidden in additionalModelRequestFields is forwarded to
|
||||
Bedrock, so the guardrail must scan it and mask it in place."""
|
||||
handler = BedrockPassthroughGuardrailHandler()
|
||||
data = _converse_data()
|
||||
data["data"]["additionalModelRequestFields"] = {"note": "ssn 123-45-6789"}
|
||||
guardrail = _make_guardrail(
|
||||
{"texts": ["You are helpful.", "Hello world", "[REDACTED]"]}
|
||||
)
|
||||
|
||||
result = await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
sent_texts = guardrail.apply_guardrail.call_args.kwargs["inputs"]["texts"]
|
||||
assert "ssn 123-45-6789" in sent_texts
|
||||
assert result["data"]["additionalModelRequestFields"]["note"] == "[REDACTED]"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_converse_endpoint_scans_full_payload(self):
|
||||
"""Invoke routes must not bypass guardrails: the full request payload is
|
||||
|
||||
Reference in New Issue
Block a user