mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
9ae1ea4704
commit
d61eda2d07
4 changed files with 225 additions and 6 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"}) == []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue