diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/__init__.py index b9679cfd42e..65508de082a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py index 166c5bfe33a..10ae49ccc51 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/chunked_evaluation.py b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/chunked_evaluation.py index cbe8b1ee140..a2adb0b9cf2 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/chunked_evaluation.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/chunked_evaluation.py @@ -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))] diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/client_cache.py b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/client_cache.py index 93a3f3c0896..a81d00c9216 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/client_cache.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/client_cache.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/credentials.py b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/credentials.py index ac5604928e6..c540b1b0e5b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/credentials.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/credentials.py @@ -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``. diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py index 0e635c78c90..a44da75c463 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py @@ -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): diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail.py index 7a12684eccd..1b171db04ef 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_apply_guardrail.py @@ -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, diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_chunked_evaluation.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_chunked_evaluation.py index 4e458980e58..dcdcbf2f721 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_chunked_evaluation.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_chunked_evaluation.py @@ -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 diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_credentials.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_credentials.py index ab142cffd04..9bb60c1c16d 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_credentials.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_credentials.py @@ -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(): diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_post_call_bridge.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_post_call_bridge.py index 5281282f572..ce1849e2ccf 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_post_call_bridge.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_post_call_bridge.py @@ -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( diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_processing.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_processing.py index c5309713af6..b5d2c2eaf27 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_processing.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_processing.py @@ -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",