feat(guardrails): add system prompt and first role skip filters to generic_guardrail_api

Adds skip_if_system_prompt_matches and skip_if_first_role_in. A request whose
instructions match a configured regex, or whose first message role is listed,
is not sent to the guardrail endpoint. Instructions are system and developer
messages, a top-level Anthropic system prompt and Responses API instructions,
and only the first 16384 characters of each are searched

The decision reads the request body through a new iter_request_messages helper,
so skip_system_message_in_guardrail and scan_only_tool_results cannot change it.
The paired response is skipped too: the request side records the reason on the
call's logging object under a key unique to the guardrail instance, as a private
type that a request body key copied into model_call_details cannot forge, and
every response call replays it. Each skipped call records a not_run guardrail
entry with that reason

Patterns compile at init, so a bad regex or a bare string value fails at boot.
Init warns that these filters read the caller-controlled body, and warns when
the mode, including tag-based modes, has no request-side hook
This commit is contained in:
Caduri Katzav 2026-09-28 18:20:45 +03:00
parent 2c9b0e00ac
commit cd786dcba0
11 changed files with 881 additions and 1 deletions

View file

@ -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``.

View file

@ -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)

View file

@ -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

View file

@ -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

View file

@ -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}"
)

View file

@ -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],

View file

@ -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"}]

View file

View file

@ -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