mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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.
This commit is contained in:
parent
e7b0c9232c
commit
0605a19709
5 changed files with 252 additions and 138 deletions
|
|
@ -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 "
|
||||
|
|
|
|||
|
|
@ -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))]
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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) ---------------
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue