mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
61c2870864
commit
2e75d23499
5 changed files with 57 additions and 117 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 [])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue