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"