mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
d852ae51f3
commit
8bceb610c8
11 changed files with 639 additions and 361 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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))]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue