From df5cd8dab3d58a100c4807bb6ed7c9533375cdf3 Mon Sep 17 00:00:00 2001 From: lior-k Date: Wed, 17 Jun 2026 18:26:54 +0300 Subject: [PATCH] fix(guardrails): Alice WonderFence scans tool definitions The chat translation layer forwards caller-supplied inputs["tools"] to the model verbatim, but apply_guardrail only scanned message text and tool-call arguments, so blocked content in tools[].function.description or nested parameter descriptions reached the model unevaluated. Extract every description string from each tool definition (top-level and recursively through the parameters JSON schema) as a request-side segment, evaluate it alongside the others, and write a MASK verdict back to the originating slot via its path. BLOCK on any tool-def segment blocks the request; the empty-input early return accounts for tool defs too. Tool definitions are scanned by default like other request content; operators who don't want their tool schemas evaluated can scope them out upstream. Regression tests: BLOCK on a tool description, BLOCK on a nested parameter description, MASK written back to function.description in place, a tools-only request still scanned, and path round-tripping for tool_definition_segments. --- .../alice_wonderfence/alice_wonderfence.py | 21 +- .../alice_wonderfence/processing.py | 184 ++++++++++++------ .../alice_wonderfence/test_apply_guardrail.py | 120 ++++++++++++ .../alice_wonderfence/test_processing.py | 52 +++++ 4 files changed, 314 insertions(+), 63 deletions(-) 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 dfb338f1a7f..b90aa5d22c6 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py @@ -29,6 +29,7 @@ from .processing import ( apply_verdicts, build_analysis_context, tool_call_arg_segments, + tool_definition_segments, ) if TYPE_CHECKING: @@ -175,9 +176,10 @@ class WonderFenceGuardrail(CustomGuardrail): """Apply WonderFence guardrail using V2 client + per-request app_id.""" texts = inputs.get("texts") or [] tool_indices, tool_segments = tool_call_arg_segments(inputs) - if not texts and not tool_segments: + tool_def_paths, tool_def_segments = tool_definition_segments(inputs) + if not texts and not tool_segments and not tool_def_segments: logger.debug( - "Alice WonderFence (apply_guardrail): no text or tool-call args to scan for %s", + "Alice WonderFence (apply_guardrail): nothing to scan for %s", input_type, ) return inputs @@ -213,11 +215,12 @@ class WonderFenceGuardrail(CustomGuardrail): custom_fields=None, ) - segments = [*texts, *tool_segments] + segments = [*texts, *tool_segments, *tool_def_segments] logger.debug( - "Alice WonderFence (apply_guardrail): evaluating %d text + %d tool-call segment(s) app_id=%s guardrail=%s input_type=%s", + "Alice WonderFence (apply_guardrail): evaluating %d text + %d tool-call + %d tool-def segment(s) app_id=%s guardrail=%s input_type=%s", len(texts), len(tool_segments), + len(tool_def_segments), app_id, self.guardrail_name, input_type, @@ -227,14 +230,18 @@ class WonderFenceGuardrail(CustomGuardrail): evaluate, max_concurrency=self._connection_pool_limit or DEFAULT_MAX_CONCURRENCY, ) + n_text = len(texts) + n_tool = len(tool_segments) apply_verdicts( inputs, - list(range(len(texts))), - verdicts[: len(texts)], + list(range(n_text)), + verdicts[:n_text], self.guardrail_name, self.block_message, tool_indices=tool_indices, - tool_verdicts=verdicts[len(texts) :], + tool_verdicts=verdicts[n_text : n_text + n_tool], + tool_def_paths=tool_def_paths, + tool_def_verdicts=verdicts[n_text + n_tool :], ) except WonderFenceBlockedError as e: diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py index 548649cb5e2..44447839fe2 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py @@ -74,6 +74,102 @@ def tool_call_arg_segments( return indices, segments +def _description_strings(obj: Any, prefix: List[Any]) -> List[Tuple[List[Any], str]]: + """Collect ``(path, text)`` for every non-blank ``description`` string under + ``obj`` (a tool's ``function`` dict). Recurses into nested JSON-schema + parameters so parameter descriptions are included, not just the top one.""" + out: List[Tuple[List[Any], str]] = [] + 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)) + elif isinstance(value, (dict, list)): + out.extend(_description_strings(value, prefix + [key])) + elif isinstance(obj, list): + for idx, item in enumerate(obj): + if isinstance(item, (dict, list)): + out.extend(_description_strings(item, prefix + [idx])) + return out + + +def tool_definition_segments( + inputs: GenericGuardrailAPIInputs, +) -> Tuple[List[List[Any]], List[str]]: + """Return (paths, texts) for free-text in tool definitions. + + 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. + """ + 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 + + +def _set_by_path(root: Any, path: List[Any], value: Any) -> None: + obj = root + for key in path[:-1]: + obj = obj[key] + obj[path[-1]] = value + + +def _block_detail( + blocked: List[SegmentVerdict], guardrail_name: str, block_message: str +) -> dict: + detections: list = [] + correlation_ids: List[str] = [] + for v in blocked: + detections.extend(v.detections) + correlation_ids.extend(v.correlation_ids) + detail: dict = { + "error": block_message, + "type": "alice_wonderfence_content_policy_violation", + "guardrail_name": guardrail_name, + "action": "BLOCK", + "wonderfence_correlation_id": correlation_ids[0] if correlation_ids else None, + "wonderfence_correlation_ids": correlation_ids, + } + if detections: + detail["detections"] = [ + d.model_dump() if hasattr(d, "model_dump") else d for d in detections + ] + return detail + + +def _masked_value( + verdict: SegmentVerdict, guardrail_name: str, label: str +) -> Optional[str]: + """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 + segment came from.""" + correlation_id = verdict.correlation_ids[0] if verdict.correlation_ids else None + if verdict.action == "MASK": + logger.info( + "Alice WonderFence (apply_guardrail): MASK applied to %s guardrail=%s correlation_id=%s", + label, + guardrail_name, + correlation_id, + ) + return verdict.masked_text if verdict.masked_text is not None else "[MASKED]" + if verdict.action == "DETECT": + logger.warning( + "Alice WonderFence (apply_guardrail): DETECT %s guardrail=%s correlation_id=%s", + label, + guardrail_name, + correlation_id, + ) + return None + + def apply_verdicts( inputs: GenericGuardrailAPIInputs, indices: List[int], @@ -82,73 +178,49 @@ def apply_verdicts( block_message: str, tool_indices: Optional[List[int]] = None, tool_verdicts: Optional[List[SegmentVerdict]] = None, + tool_def_paths: Optional[List[List[Any]]] = None, + tool_def_verdicts: Optional[List[SegmentVerdict]] = None, ) -> GenericGuardrailAPIInputs: - """Apply per-segment verdicts back onto ``inputs["texts"]`` and tool-call args. + """Apply per-segment verdicts back onto request text, tool-call args, and + tool-definition descriptions. - Any BLOCK across text or tool-call segments raises ``WonderFenceBlockedError`` - with detections/correlation ids aggregated across all blocked segments. - Otherwise each MASK verdict rewrites its mapped ``texts`` index or - ``tool_calls[i]["function"]["arguments"]`` and DETECT is logged. + 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. """ tool_indices = tool_indices or [] tool_verdicts = tool_verdicts or [] - blocked = [v for v in (*verdicts, *tool_verdicts) if v.action == "BLOCK"] + tool_def_paths = tool_def_paths or [] + tool_def_verdicts = tool_def_verdicts or [] + + blocked = [ + v + for v in (*verdicts, *tool_verdicts, *tool_def_verdicts) + if v.action == "BLOCK" + ] if blocked: - detections: list = [] - correlation_ids: List[str] = [] - for v in blocked: - detections.extend(v.detections) - correlation_ids.extend(v.correlation_ids) - detail: dict = { - "error": block_message, - "type": "alice_wonderfence_content_policy_violation", - "guardrail_name": guardrail_name, - "action": "BLOCK", - "wonderfence_correlation_id": ( - correlation_ids[0] if correlation_ids else None - ), - "wonderfence_correlation_ids": correlation_ids, - } - if detections: - detail["detections"] = [ - d.model_dump() if hasattr(d, "model_dump") else d for d in detections - ] - raise WonderFenceBlockedError(detail) + raise WonderFenceBlockedError( + _block_detail(blocked, guardrail_name, block_message) + ) texts = inputs.get("texts") or [] for idx, verdict in zip(indices, verdicts): - if verdict.action == "MASK": - texts[idx] = ( - verdict.masked_text if verdict.masked_text is not None else "[MASKED]" - ) - logger.info( - "Alice WonderFence (apply_guardrail): MASK applied guardrail=%s correlation_id=%s", - guardrail_name, - verdict.correlation_ids[0] if verdict.correlation_ids else None, - ) - elif verdict.action == "DETECT": - logger.warning( - "Alice WonderFence (apply_guardrail): DETECT guardrail=%s correlation_id=%s", - guardrail_name, - verdict.correlation_ids[0] if verdict.correlation_ids else None, - ) + masked = _masked_value(verdict, guardrail_name, "request text") + if masked is not None: + texts[idx] = masked inputs["texts"] = texts tool_calls = inputs.get("tool_calls") or [] for idx, verdict in zip(tool_indices, tool_verdicts): - if verdict.action == "MASK": - tool_calls[idx]["function"]["arguments"] = ( - verdict.masked_text if verdict.masked_text is not None else "[MASKED]" - ) - logger.info( - "Alice WonderFence (apply_guardrail): MASK applied to tool_call args guardrail=%s correlation_id=%s", - guardrail_name, - verdict.correlation_ids[0] if verdict.correlation_ids else None, - ) - elif verdict.action == "DETECT": - logger.warning( - "Alice WonderFence (apply_guardrail): DETECT tool_call args guardrail=%s correlation_id=%s", - guardrail_name, - verdict.correlation_ids[0] if verdict.correlation_ids else None, - ) + masked = _masked_value(verdict, guardrail_name, "tool_call args") + 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) + return inputs diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail.py index 67149c635ad..03405ffa792 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 @@ -715,3 +715,123 @@ async def test_malformed_override_does_not_fail_open(make_guardrail, make_reques assert exc.value.status_code == 500 assert "alice_wonderfence_app_id" in exc.value.detail["exception"] client.evaluate_prompt.assert_not_awaited() + + +def _tool_def(description="a helpful tool", param_desc=None): + fn = { + "name": "do_thing", + "description": description, + "parameters": {"type": "object", "properties": {}}, + } + if param_desc is not None: + fn["parameters"]["properties"]["city"] = { + "type": "string", + "description": param_desc, + } + return {"type": "function", "function": fn} + + +@pytest.mark.asyncio +async def test_apply_guardrail_blocks_on_tool_definition_description( + guardrail_and_client, make_request_data +): + """Blocked content in tools[].function.description must BLOCK; tool defs are + forwarded to the model but were previously unscanned.""" + guardrail, client = guardrail_and_client + + def evaluate(prompt, **kwargs): + r = Mock() + r.action = "BLOCK" if "DISALLOWED" in prompt else "NO_ACTION" + r.detections = [] + r.correlation_id = None + return r + + client.evaluate_prompt.side_effect = evaluate + + inputs = { + "texts": ["use the tool"], + "tools": [_tool_def(description="DISALLOWED instructions here")], + } + with pytest.raises(HTTPException) as exc: + await guardrail.apply_guardrail( + inputs=inputs, request_data=make_request_data(), input_type="request" + ) + assert exc.value.status_code == 400 + + +@pytest.mark.asyncio +async def test_apply_guardrail_blocks_on_tool_parameter_description( + guardrail_and_client, make_request_data +): + """Nested parameter descriptions are scanned too, not just the top-level one.""" + guardrail, client = guardrail_and_client + + def evaluate(prompt, **kwargs): + r = Mock() + r.action = "BLOCK" if "DISALLOWED" in prompt else "NO_ACTION" + r.detections = [] + r.correlation_id = None + return r + + client.evaluate_prompt.side_effect = evaluate + + inputs = { + "texts": ["hi"], + "tools": [_tool_def(description="benign", param_desc="DISALLOWED payload")], + } + with pytest.raises(HTTPException) as exc: + await guardrail.apply_guardrail( + inputs=inputs, request_data=make_request_data(), input_type="request" + ) + assert exc.value.status_code == 400 + + +@pytest.mark.asyncio +async def test_apply_guardrail_masks_tool_definition_description_in_place( + guardrail_and_client, make_request_data +): + 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.detections = [] + r.correlation_id = None + return r + + client.evaluate_prompt.side_effect = evaluate + + inputs = { + "texts": ["hi"], + "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]" + + +@pytest.mark.asyncio +async def test_apply_guardrail_scans_tools_when_no_texts_or_tool_calls( + guardrail_and_client, make_request_data +): + """A request carrying only tool definitions must still be scanned.""" + guardrail, client = guardrail_and_client + + def evaluate(prompt, **kwargs): + r = Mock() + r.action = "BLOCK" if "DISALLOWED" 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": [], "tools": [_tool_def(description="DISALLOWED")]}, + request_data=make_request_data(), + input_type="request", + ) + assert exc.value.status_code == 400 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 49a5a3e813a..317dc8d7582 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 @@ -47,3 +47,55 @@ def test_detect_and_no_action_leave_texts_unchanged(): ] out = apply_verdicts(inputs, [0, 1], verdicts, "gn", "blocked!") assert out["texts"] == ["a", "b"] + + +# --------------- tool_definition_segments --------------- + + +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": [ + { + "type": "function", + "function": { + "name": "weather", + "description": "TOP_DESC", + "parameters": { + "type": "object", + "properties": { + "city": {"type": "string", "description": "PARAM_DESC"} + }, + }, + }, + } + ] + } + 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"] == "" + + +def test_tool_definition_segments_ignores_non_dict_tools_and_blank_descriptions(): + inputs = { + "tools": [ + "not-a-dict", + {"type": "function", "function": {"name": "f", "description": " "}}, + {"type": "function"}, + ] + } + from litellm.proxy.guardrails.guardrail_hooks.alice_wonderfence.processing import ( + tool_definition_segments, + ) + + paths, segments = tool_definition_segments(inputs) + assert segments == []