mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
df5cd8dab3
commit
8c59a6f83c
6 changed files with 80 additions and 70 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -5,7 +5,6 @@ from unittest.mock import AsyncMock, Mock
|
|||
|
||||
import pytest
|
||||
|
||||
|
||||
# ----------------------------- LRU cache -----------------------------
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue