mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
8933969af4
commit
c81292bcd5
11 changed files with 108 additions and 363 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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))]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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``.
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue