mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
b4ab76ecd5
commit
abda857ed8
3 changed files with 188 additions and 12 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue