From 0605a197091a71bd617c031350b8b305fc4b3693 Mon Sep 17 00:00:00 2001 From: lior-k Date: Thu, 23 Jul 2026 18:49:18 +0300 Subject: [PATCH] feat(guardrails): overlap chunks to preempt Alice seam-mask leak + per-chunk linear reconstruction Preempts a seam-mask leak and reduces reconstruction from quadratic to linear in document size. Chunking is now overlapping: a segment over the prompt limit is split into disjoint owned regions, but each chunk is scanned with the last N chars of the previous owned region prepended as a read-only prefix. A phrase straddling an owned-region seam is therefore seen whole by one scan, so the separate boundary-window calls are gone (call volume on a large request drops from ~2*chunks-1 to ~chunks). When the service masks content that reaches into a chunk's prefix bytes (content straddling, or within N of, a seam), the masked text no longer starts with the verbatim prefix and we fail closed as BLOCK instead of the previous DETECT-and-forward, which silently let seam-straddling maskable content through un-redacted. When the prefix is intact it is stripped by its known length, so stitching needs no alignment. Mask reconstruction is now aligned per owned-region chunk (each <= the prompt limit) instead of over the whole joined document, so difflib runs on bounded inputs and only on chunks the service actually changed. Cost is O(document * chunk_size) -- linear in document size with the chunk size as the constant -- rather than quadratic in the document. SegmentVerdict carries the per-chunk (original, masked) pairs so the caller aligns one chunk at a time. RECONSTRUCT_MAX_CHARS remains only as a coarse backstop on total alignment work. --- .../alice_wonderfence/alice_wonderfence.py | 2 +- .../alice_wonderfence/chunked_evaluation.py | 151 ++++++++++-------- .../alice_wonderfence/processing.py | 121 ++++++++++---- .../test_chunked_evaluation.py | 67 ++++---- .../alice_wonderfence/test_processing.py | 49 ++++-- 5 files changed, 252 insertions(+), 138 deletions(-) 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 ba98195a028..19535cbff4d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py @@ -355,7 +355,7 @@ class WonderFenceGuardrail(CustomGuardrail): correlation_id = verdict.correlation_ids[0] if verdict.correlation_ids else None if verdict.action == "MASK": - recovered = reconstruct(pieces, verdict.masked_text or "") + recovered = reconstruct(pieces, verdict.masked_chunks) if recovered is None: logger.warning( "Alice WonderFence (apply_guardrail request): MASK reconstruction unavailable " 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 e4afeac96e8..d69f292edfb 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/chunked_evaluation.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/chunked_evaluation.py @@ -14,12 +14,13 @@ from typing import Any 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 +# Overlap: a segment longer than the prompt limit is split into disjoint "owned" +# regions, but each chunk is scanned with the last N chars of the previous owned +# region prepended as a read-only prefix. A phrase straddling an owned-region +# seam (up to N chars into the left region) is therefore seen whole by one scan, +# so it can BLOCK/DETECT and, when the service masks it, the prefix bytes change +# and we fail closed (see ``_aggregate``) rather than stitch a half-masked seam. +# This replaces the old separate boundary-window calls. Confirm sizing with the # WonderFence team alongside MAX_PROMPT_CHARS. CHUNK_OVERLAP_CHARS = 512 @@ -30,23 +31,33 @@ class SegmentVerdict: masked_text: str | None detections: list correlation_ids: list[str] + # Per owned-region ``(original, masked)`` pairs, set only on MASK. Lets the + # caller align masking back to sub-structure (e.g. joined message parts) one + # chunk at a time instead of over the whole document, so the alignment cost + # is bounded by the chunk size rather than quadratic in the segment length. + masked_chunks: list[tuple[str, str]] | None = None @dataclass(frozen=True) class WindowConfig: - """Tuning for the detection-only overlap windows. + """Tuning for the chunk-seam overlap. - ``overlap`` sizes the per-segment chunk-boundary windows (see - ``_boundary_windows``). There are no cross-segment windows: on the request - side message parts are concatenated into one joined document before - scanning (so their junctions are interior chunk seams, covered by - ``_boundary_windows``); on the response side each segment is an independent - choice or tool-call arg that the model never concatenates. + ``overlap`` sizes the read-only prefix each chunk carries from the previous + owned region (see ``_overlap_chunks``). There are no cross-segment windows: + on the request side message parts are concatenated into one joined document + before scanning (so their junctions are interior chunk seams); on the + response side each segment is an independent choice or tool-call arg that the + model never concatenates. """ overlap: int = CHUNK_OVERLAP_CHARS +# Shared default so the ``evaluate_segments`` signature has a plain-name default +# (no call in the argument default); safe to share since ``WindowConfig`` is frozen. +_DEFAULT_WINDOW_CONFIG = WindowConfig() + + def _split_text(text: str, max_chars: int) -> list[str]: """Split ``text`` into <= ``max_chars`` chunks with ``"".join(chunks) == text``. @@ -81,45 +92,59 @@ def _action_str(result: object) -> str: return action.value if hasattr(action, "value") else (action or "") -def _boundary_windows(chunks: list[str], overlap: int) -> list[str]: - """Windows spanning each adjacent chunk boundary, for detection only. +def _overlap_chunks(text: str, max_chars: int, overlap: int) -> list[tuple[str, str]]: + """Split ``text`` into overlapping scan chunks as ``(prefix, owned)`` pairs. - 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. + ``owned`` regions are disjoint and concatenate back to ``text`` (lossless); + ``prefix`` is the last ``overlap`` chars of the previous owned region (empty + for the first). The scan input for a chunk is ``prefix + owned``, giving + ``overlap`` chars of left-context so a phrase straddling the owned-region + seam is seen whole. Reassembly strips the verbatim prefix back off, so the + owned regions still rejoin losslessly. """ - if overlap <= 0: - return [] - return [chunks[i][-overlap:] + chunks[i + 1][:overlap] for i in range(len(chunks) - 1)] + owned = _split_text(text, max(1, max_chars - overlap)) + return [(owned[i - 1][-overlap:] if i and overlap > 0 else "", region) for i, region in enumerate(owned)] -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] +def _aggregate(chunks: list[tuple[str, str]], results: list[Any]) -> SegmentVerdict: + """Fold per-chunk results (scans of ``prefix + owned``) into one verdict. + + Precedence BLOCK > MASK > DETECT > NO_ACTION. A MASK whose masked text no + longer starts with its verbatim ``prefix`` means the redaction reached into + the prefix bytes -- i.e. content straddling (or sitting within ``overlap`` of) + the owned-region seam. That cannot be stitched back without double-counting + the overlap, so it fails closed (BLOCK) rather than leak the un-redacted half. + Otherwise the prefix is stripped by its known length (no alignment needed) + and the owned regions rejoin into the masked segment; the per-chunk + ``(original, masked)`` pairs are carried on the verdict for bounded caller-side + alignment. + """ + actions = [_action_str(r) for r in results] detections: list = [] correlation_ids: list[str] = [] - for r in (*chunk_results, *boundary_results): + for r in results: detections.extend(getattr(r, "detections", None) or []) cid = getattr(r, "correlation_id", None) if cid: correlation_ids.append(cid) - if "BLOCK" in chunk_actions or "BLOCK" in boundary_actions: + if "BLOCK" in actions: return SegmentVerdict("BLOCK", None, detections, correlation_ids) - 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, chunk_results) - ) - return SegmentVerdict("MASK", masked, detections, correlation_ids) - # 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): + + if "MASK" in actions: + masked_chunks: list[tuple[str, str]] = [] + for (prefix, owned), r in zip(chunks, results): + if _action_str(r) != "MASK": + masked_chunks.append((owned, owned)) + continue + masked = r.action_text if getattr(r, "action_text", None) is not None else prefix + "[MASKED]" + if not masked.startswith(prefix): + return SegmentVerdict("BLOCK", None, detections, correlation_ids) + masked_chunks.append((owned, masked[len(prefix) :])) + masked_text = "".join(m for _, m in masked_chunks) + return SegmentVerdict("MASK", masked_text, detections, correlation_ids, masked_chunks) + + if "DETECT" in actions: return SegmentVerdict("DETECT", None, detections, correlation_ids) return SegmentVerdict("", None, detections, correlation_ids) @@ -129,18 +154,18 @@ async def evaluate_segments( evaluate: Callable[[str], Awaitable[Any]], max_chars: int = MAX_PROMPT_CHARS, max_concurrency: int = DEFAULT_MAX_CONCURRENCY, - windows: WindowConfig = WindowConfig(), + windows: WindowConfig = _DEFAULT_WINDOW_CONFIG, ) -> 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``) so a phrase split across a chunk seam 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 disjoint chunks only so - the lossless rejoin holds. + Each segment is split into <= ``max_chars`` overlapping chunks (disjoint + ``owned`` regions each carrying an ``overlap``-char read-only prefix from the + previous region, see ``_overlap_chunks``) so a phrase straddling an + owned-region seam is seen whole by one scan without a separate boundary call. + Every chunk across every segment is evaluated through a single + ``asyncio.gather`` behind one shared ``Semaphore(max_concurrency)``. Results + are folded per segment (see ``_aggregate``) with precedence + BLOCK > MASK > DETECT > NO_ACTION. The request side passes a single joined document here (one segment) so the common case is one call; the response side passes one segment per choice / @@ -152,28 +177,20 @@ async def evaluate_segments( async with semaphore: return await evaluate(text) - # Keep boundary windows within the prompt limit (<= 2*ov <= max_chars). + # Keep each chunk's scan input (prefix + owned) within the prompt limit. 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] + seg_chunks = [_overlap_chunks(s, max_chars, ov) for s in segments] - index: list[tuple[str, int, int]] = [] + index: list[tuple[int, int]] = [] tasks = [] - for si in range(len(segments)): - for ci, chunk in enumerate(seg_chunks[si]): - index.append(("chunk", si, ci)) - tasks.append(run(chunk)) - for bi, window in enumerate(seg_boundaries[si]): - index.append(("bound", si, bi)) - tasks.append(run(window)) + for si, chunks in enumerate(seg_chunks): + for ci, (prefix, owned) in enumerate(chunks): + index.append((si, ci)) + tasks.append(run(prefix + owned)) 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 (kind, si, idx), res in zip(index, results): - if kind == "chunk": - chunk_res[si][idx] = res - elif kind == "bound": - bound_res[si][idx] = res + for (si, ci), res in zip(index, results): + chunk_res[si][ci] = res - return [_aggregate(seg_chunks[si], chunk_res[si], bound_res[si]) for si in range(len(segments))] + return [_aggregate(seg_chunks[si], chunk_res[si]) for si in range(len(segments))] diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py index b62e9ca51d1..bf9e974afc0 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py @@ -18,13 +18,14 @@ logger = verbose_proxy_logger.getChild("alice_wonderfence") JOINER = "\n" -# Upper bound on the document ``reconstruct`` will align. ``SequenceMatcher`` is -# O(n*m) worst case and runs synchronously on the event loop, so a large -# repetitive MASK-triggering prompt could otherwise wedge it. MASK on a document -# larger than this fails closed (block) rather than run the quadratic alignment; -# non-MASK requests of any size are unaffected. Two chunks' worth keeps the -# worst case sub-second while still covering ordinary multi-message chats. -RECONSTRUCT_MAX_CHARS = 2 * MAX_PROMPT_CHARS +# Upper bound on the document ``reconstruct`` will align. Alignment now runs +# per chunk (each <= MAX_PROMPT_CHARS) rather than over the whole document, so +# the cost is O(document / chunk * chunk^2) = O(document * chunk) -- linear in +# document size with the chunk size as the constant, instead of quadratic in the +# document. This bound is a coarse backstop on total alignment work; a MASK on a +# document larger than it fails closed (block). Non-MASK requests of any size are +# unaffected. +RECONSTRUCT_MAX_CHARS = 10 * MAX_PROMPT_CHARS def build_analysis_context( @@ -195,46 +196,106 @@ def _map_index(x: int, ops: Sequence[tuple[str, int, int, int, int]], masked_len return masked_len -def reconstruct(parts: list[str], masked: str) -> list[str] | None: - """Recover per-part masked text from the masked joined document. +_ChunkEntry = tuple[int, int, int, str, Sequence[tuple[str, int, int, int, int]] | None] + + +def _chunk_entries(masked_chunks: list[tuple[str, str]]) -> tuple[list[_ChunkEntry], str]: + """Build the per-chunk position map. + + One entry per owned region: ``(orig_start, orig_end, masked_start, + masked_owned, opcodes|None)`` with cumulative offsets in both original and + masked space. Opcodes are computed only for chunks the service actually + changed (each aligned over <= one chunk, so the alignment cost is bounded by + the chunk size); unchanged chunks map by a fixed offset with no alignment. + Returns the entries and the reassembled masked document. + """ + entries: list[_ChunkEntry] = [] + o_off = 0 + m_off = 0 + for original_owned, masked_owned in masked_chunks: + ops = ( + None + if original_owned == masked_owned + else SequenceMatcher(None, original_owned, masked_owned, autojunk=False).get_opcodes() + ) + entries.append((o_off, o_off + len(original_owned), m_off, masked_owned, ops)) + o_off += len(original_owned) + m_off += len(masked_owned) + return entries, "".join(m for _, m in masked_chunks) + + +def _map_pos(x: int, entries: list[_ChunkEntry], masked_len: int) -> int | None: + """Map original index ``x`` to its masked index via the owning chunk.""" + for o_start, o_end, m_start, masked_owned, ops in entries: + if o_start <= x < o_end: + local = x - o_start + if ops is None: + return m_start + local + r = _map_index(local, ops, len(masked_owned)) + return None if r is None else m_start + r + return masked_len + + +def _joiner_survives(j: int, entries: list[_ChunkEntry]) -> bool: + """Whether the ``JOINER`` at original index ``j`` survives the mask as an + unmodified ``\\n`` (so parts cannot merge).""" + for o_start, o_end, _m_start, masked_owned, ops in entries: + if o_start <= j < o_end: + if ops is None: + return True + local = j - o_start + return any( + tag == "equal" and i1 <= local < i2 and masked_owned[j1 + (local - i1)] == JOINER + for tag, i1, i2, j1, _j2 in ops + ) + return False + + +def reconstruct(parts: list[str], masked_chunks: list[tuple[str, str]] | None) -> list[str] | None: + """Recover per-part masked text from the per-chunk masked owned regions. ``parts`` were joined with ``JOINER`` (a plain ``"\\n"``) into the document - that was scanned; ``masked`` is the service's masked version of that same - document. We align original-vs-masked with ``difflib.SequenceMatcher`` (no - sentinel injected) and map each part's char range through the alignment. + that was scanned; ``masked_chunks`` is the list of ``(owned_original, + owned_masked)`` regions that concatenate back to that document and its masked + form. Alignment is done per chunk (see ``_chunk_entries``), so the cost is + bounded by the chunk size rather than quadratic in the whole document; each + part's char range is mapped through the owning chunk's alignment. - Fails closed (returns ``None``) when the structure is not recoverable: the - document exceeds ``RECONSTRUCT_MAX_CHARS`` (bounds the quadratic alignment - cost); any ``JOINER`` between parts does not survive the mask as an - unmodified ``\\n`` (a mask spanning a joiner would merge parts); or a part - boundary lands inside a changed block. Returns one masked string per input - part, in order; ``[]`` for no parts. Assumes masking is span substitution - that preserves the non-masked characters; if the service reflows whitespace - the joiner-survival check trips and we fail closed rather than misassign. + Fails closed (returns ``None``) when the structure is not recoverable: + ``masked_chunks`` is missing; the document exceeds ``RECONSTRUCT_MAX_CHARS``; + the owned regions do not concatenate back to the join (invariant guard); any + ``JOINER`` between parts does not survive as an unmodified ``\\n`` (a mask + spanning a joiner would merge parts); or a part boundary lands inside a + changed block. Returns one masked string per input part, in order; ``[]`` for + no parts. Assumes masking is span substitution that preserves the non-masked + characters; if the service reflows whitespace the joiner-survival check trips + and we fail closed rather than misassign. """ if not parts: return [] + if masked_chunks is None: + return None original = JOINER.join(parts) - if len(original) > RECONSTRUCT_MAX_CHARS or len(masked) > RECONSTRUCT_MAX_CHARS: + if len(original) > RECONSTRUCT_MAX_CHARS: return None + if "".join(o for o, _ in masked_chunks) != original: + return None + + entries, masked_doc = _chunk_entries(masked_chunks) + masked_len = len(masked_doc) + starts = [0, *accumulate(len(p) + len(JOINER) for p in parts)][: len(parts)] ranges = [(s, s + len(p)) for s, p in zip(starts, parts)] joiners = [end for (_s, end) in ranges[:-1]] - ops = SequenceMatcher(None, original, masked, autojunk=False).get_opcodes() - - joiner_survives = all( - any(tag == "equal" and i1 <= j < i2 and masked[j1 + (j - i1)] == JOINER for tag, i1, i2, j1, _j2 in ops) - for j in joiners - ) - if not joiner_survives: + if not all(_joiner_survives(j, entries) for j in joiners): return None - mapped = [(_map_index(s, ops, len(masked)), _map_index(e, ops, len(masked))) for s, e in ranges] + mapped = [(_map_pos(s, entries, masked_len), _map_pos(e, entries, masked_len)) for s, e in ranges] if any(ms is None or me is None or ms > me for ms, me in mapped): return None - return [masked[ms:me] for ms, me in mapped] + return [masked_doc[ms:me] for ms, me in mapped] def block_detail(blocked: list[SegmentVerdict], guardrail_name: str, block_message: str) -> dict: 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 5d126ddd2fb..951d7190084 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 @@ -87,31 +87,37 @@ async def test_block_in_non_first_chunk_blocks_whole_segment(): @pytest.mark.asyncio -async def test_mask_rejoins_per_chunk_action_text_into_full_segment(): - segment = ("ab " * 60).strip() - chunks = _split_text(segment, 50) - assert len(chunks) > 1 +async def test_mask_rejoins_masked_owned_regions_into_full_segment(): + """A multi-chunk segment where the service redacts a token in one chunk (a + real span substitution that preserves surrounding bytes) rejoins into the + fully masked segment. overlap=0 keeps the chunks disjoint for a clean check; + the per-chunk (original, masked) pairs are carried on the verdict.""" + segment = " ".join(f"w{i}" for i in range(40)) + " SECRET " + " ".join(f"v{i}" for i in range(40)) async def evaluate(text): - return _result("MASK", action_text=f"<{text}>") + return _result("MASK", action_text=text.replace("SECRET", "[X]")) if "SECRET" in text else _result("") - verdicts = await evaluate_segments([segment], evaluate, max_chars=50) + chunks = _split_text(segment, 20) + assert len(chunks) > 1 + verdicts = await evaluate_segments([segment], evaluate, max_chars=20, windows=WindowConfig(overlap=0)) assert verdicts[0].action == "MASK" - assert verdicts[0].masked_text == "".join(f"<{c}>" for c in chunks) + assert verdicts[0].masked_text == segment.replace("SECRET", "[X]") + assert verdicts[0].masked_chunks is not None + assert "".join(o for o, _ in verdicts[0].masked_chunks) == segment @pytest.mark.asyncio -async def test_unmasked_chunks_fall_back_to_original_text_on_rejoin(): +async def test_unmasked_chunks_keep_original_text_on_rejoin(): + """Chunks the service did not mask contribute their original owned text + verbatim; only the masked chunk changes.""" segment = " ".join(f"w{i}" for i in range(40)) - chunks = _split_text(segment, 20) - assert len(chunks) > 1 async def evaluate(text): - return _result("MASK" if text == chunks[0] else "", action_text="[X]") + return _result("MASK", action_text=text.replace("w0", "[X]")) if "w0 " in text else _result("") - verdicts = await evaluate_segments([segment], evaluate, max_chars=20) - expected = "[X]" + "".join(chunks[1:]) - assert verdicts[0].masked_text == expected + verdicts = await evaluate_segments([segment], evaluate, max_chars=20, windows=WindowConfig(overlap=0)) + assert verdicts[0].action == "MASK" + assert verdicts[0].masked_text == segment.replace("w0", "[X]", 1) @pytest.mark.asyncio @@ -169,14 +175,14 @@ def test_max_prompt_chars_is_positive(): assert isinstance(MAX_PROMPT_CHARS, int) and MAX_PROMPT_CHARS > 0 -# ----------------------------- boundary overlap (detection across chunk splits) ----------------------------- +# ----------------------------- seam overlap (detection / masking 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).""" +async def test_block_phrase_split_across_chunk_seam_is_detected(): + """A blocked phrase straddling an owned-region seam is caught because the + next chunk carries an overlap prefix from the previous owned region, so one + scan sees the phrase whole even though neither disjoint owned region does.""" segment = "aaaaa BLOCK ME zzzzz" chunks = _split_text(segment, 12) assert len(chunks) > 1 @@ -190,9 +196,9 @@ async def test_block_phrase_split_across_chunk_boundary_is_detected(): @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.""" +async def test_no_overlap_lets_seam_phrase_evade(): + """Control: with overlap disabled there is no prefix, so the same straddling + phrase is seen by neither owned region -- demonstrating what the overlap closes.""" segment = "aaaaa BLOCK ME zzzzz" async def evaluate(text): @@ -203,7 +209,7 @@ async def test_no_overlap_window_lets_boundary_phrase_evade(): @pytest.mark.asyncio -async def test_single_chunk_segment_evaluates_once_no_boundary_window(): +async def test_single_chunk_segment_evaluates_once(): calls = [] async def evaluate(text): @@ -215,19 +221,22 @@ async def test_single_chunk_segment_evaluates_once_no_boundary_window(): @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.""" +async def test_mask_straddling_a_seam_fails_closed_as_block(): + """A MASK whose redaction reaches into a chunk's overlap prefix (content + straddling, or within `overlap` of, an owned-region seam) cannot be stitched + without double-counting the overlap, so it fails closed as BLOCK rather than + leak the un-redacted half. This is the preempt for the seam-mask leak.""" 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("") + # The chunk that sees "SECRET HERE" whole (via its overlap prefix) masks + # it; the redaction lands in the prefix bytes -> fail closed. + return _result("MASK", action_text=text.replace("SECRET HERE", "[X]")) if "SECRET HERE" in text else _result("") verdicts = await evaluate_segments([segment], evaluate, max_chars=12, windows=WindowConfig(overlap=6)) - assert verdicts[0].action == "DETECT" + assert verdicts[0].action == "BLOCK" @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_processing.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_processing.py index bc57c5e4be2..f120c82e0ae 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_processing.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_processing.py @@ -24,24 +24,33 @@ def _block(detections=None, correlation_ids=None): return SegmentVerdict("BLOCK", None, detections or [], correlation_ids or []) -# --------------- reconstruct (masked-join alignment) --------------- +# --------------- reconstruct (per-chunk masked alignment) --------------- +# +# ``masked_chunks`` is the list of (owned_original, owned_masked) regions that +# concatenate to the joined document; for a single-chunk (<= prompt limit) +# document that is just ``[(join, masked_join)]``. + + +def _one_chunk(parts, masked): + original = JOINER.join(parts) + return [(original, masked)] def test_reconstruct_no_change_round_trips(): parts = ["alpha", "beta", "gamma"] - assert reconstruct(parts, JOINER.join(parts)) == parts + assert reconstruct(parts, _one_chunk(parts, JOINER.join(parts))) == parts def test_reconstruct_masks_a_middle_part(): parts = ["alpha", "sensitive", "gamma"] masked = JOINER.join(["alpha", "[REDACTED]", "gamma"]) - assert reconstruct(parts, masked) == ["alpha", "[REDACTED]", "gamma"] + assert reconstruct(parts, _one_chunk(parts, masked)) == ["alpha", "[REDACTED]", "gamma"] def test_reconstruct_mask_at_part_start(): parts = ["alpha", "beta", "gamma"] masked = JOINER.join(["[X]lpha", "beta", "gamma"]) - assert reconstruct(parts, masked) == ["[X]lpha", "beta", "gamma"] + assert reconstruct(parts, _one_chunk(parts, masked)) == ["[X]lpha", "beta", "gamma"] def test_reconstruct_handles_a_part_that_itself_contains_newline(): @@ -49,7 +58,7 @@ def test_reconstruct_handles_a_part_that_itself_contains_newline(): structural, not a naive split on '\\n', so this still reconstructs.""" parts = ["line1\nline1b", "second"] masked = JOINER.join(["line1\n[REDACTED]", "second"]) - assert reconstruct(parts, masked) == ["line1\n[REDACTED]", "second"] + assert reconstruct(parts, _one_chunk(parts, masked)) == ["line1\n[REDACTED]", "second"] def test_reconstruct_fails_closed_when_mask_spans_a_joiner(): @@ -57,19 +66,37 @@ def test_reconstruct_fails_closed_when_mask_spans_a_joiner(): closed (None) rather than misassign redacted text to the wrong message.""" parts = ["alpha", "beta", "gamma"] merged = "alphaXXXbeta\ngamma" # joiner between alpha|beta is gone - assert reconstruct(parts, merged) is None + assert reconstruct(parts, _one_chunk(parts, merged)) is None + + +def test_reconstruct_maps_across_multiple_chunks(): + """The document is aligned per owned-region chunk; a part living in a later + chunk is recovered through that chunk's own alignment, not a global diff.""" + parts = ["aaaa", "bbbb"] + # Two owned regions that concatenate to "aaaa\nbbbb"; the second is masked. + masked_chunks = [("aaaa\n", "aaaa\n"), ("bbbb", "XXXX")] + assert reconstruct(parts, masked_chunks) == ["aaaa", "XXXX"] + + +def test_reconstruct_fails_closed_when_chunks_do_not_match_join(): + """Invariant guard: the owned regions must concatenate back to the join.""" + parts = ["alpha", "beta"] + assert reconstruct(parts, [("alpha\nDIFFERENT", "alpha\nDIFFERENT")]) is None + + +def test_reconstruct_none_chunks_fails_closed(): + assert reconstruct(["a", "b"], None) is None def test_reconstruct_empty_parts_is_empty_list(): - assert reconstruct([], "") == [] + assert reconstruct([], None) == [] def test_reconstruct_fails_closed_when_document_exceeds_bound(): - """Reconstruction is bounded to avoid the quadratic SequenceMatcher cost - blocking the event loop; an over-bound document fails closed (None) instead - of running the alignment.""" + """Reconstruction is bounded as a coarse backstop on total alignment work; + an over-bound document fails closed (None) instead of aligning.""" big = "a" * (RECONSTRUCT_MAX_CHARS + 1) - assert reconstruct([big], big) is None + assert reconstruct([big], [(big, big)]) is None # --------------- check_scan_budget (total-work cap) ---------------