chore(guardrails): satisfy strict, PLR0913, and basedpyright gates after rebase

Reduce PLR0913 by collapsing get_or_create_client and resolve_credentials
arg lists into frozen ClientBuildSpec / CredentialConfig dataclasses, drop
ANN401 by narrowing Any to object/Callable on the SDK-facing helpers, remove
a redundant UP037 string annotation, and suppress the four net-new
reportMissingTypeStubs from the optional wonderfence_sdk imports with
pyright: ignore (the repo disables type: ignore for basedpyright).
This commit is contained in:
lior-k 2026-06-24 22:20:59 +03:00
parent 8c99df4cfc
commit 4eaea1dd8a
No known key found for this signature in database
5 changed files with 71 additions and 53 deletions

View file

@ -22,8 +22,8 @@ from litellm.types.proxy.guardrails.guardrail_hooks.alice_wonderfence import (
from litellm.types.utils import GenericGuardrailAPIInputs
from .chunked_evaluation import DEFAULT_MAX_CONCURRENCY, evaluate_segments
from .client_cache import get_or_create_client, load_sdk
from .credentials import resolve_credentials
from .client_cache import ClientBuildSpec, get_or_create_client, load_sdk
from .credentials import CredentialConfig, resolve_credentials
from .exceptions import WonderFenceBlockedError, WonderFenceMissingSecrets
from .processing import (
apply_verdicts,
@ -34,7 +34,7 @@ from .processing import (
)
if TYPE_CHECKING:
from wonderfence_sdk.client import ( # type: ignore[import-untyped]
from wonderfence_sdk.client import ( # type: ignore[import-untyped] # pyright: ignore[reportMissingTypeStubs]
WonderFenceV2Client as _WonderFenceV2Client,
)
@ -118,7 +118,7 @@ class WonderFenceGuardrail(CustomGuardrail):
if debug:
logger.setLevel(logging.DEBUG)
self._client_cache: "OrderedDict[str, _WonderFenceV2Client]" = OrderedDict()
self._client_cache: OrderedDict[str, _WonderFenceV2Client] = OrderedDict()
self._client_cache_maxsize = max_cached_clients or int(
os.environ.get("ALICE_MAX_CACHED_CLIENTS", "10")
)
@ -159,11 +159,13 @@ class WonderFenceGuardrail(CustomGuardrail):
api_key,
self._client_cache,
self._client_cache_maxsize,
self._WonderFenceV2Client,
self.api_timeout,
self.api_base,
self.platform,
self._connection_pool_limit,
ClientBuildSpec(
client_class=self._WonderFenceV2Client,
api_timeout=self.api_timeout,
api_base=self.api_base,
platform=self.platform,
connection_pool_limit=self._connection_pool_limit,
),
)
@log_guardrail_information
@ -202,9 +204,11 @@ class WonderFenceGuardrail(CustomGuardrail):
request_data,
input_type,
logging_obj,
self.guardrail_name,
self.api_key,
self.allow_request_metadata_override,
CredentialConfig(
guardrail_name=self.guardrail_name,
default_api_key=self.api_key,
allow_request_metadata_override=self.allow_request_metadata_override,
),
)
client = await self._get_client(api_key)
context = build_analysis_context(
@ -213,14 +217,14 @@ class WonderFenceGuardrail(CustomGuardrail):
if input_type == "request":
async def evaluate(text: str) -> Any:
async def evaluate(text: str) -> object:
return await client.evaluate_prompt(
app_id=app_id, prompt=text, context=context, custom_fields=None
)
else:
async def evaluate(text: str) -> Any:
async def evaluate(text: str) -> object:
return await client.evaluate_response(
app_id=app_id,
response=text,

View file

@ -61,7 +61,7 @@ def _split_text(text: str, max_chars: int) -> list[str]:
return chunks
def _action_str(result: Any) -> str:
def _action_str(result: object) -> str:
action = getattr(result, "action", "")
return action.value if hasattr(action, "value") else (action or "")
@ -130,7 +130,7 @@ async def evaluate_segments(
"""
semaphore = asyncio.Semaphore(max_concurrency)
async def run(text: str) -> Any:
async def run(text: str) -> object:
async with semaphore:
return await evaluate(text)

View file

@ -1,14 +1,26 @@
"""WonderFence SDK loader + per-api_key LRU client cache."""
from collections import OrderedDict
from typing import TYPE_CHECKING, Any
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Callable
if TYPE_CHECKING:
from wonderfence_sdk.client import ( # type: ignore[import-untyped]
from wonderfence_sdk.client import ( # type: ignore[import-untyped] # pyright: ignore[reportMissingTypeStubs]
WonderFenceV2Client as _WonderFenceV2Client,
)
@dataclass(frozen=True)
class ClientBuildSpec:
"""How to construct a WonderFenceV2Client on a cache miss."""
client_class: Callable[..., object]
api_timeout: float
api_base: str | None
platform: str | None
connection_pool_limit: int | None
def load_sdk() -> tuple[Any, Any]:
"""Lazy-import WonderFence SDK classes (``WonderFenceV2Client``, ``AnalysisContext``).
@ -18,10 +30,10 @@ def load_sdk() -> tuple[Any, Any]:
on the instance so per-call hot paths don't re-trigger the import machinery.
"""
try:
from wonderfence_sdk.client import ( # type: ignore[import-untyped]
from wonderfence_sdk.client import ( # type: ignore[import-untyped] # pyright: ignore[reportMissingTypeStubs]
WonderFenceV2Client,
)
from wonderfence_sdk.models import ( # type: ignore[import-untyped]
from wonderfence_sdk.models import ( # type: ignore[import-untyped] # pyright: ignore[reportMissingTypeStubs]
AnalysisContext,
)
except ImportError as e:
@ -35,11 +47,7 @@ def get_or_create_client(
api_key: str,
cache: "OrderedDict[str, _WonderFenceV2Client]",
cache_maxsize: int,
client_class: Any,
api_timeout: float,
api_base: str | None,
platform: str | None,
connection_pool_limit: int | None,
spec: ClientBuildSpec,
) -> "_WonderFenceV2Client":
"""LRU client lookup keyed by ``api_key``; construct on miss."""
if api_key in cache:
@ -48,16 +56,16 @@ def get_or_create_client(
client_kwargs: dict = {
"api_key": api_key,
"api_timeout": round(api_timeout),
"api_timeout": round(spec.api_timeout),
}
if api_base:
client_kwargs["base_url"] = api_base
if platform:
client_kwargs["platform"] = platform
if connection_pool_limit is not None:
client_kwargs["connection_pool_limit"] = connection_pool_limit
if spec.api_base:
client_kwargs["base_url"] = spec.api_base
if spec.platform:
client_kwargs["platform"] = spec.platform
if spec.connection_pool_limit is not None:
client_kwargs["connection_pool_limit"] = spec.connection_pool_limit
client = client_class(**client_kwargs)
client = spec.client_class(**client_kwargs)
cache[api_key] = client
if len(cache) > cache_maxsize:

View file

@ -18,12 +18,22 @@ 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
from dataclasses import dataclass
from typing import TYPE_CHECKING, Literal, Optional
from .exceptions import WonderFenceMissingSecrets
def _nonempty_str(value: Any) -> str | None:
@dataclass(frozen=True)
class CredentialConfig:
"""Per-guardrail config consulted during credential resolution."""
guardrail_name: str
default_api_key: str | None
allow_request_metadata_override: bool
def _nonempty_str(value: object) -> 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
@ -225,9 +235,7 @@ def resolve_credentials(
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional["LiteLLMLoggingObj"],
guardrail_name: str,
default_api_key: str | None,
allow_request_metadata_override: bool,
config: CredentialConfig,
) -> tuple[str, str]:
"""Resolve (api_key, app_id) for this call.
@ -241,22 +249,20 @@ def resolve_credentials(
stash for values supplied in the original request body's metadata, which
the framework drops before post_call.
"""
default_api_key = config.default_api_key
allow_override = config.allow_request_metadata_override
if input_type == "request":
api_key = resolve_api_key(
request_data, default_api_key, allow_request_metadata_override
)
app_id = resolve_app_id(request_data, allow_request_metadata_override)
stash_resolved(logging_obj, guardrail_name, api_key, app_id)
api_key = resolve_api_key(request_data, default_api_key, allow_override)
app_id = resolve_app_id(request_data, allow_override)
stash_resolved(logging_obj, config.guardrail_name, api_key, app_id)
return api_key, app_id
try:
return (
resolve_api_key(
request_data, default_api_key, allow_request_metadata_override
),
resolve_app_id(request_data, allow_request_metadata_override),
resolve_api_key(request_data, default_api_key, allow_override),
resolve_app_id(request_data, allow_override),
)
except WonderFenceMissingSecrets:
recovered = recover_resolved(logging_obj, guardrail_name)
recovered = recover_resolved(logging_obj, config.guardrail_name)
if recovered is None:
raise
return recovered

View file

@ -1,6 +1,6 @@
"""Pure transforms for Alice WonderFence: context build, user-text mapping, verdict apply."""
from typing import Any
from typing import Any, Callable
import litellm
from litellm._logging import verbose_proxy_logger
@ -16,8 +16,8 @@ logger = verbose_proxy_logger.getChild("alice_wonderfence")
def build_analysis_context(
request_data: dict,
platform: str | None,
context_class: Any,
) -> Any:
context_class: Callable[..., object],
) -> object:
"""Build WonderFence AnalysisContext from request data."""
metadata = get_metadata(request_data)
model_str = request_data.get("model", "")
@ -75,7 +75,7 @@ def tool_call_arg_segments(
def _description_strings(
root: Any, root_prefix: list[Any]
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
@ -146,7 +146,7 @@ def function_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: object) -> None:
obj = root
for key in path[:-1]:
obj = obj[key]