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:
michelligabriele 2026-09-11 14:21:10 +02:00
parent 9a715df212
commit d5578e7ae4
No known key found for this signature in database
5 changed files with 359 additions and 22 deletions

View file

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

View file

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

View file

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

View file

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

View file

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