diff --git a/litellm/__init__.py b/litellm/__init__.py index e1da202b9ee..54cc958e8a3 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -236,6 +236,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] = ( diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py index 15380f57d17..9f8993eb213 100644 --- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py @@ -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, @@ -535,6 +536,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 @@ -561,6 +563,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] @@ -586,6 +589,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) @@ -934,6 +938,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. @@ -947,6 +952,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): diff --git a/litellm/llms/base_llm/guardrail_translation/utils.py b/litellm/llms/base_llm/guardrail_translation/utils.py index 51d43436fc9..a73dfda4f55 100644 --- a/litellm/llms/base_llm/guardrail_translation/utils.py +++ b/litellm/llms/base_llm/guardrail_translation/utils.py @@ -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, ) ) diff --git a/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py b/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py index eac8afd767c..28804ac9890 100644 --- a/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py +++ b/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py @@ -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, ) @@ -83,6 +84,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. @@ -121,6 +123,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 @@ -441,8 +445,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 diff --git a/litellm/llms/openai/chat/guardrail_translation/handler.py b/litellm/llms/openai/chat/guardrail_translation/handler.py index fa5512e7bfe..3d08212257b 100644 --- a/litellm/llms/openai/chat/guardrail_translation/handler.py +++ b/litellm/llms/openai/chat/guardrail_translation/handler.py @@ -33,6 +33,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, @@ -111,6 +112,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]] = [] @@ -131,6 +133,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, ) @@ -147,6 +150,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] @@ -279,6 +283,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. @@ -289,6 +294,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 diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index ded4db6d2aa..03ff705004d 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -10330,6 +10330,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": [ { @@ -13138,6 +13150,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": [ { diff --git a/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py b/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py index 3d4aba4ac02..8c665b68926 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py +++ b/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py @@ -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, ) @@ -464,7 +465,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, @@ -473,7 +478,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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py index 2f98a9afbd8..1096e1ffd56 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py +++ b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py @@ -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())) diff --git a/litellm/proxy/guardrails/guardrail_initializers.py b/litellm/proxy/guardrails/guardrail_initializers.py index c422902d30d..be72bb451d1 100644 --- a/litellm/proxy/guardrails/guardrail_initializers.py +++ b/litellm/proxy/guardrails/guardrail_initializers.py @@ -82,6 +82,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) diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index 0dc50cd6196..f079dad757c 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -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)) diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 579a3f6322f..50fd3484395 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -940,6 +940,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=( diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_lakera_ai_v2.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_lakera_ai_v2.py index ee3f8659d51..b5502d4537b 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_lakera_ai_v2.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_lakera_ai_v2.py @@ -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(): """ diff --git a/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py b/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py index 1c1b68de6d6..fb50ef8107a 100644 --- a/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py +++ b/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py @@ -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""" diff --git a/tests/unit/llms/bedrock/passthrough/guardrail_translation/test_handler.py b/tests/unit/llms/bedrock/passthrough/guardrail_translation/test_handler.py index da7ed635dcb..62c0e3a7402 100644 --- a/tests/unit/llms/bedrock/passthrough/guardrail_translation/test_handler.py +++ b/tests/unit/llms/bedrock/passthrough/guardrail_translation/test_handler.py @@ -10,6 +10,7 @@ Validates that: """ import copy +from typing import Final import pytest from unittest.mock import AsyncMock, MagicMock @@ -31,6 +32,7 @@ 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 @@ -1212,3 +1214,41 @@ class TestDeAnonymizeConverseStream: hook_spy.assert_not_called() assert result is stream_bytes + + +@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 = MagicMock() + guardrail.apply_guardrail = AsyncMock(return_value={"texts": ["[MASKED]"]}) + guardrail.skip_system_message_in_guardrail = False + guardrail.skip_tool_message_in_guardrail = False + 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]" diff --git a/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py b/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py index 5c85faa5e13..b36c8b005c2 100644 --- a/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py +++ b/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py @@ -7,11 +7,12 @@ with guardrail transformations, including tool calls. import json from collections.abc import Mapping -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 from litellm.llms.base_llm.guardrail_translation.base_translation import StreamingScanKey from litellm.llms.openai.chat.guardrail_translation.handler import ( @@ -27,6 +28,68 @@ 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""" diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 2b8ed9aa58d..09043d53277 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -25536,6 +25536,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. @@ -34763,6 +34768,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.