From 4ad7cceb5d88d1f32c1d8112b68b35cee75b643d Mon Sep 17 00:00:00 2001 From: lior-k Date: Thu, 25 Jun 2026 10:08:34 +0300 Subject: [PATCH] fix(guardrails): Alice WonderFence detects content split across adjacent segments The chat translation layer emits each message content part as its own texts entry, but the model concatenates adjacent parts (a multimodal message's text parts join with no separator), so a blocked phrase split across two segments was seen whole by neither per-segment scan and slipped through. Extend the existing detection-only overlap approach with a window spanning each adjacent prompt-text junction, bounded to the ordered prompt texts via WindowConfig.text_segment_count so tool-call args and tool/function definitions are not falsely joined. The tuning knobs move into a frozen WindowConfig to keep evaluate_segments within the argument-count budget. --- .../alice_wonderfence/alice_wonderfence.py | 7 +- .../alice_wonderfence/chunked_evaluation.py | 78 +++++++++++++-- .../alice_wonderfence/test_apply_guardrail.py | 5 +- .../test_chunked_evaluation.py | 94 ++++++++++++++++++- 4 files changed, 168 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py index 7305680212e..166c5bfe33a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py @@ -21,7 +21,11 @@ from litellm.types.proxy.guardrails.guardrail_hooks.alice_wonderfence import ( ) from litellm.types.utils import GenericGuardrailAPIInputs -from .chunked_evaluation import DEFAULT_MAX_CONCURRENCY, evaluate_segments +from .chunked_evaluation import ( + DEFAULT_MAX_CONCURRENCY, + WindowConfig, + evaluate_segments, +) from .client_cache import ClientBuildSpec, get_or_create_client, load_sdk from .credentials import CredentialConfig, resolve_credentials from .exceptions import WonderFenceBlockedError, WonderFenceMissingSecrets @@ -252,6 +256,7 @@ class WonderFenceGuardrail(CustomGuardrail): segments, evaluate, max_concurrency=self._connection_pool_limit or DEFAULT_MAX_CONCURRENCY, + windows=WindowConfig(text_segment_count=len(texts)), ) n_text = len(texts) n_tool = len(tool_segments) diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/chunked_evaluation.py b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/chunked_evaluation.py index 23f0d26c824..cbe8b1ee140 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/chunked_evaluation.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/chunked_evaluation.py @@ -8,9 +8,9 @@ target a different backend. import asyncio import re +from collections.abc import Awaitable, Callable from dataclasses import dataclass from typing import Any -from collections.abc import Awaitable, Callable MAX_PROMPT_CHARS = 10000 # WonderFence server-side prompt limit DEFAULT_MAX_CONCURRENCY = 10 # used when the client connection_pool_limit is unset @@ -32,6 +32,19 @@ class SegmentVerdict: correlation_ids: list[str] +@dataclass(frozen=True) +class WindowConfig: + """Tuning for the detection-only overlap windows. + + ``overlap`` sizes the chunk- and segment-boundary windows; ``text_segment_count`` + is how many leading segments are ordered prompt texts the model concatenates, + bounding the cross-segment windows (see ``_cross_segment_windows``). + """ + + overlap: int = CHUNK_OVERLAP_CHARS + text_segment_count: int = 0 + + def _split_text(text: str, max_chars: int) -> list[str]: """Split ``text`` into <= ``max_chars`` chunks with ``"".join(chunks) == text``. @@ -80,6 +93,32 @@ def _boundary_windows(chunks: list[str], overlap: int) -> list[str]: ] +def _cross_segment_windows( + segments: list[str], text_segment_count: int, overlap: int +) -> list[tuple[int, str]]: + """Detection-only windows spanning each adjacent pair of prompt-text segments. + + The chat translation layer emits each message content part as its own + ``texts`` entry, but the model concatenates them (a multimodal message's text + parts join with no separator at all), so a blocked phrase split across two + adjacent segments is seen whole by neither. We also scan a window joining the + tail of one to the head of the next. Only the first ``text_segment_count`` + segments (the ordered prompt texts) are paired; tool-call args and tool / + function definitions are not concatenated into the prompt. Each window is + tagged with its left segment index so a BLOCK/DETECT folds into that + segment's verdict; windows never mask, since content cannot be redacted + across a segment boundary. + """ + if overlap <= 0: + return [] + n = min(text_segment_count, len(segments)) + return [ + (i, segments[i][-overlap:] + segments[i + 1][:overlap]) + for i in range(n - 1) + if segments[i] and segments[i + 1] + ] + + def _aggregate( chunks: list[str], chunk_results: list[Any], @@ -116,13 +155,17 @@ async def evaluate_segments( evaluate: Callable[[str], Awaitable[Any]], max_chars: int = MAX_PROMPT_CHARS, max_concurrency: int = DEFAULT_MAX_CONCURRENCY, - overlap: int = CHUNK_OVERLAP_CHARS, + windows: WindowConfig = WindowConfig(), ) -> list[SegmentVerdict]: """Evaluate every segment (chunked) in parallel; return one verdict per segment. Each segment is split into <= ``max_chars`` disjoint chunks; multi-chunk segments also get a detection-only window spanning each chunk boundary (see - ``_boundary_windows``). Every chunk and window across every segment is + ``_boundary_windows``). Adjacent prompt-text segments (the first + ``windows.text_segment_count``) additionally get a detection-only window + spanning their junction (see ``_cross_segment_windows``) so a phrase split + across two segments is still seen whole. Every chunk and window across every + segment is evaluated through a single ``asyncio.gather`` behind one shared ``Semaphore(max_concurrency)``. Results are grouped back per segment with action precedence BLOCK > MASK > DETECT > NO_ACTION; masking uses the @@ -135,27 +178,42 @@ async def evaluate_segments( return await evaluate(text) # Keep boundary windows within the prompt limit (<= 2*ov <= max_chars). - ov = min(overlap, max_chars // 2) + ov = min(windows.overlap, max_chars // 2) seg_chunks = [_split_text(s, max_chars) for s in segments] seg_boundaries = [_boundary_windows(chunks, ov) for chunks in seg_chunks] + cross_windows = _cross_segment_windows(segments, windows.text_segment_count, ov) - index: list[tuple] = [] + index: list[tuple[str, int, int]] = [] tasks = [] for si in range(len(segments)): for ci, chunk in enumerate(seg_chunks[si]): - index.append((si, False, ci)) + index.append(("chunk", si, ci)) tasks.append(run(chunk)) for bi, window in enumerate(seg_boundaries[si]): - index.append((si, True, bi)) + index.append(("bound", si, bi)) tasks.append(run(window)) + for left_idx, window in cross_windows: + index.append(("cross", left_idx, 0)) + tasks.append(run(window)) results = await asyncio.gather(*tasks) chunk_res: list[list[Any]] = [[None] * len(c) for c in seg_chunks] bound_res: list[list[Any]] = [[None] * len(b) for b in seg_boundaries] - for (si, is_boundary, idx), res in zip(index, results): - (bound_res if is_boundary else chunk_res)[si][idx] = res + for (kind, si, idx), res in zip(index, results): + if kind == "chunk": + chunk_res[si][idx] = res + elif kind == "bound": + bound_res[si][idx] = res + cross_res: list[list[Any]] = [ + [ + res + for (kind, si, _), res in zip(index, results) + if kind == "cross" and si == s + ] + for s in range(len(segments)) + ] return [ - _aggregate(seg_chunks[si], chunk_res[si], bound_res[si]) + _aggregate(seg_chunks[si], chunk_res[si], bound_res[si] + cross_res[si]) for si in range(len(segments)) ] diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail.py index 28375482d91..7a12684eccd 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail.py @@ -463,9 +463,10 @@ async def test_apply_guardrail_evaluates_every_text_without_structured_messages( request_data=make_request_data(), input_type="request", ) - assert client.evaluate_prompt.call_count == 3 prompts = {c.kwargs["prompt"] for c in client.evaluate_prompt.call_args_list} - assert prompts == {"t1", "t2", "t3"} + assert {"t1", "t2", "t3"} <= prompts + # Adjacent text segments also get a cross-segment junction window each. + assert {"t1t2", "t2t3"} <= prompts @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_chunked_evaluation.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_chunked_evaluation.py index 9357398c1f5..4e458980e58 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_chunked_evaluation.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_chunked_evaluation.py @@ -8,6 +8,7 @@ import pytest from litellm.proxy.guardrails.guardrail_hooks.alice_wonderfence.chunked_evaluation import ( MAX_PROMPT_CHARS, SegmentVerdict, + WindowConfig, _split_text, evaluate_segments, ) @@ -186,7 +187,9 @@ async def test_block_phrase_split_across_chunk_boundary_is_detected(): async def evaluate(text): return _result("BLOCK" if "BLOCK ME" in text else "") - verdicts = await evaluate_segments([segment], evaluate, max_chars=12, overlap=6) + verdicts = await evaluate_segments( + [segment], evaluate, max_chars=12, windows=WindowConfig(overlap=6) + ) assert verdicts[0].action == "BLOCK" @@ -199,7 +202,9 @@ async def test_no_overlap_window_lets_boundary_phrase_evade(): async def evaluate(text): return _result("BLOCK" if "BLOCK ME" in text else "") - verdicts = await evaluate_segments([segment], evaluate, max_chars=12, overlap=0) + verdicts = await evaluate_segments( + [segment], evaluate, max_chars=12, windows=WindowConfig(overlap=0) + ) assert verdicts[0].action == "" @@ -229,5 +234,88 @@ async def test_boundary_window_mask_is_surfaced_as_detect_not_dropped(): _result("MASK", action_text="[X]") if "SECRET HERE" in text else _result("") ) - verdicts = await evaluate_segments([segment], evaluate, max_chars=12, overlap=6) + verdicts = await evaluate_segments( + [segment], evaluate, max_chars=12, windows=WindowConfig(overlap=6) + ) assert verdicts[0].action == "DETECT" + + +# ----------------------------- cross-segment overlap (split across adjacent texts) ----------------------------- + + +@pytest.mark.asyncio +async def test_block_phrase_split_across_adjacent_text_segments_is_detected(): + """A blocked phrase split across two adjacent prompt-text segments (e.g. two + content parts of one message, which the model concatenates) is caught by the + cross-segment window even though neither segment contains it whole. Fails + without cross-segment windows -> the phrase evades scanning.""" + + async def evaluate(text): + return _result("BLOCK" if "BLOCKME" in text else "") + + verdicts = await evaluate_segments( + ["BLOCK", "ME"], evaluate, windows=WindowConfig(text_segment_count=2) + ) + assert verdicts[0].action == "BLOCK" + + +@pytest.mark.asyncio +async def test_without_text_segment_count_split_phrase_evades(): + """Control: with no declared text segments there is no cross-segment window, + so the same split phrase is seen by neither segment. Demonstrates the gap the + cross-segment window closes.""" + + async def evaluate(text): + return _result("BLOCK" if "BLOCKME" in text else "") + + verdicts = await evaluate_segments(["BLOCK", "ME"], evaluate) + assert [v.action for v in verdicts] == ["", ""] + + +@pytest.mark.asyncio +async def test_cross_segment_window_stays_within_text_segments(): + """Only the first text_segment_count segments are paired; a trailing + non-text segment (tool-call args, tool/function definition) is never joined + with the last prompt text, so a phrase straddling that junction does not + block.""" + + async def evaluate(text): + return _result("BLOCK" if "BLOCKME" in text else "") + + verdicts = await evaluate_segments( + ["BLOCK", "ME"], evaluate, windows=WindowConfig(text_segment_count=1) + ) + assert [v.action for v in verdicts] == ["", ""] + + +@pytest.mark.asyncio +async def test_cross_segment_window_surfaces_mask_as_detect_without_masking(): + """A cross-segment window cannot redact across the segment boundary, so a + MASK on it surfaces as DETECT and never rewrites the segment text.""" + + async def evaluate(text): + return ( + _result("MASK", action_text="[X]") if "SECRETHERE" in text else _result("") + ) + + verdicts = await evaluate_segments( + ["SECRET", "HERE"], evaluate, windows=WindowConfig(text_segment_count=2) + ) + assert verdicts[0].action == "DETECT" + assert verdicts[0].masked_text is None + + +@pytest.mark.asyncio +async def test_cross_segment_window_joins_segment_tail_and_head(): + """The window spans the junction (tail of one segment + head of the next), + catching a phrase that lives only across the boundary of longer segments.""" + + async def evaluate(text): + return _result("BLOCK" if "a bomb" in text else "") + + verdicts = await evaluate_segments( + ["how to make a b", "omb please"], + evaluate, + windows=WindowConfig(overlap=6, text_segment_count=2), + ) + assert verdicts[0].action == "BLOCK"