From 61c2870864930b1bce3aadd933f5f5575fd1ab59 Mon Sep 17 00:00:00 2001 From: lior-k Date: Mon, 8 Jun 2026 17:25:36 +0300 Subject: [PATCH] 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. --- .../alice_wonderfence/alice_wonderfence.py | 75 +++++--- .../alice_wonderfence/chunked_evaluation.py | 113 ++++++++++++ .../alice_wonderfence/processing.py | 145 +++++++-------- .../alice_wonderfence/test_apply_guardrail.py | 172 ++++++++++-------- .../test_chunked_evaluation.py | 170 +++++++++++++++++ .../alice_wonderfence/test_processing.py | 98 ++++++++++ 6 files changed, 598 insertions(+), 175 deletions(-) create mode 100644 litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/chunked_evaluation.py create mode 100644 tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_chunked_evaluation.py create mode 100644 tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_processing.py diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py index 86e882995da..701dc3736db 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py @@ -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: diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/chunked_evaluation.py b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/chunked_evaluation.py new file mode 100644 index 00000000000..117700bc8a2 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/chunked_evaluation.py @@ -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))] diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py index 0569bc4afea..7805faca164 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py @@ -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 diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail.py index 7ed77ba836f..1dd683667fc 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail.py @@ -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" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_chunked_evaluation.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_chunked_evaluation.py new file mode 100644 index 00000000000..631053b1969 --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_chunked_evaluation.py @@ -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 diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_processing.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_processing.py new file mode 100644 index 00000000000..87a50820a67 --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_processing.py @@ -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"]