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.
This commit is contained in:
lior-k 2026-06-17 18:47:59 +03:00
parent df5cd8dab3
commit 8c59a6f83c
No known key found for this signature in database
6 changed files with 80 additions and 70 deletions

View file

@ -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

View file

@ -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

View file

@ -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:

View file

@ -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

View file

@ -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.

View file

@ -5,7 +5,6 @@ from unittest.mock import AsyncMock, Mock
import pytest
# ----------------------------- LRU cache -----------------------------