fix(guardrails): Alice WonderFence scans every user message, not just the last turn

The messages array is fully caller-controlled and unverified, so placing
disallowed content in an earlier user turn and ending with a benign message
let it reach the model unscanned; only the last consecutive user block was
evaluated. Now every user-role message is evaluated on its own (each chunked
to the WonderFence prompt limit), all calls fan out in parallel under a
concurrency cap, and verdicts are aggregated per message with BLOCK > MASK >
DETECT precedence. Masking writes back only through texts, matching what the
chat translation layer reads.

The chunk + parallel-evaluate + aggregate logic lives in one replaceable unit
(chunked_evaluation.py) that is WonderFence-agnostic via an injected evaluate
callable.
This commit is contained in:
lior-k 2026-06-08 17:25:36 +03:00
parent 178fb74fbc
commit 61c2870864
No known key found for this signature in database
6 changed files with 598 additions and 175 deletions

View file

@ -3,7 +3,7 @@
import logging
import os
from collections import OrderedDict
from typing import TYPE_CHECKING, List, Literal, Optional, Type, Union
from typing import TYPE_CHECKING, Any, List, Literal, Optional, Type, Union
from fastapi import HTTPException
@ -21,10 +21,15 @@ 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 .client_cache import get_or_create_client, load_sdk
from .credentials import resolve_credentials
from .exceptions import WonderFenceBlockedError, WonderFenceMissingSecrets
from .processing import build_analysis_context, extract_relevant_text, handle_action
from .processing import (
apply_verdicts,
build_analysis_context,
request_user_text_indices,
)
if TYPE_CHECKING:
from wonderfence_sdk.client import ( # type: ignore[import-untyped]
@ -168,10 +173,10 @@ class WonderFenceGuardrail(CustomGuardrail):
logging_obj: Optional["LiteLLMLoggingObj"] = None,
) -> GenericGuardrailAPIInputs:
"""Apply WonderFence guardrail using V2 client + per-request app_id."""
text, text_source = extract_relevant_text(inputs, input_type)
if not text:
texts = inputs.get("texts") or []
if not texts:
logger.debug(
"Alice WonderFence (apply_guardrail): no relevant text for %s",
"Alice WonderFence (apply_guardrail): no text to scan for %s",
input_type,
)
return inputs
@ -191,32 +196,44 @@ class WonderFenceGuardrail(CustomGuardrail):
)
if input_type == "request":
logger.debug(
"Alice WonderFence (apply_guardrail): evaluating prompt app_id=%s guardrail=%s",
app_id,
self.guardrail_name,
)
result = await client.evaluate_prompt(
app_id=app_id,
prompt=text,
context=context,
custom_fields=None,
)
else:
logger.debug(
"Alice WonderFence (apply_guardrail): evaluating response app_id=%s guardrail=%s",
app_id,
self.guardrail_name,
)
result = await client.evaluate_response(
app_id=app_id,
response=text,
context=context,
custom_fields=None,
indices = request_user_text_indices(
inputs.get("structured_messages"), texts
)
handle_action(
result, inputs, text_source, self.guardrail_name, self.block_message
async def evaluate(text: str) -> Any:
return await client.evaluate_prompt(
app_id=app_id, prompt=text, context=context, custom_fields=None
)
else:
indices = list(range(len(texts)))
async def evaluate(text: str) -> Any:
return await client.evaluate_response(
app_id=app_id,
response=text,
context=context,
custom_fields=None,
)
segments = [texts[i] for i in indices]
if not segments:
return inputs
logger.debug(
"Alice WonderFence (apply_guardrail): evaluating %d segment(s) app_id=%s guardrail=%s input_type=%s",
len(segments),
app_id,
self.guardrail_name,
input_type,
)
verdicts = await evaluate_segments(
segments,
evaluate,
max_concurrency=self._connection_pool_limit or DEFAULT_MAX_CONCURRENCY,
)
apply_verdicts(
inputs, indices, verdicts, self.guardrail_name, self.block_message
)
except WonderFenceBlockedError as e:

View file

@ -0,0 +1,113 @@
"""Chunk, evaluate in parallel, and aggregate per segment.
Guardrail-agnostic: the only coupling to WonderFence is the injected
``evaluate`` callable and the result shape it returns (``action``,
``action_text``, ``detections``, ``correlation_id``). Replace ``evaluate`` to
target a different backend.
"""
import asyncio
import re
from dataclasses import dataclass
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
@dataclass
class SegmentVerdict:
action: str # "BLOCK" | "MASK" | "DETECT" | ""
masked_text: Optional[str]
detections: list
correlation_ids: List[str]
def _split_text(text: str, max_chars: int) -> List[str]:
"""Split ``text`` into <= ``max_chars`` chunks with ``"".join(chunks) == text``.
Splits at whitespace boundaries; whitespace runs are preserved as their own
tokens so the rejoin is byte-identical. A single token longer than
``max_chars`` is force-split. Always returns at least one chunk.
"""
if len(text) <= max_chars:
return [text]
tokens = re.findall(r"\S+|\s+", text)
chunks: List[str] = []
current = ""
for token in tokens:
if len(current) + len(token) <= max_chars:
current += token
continue
if current:
chunks.append(current)
current = ""
while len(token) > max_chars:
chunks.append(token[:max_chars])
token = token[max_chars:]
current = token
if current:
chunks.append(current)
return chunks
def _action_str(result: Any) -> str:
action = getattr(result, "action", "")
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]
detections: list = []
correlation_ids: List[str] = []
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 actions:
return SegmentVerdict("BLOCK", None, detections, correlation_ids)
if "MASK" in actions:
masked = "".join(
(r.action_text or "[MASKED]") if _action_str(r) == "MASK" else chunk
for chunk, r in zip(chunks, results)
)
return SegmentVerdict("MASK", masked, detections, correlation_ids)
if "DETECT" in actions:
return SegmentVerdict("DETECT", None, detections, correlation_ids)
return SegmentVerdict("", None, detections, correlation_ids)
async def evaluate_segments(
segments: List[str],
evaluate: Callable[[str], Awaitable[Any]],
max_chars: int = MAX_PROMPT_CHARS,
max_concurrency: int = DEFAULT_MAX_CONCURRENCY,
) -> 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
``Semaphore(max_concurrency)``. Results are grouped back per segment with
action precedence BLOCK > MASK > DETECT > NO_ACTION.
"""
semaphore = asyncio.Semaphore(max_concurrency)
async def run(chunk: str) -> Any:
async with semaphore:
return await evaluate(chunk)
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]
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
return [_aggregate(seg_chunks[si], per_segment[si]) for si in range(len(segments))]

View file

@ -1,19 +1,15 @@
"""Pure transforms for Alice WonderFence: context build, text extract, action dispatch."""
"""Pure transforms for Alice WonderFence: context build, user-text mapping, verdict apply."""
from typing import Any, Literal, Optional, Tuple
from typing import Any, List, Optional
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.prompt_templates.common_utils import (
get_last_user_message,
set_last_user_message,
)
from litellm.types.utils import GenericGuardrailAPIInputs
from .chunked_evaluation import SegmentVerdict
from .credentials import get_metadata
from .exceptions import WonderFenceBlockedError
logger = verbose_proxy_logger.getChild("alice_wonderfence")
@ -56,88 +52,93 @@ def build_analysis_context(
)
def extract_relevant_text(
inputs: GenericGuardrailAPIInputs,
input_type: Literal["request", "response"],
) -> Tuple[Optional[str], Optional[Literal["structured_messages", "texts"]]]:
"""Extract latest user message (request) or latest assistant message (response).
def request_user_text_indices(
structured_messages: Optional[List[Any]],
texts: List[str],
) -> List[int]:
"""Return indices into ``texts`` that came from user-role messages.
Returns (text, source) — ``source`` identifies which slot the text came
from so MASK can write the redacted version back to the same place.
Replays the same flatten the translation layer uses to build ``texts``
(string content -> one entry; list content -> one entry per item with a
``text`` field) over ``structured_messages`` and tags each entry's role. If
``structured_messages`` is absent or the replayed count diverges from
``len(texts)``, every index is returned: over-scanning is safe, mis-mapping
a mask onto a non-user slot is not.
"""
if input_type == "request":
structured_messages = inputs.get("structured_messages", [])
if structured_messages:
return (
get_last_user_message(structured_messages),
"structured_messages",
)
texts = inputs.get("texts", [])
return (texts[-1] if texts else None), ("texts" if texts else None)
texts = inputs.get("texts", [])
return (texts[-1] if texts else None), ("texts" if texts else None)
n = len(texts)
if not structured_messages:
return list(range(n))
roles: List[str] = []
for message in structured_messages:
role = str(message.get("role") or "").lower()
content = message.get("content", None)
if content is None:
continue
if isinstance(content, str):
roles.append(role)
elif isinstance(content, list):
for item in content:
if item.get("text", None) is not None:
roles.append(role)
if len(roles) != n:
return list(range(n))
return [i for i, role in enumerate(roles) if role == "user"]
def handle_action(
result: Any,
def apply_verdicts(
inputs: GenericGuardrailAPIInputs,
text_source: Optional[Literal["structured_messages", "texts"]],
indices: List[int],
verdicts: List[SegmentVerdict],
guardrail_name: str,
block_message: str,
) -> None:
"""Dispatch BLOCK/MASK/DETECT/NO_ACTION. Raises ``WonderFenceBlockedError`` on BLOCK.
) -> GenericGuardrailAPIInputs:
"""Apply per-segment verdicts back onto ``inputs["texts"]``.
``text_source`` identifies which inputs slot supplied the analyzed text;
MASK writes the redacted value back to the same slot.
Any BLOCK raises ``WonderFenceBlockedError`` with detections/correlation ids
aggregated across all blocked segments. Otherwise each MASK verdict rewrites
its mapped ``texts`` index and DETECT is logged.
"""
action = result.action.value if hasattr(result.action, "value") else result.action
correlation_id = getattr(result, "correlation_id", None)
if action == "BLOCK":
blocked = [v for v in verdicts if v.action == "BLOCK"]
if blocked:
detections: list = []
correlation_ids: List[str] = []
for v in blocked:
detections.extend(v.detections)
correlation_ids.extend(v.correlation_ids)
detail: dict = {
"error": block_message,
"type": "alice_wonderfence_content_policy_violation",
"guardrail_name": guardrail_name,
"action": "BLOCK",
"wonderfence_correlation_id": correlation_id,
"wonderfence_correlation_id": (
correlation_ids[0] if correlation_ids else None
),
"wonderfence_correlation_ids": correlation_ids,
}
if hasattr(result, "detections") and result.detections:
if detections:
detail["detections"] = [
d.model_dump() if hasattr(d, "model_dump") else str(d)
for d in result.detections
d.model_dump() if hasattr(d, "model_dump") else d for d in detections
]
raise WonderFenceBlockedError(detail)
if action == "MASK":
masked_text = result.action_text or "[MASKED]"
wrote = False
if text_source == "structured_messages":
inputs["structured_messages"] = set_last_user_message(
inputs.get("structured_messages", []), masked_text
texts = inputs.get("texts") or []
for idx, verdict in zip(indices, verdicts):
if verdict.action == "MASK":
texts[idx] = (
verdict.masked_text if verdict.masked_text is not None else "[MASKED]"
)
wrote = True
# Always also overwrite texts[-1] when texts is populated. The OpenAI
# chat translation layer reads back only ``texts`` after
# apply_guardrail returns and maps it onto messages — masking only
# ``structured_messages`` lets the unmasked ``texts`` slot win and the
# original prompt reaches the LLM.
texts = inputs.get("texts")
if texts:
texts[-1] = masked_text
inputs["texts"] = texts
wrote = True
if not wrote: # pragma: no cover
raise RuntimeError(
"Alice WonderFence MASK requested but no text source — refusing "
"to silently no-op."
logger.info(
"Alice WonderFence (apply_guardrail): MASK applied guardrail=%s correlation_id=%s",
guardrail_name,
verdict.correlation_ids[0] if verdict.correlation_ids else None,
)
logger.info(
"Alice WonderFence (apply_guardrail): MASK applied guardrail=%s correlation_id=%s",
guardrail_name,
correlation_id,
)
elif action == "DETECT":
logger.warning(
"Alice WonderFence (apply_guardrail): DETECT guardrail=%s correlation_id=%s",
guardrail_name,
correlation_id,
)
elif verdict.action == "DETECT":
logger.warning(
"Alice WonderFence (apply_guardrail): DETECT guardrail=%s correlation_id=%s",
guardrail_name,
verdict.correlation_ids[0] if verdict.correlation_ids else None,
)
inputs["texts"] = texts
return inputs

View file

@ -6,7 +6,6 @@ from unittest.mock import Mock
import pytest
from fastapi import HTTPException
# ----------------------------- BLOCK -----------------------------
@ -80,7 +79,7 @@ async def test_block_not_bypassed_by_fail_open(make_guardrail, make_request_data
@pytest.mark.asyncio
async def test_apply_guardrail_mask_replaces_last_text(
async def test_apply_guardrail_mask_replaces_scanned_text(
guardrail_and_client, make_request_data
):
guardrail, client = guardrail_and_client
@ -92,60 +91,31 @@ async def test_apply_guardrail_mask_replaces_last_text(
client.evaluate_prompt.return_value = result_obj
out = await guardrail.apply_guardrail(
inputs={"texts": ["a", "b", "c"]},
inputs={"texts": ["sensitive"]},
request_data=make_request_data(),
input_type="request",
)
assert out["texts"] == ["a", "b", "[REDACTED]"]
assert out["texts"] == ["[REDACTED]"]
@pytest.mark.asyncio
async def test_apply_guardrail_mask_replaces_structured_messages(
async def test_apply_guardrail_mask_targets_correct_user_slot(
guardrail_and_client, make_request_data
):
"""MASK on the request path must rewrite structured_messages when that's
the source of the extracted text. Otherwise the user's prompt reaches the
LLM unredacted while the header still claims the guardrail applied."""
"""MASK must rewrite the ``texts`` entry of the offending user message in
place; assistant/system entries are never sent for evaluation, so they must
survive untouched. Confirms the positional mapping is correct."""
guardrail, client = guardrail_and_client
result_obj = Mock()
result_obj.action = "MASK"
result_obj.action_text = "[REDACTED]"
result_obj.detections = []
result_obj.correlation_id = None
client.evaluate_prompt.return_value = result_obj
inputs = {
"structured_messages": [
{"role": "user", "content": "first"},
{"role": "assistant", "content": "ack"},
{"role": "user", "content": "sensitive content"},
],
}
out = await guardrail.apply_guardrail(
inputs=inputs,
request_data=make_request_data(),
input_type="request",
)
last_user = [m for m in out["structured_messages"] if m.get("role") == "user"][-1]
assert last_user["content"] == "[REDACTED]"
def evaluate(prompt, **kwargs):
r = Mock()
r.action = "MASK" if prompt == "sensitive content" else "NO_ACTION"
r.action_text = "[REDACTED]"
r.detections = []
r.correlation_id = None
return r
@pytest.mark.asyncio
async def test_apply_guardrail_mask_rewrites_texts_when_both_slots_present(
guardrail_and_client, make_request_data
):
"""OpenAI chat translation populates both ``structured_messages`` and ``texts``,
then reads back only ``texts``. MASK must overwrite ``texts[-1]`` even when
the analyzed text was extracted from ``structured_messages``, otherwise the
unmasked ``texts`` slot wins downstream and the original prompt reaches the
LLM while the response header still claims the guardrail applied."""
guardrail, client = guardrail_and_client
result_obj = Mock()
result_obj.action = "MASK"
result_obj.action_text = "[REDACTED]"
result_obj.detections = []
result_obj.correlation_id = None
client.evaluate_prompt.return_value = result_obj
client.evaluate_prompt.side_effect = evaluate
inputs = {
"structured_messages": [
@ -161,12 +131,13 @@ async def test_apply_guardrail_mask_rewrites_texts_when_both_slots_present(
input_type="request",
)
assert out["texts"] == ["first", "ack", "[REDACTED]"]
last_user = [m for m in out["structured_messages"] if m.get("role") == "user"][-1]
assert last_user["content"] == "[REDACTED]"
evaluated = {c.kwargs["prompt"] for c in client.evaluate_prompt.call_args_list}
assert evaluated == {"first", "sensitive content"}
assert "ack" not in evaluated
@pytest.mark.asyncio
async def test_apply_guardrail_mask_replaces_last_text_response(
async def test_apply_guardrail_mask_replaces_scanned_text_response(
guardrail_and_client, make_request_data
):
guardrail, client = guardrail_and_client
@ -178,11 +149,11 @@ async def test_apply_guardrail_mask_replaces_last_text_response(
client.evaluate_response.return_value = result_obj
out = await guardrail.apply_guardrail(
inputs={"texts": ["a", "b", "c"]},
inputs={"texts": ["model output"]},
request_data=make_request_data(),
input_type="response",
)
assert out["texts"] == ["a", "b", "[REDACTED]"]
assert out["texts"] == ["[REDACTED]"]
@pytest.mark.asyncio
@ -198,11 +169,11 @@ async def test_apply_guardrail_mask_fallback_when_action_text_is_none(
client.evaluate_prompt.return_value = result_obj
out = await guardrail.apply_guardrail(
inputs={"texts": ["a", "b", "c"]},
inputs={"texts": ["a"]},
request_data=make_request_data(),
input_type="request",
)
assert out["texts"] == ["a", "b", "[MASKED]"]
assert out["texts"] == ["[MASKED]"]
# ----------------------------- DETECT / NO_ACTION -----------------------------
@ -301,9 +272,11 @@ async def test_apply_guardrail_response_path_passes_app_id(
@pytest.mark.asyncio
async def test_apply_guardrail_evaluates_only_last_text(
async def test_apply_guardrail_evaluates_every_text_without_structured_messages(
guardrail_and_client, make_request_data
):
"""With no structured_messages to identify roles, every text entry is
scanned (over-scan is safe); the old code scanned only the last."""
guardrail, client = guardrail_and_client
result_obj = Mock()
result_obj.action = "NO_ACTION"
@ -316,8 +289,78 @@ async def test_apply_guardrail_evaluates_only_last_text(
request_data=make_request_data(),
input_type="request",
)
assert client.evaluate_prompt.call_count == 1
assert client.evaluate_prompt.call_args.kwargs["prompt"] == "t3"
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"}
@pytest.mark.asyncio
async def test_apply_guardrail_blocks_on_earlier_user_turn(
guardrail_and_client, make_request_data
):
"""Bypass regression: disallowed content in an earlier user turn followed by
a benign final turn must still BLOCK. The old last-only path only saw the
benign final message and let the request through."""
guardrail, client = guardrail_and_client
def evaluate(prompt, **kwargs):
r = Mock()
r.action = "BLOCK" if prompt == "disallowed" else "NO_ACTION"
r.detections = []
r.correlation_id = None
return r
client.evaluate_prompt.side_effect = evaluate
inputs = {
"structured_messages": [
{"role": "user", "content": "disallowed"},
{"role": "assistant", "content": "ok"},
{"role": "user", "content": "hello"},
],
"texts": ["disallowed", "ok", "hello"],
}
with pytest.raises(HTTPException) as exc:
await guardrail.apply_guardrail(
inputs=inputs,
request_data=make_request_data(),
input_type="request",
)
assert exc.value.status_code == 400
assert exc.value.detail["action"] == "BLOCK"
@pytest.mark.asyncio
async def test_apply_guardrail_blocks_when_oversized_message_trips_in_late_chunk(
guardrail_and_client, make_request_data
):
"""A single user message over the prompt limit is chunked; a BLOCK in a
non-first chunk still blocks the request."""
from litellm.proxy.guardrails.guardrail_hooks.alice_wonderfence import (
chunked_evaluation,
)
guardrail, client = guardrail_and_client
long_prompt = ("safe " * 5000) + "TRIPWIRE"
assert len(long_prompt) > chunked_evaluation.MAX_PROMPT_CHARS
def evaluate(prompt, **kwargs):
r = Mock()
r.action = "BLOCK" if "TRIPWIRE" in prompt else "NO_ACTION"
r.detections = []
r.correlation_id = None
return r
client.evaluate_prompt.side_effect = evaluate
with pytest.raises(HTTPException) as exc:
await guardrail.apply_guardrail(
inputs={"texts": [long_prompt]},
request_data=make_request_data(),
input_type="request",
)
assert exc.value.status_code == 400
assert client.evaluate_prompt.call_count > 1
@pytest.mark.asyncio
@ -475,22 +518,3 @@ def test_build_analysis_context_falls_back_to_slash_split(monkeypatch, make_guar
kwargs = AnalysisContext.call_args.kwargs
assert kwargs["provider"] == "myorg"
assert kwargs["model_name"] == "custom-llm"
def test_extract_relevant_text_uses_structured_messages():
"""Request path with structured_messages routes through get_last_user_message."""
from litellm.proxy.guardrails.guardrail_hooks.alice_wonderfence.processing import (
extract_relevant_text,
)
inputs = {
"structured_messages": [
{"role": "user", "content": "first"},
{"role": "assistant", "content": "ack"},
{"role": "user", "content": "latest user msg"},
],
"texts": ["unused-fallback"],
}
text, source = extract_relevant_text(inputs, input_type="request") # type: ignore[arg-type]
assert text == "latest user msg"
assert source == "structured_messages"

View file

@ -0,0 +1,170 @@
"""Tests for the WonderFence-agnostic chunk + parallel-evaluate + aggregate unit."""
import asyncio
from unittest.mock import Mock
import pytest
from litellm.proxy.guardrails.guardrail_hooks.alice_wonderfence.chunked_evaluation import (
MAX_PROMPT_CHARS,
SegmentVerdict,
_split_text,
evaluate_segments,
)
def _result(action, action_text=None, detections=None, correlation_id=None):
r = Mock()
r.action = action
r.action_text = action_text
r.detections = detections or []
r.correlation_id = correlation_id
return r
# ----------------------------- _split_text -----------------------------
def test_split_short_text_is_single_chunk():
assert _split_text("hello world", 10000) == ["hello world"]
def test_split_long_text_is_lossless():
text = " ".join(f"word{i}" for i in range(5000))
chunks = _split_text(text, 100)
assert len(chunks) > 1
assert all(len(c) <= 100 for c in chunks)
assert "".join(chunks) == text
def test_split_force_splits_oversized_single_token():
text = "x" * 250
chunks = _split_text(text, 100)
assert all(len(c) <= 100 for c in chunks)
assert "".join(chunks) == text
# ----------------------------- evaluate_segments alignment -----------------------------
@pytest.mark.asyncio
async def test_verdicts_align_one_to_one_with_segments():
actions = {"a": "BLOCK", "b": "MASK", "c": ""}
async def evaluate(text):
return _result(
actions[text], action_text="[M]" if actions[text] == "MASK" else None
)
verdicts = await evaluate_segments(["a", "b", "c"], evaluate)
assert [v.action for v in verdicts] == ["BLOCK", "MASK", ""]
assert isinstance(verdicts[0], SegmentVerdict)
@pytest.mark.asyncio
async def test_mask_verdict_carries_masked_text():
async def evaluate(text):
return _result("MASK", action_text="[REDACTED]")
verdicts = await evaluate_segments(["secret"], evaluate)
assert verdicts[0].action == "MASK"
assert verdicts[0].masked_text == "[REDACTED]"
# ----------------------------- chunking precedence -----------------------------
@pytest.mark.asyncio
async def test_block_in_non_first_chunk_blocks_whole_segment():
"""A segment split into chunks where only a later chunk trips BLOCK must
still produce a BLOCK verdict; the old last-only path never saw earlier text."""
segment = ("safe " * 30) + "TRIPWIRE"
async def evaluate(text):
return _result("BLOCK" if "TRIPWIRE" in text else "")
verdicts = await evaluate_segments([segment], evaluate, max_chars=50)
assert verdicts[0].action == "BLOCK"
@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 evaluate(text):
return _result("MASK", action_text=f"<{text}>")
verdicts = await evaluate_segments([segment], evaluate, max_chars=50)
assert verdicts[0].action == "MASK"
assert verdicts[0].masked_text == "".join(f"<{c}>" for c in chunks)
@pytest.mark.asyncio
async def test_unmasked_chunks_fall_back_to_original_text_on_rejoin():
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]")
verdicts = await evaluate_segments([segment], evaluate, max_chars=20)
expected = "[X]" + "".join(chunks[1:])
assert verdicts[0].masked_text == expected
@pytest.mark.asyncio
async def test_block_beats_mask_within_segment():
chunks_seen = []
async def evaluate(text):
chunks_seen.append(text)
return _result("BLOCK" if "B" in text else "MASK", action_text="[m]")
segment = "aaa B"
verdicts = await evaluate_segments([segment], evaluate, max_chars=2)
assert verdicts[0].action == "BLOCK"
# ----------------------------- aggregation of detections/correlation ids -----------------------------
@pytest.mark.asyncio
async def test_block_verdict_aggregates_detections_and_correlation_ids():
d1, d2 = Mock(), Mock()
async def evaluate(text):
if "x" in text:
return _result("BLOCK", detections=[d1], correlation_id="c1")
return _result("BLOCK", detections=[d2], correlation_id="c2")
verdicts = await evaluate_segments(["x", "y"], evaluate)
assert verdicts[0].detections == [d1]
assert verdicts[0].correlation_ids == ["c1"]
assert verdicts[1].correlation_ids == ["c2"]
# ----------------------------- concurrency cap -----------------------------
@pytest.mark.asyncio
async def test_evaluations_run_in_parallel_under_a_cap():
state = {"current": 0, "max_seen": 0}
async def evaluate(text):
state["current"] += 1
state["max_seen"] = max(state["max_seen"], state["current"])
await asyncio.sleep(0.01)
state["current"] -= 1
return _result("")
segments = [f"s{i}" for i in range(12)]
await evaluate_segments(segments, evaluate, max_concurrency=3)
assert state["max_seen"] > 1, "evaluations did not run concurrently"
assert state["max_seen"] <= 3, "concurrency cap exceeded"
def test_max_prompt_chars_is_positive():
assert isinstance(MAX_PROMPT_CHARS, int) and MAX_PROMPT_CHARS > 0

View file

@ -0,0 +1,98 @@
"""Tests for processing.py pure transforms: user-text mapping and verdict apply."""
import pytest
from litellm.proxy.guardrails.guardrail_hooks.alice_wonderfence.chunked_evaluation import (
SegmentVerdict,
)
from litellm.proxy.guardrails.guardrail_hooks.alice_wonderfence.exceptions import (
WonderFenceBlockedError,
)
from litellm.proxy.guardrails.guardrail_hooks.alice_wonderfence.processing import (
apply_verdicts,
request_user_text_indices,
)
# ----------------------------- request_user_text_indices -----------------------------
def test_only_user_string_messages_are_indexed():
messages = [
{"role": "user", "content": "a"},
{"role": "assistant", "content": "b"},
{"role": "user", "content": "c"},
]
assert request_user_text_indices(messages, ["a", "b", "c"]) == [0, 2]
def test_system_message_excluded_even_when_present_in_texts():
messages = [
{"role": "system", "content": "sys"},
{"role": "user", "content": "hi"},
]
assert request_user_text_indices(messages, ["sys", "hi"]) == [1]
def test_list_content_yields_one_index_per_text_part():
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "x"},
{"type": "image_url", "image_url": {"url": "http://img"}},
{"type": "text", "text": "y"},
],
},
]
# texts flattens to the two text parts (image contributes no text entry)
assert request_user_text_indices(messages, ["x", "y"]) == [0, 1]
def test_absent_structured_messages_scans_all_indices():
assert request_user_text_indices(None, ["a", "b", "c"]) == [0, 1, 2]
def test_count_mismatch_falls_back_to_scanning_all():
"""If the replayed flatten count diverges from len(texts), over-scan rather
than risk mis-mapping a mask onto the wrong slot."""
messages = [{"role": "user", "content": "a"}]
assert request_user_text_indices(messages, ["a", "b"]) == [0, 1]
# ----------------------------- apply_verdicts -----------------------------
def _block(detections=None, correlation_ids=None):
return SegmentVerdict("BLOCK", None, detections or [], correlation_ids or [])
def test_block_verdict_raises_with_aggregated_detections():
inputs = {"texts": ["bad", "ok"]}
d = {"policy_name": "p"}
verdicts = [
_block(detections=[d], correlation_ids=["c1"]),
SegmentVerdict("", None, [], []),
]
with pytest.raises(WonderFenceBlockedError) as exc:
apply_verdicts(inputs, [0, 1], verdicts, "gn", "blocked!")
assert exc.value.detail["error"] == "blocked!"
assert exc.value.detail["action"] == "BLOCK"
assert exc.value.detail["detections"] == [d]
assert exc.value.detail["wonderfence_correlation_id"] == "c1"
def test_mask_writes_to_the_mapped_text_index_only():
inputs = {"texts": ["keep", "MASK_ME", "keep2"]}
verdicts = [SegmentVerdict("MASK", "[R]", [], [])]
out = apply_verdicts(inputs, [1], verdicts, "gn", "blocked!")
assert out["texts"] == ["keep", "[R]", "keep2"]
def test_detect_and_no_action_leave_texts_unchanged():
inputs = {"texts": ["a", "b"]}
verdicts = [
SegmentVerdict("DETECT", None, [], []),
SegmentVerdict("", None, [], []),
]
out = apply_verdicts(inputs, [0, 1], verdicts, "gn", "blocked!")
assert out["texts"] == ["a", "b"]