diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py index f66e59dd9a2..dfb338f1a7f 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py @@ -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: diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py index d33134a436b..548649cb5e2 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py @@ -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": ""}}``; 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 diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail.py index 2a77dd34765..1341fa83880 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail.py @@ -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