fix(guardrails): Alice WonderFence scans every request segment regardless of role

Filtering the request side back down to user-role messages reopened the same
class of bypass for non-user content: disallowed text placed in a system,
assistant (prefill), or tool message went unscanned while still reaching the
model. Evaluate every segment the translation layer hands us in inputs["texts"]
instead of re-filtering by role; whether a role is included is already governed
upstream by skip_system_message_in_guardrail / skip_tool_message_in_guardrail,
so the hook should not hardcode its own role policy. This also removes the
role-mapping replay and its count-mismatch fallback entirely.

example_config sets skip_system_message_in_guardrail: true so admin-controlled
system prompts are excluded by default, which avoids false positives while still
scanning the caller-controllable assistant and tool segments.
This commit is contained in:
lior-k 2026-06-08 18:12:28 +03:00
parent 61c2870864
commit 2e75d23499
No known key found for this signature in database
5 changed files with 57 additions and 117 deletions

View file

@ -25,11 +25,7 @@ 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 (
apply_verdicts,
build_analysis_context,
request_user_text_indices,
)
from .processing import apply_verdicts, build_analysis_context
if TYPE_CHECKING:
from wonderfence_sdk.client import ( # type: ignore[import-untyped]
@ -196,9 +192,6 @@ class WonderFenceGuardrail(CustomGuardrail):
)
if input_type == "request":
indices = request_user_text_indices(
inputs.get("structured_messages"), texts
)
async def evaluate(text: str) -> Any:
return await client.evaluate_prompt(
@ -206,7 +199,6 @@ class WonderFenceGuardrail(CustomGuardrail):
)
else:
indices = list(range(len(texts)))
async def evaluate(text: str) -> Any:
return await client.evaluate_response(
@ -216,24 +208,24 @@ class WonderFenceGuardrail(CustomGuardrail):
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),
len(texts),
app_id,
self.guardrail_name,
input_type,
)
verdicts = await evaluate_segments(
segments,
texts,
evaluate,
max_concurrency=self._connection_pool_limit or DEFAULT_MAX_CONCURRENCY,
)
apply_verdicts(
inputs, indices, verdicts, self.guardrail_name, self.block_message
inputs,
list(range(len(texts))),
verdicts,
self.guardrail_name,
self.block_message,
)
except WonderFenceBlockedError as e:

View file

@ -41,6 +41,14 @@ guardrails:
max_cached_clients: 10
block_message: "Content violates our policies and has been blocked by Alice WonderFence"
# Every remaining message segment is evaluated (user, assistant, tool),
# not just the last user turn, so disallowed content placed in an earlier
# turn or an assistant prefill cannot slip past. System prompts are
# admin-controlled and excluded by default to avoid false positives;
# set this to false to scan them too. skip_tool_message_in_guardrail is
# the matching knob for tool messages.
skip_system_message_in_guardrail: true
# connection_pool_limit: 20
# Enable only for trusted-gateway deployments that need to forward a

View file

@ -52,41 +52,6 @@ def build_analysis_context(
)
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.
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.
"""
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 apply_verdicts(
inputs: GenericGuardrailAPIInputs,
indices: List[int],

View file

@ -99,12 +99,11 @@ async def test_apply_guardrail_mask_replaces_scanned_text(
@pytest.mark.asyncio
async def test_apply_guardrail_mask_targets_correct_user_slot(
async def test_apply_guardrail_mask_targets_only_the_flagged_slot(
guardrail_and_client, make_request_data
):
"""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."""
"""MASK rewrites the ``texts`` entry of the flagged segment in place; the
other scanned entries survive untouched. Confirms positional 1:1 mapping."""
guardrail, client = guardrail_and_client
def evaluate(prompt, **kwargs):
@ -117,23 +116,48 @@ async def test_apply_guardrail_mask_targets_correct_user_slot(
client.evaluate_prompt.side_effect = evaluate
inputs = {
"structured_messages": [
{"role": "user", "content": "first"},
{"role": "assistant", "content": "ack"},
{"role": "user", "content": "sensitive content"},
],
"texts": ["first", "ack", "sensitive content"],
}
out = await guardrail.apply_guardrail(
inputs=inputs,
inputs={"texts": ["first", "ack", "sensitive content"]},
request_data=make_request_data(),
input_type="request",
)
assert out["texts"] == ["first", "ack", "[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_scans_non_user_role_segments(
guardrail_and_client, make_request_data
):
"""Bypass regression: blocked content in a system/assistant/tool message
must still BLOCK. The translation layer already strips system/tool when the
guardrail is configured to skip them, so whatever remains in ``texts`` is
scanned regardless of role; the hook must not re-filter to user-only."""
guardrail, client = guardrail_and_client
def evaluate(prompt, **kwargs):
r = Mock()
r.action = "BLOCK" if prompt == "disallowed system instruction" else "NO_ACTION"
r.detections = []
r.correlation_id = None
return r
client.evaluate_prompt.side_effect = evaluate
inputs = {
"structured_messages": [
{"role": "system", "content": "disallowed system instruction"},
{"role": "user", "content": "hello"},
],
"texts": ["disallowed system instruction", "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

View file

@ -1,4 +1,4 @@
"""Tests for processing.py pure transforms: user-text mapping and verdict apply."""
"""Tests for processing.py pure transforms: verdict apply."""
import pytest
@ -10,57 +10,8 @@ from litellm.proxy.guardrails.guardrail_hooks.alice_wonderfence.exceptions impor
)
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 [])