From d852ae51f3573b9a777737be43600f6368007c62 Mon Sep 17 00:00:00 2001 From: lior-k Date: Mon, 6 Jul 2026 11:02:20 +0300 Subject: [PATCH] test(guardrails): split test_apply_guardrail into themed files under 500 LOC Break the 899-line test_apply_guardrail.py into three focused files: text-side actions stay in test_apply_guardrail.py, tool-call / tool-definition / legacy functions[] scanning moves to test_apply_guardrail_tools.py, and fail-open/closed plus missing-secrets and helpers move to test_apply_guardrail_failmodes.py. No test bodies changed; every alice file is now under 500 added LOC. --- .../alice_wonderfence/test_apply_guardrail.py | 533 +----------------- .../test_apply_guardrail_failmodes.py | 147 +++++ .../test_apply_guardrail_tools.py | 385 +++++++++++++ 3 files changed, 538 insertions(+), 527 deletions(-) create mode 100644 tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail_failmodes.py create mode 100644 tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail_tools.py 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 1b171db04ef..cb3436d281a 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 @@ -1,6 +1,10 @@ -"""Tests for ``apply_guardrail`` BLOCK/MASK/DETECT/NO_ACTION + fail modes + helpers.""" +"""Tests for ``apply_guardrail`` text-side BLOCK/MASK/DETECT/NO_ACTION + core scanning path. + +Tool-call / tool-definition / legacy functions[] scanning lives in +``test_apply_guardrail_tools.py``; fail-open/closed, missing secrets, and helpers +live in ``test_apply_guardrail_failmodes.py``. +""" -import sys from unittest.mock import Mock import pytest @@ -150,146 +154,6 @@ async def test_apply_guardrail_scans_non_user_role_segments(guardrail_and_client 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_detect_on_tool_call_args_passes_through(guardrail_and_client, make_request_data): - """A DETECT verdict on a tool-call argument logs but does not block or mutate - the arguments (symmetric with the text-side DETECT behavior).""" - guardrail, client = guardrail_and_client - - def evaluate(prompt, **kwargs): - r = Mock() - r.action = "DETECT" if "watch me" in prompt else "NO_ACTION" - r.action_text = None - r.detections = [] - r.correlation_id = "corr-detect" - return r - - client.evaluate_prompt.side_effect = evaluate - - inputs = {"texts": ["benign"], "tool_calls": [_tool_call('{"x": "watch me"}')]} - out = await guardrail.apply_guardrail( - inputs=inputs, - request_data=make_request_data(), - input_type="request", - ) - assert out["tool_calls"][0]["function"]["arguments"] == '{"x": "watch me"}' - 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): guardrail, client = guardrail_and_client @@ -512,388 +376,3 @@ async def test_apply_guardrail_no_text_short_circuits(guardrail_and_client, make assert out == {"texts": []} client.evaluate_prompt.assert_not_awaited() client.evaluate_response.assert_not_awaited() - - -# ----------------------------- fail modes ----------------------------- - - -@pytest.mark.asyncio -async def test_apply_guardrail_missing_app_id_fail_closed_returns_500(guardrail_and_client, make_request_data): - """Missing app_id follows the fail_open pattern: fail_open=False → HTTP 500.""" - guardrail, _ = guardrail_and_client - with pytest.raises(HTTPException) as exc: - await guardrail.apply_guardrail( - inputs={"texts": ["hi"]}, - request_data=make_request_data(metadata={}), - input_type="request", - ) - assert exc.value.status_code == 500 - assert "Error in Alice WonderFence Guardrail" in exc.value.detail["error"] - assert "alice_wonderfence_app_id" in exc.value.detail["exception"] - - -@pytest.mark.asyncio -async def test_apply_guardrail_missing_api_key_fail_closed_returns_500(monkeypatch, make_guardrail, make_request_data): - """Missing api_key follows the fail_open pattern: fail_open=False → HTTP 500.""" - monkeypatch.delenv("ALICE_API_KEY", raising=False) - guardrail, _ = make_guardrail(api_key=None) - with pytest.raises(HTTPException) as exc: - await guardrail.apply_guardrail( - inputs={"texts": ["hi"]}, - request_data=make_request_data(), - input_type="request", - ) - assert exc.value.status_code == 500 - assert "Error in Alice WonderFence Guardrail" in exc.value.detail["error"] - assert "alice_wonderfence_api_key" in exc.value.detail["exception"] - - -@pytest.mark.asyncio -async def test_apply_guardrail_missing_app_id_fail_open_returns_500(make_guardrail, make_request_data): - """Missing app_id is a config error: never fail-open, even with fail_open=True.""" - guardrail, _ = make_guardrail(fail_open=True) - with pytest.raises(HTTPException) as exc: - await guardrail.apply_guardrail( - inputs={"texts": ["hi"]}, - request_data=make_request_data(metadata={}), - input_type="request", - ) - assert exc.value.status_code == 500 - assert "alice_wonderfence_app_id" in exc.value.detail["exception"] - - -@pytest.mark.asyncio -async def test_apply_guardrail_missing_api_key_fail_open_returns_500(monkeypatch, make_guardrail, make_request_data): - """Missing api_key is a config error: never fail-open, even with fail_open=True.""" - monkeypatch.delenv("ALICE_API_KEY", raising=False) - guardrail, _ = make_guardrail(api_key=None, fail_open=True) - with pytest.raises(HTTPException) as exc: - await guardrail.apply_guardrail( - inputs={"texts": ["hi"]}, - request_data=make_request_data(), - input_type="request", - ) - assert exc.value.status_code == 500 - assert "alice_wonderfence_api_key" in exc.value.detail["exception"] - - -@pytest.mark.asyncio -async def test_apply_guardrail_fail_open_swallows_transport_error(make_guardrail, make_request_data): - guardrail, client = make_guardrail(fail_open=True) - guardrail._client_cache["default-api-key"] = client - client.evaluate_prompt.side_effect = RuntimeError("network down") - - inputs = {"texts": ["original"]} - out = await guardrail.apply_guardrail( - inputs=inputs, - request_data=make_request_data(), - input_type="request", - ) - assert out["texts"] == ["original"] - - -@pytest.mark.asyncio -async def test_apply_guardrail_fail_closed_returns_500(guardrail_and_client, make_request_data): - guardrail, client = guardrail_and_client - client.evaluate_prompt.side_effect = RuntimeError("network down") - - with pytest.raises(HTTPException) as exc: - await guardrail.apply_guardrail( - inputs={"texts": ["hi"]}, - request_data=make_request_data(), - input_type="request", - ) - assert exc.value.status_code == 500 - assert "Error in Alice WonderFence Guardrail" in exc.value.detail["error"] - - -# ----------------------------- helpers ----------------------------- - - -def test_get_config_model(make_guardrail): - from litellm.types.proxy.guardrails.guardrail_hooks.alice_wonderfence import ( - WonderFenceGuardrailConfigModel, - ) - - guardrail, _ = make_guardrail() - assert guardrail.get_config_model() is WonderFenceGuardrailConfigModel - - -def test_build_analysis_context_falls_back_to_slash_split(monkeypatch, make_guardrail): - """When ``litellm.get_llm_provider`` raises, fall back to ``provider/model`` split.""" - import litellm - - from litellm.proxy.guardrails.guardrail_hooks.alice_wonderfence.processing import ( - build_analysis_context, - ) - - guardrail, _ = make_guardrail() - - def boom(model): - raise ValueError("unknown provider") - - monkeypatch.setattr(litellm, "get_llm_provider", boom) - build_analysis_context({"model": "myorg/custom-llm"}, guardrail.platform, guardrail._AnalysisContext) - - AnalysisContext = sys.modules["wonderfence_sdk.models"].AnalysisContext - kwargs = AnalysisContext.call_args.kwargs - assert kwargs["provider"] == "myorg" - assert kwargs["model_name"] == "custom-llm" - - -@pytest.mark.asyncio -async def test_malformed_override_does_not_fail_open(make_guardrail, make_request_data): - """A non-string request-metadata app_id override must not slip through under - fail_open: it resolves to a config error (500), not a swallowed exception - that skips scanning. The SDK is never called with a malformed value.""" - guardrail, client = make_guardrail(fail_open=True, allow_request_metadata_override=True) - guardrail._client_cache["default-api-key"] = client - - with pytest.raises(HTTPException) as exc: - await guardrail.apply_guardrail( - inputs={"texts": ["hi"]}, - request_data=make_request_data(metadata={"alice_wonderfence_app_id": ["not", "a", "string"]}), - input_type="request", - ) - assert exc.value.status_code == 500 - assert "alice_wonderfence_app_id" in exc.value.detail["exception"] - client.evaluate_prompt.assert_not_awaited() - - -def _tool_def(description="a helpful tool", 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 {"type": "function", "function": fn} - - -@pytest.mark.asyncio -async def test_apply_guardrail_blocks_on_tool_definition_description(guardrail_and_client, make_request_data): - """Blocked content in tools[].function.description must BLOCK; tool defs are - forwarded to the model but were previously unscanned.""" - 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": ["use the tool"], - "tools": [_tool_def(description="DISALLOWED instructions here")], - } - 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 - - -@pytest.mark.asyncio -async def test_apply_guardrail_blocks_on_tool_parameter_description(guardrail_and_client, make_request_data): - """Nested parameter descriptions are scanned too, not just the top-level one.""" - 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": ["hi"], - "tools": [_tool_def(description="benign", param_desc="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 - - -@pytest.mark.asyncio -async def test_apply_guardrail_masks_tool_definition_description_in_place(guardrail_and_client, make_request_data): - 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 - - inputs = { - "texts": ["hi"], - "tools": [_tool_def(description="contains secret stuff")], - } - out = await guardrail.apply_guardrail(inputs=inputs, request_data=make_request_data(), input_type="request") - assert out["tools"][0]["function"]["description"] == "[REDACTED]" - - -@pytest.mark.asyncio -async def test_apply_guardrail_scans_tools_when_no_texts_or_tool_calls(guardrail_and_client, make_request_data): - """A request carrying only tool definitions must still be 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": [], "tools": [_tool_def(description="DISALLOWED")]}, - request_data=make_request_data(), - 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_detect_does_not_mutate(guardrail_and_client, make_request_data): - """A DETECT verdict on a function definition logs but does not rewrite it.""" - guardrail, client = guardrail_and_client - - def evaluate(prompt, **kwargs): - r = Mock() - r.action = "DETECT" if "watch" in prompt else "NO_ACTION" - r.action_text = None - r.detections = [] - r.correlation_id = None - return r - - client.evaluate_prompt.side_effect = evaluate - - request_data = make_request_data(functions=[_legacy_function(description="watch this")]) - out = await guardrail.apply_guardrail( - inputs={"texts": ["hi"]}, - request_data=request_data, - input_type="request", - ) - assert out is not None - assert request_data["functions"][0]["description"] == "watch this" - - -@pytest.mark.asyncio -async def test_apply_guardrail_masks_legacy_function_description_in_place(guardrail_and_client, make_request_data): - """A MASK verdict on a functions[] description must be written back into - request_data['functions'], not left as the original unredacted text.""" - 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")]) - await guardrail.apply_guardrail( - inputs={"texts": ["hi"]}, - request_data=request_data, - input_type="request", - ) - assert request_data["functions"][0]["description"] == "[REDACTED]" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail_failmodes.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail_failmodes.py new file mode 100644 index 00000000000..8238e55f6c1 --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail_failmodes.py @@ -0,0 +1,147 @@ +"""Tests for ``apply_guardrail`` fail-open/fail-closed behavior, missing secrets, and helpers.""" + +import sys +from unittest.mock import Mock + +import pytest +from fastapi import HTTPException + + +@pytest.mark.asyncio +async def test_apply_guardrail_missing_app_id_fail_closed_returns_500(guardrail_and_client, make_request_data): + """Missing app_id follows the fail_open pattern: fail_open=False → HTTP 500.""" + guardrail, _ = guardrail_and_client + with pytest.raises(HTTPException) as exc: + await guardrail.apply_guardrail( + inputs={"texts": ["hi"]}, + request_data=make_request_data(metadata={}), + input_type="request", + ) + assert exc.value.status_code == 500 + assert "Error in Alice WonderFence Guardrail" in exc.value.detail["error"] + assert "alice_wonderfence_app_id" in exc.value.detail["exception"] + + +@pytest.mark.asyncio +async def test_apply_guardrail_missing_api_key_fail_closed_returns_500(monkeypatch, make_guardrail, make_request_data): + """Missing api_key follows the fail_open pattern: fail_open=False → HTTP 500.""" + monkeypatch.delenv("ALICE_API_KEY", raising=False) + guardrail, _ = make_guardrail(api_key=None) + with pytest.raises(HTTPException) as exc: + await guardrail.apply_guardrail( + inputs={"texts": ["hi"]}, + request_data=make_request_data(), + input_type="request", + ) + assert exc.value.status_code == 500 + assert "Error in Alice WonderFence Guardrail" in exc.value.detail["error"] + assert "alice_wonderfence_api_key" in exc.value.detail["exception"] + + +@pytest.mark.asyncio +async def test_apply_guardrail_missing_app_id_fail_open_returns_500(make_guardrail, make_request_data): + """Missing app_id is a config error: never fail-open, even with fail_open=True.""" + guardrail, _ = make_guardrail(fail_open=True) + with pytest.raises(HTTPException) as exc: + await guardrail.apply_guardrail( + inputs={"texts": ["hi"]}, + request_data=make_request_data(metadata={}), + input_type="request", + ) + assert exc.value.status_code == 500 + assert "alice_wonderfence_app_id" in exc.value.detail["exception"] + + +@pytest.mark.asyncio +async def test_apply_guardrail_missing_api_key_fail_open_returns_500(monkeypatch, make_guardrail, make_request_data): + """Missing api_key is a config error: never fail-open, even with fail_open=True.""" + monkeypatch.delenv("ALICE_API_KEY", raising=False) + guardrail, _ = make_guardrail(api_key=None, fail_open=True) + with pytest.raises(HTTPException) as exc: + await guardrail.apply_guardrail( + inputs={"texts": ["hi"]}, + request_data=make_request_data(), + input_type="request", + ) + assert exc.value.status_code == 500 + assert "alice_wonderfence_api_key" in exc.value.detail["exception"] + + +@pytest.mark.asyncio +async def test_apply_guardrail_fail_open_swallows_transport_error(make_guardrail, make_request_data): + guardrail, client = make_guardrail(fail_open=True) + guardrail._client_cache["default-api-key"] = client + client.evaluate_prompt.side_effect = RuntimeError("network down") + + inputs = {"texts": ["original"]} + out = await guardrail.apply_guardrail( + inputs=inputs, + request_data=make_request_data(), + input_type="request", + ) + assert out["texts"] == ["original"] + + +@pytest.mark.asyncio +async def test_apply_guardrail_fail_closed_returns_500(guardrail_and_client, make_request_data): + guardrail, client = guardrail_and_client + client.evaluate_prompt.side_effect = RuntimeError("network down") + + with pytest.raises(HTTPException) as exc: + await guardrail.apply_guardrail( + inputs={"texts": ["hi"]}, + request_data=make_request_data(), + input_type="request", + ) + assert exc.value.status_code == 500 + assert "Error in Alice WonderFence Guardrail" in exc.value.detail["error"] + + +@pytest.mark.asyncio +async def test_malformed_override_does_not_fail_open(make_guardrail, make_request_data): + """A non-string request-metadata app_id override must not slip through under + fail_open: it resolves to a config error (500), not a swallowed exception + that skips scanning. The SDK is never called with a malformed value.""" + guardrail, client = make_guardrail(fail_open=True, allow_request_metadata_override=True) + guardrail._client_cache["default-api-key"] = client + + with pytest.raises(HTTPException) as exc: + await guardrail.apply_guardrail( + inputs={"texts": ["hi"]}, + request_data=make_request_data(metadata={"alice_wonderfence_app_id": ["not", "a", "string"]}), + input_type="request", + ) + assert exc.value.status_code == 500 + assert "alice_wonderfence_app_id" in exc.value.detail["exception"] + client.evaluate_prompt.assert_not_awaited() + + +def test_get_config_model(make_guardrail): + from litellm.types.proxy.guardrails.guardrail_hooks.alice_wonderfence import ( + WonderFenceGuardrailConfigModel, + ) + + guardrail, _ = make_guardrail() + assert guardrail.get_config_model() is WonderFenceGuardrailConfigModel + + +def test_build_analysis_context_falls_back_to_slash_split(monkeypatch, make_guardrail): + """When ``litellm.get_llm_provider`` raises, fall back to ``provider/model`` split.""" + import litellm + + from litellm.proxy.guardrails.guardrail_hooks.alice_wonderfence.processing import ( + build_analysis_context, + ) + + guardrail, _ = make_guardrail() + + def boom(model): + raise ValueError("unknown provider") + + monkeypatch.setattr(litellm, "get_llm_provider", boom) + build_analysis_context({"model": "myorg/custom-llm"}, guardrail.platform, guardrail._AnalysisContext) + + AnalysisContext = sys.modules["wonderfence_sdk.models"].AnalysisContext + kwargs = AnalysisContext.call_args.kwargs + assert kwargs["provider"] == "myorg" + assert kwargs["model_name"] == "custom-llm" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail_tools.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail_tools.py new file mode 100644 index 00000000000..5a1fe6a5e7e --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail_tools.py @@ -0,0 +1,385 @@ +"""Tests for ``apply_guardrail`` scanning of tool calls, tool definitions, and legacy functions[].""" + +from unittest.mock import Mock + +import pytest +from fastapi import HTTPException + + +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_detect_on_tool_call_args_passes_through(guardrail_and_client, make_request_data): + """A DETECT verdict on a tool-call argument logs but does not block or mutate + the arguments (symmetric with the text-side DETECT behavior).""" + guardrail, client = guardrail_and_client + + def evaluate(prompt, **kwargs): + r = Mock() + r.action = "DETECT" if "watch me" in prompt else "NO_ACTION" + r.action_text = None + r.detections = [] + r.correlation_id = "corr-detect" + return r + + client.evaluate_prompt.side_effect = evaluate + + inputs = {"texts": ["benign"], "tool_calls": [_tool_call('{"x": "watch me"}')]} + out = await guardrail.apply_guardrail( + inputs=inputs, + request_data=make_request_data(), + input_type="request", + ) + assert out["tool_calls"][0]["function"]["arguments"] == '{"x": "watch me"}' + 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 + + +def _tool_def(description="a helpful tool", 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 {"type": "function", "function": fn} + + +@pytest.mark.asyncio +async def test_apply_guardrail_blocks_on_tool_definition_description(guardrail_and_client, make_request_data): + """Blocked content in tools[].function.description must BLOCK; tool defs are + forwarded to the model but were previously unscanned.""" + 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": ["use the tool"], + "tools": [_tool_def(description="DISALLOWED instructions here")], + } + 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 + + +@pytest.mark.asyncio +async def test_apply_guardrail_blocks_on_tool_parameter_description(guardrail_and_client, make_request_data): + """Nested parameter descriptions are scanned too, not just the top-level one.""" + 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": ["hi"], + "tools": [_tool_def(description="benign", param_desc="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 + + +@pytest.mark.asyncio +async def test_apply_guardrail_masks_tool_definition_description_in_place(guardrail_and_client, make_request_data): + 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 + + inputs = { + "texts": ["hi"], + "tools": [_tool_def(description="contains secret stuff")], + } + out = await guardrail.apply_guardrail(inputs=inputs, request_data=make_request_data(), input_type="request") + assert out["tools"][0]["function"]["description"] == "[REDACTED]" + + +@pytest.mark.asyncio +async def test_apply_guardrail_scans_tools_when_no_texts_or_tool_calls(guardrail_and_client, make_request_data): + """A request carrying only tool definitions must still be 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": [], "tools": [_tool_def(description="DISALLOWED")]}, + request_data=make_request_data(), + 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_detect_does_not_mutate(guardrail_and_client, make_request_data): + """A DETECT verdict on a function definition logs but does not rewrite it.""" + guardrail, client = guardrail_and_client + + def evaluate(prompt, **kwargs): + r = Mock() + r.action = "DETECT" if "watch" in prompt else "NO_ACTION" + r.action_text = None + r.detections = [] + r.correlation_id = None + return r + + client.evaluate_prompt.side_effect = evaluate + + request_data = make_request_data(functions=[_legacy_function(description="watch this")]) + out = await guardrail.apply_guardrail( + inputs={"texts": ["hi"]}, + request_data=request_data, + input_type="request", + ) + assert out is not None + assert request_data["functions"][0]["description"] == "watch this" + + +@pytest.mark.asyncio +async def test_apply_guardrail_masks_legacy_function_description_in_place(guardrail_and_client, make_request_data): + """A MASK verdict on a functions[] description must be written back into + request_data['functions'], not left as the original unredacted text.""" + 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")]) + await guardrail.apply_guardrail( + inputs={"texts": ["hi"]}, + request_data=request_data, + input_type="request", + ) + assert request_data["functions"][0]["description"] == "[REDACTED]"