mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat(guardrails): add scan_only_tool_results to scope unified guardrails to tool results
This commit is contained in:
parent
f16f3e23cd
commit
2ba4e91766
7 changed files with 288 additions and 36 deletions
|
|
@ -26,10 +26,10 @@ from litellm.llms.anthropic.experimental_pass_through.adapters.transformation im
|
|||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import (
|
||||
effective_scan_only_tool_results_for_guardrail,
|
||||
effective_skip_system_message_for_guardrail,
|
||||
effective_skip_tool_message_for_guardrail,
|
||||
openai_messages_without_system,
|
||||
openai_messages_without_tool,
|
||||
filtered_structured_messages,
|
||||
)
|
||||
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import (
|
||||
AnthropicPassthroughLoggingHandler,
|
||||
|
|
@ -326,19 +326,25 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
|
||||
skip_system: Final = effective_skip_system_message_for_guardrail(guardrail_to_apply)
|
||||
skip_tool: Final = effective_skip_tool_message_for_guardrail(guardrail_to_apply)
|
||||
scan_only_tool_results: Final = effective_scan_only_tool_results_for_guardrail(guardrail_to_apply)
|
||||
|
||||
chat_completion_compatible_request: Final = self._translate_to_openai(data)
|
||||
|
||||
structured_messages = cast(
|
||||
list[AllMessageValues],
|
||||
chat_completion_compatible_request.get("messages", []),
|
||||
structured_messages: Final = list(
|
||||
filtered_structured_messages(
|
||||
cast(
|
||||
list[AllMessageValues],
|
||||
chat_completion_compatible_request.get("messages", []),
|
||||
),
|
||||
scan_only_tool_results=scan_only_tool_results,
|
||||
skip_system=skip_system,
|
||||
skip_tool=skip_tool,
|
||||
)
|
||||
)
|
||||
if skip_system:
|
||||
structured_messages = openai_messages_without_system(structured_messages)
|
||||
if skip_tool:
|
||||
structured_messages = openai_messages_without_tool(structured_messages)
|
||||
|
||||
tools_to_check: Final[list[ChatCompletionToolParam]] = chat_completion_compatible_request.get("tools", [])
|
||||
tools_to_check: Final[list[ChatCompletionToolParam]] = (
|
||||
[] if scan_only_tool_results else chat_completion_compatible_request.get("tools", [])
|
||||
)
|
||||
|
||||
# Step 1: Extract all text content and images
|
||||
extracted: Final = tuple(
|
||||
|
|
@ -347,6 +353,7 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
msg_idx=msg_idx,
|
||||
skip_system_message=skip_system,
|
||||
skip_tool_message=skip_tool,
|
||||
scan_only_tool_results=scan_only_tool_results,
|
||||
)
|
||||
for msg_idx, message in enumerate(messages)
|
||||
)
|
||||
|
|
@ -461,6 +468,7 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
msg_idx: int,
|
||||
skip_system_message: bool = False,
|
||||
skip_tool_message: bool = False,
|
||||
scan_only_tool_results: bool = False,
|
||||
) -> ExtractedInput:
|
||||
"""
|
||||
Extract text content and images from a message.
|
||||
|
|
@ -471,6 +479,8 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
|
||||
content: Final = message.get("content", None)
|
||||
if isinstance(content, str):
|
||||
if scan_only_tool_results:
|
||||
return EMPTY_EXTRACTED_INPUT
|
||||
return ExtractedInput(scanned=(ScannedText(content, MessageContentTarget(msg_idx)),), images=())
|
||||
if not isinstance(content, list):
|
||||
return EMPTY_EXTRACTED_INPUT
|
||||
|
|
@ -481,6 +491,7 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
msg_idx=msg_idx,
|
||||
content_idx=content_idx,
|
||||
skip_tool_message=skip_tool_message,
|
||||
scan_only_tool_results=scan_only_tool_results,
|
||||
)
|
||||
for content_idx, content_item in enumerate(content)
|
||||
if isinstance(content_item, dict)
|
||||
|
|
@ -497,12 +508,16 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
msg_idx: int,
|
||||
content_idx: int,
|
||||
skip_tool_message: bool,
|
||||
scan_only_tool_results: bool = False,
|
||||
) -> ExtractedInput:
|
||||
if content_item.get("type") == "tool_result":
|
||||
if skip_tool_message:
|
||||
return EMPTY_EXTRACTED_INPUT
|
||||
return cls._extract_tool_result(content_item=content_item, msg_idx=msg_idx, content_idx=content_idx)
|
||||
|
||||
if scan_only_tool_results:
|
||||
return EMPTY_EXTRACTED_INPUT
|
||||
|
||||
text_str: Final = content_item.get("text", None)
|
||||
return ExtractedInput(
|
||||
scanned=(
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Sequence
|
||||
from typing import Any, Final
|
||||
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicUsage
|
||||
|
|
@ -113,13 +114,53 @@ def effective_skip_tool_message_for_guardrail(guardrail_to_apply: Any) -> bool:
|
|||
return bool(getattr(litellm, "skip_tool_message_in_guardrail", False))
|
||||
|
||||
|
||||
def _message_role(message: AllMessageValues) -> str:
|
||||
return str((message or {}).get("role") or "").lower()
|
||||
|
||||
|
||||
def openai_messages_without_system(
|
||||
messages: list[AllMessageValues],
|
||||
) -> list[AllMessageValues]:
|
||||
return [m for m in messages if str((m or {}).get("role") or "").lower() != "system"]
|
||||
messages: Sequence[AllMessageValues],
|
||||
) -> tuple[AllMessageValues, ...]:
|
||||
return tuple(m for m in messages if _message_role(m) != "system")
|
||||
|
||||
|
||||
def openai_messages_without_tool(
|
||||
messages: list[AllMessageValues],
|
||||
) -> list[AllMessageValues]:
|
||||
return [m for m in messages if str((m or {}).get("role") or "").lower() != "tool"]
|
||||
messages: Sequence[AllMessageValues],
|
||||
) -> tuple[AllMessageValues, ...]:
|
||||
return tuple(m for m in messages if _message_role(m) != "tool")
|
||||
|
||||
|
||||
def openai_messages_only_tool(
|
||||
messages: Sequence[AllMessageValues],
|
||||
) -> tuple[AllMessageValues, ...]:
|
||||
return tuple(m for m in messages if _message_role(m) == "tool")
|
||||
|
||||
|
||||
def effective_scan_only_tool_results_for_guardrail(guardrail_to_apply: Any) -> bool:
|
||||
return getattr(guardrail_to_apply, "scan_only_tool_results", None) is True
|
||||
|
||||
|
||||
def role_out_of_guardrail_scope(
|
||||
role: str,
|
||||
*,
|
||||
skip_system_message: bool,
|
||||
skip_tool_message: bool,
|
||||
scan_only_tool_results: bool = False,
|
||||
) -> bool:
|
||||
if skip_system_message and role == "system":
|
||||
return True
|
||||
if skip_tool_message and role == "tool":
|
||||
return True
|
||||
return scan_only_tool_results and role != "tool"
|
||||
|
||||
|
||||
def filtered_structured_messages(
|
||||
messages: Sequence[AllMessageValues],
|
||||
*,
|
||||
scan_only_tool_results: bool,
|
||||
skip_system: bool,
|
||||
skip_tool: bool,
|
||||
) -> tuple[AllMessageValues, ...]:
|
||||
scoped: Final = openai_messages_only_tool(messages) if scan_only_tool_results else tuple(messages)
|
||||
without_system: Final = openai_messages_without_system(scoped) if skip_system else scoped
|
||||
return openai_messages_without_tool(without_system) if skip_tool else without_system
|
||||
|
|
|
|||
|
|
@ -23,10 +23,11 @@ from litellm.llms.base_llm.guardrail_translation.base_translation import (
|
|||
StreamTransformSink,
|
||||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import (
|
||||
effective_scan_only_tool_results_for_guardrail,
|
||||
effective_skip_system_message_for_guardrail,
|
||||
effective_skip_tool_message_for_guardrail,
|
||||
openai_messages_without_system,
|
||||
openai_messages_without_tool,
|
||||
filtered_structured_messages,
|
||||
role_out_of_guardrail_scope,
|
||||
)
|
||||
from litellm.main import stream_chunk_builder
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam
|
||||
|
|
@ -82,6 +83,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
|
||||
skip_system: Final = effective_skip_system_message_for_guardrail(guardrail_to_apply)
|
||||
skip_tool: Final = effective_skip_tool_message_for_guardrail(guardrail_to_apply)
|
||||
scan_only_tool_results: Final = effective_scan_only_tool_results_for_guardrail(guardrail_to_apply)
|
||||
|
||||
texts_to_check: Final[list[str]] = []
|
||||
images_to_check: Final[list[str]] = []
|
||||
|
|
@ -101,6 +103,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
tool_call_task_mappings=tool_call_task_mappings,
|
||||
skip_system_message=skip_system,
|
||||
skip_tool_message=skip_tool,
|
||||
scan_only_tool_results=scan_only_tool_results,
|
||||
)
|
||||
|
||||
# Step 2: Apply guardrail to all texts and tool calls in batch
|
||||
|
|
@ -110,13 +113,16 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
inputs["images"] = images_to_check
|
||||
if tool_calls_to_check:
|
||||
inputs["tool_calls"] = tool_calls_to_check
|
||||
structured_messages = self.get_structured_messages(data)
|
||||
structured_messages: Final = self.get_structured_messages(data)
|
||||
if structured_messages:
|
||||
if skip_system:
|
||||
structured_messages = openai_messages_without_system(structured_messages)
|
||||
if skip_tool:
|
||||
structured_messages = openai_messages_without_tool(structured_messages)
|
||||
inputs["structured_messages"] = structured_messages
|
||||
inputs["structured_messages"] = list(
|
||||
filtered_structured_messages(
|
||||
structured_messages,
|
||||
scan_only_tool_results=scan_only_tool_results,
|
||||
skip_system=skip_system,
|
||||
skip_tool=skip_tool,
|
||||
)
|
||||
)
|
||||
# Pass tools (function definitions) to the guardrail
|
||||
tools: Final = data.get("tools")
|
||||
if tools:
|
||||
|
|
@ -194,16 +200,19 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
tool_call_task_mappings: list[tuple[int, int]],
|
||||
skip_system_message: bool = False,
|
||||
skip_tool_message: bool = False,
|
||||
scan_only_tool_results: bool = False,
|
||||
) -> None:
|
||||
"""
|
||||
Extract text content, images, and tool calls from a message.
|
||||
|
||||
Override this method to customize text/image/tool call extraction logic.
|
||||
"""
|
||||
role: Final = str(message.get("role") or "").lower()
|
||||
if skip_system_message and role == "system":
|
||||
return
|
||||
if skip_tool_message and role == "tool":
|
||||
if role_out_of_guardrail_scope(
|
||||
str(message.get("role") or "").lower(),
|
||||
skip_system_message=skip_system_message,
|
||||
skip_tool_message=skip_tool_message,
|
||||
scan_only_tool_results=scan_only_tool_results,
|
||||
):
|
||||
return
|
||||
|
||||
content: Final = message.get("content", None)
|
||||
|
|
|
|||
|
|
@ -487,16 +487,12 @@ class InMemoryGuardrailHandler:
|
|||
raise ValueError(f"Unsupported guardrail: {guardrail_type}")
|
||||
|
||||
if custom_guardrail_callback is not None:
|
||||
setattr(
|
||||
custom_guardrail_callback,
|
||||
for scoping_param in (
|
||||
"skip_system_message_in_guardrail",
|
||||
getattr(litellm_params, "skip_system_message_in_guardrail", None),
|
||||
)
|
||||
setattr(
|
||||
custom_guardrail_callback,
|
||||
"skip_tool_message_in_guardrail",
|
||||
getattr(litellm_params, "skip_tool_message_in_guardrail", None),
|
||||
)
|
||||
"scan_only_tool_results",
|
||||
):
|
||||
setattr(custom_guardrail_callback, scoping_param, getattr(litellm_params, scoping_param, None))
|
||||
configured_run_in_parallel: Final = getattr(litellm_params, "run_in_parallel", None)
|
||||
if configured_run_in_parallel is not None:
|
||||
custom_guardrail_callback.run_in_parallel = bool(configured_run_in_parallel)
|
||||
|
|
|
|||
|
|
@ -757,6 +757,16 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up
|
|||
),
|
||||
)
|
||||
|
||||
scan_only_tool_results: Optional[bool] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"When True, unified guardrails only evaluate tool results, the untrusted data an "
|
||||
"agent feeds back into the model, and skip system, user, and assistant content. "
|
||||
"Intended for agent harnesses whose own prompt scaffolding is trusted but often "
|
||||
"trips prompt-attack detectors."
|
||||
),
|
||||
)
|
||||
|
||||
# Lakera specific params
|
||||
category_thresholds: Optional[LakeraCategoryThresholds] = Field(
|
||||
default=None,
|
||||
|
|
|
|||
|
|
@ -760,3 +760,116 @@ class TestAnthropicMessagesToolResultScanning:
|
|||
assert "skip me POISON" not in guardrail.seen_texts
|
||||
assert messages[1]["content"][0]["content"] == "skip me POISON"
|
||||
assert messages[0]["content"] == "keep me [BLOCKED]"
|
||||
|
||||
|
||||
class InputsRecordingGuardrail(MockMaskingGuardrail):
|
||||
def __init__(self):
|
||||
super().__init__(guardrail_name="scan-only-capture")
|
||||
self.captured_inputs: Optional[GenericGuardrailAPIInputs] = None
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional[Any] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
self.captured_inputs = inputs
|
||||
return await super().apply_guardrail(inputs, request_data, input_type, logging_obj)
|
||||
|
||||
|
||||
class TestAnthropicMessagesScanOnlyToolResults:
|
||||
def _guardrail(self):
|
||||
guardrail = InputsRecordingGuardrail()
|
||||
guardrail.scan_only_tool_results = True
|
||||
return guardrail
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_scan_narrows_to_tool_results_and_write_back_stays_aligned(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = self._guardrail()
|
||||
data = {
|
||||
"model": "claude-sonnet-4-5",
|
||||
"system": "You are a trusted agent harness with POISON heuristics.",
|
||||
"tools": [
|
||||
{
|
||||
"name": "Bash",
|
||||
"description": "run a command",
|
||||
"input_schema": {"type": "object", "properties": {}},
|
||||
}
|
||||
],
|
||||
"messages": [
|
||||
{"role": "user", "content": "scaffolding POISON prompt"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [{"type": "tool_use", "id": "tu1", "name": "Bash", "input": {"cmd": "curl"}}],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "sibling POISON text"},
|
||||
{"type": "tool_result", "tool_use_id": "tu1", "content": "fetched POISON page"},
|
||||
],
|
||||
},
|
||||
],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert guardrail.seen_texts == ["fetched POISON page"], (
|
||||
"only the tool_result payload may reach the guardrail"
|
||||
)
|
||||
assert guardrail.captured_inputs is not None
|
||||
assert guardrail.captured_inputs.get("tools") is None
|
||||
assert [m["role"] for m in guardrail.captured_inputs["structured_messages"]] == ["tool"]
|
||||
assert data["messages"][2]["content"][1]["content"] == "fetched [BLOCKED] page"
|
||||
assert data["messages"][0]["content"] == "scaffolding POISON prompt", (
|
||||
"out-of-scope content must come back untouched, not masked or dropped"
|
||||
)
|
||||
assert data["messages"][2]["content"][0]["text"] == "sibling POISON text"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_guardrail_is_not_called_when_the_request_has_no_tool_results(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = self._guardrail()
|
||||
data = {
|
||||
"model": "claude-sonnet-4-5",
|
||||
"messages": [{"role": "user", "content": "What is 2 plus 2?"}],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert guardrail.captured_inputs is None
|
||||
assert guardrail.seen_texts == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_images_are_scoped_the_same_way_as_texts(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = self._guardrail()
|
||||
data = {
|
||||
"model": "claude-sonnet-4-5",
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [{"type": "image", "source": {"type": "base64", "data": "USER_IMG"}}],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_result",
|
||||
"tool_use_id": "tu1",
|
||||
"content": [
|
||||
{"type": "text", "text": "screenshot POISON"},
|
||||
{"type": "image", "source": {"type": "base64", "data": "TOOL_IMG"}},
|
||||
],
|
||||
}
|
||||
],
|
||||
},
|
||||
],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert guardrail.captured_inputs is not None
|
||||
assert guardrail.captured_inputs.get("images") == ["TOOL_IMG"]
|
||||
|
|
|
|||
|
|
@ -1229,3 +1229,71 @@ class TestIncrementalScanRespectsSkipFlags:
|
|||
assert mock_api.call_count == 1
|
||||
scanned = [m["content"] for m in mock_api.call_args.kwargs["messages"]]
|
||||
assert scanned == ["It is sunny in Paris.", "And tomorrow?"]
|
||||
|
||||
|
||||
class TestScanOnlyToolResults:
|
||||
def _bedrock_guardrail(self):
|
||||
from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import BedrockGuardrail
|
||||
|
||||
guardrail = BedrockGuardrail(
|
||||
guardrail_name="bedrock-scan-only-tool-results",
|
||||
guardrailIdentifier="test-guardrail",
|
||||
guardrailVersion="DRAFT",
|
||||
default_on=True,
|
||||
)
|
||||
guardrail.scan_only_tool_results = True
|
||||
return guardrail
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_only_tool_role_content_is_scanned(self):
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
guardrail = self._bedrock_guardrail()
|
||||
data = {
|
||||
"messages": [
|
||||
{"role": "system", "content": "SYSTEM-PROMPT-not-scanned"},
|
||||
{"role": "user", "content": "USER-PROMPT-not-scanned"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "ASSISTANT-not-scanned",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": {"name": "read_file", "arguments": '{"path": "report.html"}'},
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "tool", "tool_call_id": "call_1", "content": "TOOL-RESULT-scanned"},
|
||||
]
|
||||
}
|
||||
with patch.object(guardrail, "make_bedrock_api_request", new_callable=AsyncMock) as mock_api:
|
||||
mock_api.return_value = {"action": "NONE", "output": [], "outputs": []}
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
assert mock_api.call_count == 1
|
||||
scanned = [m["content"] for m in mock_api.call_args.kwargs["messages"]]
|
||||
assert scanned == ["TOOL-RESULT-scanned"]
|
||||
|
||||
@pytest.mark.parametrize("flag_value", [None, "false", 0, object()])
|
||||
@pytest.mark.asyncio
|
||||
async def test_scope_narrows_only_when_the_flag_is_actually_true(self, flag_value):
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
guardrail = self._bedrock_guardrail()
|
||||
guardrail.scan_only_tool_results = flag_value
|
||||
data = {
|
||||
"messages": [
|
||||
{"role": "user", "content": "USER-PROMPT"},
|
||||
{"role": "tool", "tool_call_id": "call_1", "content": "TOOL-RESULT"},
|
||||
]
|
||||
}
|
||||
with patch.object(guardrail, "make_bedrock_api_request", new_callable=AsyncMock) as mock_api:
|
||||
mock_api.return_value = {"action": "NONE", "output": [], "outputs": []}
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
assert mock_api.call_count == 1
|
||||
scanned = [m["content"] for m in mock_api.call_args.kwargs["messages"]]
|
||||
assert scanned == ["USER-PROMPT", "TOOL-RESULT"], (
|
||||
"anything but an explicit True must leave the whole request in scope"
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue