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