style(guardrails): apply ruff format to Alice WonderFence files

The base adopted ruff format (line-length 120) as the CI formatter; reflow the
Alice WonderFence module and its tests so ruff format --check passes.
This commit is contained in:
lior-k 2026-07-05 16:10:52 +03:00
parent 8933969af4
commit c81292bcd5
No known key found for this signature in database
11 changed files with 108 additions and 363 deletions

View file

@ -14,9 +14,7 @@ if TYPE_CHECKING:
from litellm.types.guardrails import Guardrail, LitellmParams
def initialize_guardrail(
litellm_params: "LitellmParams", guardrail: "Guardrail"
) -> WonderFenceGuardrail:
def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail") -> WonderFenceGuardrail:
import litellm
guardrail_name = guardrail.get("guardrail_name")
@ -34,9 +32,7 @@ def initialize_guardrail(
"max_cached_clients": litellm_params.max_cached_clients,
"connection_pool_limit": litellm_params.connection_pool_limit,
"event_hook": litellm_params.mode,
"default_on": (
litellm_params.default_on if litellm_params.default_on is not None else True
),
"default_on": (litellm_params.default_on if litellm_params.default_on is not None else True),
}
if litellm_params.api_timeout is not None:
init_kwargs["api_timeout"] = litellm_params.api_timeout
@ -47,9 +43,7 @@ def initialize_guardrail(
if litellm_params.debug is not None:
init_kwargs["debug"] = litellm_params.debug
if litellm_params.allow_request_metadata_override is not None:
init_kwargs["allow_request_metadata_override"] = (
litellm_params.allow_request_metadata_override
)
init_kwargs["allow_request_metadata_override"] = litellm_params.allow_request_metadata_override
wonderfence_guardrail = WonderFenceGuardrail(**init_kwargs)

View file

@ -77,9 +77,7 @@ class WonderFenceGuardrail(CustomGuardrail):
max_cached_clients: int | None = None,
connection_pool_limit: int | None = None,
allow_request_metadata_override: bool = False,
event_hook: (
Union[GuardrailEventHooks, list[GuardrailEventHooks], Mode] | None
) = None,
event_hook: (Union[GuardrailEventHooks, list[GuardrailEventHooks], Mode] | None) = None,
default_on: bool = True,
**kwargs: Any,
) -> None:
@ -123,14 +121,10 @@ class WonderFenceGuardrail(CustomGuardrail):
logger.setLevel(logging.DEBUG)
self._client_cache: OrderedDict[str, _WonderFenceV2Client] = OrderedDict()
self._client_cache_maxsize = max_cached_clients or int(
os.environ.get("ALICE_MAX_CACHED_CLIENTS", "10")
)
self._client_cache_maxsize = max_cached_clients or int(os.environ.get("ALICE_MAX_CACHED_CLIENTS", "10"))
env_pool = os.environ.get("ALICE_CONNECTION_POOL_LIMIT")
self._connection_pool_limit: int | None = (
connection_pool_limit
if connection_pool_limit is not None
else (int(env_pool) if env_pool else None)
connection_pool_limit if connection_pool_limit is not None else (int(env_pool) if env_pool else None)
)
supported_event_hooks = [
@ -187,16 +181,9 @@ class WonderFenceGuardrail(CustomGuardrail):
# Legacy top-level functions[] only exist on the request body; the
# translation layer does not surface them in inputs, so read request_data.
function_def_paths, function_def_segments = (
function_definition_segments(request_data)
if input_type == "request"
else ([], [])
function_definition_segments(request_data) if input_type == "request" else ([], [])
)
if (
not texts
and not tool_segments
and not tool_def_segments
and not function_def_segments
):
if not texts and not tool_segments and not tool_def_segments and not function_def_segments:
logger.debug(
"Alice WonderFence (apply_guardrail): nothing to scan for %s",
input_type,
@ -215,16 +202,12 @@ class WonderFenceGuardrail(CustomGuardrail):
),
)
client = await self._get_client(api_key)
context = build_analysis_context(
request_data, self.platform, self._AnalysisContext
)
context = build_analysis_context(request_data, self.platform, self._AnalysisContext)
if input_type == "request":
async def evaluate(text: str) -> object:
return await client.evaluate_prompt(
app_id=app_id, prompt=text, context=context, custom_fields=None
)
return await client.evaluate_prompt(app_id=app_id, prompt=text, context=context, custom_fields=None)
else:
@ -270,9 +253,7 @@ class WonderFenceGuardrail(CustomGuardrail):
tool_indices=tool_indices,
tool_verdicts=verdicts[n_text : n_text + n_tool],
tool_def_paths=tool_def_paths,
tool_def_verdicts=verdicts[
n_text + n_tool : n_text + n_tool + n_tool_def
],
tool_def_verdicts=verdicts[n_text + n_tool : n_text + n_tool + n_tool_def],
function_def_paths=function_def_paths,
function_def_verdicts=verdicts[n_text + n_tool + n_tool_def :],
function_def_request_data=request_data,
@ -326,9 +307,7 @@ class WonderFenceGuardrail(CustomGuardrail):
},
) from e
add_guardrail_to_applied_guardrails_header(
request_data=request_data, guardrail_name=self.guardrail_name
)
add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name)
return inputs
@staticmethod

View file

@ -88,14 +88,10 @@ def _boundary_windows(chunks: list[str], overlap: int) -> list[str]:
"""
if overlap <= 0:
return []
return [
chunks[i][-overlap:] + chunks[i + 1][:overlap] for i in range(len(chunks) - 1)
]
return [chunks[i][-overlap:] + chunks[i + 1][:overlap] for i in range(len(chunks) - 1)]
def _cross_segment_windows(
segments: list[str], text_segment_count: int, overlap: int
) -> list[tuple[int, str]]:
def _cross_segment_windows(segments: list[str], text_segment_count: int, overlap: int) -> list[tuple[int, str]]:
"""Detection-only windows spanning each adjacent pair of prompt-text segments.
The chat translation layer emits each message content part as its own
@ -113,9 +109,7 @@ def _cross_segment_windows(
return []
n = min(text_segment_count, len(segments))
return [
(i, segments[i][-overlap:] + segments[i + 1][:overlap])
for i in range(n - 1)
if segments[i] and segments[i + 1]
(i, segments[i][-overlap:] + segments[i + 1][:overlap]) for i in range(n - 1) if segments[i] and segments[i + 1]
]
@ -205,15 +199,7 @@ async def evaluate_segments(
elif kind == "bound":
bound_res[si][idx] = res
cross_res: list[list[Any]] = [
[
res
for (kind, si, _), res in zip(index, results)
if kind == "cross" and si == s
]
for s in range(len(segments))
[res for (kind, si, _), res in zip(index, results) if kind == "cross" and si == s] for s in range(len(segments))
]
return [
_aggregate(seg_chunks[si], chunk_res[si], bound_res[si] + cross_res[si])
for si in range(len(segments))
]
return [_aggregate(seg_chunks[si], chunk_res[si], bound_res[si] + cross_res[si]) for si in range(len(segments))]

View file

@ -37,9 +37,7 @@ def load_sdk() -> tuple[Any, Any]:
AnalysisContext,
)
except ImportError as e:
raise ImportError(
"Alice WonderFence SDK not installed. Install with: pip install wonderfence-sdk"
) from e
raise ImportError("Alice WonderFence SDK not installed. Install with: pip install wonderfence-sdk") from e
return WonderFenceV2Client, AnalysisContext

View file

@ -212,9 +212,7 @@ def stash_resolved(
setattr(logging_obj, _stash_attr(guardrail_name), (api_key, app_id))
def recover_resolved(
logging_obj: Optional["LiteLLMLoggingObj"], guardrail_name: str
) -> tuple[str, str] | None:
def recover_resolved(logging_obj: Optional["LiteLLMLoggingObj"], guardrail_name: str) -> tuple[str, str] | None:
"""Look up the (api_key, app_id) this guardrail stashed earlier in this
request, or ``None``.

View file

@ -31,16 +31,10 @@ def build_analysis_context(
if "/" in model_str:
provider, model_name = model_str.split("/", 1)
user_id = (
metadata.get("user_api_key_end_user_id")
or metadata.get("end_user_id")
or metadata.get("user_id")
)
user_id = metadata.get("user_api_key_end_user_id") or metadata.get("end_user_id") or metadata.get("user_id")
session_id = (
request_data.get("litellm_session_id")
or metadata.get("litellm_session_id")
or metadata.get("session_id")
request_data.get("litellm_session_id") or metadata.get("litellm_session_id") or metadata.get("session_id")
)
return context_class(
@ -74,9 +68,7 @@ def tool_call_arg_segments(
return indices, segments
def _description_strings(
root: object, root_prefix: list[Any]
) -> list[tuple[list[Any], str]]:
def _description_strings(root: object, root_prefix: list[Any]) -> list[tuple[list[Any], str]]:
"""Collect ``(path, text)`` for every non-blank ``description`` string under
``root`` (a tool's ``function`` dict), walking nested JSON-schema parameters
so parameter descriptions are included, not just the top one.
@ -153,9 +145,7 @@ def _set_by_path(root: Any, path: list[Any], value: object) -> None:
obj[path[-1]] = value
def _block_detail(
blocked: list[SegmentVerdict], guardrail_name: str, block_message: str
) -> dict:
def _block_detail(blocked: list[SegmentVerdict], guardrail_name: str, block_message: str) -> dict:
detections: list = []
correlation_ids: list[str] = []
for v in blocked:
@ -170,15 +160,11 @@ def _block_detail(
"wonderfence_correlation_ids": correlation_ids,
}
if detections:
detail["detections"] = [
d.model_dump() if hasattr(d, "model_dump") else d for d in detections
]
detail["detections"] = [d.model_dump() if hasattr(d, "model_dump") else d for d in detections]
return detail
def _masked_value(
verdict: SegmentVerdict, guardrail_name: str, label: str
) -> str | None:
def _masked_value(verdict: SegmentVerdict, guardrail_name: str, label: str) -> str | None:
"""Return the replacement string for a MASK verdict (logging as a side
effect), or None for DETECT/NO_ACTION. The caller writes it to the slot the
segment came from."""
@ -241,9 +227,7 @@ def apply_verdicts(
if v.action == "BLOCK"
]
if blocked:
raise WonderFenceBlockedError(
_block_detail(blocked, guardrail_name, block_message)
)
raise WonderFenceBlockedError(_block_detail(blocked, guardrail_name, block_message))
texts = inputs.get("texts") or []
for idx, verdict in zip(indices, verdicts):

View file

@ -29,16 +29,12 @@ async def test_apply_guardrail_block_action(guardrail_and_client, make_request_d
assert exc.value.status_code == 400
assert exc.value.detail["action"] == "BLOCK"
assert exc.value.detail["wonderfence_correlation_id"] == "corr-1"
assert exc.value.detail["error"] == (
"Content violates our policies and has been blocked"
)
assert exc.value.detail["error"] == ("Content violates our policies and has been blocked")
assert exc.value.detail["detections"][0]["policy_name"] == "x"
@pytest.mark.asyncio
async def test_apply_guardrail_block_uses_custom_block_message(
make_guardrail, make_request_data
):
async def test_apply_guardrail_block_uses_custom_block_message(make_guardrail, make_request_data):
guardrail, client = make_guardrail(block_message="custom blocked text")
guardrail._client_cache["default-api-key"] = client
result_obj = Mock()
@ -79,9 +75,7 @@ async def test_block_not_bypassed_by_fail_open(make_guardrail, make_request_data
@pytest.mark.asyncio
async def test_apply_guardrail_mask_replaces_scanned_text(
guardrail_and_client, make_request_data
):
async def test_apply_guardrail_mask_replaces_scanned_text(guardrail_and_client, make_request_data):
guardrail, client = guardrail_and_client
result_obj = Mock()
result_obj.action = "MASK"
@ -99,9 +93,7 @@ async def test_apply_guardrail_mask_replaces_scanned_text(
@pytest.mark.asyncio
async def test_apply_guardrail_mask_targets_only_the_flagged_slot(
guardrail_and_client, make_request_data
):
async def test_apply_guardrail_mask_targets_only_the_flagged_slot(guardrail_and_client, make_request_data):
"""MASK rewrites the ``texts`` entry of the flagged segment in place; the
other scanned entries survive untouched. Confirms positional 1:1 mapping."""
guardrail, client = guardrail_and_client
@ -125,9 +117,7 @@ async def test_apply_guardrail_mask_targets_only_the_flagged_slot(
@pytest.mark.asyncio
async def test_apply_guardrail_scans_non_user_role_segments(
guardrail_and_client, make_request_data
):
async def test_apply_guardrail_scans_non_user_role_segments(guardrail_and_client, make_request_data):
"""Bypass regression: blocked content in a system/assistant/tool message
must still BLOCK. The translation layer already strips system/tool when the
guardrail is configured to skip them, so whatever remains in ``texts`` is
@ -169,9 +159,7 @@ def _tool_call(arguments, name="send_email"):
@pytest.mark.asyncio
async def test_apply_guardrail_blocks_on_tool_call_arguments(
guardrail_and_client, make_request_data
):
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
@ -200,9 +188,7 @@ async def test_apply_guardrail_blocks_on_tool_call_arguments(
@pytest.mark.asyncio
async def test_apply_guardrail_masks_tool_call_arguments_in_place(
guardrail_and_client, make_request_data
):
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
@ -231,9 +217,7 @@ async def test_apply_guardrail_masks_tool_call_arguments_in_place(
@pytest.mark.asyncio
async def test_apply_guardrail_detect_on_tool_call_args_passes_through(
guardrail_and_client, make_request_data
):
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
@ -259,9 +243,7 @@ async def test_apply_guardrail_detect_on_tool_call_args_passes_through(
@pytest.mark.asyncio
async def test_apply_guardrail_scans_tool_calls_when_no_texts(
guardrail_and_client, make_request_data
):
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)."""
@ -286,9 +268,7 @@ async def test_apply_guardrail_scans_tool_calls_when_no_texts(
@pytest.mark.asyncio
async def test_apply_guardrail_blocks_on_response_tool_call_arguments(
guardrail_and_client, make_request_data
):
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
@ -311,9 +291,7 @@ async def test_apply_guardrail_blocks_on_response_tool_call_arguments(
@pytest.mark.asyncio
async def test_apply_guardrail_mask_replaces_scanned_text_response(
guardrail_and_client, make_request_data
):
async def test_apply_guardrail_mask_replaces_scanned_text_response(guardrail_and_client, make_request_data):
guardrail, client = guardrail_and_client
result_obj = Mock()
result_obj.action = "MASK"
@ -331,9 +309,7 @@ async def test_apply_guardrail_mask_replaces_scanned_text_response(
@pytest.mark.asyncio
async def test_apply_guardrail_mask_fallback_when_action_text_is_none(
guardrail_and_client, make_request_data
):
async def test_apply_guardrail_mask_fallback_when_action_text_is_none(guardrail_and_client, make_request_data):
guardrail, client = guardrail_and_client
result_obj = Mock()
result_obj.action = "MASK"
@ -354,9 +330,7 @@ async def test_apply_guardrail_mask_fallback_when_action_text_is_none(
@pytest.mark.asyncio
async def test_apply_guardrail_no_action_passthrough(
guardrail_and_client, make_request_data
):
async def test_apply_guardrail_no_action_passthrough(guardrail_and_client, make_request_data):
guardrail, client = guardrail_and_client
result_obj = Mock()
result_obj.action = "NO_ACTION"
@ -374,9 +348,7 @@ async def test_apply_guardrail_no_action_passthrough(
@pytest.mark.asyncio
async def test_apply_guardrail_detect_action_passes_through(
guardrail_and_client, make_request_data
):
async def test_apply_guardrail_detect_action_passes_through(guardrail_and_client, make_request_data):
"""DETECT action logs a warning but does not block or mutate inputs."""
guardrail, client = guardrail_and_client
result_obj = Mock()
@ -398,9 +370,7 @@ async def test_apply_guardrail_detect_action_passes_through(
@pytest.mark.asyncio
async def test_apply_guardrail_passes_app_id_per_call(
guardrail_and_client, make_request_data
):
async def test_apply_guardrail_passes_app_id_per_call(guardrail_and_client, make_request_data):
guardrail, client = guardrail_and_client
result_obj = Mock()
result_obj.action = "NO_ACTION"
@ -410,9 +380,7 @@ async def test_apply_guardrail_passes_app_id_per_call(
await guardrail.apply_guardrail(
inputs={"texts": ["hi"]},
request_data=make_request_data(
metadata={"user_api_key_metadata": {"alice_wonderfence_app_id": "tenant-A"}}
),
request_data=make_request_data(metadata={"user_api_key_metadata": {"alice_wonderfence_app_id": "tenant-A"}}),
input_type="request",
)
kwargs = client.evaluate_prompt.call_args.kwargs
@ -422,9 +390,7 @@ async def test_apply_guardrail_passes_app_id_per_call(
@pytest.mark.asyncio
async def test_apply_guardrail_response_path_passes_app_id(
make_guardrail, make_request_data
):
async def test_apply_guardrail_response_path_passes_app_id(make_guardrail, make_request_data):
guardrail, client = make_guardrail()
guardrail._client_cache["default-api-key"] = client
result_obj = Mock()
@ -435,9 +401,7 @@ async def test_apply_guardrail_response_path_passes_app_id(
await guardrail.apply_guardrail(
inputs={"texts": ["resp"]},
request_data=make_request_data(
metadata={"user_api_key_metadata": {"alice_wonderfence_app_id": "tenant-B"}}
),
request_data=make_request_data(metadata={"user_api_key_metadata": {"alice_wonderfence_app_id": "tenant-B"}}),
input_type="response",
)
kwargs = client.evaluate_response.call_args.kwargs
@ -470,9 +434,7 @@ async def test_apply_guardrail_evaluates_every_text_without_structured_messages(
@pytest.mark.asyncio
async def test_apply_guardrail_blocks_on_earlier_user_turn(
guardrail_and_client, make_request_data
):
async def test_apply_guardrail_blocks_on_earlier_user_turn(guardrail_and_client, make_request_data):
"""Bypass regression: disallowed content in an earlier user turn followed by
a benign final turn must still BLOCK. The old last-only path only saw the
benign final message and let the request through."""
@ -539,9 +501,7 @@ async def test_apply_guardrail_blocks_when_oversized_message_trips_in_late_chunk
@pytest.mark.asyncio
async def test_apply_guardrail_no_text_short_circuits(
guardrail_and_client, make_request_data
):
async def test_apply_guardrail_no_text_short_circuits(guardrail_and_client, make_request_data):
"""Empty inputs must skip the SDK call and return inputs unchanged."""
guardrail, client = guardrail_and_client
out = await guardrail.apply_guardrail(
@ -558,9 +518,7 @@ async def test_apply_guardrail_no_text_short_circuits(
@pytest.mark.asyncio
async def test_apply_guardrail_missing_app_id_fail_closed_returns_500(
guardrail_and_client, make_request_data
):
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:
@ -575,9 +533,7 @@ async def test_apply_guardrail_missing_app_id_fail_closed_returns_500(
@pytest.mark.asyncio
async def test_apply_guardrail_missing_api_key_fail_closed_returns_500(
monkeypatch, make_guardrail, make_request_data
):
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)
@ -593,9 +549,7 @@ async def test_apply_guardrail_missing_api_key_fail_closed_returns_500(
@pytest.mark.asyncio
async def test_apply_guardrail_missing_app_id_fail_open_returns_500(
make_guardrail, make_request_data
):
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:
@ -609,9 +563,7 @@ async def test_apply_guardrail_missing_app_id_fail_open_returns_500(
@pytest.mark.asyncio
async def test_apply_guardrail_missing_api_key_fail_open_returns_500(
monkeypatch, make_guardrail, make_request_data
):
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)
@ -626,9 +578,7 @@ async def test_apply_guardrail_missing_api_key_fail_open_returns_500(
@pytest.mark.asyncio
async def test_apply_guardrail_fail_open_swallows_transport_error(
make_guardrail, make_request_data
):
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")
@ -643,9 +593,7 @@ async def test_apply_guardrail_fail_open_swallows_transport_error(
@pytest.mark.asyncio
async def test_apply_guardrail_fail_closed_returns_500(
guardrail_and_client, make_request_data
):
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")
@ -685,9 +633,7 @@ def test_build_analysis_context_falls_back_to_slash_split(monkeypatch, make_guar
raise ValueError("unknown provider")
monkeypatch.setattr(litellm, "get_llm_provider", boom)
build_analysis_context(
{"model": "myorg/custom-llm"}, guardrail.platform, guardrail._AnalysisContext
)
build_analysis_context({"model": "myorg/custom-llm"}, guardrail.platform, guardrail._AnalysisContext)
AnalysisContext = sys.modules["wonderfence_sdk.models"].AnalysisContext
kwargs = AnalysisContext.call_args.kwargs
@ -700,17 +646,13 @@ async def test_malformed_override_does_not_fail_open(make_guardrail, make_reques
"""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 = 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"]}
),
request_data=make_request_data(metadata={"alice_wonderfence_app_id": ["not", "a", "string"]}),
input_type="request",
)
assert exc.value.status_code == 500
@ -733,9 +675,7 @@ def _tool_def(description="a helpful tool", param_desc=None):
@pytest.mark.asyncio
async def test_apply_guardrail_blocks_on_tool_definition_description(
guardrail_and_client, make_request_data
):
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
@ -754,16 +694,12 @@ async def test_apply_guardrail_blocks_on_tool_definition_description(
"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"
)
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
):
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
@ -781,16 +717,12 @@ async def test_apply_guardrail_blocks_on_tool_parameter_description(
"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"
)
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
):
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):
@ -807,16 +739,12 @@ async def test_apply_guardrail_masks_tool_definition_description_in_place(
"texts": ["hi"],
"tools": [_tool_def(description="contains secret stuff")],
}
out = await guardrail.apply_guardrail(
inputs=inputs, request_data=make_request_data(), input_type="request"
)
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
):
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
@ -853,9 +781,7 @@ def _legacy_function(description="a function", param_desc=None):
@pytest.mark.asyncio
async def test_apply_guardrail_blocks_on_legacy_function_description(
guardrail_and_client, make_request_data
):
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
@ -872,18 +798,14 @@ async def test_apply_guardrail_blocks_on_legacy_function_description(
with pytest.raises(HTTPException) as exc:
await guardrail.apply_guardrail(
inputs={"texts": ["hi"]},
request_data=make_request_data(
functions=[_legacy_function(description="DISALLOWED instructions")]
),
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
):
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):
@ -898,18 +820,14 @@ async def test_apply_guardrail_blocks_on_legacy_function_parameter_description(
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")]
),
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
):
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
@ -925,18 +843,14 @@ async def test_apply_guardrail_scans_legacy_functions_when_no_other_content(
with pytest.raises(HTTPException) as exc:
await guardrail.apply_guardrail(
inputs={"texts": []},
request_data=make_request_data(
functions=[_legacy_function(description="DISALLOWED")]
),
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
):
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
@ -950,9 +864,7 @@ async def test_apply_guardrail_legacy_function_detect_does_not_mutate(
client.evaluate_prompt.side_effect = evaluate
request_data = make_request_data(
functions=[_legacy_function(description="watch this")]
)
request_data = make_request_data(functions=[_legacy_function(description="watch this")])
out = await guardrail.apply_guardrail(
inputs={"texts": ["hi"]},
request_data=request_data,
@ -963,9 +875,7 @@ async def test_apply_guardrail_legacy_function_detect_does_not_mutate(
@pytest.mark.asyncio
async def test_apply_guardrail_masks_legacy_function_description_in_place(
guardrail_and_client, make_request_data
):
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
@ -980,9 +890,7 @@ async def test_apply_guardrail_masks_legacy_function_description_in_place(
client.evaluate_prompt.side_effect = evaluate
request_data = make_request_data(
functions=[_legacy_function(description="contains secret stuff")]
)
request_data = make_request_data(functions=[_legacy_function(description="contains secret stuff")])
await guardrail.apply_guardrail(
inputs={"texts": ["hi"]},
request_data=request_data,

View file

@ -53,9 +53,7 @@ async def test_verdicts_align_one_to_one_with_segments():
actions = {"a": "BLOCK", "b": "MASK", "c": ""}
async def evaluate(text):
return _result(
actions[text], action_text="[M]" if actions[text] == "MASK" else None
)
return _result(actions[text], action_text="[M]" if actions[text] == "MASK" else None)
verdicts = await evaluate_segments(["a", "b", "c"], evaluate)
assert [v.action for v in verdicts] == ["BLOCK", "MASK", ""]
@ -187,9 +185,7 @@ async def test_block_phrase_split_across_chunk_boundary_is_detected():
async def evaluate(text):
return _result("BLOCK" if "BLOCK ME" in text else "")
verdicts = await evaluate_segments(
[segment], evaluate, max_chars=12, windows=WindowConfig(overlap=6)
)
verdicts = await evaluate_segments([segment], evaluate, max_chars=12, windows=WindowConfig(overlap=6))
assert verdicts[0].action == "BLOCK"
@ -202,9 +198,7 @@ async def test_no_overlap_window_lets_boundary_phrase_evade():
async def evaluate(text):
return _result("BLOCK" if "BLOCK ME" in text else "")
verdicts = await evaluate_segments(
[segment], evaluate, max_chars=12, windows=WindowConfig(overlap=0)
)
verdicts = await evaluate_segments([segment], evaluate, max_chars=12, windows=WindowConfig(overlap=0))
assert verdicts[0].action == ""
@ -230,13 +224,9 @@ async def test_boundary_window_mask_is_surfaced_as_detect_not_dropped():
async def evaluate(text):
# Only the boundary window sees the full "SECRET HERE".
return (
_result("MASK", action_text="[X]") if "SECRET HERE" in text else _result("")
)
return _result("MASK", action_text="[X]") if "SECRET HERE" in text else _result("")
verdicts = await evaluate_segments(
[segment], evaluate, max_chars=12, windows=WindowConfig(overlap=6)
)
verdicts = await evaluate_segments([segment], evaluate, max_chars=12, windows=WindowConfig(overlap=6))
assert verdicts[0].action == "DETECT"
@ -253,9 +243,7 @@ async def test_block_phrase_split_across_adjacent_text_segments_is_detected():
async def evaluate(text):
return _result("BLOCK" if "BLOCKME" in text else "")
verdicts = await evaluate_segments(
["BLOCK", "ME"], evaluate, windows=WindowConfig(text_segment_count=2)
)
verdicts = await evaluate_segments(["BLOCK", "ME"], evaluate, windows=WindowConfig(text_segment_count=2))
assert verdicts[0].action == "BLOCK"
@ -282,9 +270,7 @@ async def test_cross_segment_window_stays_within_text_segments():
async def evaluate(text):
return _result("BLOCK" if "BLOCKME" in text else "")
verdicts = await evaluate_segments(
["BLOCK", "ME"], evaluate, windows=WindowConfig(text_segment_count=1)
)
verdicts = await evaluate_segments(["BLOCK", "ME"], evaluate, windows=WindowConfig(text_segment_count=1))
assert [v.action for v in verdicts] == ["", ""]
@ -294,13 +280,9 @@ async def test_cross_segment_window_surfaces_mask_as_detect_without_masking():
MASK on it surfaces as DETECT and never rewrites the segment text."""
async def evaluate(text):
return (
_result("MASK", action_text="[X]") if "SECRETHERE" in text else _result("")
)
return _result("MASK", action_text="[X]") if "SECRETHERE" in text else _result("")
verdicts = await evaluate_segments(
["SECRET", "HERE"], evaluate, windows=WindowConfig(text_segment_count=2)
)
verdicts = await evaluate_segments(["SECRET", "HERE"], evaluate, windows=WindowConfig(text_segment_count=2))
assert verdicts[0].action == "DETECT"
assert verdicts[0].masked_text is None

View file

@ -110,24 +110,14 @@ def test_resolve_app_id_missing_raises():
def test_resolve_api_key_from_request_metadata_requires_override_flag():
data = _data(metadata={"alice_wonderfence_api_key": "from-req"})
assert (
resolve_api_key(
data, default_api_key="default", allow_request_metadata_override=True
)
== "from-req"
)
assert resolve_api_key(data, default_api_key="default", allow_request_metadata_override=True) == "from-req"
def test_resolve_api_key_request_metadata_ignored_when_override_disabled():
"""With override off, a caller-supplied api_key must not be honored;
falls back to the configured default instead."""
data = _data(metadata={"alice_wonderfence_api_key": "from-req"})
assert (
resolve_api_key(
data, default_api_key="default", allow_request_metadata_override=False
)
== "default"
)
assert resolve_api_key(data, default_api_key="default", allow_request_metadata_override=False) == "default"
def test_resolve_api_key_key_beats_request_even_when_override_enabled():
@ -139,12 +129,7 @@ def test_resolve_api_key_key_beats_request_even_when_override_enabled():
"user_api_key_metadata": {"alice_wonderfence_api_key": "from-key"},
}
)
assert (
resolve_api_key(
data, default_api_key="default", allow_request_metadata_override=True
)
== "from-key"
)
assert resolve_api_key(data, default_api_key="default", allow_request_metadata_override=True) == "from-key"
def test_resolve_api_key_from_key_metadata():
@ -153,12 +138,7 @@ def test_resolve_api_key_from_key_metadata():
"user_api_key_metadata": {"alice_wonderfence_api_key": "from-key"},
}
)
assert (
resolve_api_key(
data, default_api_key="default", allow_request_metadata_override=False
)
== "from-key"
)
assert resolve_api_key(data, default_api_key="default", allow_request_metadata_override=False) == "from-key"
def test_resolve_api_key_from_team_metadata():
@ -167,30 +147,18 @@ def test_resolve_api_key_from_team_metadata():
"user_api_key_team_metadata": {"alice_wonderfence_api_key": "from-team"},
}
)
assert (
resolve_api_key(
data, default_api_key="default", allow_request_metadata_override=False
)
== "from-team"
)
assert resolve_api_key(data, default_api_key="default", allow_request_metadata_override=False) == "from-team"
def test_resolve_api_key_falls_back_to_default():
data = _data(metadata={})
assert (
resolve_api_key(
data, default_api_key="default-key", allow_request_metadata_override=False
)
== "default-key"
)
assert resolve_api_key(data, default_api_key="default-key", allow_request_metadata_override=False) == "default-key"
def test_resolve_api_key_missing_everywhere_raises():
data = _data(metadata={})
with pytest.raises(WonderFenceMissingSecrets):
resolve_api_key(
data, default_api_key=None, allow_request_metadata_override=False
)
resolve_api_key(data, default_api_key=None, allow_request_metadata_override=False)
# ----------------------------- metadata fallback -----------------------------
@ -202,13 +170,9 @@ def test_resolve_reads_litellm_metadata_when_metadata_absent():
needing the request-override flag."""
data = {
"model": "gpt-4",
"litellm_metadata": {
"user_api_key_metadata": {"alice_wonderfence_app_id": "from-litellm-md"}
},
"litellm_metadata": {"user_api_key_metadata": {"alice_wonderfence_app_id": "from-litellm-md"}},
}
assert (
resolve_app_id(data, allow_request_metadata_override=False) == "from-litellm-md"
)
assert resolve_app_id(data, allow_request_metadata_override=False) == "from-litellm-md"
def test_get_metadata_merges_with_litellm_metadata_winning():
@ -241,13 +205,9 @@ def test_get_metadata_ignores_non_dict_caller_metadata():
(carrying the admin pins) is preserved."""
data = {
"metadata": "not-a-dict",
"litellm_metadata": {
"user_api_key_metadata": {"alice_wonderfence_app_id": "admin-pinned"}
},
}
assert get_metadata(data) == {
"user_api_key_metadata": {"alice_wonderfence_app_id": "admin-pinned"}
"litellm_metadata": {"user_api_key_metadata": {"alice_wonderfence_app_id": "admin-pinned"}},
}
assert get_metadata(data) == {"user_api_key_metadata": {"alice_wonderfence_app_id": "admin-pinned"}}
def test_non_dict_caller_metadata_does_not_bypass_resolution():
@ -266,12 +226,7 @@ def test_non_dict_caller_metadata_does_not_bypass_resolution():
},
}
assert resolve_app_id(data, allow_request_metadata_override=True) == "admin-pinned"
assert (
resolve_api_key(
data, default_api_key=None, allow_request_metadata_override=True
)
== "admin-key"
)
assert resolve_api_key(data, default_api_key=None, allow_request_metadata_override=True) == "admin-key"
def test_responses_route_admin_pin_beats_caller_metadata():
@ -282,9 +237,7 @@ def test_responses_route_admin_pin_beats_caller_metadata():
data = {
"model": "gpt-4",
"metadata": {"alice_wonderfence_app_id": "caller-override"},
"litellm_metadata": {
"user_api_key_metadata": {"alice_wonderfence_app_id": "admin-pinned"}
},
"litellm_metadata": {"user_api_key_metadata": {"alice_wonderfence_app_id": "admin-pinned"}},
}
assert resolve_app_id(data, allow_request_metadata_override=True) == "admin-pinned"
@ -295,16 +248,9 @@ def test_responses_route_admin_pin_beats_caller_metadata_api_key():
data = {
"model": "gpt-4",
"metadata": {"alice_wonderfence_api_key": "caller-override"},
"litellm_metadata": {
"user_api_key_metadata": {"alice_wonderfence_api_key": "admin-pinned"}
},
"litellm_metadata": {"user_api_key_metadata": {"alice_wonderfence_api_key": "admin-pinned"}},
}
assert (
resolve_api_key(
data, default_api_key="default", allow_request_metadata_override=True
)
== "admin-pinned"
)
assert resolve_api_key(data, default_api_key="default", allow_request_metadata_override=True) == "admin-pinned"
# --------------- stash storage: secret must not leak to logged payload ---------------
@ -357,12 +303,7 @@ def test_resolve_api_key_ignores_non_string_request_override():
"""A truthy non-string request override must not be returned (it would reach
the SDK and raise, which fail_open could swallow); fall back to default."""
data = _data(metadata={"alice_wonderfence_api_key": ["not", "a", "string"]})
assert (
resolve_api_key(
data, default_api_key="default", allow_request_metadata_override=True
)
== "default"
)
assert resolve_api_key(data, default_api_key="default", allow_request_metadata_override=True) == "default"
def test_resolve_app_id_non_string_request_override_raises():
@ -373,12 +314,7 @@ def test_resolve_app_id_non_string_request_override_raises():
def test_resolve_api_key_ignores_blank_string_override():
data = _data(metadata={"alice_wonderfence_api_key": " "})
assert (
resolve_api_key(
data, default_api_key="default", allow_request_metadata_override=True
)
== "default"
)
assert resolve_api_key(data, default_api_key="default", allow_request_metadata_override=True) == "default"
def test_resolve_app_id_non_string_key_metadata_falls_through():

View file

@ -7,9 +7,7 @@ from fastapi import HTTPException
@pytest.mark.asyncio
async def test_post_call_recovers_app_id_via_logging_obj_stash(
make_guardrail, make_request_data, make_logging_obj
):
async def test_post_call_recovers_app_id_via_logging_obj_stash(make_guardrail, make_request_data, make_logging_obj):
"""Reproduces the framework gap: request body metadata is dropped before
post_call. The logging_obj stash from the prior ``input_type="request"``
call must be used to resolve app_id."""
@ -32,9 +30,7 @@ async def test_post_call_recovers_app_id_via_logging_obj_stash(
# metadata — this is where the stash happens.
await guardrail.apply_guardrail(
inputs={"texts": ["hello"]},
request_data=make_request_data(
metadata={"alice_wonderfence_app_id": "tenant-X"}
),
request_data=make_request_data(metadata={"alice_wonderfence_app_id": "tenant-X"}),
input_type="request",
logging_obj=logging_obj,
)
@ -54,9 +50,7 @@ async def test_post_call_recovers_app_id_via_logging_obj_stash(
@pytest.mark.asyncio
async def test_post_call_prefers_request_data_over_stash(
make_guardrail, make_request_data, make_logging_obj
):
async def test_post_call_prefers_request_data_over_stash(make_guardrail, make_request_data, make_logging_obj):
"""If post_call's request_data still resolves (e.g. app_id from key/team
metadata), use it — don't fall back to the stash."""
guardrail, client = make_guardrail(allow_request_metadata_override=True)
@ -77,9 +71,7 @@ async def test_post_call_prefers_request_data_over_stash(
# Stash a different app_id during the request phase.
await guardrail.apply_guardrail(
inputs={"texts": ["hi"]},
request_data=make_request_data(
metadata={"alice_wonderfence_app_id": "stashed-app"}
),
request_data=make_request_data(metadata={"alice_wonderfence_app_id": "stashed-app"}),
input_type="request",
logging_obj=logging_obj,
)
@ -90,9 +82,7 @@ async def test_post_call_prefers_request_data_over_stash(
inputs={"texts": ["resp"]},
request_data={
"model": "gpt-4",
"metadata": {
"user_api_key_metadata": {"alice_wonderfence_app_id": "key-app"}
},
"metadata": {"user_api_key_metadata": {"alice_wonderfence_app_id": "key-app"}},
},
input_type="response",
logging_obj=logging_obj,
@ -122,9 +112,7 @@ async def test_post_call_without_prior_stash_raises(make_guardrail, make_logging
@pytest.mark.asyncio
async def test_post_call_does_not_borrow_sibling_stash(
make_guardrail, make_request_data, make_logging_obj
):
async def test_post_call_does_not_borrow_sibling_stash(make_guardrail, make_request_data, make_logging_obj):
"""A stricter instance must NOT inherit a sibling's stashed credentials.
Exploit being closed: a permissive writer (allow_request_metadata_override
@ -155,9 +143,7 @@ async def test_post_call_does_not_borrow_sibling_stash(
# Writer stashes caller-supplied request-body app_id (override allowed).
await g_writer.apply_guardrail(
inputs={"texts": ["hi"]},
request_data=make_request_data(
metadata={"alice_wonderfence_app_id": "caller-supplied-app"}
),
request_data=make_request_data(metadata={"alice_wonderfence_app_id": "caller-supplied-app"}),
input_type="request",
logging_obj=logging_obj,
)
@ -177,9 +163,7 @@ async def test_post_call_does_not_borrow_sibling_stash(
@pytest.mark.asyncio
async def test_stash_keyed_per_guardrail_name(
make_guardrail, make_request_data, make_logging_obj
):
async def test_stash_keyed_per_guardrail_name(make_guardrail, make_request_data, make_logging_obj):
"""Two alice_wonderfence instances on the same logging_obj must not
overwrite each other's stash — they're keyed by guardrail_name."""
g1, c1 = make_guardrail(

View file

@ -67,9 +67,7 @@ def test_tool_definition_segments_extracts_description_and_param_descriptions():
"description": "TOP_DESC",
"parameters": {
"type": "object",
"properties": {
"city": {"type": "string", "description": "PARAM_DESC"}
},
"properties": {"city": {"type": "string", "description": "PARAM_DESC"}},
},
},
}
@ -117,9 +115,7 @@ def test_function_definition_segments_extracts_descriptions_and_paths():
"description": "TOP_DESC",
"parameters": {
"type": "object",
"properties": {
"city": {"type": "string", "description": "PARAM_DESC"}
},
"properties": {"city": {"type": "string", "description": "PARAM_DESC"}},
},
},
"not-a-dict",