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 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:
parent
4eaea1dd8a
commit
4ad7cceb5d
4 changed files with 168 additions and 16 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue