mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(guardrails): Alice WonderFence detects content split across chunk boundaries
Splitting an oversized segment into disjoint <=MAX_PROMPT_CHARS chunks let a blocked phrase straddle a boundary so neither chunk saw it whole. Multi-chunk segments now also evaluate a detection-only window spanning each boundary (last N + first N chars, N=CHUNK_OVERLAP_CHARS, clamped to max_chars/2), feeding BLOCK/DETECT so a phrase up to ~2N chars can't slip through the split. Masking still uses the disjoint chunks so the lossless rejoin holds; a boundary window that flags maskable content surfaces as DETECT since it can't be redacted across chunks. Single-chunk segments add no extra calls. Regression test: a phrase straddling the boundary blocks with overlap and evades with overlap=0.
This commit is contained in:
parent
caa5b8c6b3
commit
f234039d62
2 changed files with 130 additions and 20 deletions
|
|
@ -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))
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue