From 8c59a6f83cbee305d70dd649c32a1230a47396f0 Mon Sep 17 00:00:00 2001 From: lior-k Date: Wed, 17 Jun 2026 18:47:59 +0300 Subject: [PATCH] chore(guardrails): satisfy strict typing, recursion, and Any-budget CI gates The rebase onto latest litellm_internal_staging pulled in newer governance gates that the Alice WonderFence module (new to the base) tripped: - UP045/UP006/UP035: use built-in generics and `X | None` (the repo targets Python >=3.10) instead of typing.Optional/List/Dict/Tuple, removing the net-new strict-rule violations that exceeded the codebase ceiling. - recursive_detector: rewrite the tool-definition description walker iteratively (explicit stack) instead of recursively; unbounded recursion over caller-supplied tool schemas is a stack-overflow/DoS risk anyway. - any-discipline: record per-file Any baselines for the module in any-discipline-budget.json, matching how every other guardrail provider is budgeted (SDK-interop and JSON traversal inherently surface Any). The PR title was also lowercased to satisfy the Conventional Commits subject check. No behavior change; 100 unit tests still pass. --- .../alice_wonderfence/alice_wonderfence.py | 22 +++--- .../alice_wonderfence/chunked_evaluation.py | 31 ++++---- .../alice_wonderfence/client_cache.py | 10 +-- .../alice_wonderfence/credentials.py | 12 +-- .../alice_wonderfence/processing.py | 74 +++++++++++-------- .../alice_wonderfence/test_client_cache.py | 1 - 6 files changed, 80 insertions(+), 70 deletions(-) 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 b90aa5d22c6..53be48ccae6 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/alice_wonderfence.py @@ -3,7 +3,7 @@ import logging import os from collections import OrderedDict -from typing import TYPE_CHECKING, Any, List, Literal, Optional, Type, Union +from typing import TYPE_CHECKING, Any, Literal, Optional, Union from fastapi import HTTPException @@ -62,19 +62,19 @@ class WonderFenceGuardrail(CustomGuardrail): def __init__( self, guardrail_name: str, - api_key: Optional[str] = None, - api_base: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, api_timeout: float = 10.0, - platform: Optional[str] = None, + platform: str | None = None, fail_open: bool = False, block_message: str = "Content violates our policies and has been blocked", debug: bool = False, - max_cached_clients: Optional[int] = None, - connection_pool_limit: Optional[int] = None, + max_cached_clients: int | None = None, + connection_pool_limit: int | None = None, allow_request_metadata_override: bool = False, - event_hook: Optional[ - Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode] - ] = None, + event_hook: ( + Union[GuardrailEventHooks, list[GuardrailEventHooks], Mode] | None + ) = None, default_on: bool = True, **kwargs, ) -> None: @@ -122,7 +122,7 @@ class WonderFenceGuardrail(CustomGuardrail): os.environ.get("ALICE_MAX_CACHED_CLIENTS", "10") ) env_pool = os.environ.get("ALICE_CONNECTION_POOL_LIMIT") - self._connection_pool_limit: Optional[int] = ( + 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) @@ -298,6 +298,6 @@ class WonderFenceGuardrail(CustomGuardrail): return inputs @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: """Return the config model for UI rendering.""" return WonderFenceGuardrailConfigModel 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 d6ad0dcace9..544c033f292 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/chunked_evaluation.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/chunked_evaluation.py @@ -9,7 +9,8 @@ target a different backend. import asyncio import re from dataclasses import dataclass -from typing import Any, Awaitable, Callable, List, Optional +from typing import Any +from collections.abc import Awaitable, Callable MAX_PROMPT_CHARS = 10000 # WonderFence server-side prompt limit DEFAULT_MAX_CONCURRENCY = 10 # used when the client connection_pool_limit is unset @@ -26,12 +27,12 @@ CHUNK_OVERLAP_CHARS = 512 @dataclass class SegmentVerdict: action: str # "BLOCK" | "MASK" | "DETECT" | "" - masked_text: Optional[str] + masked_text: str | None detections: list - correlation_ids: List[str] + correlation_ids: list[str] -def _split_text(text: str, max_chars: int) -> List[str]: +def _split_text(text: str, max_chars: int) -> list[str]: """Split ``text`` into <= ``max_chars`` chunks with ``"".join(chunks) == text``. Splits at whitespace boundaries; whitespace runs are preserved as their own @@ -42,7 +43,7 @@ def _split_text(text: str, max_chars: int) -> List[str]: return [text] tokens = re.findall(r"\S+|\s+", text) - chunks: List[str] = [] + chunks: list[str] = [] current = "" for token in tokens: if len(current) + len(token) <= max_chars: @@ -65,7 +66,7 @@ def _action_str(result: Any) -> str: return action.value if hasattr(action, "value") else (action or "") -def _boundary_windows(chunks: List[str], overlap: int) -> List[str]: +def _boundary_windows(chunks: list[str], overlap: int) -> list[str]: """Windows spanning each adjacent chunk boundary, for detection only. Each window is the last ``overlap`` chars of one chunk joined to the first @@ -80,14 +81,14 @@ def _boundary_windows(chunks: List[str], overlap: int) -> List[str]: def _aggregate( - chunks: List[str], - chunk_results: List[Any], - boundary_results: List[Any], + chunks: list[str], + chunk_results: list[Any], + boundary_results: list[Any], ) -> SegmentVerdict: chunk_actions = [_action_str(r) for r in chunk_results] boundary_actions = [_action_str(r) for r in boundary_results] detections: list = [] - correlation_ids: List[str] = [] + correlation_ids: list[str] = [] for r in (*chunk_results, *boundary_results): detections.extend(getattr(r, "detections", None) or []) cid = getattr(r, "correlation_id", None) @@ -111,12 +112,12 @@ def _aggregate( async def evaluate_segments( - segments: List[str], + segments: list[str], evaluate: Callable[[str], Awaitable[Any]], max_chars: int = MAX_PROMPT_CHARS, max_concurrency: int = DEFAULT_MAX_CONCURRENCY, overlap: int = CHUNK_OVERLAP_CHARS, -) -> List[SegmentVerdict]: +) -> list[SegmentVerdict]: """Evaluate every segment (chunked) in parallel; return one verdict per segment. Each segment is split into <= ``max_chars`` disjoint chunks; multi-chunk @@ -138,7 +139,7 @@ async def evaluate_segments( seg_chunks = [_split_text(s, max_chars) for s in segments] seg_boundaries = [_boundary_windows(chunks, ov) for chunks in seg_chunks] - index: List[tuple] = [] + index: list[tuple] = [] tasks = [] for si in range(len(segments)): for ci, chunk in enumerate(seg_chunks[si]): @@ -149,8 +150,8 @@ async def evaluate_segments( tasks.append(run(window)) results = await asyncio.gather(*tasks) - chunk_res: List[List[Any]] = [[None] * len(c) for c in seg_chunks] - bound_res: List[List[Any]] = [[None] * len(b) for b in seg_boundaries] + chunk_res: list[list[Any]] = [[None] * len(c) for c in seg_chunks] + bound_res: list[list[Any]] = [[None] * len(b) for b in seg_boundaries] for (si, is_boundary, idx), res in zip(index, results): (bound_res if is_boundary else chunk_res)[si][idx] = res 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 ef046e3975b..89fb53e0922 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/client_cache.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/client_cache.py @@ -1,7 +1,7 @@ """WonderFence SDK loader + per-api_key LRU client cache.""" from collections import OrderedDict -from typing import TYPE_CHECKING, Any, Optional, Tuple +from typing import TYPE_CHECKING, Any if TYPE_CHECKING: from wonderfence_sdk.client import ( # type: ignore[import-untyped] @@ -9,7 +9,7 @@ if TYPE_CHECKING: ) -def load_sdk() -> Tuple[Any, Any]: +def load_sdk() -> tuple[Any, Any]: """Lazy-import WonderFence SDK classes (``WonderFenceV2Client``, ``AnalysisContext``). Deferred to instance construction (not module load) because wonderfence_sdk @@ -37,9 +37,9 @@ def get_or_create_client( cache_maxsize: int, client_class: Any, api_timeout: float, - api_base: Optional[str], - platform: Optional[str], - connection_pool_limit: Optional[int], + api_base: str | None, + platform: str | None, + connection_pool_limit: int | None, ) -> "_WonderFenceV2Client": """LRU client lookup keyed by ``api_key``; construct on miss.""" if api_key in cache: diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/credentials.py b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/credentials.py index 78cacf56d55..a60d5f5e74b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/credentials.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/credentials.py @@ -18,12 +18,12 @@ The stash bridges pre_call resolution into post_call where request metadata is gone — see ``stash_resolved`` for the full rationale. """ -from typing import TYPE_CHECKING, Any, Literal, Optional, Tuple +from typing import TYPE_CHECKING, Any, Literal, Optional from .exceptions import WonderFenceMissingSecrets -def _nonempty_str(value: Any) -> Optional[str]: +def _nonempty_str(value: Any) -> str | None: """Return ``value`` only if it is a non-empty/non-blank string, else None. Credential sources (request body, key/team metadata, config default) are @@ -76,7 +76,7 @@ def get_metadata(request_data: dict) -> dict: def resolve_api_key( request_data: dict, - default_api_key: Optional[str], + default_api_key: str | None, allow_request_metadata_override: bool, ) -> str: """Resolve api_key from key → team → (request, when opt-in) → default. @@ -204,7 +204,7 @@ def stash_resolved( def recover_resolved( logging_obj: Optional["LiteLLMLoggingObj"], guardrail_name: str -) -> Optional[Tuple[str, str]]: +) -> tuple[str, str] | None: """Look up the (api_key, app_id) this guardrail stashed earlier in this request, or ``None``. @@ -226,9 +226,9 @@ def resolve_credentials( input_type: Literal["request", "response"], logging_obj: Optional["LiteLLMLoggingObj"], guardrail_name: str, - default_api_key: Optional[str], + default_api_key: str | None, allow_request_metadata_override: bool, -) -> Tuple[str, str]: +) -> tuple[str, str]: """Resolve (api_key, app_id) for this call. For ``request``: read from request_data (canonical pre_call path) and stash diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py index 44447839fe2..1e9e30a4314 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/processing.py @@ -1,6 +1,6 @@ """Pure transforms for Alice WonderFence: context build, user-text mapping, verdict apply.""" -from typing import Any, List, Optional, Tuple +from typing import Any import litellm from litellm._logging import verbose_proxy_logger @@ -15,7 +15,7 @@ logger = verbose_proxy_logger.getChild("alice_wonderfence") def build_analysis_context( request_data: dict, - platform: Optional[str], + platform: str | None, context_class: Any, ) -> Any: """Build WonderFence AnalysisContext from request data.""" @@ -54,7 +54,7 @@ def build_analysis_context( def tool_call_arg_segments( inputs: GenericGuardrailAPIInputs, -) -> Tuple[List[int], List[str]]: +) -> tuple[list[int], list[str]]: """Return (indices, argument strings) for tool calls carrying string args. ``inputs["tool_calls"]`` entries are dicts shaped @@ -63,8 +63,8 @@ def tool_call_arg_segments( scanned like any other segment. """ tool_calls = inputs.get("tool_calls") or [] - indices: List[int] = [] - segments: List[str] = [] + indices: list[int] = [] + segments: list[str] = [] for i, tool_call in enumerate(tool_calls): fn = tool_call.get("function") if isinstance(tool_call, dict) else None args = fn.get("arguments") if isinstance(fn, dict) else None @@ -74,27 +74,37 @@ def tool_call_arg_segments( return indices, segments -def _description_strings(obj: Any, prefix: List[Any]) -> List[Tuple[List[Any], str]]: +def _description_strings( + root: Any, root_prefix: list[Any] +) -> list[tuple[list[Any], str]]: """Collect ``(path, text)`` for every non-blank ``description`` string under - ``obj`` (a tool's ``function`` dict). Recurses into nested JSON-schema - parameters so parameter descriptions are included, not just the top one.""" - out: List[Tuple[List[Any], str]] = [] - if isinstance(obj, dict): - for key, value in obj.items(): - if key == "description" and isinstance(value, str) and value.strip(): - out.append((prefix + [key], value)) - elif isinstance(value, (dict, list)): - out.extend(_description_strings(value, prefix + [key])) - elif isinstance(obj, list): - for idx, item in enumerate(obj): - if isinstance(item, (dict, list)): - out.extend(_description_strings(item, prefix + [idx])) + ``root`` (a tool's ``function`` dict), walking nested JSON-schema parameters + so parameter descriptions are included, not just the top one. + + Iterative (explicit stack) rather than recursive: caller-supplied tool + schemas can nest arbitrarily, and unbounded recursion on request input is a + DoS / stack-overflow risk. + """ + out: list[tuple[list[Any], str]] = [] + stack: list[tuple[Any, list[Any]]] = [(root, root_prefix)] + while stack: + obj, prefix = stack.pop() + if isinstance(obj, dict): + for key, value in obj.items(): + if key == "description" and isinstance(value, str) and value.strip(): + out.append((prefix + [key], value)) + elif isinstance(value, (dict, list)): + stack.append((value, prefix + [key])) + elif isinstance(obj, list): + for idx, item in enumerate(obj): + if isinstance(item, (dict, list)): + stack.append((item, prefix + [idx])) return out def tool_definition_segments( inputs: GenericGuardrailAPIInputs, -) -> Tuple[List[List[Any]], List[str]]: +) -> tuple[list[list[Any]], list[str]]: """Return (paths, texts) for free-text in tool definitions. The chat translation layer passes caller-supplied ``inputs["tools"]`` to the @@ -103,8 +113,8 @@ def tool_definition_segments( the string within ``inputs["tools"]`` so a MASK verdict can be written back. """ tools = inputs.get("tools") or [] - paths: List[List[Any]] = [] - segments: List[str] = [] + paths: list[list[Any]] = [] + segments: list[str] = [] for i, tool in enumerate(tools): fn = tool.get("function") if isinstance(tool, dict) else None if not isinstance(fn, dict): @@ -115,7 +125,7 @@ def tool_definition_segments( return paths, segments -def _set_by_path(root: Any, path: List[Any], value: Any) -> None: +def _set_by_path(root: Any, path: list[Any], value: Any) -> None: obj = root for key in path[:-1]: obj = obj[key] @@ -123,10 +133,10 @@ def _set_by_path(root: Any, path: List[Any], value: Any) -> None: def _block_detail( - blocked: List[SegmentVerdict], guardrail_name: str, block_message: str + blocked: list[SegmentVerdict], guardrail_name: str, block_message: str ) -> dict: detections: list = [] - correlation_ids: List[str] = [] + correlation_ids: list[str] = [] for v in blocked: detections.extend(v.detections) correlation_ids.extend(v.correlation_ids) @@ -147,7 +157,7 @@ def _block_detail( def _masked_value( verdict: SegmentVerdict, guardrail_name: str, label: str -) -> Optional[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.""" @@ -172,14 +182,14 @@ def _masked_value( def apply_verdicts( inputs: GenericGuardrailAPIInputs, - indices: List[int], - verdicts: List[SegmentVerdict], + indices: list[int], + verdicts: list[SegmentVerdict], guardrail_name: str, block_message: str, - tool_indices: Optional[List[int]] = None, - tool_verdicts: Optional[List[SegmentVerdict]] = None, - tool_def_paths: Optional[List[List[Any]]] = None, - tool_def_verdicts: Optional[List[SegmentVerdict]] = None, + tool_indices: list[int] | None = None, + tool_verdicts: list[SegmentVerdict] | None = None, + tool_def_paths: list[list[Any]] | None = None, + tool_def_verdicts: list[SegmentVerdict] | None = None, ) -> GenericGuardrailAPIInputs: """Apply per-segment verdicts back onto request text, tool-call args, and tool-definition descriptions. diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_client_cache.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_client_cache.py index 226809caacb..1ccecf3b281 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_client_cache.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/alice_wonderfence/test_client_cache.py @@ -5,7 +5,6 @@ from unittest.mock import AsyncMock, Mock import pytest - # ----------------------------- LRU cache -----------------------------