mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge cd786dcba0 into 3930c5bab6
This commit is contained in:
commit
e50fbb6044
9 changed files with 881 additions and 1 deletions
|
|
@ -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``.
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
)
|
||||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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"}]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue