feat(guardrails): join Alice WonderFence request scan into one call + total-work cap

WonderFence has no batch API (one HTTP POST per string), so the request side
previously issued one call per message part plus tool-call args, tool defs and
legacy function defs, and cross-segment windows on top; call volume scaled with
message count and nothing bounded the total (the open veria "unbounded upstream
request amplification" finding).

Request side now joins all scan pieces into one document and scans it in ~1
call (chunked only when it exceeds the size limit), matching the dominant
one-call-per-direction pattern of the other join-style guardrails; call volume
scales with total size, not message count. On MASK we reconstruct per-part
masked text by aligning the join against the masked document with
difflib.SequenceMatcher (plain "\n" joiner, no sentinel) and write the recovered
message-text parts back to inputs["texts"] positionally so the handler maps them
onto the right message parts; if a joiner or a part boundary lands inside a
masked span we fail closed rather than misassign. Tool-call args and tool /
function descriptions are appended as detection-only pieces (raw strings, as
lakera/panw do) since the joined form is not the wire format and a redaction
cannot be spliced back into arguments/schema.

Adds a fail-closed total-work cap (max_scan_chars / max_scan_segments, with env
overrides) rejected with HTTP 400 before any provider call and never subject to
fail_open, which resolves the amplification finding.

Response side stays per-segment (independent choices / model tool-call args)
because the handler's response write-back is purely positional and has no
structured_messages path. Cross-segment windows and WindowConfig.text_segment_count
are removed (message-part junctions are now interior chunk seams; response
choices are never concatenated).
This commit is contained in:
lior-k 2026-07-23 11:11:43 +03:00
parent d852ae51f3
commit 8bceb610c8
No known key found for this signature in database
11 changed files with 639 additions and 361 deletions

View file

@ -44,6 +44,10 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
init_kwargs["debug"] = litellm_params.debug
if litellm_params.allow_request_metadata_override is not None:
init_kwargs["allow_request_metadata_override"] = litellm_params.allow_request_metadata_override
if litellm_params.max_scan_chars is not None:
init_kwargs["max_scan_chars"] = litellm_params.max_scan_chars
if litellm_params.max_scan_segments is not None:
init_kwargs["max_scan_segments"] = litellm_params.max_scan_segments
wonderfence_guardrail = WonderFenceGuardrail(**init_kwargs)

View file

@ -3,6 +3,7 @@
import logging
import os
from collections import OrderedDict
from collections.abc import Awaitable, Callable
from typing import TYPE_CHECKING, Any, Literal, Optional
from fastapi import HTTPException
@ -23,16 +24,24 @@ from litellm.types.utils import GenericGuardrailAPIInputs
from .chunked_evaluation import (
DEFAULT_MAX_CONCURRENCY,
WindowConfig,
evaluate_segments,
)
from .client_cache import ClientBuildSpec, get_or_create_client, load_sdk
from .credentials import CredentialConfig, resolve_credentials
from .exceptions import WonderFenceBlockedError, WonderFenceMissingSecrets
from .exceptions import (
WonderFenceBlockedError,
WonderFenceMissingSecrets,
WonderFenceScanBudgetExceeded,
)
from .processing import (
apply_verdicts,
JOINER,
apply_response_verdicts,
block_detail,
build_analysis_context,
check_scan_budget,
function_definition_segments,
raise_if_blocked,
reconstruct,
tool_call_arg_segments,
tool_definition_segments,
)
@ -77,6 +86,8 @@ class WonderFenceGuardrail(CustomGuardrail):
max_cached_clients: int | None = None,
connection_pool_limit: int | None = None,
allow_request_metadata_override: bool = False,
max_scan_chars: int | None = None,
max_scan_segments: int | None = None,
event_hook: (GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None) = None,
default_on: bool = True,
**kwargs: Any,
@ -102,6 +113,12 @@ class WonderFenceGuardrail(CustomGuardrail):
``metadata.alice_wonderfence_app_id`` as a last-resort source
(after API-key and team metadata). Defaults to False so
caller-controlled fields cannot bypass admin-pinned credentials.
max_scan_chars: Fail-closed total-work cap on combined scan
characters per request/response. Default 1_000_000. Env:
ALICE_MAX_SCAN_CHARS.
max_scan_segments: Fail-closed total-work cap on scan segment count
per request/response. Default 1_000. Env:
ALICE_MAX_SCAN_SEGMENTS.
event_hook: Event hook mode.
default_on: Whether the guardrail is enabled by default.
"""
@ -116,6 +133,16 @@ class WonderFenceGuardrail(CustomGuardrail):
self.fail_open = fail_open
self.block_message = block_message
self.allow_request_metadata_override = allow_request_metadata_override
env_max_chars = os.environ.get("ALICE_MAX_SCAN_CHARS")
self.max_scan_chars: int | None = (
max_scan_chars if max_scan_chars is not None else (int(env_max_chars) if env_max_chars else 1_000_000)
)
env_max_segments = os.environ.get("ALICE_MAX_SCAN_SEGMENTS")
self.max_scan_segments: int | None = (
max_scan_segments
if max_scan_segments is not None
else (int(env_max_segments) if env_max_segments else 1_000)
)
if debug:
logger.setLevel(logging.DEBUG)
@ -174,16 +201,29 @@ class WonderFenceGuardrail(CustomGuardrail):
input_type: Literal["request", "response"],
logging_obj: Optional["LiteLLMLoggingObj"] = None,
) -> GenericGuardrailAPIInputs:
"""Apply WonderFence guardrail using V2 client + per-request app_id."""
"""Apply WonderFence guardrail using V2 client + per-request app_id.
Request side joins all scan pieces (message text plus detection-only
tool-call args and tool/function descriptions) into one document and
scans it in ~1 call (chunked only when it exceeds the size limit), so
call volume scales with total size rather than message count. Response
side stays per-segment (independent choices / model tool-call args)
because the handler's response write-back is purely positional and has
no ``structured_messages`` path.
"""
texts = inputs.get("texts") or []
tool_indices, tool_segments = tool_call_arg_segments(inputs)
tool_def_paths, tool_def_segments = tool_definition_segments(inputs)
tool_indices, tool_arg_segments = tool_call_arg_segments(inputs)
tool_def_texts = tool_definition_segments(inputs)
# Legacy top-level functions[] only exist on the request body; the
# translation layer does not surface them in inputs, so read request_data.
function_def_paths, function_def_segments = (
function_definition_segments(request_data) if input_type == "request" else ([], [])
)
if not texts and not tool_segments and not tool_def_segments and not function_def_segments:
function_def_texts = function_definition_segments(request_data) if input_type == "request" else []
if input_type == "request":
scan_pieces = [*texts, *tool_arg_segments, *tool_def_texts, *function_def_texts]
else:
scan_pieces = [*texts, *tool_arg_segments]
if not scan_pieces:
logger.debug(
"Alice WonderFence (apply_guardrail): nothing to scan for %s",
input_type,
@ -191,6 +231,7 @@ class WonderFenceGuardrail(CustomGuardrail):
return inputs
try:
check_scan_budget(scan_pieces, self.max_scan_chars, self.max_scan_segments)
api_key, app_id = resolve_credentials(
request_data,
input_type,
@ -203,12 +244,14 @@ class WonderFenceGuardrail(CustomGuardrail):
)
client = await self._get_client(api_key)
context = build_analysis_context(request_data, self.platform, self._AnalysisContext)
max_concurrency = self._connection_pool_limit or DEFAULT_MAX_CONCURRENCY
if input_type == "request":
async def evaluate(text: str) -> object:
return await client.evaluate_prompt(app_id=app_id, prompt=text, context=context, custom_fields=None)
await self._scan_request(inputs, scan_pieces, len(texts), evaluate, max_concurrency, app_id)
else:
async def evaluate(text: str) -> object:
@ -219,46 +262,12 @@ class WonderFenceGuardrail(CustomGuardrail):
custom_fields=None,
)
segments = [
*texts,
*tool_segments,
*tool_def_segments,
*function_def_segments,
]
logger.debug(
"Alice WonderFence (apply_guardrail): evaluating %d text + %d tool-call + %d tool-def + %d function-def segment(s) app_id=%s guardrail=%s input_type=%s",
len(texts),
len(tool_segments),
len(tool_def_segments),
len(function_def_segments),
app_id,
self.guardrail_name,
input_type,
)
verdicts = await evaluate_segments(
segments,
evaluate,
max_concurrency=self._connection_pool_limit or DEFAULT_MAX_CONCURRENCY,
windows=WindowConfig(text_segment_count=len(texts)),
)
n_text = len(texts)
n_tool = len(tool_segments)
n_tool_def = len(tool_def_segments)
apply_verdicts(
inputs,
list(range(n_text)),
verdicts[:n_text],
self.guardrail_name,
self.block_message,
tool_indices=tool_indices,
tool_verdicts=verdicts[n_text : n_text + n_tool],
tool_def_paths=tool_def_paths,
tool_def_verdicts=verdicts[n_text + n_tool : n_text + n_tool + n_tool_def],
function_def_paths=function_def_paths,
function_def_verdicts=verdicts[n_text + n_tool + n_tool_def :],
function_def_request_data=request_data,
)
await self._scan_response(inputs, texts, tool_indices, tool_arg_segments, evaluate, max_concurrency)
except WonderFenceScanBudgetExceeded as e:
# Fail-closed config/abuse guard: reject before any provider call and
# never fall through to the fail_open path below.
raise HTTPException(status_code=400, detail=e.detail)
except WonderFenceBlockedError as e:
raise HTTPException(status_code=400, detail=e.detail)
except WonderFenceMissingSecrets as e:
@ -310,6 +319,92 @@ class WonderFenceGuardrail(CustomGuardrail):
add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name)
return inputs
async def _scan_request(
self,
inputs: GenericGuardrailAPIInputs,
pieces: list[str],
n_text: int,
evaluate: "Callable[[str], Awaitable[object]]",
max_concurrency: int,
app_id: str,
) -> None:
"""Scan the joined request document and write MASK back to ``texts``.
The first ``n_text`` pieces are maskable message-text parts; the rest are
detection-only tool-call args and tool/function descriptions. On MASK we
reconstruct per-part masked text by aligning the join against the masked
document and write the recovered message-text parts back to
``inputs["texts"]`` (positional write-back / "Path B"): the handler
already maps that list onto the right message parts, so there is no need
to rebuild ``structured_messages``. Reconstruction failure fails closed
(block) rather than misassigning a redaction.
"""
document = JOINER.join(pieces)
logger.debug(
"Alice WonderFence (apply_guardrail request): scanning joined document of %d piece(s) "
"(%d text + %d detection-only), %d chars, guardrail=%s app_id=%s",
len(pieces),
n_text,
len(pieces) - n_text,
len(document),
self.guardrail_name,
app_id,
)
verdict = (await evaluate_segments([document], evaluate, max_concurrency=max_concurrency))[0]
raise_if_blocked([verdict], self.guardrail_name, self.block_message)
correlation_id = verdict.correlation_ids[0] if verdict.correlation_ids else None
if verdict.action == "MASK":
recovered = reconstruct(pieces, verdict.masked_text or "")
if recovered is None:
logger.warning(
"Alice WonderFence (apply_guardrail request): MASK reconstruction failed "
"(a joiner or part boundary landed inside a masked span); failing closed. guardrail=%s correlation_id=%s",
self.guardrail_name,
correlation_id,
)
raise WonderFenceBlockedError(block_detail([verdict], self.guardrail_name, self.block_message))
inputs["texts"] = recovered[:n_text]
logger.info(
"Alice WonderFence (apply_guardrail request): MASK applied to request text guardrail=%s correlation_id=%s",
self.guardrail_name,
correlation_id,
)
elif verdict.action == "DETECT":
logger.warning(
"Alice WonderFence (apply_guardrail request): DETECT on joined document guardrail=%s correlation_id=%s",
self.guardrail_name,
correlation_id,
)
async def _scan_response(
self,
inputs: GenericGuardrailAPIInputs,
texts: list[str],
tool_indices: list[int],
tool_arg_segments: list[str],
evaluate: "Callable[[str], Awaitable[object]]",
max_concurrency: int,
) -> None:
"""Scan response segments per-index and write masks back in place."""
segments = [*texts, *tool_arg_segments]
logger.debug(
"Alice WonderFence (apply_guardrail response): evaluating %d text + %d tool-call segment(s) guardrail=%s",
len(texts),
len(tool_arg_segments),
self.guardrail_name,
)
verdicts = await evaluate_segments(segments, evaluate, max_concurrency=max_concurrency)
n_text = len(texts)
apply_response_verdicts(
inputs,
verdicts[:n_text],
tool_indices,
verdicts[n_text:],
self.guardrail_name,
self.block_message,
)
@staticmethod
def get_config_model() -> type["GuardrailConfigModel"] | None:
"""Return the config model for UI rendering."""

View file

@ -36,13 +36,15 @@ class SegmentVerdict:
class WindowConfig:
"""Tuning for the detection-only overlap windows.
``overlap`` sizes the chunk- and segment-boundary windows; ``text_segment_count``
is how many leading segments are ordered prompt texts the model concatenates,
bounding the cross-segment windows (see ``_cross_segment_windows``).
``overlap`` sizes the per-segment chunk-boundary windows (see
``_boundary_windows``). There are no cross-segment windows: on the request
side message parts are concatenated into one joined document before
scanning (so their junctions are interior chunk seams, covered by
``_boundary_windows``); on the response side each segment is an independent
choice or tool-call arg that the model never concatenates.
"""
overlap: int = CHUNK_OVERLAP_CHARS
text_segment_count: int = 0
def _split_text(text: str, max_chars: int) -> list[str]:
@ -91,28 +93,6 @@ def _boundary_windows(chunks: list[str], overlap: int) -> list[str]:
return [chunks[i][-overlap:] + chunks[i + 1][:overlap] for i in range(len(chunks) - 1)]
def _cross_segment_windows(segments: list[str], text_segment_count: int, overlap: int) -> list[tuple[int, str]]:
"""Detection-only windows spanning each adjacent pair of prompt-text segments.
The chat translation layer emits each message content part as its own
``texts`` entry, but the model concatenates them (a multimodal message's text
parts join with no separator at all), so a blocked phrase split across two
adjacent segments is seen whole by neither. We also scan a window joining the
tail of one to the head of the next. Only the first ``text_segment_count``
segments (the ordered prompt texts) are paired; tool-call args and tool /
function definitions are not concatenated into the prompt. Each window is
tagged with its left segment index so a BLOCK/DETECT folds into that
segment's verdict; windows never mask, since content cannot be redacted
across a segment boundary.
"""
if overlap <= 0:
return []
n = min(text_segment_count, len(segments))
return [
(i, segments[i][-overlap:] + segments[i + 1][:overlap]) for i in range(n - 1) if segments[i] and segments[i + 1]
]
def _aggregate(
chunks: list[str],
chunk_results: list[Any],
@ -155,15 +135,16 @@ async def evaluate_segments(
Each segment is split into <= ``max_chars`` disjoint chunks; multi-chunk
segments also get a detection-only window spanning each chunk boundary (see
``_boundary_windows``). Adjacent prompt-text segments (the first
``windows.text_segment_count``) additionally get a detection-only window
spanning their junction (see ``_cross_segment_windows``) so a phrase split
across two segments is still seen whole. Every chunk and window 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; masking uses the
disjoint chunks only so the lossless rejoin holds.
``_boundary_windows``) so a phrase split across a chunk seam is still seen
whole. Every chunk and window 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; masking uses the disjoint chunks only so
the lossless rejoin holds.
The request side passes a single joined document here (one segment) so the
common case is one call; the response side passes one segment per choice /
tool-call arg. There is no cross-segment window (see ``WindowConfig``).
"""
semaphore = asyncio.Semaphore(max_concurrency)
@ -175,7 +156,6 @@ async def evaluate_segments(
ov = min(windows.overlap, max_chars // 2)
seg_chunks = [_split_text(s, max_chars) for s in segments]
seg_boundaries = [_boundary_windows(chunks, ov) for chunks in seg_chunks]
cross_windows = _cross_segment_windows(segments, windows.text_segment_count, ov)
index: list[tuple[str, int, int]] = []
tasks = []
@ -186,9 +166,6 @@ async def evaluate_segments(
for bi, window in enumerate(seg_boundaries[si]):
index.append(("bound", si, bi))
tasks.append(run(window))
for left_idx, window in cross_windows:
index.append(("cross", left_idx, 0))
tasks.append(run(window))
results = await asyncio.gather(*tasks)
chunk_res: list[list[Any]] = [[None] * len(c) for c in seg_chunks]
@ -198,8 +175,5 @@ async def evaluate_segments(
chunk_res[si][idx] = res
elif kind == "bound":
bound_res[si][idx] = res
cross_res: list[list[Any]] = [
[res for (kind, si, _), res in zip(index, results) if kind == "cross" and si == s] for s in range(len(segments))
]
return [_aggregate(seg_chunks[si], chunk_res[si], bound_res[si] + cross_res[si]) for si in range(len(segments))]
return [_aggregate(seg_chunks[si], chunk_res[si], bound_res[si]) for si in range(len(segments))]

View file

@ -7,6 +7,8 @@
# ALICE_API_KEY - Default WonderFence API key (overridable per request)
# ALICE_MAX_CACHED_CLIENTS - Optional: max cached V2 SDK clients (default 10)
# ALICE_CONNECTION_POOL_LIMIT - Optional: HTTP pool size per client
# ALICE_MAX_SCAN_CHARS - Optional: total-work cap, max scan chars (default 1000000)
# ALICE_MAX_SCAN_SEGMENTS - Optional: total-work cap, max scan segments (default 1000)
# OPENAI_API_KEY - API key for OpenAI
#
# Per-key / per-team metadata keys (admin-controlled):
@ -49,6 +51,13 @@ guardrails:
# the matching knob for tool messages.
skip_system_message_in_guardrail: true
# Fail-closed total-work cap (DoS backstop). WonderFence has no batch API,
# so the request side joins all scan pieces into one document and scans it
# in ~1 call (chunked only past the size limit); these bounds reject an
# abusive request before any provider call and are never fail-open.
# max_scan_chars: 1000000
# max_scan_segments: 1000
# connection_pool_limit: 20
# Enable only for trusted-gateway deployments that need to forward a

View file

@ -11,3 +11,17 @@ class WonderFenceBlockedError(Exception):
def __init__(self, detail: dict):
self.detail = detail
super().__init__(detail.get("error", "Blocked by Alice WonderFence guardrail"))
class WonderFenceScanBudgetExceeded(Exception):
"""Raised when a request/response exceeds the configured total-work cap.
A fail-closed configuration/abuse guard, never a transport failure: it is
mapped to HTTP 400 before any WonderFence call and is not subject to
``fail_open`` (a caller must not be able to bypass scanning by overflowing
the cap).
"""
def __init__(self, detail: dict):
self.detail = detail
super().__init__(detail.get("error", "Alice WonderFence scan budget exceeded"))

View file

@ -1,6 +1,10 @@
"""Pure transforms for Alice WonderFence: context build, user-text mapping, verdict apply."""
"""Pure transforms for Alice WonderFence: context build, scan-piece gathering,
joined-document masked reconstruction, response-side verdict apply, total-work cap."""
from typing import Any, Callable
from collections.abc import Sequence
from difflib import SequenceMatcher
from itertools import accumulate
from typing import Callable
import litellm
from litellm._logging import verbose_proxy_logger
@ -8,10 +12,12 @@ from litellm.types.utils import GenericGuardrailAPIInputs
from .chunked_evaluation import SegmentVerdict
from .credentials import get_metadata
from .exceptions import WonderFenceBlockedError
from .exceptions import WonderFenceBlockedError, WonderFenceScanBudgetExceeded
logger = verbose_proxy_logger.getChild("alice_wonderfence")
JOINER = "\n"
def build_analysis_context(
request_data: dict,
@ -54,7 +60,10 @@ def tool_call_arg_segments(
``inputs["tool_calls"]`` entries are dicts shaped
``{"function": {"arguments": "<json string>"}}``; the argument string is the
caller- or model-controlled payload that reaches the model/client, so it is
scanned like any other segment.
scanned. ``indices`` is only used on the response side, where a MASK verdict
is written back in place; on the request side the argument strings are
appended to the joined document as detection-only pieces (see
``apply_guardrail``).
"""
tool_calls = inputs.get("tool_calls") or []
indices: list[int] = []
@ -68,84 +77,156 @@ def tool_call_arg_segments(
return indices, segments
def _description_strings(root: object, root_prefix: list[Any]) -> list[tuple[list[Any], str]]:
"""Collect ``(path, text)`` for every non-blank ``description`` string under
``root`` (a tool's ``function`` dict), walking nested JSON-schema parameters
so parameter descriptions are included, not just the top one.
def _description_texts(fn: object) -> list[str]:
"""Collect every non-blank ``description`` string under a tool's ``function``
dict, walking nested JSON-schema parameters so parameter descriptions are
included, not just the top one.
Iterative (explicit stack) rather than recursive: caller-supplied tool
schemas can nest arbitrarily, and unbounded recursion on request input is a
DoS / stack-overflow risk.
DoS / stack-overflow risk. Only the strings are returned (no write-back
paths): tool/function descriptions are scanned detection-only, so there is
nothing to mask back into the schema.
"""
out: list[tuple[list[Any], str]] = []
stack: list[tuple[Any, list[Any]]] = [(root, root_prefix)]
out: list[str] = []
stack: list[object] = [fn]
while stack:
obj, prefix = stack.pop()
obj = stack.pop()
if isinstance(obj, dict):
for key, value in obj.items():
if key == "description" and isinstance(value, str) and value.strip():
out.append((prefix + [key], value))
out.append(value)
elif isinstance(value, (dict, list)):
stack.append((value, prefix + [key]))
stack.append(value)
elif isinstance(obj, list):
for idx, item in enumerate(obj):
if isinstance(item, (dict, list)):
stack.append((item, prefix + [idx]))
stack.extend(item for item in obj if isinstance(item, (dict, list)))
return out
def tool_definition_segments(
inputs: GenericGuardrailAPIInputs,
) -> tuple[list[list[Any]], list[str]]:
"""Return (paths, texts) for free-text in tool definitions.
def tool_definition_segments(inputs: GenericGuardrailAPIInputs) -> list[str]:
"""Return description texts from ``inputs["tools"]`` (detection-only).
The chat translation layer passes caller-supplied ``inputs["tools"]`` to the
model verbatim, so a tool's ``function.description`` and its nested parameter
descriptions are scanned like any other request segment. Each path locates
the string within ``inputs["tools"]`` so a MASK verdict can be written back.
descriptions are scanned. Detection-only: they are rendered into the joined
document as extra pieces and can BLOCK/DETECT but are never masked back
(there is no faithful place to splice a redaction into a schema).
"""
tools = inputs.get("tools") or []
paths: list[list[Any]] = []
segments: list[str] = []
for i, tool in enumerate(tools):
fn = tool.get("function") if isinstance(tool, dict) else None
if not isinstance(fn, dict):
continue
for sub_path, text in _description_strings(fn, ["function"]):
paths.append([i, *sub_path])
segments.append(text)
return paths, segments
return [
text
for tool in tools
if isinstance(tool, dict) and isinstance(tool.get("function"), dict)
for text in _description_texts(tool["function"])
]
def function_definition_segments(
request_data: dict,
) -> tuple[list[list[Any]], list[str]]:
"""Description paths and texts from the deprecated ``functions[]`` parameter.
def function_definition_segments(request_data: dict) -> list[str]:
"""Return description texts from the deprecated top-level ``functions[]``.
Each entry is shaped like a tool's ``function`` object so the same
description walker applies. Returns ``(paths, segments)`` so MASK verdicts
can be written back into ``request_data["functions"]`` the same way
``tool_def_paths`` are used for ``inputs["tools"]``.
Each entry is shaped like a tool's ``function`` object, so the same
description walker applies. Detection-only, same rationale as
``tool_definition_segments``.
"""
functions = request_data.get("functions") or []
paths: list[list[Any]] = []
segments: list[str] = []
for i, fn in enumerate(functions):
if isinstance(fn, dict):
for sub_path, text in _description_strings(fn, []):
paths.append([i, *sub_path])
segments.append(text)
return paths, segments
return [text for fn in functions if isinstance(fn, dict) for text in _description_texts(fn)]
def _set_by_path(root: Any, path: list[Any], value: object) -> None:
obj = root
for key in path[:-1]:
obj = obj[key]
obj[path[-1]] = value
def check_scan_budget(
segments: list[str],
max_scan_chars: int | None,
max_scan_segments: int | None,
) -> None:
"""Fail-closed total-work cap; raises ``WonderFenceScanBudgetExceeded`` when
the combined scan characters or segment count exceed the configured limits.
WonderFence has no batch API (one HTTP POST per string), so without a cap a
single crafted request with thousands of tiny parts or huge content could
amplify into an unbounded number of upstream calls. This check runs before
any WonderFence call and is never subject to ``fail_open`` — a caller must
not be able to bypass scanning by overflowing the cap.
"""
n = len(segments)
if max_scan_segments is not None and n > max_scan_segments:
raise WonderFenceScanBudgetExceeded(
{
"error": f"Alice WonderFence scan budget exceeded: {n} segments > max_scan_segments={max_scan_segments}",
"type": "alice_wonderfence_scan_budget_exceeded",
"limit": "max_scan_segments",
"max_scan_segments": max_scan_segments,
"segments": n,
}
)
total_chars = sum(len(s) for s in segments)
if max_scan_chars is not None and total_chars > max_scan_chars:
raise WonderFenceScanBudgetExceeded(
{
"error": f"Alice WonderFence scan budget exceeded: {total_chars} chars > max_scan_chars={max_scan_chars}",
"type": "alice_wonderfence_scan_budget_exceeded",
"limit": "max_scan_chars",
"max_scan_chars": max_scan_chars,
"chars": total_chars,
}
)
def _block_detail(blocked: list[SegmentVerdict], guardrail_name: str, block_message: str) -> dict:
def _map_index(x: int, ops: Sequence[tuple[str, int, int, int, int]], masked_len: int) -> int | None:
"""Map an index in the original joined document to its index in ``masked``.
Uses the ``SequenceMatcher`` opcodes: an index inside (or at the end of) an
``equal`` block maps positionally; a boundary that lands at the very start
of a changed block still maps (the range simply begins there); a boundary
that lands *inside* a changed block is ambiguous and returns ``None`` so the
caller fails closed rather than misassigning.
"""
for tag, i1, i2, j1, _j2 in ops:
if i1 <= x < i2 or (x == i2 and tag == "equal"):
if tag == "equal":
return j1 + (x - i1)
return j1 if x == i1 else None
return masked_len
def reconstruct(parts: list[str], masked: str) -> list[str] | None:
"""Recover per-part masked text from the masked joined document.
``parts`` were joined with ``JOINER`` (a plain ``"\\n"``) into the document
that was scanned; ``masked`` is the service's masked version of that same
document. We align original-vs-masked with ``difflib.SequenceMatcher`` (no
sentinel injected) and map each part's char range through the alignment.
Fails closed (returns ``None``) when the structure is not recoverable: every
``JOINER`` between parts must survive the mask as an unmodified ``\\n`` (a
mask spanning a joiner would merge parts), and no part boundary may land
inside a changed block. Returns one masked string per input part, in order;
``[]`` for no parts. Assumes masking is span substitution that preserves the
non-masked characters; if the service reflows whitespace the joiner-survival
check trips and we fail closed rather than misassign.
"""
if not parts:
return []
original = JOINER.join(parts)
starts = [0, *accumulate(len(p) + len(JOINER) for p in parts)][: len(parts)]
ranges = [(s, s + len(p)) for s, p in zip(starts, parts)]
joiners = [end for (_s, end) in ranges[:-1]]
ops = SequenceMatcher(None, original, masked, autojunk=False).get_opcodes()
joiner_survives = all(
any(tag == "equal" and i1 <= j < i2 and masked[j1 + (j - i1)] == JOINER for tag, i1, i2, j1, _j2 in ops)
for j in joiners
)
if not joiner_survives:
return None
mapped = [(_map_index(s, ops, len(masked)), _map_index(e, ops, len(masked))) for s, e in ranges]
if any(ms is None or me is None or ms > me for ms, me in mapped):
return None
return [masked[ms:me] for ms, me in mapped]
def block_detail(blocked: list[SegmentVerdict], guardrail_name: str, block_message: str) -> dict:
detections: list = []
correlation_ids: list[str] = []
for v in blocked:
@ -164,6 +245,14 @@ def _block_detail(blocked: list[SegmentVerdict], guardrail_name: str, block_mess
return detail
def raise_if_blocked(verdicts: list[SegmentVerdict], guardrail_name: str, block_message: str) -> None:
"""Raise ``WonderFenceBlockedError`` if any verdict is BLOCK, aggregating
detections / correlation ids across all blocked verdicts."""
blocked = [v for v in verdicts if v.action == "BLOCK"]
if blocked:
raise WonderFenceBlockedError(block_detail(blocked, guardrail_name, block_message))
def _masked_value(verdict: SegmentVerdict, guardrail_name: str, label: str) -> str | None:
"""Return the replacement string for a MASK verdict (logging as a side
effect), or None for DETECT/NO_ACTION. The caller writes it to the slot the
@ -187,51 +276,28 @@ def _masked_value(verdict: SegmentVerdict, guardrail_name: str, label: str) -> s
return None
def apply_verdicts(
def apply_response_verdicts(
inputs: GenericGuardrailAPIInputs,
indices: list[int],
verdicts: list[SegmentVerdict],
text_verdicts: list[SegmentVerdict],
tool_indices: list[int],
tool_verdicts: list[SegmentVerdict],
guardrail_name: str,
block_message: str,
tool_indices: list[int] | None = None,
tool_verdicts: list[SegmentVerdict] | None = None,
tool_def_paths: list[list[Any]] | None = None,
tool_def_verdicts: list[SegmentVerdict] | None = None,
function_def_paths: list[list[Any]] | None = None,
function_def_verdicts: list[SegmentVerdict] | None = None,
function_def_request_data: dict | None = None,
) -> GenericGuardrailAPIInputs:
"""Apply per-segment verdicts back onto request text, tool-call args,
tool-definition descriptions, and legacy function-definition descriptions.
"""Response-side write-back: index-aligned MASK into ``texts`` (per choice)
and ``tool_calls[i].function.arguments`` (model-generated).
Any BLOCK across any group raises ``WonderFenceBlockedError`` with
detections/correlation ids aggregated across all blocked segments. Otherwise
each MASK verdict rewrites the slot its segment came from and DETECT is
logged.
BLOCK across any segment raises first. Response text is written per-index
(never joined) because the handler's response write-back is purely
positional over the returned ``texts`` list and has no ``structured_messages``
path, so collapsing choices into one string would dump every choice's text
into choice 0.
"""
tool_indices = tool_indices or []
tool_verdicts = tool_verdicts or []
tool_def_paths = tool_def_paths or []
tool_def_verdicts = tool_def_verdicts or []
function_def_paths = function_def_paths or []
function_def_verdicts = function_def_verdicts or []
blocked = [
v
for v in (
*verdicts,
*tool_verdicts,
*tool_def_verdicts,
*function_def_verdicts,
)
if v.action == "BLOCK"
]
if blocked:
raise WonderFenceBlockedError(_block_detail(blocked, guardrail_name, block_message))
raise_if_blocked([*text_verdicts, *tool_verdicts], guardrail_name, block_message)
texts = inputs.get("texts") or []
for idx, verdict in zip(indices, verdicts):
masked = _masked_value(verdict, guardrail_name, "request text")
for idx, verdict in enumerate(text_verdicts):
masked = _masked_value(verdict, guardrail_name, "response text")
if masked is not None:
texts[idx] = masked
inputs["texts"] = texts
@ -242,16 +308,4 @@ def apply_verdicts(
if masked is not None:
tool_calls[idx]["function"]["arguments"] = masked
tools = inputs.get("tools") or []
for path, verdict in zip(tool_def_paths, tool_def_verdicts):
masked = _masked_value(verdict, guardrail_name, "tool definition")
if masked is not None:
_set_by_path(tools, path, masked)
functions = (function_def_request_data or {}).get("functions") or []
for path, verdict in zip(function_def_paths, function_def_verdicts):
masked = _masked_value(verdict, guardrail_name, "function definition")
if masked is not None and functions:
_set_by_path(functions, path, masked)
return inputs

View file

@ -62,6 +62,14 @@ class WonderFenceGuardrailConfigModel(GuardrailConfigModel):
default=None,
description="Max connections per SDK client HTTP pool. Env: ALICE_CONNECTION_POOL_LIMIT.",
)
max_scan_chars: Optional[int] = Field(
default=1_000_000,
description="Total-work cap (fail-closed DoS backstop): reject a request/response whose combined scan characters (message text plus tool-call args and tool/function descriptions) exceed this before any WonderFence call. Bounds upstream call amplification since WonderFence has no batch API. Env: ALICE_MAX_SCAN_CHARS.",
)
max_scan_segments: Optional[int] = Field(
default=1_000,
description="Total-work cap (fail-closed DoS backstop): reject a request/response carrying more than this many scan segments (message text parts, tool-call args, tool/function descriptions) before any WonderFence call. Env: ALICE_MAX_SCAN_SEGMENTS.",
)
@staticmethod
def ui_friendly_name() -> str:

View file

@ -97,15 +97,17 @@ async def test_apply_guardrail_mask_replaces_scanned_text(guardrail_and_client,
@pytest.mark.asyncio
async def test_apply_guardrail_mask_targets_only_the_flagged_slot(guardrail_and_client, make_request_data):
"""MASK rewrites the ``texts`` entry of the flagged segment in place; the
other scanned entries survive untouched. Confirms positional 1:1 mapping."""
async def test_apply_guardrail_mask_reconstructs_only_the_flagged_message_part(guardrail_and_client, make_request_data):
"""Request side joins the message parts into one document, scans once, and on
MASK reconstructs per-part masked text by aligning the join against the
masked document. Only the flagged part changes; the others survive."""
guardrail, client = guardrail_and_client
def evaluate(prompt, **kwargs):
r = Mock()
r.action = "MASK" if prompt == "sensitive content" else "NO_ACTION"
r.action_text = "[REDACTED]"
# One joined call: mask just the sensitive part inside the joined doc.
r.action = "MASK"
r.action_text = prompt.replace("sensitive content", "[REDACTED]")
r.detections = []
r.correlation_id = None
return r
@ -118,6 +120,7 @@ async def test_apply_guardrail_mask_targets_only_the_flagged_slot(guardrail_and_
input_type="request",
)
assert out["texts"] == ["first", "ack", "[REDACTED]"]
client.evaluate_prompt.assert_awaited_once()
@pytest.mark.asyncio
@ -130,7 +133,7 @@ async def test_apply_guardrail_scans_non_user_role_segments(guardrail_and_client
def evaluate(prompt, **kwargs):
r = Mock()
r.action = "BLOCK" if prompt == "disallowed system instruction" else "NO_ACTION"
r.action = "BLOCK" if "disallowed system instruction" in prompt else "NO_ACTION"
r.detections = []
r.correlation_id = None
return r
@ -274,11 +277,10 @@ async def test_apply_guardrail_response_path_passes_app_id(make_guardrail, make_
@pytest.mark.asyncio
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."""
async def test_apply_guardrail_joins_all_message_parts_into_one_call(guardrail_and_client, make_request_data):
"""Every message part is scanned, but as a single joined document in ONE
Alice call (call volume scales with size, not message count). The join uses
a plain newline so cross-part content is seen whole."""
guardrail, client = guardrail_and_client
result_obj = Mock()
result_obj.action = "NO_ACTION"
@ -291,10 +293,8 @@ async def test_apply_guardrail_evaluates_every_text_without_structured_messages(
request_data=make_request_data(),
input_type="request",
)
prompts = {c.kwargs["prompt"] for c in client.evaluate_prompt.call_args_list}
assert {"t1", "t2", "t3"} <= prompts
# Adjacent text segments also get a cross-segment junction window each.
assert {"t1t2", "t2t3"} <= prompts
client.evaluate_prompt.assert_awaited_once()
assert client.evaluate_prompt.call_args.kwargs["prompt"] == "t1\nt2\nt3"
@pytest.mark.asyncio
@ -306,7 +306,7 @@ async def test_apply_guardrail_blocks_on_earlier_user_turn(guardrail_and_client,
def evaluate(prompt, **kwargs):
r = Mock()
r.action = "BLOCK" if prompt == "disallowed" else "NO_ACTION"
r.action = "BLOCK" if "disallowed" in prompt else "NO_ACTION"
r.detections = []
r.correlation_id = None
return r
@ -376,3 +376,125 @@ async def test_apply_guardrail_no_text_short_circuits(guardrail_and_client, make
assert out == {"texts": []}
client.evaluate_prompt.assert_not_awaited()
client.evaluate_response.assert_not_awaited()
# ----------------------------- join: cross-part visibility -----------------------------
@pytest.mark.asyncio
async def test_apply_guardrail_blocks_phrase_split_across_message_parts(guardrail_and_client, make_request_data):
"""Two content parts that individually look benign are joined into one
document, so a phrase split across the part boundary is seen in a single
scan and still BLOCKs (the join replaces the old cross-segment windows)."""
guardrail, client = guardrail_and_client
def evaluate(prompt, **kwargs):
r = Mock()
# Neither part alone contains the whole phrase; the joined document does.
r.action = "BLOCK" if ("make a b" in prompt and "omb" 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": ["how to make a b", "omb please"]},
request_data=make_request_data(),
input_type="request",
)
assert exc.value.status_code == 400
client.evaluate_prompt.assert_awaited_once()
# ----------------------------- join: MASK reconstruction write-back -----------------------------
@pytest.mark.asyncio
async def test_apply_guardrail_mask_writes_back_to_the_correct_message(guardrail_and_client, make_request_data):
"""A real PII MASK on the joined document reconstructs per-part masked text
and writes it back to the message that carried the PII, leaving the others
intact."""
guardrail, client = guardrail_and_client
def evaluate(prompt, **kwargs):
r = Mock()
r.action = "MASK"
r.action_text = prompt.replace("john@example.com", "[EMAIL]")
r.detections = []
r.correlation_id = "corr-mask"
return r
client.evaluate_prompt.side_effect = evaluate
out = await guardrail.apply_guardrail(
inputs={"texts": ["hello there", "my email is john@example.com", "thanks"]},
request_data=make_request_data(),
input_type="request",
)
assert out["texts"] == ["hello there", "my email is [EMAIL]", "thanks"]
client.evaluate_prompt.assert_awaited_once()
@pytest.mark.asyncio
async def test_apply_guardrail_mask_reconstruction_failure_fails_closed(guardrail_and_client, make_request_data):
"""If masking destroys a joiner (parts would merge), reconstruction cannot
safely attribute the redaction, so the request is blocked rather than
silently misassigned or passed through unmasked."""
guardrail, client = guardrail_and_client
def evaluate(prompt, **kwargs):
r = Mock()
r.action = "MASK"
r.action_text = prompt.replace("\n", "") # destroys the joiner -> parts merge
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": ["alpha", "beta"]},
request_data=make_request_data(),
input_type="request",
)
assert exc.value.status_code == 400
# ----------------------------- total-work cap -----------------------------
@pytest.mark.asyncio
async def test_apply_guardrail_rejects_over_segment_cap_without_scanning(make_guardrail, make_request_data):
guardrail, client = make_guardrail(max_scan_segments=3)
guardrail._client_cache["default-api-key"] = client
with pytest.raises(HTTPException) as exc:
await guardrail.apply_guardrail(
inputs={"texts": ["a", "b", "c", "d"]},
request_data=make_request_data(),
input_type="request",
)
assert exc.value.status_code == 400
assert exc.value.detail["limit"] == "max_scan_segments"
client.evaluate_prompt.assert_not_awaited()
@pytest.mark.asyncio
async def test_apply_guardrail_cap_is_not_bypassed_by_fail_open(make_guardrail, make_request_data):
"""The cap is a config/abuse guard, never fail-open: an oversized request
is rejected 400 even with fail_open=True and the SDK is never called."""
guardrail, client = make_guardrail(max_scan_chars=10, fail_open=True)
guardrail._client_cache["default-api-key"] = client
with pytest.raises(HTTPException) as exc:
await guardrail.apply_guardrail(
inputs={"texts": ["x" * 50]},
request_data=make_request_data(),
input_type="request",
)
assert exc.value.status_code == 400
assert exc.value.detail["limit"] == "max_scan_chars"
client.evaluate_prompt.assert_not_awaited()

View file

@ -44,15 +44,17 @@ async def test_apply_guardrail_blocks_on_tool_call_arguments(guardrail_and_clien
@pytest.mark.asyncio
async def test_apply_guardrail_masks_tool_call_arguments_in_place(guardrail_and_client, make_request_data):
"""MASK on a tool-call argument string rewrites
inputs['tool_calls'][i]['function']['arguments']."""
async def test_apply_guardrail_request_tool_call_args_are_detection_only(guardrail_and_client, make_request_data):
"""On the request side, tool-call args are rendered into the joined document
as detection-only pieces: they can BLOCK/DETECT but a MASK is never spliced
back into the arguments string (the joined form is not the wire format).
Message text still masks; the args survive untouched."""
guardrail, client = guardrail_and_client
def evaluate(prompt, **kwargs):
r = Mock()
r.action = "MASK" if "secret" in prompt else "NO_ACTION"
r.action_text = '{"body": "[REDACTED]"}'
r.action = "MASK"
r.action_text = prompt.replace("secret value", "[REDACTED]")
r.detections = []
r.correlation_id = None
return r
@ -68,8 +70,9 @@ async def test_apply_guardrail_masks_tool_call_arguments_in_place(guardrail_and_
request_data=make_request_data(),
input_type="request",
)
assert out["tool_calls"][0]["function"]["arguments"] == '{"body": "[REDACTED]"}'
assert out["tool_calls"][0]["function"]["arguments"] == '{"body": "secret value"}'
assert out["texts"] == ["benign"]
client.evaluate_prompt.assert_awaited_once()
@pytest.mark.asyncio
@ -208,13 +211,15 @@ async def test_apply_guardrail_blocks_on_tool_parameter_description(guardrail_an
@pytest.mark.asyncio
async def test_apply_guardrail_masks_tool_definition_description_in_place(guardrail_and_client, make_request_data):
async def test_apply_guardrail_tool_definitions_are_detection_only(guardrail_and_client, make_request_data):
"""Tool definitions are scanned detection-only (they can BLOCK/DETECT) but a
MASK is never written back into the schema; the description survives."""
guardrail, client = guardrail_and_client
def evaluate(prompt, **kwargs):
r = Mock()
r.action = "MASK" if "secret" in prompt else "NO_ACTION"
r.action_text = "[REDACTED]"
r.action = "MASK"
r.action_text = prompt.replace("secret", "[REDACTED]")
r.detections = []
r.correlation_id = None
return r
@ -226,7 +231,7 @@ async def test_apply_guardrail_masks_tool_definition_description_in_place(guardr
"tools": [_tool_def(description="contains secret stuff")],
}
out = await guardrail.apply_guardrail(inputs=inputs, request_data=make_request_data(), input_type="request")
assert out["tools"][0]["function"]["description"] == "[REDACTED]"
assert out["tools"][0]["function"]["description"] == "contains secret stuff"
@pytest.mark.asyncio
@ -361,15 +366,15 @@ async def test_apply_guardrail_legacy_function_detect_does_not_mutate(guardrail_
@pytest.mark.asyncio
async def test_apply_guardrail_masks_legacy_function_description_in_place(guardrail_and_client, make_request_data):
"""A MASK verdict on a functions[] description must be written back into
request_data['functions'], not left as the original unredacted text."""
async def test_apply_guardrail_legacy_function_definitions_are_detection_only(guardrail_and_client, make_request_data):
"""Legacy functions[] descriptions are scanned detection-only; a MASK is
never spliced back into request_data['functions']."""
guardrail, client = guardrail_and_client
def evaluate(prompt, **kwargs):
r = Mock()
r.action = "MASK" if "secret" in prompt else "NO_ACTION"
r.action_text = "[REDACTED]"
r.action = "MASK"
r.action_text = prompt.replace("secret", "[REDACTED]")
r.detections = []
r.correlation_id = None
return r
@ -382,4 +387,4 @@ async def test_apply_guardrail_masks_legacy_function_description_in_place(guardr
request_data=request_data,
input_type="request",
)
assert request_data["functions"][0]["description"] == "[REDACTED]"
assert request_data["functions"][0]["description"] == "contains secret stuff"

View file

@ -230,74 +230,15 @@ async def test_boundary_window_mask_is_surfaced_as_detect_not_dropped():
assert verdicts[0].action == "DETECT"
# ----------------------------- cross-segment overlap (split across adjacent texts) -----------------------------
@pytest.mark.asyncio
async def test_block_phrase_split_across_adjacent_text_segments_is_detected():
"""A blocked phrase split across two adjacent prompt-text segments (e.g. two
content parts of one message, which the model concatenates) is caught by the
cross-segment window even though neither segment contains it whole. Fails
without cross-segment windows -> the phrase evades scanning."""
async def evaluate(text):
return _result("BLOCK" if "BLOCKME" in text else "")
verdicts = await evaluate_segments(["BLOCK", "ME"], evaluate, windows=WindowConfig(text_segment_count=2))
assert verdicts[0].action == "BLOCK"
@pytest.mark.asyncio
async def test_without_text_segment_count_split_phrase_evades():
"""Control: with no declared text segments there is no cross-segment window,
so the same split phrase is seen by neither segment. Demonstrates the gap the
cross-segment window closes."""
async def test_independent_segments_are_not_concatenated():
"""Segments are scanned independently (no cross-segment window): a phrase
split across two segments is NOT joined. On the request side, message parts
are joined into one document *before* reaching here; the response side has
independent choices that the model never concatenates."""
async def evaluate(text):
return _result("BLOCK" if "BLOCKME" in text else "")
verdicts = await evaluate_segments(["BLOCK", "ME"], evaluate)
assert [v.action for v in verdicts] == ["", ""]
@pytest.mark.asyncio
async def test_cross_segment_window_stays_within_text_segments():
"""Only the first text_segment_count segments are paired; a trailing
non-text segment (tool-call args, tool/function definition) is never joined
with the last prompt text, so a phrase straddling that junction does not
block."""
async def evaluate(text):
return _result("BLOCK" if "BLOCKME" in text else "")
verdicts = await evaluate_segments(["BLOCK", "ME"], evaluate, windows=WindowConfig(text_segment_count=1))
assert [v.action for v in verdicts] == ["", ""]
@pytest.mark.asyncio
async def test_cross_segment_window_surfaces_mask_as_detect_without_masking():
"""A cross-segment window cannot redact across the segment boundary, so a
MASK on it surfaces as DETECT and never rewrites the segment text."""
async def evaluate(text):
return _result("MASK", action_text="[X]") if "SECRETHERE" in text else _result("")
verdicts = await evaluate_segments(["SECRET", "HERE"], evaluate, windows=WindowConfig(text_segment_count=2))
assert verdicts[0].action == "DETECT"
assert verdicts[0].masked_text is None
@pytest.mark.asyncio
async def test_cross_segment_window_joins_segment_tail_and_head():
"""The window spans the junction (tail of one segment + head of the next),
catching a phrase that lives only across the boundary of longer segments."""
async def evaluate(text):
return _result("BLOCK" if "a bomb" in text else "")
verdicts = await evaluate_segments(
["how to make a b", "omb please"],
evaluate,
windows=WindowConfig(overlap=6, text_segment_count=2),
)
assert verdicts[0].action == "BLOCK"

View file

@ -1,4 +1,4 @@
"""Tests for processing.py pure transforms: verdict apply."""
"""Tests for processing.py pure transforms: reconstruction, extractors, cap, verdict apply."""
import pytest
@ -7,9 +7,15 @@ from litellm.proxy.guardrails.guardrail_hooks.alice_wonderfence.chunked_evaluati
)
from litellm.proxy.guardrails.guardrail_hooks.alice_wonderfence.exceptions import (
WonderFenceBlockedError,
WonderFenceScanBudgetExceeded,
)
from litellm.proxy.guardrails.guardrail_hooks.alice_wonderfence.processing import (
apply_verdicts,
JOINER,
apply_response_verdicts,
check_scan_budget,
function_definition_segments,
reconstruct,
tool_definition_segments,
)
@ -17,47 +23,75 @@ 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"
# --------------- reconstruct (masked-join alignment) ---------------
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_reconstruct_no_change_round_trips():
parts = ["alpha", "beta", "gamma"]
assert reconstruct(parts, JOINER.join(parts)) == parts
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"]
def test_reconstruct_masks_a_middle_part():
parts = ["alpha", "sensitive", "gamma"]
masked = JOINER.join(["alpha", "[REDACTED]", "gamma"])
assert reconstruct(parts, masked) == ["alpha", "[REDACTED]", "gamma"]
# --------------- tool_definition_segments ---------------
def test_reconstruct_mask_at_part_start():
parts = ["alpha", "beta", "gamma"]
masked = JOINER.join(["[X]lpha", "beta", "gamma"])
assert reconstruct(parts, masked) == ["[X]lpha", "beta", "gamma"]
def test_reconstruct_handles_a_part_that_itself_contains_newline():
"""A message part can itself contain the joiner char; alignment is
structural, not a naive split on '\\n', so this still reconstructs."""
parts = ["line1\nline1b", "second"]
masked = JOINER.join(["line1\n[REDACTED]", "second"])
assert reconstruct(parts, masked) == ["line1\n[REDACTED]", "second"]
def test_reconstruct_fails_closed_when_mask_spans_a_joiner():
"""If the mask swallows a joiner (parts merged), reconstruction must fail
closed (None) rather than misassign redacted text to the wrong message."""
parts = ["alpha", "beta", "gamma"]
merged = "alphaXXXbeta\ngamma" # joiner between alpha|beta is gone
assert reconstruct(parts, merged) is None
def test_reconstruct_empty_parts_is_empty_list():
assert reconstruct([], "") == []
# --------------- check_scan_budget (total-work cap) ---------------
def test_check_scan_budget_passes_within_limits():
check_scan_budget(["a", "b", "c"], max_scan_chars=100, max_scan_segments=100)
def test_check_scan_budget_rejects_too_many_segments():
with pytest.raises(WonderFenceScanBudgetExceeded) as exc:
check_scan_budget(["x"] * 11, max_scan_chars=10_000, max_scan_segments=10)
assert exc.value.detail["limit"] == "max_scan_segments"
assert exc.value.detail["max_scan_segments"] == 10
def test_check_scan_budget_rejects_too_many_chars():
with pytest.raises(WonderFenceScanBudgetExceeded) as exc:
check_scan_budget(["x" * 50, "y" * 60], max_scan_chars=100, max_scan_segments=100)
assert exc.value.detail["limit"] == "max_scan_chars"
assert exc.value.detail["chars"] == 110
def test_check_scan_budget_none_limits_disable_the_cap():
check_scan_budget(["x" * 10_000] * 100, max_scan_chars=None, max_scan_segments=None)
# --------------- tool/function definition extractors (detection-only, list of texts) ---------------
def test_tool_definition_segments_extracts_description_and_param_descriptions():
from litellm.proxy.guardrails.guardrail_hooks.alice_wonderfence.processing import (
_set_by_path,
tool_definition_segments,
)
inputs = {
"tools": [
{
@ -73,14 +107,7 @@ def test_tool_definition_segments_extracts_description_and_param_descriptions():
}
]
}
paths, segments = tool_definition_segments(inputs)
assert set(segments) == {"TOP_DESC", "PARAM_DESC"}
# each path round-trips: writing via the path updates the right slot
for path, text in zip(paths, segments):
_set_by_path(inputs["tools"], path, f"<{text}>")
fn = inputs["tools"][0]["function"]
assert fn["description"] == "<TOP_DESC>"
assert fn["parameters"]["properties"]["city"]["description"] == "<PARAM_DESC>"
assert set(tool_definition_segments(inputs)) == {"TOP_DESC", "PARAM_DESC"}
def test_tool_definition_segments_ignores_non_dict_tools_and_blank_descriptions():
@ -91,23 +118,10 @@ def test_tool_definition_segments_ignores_non_dict_tools_and_blank_descriptions(
{"type": "function"},
]
}
from litellm.proxy.guardrails.guardrail_hooks.alice_wonderfence.processing import (
tool_definition_segments,
)
paths, segments = tool_definition_segments(inputs)
assert segments == []
assert tool_definition_segments(inputs) == []
# --------------- function_definition_segments (legacy functions[]) ---------------
def test_function_definition_segments_extracts_descriptions_and_paths():
from litellm.proxy.guardrails.guardrail_hooks.alice_wonderfence.processing import (
_set_by_path,
function_definition_segments,
)
def test_function_definition_segments_extracts_descriptions():
request_data = {
"functions": [
{
@ -122,19 +136,57 @@ def test_function_definition_segments_extracts_descriptions_and_paths():
{"name": "f", "description": " "},
]
}
paths, segments = function_definition_segments(request_data)
assert set(segments) == {"TOP_DESC", "PARAM_DESC"}
for path, text in zip(paths, segments):
_set_by_path(request_data["functions"], path, f"<{text}>")
fn = request_data["functions"][0]
assert fn["description"] == "<TOP_DESC>"
assert fn["parameters"]["properties"]["city"]["description"] == "<PARAM_DESC>"
assert set(function_definition_segments(request_data)) == {"TOP_DESC", "PARAM_DESC"}
def test_function_definition_segments_empty_when_absent():
from litellm.proxy.guardrails.guardrail_hooks.alice_wonderfence.processing import (
function_definition_segments,
)
assert function_definition_segments({"model": "gpt-4"}) == []
paths, segs = function_definition_segments({"model": "gpt-4"})
assert paths == [] and segs == []
# --------------- apply_response_verdicts ---------------
def test_response_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_response_verdicts(inputs, 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_response_mask_writes_to_the_mapped_text_index_only():
inputs = {"texts": ["keep", "MASK_ME", "keep2"]}
verdicts = [
SegmentVerdict("", None, [], []),
SegmentVerdict("MASK", "[R]", [], []),
SegmentVerdict("", None, [], []),
]
out = apply_response_verdicts(inputs, verdicts, [], [], "gn", "blocked!")
assert out["texts"] == ["keep", "[R]", "keep2"]
def test_response_mask_writes_tool_call_arguments_in_place():
inputs = {
"texts": ["ok"],
"tool_calls": [{"function": {"arguments": '{"x": "secret"}'}}],
}
out = apply_response_verdicts(
inputs,
[SegmentVerdict("", None, [], [])],
[0],
[SegmentVerdict("MASK", '{"x": "[R]"}', [], [])],
"gn",
"blocked!",
)
assert out["tool_calls"][0]["function"]["arguments"] == '{"x": "[R]"}'
def test_response_detect_and_no_action_leave_texts_unchanged():
inputs = {"texts": ["a", "b"]}
verdicts = [SegmentVerdict("DETECT", None, [], []), SegmentVerdict("", None, [], [])]
out = apply_response_verdicts(inputs, verdicts, [], [], "gn", "blocked!")
assert out["texts"] == ["a", "b"]