fix(guardrails): Alice WonderFence scans tool-call arguments

tool_calls reach the model (request side, from assistant messages) and the
client (response side, model-generated), and the translation layer threads them
through inputs["tool_calls"] and writes mutations back, but apply_guardrail only
looked at inputs["texts"]. Disallowed content placed in
tool_calls[].function.arguments therefore went unscanned. The early return also
skipped requests whose only content was a tool call (empty texts).

Each tool-call argument string is now evaluated as a segment alongside the text
segments through the same WonderFence call; BLOCK raises, MASK rewrites
inputs["tool_calls"][i]["function"]["arguments"] in place (the translation layer
writes it back), DETECT logs. The empty-texts early return now also accounts for
tool-call args. Regression tests cover request/response BLOCK on tool args, MASK
write-back, and the tool-calls-without-texts case; all fail on the prior code.
This commit is contained in:
lior-k 2026-06-09 13:06:25 +03:00
parent b4ab76ecd5
commit abda857ed8
No known key found for this signature in database
3 changed files with 188 additions and 12 deletions

View file

@ -25,7 +25,11 @@ from .chunked_evaluation import DEFAULT_MAX_CONCURRENCY, evaluate_segments
from .client_cache import get_or_create_client, load_sdk
from .credentials import resolve_credentials
from .exceptions import WonderFenceBlockedError, WonderFenceMissingSecrets
from .processing import apply_verdicts, build_analysis_context
from .processing import (
apply_verdicts,
build_analysis_context,
tool_call_arg_segments,
)
if TYPE_CHECKING:
from wonderfence_sdk.client import ( # type: ignore[import-untyped]
@ -170,9 +174,10 @@ class WonderFenceGuardrail(CustomGuardrail):
) -> GenericGuardrailAPIInputs:
"""Apply WonderFence guardrail using V2 client + per-request app_id."""
texts = inputs.get("texts") or []
if not texts:
tool_indices, tool_segments = tool_call_arg_segments(inputs)
if not texts and not tool_segments:
logger.debug(
"Alice WonderFence (apply_guardrail): no text to scan for %s",
"Alice WonderFence (apply_guardrail): no text or tool-call args to scan for %s",
input_type,
)
return inputs
@ -208,24 +213,28 @@ class WonderFenceGuardrail(CustomGuardrail):
custom_fields=None,
)
segments = [*texts, *tool_segments]
logger.debug(
"Alice WonderFence (apply_guardrail): evaluating %d segment(s) app_id=%s guardrail=%s input_type=%s",
"Alice WonderFence (apply_guardrail): evaluating %d text + %d tool-call segment(s) app_id=%s guardrail=%s input_type=%s",
len(texts),
len(tool_segments),
app_id,
self.guardrail_name,
input_type,
)
verdicts = await evaluate_segments(
texts,
segments,
evaluate,
max_concurrency=self._connection_pool_limit or DEFAULT_MAX_CONCURRENCY,
)
apply_verdicts(
inputs,
list(range(len(texts))),
verdicts,
verdicts[: len(texts)],
self.guardrail_name,
self.block_message,
tool_indices=tool_indices,
tool_verdicts=verdicts[len(texts) :],
)
except WonderFenceBlockedError as e:

View file

