mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(guardrails): honor only_scan_new_messages on litellm_content_filter
The content filter accepted only_scan_new_messages on its config, validated it, and then ignored it: initialize_guardrail builds ContentFilterGuardrail from an explicit kwarg list that never included the parameter, so the constructor default of False always won. Every request re-scanned the whole conversation against every compiled pattern and keyword, which is O(history x rules) growth on long sessions with no way to opt out. Forward the parameter from the initializer and scan only the per-session diff on request-side scans, mirroring BedrockGuardrail's incremental path. The cache resolver moves from BedrockGuardrail up to CustomGuardrail so both guardrails share one implementation, as the note at the Bedrock call site anticipated. Two properties shape the implementation. inputs["texts"] is returned untouched on the incremental path because the chat-completions handler writes the returned texts back into the request positionally, so returning the scanned subset would overwrite earlier messages with later ones. And the feature disables itself at init, with a warning, when any rule carries a MASK action: skipping an already-seen segment under a MASK rule would forward it to the provider unmasked. Scanned state is recorded only after the guardrail allows the request, so a blocked turn is never marked and a retry is checked again. Default is unchanged (False), so existing deployments behave identically.
This commit is contained in:
parent
9a715df212
commit
d5578e7ae4
5 changed files with 359 additions and 22 deletions
|
|
@ -441,6 +441,23 @@ class CustomGuardrail(CustomLogger):
|
|||
def _scanned_texts_cache_key(self, session_id: str) -> str:
|
||||
return f"guardrail_scanned_texts:{self.guardrail_name}:{session_id}"
|
||||
|
||||
@staticmethod
|
||||
def _incremental_scan_cache() -> DualCache:
|
||||
"""Resolve the cache used to remember which segments a session already scanned.
|
||||
|
||||
Prefers the proxy's shared cache (``internal_usage_cache.dual_cache``), which is
|
||||
backed by Redis when the deployment configures it, so incremental state is shared
|
||||
across proxy instances. Falls back to a process-local ``DualCache`` singleton when
|
||||
the proxy is not running (e.g. unit tests), where sharing does not apply.
|
||||
"""
|
||||
try:
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj as _proxy_logging
|
||||
except Exception: # noqa: BLE001 # proxy not importable outside the server; use local fallback
|
||||
return dc
|
||||
if _proxy_logging is not None:
|
||||
return _proxy_logging.internal_usage_cache.dual_cache
|
||||
return dc
|
||||
|
||||
async def filter_new_texts_for_session(
|
||||
self,
|
||||
texts: list[str] | None,
|
||||
|
|
|
|||
|
|
@ -3042,25 +3042,6 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
masking_index += 1
|
||||
verbose_proxy_logger.debug("Applied masking to choice text content")
|
||||
|
||||
@staticmethod
|
||||
def _incremental_scan_cache() -> DualCache:
|
||||
"""Resolve the cache used to remember which segments a session already scanned.
|
||||
|
||||
Prefers the proxy's shared cache (``internal_usage_cache.dual_cache``), which is
|
||||
backed by Redis when the deployment configures it, so incremental state is shared
|
||||
across proxy instances. Falls back to a process-local ``DualCache`` singleton when
|
||||
the proxy is not running (e.g. unit tests), where sharing does not apply.
|
||||
"""
|
||||
from litellm.integrations.custom_guardrail import dc as fallback_cache
|
||||
|
||||
try:
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj as _proxy_logging
|
||||
except Exception: # noqa: BLE001 # proxy not importable outside the server; use local fallback
|
||||
return fallback_cache
|
||||
if _proxy_logging is not None:
|
||||
return _proxy_logging.internal_usage_cache.dual_cache
|
||||
return fallback_cache
|
||||
|
||||
def _bedrock_response_has_masked_output(self, response: BedrockGuardrailResponse) -> bool:
|
||||
"""Return True if the guardrail rewrote (masked/anonymized) any scanned text.
|
||||
|
||||
|
|
|
|||
|
|
@ -48,6 +48,7 @@ def initialize_guardrail(
|
|||
end_session_after_n_fails=getattr(litellm_params, "end_session_after_n_fails", None),
|
||||
on_violation=getattr(litellm_params, "on_violation", None),
|
||||
realtime_violation_message=getattr(litellm_params, "realtime_violation_message", None),
|
||||
only_scan_new_messages=litellm_params.only_scan_new_messages or False,
|
||||
)
|
||||
|
||||
litellm.logging_callback_manager.add_litellm_callback(content_filter_guardrail)
|
||||
|
|
|
|||
|
|
@ -280,6 +280,17 @@ 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.
|
||||
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 "
|
||||
"(skipped text cannot be masked); scanning the full context on every request.",
|
||||
self.guardrail_name,
|
||||
)
|
||||
self.only_scan_new_messages = False
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
"ContentFilterGuardrail initialized with %s patterns and %s blocked words",
|
||||
len(self.compiled_patterns),
|
||||
|
|
@ -332,6 +343,20 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
result.append(word)
|
||||
return result
|
||||
|
||||
def _has_mask_action(self) -> bool:
|
||||
"""Whether any configured rule rewrites text rather than blocking it."""
|
||||
return (
|
||||
any(entry["action"] == ContentFilterAction.MASK for entry in self.compiled_patterns)
|
||||
or any(action == ContentFilterAction.MASK for action, _ in self.blocked_words.values())
|
||||
or any(
|
||||
action == ContentFilterAction.MASK
|
||||
for _, _, action in (
|
||||
*self.category_keywords.values(),
|
||||
*self.always_block_category_keywords.values(),
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _category_config_view(cat_config: ContentFilterCategoryConfig) -> _CategoryConfigView:
|
||||
return {
|
||||
|
|
@ -1864,6 +1889,16 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
if filtered_arguments != arguments:
|
||||
self._set_tool_call_arguments(tool_call, filtered_arguments)
|
||||
|
||||
async def _filter_new_request_texts(self, texts: list[str], request_data: dict) -> list[str] | None:
|
||||
return await self.filter_new_texts_for_session(
|
||||
texts=texts, request_data=request_data, cache=self._incremental_scan_cache()
|
||||
)
|
||||
|
||||
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()
|
||||
)
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: "GenericGuardrailAPIInputs",
|
||||
|
|
@ -1902,11 +1937,20 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
# Process images if present
|
||||
await self._process_images(images, detections)
|
||||
|
||||
new_texts: Final = (
|
||||
await self._filter_new_request_texts(texts=texts, request_data=request_data)
|
||||
if self.only_scan_new_messages and input_type == "request"
|
||||
else None
|
||||
)
|
||||
texts_to_scan: Final = texts if new_texts is None else new_texts
|
||||
|
||||
# Process texts
|
||||
verbose_proxy_logger.debug("ContentFilterGuardrail: Applying guardrail to %s text(s)", len(texts))
|
||||
verbose_proxy_logger.debug(
|
||||
"ContentFilterGuardrail: Applying guardrail to %s of %s text(s)", len(texts_to_scan), len(texts)
|
||||
)
|
||||
|
||||
processed_texts: Final = []
|
||||
for text in texts:
|
||||
for text in texts_to_scan:
|
||||
# Competitor intent check first (optional; may refuse/reframe)
|
||||
if self._competitor_intent_checker and text:
|
||||
intent_result = self._competitor_intent_checker.run(text)
|
||||
|
|
@ -1916,7 +1960,13 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
processed_texts.append(filtered_text)
|
||||
|
||||
verbose_proxy_logger.debug("ContentFilterGuardrail: Guardrail applied successfully")
|
||||
inputs["texts"] = processed_texts
|
||||
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.
|
||||
await self._mark_request_texts_scanned(texts=texts, request_data=request_data)
|
||||
|
||||
self._scan_tool_call_arguments(inputs=inputs, detections=detections)
|
||||
|
||||
|
|
|
|||
|
|
@ -3068,3 +3068,291 @@ class TestContentFilterToolCallArguments:
|
|||
request_data={},
|
||||
input_type="response",
|
||||
)
|
||||
|
||||
|
||||
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.
|
||||
"""
|
||||
|
||||
BLOCKED_KEYWORD = "hunter2"
|
||||
|
||||
def _guardrail(self, only_scan_new_messages: bool = True) -> ContentFilterGuardrail:
|
||||
return ContentFilterGuardrail(
|
||||
guardrail_name="content-filter-incremental",
|
||||
blocked_words=[
|
||||
BlockedWord(keyword=self.BLOCKED_KEYWORD, action=ContentFilterAction.BLOCK),
|
||||
],
|
||||
default_on=True,
|
||||
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
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_second_turn_scans_only_new_texts(self):
|
||||
guardrail = self._guardrail()
|
||||
session = {"litellm_session_id": "cf-incremental-diff"}
|
||||
scanned = self._record_scanned_texts(guardrail)
|
||||
|
||||
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"
|
||||
|
||||
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"
|
||||
|
||||
@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.
|
||||
"""
|
||||
guardrail = self._guardrail()
|
||||
session = {"litellm_session_id": "cf-incremental-writeback"}
|
||||
|
||||
await guardrail.apply_guardrail(
|
||||
inputs={"texts": ["be helpful", "first question"]},
|
||||
request_data=session,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
history = ["be helpful", "first question", "first answer", "second question"]
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs={"texts": list(history)}, request_data=session, input_type="request"
|
||||
)
|
||||
assert result["texts"] == history
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_second_turn_through_handler_leaves_messages_intact(self):
|
||||
"""End-to-end form of the positional-writeback guard: the live request must survive."""
|
||||
from litellm.llms.openai.chat.guardrail_translation.handler import (
|
||||
OpenAIChatCompletionsHandler,
|
||||
)
|
||||
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
guardrail = self._guardrail()
|
||||
session = "cf-incremental-handler"
|
||||
first_turn = [
|
||||
{"role": "system", "content": "be helpful"},
|
||||
{"role": "user", "content": "first question"},
|
||||
]
|
||||
second_turn = first_turn + [
|
||||
{"role": "assistant", "content": "first answer"},
|
||||
{"role": "user", "content": "second question"},
|
||||
]
|
||||
|
||||
await handler.process_input_messages(
|
||||
data={"messages": [dict(m) for m in first_turn], "litellm_session_id": session},
|
||||
guardrail_to_apply=guardrail,
|
||||
)
|
||||
|
||||
data = {"messages": [dict(m) for m in second_turn], "litellm_session_id": session}
|
||||
result = await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert [m["content"] for m in result["messages"]] == [m["content"] for m in second_turn]
|
||||
assert [m["role"] for m in result["messages"]] == [m["role"] for m in second_turn]
|
||||
|
||||
@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.
|
||||
"""
|
||||
guardrail = self._guardrail()
|
||||
session = {"litellm_session_id": "cf-incremental-edited"}
|
||||
|
||||
await guardrail.apply_guardrail(
|
||||
inputs={"texts": ["be helpful", "first question"]},
|
||||
request_data=session,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await guardrail.apply_guardrail(
|
||||
inputs={"texts": ["be helpful", f"first question {self.BLOCKED_KEYWORD}", "second question"]},
|
||||
request_data=session,
|
||||
input_type="request",
|
||||
)
|
||||
assert exc.value.status_code == 400
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_blocked_turn_is_not_marked_scanned(self):
|
||||
"""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"
|
||||
)
|
||||
|
||||
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"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_session_id_scans_full_context(self):
|
||||
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
|
||||
|
||||
@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.
|
||||
"""
|
||||
guardrail = ContentFilterGuardrail(
|
||||
guardrail_name="content-filter-incremental-mask",
|
||||
patterns=[
|
||||
ContentFilterPattern(
|
||||
pattern_type="prebuilt",
|
||||
pattern_name="email",
|
||||
action=ContentFilterAction.MASK,
|
||||
)
|
||||
],
|
||||
only_scan_new_messages=True,
|
||||
)
|
||||
assert guardrail.only_scan_new_messages is False, "a MASK rule must switch the feature off at init"
|
||||
|
||||
session = {"litellm_session_id": "cf-incremental-mask"}
|
||||
for _ in range(2):
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs={"texts": ["mail me at victim@example.com"]},
|
||||
request_data=session,
|
||||
input_type="request",
|
||||
)
|
||||
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_word_mask = ContentFilterGuardrail(
|
||||
guardrail_name="cf-mask-blocked-word",
|
||||
blocked_words=[BlockedWord(keyword="acme", action=ContentFilterAction.MASK)],
|
||||
only_scan_new_messages=True,
|
||||
)
|
||||
category_mask = ContentFilterGuardrail(
|
||||
guardrail_name="cf-mask-category",
|
||||
categories=[
|
||||
ContentFilterCategoryConfig(
|
||||
category="harm_toxic_abuse",
|
||||
enabled=True,
|
||||
action=ContentFilterAction.MASK,
|
||||
)
|
||||
],
|
||||
only_scan_new_messages=True,
|
||||
)
|
||||
|
||||
assert blocked_word_mask.only_scan_new_messages is False
|
||||
assert category_mask.only_scan_new_messages is False
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_response_scans_are_never_incremental(self):
|
||||
"""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"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_flag_off_scans_full_context(self):
|
||||
"""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
|
||||
|
||||
|
||||
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.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def _initialize(**litellm_params_kwargs) -> ContentFilterGuardrail:
|
||||
import litellm
|
||||
from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter import (
|
||||
initialize_guardrail,
|
||||
)
|
||||
from litellm.types.guardrails import LitellmParams
|
||||
|
||||
callbacks_snapshot = list(litellm.callbacks)
|
||||
try:
|
||||
return initialize_guardrail(
|
||||
litellm_params=LitellmParams(
|
||||
guardrail="litellm_content_filter",
|
||||
mode="pre_call",
|
||||
blocked_words=[BlockedWord(keyword="hunter2", action=ContentFilterAction.BLOCK)],
|
||||
**litellm_params_kwargs,
|
||||
),
|
||||
guardrail={"guardrail_name": "cf-initializer-propagation"},
|
||||
)
|
||||
finally:
|
||||
litellm.callbacks[:] = callbacks_snapshot
|
||||
|
||||
def test_configured_true_reaches_the_instance(self):
|
||||
assert self._initialize(only_scan_new_messages=True).only_scan_new_messages is True
|
||||
|
||||
def test_defaults_to_false(self):
|
||||
assert self._initialize().only_scan_new_messages is False
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue