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:
lior-k 2026-06-17 16:36:17 +03:00
parent caa5b8c6b3
commit f234039d62
No known key found for this signature in database
2 changed files with 130 additions and 20 deletions

View file

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

View file

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