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 117700bc8a2..d6ad0dcace9 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/chunked_evaluation.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/chunked_evaluation.py @@ -13,6 +13,14 @@ from typing import Any, Awaitable, Callable, List, Optional MAX_PROMPT_CHARS = 10000 # WonderFence server-side prompt limit DEFAULT_MAX_CONCURRENCY = 10 # used when the client connection_pool_limit is unset +# Detection-only overlap: when a segment is split into multiple chunks, content +# straddling a chunk boundary would be seen whole by neither chunk. We also +# evaluate a window spanning each boundary (last N chars of one chunk + first N +# of the next) so a blocked phrase up to ~2N chars long can't slip through the +# split. These windows feed BLOCK/DETECT only; masking still uses the disjoint +# chunks so the lossless rejoin invariant holds. Confirm sizing with the +# WonderFence team alongside MAX_PROMPT_CHARS. +CHUNK_OVERLAP_CHARS = 512 @dataclass @@ -57,25 +65,47 @@ def _action_str(result: Any) -> str: return action.value if hasattr(action, "value") else (action or "") -def _aggregate(chunks: List[str], results: List[Any]) -> SegmentVerdict: - actions = [_action_str(r) for r in results] +def _boundary_windows(chunks: List[str], overlap: int) -> List[str]: + """Windows spanning each adjacent chunk boundary, for detection only. + + Each window is the last ``overlap`` chars of one chunk joined to the first + ``overlap`` chars of the next, so a phrase split across the boundary is seen + whole by the window (up to ~2*overlap long). Empty when there is one chunk. + """ + if overlap <= 0: + return [] + return [ + chunks[i][-overlap:] + chunks[i + 1][:overlap] for i in range(len(chunks) - 1) + ] + + +def _aggregate( + chunks: List[str], + chunk_results: List[Any], + boundary_results: List[Any], +) -> SegmentVerdict: + chunk_actions = [_action_str(r) for r in chunk_results] + boundary_actions = [_action_str(r) for r in boundary_results] detections: list = [] correlation_ids: List[str] = [] - for r in results: + for r in (*chunk_results, *boundary_results): detections.extend(getattr(r, "detections", None) or []) cid = getattr(r, "correlation_id", None) if cid: correlation_ids.append(cid) - if "BLOCK" in actions: + if "BLOCK" in chunk_actions or "BLOCK" in boundary_actions: return SegmentVerdict("BLOCK", None, detections, correlation_ids) - if "MASK" in actions: + if "MASK" in chunk_actions: masked = "".join( (r.action_text or "[MASKED]") if _action_str(r) == "MASK" else chunk - for chunk, r in zip(chunks, results) + for chunk, r in zip(chunks, chunk_results) ) return SegmentVerdict("MASK", masked, detections, correlation_ids) - if "DETECT" in actions: + # A boundary window can only flag content that straddles a chunk split; we + # cannot redact it across disjoint chunks, so surface it as DETECT rather + # than dropping it. Per-chunk DETECT is folded in here too. + if "DETECT" in chunk_actions or {"MASK", "DETECT"} & set(boundary_actions): return SegmentVerdict("DETECT", None, detections, correlation_ids) return SegmentVerdict("", None, detections, correlation_ids) @@ -85,29 +115,46 @@ 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, ) -> List[SegmentVerdict]: """Evaluate every segment (chunked) in parallel; return one verdict per segment. - Each segment is split into <= ``max_chars`` chunks; every chunk across every - segment is evaluated through a single ``asyncio.gather`` behind one shared + 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 + 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. + action precedence BLOCK > MASK > DETECT > NO_ACTION; masking uses the + disjoint chunks only so the lossless rejoin holds. """ semaphore = asyncio.Semaphore(max_concurrency) - async def run(chunk: str) -> Any: + async def run(text: str) -> Any: async with semaphore: - return await evaluate(chunk) + return await evaluate(text) + # Keep boundary windows within the prompt limit (<= 2*ov <= max_chars). + ov = min(overlap, max_chars // 2) seg_chunks = [_split_text(s, max_chars) for s in segments] - flat_index = [ - (si, ci) for si, chunks in enumerate(seg_chunks) for ci in range(len(chunks)) - ] - tasks = [run(seg_chunks[si][ci]) for si, ci in flat_index] + seg_boundaries = [_boundary_windows(chunks, ov) for chunks in seg_chunks] + + index: List[tuple] = [] + tasks = [] + for si in range(len(segments)): + for ci, chunk in enumerate(seg_chunks[si]): + index.append((si, False, ci)) + tasks.append(run(chunk)) + for bi, window in enumerate(seg_boundaries[si]): + index.append((si, True, bi)) + tasks.append(run(window)) results = await asyncio.gather(*tasks) - per_segment: List[List[Any]] = [[None] * len(chunks) for chunks in seg_chunks] - for (si, ci), res in zip(flat_index, results): - per_segment[si][ci] = res + 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 - return [_aggregate(seg_chunks[si], per_segment[si]) for si in range(len(segments))] + return [ + _aggregate(seg_chunks[si], chunk_res[si], bound_res[si]) + for si in range(len(segments)) + ] 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 631053b1969..9357398c1f5 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 @@ -168,3 +168,66 @@ async def test_evaluations_run_in_parallel_under_a_cap(): def test_max_prompt_chars_is_positive(): assert isinstance(MAX_PROMPT_CHARS, int) and MAX_PROMPT_CHARS > 0 + + +# ----------------------------- boundary overlap (detection across chunk splits) ----------------------------- + + +@pytest.mark.asyncio +async def test_block_phrase_split_across_chunk_boundary_is_detected(): + """A blocked phrase straddling the chunk boundary is caught by the overlap + window even though neither disjoint chunk contains it whole. Fails on the + pre-overlap implementation (no boundary windows -> phrase evades).""" + segment = "aaaaa BLOCK ME zzzzz" + chunks = _split_text(segment, 12) + assert len(chunks) > 1 + assert all("BLOCK ME" not in c for c in chunks) + + async def evaluate(text): + return _result("BLOCK" if "BLOCK ME" in text else "") + + verdicts = await evaluate_segments([segment], evaluate, max_chars=12, overlap=6) + assert verdicts[0].action == "BLOCK" + + +@pytest.mark.asyncio +async def test_no_overlap_window_lets_boundary_phrase_evade(): + """Control: with overlap disabled the same straddling phrase is not seen by + any disjoint chunk, demonstrating what the overlap window closes.""" + segment = "aaaaa BLOCK ME zzzzz" + + async def evaluate(text): + return _result("BLOCK" if "BLOCK ME" in text else "") + + verdicts = await evaluate_segments([segment], evaluate, max_chars=12, overlap=0) + assert verdicts[0].action == "" + + +@pytest.mark.asyncio +async def test_single_chunk_segment_evaluates_once_no_boundary_window(): + calls = [] + + async def evaluate(text): + calls.append(text) + return _result("") + + await evaluate_segments(["short benign text"], evaluate, max_chars=10000) + assert calls == ["short benign text"] + + +@pytest.mark.asyncio +async def test_boundary_window_mask_is_surfaced_as_detect_not_dropped(): + """A boundary window can flag content we cannot redact across disjoint + chunks; it must surface as DETECT rather than pass silently.""" + segment = "aaaaa SECRET HERE zzzzz" + chunks = _split_text(segment, 12) + assert len(chunks) > 1 + + async def evaluate(text): + # Only the boundary window sees the full "SECRET HERE". + return ( + _result("MASK", action_text="[X]") if "SECRET HERE" in text else _result("") + ) + + verdicts = await evaluate_segments([segment], evaluate, max_chars=12, overlap=6) + assert verdicts[0].action == "DETECT"