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.
This commit is contained in:
lior-k 2026-06-25 10:08:34 +03:00
parent 4eaea1dd8a
commit 4ad7cceb5d
No known key found for this signature in database
4 changed files with 168 additions and 16 deletions

View file

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

View file

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

View file

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

View file

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