mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
178fb74fbc
commit
61c2870864
6 changed files with 598 additions and 175 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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))]
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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"]
|
||||
Loading…
Add table
Reference in a new issue