mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
8c99df4cfc
commit
4eaea1dd8a
5 changed files with 71 additions and 53 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue