diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/__init__.py index 65508de082a..813cf1bc57d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py index 585687cacf8..ea28b394a66 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py @@ -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.""" diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/chunked_evaluation.py b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/chunked_evaluation.py index a2adb0b9cf2..e4afeac96e8 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/chunked_evaluation.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/chunked_evaluation.py @@ -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))] diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/example_config.yaml b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/example_config.yaml index cab1e683911..3254dc65a51 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/example_config.yaml +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/example_config.yaml @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/exceptions.py b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/exceptions.py index 970a9d26fe5..2ec9b85153c 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/exceptions.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/exceptions.py @@ -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")) diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py index c63945b5bac..d4ee297a05a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py @@ -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": ""}}``; 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 diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/alice_wonderfence.py b/litellm/types/proxy/guardrails/guardrail_hooks/alice_wonderfence.py index 6785fc5fb08..e6655284fd1 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/alice_wonderfence.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/alice_wonderfence.py @@ -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: diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail.py index cb3436d281a..53d5a25d88d 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail.py @@ -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() diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail_tools.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail_tools.py index 5a1fe6a5e7e..1506236859c 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail_tools.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail_tools.py @@ -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" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_chunked_evaluation.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_chunked_evaluation.py index dcdcbf2f721..5d126ddd2fb 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_chunked_evaluation.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_chunked_evaluation.py @@ -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" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_processing.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_processing.py index b5d2c2eaf27..e8570d3f05a 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_processing.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_processing.py @@ -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"] == "" - assert fn["parameters"]["properties"]["city"]["description"] == "" + 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"] == "" - assert fn["parameters"]["properties"]["city"]["description"] == "" + 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"]