mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
feat(guardrails): support skipping assistant request messages
This commit is contained in:
parent
c8114ba41f
commit
556ce4e04f
17 changed files with 278 additions and 9 deletions
|
|
@ -232,6 +232,7 @@ bedrock_request_metadata_fields: Optional[Sequence[str]] = (
|
|||
store_audit_logs: bool | None = None
|
||||
skip_system_message_in_guardrail: bool = False
|
||||
skip_tool_message_in_guardrail: bool = False
|
||||
skip_assistant_message_in_guardrail: bool = False
|
||||
### end of callbacks #############
|
||||
|
||||
email: Optional[str] = (
|
||||
|
|
|
|||
|
|
@ -38,6 +38,7 @@ from litellm.llms.base_llm.guardrail_translation.utils import (
|
|||
anthropic_tool_name,
|
||||
anthropic_tool_names,
|
||||
effective_scan_only_tool_results_for_guardrail,
|
||||
effective_skip_assistant_message_for_guardrail,
|
||||
effective_skip_system_message_for_guardrail,
|
||||
effective_skip_tool_message_for_guardrail,
|
||||
merge_guardrailed_scoped_messages,
|
||||
|
|
@ -542,6 +543,7 @@ 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)
|
||||
skip_assistant: Final = effective_skip_assistant_message_for_guardrail(guardrail_to_apply)
|
||||
scan_only_tool_results: Final = effective_scan_only_tool_results_for_guardrail(guardrail_to_apply)
|
||||
|
||||
# The top-level prompt is translated on its own below so it can be hoisted in front of
|
||||
|
|
@ -568,6 +570,7 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
scan_only_tool_results=scan_only_tool_results,
|
||||
skip_system=False,
|
||||
skip_tool=skip_tool,
|
||||
skip_assistant=skip_assistant,
|
||||
)
|
||||
structured_messages: Final = [full_structured_messages[index] for index in scoped_message_indices]
|
||||
|
||||
|
|
@ -593,6 +596,7 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
msg_idx=msg_idx,
|
||||
skip_system_message=skip_system,
|
||||
skip_tool_message=skip_tool,
|
||||
skip_assistant_message=skip_assistant,
|
||||
scan_only_tool_results=scan_only_tool_results,
|
||||
)
|
||||
for msg_idx, message in enumerate(messages)
|
||||
|
|
@ -947,6 +951,7 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
skip_system_message: bool = False,
|
||||
skip_tool_message: bool = False,
|
||||
scan_only_tool_results: bool = False,
|
||||
skip_assistant_message: bool = False,
|
||||
) -> ExtractedInput:
|
||||
"""Extract text content and images from a message.
|
||||
|
||||
|
|
@ -960,6 +965,8 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
return cls._extract_midturn_system_text(message=message, msg_idx=msg_idx)
|
||||
if skip_tool_message and role.lower() == "tool":
|
||||
return EMPTY_EXTRACTED_INPUT
|
||||
if skip_assistant_message and role.lower() == "assistant":
|
||||
return EMPTY_EXTRACTED_INPUT
|
||||
|
||||
content: Final = message.get("content", None)
|
||||
if isinstance(content, str):
|
||||
|
|
|
|||
26
litellm/llms/base_llm/guardrail_translation/README.md
Normal file
26
litellm/llms/base_llm/guardrail_translation/README.md
Normal file
|
|
@ -0,0 +1,26 @@
|
|||
# Request message scoping
|
||||
|
||||
Set `skip_assistant_message_in_guardrail` to exclude assistant messages in request history from guardrail evaluation. This excludes their text, images and tool calls from the supported request handlers, and excludes the messages from structured guardrail inputs. The model still receives the original assistant messages, and post-call guardrails still evaluate newly generated replies
|
||||
|
||||
Enable it for the proxy in `config.yaml`:
|
||||
|
||||
```yaml
|
||||
litellm_settings:
|
||||
skip_assistant_message_in_guardrail: true
|
||||
```
|
||||
|
||||
Or set it for one guardrail:
|
||||
|
||||
```yaml
|
||||
guardrails:
|
||||
- guardrail_name: content-check
|
||||
litellm_params:
|
||||
guardrail: generic_guardrail_api
|
||||
mode: pre_call
|
||||
api_base: https://your-guardrail-api.example/check
|
||||
skip_assistant_message_in_guardrail: true
|
||||
```
|
||||
|
||||
A per-guardrail `true` or `false` takes precedence over the global value. Omitting the per-guardrail value inherits the global setting, which defaults to `false`
|
||||
|
||||
The setting follows the existing request-role filters for unified Chat Completions, Anthropic Messages, Bedrock Converse passthrough and Lakera v2. The unified Responses API handler does not currently implement these request-role filters
|
||||
|
|
@ -213,6 +213,15 @@ def _message_role(message: AllMessageValues) -> str:
|
|||
return str((message or {}).get("role") or "").lower()
|
||||
|
||||
|
||||
def effective_skip_assistant_message_for_guardrail(guardrail_to_apply: object) -> bool:
|
||||
per: Final = getattr(guardrail_to_apply, "skip_assistant_message_in_guardrail", None)
|
||||
if isinstance(per, bool):
|
||||
return per
|
||||
import litellm
|
||||
|
||||
return litellm.skip_assistant_message_in_guardrail
|
||||
|
||||
|
||||
def openai_messages_without_system(
|
||||
messages: Sequence[AllMessageValues],
|
||||
) -> tuple[AllMessageValues, ...]:
|
||||
|
|
@ -228,16 +237,21 @@ def openai_messages_without_tool(
|
|||
def filter_messages_by_skip_flags(
|
||||
guardrail_to_apply: object, messages: Sequence[AllMessageValues]
|
||||
) -> tuple[tuple[AllMessageValues, ...], bool]:
|
||||
system_filtered = (
|
||||
system_filtered: Final = (
|
||||
openai_messages_without_system(messages)
|
||||
if effective_skip_system_message_for_guardrail(guardrail_to_apply)
|
||||
else tuple(messages)
|
||||
)
|
||||
fully_filtered = (
|
||||
tool_filtered: Final = (
|
||||
openai_messages_without_tool(system_filtered)
|
||||
if effective_skip_tool_message_for_guardrail(guardrail_to_apply)
|
||||
else system_filtered
|
||||
)
|
||||
fully_filtered: Final = (
|
||||
tuple(message for message in tool_filtered if _message_role(message) != "assistant")
|
||||
if effective_skip_assistant_message_for_guardrail(guardrail_to_apply)
|
||||
else tool_filtered
|
||||
)
|
||||
return fully_filtered, len(fully_filtered) != len(messages)
|
||||
|
||||
|
||||
|
|
@ -251,11 +265,14 @@ def role_out_of_guardrail_scope(
|
|||
skip_system_message: bool,
|
||||
skip_tool_message: bool,
|
||||
scan_only_tool_results: bool = False,
|
||||
skip_assistant_message: bool = False,
|
||||
) -> bool:
|
||||
if skip_system_message and role == "system":
|
||||
return True
|
||||
if skip_tool_message and role == "tool":
|
||||
return True
|
||||
if skip_assistant_message and role == "assistant":
|
||||
return True
|
||||
return scan_only_tool_results and role not in ("tool", "function")
|
||||
|
||||
|
||||
|
|
@ -265,6 +282,7 @@ def scoped_structured_message_indices(
|
|||
scan_only_tool_results: bool,
|
||||
skip_system: bool,
|
||||
skip_tool: bool,
|
||||
skip_assistant: bool = False,
|
||||
) -> tuple[int, ...]:
|
||||
return tuple(
|
||||
index
|
||||
|
|
@ -273,6 +291,7 @@ def scoped_structured_message_indices(
|
|||
_message_role(message),
|
||||
skip_system_message=skip_system,
|
||||
skip_tool_message=skip_tool,
|
||||
skip_assistant_message=skip_assistant,
|
||||
scan_only_tool_results=scan_only_tool_results,
|
||||
)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ from typing_extensions import ReadOnly, TypedDict
|
|||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import (
|
||||
effective_skip_assistant_message_for_guardrail,
|
||||
effective_skip_system_message_for_guardrail,
|
||||
effective_skip_tool_message_for_guardrail,
|
||||
)
|
||||
|
|
@ -80,6 +81,7 @@ def _extract_converse_texts(
|
|||
body: dict,
|
||||
skip_system: bool,
|
||||
skip_tool: bool,
|
||||
skip_assistant: bool = False,
|
||||
) -> tuple[list[str], list[_StringHolder]]:
|
||||
"""
|
||||
Walk a Bedrock Converse request body and collect text content.
|
||||
|
|
@ -118,6 +120,8 @@ def _extract_converse_texts(
|
|||
for message in body.get("messages") or []:
|
||||
if not isinstance(message, dict):
|
||||
continue
|
||||
if skip_assistant and "role" in message and message["role"] == "assistant":
|
||||
continue
|
||||
for block in message.get("content") or []:
|
||||
if not isinstance(block, dict):
|
||||
continue
|
||||
|
|
@ -438,8 +442,9 @@ class BedrockPassthroughGuardrailHandler(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)
|
||||
skip_assistant: Final = effective_skip_assistant_message_for_guardrail(guardrail_to_apply)
|
||||
|
||||
texts, holders = _extract_converse_texts(body, skip_system, skip_tool)
|
||||
texts, holders = _extract_converse_texts(body, skip_system, skip_tool, skip_assistant)
|
||||
|
||||
if not texts:
|
||||
return data
|
||||
|
|
|
|||
|
|
@ -32,6 +32,7 @@ from litellm.llms.base_llm.guardrail_translation.base_translation import (
|
|||
from litellm.llms.base_llm.guardrail_translation.utils import (
|
||||
blocked_chat_stream_usage,
|
||||
effective_scan_only_tool_results_for_guardrail,
|
||||
effective_skip_assistant_message_for_guardrail,
|
||||
effective_skip_system_message_for_guardrail,
|
||||
effective_skip_tool_message_for_guardrail,
|
||||
merge_guardrailed_scoped_messages,
|
||||
|
|
@ -110,6 +111,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)
|
||||
skip_assistant: Final = effective_skip_assistant_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]] = []
|
||||
|
|
@ -130,6 +132,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
tool_call_task_mappings=tool_call_task_mappings,
|
||||
skip_system_message=skip_system,
|
||||
skip_tool_message=skip_tool,
|
||||
skip_assistant_message=skip_assistant,
|
||||
scan_only_tool_results=scan_only_tool_results,
|
||||
)
|
||||
|
||||
|
|
@ -146,6 +149,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
scan_only_tool_results=scan_only_tool_results,
|
||||
skip_system=skip_system,
|
||||
skip_tool=skip_tool,
|
||||
skip_assistant=skip_assistant,
|
||||
)
|
||||
if structured_messages:
|
||||
inputs["structured_messages"] = [structured_messages[index] for index in scoped_message_indices]
|
||||
|
|
@ -278,6 +282,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
skip_system_message: bool = False,
|
||||
skip_tool_message: bool = False,
|
||||
scan_only_tool_results: bool = False,
|
||||
skip_assistant_message: bool = False,
|
||||
) -> None:
|
||||
"""
|
||||
Extract text content, images, and tool calls from a message.
|
||||
|
|
@ -288,6 +293,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
str(message.get("role") or "").lower(),
|
||||
skip_system_message=skip_system_message,
|
||||
skip_tool_message=skip_tool_message,
|
||||
skip_assistant_message=skip_assistant_message,
|
||||
scan_only_tool_results=scan_only_tool_results,
|
||||
):
|
||||
return
|
||||
|
|
|
|||
|
|
@ -9910,6 +9910,18 @@
|
|||
"description": "Minimum severity to block (high, medium, low)",
|
||||
"title": "Severity Threshold"
|
||||
},
|
||||
"skip_assistant_message_in_guardrail": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "boolean"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"description": "When True, skip assistant-role messages in request history when building guardrail evaluation inputs. When False, include them even if the global litellm.skip_assistant_message_in_guardrail setting is True. When None, inherit the global setting. Does not skip checks on newly generated responses.",
|
||||
"title": "Skip Assistant Message In Guardrail"
|
||||
},
|
||||
"skip_system_message_in_guardrail": {
|
||||
"anyOf": [
|
||||
{
|
||||
|
|
@ -12671,6 +12683,18 @@
|
|||
"description": "The Singulr Guardrail ID. Get guardrail ID from Singulr Platform.",
|
||||
"title": "Singulr Guardrail Id"
|
||||
},
|
||||
"skip_assistant_message_in_guardrail": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "boolean"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"description": "When True, skip assistant-role messages in request history when building guardrail evaluation inputs. When False, include them even if the global litellm.skip_assistant_message_in_guardrail setting is True. When None, inherit the global setting. Does not skip checks on newly generated responses.",
|
||||
"title": "Skip Assistant Message In Guardrail"
|
||||
},
|
||||
"skip_system_message_in_guardrail": {
|
||||
"anyOf": [
|
||||
{
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ from litellm.integrations.custom_guardrail import (
|
|||
log_guardrail_information,
|
||||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import (
|
||||
effective_skip_assistant_message_for_guardrail,
|
||||
effective_skip_system_message_for_guardrail,
|
||||
effective_skip_tool_message_for_guardrail,
|
||||
)
|
||||
|
|
@ -458,7 +459,11 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
|
|||
|
||||
@override
|
||||
def structured_messages_cover_full_request(self) -> bool:
|
||||
return effective_skip_system_message_for_guardrail(self) or effective_skip_tool_message_for_guardrail(self)
|
||||
return (
|
||||
effective_skip_system_message_for_guardrail(self)
|
||||
or effective_skip_tool_message_for_guardrail(self)
|
||||
or effective_skip_assistant_message_for_guardrail(self)
|
||||
)
|
||||
|
||||
def _writeback_messages(
|
||||
self,
|
||||
|
|
@ -467,7 +472,7 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
|
|||
sent_indices: tuple[int, ...],
|
||||
request_data: dict[str, object],
|
||||
) -> list[AllMessageValues] | None:
|
||||
if effective_skip_system_message_for_guardrail(self) or effective_skip_tool_message_for_guardrail(self):
|
||||
if self.structured_messages_cover_full_request():
|
||||
request_messages: Final = request_data.get("messages")
|
||||
full_messages = (
|
||||
cast("list[AllMessageValues]", request_messages)
|
||||
|
|
|
|||
|
|
@ -15,10 +15,12 @@ from litellm.integrations.custom_guardrail import (
|
|||
CustomGuardrail,
|
||||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import (
|
||||
effective_skip_assistant_message_for_guardrail,
|
||||
effective_skip_system_message_for_guardrail,
|
||||
effective_skip_tool_message_for_guardrail,
|
||||
filter_messages_by_skip_flags,
|
||||
merge_guardrailed_scoped_messages,
|
||||
role_out_of_guardrail_scope,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client,
|
||||
|
|
@ -102,14 +104,19 @@ def _pre_masking_scope_indices(
|
|||
on length and the caller's strict positional zip raises."""
|
||||
skip_system: Final = effective_skip_system_message_for_guardrail(guardrail)
|
||||
skip_tool: Final = effective_skip_tool_message_for_guardrail(guardrail)
|
||||
skip_assistant: Final = effective_skip_assistant_message_for_guardrail(guardrail)
|
||||
return tuple(
|
||||
idx
|
||||
for idx, message in enumerate(messages)
|
||||
if isinstance(message, dict)
|
||||
and isinstance(message.get("content"), str)
|
||||
and message["content"]
|
||||
and not (skip_system and str(message.get("role") or "").lower() == "system")
|
||||
and not (skip_tool and str(message.get("role") or "").lower() == "tool")
|
||||
and not role_out_of_guardrail_scope(
|
||||
str(message.get("role") or "").lower(),
|
||||
skip_system_message=skip_system,
|
||||
skip_tool_message=skip_tool,
|
||||
skip_assistant_message=skip_assistant,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -257,6 +264,7 @@ class LakeraAIGuardrail(CustomGuardrail):
|
|||
skip_system_message_in_guardrail: bool | None = None,
|
||||
skip_tool_message_in_guardrail: bool | None = None,
|
||||
advisory_system_message: str | None = None,
|
||||
skip_assistant_message_in_guardrail: bool | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
|
|
@ -293,6 +301,7 @@ class LakeraAIGuardrail(CustomGuardrail):
|
|||
self.dev_info: bool | None = dev_info
|
||||
self.skip_system_message_in_guardrail = skip_system_message_in_guardrail
|
||||
self.skip_tool_message_in_guardrail = skip_tool_message_in_guardrail
|
||||
self.skip_assistant_message_in_guardrail = skip_assistant_message_in_guardrail
|
||||
self.on_flagged = on_flagged or "block"
|
||||
self.advisory_system_message = advisory_system_message
|
||||
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
|
||||
|
|
|
|||
|
|
@ -81,6 +81,7 @@ def initialize_lakera_v2(litellm_params: LitellmParams, guardrail: Guardrail):
|
|||
on_flagged=litellm_params.on_flagged,
|
||||
skip_system_message_in_guardrail=litellm_params.skip_system_message_in_guardrail,
|
||||
skip_tool_message_in_guardrail=litellm_params.skip_tool_message_in_guardrail,
|
||||
skip_assistant_message_in_guardrail=litellm_params.skip_assistant_message_in_guardrail,
|
||||
advisory_system_message=litellm_params.advisory_system_message,
|
||||
)
|
||||
litellm.logging_callback_manager.add_litellm_callback(_lakera_v2_callback)
|
||||
|
|
|
|||
|
|
@ -442,6 +442,7 @@ def _configure_callback_scoping(
|
|||
for scoping_param in (
|
||||
"skip_system_message_in_guardrail",
|
||||
"skip_tool_message_in_guardrail",
|
||||
"skip_assistant_message_in_guardrail",
|
||||
"scan_only_tool_results",
|
||||
):
|
||||
setattr(custom_guardrail_callback, scoping_param, getattr(litellm_params, scoping_param, None))
|
||||
|
|
|
|||
|
|
@ -836,6 +836,16 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up
|
|||
),
|
||||
)
|
||||
|
||||
skip_assistant_message_in_guardrail: bool | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"When True, skip assistant-role messages in request history when building "
|
||||
"guardrail evaluation inputs. When False, include them even if the global "
|
||||
"litellm.skip_assistant_message_in_guardrail setting is True. When None, "
|
||||
"inherit the global setting. Does not skip checks on newly generated responses."
|
||||
),
|
||||
)
|
||||
|
||||
scan_only_tool_results: bool | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
|
|
|
|||
|
|
@ -6,7 +6,8 @@ with guardrail transformations, specifically testing edge cases with empty choic
|
|||
"""
|
||||
|
||||
import json
|
||||
from typing import Any, Literal, Optional
|
||||
from copy import deepcopy
|
||||
from typing import Any, Final, Literal, Optional
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -21,6 +22,45 @@ from litellm.llms.anthropic.chat.guardrail_translation.handler import (
|
|||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
|
||||
@pytest.mark.parametrize("skip_assistant", [False, True])
|
||||
@pytest.mark.asyncio
|
||||
async def test_skip_assistant_keeps_tool_results_and_mask_writeback(skip_assistant: bool) -> None:
|
||||
guardrail: Final = MockMaskingGuardrail()
|
||||
guardrail.skip_assistant_message_in_guardrail = skip_assistant
|
||||
data: Final = {
|
||||
"model": "test-model",
|
||||
"messages": [
|
||||
{"role": "user", "content": "hello"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{"type": "text", "text": "previous reply"},
|
||||
{"type": "tool_use", "id": "call_1", "name": "search", "input": {"q": "old"}},
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "tool_result", "tool_use_id": "call_1", "content": "tool result"},
|
||||
{"type": "text", "text": "prohibited correction"},
|
||||
],
|
||||
},
|
||||
],
|
||||
}
|
||||
original_assistant: Final = deepcopy(data["messages"][1])
|
||||
|
||||
await AnthropicMessagesHandler().process_input_messages(data, guardrail)
|
||||
|
||||
assert guardrail.inputs is not None
|
||||
assert ("previous reply" in guardrail.inputs["texts"]) is not skip_assistant
|
||||
assert "tool result" in guardrail.inputs["texts"]
|
||||
assert bool(guardrail.inputs.get("tool_calls")) is not skip_assistant
|
||||
assert any(message["role"] == "assistant" for message in guardrail.inputs["structured_messages"]) is not skip_assistant
|
||||
assert data["messages"][1] == original_assistant
|
||||
assert data["messages"][2]["content"][0]["content"] == "tool result"
|
||||
assert data["messages"][2]["content"][1]["text"] == "[MASKED]"
|
||||
|
||||
|
||||
class MockPassThroughGuardrail(CustomGuardrail):
|
||||
"""Mock guardrail that passes through without blocking - for testing streaming fallback behavior"""
|
||||
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ Validates that:
|
|||
"""
|
||||
|
||||
import copy
|
||||
from typing import Final
|
||||
import pytest
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
|
|
@ -31,9 +32,31 @@ def _make_guardrail(apply_result: dict) -> MagicMock:
|
|||
g.apply_guardrail = AsyncMock(return_value=apply_result)
|
||||
g.skip_system_message_in_guardrail = False
|
||||
g.skip_tool_message_in_guardrail = False
|
||||
g.skip_assistant_message_in_guardrail = False
|
||||
return g
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_skip_assistant_preserves_converse_history_and_masks_user() -> None:
|
||||
assistant: Final = {
|
||||
"role": "assistant",
|
||||
"content": [{"text": "old reply"}, {"toolUse": {"toolUseId": "t1", "name": "search", "input": {"q": "old"}}}],
|
||||
}
|
||||
data: Final = {
|
||||
"endpoint": "model/test/converse",
|
||||
"data": {"messages": [assistant, {"role": "user", "content": [{"text": "private"}]}]},
|
||||
}
|
||||
original_assistant: Final = copy.deepcopy(assistant)
|
||||
guardrail: Final = _make_guardrail({"texts": ["[MASKED]"]})
|
||||
guardrail.skip_assistant_message_in_guardrail = True
|
||||
|
||||
await BedrockPassthroughGuardrailHandler().process_input_messages(data, guardrail)
|
||||
|
||||
assert guardrail.apply_guardrail.call_args.kwargs["inputs"]["texts"] == ["private"]
|
||||
assert data["data"]["messages"][0] == original_assistant
|
||||
assert data["data"]["messages"][1]["content"][0]["text"] == "[MASKED]"
|
||||
|
||||
|
||||
def _converse_data(endpoint: str = "model/anthropic.claude-3-sonnet/converse") -> dict:
|
||||
return {
|
||||
"endpoint": endpoint,
|
||||
|
|
|
|||
|
|
@ -6,9 +6,11 @@ with guardrail transformations, including tool calls.
|
|||
"""
|
||||
|
||||
import json
|
||||
from typing import Any, Literal, Optional
|
||||
from copy import deepcopy
|
||||
from typing import Any, Final, Literal, Optional
|
||||
|
||||
import pytest
|
||||
import litellm
|
||||
|
||||
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
|
|
@ -26,6 +28,66 @@ from litellm.types.utils import (
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"global_skip,per_guardrail_skip,expected_skip",
|
||||
[(False, None, False), (True, None, True), (False, True, True), (True, False, False)],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_skip_assistant_preserves_history_and_scans_other_roles(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
global_skip: bool,
|
||||
per_guardrail_skip: bool | None,
|
||||
expected_skip: bool,
|
||||
) -> None:
|
||||
monkeypatch.setattr(litellm, "skip_assistant_message_in_guardrail", global_skip, raising=False)
|
||||
guardrail: Final = MockGuardrail()
|
||||
guardrail.skip_assistant_message_in_guardrail = per_guardrail_skip
|
||||
data: Final = {
|
||||
"messages": [
|
||||
{"role": "user", "content": "hello"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "previous reply",
|
||||
"tool_calls": [{"id": "call_1", "type": "function", "function": {"name": "f", "arguments": '{"q":"old"}'}}],
|
||||
},
|
||||
{"role": "tool", "content": "tool result", "tool_call_id": "call_1"},
|
||||
],
|
||||
}
|
||||
original_assistant: Final = deepcopy(data["messages"][1])
|
||||
|
||||
result: Final = await OpenAIChatCompletionsHandler().process_input_messages(data, guardrail)
|
||||
|
||||
assert guardrail.last_inputs is not None
|
||||
assert guardrail.last_inputs["texts"] == (
|
||||
["hello", "tool result"] if expected_skip else ["hello", "previous reply", "tool result"]
|
||||
)
|
||||
assert [message["role"] for message in guardrail.last_inputs["structured_messages"]] == (
|
||||
["user", "tool"] if expected_skip else ["user", "assistant", "tool"]
|
||||
)
|
||||
assert guardrail.tool_calls_modified is not expected_skip
|
||||
assert result["messages"][0]["content"] == "HELLO"
|
||||
assert result["messages"][2]["content"] == "TOOL RESULT"
|
||||
if expected_skip:
|
||||
assert result["messages"][1] == original_assistant
|
||||
else:
|
||||
assert result["messages"][1]["content"] == "PREVIOUS REPLY"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_skip_assistant_history_does_not_skip_new_output() -> None:
|
||||
guardrail: Final = MockGuardrail()
|
||||
guardrail.skip_assistant_message_in_guardrail = True
|
||||
handler: Final = OpenAIChatCompletionsHandler()
|
||||
data: Final = {"messages": [{"role": "assistant", "content": "old reply"}]}
|
||||
|
||||
assert await handler.process_input_messages(data, guardrail) == data
|
||||
assert guardrail.last_inputs is None
|
||||
|
||||
response: Final = ModelResponse(choices=[Choices(message=Message(role="assistant", content="new reply"))])
|
||||
result: Final = await handler.process_output_response(response, guardrail)
|
||||
assert result.choices[0].message.content == "NEW REPLY"
|
||||
|
||||
|
||||
class MockGuardrail(CustomGuardrail):
|
||||
"""Mock guardrail for testing that transforms text and tool calls"""
|
||||
|
||||
|
|
|
|||
|
|
@ -6,6 +6,8 @@ Additional tests live in tests/guardrails_tests/test_lakera_v2.py.
|
|||
"""
|
||||
|
||||
import logging
|
||||
from copy import deepcopy
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -19,6 +21,7 @@ from litellm.llms.base_llm.guardrail_translation.utils import (
|
|||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.guardrails.guardrail_hooks.lakera_ai_v2 import (
|
||||
LakeraAIGuardrail,
|
||||
_apply_redacted_messages_back_preserving_fields,
|
||||
_build_lakera_inspection_messages,
|
||||
humanize_lakera_block_reasons,
|
||||
)
|
||||
|
|
@ -26,6 +29,23 @@ from litellm.types.guardrails import LitellmParams, Mode
|
|||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
||||
def test_skip_assistant_filters_inspection_and_preserves_masking_positions() -> None:
|
||||
guardrail: Final = LakeraAIGuardrail(api_key="test_key", skip_assistant_message_in_guardrail=True)
|
||||
messages: Final = [
|
||||
{"role": "assistant", "content": "old reply", "name": "helper"},
|
||||
{"role": "user", "content": "private"},
|
||||
]
|
||||
original_assistant: Final = deepcopy(messages[0])
|
||||
filtered, was_skipped = guardrail._filter_skipped_messages(messages)
|
||||
assert filtered == (messages[1],)
|
||||
assert was_skipped is True
|
||||
data: Final = {"messages": messages}
|
||||
|
||||
_apply_redacted_messages_back_preserving_fields(guardrail, data, [{"role": "user", "content": "[MASKED]"}])
|
||||
|
||||
assert data["messages"] == [original_assistant, {"role": "user", "content": "[MASKED]"}]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_lakera_post_call_success_hook_returns_model_response_when_pii_masked():
|
||||
"""
|
||||
|
|
|
|||
10
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
10
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -23977,6 +23977,11 @@ export interface components {
|
|||
* @description Minimum severity to block (high, medium, low)
|
||||
*/
|
||||
severity_threshold?: string | null;
|
||||
/**
|
||||
* Skip Assistant Message In Guardrail
|
||||
* @description When True, skip assistant-role messages in request history when building guardrail evaluation inputs. When False, include them even if the global litellm.skip_assistant_message_in_guardrail setting is True. When None, inherit the global setting. Does not skip checks on newly generated responses.
|
||||
*/
|
||||
skip_assistant_message_in_guardrail?: boolean | null;
|
||||
/**
|
||||
* Skip System Message In Guardrail
|
||||
* @description When True, unified guardrails skip system-role messages when building evaluation inputs (texts and structured_messages). When False, system messages are included even if litellm_settings sets a global skip. When None, use the global litellm.skip_system_message_in_guardrail setting. For Anthropic /v1/messages, the flag applies only to the trusted top-level system prompt. In-sequence system entries are untrusted client input and remain in texts and structured_messages.
|
||||
|
|
@ -31501,6 +31506,11 @@ export interface components {
|
|||
* @description The Singulr Guardrail ID. Get guardrail ID from Singulr Platform.
|
||||
*/
|
||||
singulr_guardrail_id?: string | null;
|
||||
/**
|
||||
* Skip Assistant Message In Guardrail
|
||||
* @description When True, skip assistant-role messages in request history when building guardrail evaluation inputs. When False, include them even if the global litellm.skip_assistant_message_in_guardrail setting is True. When None, inherit the global setting. Does not skip checks on newly generated responses.
|
||||
*/
|
||||
skip_assistant_message_in_guardrail?: boolean | null;
|
||||
/**
|
||||
* Skip System Message In Guardrail
|
||||
* @description When True, unified guardrails skip system-role messages when building evaluation inputs (texts and structured_messages). When False, system messages are included even if litellm_settings sets a global skip. When None, use the global litellm.skip_system_message_in_guardrail setting. For Anthropic /v1/messages, the flag applies only to the trusted top-level system prompt. In-sequence system entries are untrusted client input and remain in texts and structured_messages.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue