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.
This commit is contained in:
lior-k 2026-06-17 18:26:54 +03:00
parent f234039d62
commit df5cd8dab3
No known key found for this signature in database
4 changed files with 314 additions and 63 deletions

View file

@ -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:

View file

@ -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

View file

@ -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

View file

@ -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"] == "<TOP_DESC>"
assert fn["parameters"]["properties"]["city"]["description"] == "<PARAM_DESC>"
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 == []