fix(guardrails): Alice WonderFence scans legacy functions[] definitions

The deprecated top-level functions[] request parameter is forwarded to providers
(litellm converts it to tools only later, during the LLM call, after the
guardrail runs), so blocked content in functions[].description or nested
parameter descriptions reached the model unscanned. Each functions[] entry is
shaped like a tool's function object, so its descriptions are now extracted
(reusing the tool-definition walker) and evaluated as request-side segments.

Read from request_data because the chat translation layer surfaces tools but not
functions in inputs. Detection only: BLOCK raises and DETECT logs; functions has
no inputs write-back path so it is not masked (matching the precedent set by the
cisco_ai_defense guardrail, which scans both tools and functions for detection).

Regression tests: BLOCK on a function description, BLOCK on a nested parameter
description, a functions-only request still scanned, and detection-without-mask
leaving request_data["functions"] untouched; the BLOCK cases fail on prior code.
This commit is contained in:
lior-k 2026-06-21 21:03:29 +03:00
parent 9ae1ea4704
commit d61eda2d07
No known key found for this signature in database
4 changed files with 225 additions and 6 deletions

View file

@ -28,6 +28,7 @@ from .exceptions import WonderFenceBlockedError, WonderFenceMissingSecrets
from .processing import (
apply_verdicts,
build_analysis_context,
function_definition_segments,
tool_call_arg_segments,
tool_definition_segments,
)
@ -177,7 +178,19 @@ class WonderFenceGuardrail(CustomGuardrail):
texts = inputs.get("texts") or []
tool_indices, tool_segments = tool_call_arg_segments(inputs)
tool_def_paths, tool_def_segments = tool_definition_segments(inputs)
if not texts and not tool_segments and not tool_def_segments:
# 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_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
):
logger.debug(
"Alice WonderFence (apply_guardrail): nothing to scan for %s",
input_type,
@ -215,12 +228,18 @@ class WonderFenceGuardrail(CustomGuardrail):
custom_fields=None,
)
segments = [*texts, *tool_segments, *tool_def_segments]
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 segment(s) app_id=%s guardrail=%s input_type=%s",
"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,
@ -232,6 +251,7 @@ class WonderFenceGuardrail(CustomGuardrail):
)
n_text = len(texts)
n_tool = len(tool_segments)
n_tool_def = len(tool_def_segments)
apply_verdicts(
inputs,
list(range(n_text)),
@ -241,7 +261,10 @@ class WonderFenceGuardrail(CustomGuardrail):
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 :],
tool_def_verdicts=verdicts[
n_text + n_tool : n_text + n_tool + n_tool_def
],
function_def_verdicts=verdicts[n_text + n_tool + n_tool_def :],
)
except WonderFenceBlockedError as e:

View file