@ -1,6 +1,6 @@
"""Pure transforms for Alice WonderFence: context build, user-text mapping, verdict apply."""
from typing import Any, List, Optional
from typing import Any, List, Optional, Tuple
import litellm
from litellm._logging import verbose_proxy_logger
@ -52,20 +52,47 @@ def build_analysis_context(
)
def tool_call_arg_segments(
inputs: GenericGuardrailAPIInputs,
) -> Tuple[List[int], List[str]]:
"""Return (indices, argument strings) for tool calls carrying string args.
``inputs["tool_calls"]`` entries are dicts shaped
``{"function": {"arguments": "<json string>"}}``; the argument string is the
caller- or model-controlled payload that reaches the model/client, so it is
scanned like any other segment.
"""
tool_calls = inputs.get("tool_calls") or []
indices: List[int] = []
segments: List[str] = []
for i, tool_call in enumerate(tool_calls):
fn = tool_call.get("function") if isinstance(tool_call, dict) else None
args = fn.get("arguments") if isinstance(fn, dict) else None
if isinstance(args, str) and args.strip():
indices.append(i)
segments.append(args)
return indices, segments
def apply_verdicts(
inputs: GenericGuardrailAPIInputs,
indices: List[int],
verdicts: List[SegmentVerdict],
guardrail_name: str,
block_message: str,
tool_indices: Optional[List[int]] = None,
tool_verdicts: Optional[List[SegmentVerdict]] = None,
) -> GenericGuardrailAPIInputs:
"""Apply per-segment verdicts back onto ``inputs["texts"]``.
"""Apply per-segment verdicts back onto ``inputs["texts"]`` and tool-call args.
Any BLOCK raises ``WonderFenceBlockedError`` with detections/correlation ids
aggregated across all blocked segments. Otherwise each MASK verdict rewrites
its mapped ``texts`` index and DETECT is logged.
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.
"""
blocked = [v for v in verdicts if v.action == "BLOCK"]
tool_indices = tool_indices or []
tool_verdicts = tool_verdicts or []
blocked = [v for v in (*verdicts, *tool_verdicts) if v.action == "BLOCK"]
if blocked:
detections: list = []
correlation_ids: List[str] = []
@ -106,4 +133,22 @@ def apply_verdicts(
verdict.correlation_ids[0] if verdict.correlation_ids else None,
)
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,
)
return inputs

View file

@ -160,6 +160,128 @@ async def test_apply_guardrail_scans_non_user_role_segments(
assert exc.value.detail["action"] == "BLOCK"
def _tool_call(arguments, name="send_email"):
return {
"id": "call_1",
"type": "function",
"function": {"name": name, "arguments": arguments},
}
@pytest.mark.asyncio
async def test_apply_guardrail_blocks_on_tool_call_arguments(
guardrail_and_client, make_request_data
):
"""Bypass regression: blocked content in tool_calls[].function.arguments must
BLOCK. tool_calls reach the model but were never scanned (texts-only)."""
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": ["please run the tool"],
"tool_calls": [_tool_call('{"body": "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
assert exc.value.detail["action"] == "BLOCK"
@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']."""
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.detections = []
r.correlation_id = None
return r
client.evaluate_prompt.side_effect = evaluate
inputs = {
"texts": ["benign"],
"tool_calls": [_tool_call('{"body": "secret value"}')],
}
out = await guardrail.apply_guardrail(
inputs=inputs,
request_data=make_request_data(),
input_type="request",
)
assert out["tool_calls"][0]["function"]["arguments"] == '{"body": "[REDACTED]"}'
assert out["texts"] == ["benign"]
@pytest.mark.asyncio
async def test_apply_guardrail_scans_tool_calls_when_no_texts(
guardrail_and_client, make_request_data
):
"""An assistant message can carry tool_calls with no text content, so texts
is empty; the hook must still scan the tool-call arguments (the old
empty-texts early return skipped them)."""
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": [], "tool_calls": [_tool_call('{"x": "DISALLOWED"}')]},
request_data=make_request_data(),
input_type="request",
)
assert exc.value.status_code == 400
@pytest.mark.asyncio
async def test_apply_guardrail_blocks_on_response_tool_call_arguments(
guardrail_and_client, make_request_data
):
"""Model-generated tool-call arguments on the response side are scanned too."""
guardrail, client = guardrail_and_client
def evaluate(response, **kwargs):
r = Mock()
r.action = "BLOCK" if "DISALLOWED" in response else "NO_ACTION"
r.detections = []
r.correlation_id = None
return r
client.evaluate_response.side_effect = evaluate
with pytest.raises(HTTPException) as exc:
await guardrail.apply_guardrail(
inputs={"texts": [], "tool_calls": [_tool_call('{"x": "DISALLOWED"}')]},
request_data=make_request_data(),
input_type="response",
)
assert exc.value.status_code == 400
@pytest.mark.asyncio
async def test_apply_guardrail_mask_replaces_scanned_text_response(
guardrail_and_client, make_request_data