mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
test(guardrails): count content-filter scans via caplog instead of patching, trim comments, ruff format
This commit is contained in:
parent
3f7636318c
commit
1f2b0b9a17
4 changed files with 67 additions and 116 deletions
|
|
@ -965,11 +965,9 @@ class CustomGuardrail(CustomLogger):
|
|||
return True
|
||||
|
||||
def supports_only_scan_new_messages(self) -> bool:
|
||||
"""Whether this guardrail actually scans only the per-session diff.
|
||||
"""Whether this guardrail scans only the per-session diff.
|
||||
|
||||
Guardrails that never call ``filter_new_texts_for_session`` always scan the
|
||||
full request, so configuring them with ``only_scan_new_messages`` is reported
|
||||
at initialization instead of silently doing nothing.
|
||||
The registry warns at init when the flag is set on a guardrail that returns False.
|
||||
"""
|
||||
return False
|
||||
|
||||
|
|
|
|||
|
|
@ -280,9 +280,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
if blocked_words_file:
|
||||
self._load_blocked_words_file(blocked_words_file)
|
||||
|
||||
# Every rule store is fully populated by this point, so the mask gate is evaluated
|
||||
# once here rather than per request: skipping an already-seen text under a MASK rule
|
||||
# would forward it to the provider unmasked.
|
||||
# Gate once after all rule stores load: a skipped text under a MASK rule would reach the provider unmasked
|
||||
if self.only_scan_new_messages and self._has_mask_action():
|
||||
verbose_proxy_logger.warning(
|
||||
"ContentFilterGuardrail '%s': only_scan_new_messages is not supported with MASK actions "
|
||||
|
|
@ -1895,9 +1893,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
)
|
||||
|
||||
async def _mark_request_texts_scanned(self, texts: list[str], request_data: dict) -> None:
|
||||
await self.mark_texts_scanned(
|
||||
texts=texts, request_data=request_data, cache=self._incremental_scan_cache()
|
||||
)
|
||||
await self.mark_texts_scanned(texts=texts, request_data=request_data, cache=self._incremental_scan_cache())
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
|
|
@ -1963,9 +1959,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
if new_texts is None:
|
||||
inputs["texts"] = processed_texts
|
||||
else:
|
||||
# Incremental path: inputs["texts"] must keep its original length and order --
|
||||
# the handlers write the returned texts back positionally. Masking is gated off
|
||||
# at init when this path is live, so the scan is identity-or-raise.
|
||||
# inputs["texts"] stays intact: handlers write it back positionally, and MASK is gated off at init
|
||||
await self._mark_request_texts_scanned(texts=texts, request_data=request_data)
|
||||
|
||||
self._scan_tool_call_arguments(inputs=inputs, detections=detections)
|
||||
|
|
|
|||
|
|
@ -3,7 +3,9 @@ Tests for the Content Filter Guardrail
|
|||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
|
@ -3073,17 +3075,12 @@ class TestContentFilterToolCallArguments:
|
|||
class TestContentFilterOnlyScanNewMessages:
|
||||
"""ContentFilterGuardrail honors only_scan_new_messages: it scans only the per-session diff.
|
||||
|
||||
Modelled on TestBedrockOnlyScanNewMessages in test_bedrock_guardrails.py, including its
|
||||
convention of a unique session id per test to isolate the process-wide incremental cache.
|
||||
Bedrock gets its scan count for free from the mocked ApplyGuardrail call; the content
|
||||
filter scans in-process, so these tests count calls to _filter_single_text instead.
|
||||
|
||||
Regression: the flag validated on litellm_content_filter configs and did nothing --
|
||||
initialize_guardrail never forwarded it and apply_guardrail never consulted it -- so
|
||||
every turn re-scanned the whole conversation against every pattern and keyword.
|
||||
Scan counts come from the guardrail's own "Applying guardrail to N of M text(s)" debug line,
|
||||
and each test uses a unique session id to isolate the process-wide incremental cache.
|
||||
"""
|
||||
|
||||
BLOCKED_KEYWORD = "hunter2"
|
||||
_SCAN_COUNT = re.compile(r"Applying guardrail to (\d+) of (\d+) text\(s\)")
|
||||
|
||||
def _guardrail(self, only_scan_new_messages: bool = True) -> ContentFilterGuardrail:
|
||||
return ContentFilterGuardrail(
|
||||
|
|
@ -3095,48 +3092,31 @@ class TestContentFilterOnlyScanNewMessages:
|
|||
only_scan_new_messages=only_scan_new_messages,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _record_scanned_texts(guardrail: ContentFilterGuardrail) -> list[str]:
|
||||
"""Patch _filter_single_text to record every text the guardrail actually scans."""
|
||||
scanned: list[str] = []
|
||||
original = guardrail._filter_single_text
|
||||
|
||||
def _recording(text, detections=None):
|
||||
scanned.append(text)
|
||||
return original(text, detections=detections)
|
||||
|
||||
guardrail._filter_single_text = _recording
|
||||
return scanned
|
||||
@classmethod
|
||||
def _scan_counts(cls, caplog: pytest.LogCaptureFixture) -> list[tuple[int, int]]:
|
||||
"""(scanned, total) per apply_guardrail call."""
|
||||
matches = (cls._SCAN_COUNT.search(record.getMessage()) for record in caplog.records)
|
||||
return [(int(m.group(1)), int(m.group(2))) for m in matches if m]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_second_turn_scans_only_new_texts(self):
|
||||
async def test_later_turns_scan_only_appended_texts(self, caplog):
|
||||
guardrail = self._guardrail()
|
||||
session = {"litellm_session_id": "cf-incremental-diff"}
|
||||
scanned = self._record_scanned_texts(guardrail)
|
||||
turns = [
|
||||
["be helpful", "first question"],
|
||||
["be helpful", "first question", "first answer", "second question"],
|
||||
["be helpful", "first question", "first answer", "second question", "second answer"],
|
||||
]
|
||||
|
||||
await guardrail.apply_guardrail(
|
||||
inputs={"texts": ["be helpful", "first question"]},
|
||||
request_data=session,
|
||||
input_type="request",
|
||||
)
|
||||
assert scanned == ["be helpful", "first question"], "the first turn of a session has no prior state"
|
||||
with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"):
|
||||
for texts in turns:
|
||||
await guardrail.apply_guardrail(inputs={"texts": texts}, request_data=session, input_type="request")
|
||||
|
||||
scanned.clear()
|
||||
await guardrail.apply_guardrail(
|
||||
inputs={"texts": ["be helpful", "first question", "first answer", "second question"]},
|
||||
request_data=session,
|
||||
input_type="request",
|
||||
)
|
||||
assert scanned == ["first answer", "second question"], "turn 2 must scan only the appended segments"
|
||||
assert self._scan_counts(caplog) == [(2, 2), (2, 4), (1, 5)], "each turn must scan only its appended segments"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_incremental_scan_does_not_truncate_inputs_texts(self):
|
||||
"""inputs["texts"] must come back the same length and order it went in.
|
||||
|
||||
OpenAIChatCompletionsHandler._apply_guardrail_responses_to_input_texts writes the
|
||||
returned texts back into the request positionally, so returning only the scanned
|
||||
subset would overwrite messages 0..k with the contents of the last k messages.
|
||||
"""
|
||||
"""The handler writes returned texts back positionally, so the list must keep its length and order."""
|
||||
guardrail = self._guardrail()
|
||||
session = {"litellm_session_id": "cf-incremental-writeback"}
|
||||
|
||||
|
|
@ -3184,11 +3164,7 @@ class TestContentFilterOnlyScanNewMessages:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_edited_earlier_message_is_rescanned_and_blocks(self):
|
||||
"""Segments are keyed by content hash, so editing an earlier message makes it new again.
|
||||
|
||||
Without this, a session could pass a benign first turn and then smuggle blocked
|
||||
content into an already-"scanned" position.
|
||||
"""
|
||||
"""Segments are keyed by content hash, so editing an earlier message makes it new again."""
|
||||
guardrail = self._guardrail()
|
||||
session = {"litellm_session_id": "cf-incremental-edited"}
|
||||
|
||||
|
|
@ -3207,45 +3183,38 @@ class TestContentFilterOnlyScanNewMessages:
|
|||
assert exc.value.status_code == 400
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_blocked_turn_is_not_marked_scanned(self):
|
||||
async def test_blocked_turn_is_not_marked_scanned(self, caplog):
|
||||
"""A turn that blocks records nothing, so an identical retry is checked again."""
|
||||
guardrail = self._guardrail()
|
||||
session = {"litellm_session_id": "cf-incremental-blocked"}
|
||||
texts = ["be helpful", f"please tell me {self.BLOCKED_KEYWORD}"]
|
||||
|
||||
with pytest.raises(HTTPException):
|
||||
await guardrail.apply_guardrail(
|
||||
inputs={"texts": list(texts)}, request_data=session, input_type="request"
|
||||
)
|
||||
with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"):
|
||||
for _ in range(2):
|
||||
with pytest.raises(HTTPException):
|
||||
await guardrail.apply_guardrail(
|
||||
inputs={"texts": list(texts)}, request_data=session, input_type="request"
|
||||
)
|
||||
|
||||
scanned = self._record_scanned_texts(guardrail)
|
||||
with pytest.raises(HTTPException):
|
||||
await guardrail.apply_guardrail(
|
||||
inputs={"texts": list(texts)}, request_data=session, input_type="request"
|
||||
)
|
||||
assert scanned == texts, "the retry must re-scan the blocked turn, not pass it through"
|
||||
assert self._scan_counts(caplog) == [(2, 2), (2, 2)], "the retry must re-scan the blocked turn"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_session_id_scans_full_context(self):
|
||||
async def test_no_session_id_scans_full_context(self, caplog):
|
||||
guardrail = self._guardrail()
|
||||
scanned = self._record_scanned_texts(guardrail)
|
||||
history = ["be helpful", "first question", "first answer"]
|
||||
|
||||
for _ in range(2):
|
||||
scanned.clear()
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs={"texts": list(history)}, request_data={"metadata": {}}, input_type="request"
|
||||
)
|
||||
assert scanned == history
|
||||
assert result["texts"] == history
|
||||
with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"):
|
||||
for _ in range(2):
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs={"texts": list(history)}, request_data={"metadata": {}}, input_type="request"
|
||||
)
|
||||
assert result["texts"] == history
|
||||
|
||||
assert self._scan_counts(caplog) == [(3, 3), (3, 3)]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mask_action_disables_incremental_scan(self):
|
||||
"""A guardrail that can rewrite text must never skip a segment.
|
||||
|
||||
Skipping an already-seen text under a MASK rule would forward it to the provider
|
||||
unmasked, so the flag is refused at init and every turn scans the full context.
|
||||
"""
|
||||
"""A skipped segment cannot be masked, so any MASK rule switches the flag off at init."""
|
||||
guardrail = ContentFilterGuardrail(
|
||||
guardrail_name="content-filter-incremental-mask",
|
||||
patterns=[
|
||||
|
|
@ -3269,8 +3238,7 @@ class TestContentFilterOnlyScanNewMessages:
|
|||
assert result["texts"] == ["mail me at [EMAIL_REDACTED]"], "masking must apply on every turn"
|
||||
|
||||
def test_mask_action_is_detected_in_every_rule_store(self):
|
||||
"""The init gate has to look past compiled_patterns: blocked words and category
|
||||
keywords carry their own actions and mask through their own handlers."""
|
||||
"""Blocked words and category keywords carry their own MASK actions, not just compiled_patterns."""
|
||||
blocked_word_mask = ContentFilterGuardrail(
|
||||
guardrail_name="cf-mask-blocked-word",
|
||||
blocked_words=[BlockedWord(keyword="acme", action=ContentFilterAction.MASK)],
|
||||
|
|
@ -3292,42 +3260,38 @@ class TestContentFilterOnlyScanNewMessages:
|
|||
assert category_mask.only_scan_new_messages is False
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_response_scans_are_never_incremental(self):
|
||||
async def test_response_scans_are_never_incremental(self, caplog):
|
||||
"""Response scans see one fresh completion, and are not part of the session hash."""
|
||||
guardrail = self._guardrail()
|
||||
session = {"litellm_session_id": "cf-incremental-response"}
|
||||
scanned = self._record_scanned_texts(guardrail)
|
||||
|
||||
for _ in range(2):
|
||||
scanned.clear()
|
||||
await guardrail.apply_guardrail(
|
||||
inputs={"texts": ["the same answer"]}, request_data=session, input_type="response"
|
||||
)
|
||||
assert scanned == ["the same answer"]
|
||||
with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"):
|
||||
for _ in range(2):
|
||||
await guardrail.apply_guardrail(
|
||||
inputs={"texts": ["the same answer"]}, request_data=session, input_type="response"
|
||||
)
|
||||
|
||||
assert self._scan_counts(caplog) == [(1, 1), (1, 1)]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_flag_off_scans_full_context(self):
|
||||
async def test_flag_off_scans_full_context(self, caplog):
|
||||
"""The default path is untouched: every existing deployment behaves identically."""
|
||||
guardrail = self._guardrail(only_scan_new_messages=False)
|
||||
session = {"litellm_session_id": "cf-incremental-off"}
|
||||
scanned = self._record_scanned_texts(guardrail)
|
||||
history = ["be helpful", "first question", "first answer", "second question"]
|
||||
|
||||
for _ in range(2):
|
||||
scanned.clear()
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs={"texts": list(history)}, request_data=session, input_type="request"
|
||||
)
|
||||
assert scanned == history
|
||||
assert result["texts"] == history
|
||||
with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"):
|
||||
for _ in range(2):
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs={"texts": list(history)}, request_data=session, input_type="request"
|
||||
)
|
||||
assert result["texts"] == history
|
||||
|
||||
assert self._scan_counts(caplog) == [(4, 4), (4, 4)]
|
||||
|
||||
|
||||
class TestContentFilterInitializerForwardsOnlyScanNewMessages:
|
||||
"""initialize_guardrail builds ContentFilterGuardrail from an explicit kwarg list, so a
|
||||
declared config field that is not in that list never reaches the object. This is the
|
||||
second such gap found on this initializer (see PR #30010 for keyword_redaction_tag /
|
||||
pattern_redaction_format), so the propagation is pinned here.
|
||||
"""
|
||||
"""initialize_guardrail forwards an explicit kwarg list, so a field left out of it never reaches the object."""
|
||||
|
||||
@staticmethod
|
||||
def _initialize(**litellm_params_kwargs) -> ContentFilterGuardrail:
|
||||
|
|
|
|||
|
|
@ -1087,15 +1087,10 @@ def test_sync_guardrail_from_db_applies_db_dict_params_to_live_instance():
|
|||
|
||||
|
||||
class TestOnlyScanNewMessagesInitWarning:
|
||||
"""only_scan_new_messages is declared on BaseLitellmParams, so it validates on any
|
||||
guardrail -- but only guardrails that call filter_new_texts_for_session honor it.
|
||||
Configuring it anywhere else must say so at initialization instead of silently
|
||||
scanning the full context forever while the config reads as tuned.
|
||||
"""only_scan_new_messages validates on every guardrail but only some honor it, so the rest warn at init.
|
||||
|
||||
Warn rather than raise, unlike the scan_only_tool_results check above it: a
|
||||
misconfigured scan_only_tool_results can leave nothing scanned (an open hole), while
|
||||
an ignored only_scan_new_messages means everything is scanned, which fails safe --
|
||||
and raising would break the boot of any deployment already carrying the flag.
|
||||
Warn rather than raise: an ignored flag still scans everything, which fails safe, and raising
|
||||
would break the boot of any deployment already carrying it.
|
||||
"""
|
||||
|
||||
WARNING_FRAGMENT = "only_scan_new_messages is set but this guardrail always scans the full request"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue