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.
This commit is contained in:
lior-k 2026-07-06 11:02:20 +03:00
parent 73dfe12839
commit d852ae51f3
No known key found for this signature in database
3 changed files with 538 additions and 527 deletions

View file

@ -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]"

View file

@ -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"

View file

@ -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]"