diff --git a/litellm/proxy/guardrails/_content_utils.py b/litellm/proxy/guardrails/_content_utils.py index 7529fe99f52..574e3c765c8 100644 --- a/litellm/proxy/guardrails/_content_utils.py +++ b/litellm/proxy/guardrails/_content_utils.py @@ -146,6 +146,27 @@ def iter_message_text(data: Mapping[str, object]) -> Iterator[str]: yield from _iter_text_parts_in_content(message.get("content")) +def iter_request_messages(data: Mapping[str, object]) -> Iterator[Mapping[str, object]]: + """Yield the request's messages in prompt order. + + A top-level ``system`` prompt (Anthropic) and Responses-API ``instructions`` + that carry text (a string or a list of parts) come first, each as a ``system`` + message, then ``messages`` and ``input``. A body in any other shape yields nothing. + """ + top_level_prompts: Final = tuple( + prompt + for prompt in (data.get("system"), data.get("instructions")) + if isinstance(prompt, (str, list)) and prompt + ) + yield from ({"role": "system", "content": prompt} for prompt in top_level_prompts) + yield from (message for message in _iter_inspection_messages(data) if isinstance(message, dict)) + + +def message_text(message: Mapping[str, object]) -> str: + """Join a message's text fragments, skipping images and other non-text parts.""" + return "\n".join(_iter_text_parts_in_content(message.get("content"))) + + def walk_user_text(data: dict[str, Any], visit: Callable[[str], str]) -> int: """Rewrite every text fragment in place via ``visit``. diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py index e3511d46544..503e80d250c 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py @@ -39,6 +39,10 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" streaming_end_of_stream_only=_get_config_value(litellm_params, optional_params, "streaming_end_of_stream_only"), streaming_sampling_rate=_get_config_value(litellm_params, optional_params, "streaming_sampling_rate"), streaming_transform_mode=_get_config_value(litellm_params, optional_params, "streaming_transform_mode"), + skip_if_system_prompt_matches=_get_config_value( + litellm_params, optional_params, "skip_if_system_prompt_matches" + ), + skip_if_first_role_in=_get_config_value(litellm_params, optional_params, "skip_if_first_role_in"), ) litellm.logging_callback_manager.add_litellm_callback(_generic_guardrail_api_callback) diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/config_parsing.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/config_parsing.py new file mode 100644 index 00000000000..10934c8e244 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/config_parsing.py @@ -0,0 +1,15 @@ +import re +from collections.abc import Sequence + + +def config_values(raw: Sequence[str] | None, *, option_name: str) -> tuple[str, ...]: + if isinstance(raw, str): + raise ValueError(f"{option_name} must be a list of strings, got the single string {raw!r}") + return tuple(raw or ()) + + +def compile_patterns(raw: Sequence[str] | None, *, option_name: str) -> tuple[re.Pattern[str], ...]: + try: + return tuple(re.compile(pattern) for pattern in config_values(raw, option_name=option_name)) + except re.error as e: + raise ValueError(f"{option_name} contains an invalid regex: {e}") from e diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py index 3d1a173635e..8e9995b3629 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py @@ -20,6 +20,7 @@ from litellm.integrations.custom_guardrail import ( log_guardrail_information, ) from litellm.llms.custom_httpx.http_handler import ( + AsyncHTTPHandler, get_async_httpx_client, httpxSpecialProvider, ) @@ -33,6 +34,8 @@ from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import ) from litellm.types.utils import GenericGuardrailAPIInputs +from .message_filter import build_message_skip_filter + if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel @@ -170,6 +173,10 @@ def _structured_rows_to_write_back( ) +def _passthrough_inputs(inputs: GenericGuardrailAPIInputs) -> GenericGuardrailAPIInputs: + return GenericGuardrailAPIInputs(**inputs) + + class GenericGuardrailAPI(CustomGuardrail): """ Generic Guardrail API integration for LiteLLM. @@ -204,9 +211,14 @@ class GenericGuardrailAPI(CustomGuardrail): streaming_end_of_stream_only: bool | None = None, streaming_sampling_rate: int | None = None, streaming_transform_mode: Literal["block_only", "incremental_diff"] | None = None, + skip_if_system_prompt_matches: Sequence[str] | None = None, + skip_if_first_role_in: Sequence[str] | None = None, + async_handler: AsyncHTTPHandler | None = None, **kwargs, ): - self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback) + self.async_handler = async_handler or get_async_httpx_client( + llm_provider=httpxSpecialProvider.GuardrailCallback + ) self.headers = headers or {} self.extra_headers = extra_headers or [] @@ -256,6 +268,13 @@ class GenericGuardrailAPI(CustomGuardrail): super().__init__(**kwargs) + self._message_skip_filter: Final = build_message_skip_filter( + skip_if_system_prompt_matches=skip_if_system_prompt_matches, + skip_if_first_role_in=skip_if_first_role_in, + guardrail_name=self.guardrail_name, + event_hook=self.event_hook, + ) + verbose_proxy_logger.debug("Generic Guardrail API initialized with api_base: %s", self.api_base) def _extract_user_api_key_metadata(self, request_data: dict) -> GenericGuardrailAPIMetadata: @@ -432,6 +451,19 @@ class GenericGuardrailAPI(CustomGuardrail): if request_data is None: request_data = {} + skip_reason: Final = self._message_skip_filter.skip_reason( + input_type=input_type, + request_data=request_data, + logging_obj=logging_obj, + ) + if skip_reason is not None: + self.add_standard_logging_guardrail_information_to_request_data( + guardrail_json_response=skip_reason, + request_data=request_data, + guardrail_status="not_run", + ) + return _passthrough_inputs(inputs) + request_body: Final = request_data.get("body") or {} # Merge additional provider specific params from config and dynamic params diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/message_filter.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/message_filter.py new file mode 100644 index 00000000000..1b3cc453f93 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/message_filter.py @@ -0,0 +1,154 @@ +import re +import uuid +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from itertools import chain +from typing import TYPE_CHECKING, Final, Literal + +from litellm._logging import verbose_proxy_logger +from litellm.proxy.guardrails._content_utils import iter_request_messages, message_text +from litellm.types.guardrails import GuardrailEventHooks, Mode + +from .config_parsing import compile_patterns, config_values + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + +INSTRUCTION_ROLES: Final = frozenset({"system", "developer"}) + +MAX_INSTRUCTION_SEARCH_CHARS: Final = 16 * 1024 + +SYSTEM_PROMPT_SKIP_REASON: Final = "skipped: skip_if_system_prompt_matches" +FIRST_ROLE_SKIP_REASON: Final = "skipped: skip_if_first_role_in" + +_REQUEST_SIDE_HOOKS: Final = frozenset( + hook.value + for hook in ( + GuardrailEventHooks.pre_call, + GuardrailEventHooks.during_call, + GuardrailEventHooks.pre_mcp_call, + GuardrailEventHooks.during_mcp_call, + ) +) + + +def _configured_hooks(event_hook: str | Sequence[str] | Mode | None) -> tuple[str, ...]: + if event_hook is None: + return () + if isinstance(event_hook, Mode): + mode_values: Final = (event_hook.default, *event_hook.tags.values()) + return tuple(chain.from_iterable(_configured_hooks(value) for value in mode_values)) + if isinstance(event_hook, str): + return (_hook_value(event_hook),) + return tuple(_hook_value(hook) for hook in event_hook) + + +def _hook_value(hook: str) -> str: + return hook.value if isinstance(hook, GuardrailEventHooks) else hook + + +def _has_request_side_hook(event_hook: str | Sequence[str] | Mode | None) -> bool: + return event_hook is None or any(hook in _REQUEST_SIDE_HOOKS for hook in _configured_hooks(event_hook)) + + +def _role(message: Mapping[str, object]) -> str | None: + role: Final = message.get("role") + return role if isinstance(role, str) else None + + +@dataclass(frozen=True, slots=True) +class MessageSkipPolicy: + system_prompt_patterns: tuple[re.Pattern[str], ...] = () + first_role_in: frozenset[str] = frozenset() + + @property + def enabled(self) -> bool: + return bool(self.system_prompt_patterns or self.first_role_in) + + def skip_reason(self, request_data: Mapping[str, object]) -> str | None: + messages: Final = tuple(iter_request_messages(request_data)) + if not messages: + return None + if _role(messages[0]) in self.first_role_in: + return FIRST_ROLE_SKIP_REASON + instructions: Final = (message for message in messages if _role(message) in INSTRUCTION_ROLES) + return ( + SYSTEM_PROMPT_SKIP_REASON if any(self._instruction_matches(message) for message in instructions) else None + ) + + def _instruction_matches(self, message: Mapping[str, object]) -> bool: + searched: Final = message_text(message)[:MAX_INSTRUCTION_SEARCH_CHARS] + return any(pattern.search(searched) is not None for pattern in self.system_prompt_patterns) + + +@dataclass(frozen=True, slots=True) +class _SkipDecision: + reason: str + + +@dataclass(frozen=True, slots=True) +class MessageSkipFilter: + """Skips a matching request and replays that decision on its response, which has no system prompt to match. + + The decision rides on the call's logging object, which every request and response hook of the call shares. + The key is unique per filter instance, since config allows two guardrails with the same name. The value is + a ``_SkipDecision``: request body keys can reach ``model_call_details`` through ``optional_params``, but only + as JSON values, so they cannot forge one. + """ + + policy: MessageSkipPolicy + marker: str + + def skip_reason( + self, + *, + input_type: Literal["request", "response"], + request_data: Mapping[str, object], + logging_obj: "LiteLLMLoggingObj | None", + ) -> str | None: + if not self.policy.enabled: + return None + if input_type == "response": + return self._recorded_reason(logging_obj) + reason: Final = self.policy.skip_reason(request_data) + if reason is not None and logging_obj is not None: + logging_obj.model_call_details[self.marker] = _SkipDecision(reason) + return reason + + def _recorded_reason(self, logging_obj: "LiteLLMLoggingObj | None") -> str | None: + recorded: Final = logging_obj.model_call_details.get(self.marker) if logging_obj is not None else None + return recorded.reason if isinstance(recorded, _SkipDecision) else None + + +def build_message_skip_filter( + *, + skip_if_system_prompt_matches: Sequence[str] | None, + skip_if_first_role_in: Sequence[str] | None, + guardrail_name: str | None, + event_hook: str | Sequence[str] | Mode | None, +) -> MessageSkipFilter: + policy: Final = MessageSkipPolicy( + system_prompt_patterns=compile_patterns( + skip_if_system_prompt_matches, option_name="skip_if_system_prompt_matches" + ), + first_role_in=frozenset(config_values(skip_if_first_role_in, option_name="skip_if_first_role_in")), + ) + if policy.enabled: + verbose_proxy_logger.warning( + "Generic Guardrail API (%s): skip_if_system_prompt_matches / skip_if_first_role_in match on the " + "request body, which the caller controls, so a caller that knows the configured value can exempt " + "itself from this guardrail. Use skip_if_key_alias_in / skip_if_team_id_in when the exemption must " + "hold against the caller.", + guardrail_name, + ) + if policy.enabled and not _has_request_side_hook(event_hook): + verbose_proxy_logger.warning( + "Generic Guardrail API (%s): skip_if_system_prompt_matches / skip_if_first_role_in need a " + "request-side hook (pre_call, during_call, pre_mcp_call or during_mcp_call) to decide anything. " + "mode=%s only sees responses, so nothing will be skipped.", + guardrail_name, + event_hook, + ) + return MessageSkipFilter( + policy=policy, marker=f"generic_guardrail_api_message_skip::{guardrail_name}::{uuid.uuid4().hex}" + ) diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py index 44e2cc2404f..825d426d6d3 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py @@ -103,6 +103,31 @@ class GenericGuardrailAPIOptionalParams(BaseModel): ), ) + skip_if_system_prompt_matches: tuple[str, ...] | None = Field( + default=None, + description=( + "Regex patterns searched in the request's instructions: system and developer messages, a top-level " + "system prompt (Anthropic) and Responses API instructions, never user text. Only " + "the first 16384 characters of each instruction message are searched. On a match the guardrail is " + "skipped for the call: nothing is sent for the request or for its paired response. The decision reads " + "the full request, so skip_system_message_in_guardrail and scan_only_tool_results do not change it. " + "Needs a request-side hook (pre_call or during_call) in mode. An invalid regex fails at startup. The " + "caller controls its own messages, so a caller that knows a pattern can exempt itself: treat this as " + "traffic scoping, not enforcement. Patterns run on caller-supplied text and Python's re has no " + "timeout, so keep them linear-time (no nested quantifiers such as (a+)+)." + ), + ) + + skip_if_first_role_in: tuple[str, ...] | None = Field( + default=None, + description=( + "If the role of the request's first message is in this list (e.g. ['developer']), the guardrail is " + "skipped for the call, request and paired response alike. A top-level system prompt or Responses API " + "instructions count as a leading system message. Same mode requirement and same trust boundary as " + "skip_if_system_prompt_matches: the caller chooses the roles it sends." + ), + ) + class GenericGuardrailAPIConfigModel( GuardrailConfigModel[GenericGuardrailAPIOptionalParams], diff --git a/tests/test_litellm/proxy/guardrails/test_content_utils.py b/tests/test_litellm/proxy/guardrails/test_content_utils.py index 920ffc77095..858efb20da7 100644 --- a/tests/test_litellm/proxy/guardrails/test_content_utils.py +++ b/tests/test_litellm/proxy/guardrails/test_content_utils.py @@ -7,6 +7,8 @@ from litellm.proxy.guardrails._content_utils import ( is_non_conversational_call_type, is_string_batch_input, iter_message_text, + iter_request_messages, + message_text, walk_user_text, ) @@ -741,3 +743,58 @@ def test_is_non_conversational_call_type_defaults_to_inspecting_unknown_call_typ """A call type this module has never heard of must still be inspected — failing closed is the point of the deny-list.""" assert is_non_conversational_call_type("some_future_call_type") is False + + +# ── iter_request_messages / message_text ───────────────────────────────────────── + + +def test_iter_request_messages_puts_a_top_level_system_prompt_first(): + data = { + "system": [{"type": "text", "text": "anthropic rules"}], + "messages": [{"role": "user", "content": "hi"}, "not a message", {"role": "assistant", "content": "yo"}], + } + assert list(iter_request_messages(data)) == [ + {"role": "system", "content": [{"type": "text", "text": "anthropic rules"}]}, + {"role": "user", "content": "hi"}, + {"role": "assistant", "content": "yo"}, + ] + + +def test_iter_request_messages_puts_responses_instructions_before_input(): + data = {"instructions": "be brief", "input": [{"role": "developer", "content": "id=agent"}, "question"]} + assert list(iter_request_messages(data)) == [ + {"role": "system", "content": "be brief"}, + {"role": "developer", "content": "id=agent"}, + {"role": "user", "content": "question"}, + ] + + +def test_iter_request_messages_yields_nothing_for_an_unknown_shape(): + data = {"contents": [{"role": "user", "parts": [{"text": "hi"}]}], "system": "", "instructions": None} + assert list(iter_request_messages(data)) == [] + + +def test_message_text_joins_text_parts_only(): + message = { + "role": "system", + "content": [ + {"type": "text", "text": "first"}, + {"type": "image_url", "image_url": {"url": "data:image/png;base64,AAAA"}}, + "bare string part", + {"text": "bedrock converse block"}, + ], + } + assert message_text(message) == "first\nbare string part\nbedrock converse block" + + +def test_message_text_of_a_message_without_content_is_empty(): + assert message_text({"role": "system"}) == "" + + +def test_iter_request_messages_ignores_top_level_prompts_that_carry_no_text(): + data = { + "instructions": 1, + "system": {"not": "text"}, + "messages": [{"role": "developer", "content": "id=agent"}], + } + assert list(iter_request_messages(data)) == [{"role": "developer", "content": "id=agent"}] diff --git a/tests/unit/proxy/guardrails/__init__.py b/tests/unit/proxy/guardrails/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/__init__.py b/tests/unit/proxy/guardrails/guardrail_hooks/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_message_filter.py b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_message_filter.py new file mode 100644 index 00000000000..486b77605da --- /dev/null +++ b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_message_filter.py @@ -0,0 +1,572 @@ +import json +import logging +from collections.abc import Callable +from datetime import datetime, timezone +from typing import Final + +import httpx +import pytest + +from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.llms.anthropic.chat.guardrail_translation.handler import AnthropicMessagesHandler +from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler +from litellm.llms.openai.chat.guardrail_translation.handler import OpenAIChatCompletionsHandler +from litellm.llms.openai.responses.guardrail_translation.handler import OpenAIResponsesHandler +from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI, initialize_guardrail +from litellm.proxy.guardrails.guardrail_registry import _configure_callback_scoping +from litellm.types.guardrails import Guardrail, LitellmParams, Mode +from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPIOptionalParams + +MARKER: Final = "internal-agent-7f3c" +TRUST_BOUNDARY_WARNING: Final = "which the caller controls" +NO_REQUEST_HOOK_WARNING: Final = "need a request-side hook" +DOCUMENTED_SEARCH_LIMIT: Final = 16384 + + +class _Endpoint: + def __init__(self, response_body: dict[str, object] | None = None) -> None: + self.received: Final[list[dict[str, object]]] = [] # mutable-ok: records what the endpoint was sent + self._response_body: Final = response_body or {"action": "NONE"} + self.handler: Final = AsyncHTTPHandler(transport=httpx.MockTransport(self._respond)) + + def _respond(self, request: httpx.Request) -> httpx.Response: + self.received.append(json.loads(request.content)) + return httpx.Response(200, json=self._response_body, request=request) + + @property + def input_types(self) -> list[object]: + return [payload["input_type"] for payload in self.received] + + +def _guardrail( + endpoint: _Endpoint, + *, + name: str = "gg-skip", + event_hook: str | list[str] | Mode | None = None, + **options: object, +) -> GenericGuardrailAPI: + return GenericGuardrailAPI( + api_base="https://guardrail.test", + guardrail_name=name, + event_hook=event_hook or ["pre_call", "post_call"], + default_on=True, + async_handler=endpoint.handler, + **options, + ) + + +def _logging_obj(call_id: str = "call-1") -> Logging: + return Logging( + model="gpt-4o", + messages=[], + stream=False, + call_type="acompletion", + start_time=datetime(2026, 1, 1, tzinfo=timezone.utc), + litellm_call_id=call_id, + function_id=call_id, + ) + + +async def _request( + guardrail: GenericGuardrailAPI, body: dict[str, object], *, logging_obj: Logging | None +) -> dict[str, object]: + return dict( + await guardrail.apply_guardrail( + inputs={"texts": ["hello"]}, + request_data=body, + input_type="request", + logging_obj=logging_obj, + ) + ) + + +async def _response( + guardrail: GenericGuardrailAPI, *, logging_obj: Logging | None, request_data: dict[str, object] | None = None +) -> dict[str, object]: + return dict( + await guardrail.apply_guardrail( + inputs={"texts": ["model output"]}, + request_data={} if request_data is None else request_data, + input_type="response", + logging_obj=logging_obj, + ) + ) + + +def _chat(*messages: dict[str, object]) -> dict[str, object]: + return {"model": "gpt-4o", "messages": list(messages)} + + +def _system_marked_chat() -> dict[str, object]: + return _chat( + {"role": "system", "content": f"you are {MARKER}, tracked elsewhere"}, + {"role": "user", "content": "hello"}, + ) + + +@pytest.mark.asyncio +async def test_matching_system_prompt_skips_request_and_paired_response(): + endpoint: Final = _Endpoint(response_body={"action": "BLOCKED", "blocked_reason": "would block"}) + guardrail: Final = _guardrail(endpoint, skip_if_system_prompt_matches=[MARKER]) + logging_obj: Final = _logging_obj() + + request_result: Final = await _request(guardrail, _system_marked_chat(), logging_obj=logging_obj) + response_result: Final = await _response(guardrail, logging_obj=logging_obj) + + assert endpoint.received == [] + assert request_result == {"texts": ["hello"]} + assert response_result == {"texts": ["model output"]} + + +def _recorded(request_data: dict[str, object]) -> list[tuple[object, object]]: + metadata: Final = request_data["metadata"] + assert isinstance(metadata, dict) + return [ + (entry["guardrail_status"], entry["guardrail_response"]) + for entry in metadata["standard_logging_guardrail_information"] + ] + + +@pytest.mark.parametrize( + ("options", "build_body", "reason"), + [ + pytest.param( + {"skip_if_system_prompt_matches": [MARKER]}, + _system_marked_chat, + "skipped: skip_if_system_prompt_matches", + id="system_prompt", + ), + pytest.param( + {"skip_if_first_role_in": ["developer"]}, + lambda: _chat({"role": "developer", "content": "instructions"}, {"role": "user", "content": "hello"}), + "skipped: skip_if_first_role_in", + id="first_role", + ), + ], +) +@pytest.mark.asyncio +async def test_a_skipped_call_records_one_not_run_entry_on_each_side( + options: dict[str, object], build_body: Callable[[], dict[str, object]], reason: str +): + guardrail: Final = _guardrail(_Endpoint(), **options) + logging_obj: Final = _logging_obj() + request_data: Final = build_body() + response_data: Final[dict[str, object]] = {"model": "gpt-4o"} + + await _request(guardrail, request_data, logging_obj=logging_obj) + await _response(guardrail, logging_obj=logging_obj, request_data=response_data) + + assert _recorded(request_data) == [("not_run", reason)] + assert _recorded(response_data) == [("not_run", reason)] + + +@pytest.mark.asyncio +async def test_a_scanned_call_still_records_success(): + guardrail: Final = _guardrail(_Endpoint(), skip_if_system_prompt_matches=[MARKER]) + logging_obj: Final = _logging_obj() + request_data: Final = _chat({"role": "system", "content": "plain"}, {"role": "user", "content": "hello"}) + response_data: Final[dict[str, object]] = {"model": "gpt-4o"} + + await _request(guardrail, request_data, logging_obj=logging_obj) + await _response(guardrail, logging_obj=logging_obj, request_data=response_data) + + assert [status for status, _ in _recorded(request_data)] == ["success"] + assert [status for status, _ in _recorded(response_data)] == ["success"] + + +@pytest.mark.parametrize( + "build_body", + [ + pytest.param(lambda: _chat({"role": "system", "content": "plain"}), id="scanned_request"), + pytest.param(_system_marked_chat, id="skipped_request"), + ], +) +@pytest.mark.asyncio +async def test_a_body_key_named_like_the_skip_marker_cannot_skip_the_response( + build_body: Callable[[], dict[str, object]], +): + endpoint: Final = _Endpoint() + guardrail: Final = _guardrail(endpoint, name="gg-skip", skip_if_system_prompt_matches=[MARKER]) + logging_obj: Final = _logging_obj() + probe: Final = _logging_obj("call-probe") + await _request(guardrail, _system_marked_chat(), logging_obj=probe) + (forged_key,) = (key for key in probe.model_call_details if key.startswith("generic_guardrail_api_message_skip::")) + + await _request(guardrail, build_body(), logging_obj=logging_obj) + logging_obj.update_environment_variables( + litellm_params={}, optional_params={forged_key: "skipped: skip_if_system_prompt_matches"} + ) + await _response(guardrail, logging_obj=logging_obj) + + assert logging_obj.model_call_details[forged_key] == "skipped: skip_if_system_prompt_matches" + assert endpoint.input_types[-1:] == ["response"], "a provider param copied into model_call_details forged a skip" + + +@pytest.mark.asyncio +async def test_every_response_call_of_a_skipped_request_stays_skipped(): + endpoint: Final = _Endpoint() + guardrail: Final = _guardrail(endpoint, skip_if_system_prompt_matches=[MARKER]) + logging_obj: Final = _logging_obj() + + await _request(guardrail, _system_marked_chat(), logging_obj=logging_obj) + for _ in range(3): + await _response(guardrail, logging_obj=logging_obj) + + assert endpoint.received == [], "sampled stream chunks and the end-of-stream call all replay the skip" + + +@pytest.mark.asyncio +async def test_marker_in_user_message_is_still_scanned(): + endpoint: Final = _Endpoint() + guardrail: Final = _guardrail(endpoint, skip_if_system_prompt_matches=[MARKER]) + logging_obj: Final = _logging_obj() + + await _request( + guardrail, + _chat({"role": "system", "content": "you are helpful"}, {"role": "user", "content": f"what is {MARKER}?"}), + logging_obj=logging_obj, + ) + await _response(guardrail, logging_obj=logging_obj) + + assert endpoint.input_types == ["request", "response"] + + +@pytest.mark.asyncio +async def test_developer_message_counts_as_instructions(): + endpoint: Final = _Endpoint() + guardrail: Final = _guardrail(endpoint, skip_if_system_prompt_matches=[r"internal-agent-[0-9a-f]{4}\b"]) + + await _request( + guardrail, + _chat( + {"role": "user", "content": "hello"}, + {"role": "developer", "content": [{"type": "text", "text": f"id={MARKER}"}]}, + ), + logging_obj=_logging_obj(), + ) + + assert endpoint.received == [] + + +@pytest.mark.parametrize( + ("padding", "skipped"), + [(DOCUMENTED_SEARCH_LIMIT - len(MARKER), True), (1 + DOCUMENTED_SEARCH_LIMIT - len(MARKER), False)], +) +@pytest.mark.asyncio +async def test_only_the_start_of_each_instruction_message_is_searched(padding: int, skipped: bool): + endpoint: Final = _Endpoint() + guardrail: Final = _guardrail(endpoint, skip_if_system_prompt_matches=[MARKER]) + + await _request( + guardrail, + _chat({"role": "system", "content": "x" * padding + MARKER}, {"role": "user", "content": "hello"}), + logging_obj=_logging_obj(), + ) + + assert endpoint.input_types == ([] if skipped else ["request"]) + + +@pytest.mark.asyncio +async def test_a_malformed_role_is_scanned_instead_of_failing_the_request(): + endpoint: Final = _Endpoint() + guardrail: Final = _guardrail(endpoint, skip_if_system_prompt_matches=[MARKER], skip_if_first_role_in=["developer"]) + + await _request(guardrail, _chat({"role": ["developer"], "content": MARKER}), logging_obj=_logging_obj()) + + assert endpoint.input_types == ["request"] + + +@pytest.mark.asyncio +async def test_non_matching_system_prompt_leaves_both_sides_scanned(): + endpoint: Final = _Endpoint() + guardrail: Final = _guardrail(endpoint, skip_if_system_prompt_matches=[MARKER]) + logging_obj: Final = _logging_obj() + + await _request(guardrail, _chat({"role": "system", "content": "plain"}), logging_obj=logging_obj) + await _response(guardrail, logging_obj=logging_obj) + + assert endpoint.input_types == ["request", "response"] + + +@pytest.mark.asyncio +async def test_first_role_match_skips_request_and_paired_response(): + endpoint: Final = _Endpoint() + guardrail: Final = _guardrail(endpoint, skip_if_first_role_in=["developer"]) + skipped_call: Final = _logging_obj("call-skipped") + scanned_call: Final = _logging_obj("call-scanned") + + await _request( + guardrail, + _chat({"role": "developer", "content": "instructions"}, {"role": "user", "content": "hello"}), + logging_obj=skipped_call, + ) + await _response(guardrail, logging_obj=skipped_call) + await _request( + guardrail, + _chat({"role": "user", "content": "hello"}, {"role": "developer", "content": "instructions"}), + logging_obj=scanned_call, + ) + + assert endpoint.input_types == ["request"] + assert endpoint.received[0]["litellm_call_id"] == "call-scanned" + + +@pytest.mark.asyncio +async def test_skip_decision_does_not_leak_to_another_guardrail_with_the_same_name(): + skipping_endpoint: Final = _Endpoint() + other_endpoint: Final = _Endpoint() + skipping: Final = _guardrail(skipping_endpoint, name="gg", skip_if_system_prompt_matches=[MARKER]) + other: Final = _guardrail(other_endpoint, name="gg", skip_if_system_prompt_matches=["never-matches"]) + logging_obj: Final = _logging_obj() + + await _request(skipping, _system_marked_chat(), logging_obj=logging_obj) + await _request(other, _system_marked_chat(), logging_obj=logging_obj) + await _response(other, logging_obj=logging_obj) + await _response(skipping, logging_obj=logging_obj) + + assert other_endpoint.input_types == ["request", "response"] + assert skipping_endpoint.received == [] + + +@pytest.mark.asyncio +async def test_skip_decision_does_not_leak_to_another_guardrail(): + endpoint: Final = _Endpoint() + skipping: Final = _guardrail(endpoint, name="skipping", skip_if_system_prompt_matches=[MARKER]) + other: Final = _guardrail(endpoint, name="other", skip_if_system_prompt_matches=["something-else"]) + logging_obj: Final = _logging_obj() + + await _request(skipping, _system_marked_chat(), logging_obj=logging_obj) + await _response(other, logging_obj=logging_obj) + + assert endpoint.input_types == ["response"] + + +@pytest.mark.asyncio +async def test_defaults_scan_everything(): + endpoint: Final = _Endpoint() + guardrail: Final = _guardrail(endpoint) + logging_obj: Final = _logging_obj() + + await _request( + guardrail, + _chat({"role": "developer", "content": f"you are {MARKER}"}, {"role": "user", "content": "hello"}), + logging_obj=logging_obj, + ) + await _response(guardrail, logging_obj=logging_obj) + + assert endpoint.input_types == ["request", "response"] + + +def _chat_request(system_prompt: str) -> dict[str, object]: + return _chat( + {"role": "system", "content": system_prompt}, + {"role": "developer", "content": "be brief"}, + {"role": "user", "content": "look it up"}, + { + "role": "assistant", + "content": None, + "tool_calls": [{"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}}], + }, + {"role": "tool", "tool_call_id": "call_1", "content": "tool result"}, + ) + + +def _anthropic_request(system_prompt: str) -> dict[str, object]: + return { + "model": "claude-sonnet-4-5", + "max_tokens": 16, + "system": system_prompt, + "messages": [ + {"role": "user", "content": "look it up"}, + {"role": "assistant", "content": [{"type": "tool_use", "id": "toolu_1", "name": "lookup", "input": {}}]}, + {"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_1", "content": "tool result"}]}, + ], + } + + +def _responses_request(system_prompt: str) -> dict[str, object]: + return { + "model": "gpt-4o", + "instructions": system_prompt, + "input": [ + {"role": "developer", "content": "be brief"}, + {"role": "user", "content": "look it up"}, + {"type": "function_call", "call_id": "call_1", "name": "lookup", "arguments": "{}"}, + {"type": "function_call_output", "call_id": "call_1", "output": "tool result"}, + ], + } + + +REQUEST_SHAPES: Final = pytest.mark.parametrize( + ("translation", "build_request", "role_after_system_prompt"), + [ + pytest.param(OpenAIChatCompletionsHandler, _chat_request, "developer", id="chat"), + pytest.param(AnthropicMessagesHandler, _anthropic_request, "user", id="anthropic_messages"), + pytest.param(OpenAIResponsesHandler, _responses_request, "developer", id="responses"), + ], +) +SCOPING: Final = pytest.mark.parametrize( + "scoping", + [ + pytest.param({}, id="unscoped"), + pytest.param({"skip_system_message_in_guardrail": True}, id="skip_system_message"), + pytest.param({"scan_only_tool_results": True}, id="scan_only_tool_results"), + ], +) + + +async def _run_request_hook( + translation: type[BaseTranslation], + body: dict[str, object], + guardrail: GenericGuardrailAPI, + scoping: dict[str, bool], +) -> None: + _configure_callback_scoping( + guardrail, + guardrail.guardrail_name or "", + LitellmParams(guardrail="generic_guardrail_api", mode="pre_call", **scoping), + ) + await translation().process_input_messages( + data=body, guardrail_to_apply=guardrail, litellm_logging_obj=_logging_obj() + ) + + +@REQUEST_SHAPES +@SCOPING +@pytest.mark.asyncio +async def test_system_prompt_decision_ignores_guardrail_scoping( + translation: type[BaseTranslation], + build_request: Callable[[str], dict[str, object]], + role_after_system_prompt: str, + scoping: dict[str, bool], +): + unfiltered_endpoint: Final = _Endpoint() + filtered_endpoint: Final = _Endpoint() + marked_prompt: Final = f"you are {MARKER}" + + await _run_request_hook(translation, build_request(marked_prompt), _guardrail(unfiltered_endpoint), scoping) + await _run_request_hook( + translation, + build_request(marked_prompt), + _guardrail(filtered_endpoint, skip_if_system_prompt_matches=[MARKER]), + scoping, + ) + + assert unfiltered_endpoint.input_types == ["request"], "the handler must reach the guardrail for this request" + assert filtered_endpoint.received == [] + + +@REQUEST_SHAPES +@SCOPING +@pytest.mark.asyncio +async def test_first_role_decision_ignores_guardrail_scoping( + translation: type[BaseTranslation], + build_request: Callable[[str], dict[str, object]], + role_after_system_prompt: str, + scoping: dict[str, bool], +): + endpoint: Final = _Endpoint() + + await _run_request_hook( + translation, + build_request("you are helpful"), + _guardrail(endpoint, skip_if_first_role_in=[role_after_system_prompt]), + scoping, + ) + + assert endpoint.input_types == ["request"], "the system prompt leads the request, so its first role is system" + + +@pytest.mark.parametrize( + "options", + [ + {"skip_if_system_prompt_matches": ["(unclosed"]}, + {"skip_if_system_prompt_matches": MARKER}, + {"skip_if_first_role_in": "developer"}, + ], +) +def test_invalid_config_fails_at_init(options: dict[str, object]): + with pytest.raises(ValueError, match="skip_if_"): + _guardrail(_Endpoint(), **options) + + +@pytest.mark.parametrize( + ("option", "invalid_value"), + [("skip_if_system_prompt_matches", ("(unclosed",)), ("skip_if_first_role_in", "developer")], +) +@pytest.mark.parametrize("via_optional_params", [False, True]) +def test_config_options_reach_the_guardrail(option: str, invalid_value: object, via_optional_params: bool): + top_level: Final = {} if via_optional_params else {option: invalid_value} + params: Final = LitellmParams( + guardrail="generic_guardrail_api", mode="pre_call", api_base="https://guardrail.test", **top_level + ) + if via_optional_params: + params.optional_params = GenericGuardrailAPIOptionalParams.model_construct(**{option: invalid_value}) + with pytest.raises(ValueError, match=option): + initialize_guardrail(params, Guardrail(guardrail_name="gg-config", litellm_params=params)) + + +@pytest.mark.parametrize( + ("event_hook", "options", "expected"), + [ + ("pre_call", {}, set()), + ("post_call", {}, set()), + ("pre_call", {"skip_if_system_prompt_matches": [MARKER]}, {TRUST_BOUNDARY_WARNING}), + ("during_call", {"skip_if_first_role_in": ["developer"]}, {TRUST_BOUNDARY_WARNING}), + (["pre_call", "post_call"], {"skip_if_system_prompt_matches": [MARKER]}, {TRUST_BOUNDARY_WARNING}), + ( + "post_call", + {"skip_if_system_prompt_matches": [MARKER]}, + {TRUST_BOUNDARY_WARNING, NO_REQUEST_HOOK_WARNING}, + ), + ( + ["post_call"], + {"skip_if_first_role_in": ["developer"]}, + {TRUST_BOUNDARY_WARNING, NO_REQUEST_HOOK_WARNING}, + ), + ( + Mode(tags={"env": "post_call"}, default="post_call"), + {"skip_if_system_prompt_matches": [MARKER]}, + {TRUST_BOUNDARY_WARNING, NO_REQUEST_HOOK_WARNING}, + ), + ( + Mode(tags={"env": ["post_call", "pre_call"]}, default="post_call"), + {"skip_if_system_prompt_matches": [MARKER]}, + {TRUST_BOUNDARY_WARNING}, + ), + ( + Mode(tags={"env": "post_call"}, default=["during_call"]), + {"skip_if_system_prompt_matches": [MARKER]}, + {TRUST_BOUNDARY_WARNING}, + ), + ], +) +def test_init_warnings( + caplog: pytest.LogCaptureFixture, + event_hook: str | list[str] | Mode, + options: dict[str, object], + expected: set[str], +): + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + _guardrail(_Endpoint(), event_hook=event_hook, **options) + + warned: Final = { + warning + for warning in (TRUST_BOUNDARY_WARNING, NO_REQUEST_HOOK_WARNING) + if any(warning in record.getMessage() for record in caplog.records) + } + assert warned == expected + + +@pytest.mark.parametrize("event_hook", ["pre_mcp_call", ["post_call", "during_mcp_call"]]) +def test_mcp_request_hooks_count_as_request_side( + monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture, event_hook: str | list[str] +): + monkeypatch.setenv("LITELLM_STRICT_GUARDRAIL_MODES", "false") + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + _guardrail(_Endpoint(), event_hook=event_hook, skip_if_system_prompt_matches=[MARKER]) + + messages: Final = [record.getMessage() for record in caplog.records] + assert any(TRUST_BOUNDARY_WARNING in message for message in messages) + assert not any(NO_REQUEST_HOOK_WARNING in message for message in messages), messages