diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 77bf4820a1a..f4fd68b600b 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -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, diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 6a7ac4361b9..b69a223f7d1 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -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. diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/__init__.py index bbe3ded791d..55711073cd5 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py index 722f96ef814..8d3ebd93475 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py @@ -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) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py index be55ac47bde..822aaf70e15 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py @@ -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