@ -125,6 +125,25 @@ def tool_definition_segments(
return paths, segments
def function_definition_segments(request_data: dict) -> list[str]:
"""Description strings from the deprecated top-level ``functions[]`` request
parameter.
Each entry is shaped like a tool's ``function`` object
(``{name, description, parameters}``) and LiteLLM forwards it to providers,
so its descriptions are scanned. Read from ``request_data`` because the chat
translation layer surfaces ``tools`` but not ``functions`` in ``inputs``.
Detection only (BLOCK/DETECT) -- ``functions`` has no inputs write-back path,
so it is not masked.
"""
functions = request_data.get("functions") or []
segments: list[str] = []
for fn in functions:
if isinstance(fn, dict):
segments.extend(text for _path, text in _description_strings(fn, []))
return segments
def _set_by_path(root: Any, path: list[Any], value: Any) -> None:
obj = root
for key in path[:-1]:
@ -190,6 +209,7 @@ def apply_verdicts(
tool_verdicts: list[SegmentVerdict] | None = None,
tool_def_paths: list[list[Any]] | None = None,
tool_def_verdicts: list[SegmentVerdict] | None = None,
function_def_verdicts: list[SegmentVerdict] | None = None,
) -> GenericGuardrailAPIInputs:
"""Apply per-segment verdicts back onto request text, tool-call args, and
tool-definition descriptions.
@ -197,16 +217,23 @@ def apply_verdicts(
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.
logged. ``function_def_verdicts`` (legacy ``functions[]``) are detection
only: BLOCK raises, anything else is logged, never masked.
"""
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_verdicts = function_def_verdicts or []
blocked = [
v
for v in (*verdicts, *tool_verdicts, *tool_def_verdicts)
for v in (
*verdicts,
*tool_verdicts,
*tool_def_verdicts,
*function_def_verdicts,
)
if v.action == "BLOCK"
]
if blocked:
@ -233,4 +260,12 @@ def apply_verdicts(
if masked is not None:
_set_by_path(tools, path, masked)
for verdict in function_def_verdicts:
if verdict.action in ("MASK", "DETECT"):
logger.warning(
"Alice WonderFence (apply_guardrail): DETECT function definition guardrail=%s correlation_id=%s",
guardrail_name,
verdict.correlation_ids[0] if verdict.correlation_ids else None,
)
return inputs

View file

@ -835,3 +835,129 @@ async def test_apply_guardrail_scans_tools_when_no_texts_or_tool_calls(
input_type="request",
)
assert exc.value.status_code == 400
def _legacy_function(description="a function", 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 fn
@pytest.mark.asyncio
async def test_apply_guardrail_blocks_on_legacy_function_description(
guardrail_and_client, make_request_data
):
"""Blocked content in the deprecated functions[].description (read from
request_data, not inputs) must BLOCK."""
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": ["hi"]},
request_data=make_request_data(
functions=[_legacy_function(description="DISALLOWED instructions")]
),
input_type="request",
)
assert exc.value.status_code == 400
@pytest.mark.asyncio
async def test_apply_guardrail_blocks_on_legacy_function_parameter_description(
guardrail_and_client, make_request_data
):
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": ["hi"]},
request_data=make_request_data(
functions=[_legacy_function(description="ok", param_desc="DISALLOWED")]
),
input_type="request",
)
assert exc.value.status_code == 400
@pytest.mark.asyncio
async def test_apply_guardrail_scans_legacy_functions_when_no_other_content(
guardrail_and_client, make_request_data
):
"""A request whose only scannable content is functions[] is still 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": []},
request_data=make_request_data(
functions=[_legacy_function(description="DISALLOWED")]
),
input_type="request",
)
assert exc.value.status_code == 400
@pytest.mark.asyncio
async def test_apply_guardrail_legacy_function_not_masked_only_detected(
guardrail_and_client, make_request_data
):
"""A non-BLOCK verdict on a function definition passes through without
mutating request_data['functions'] (detection only, no mask write-back)."""
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
request_data = make_request_data(
functions=[_legacy_function(description="contains secret stuff")]
)
out = await guardrail.apply_guardrail(
inputs={"texts": ["hi"]},
request_data=request_data,
input_type="request",
)
assert out is not None
# functions left untouched (no mask write-back)
assert request_data["functions"][0]["description"] == "contains secret stuff"

View file

@ -99,3 +99,38 @@ def test_tool_definition_segments_ignores_non_dict_tools_and_blank_descriptions(
paths, segments = tool_definition_segments(inputs)
assert segments == []
# --------------- function_definition_segments (legacy functions[]) ---------------
def test_function_definition_segments_extracts_descriptions():
from litellm.proxy.guardrails.guardrail_hooks.alice_wonderfence.processing import (
function_definition_segments,
)
request_data = {
"functions": [
{
"name": "weather",
"description": "TOP_DESC",
"parameters": {
"type": "object",
"properties": {
"city": {"type": "string", "description": "PARAM_DESC"}
},
},
},
"not-a-dict",
{"name": "f", "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"}) == []