mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
style: format litellm/proxy/guardrails/guardrail_registry.py with ruff to pass CI lint gate
This commit is contained in:
parent
209f222231
commit
9b721728b8
1 changed files with 26 additions and 10 deletions
|
|
@ -8,8 +8,6 @@ from datetime import datetime, timezone
|
|||
from itertools import chain, count
|
||||
from typing import TYPE_CHECKING, Final, Literal, Optional, Protocol, TypeAlias, cast
|
||||
|
||||
from pydantic import ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -52,6 +50,7 @@ from litellm.types.guardrails import (
|
|||
LitellmParams,
|
||||
SupportedGuardrailIntegrations,
|
||||
)
|
||||
from pydantic import ValidationError
|
||||
|
||||
from .guardrail_hooks.llm_as_a_judge import (
|
||||
initialize_guardrail as initialize_llm_as_a_judge,
|
||||
|
|
@ -76,7 +75,9 @@ class _GuardrailRowLike(Protocol):
|
|||
def __iter__(self) -> Iterator[tuple[str, object]]: ...
|
||||
|
||||
|
||||
def _guardrail_table(prisma_client: PrismaClient) -> "TableActions[prisma_models.LiteLLM_GuardrailsTable]":
|
||||
def _guardrail_table(
|
||||
prisma_client: PrismaClient,
|
||||
) -> "TableActions[prisma_models.LiteLLM_GuardrailsTable]":
|
||||
"""Typed view of the guardrails table actions exposed by the Prisma repository."""
|
||||
return GuardrailsRepository(prisma_client).table
|
||||
|
||||
|
|
@ -157,7 +158,9 @@ def get_guardrail_initializer_from_hooks():
|
|||
if isinstance(registry, dict):
|
||||
discovered_initializers.update(registry)
|
||||
verbose_proxy_logger.debug(
|
||||
"Found guardrail_initializer_registry in %s: %s", module_path, list(registry.keys())
|
||||
"Found guardrail_initializer_registry in %s: %s",
|
||||
module_path,
|
||||
list(registry.keys()),
|
||||
)
|
||||
|
||||
# Check for standalone initialize_guardrail function (fallback for directory-based guardrails)
|
||||
|
|
@ -443,14 +446,20 @@ def _as_callback_tuple(
|
|||
|
||||
|
||||
def _configure_callback_scoping(
|
||||
custom_guardrail_callback: CustomGuardrail, guardrail_name: str, litellm_params: LitellmParams
|
||||
custom_guardrail_callback: CustomGuardrail,
|
||||
guardrail_name: str,
|
||||
litellm_params: LitellmParams,
|
||||
) -> None:
|
||||
for scoping_param in (
|
||||
"skip_system_message_in_guardrail",
|
||||
"skip_tool_message_in_guardrail",
|
||||
"scan_only_tool_results",
|
||||
):
|
||||
setattr(custom_guardrail_callback, scoping_param, getattr(litellm_params, scoping_param, None))
|
||||
setattr(
|
||||
custom_guardrail_callback,
|
||||
scoping_param,
|
||||
getattr(litellm_params, scoping_param, None),
|
||||
)
|
||||
scan_only_tool_results_enabled: Final = effective_scan_only_tool_results_for_guardrail(custom_guardrail_callback)
|
||||
if scan_only_tool_results_enabled and not custom_guardrail_callback.supports_scan_only_tool_results():
|
||||
raise ValueError(
|
||||
|
|
@ -493,7 +502,10 @@ class InMemoryGuardrailHandler:
|
|||
"""
|
||||
|
||||
def _stable_guardrail_id(self, guardrail_name: str) -> str:
|
||||
seeds: Final = chain((guardrail_name,), (f"{guardrail_name}:{occurrence}" for occurrence in count(1)))
|
||||
seeds: Final = chain(
|
||||
(guardrail_name,),
|
||||
(f"{guardrail_name}:{occurrence}" for occurrence in count(1)),
|
||||
)
|
||||
candidate_ids: Final = (str(uuid.uuid5(CONFIG_GUARDRAIL_ID_NAMESPACE, seed.encode("utf-8"))) for seed in seeds)
|
||||
return next(candidate_id for candidate_id in candidate_ids if candidate_id not in self.IN_MEMORY_GUARDRAILS)
|
||||
|
||||
|
|
@ -527,7 +539,7 @@ class InMemoryGuardrailHandler:
|
|||
else:
|
||||
litellm_params = litellm_params_data
|
||||
|
||||
if "category_thresholds" in litellm_params_data and litellm_params_data["category_thresholds"]:
|
||||
if litellm_params_data.get("category_thresholds"):
|
||||
lakera_category_thresholds: Final = LakeraCategoryThresholds(**litellm_params_data["category_thresholds"])
|
||||
litellm_params.category_thresholds = lakera_category_thresholds
|
||||
|
||||
|
|
@ -857,7 +869,9 @@ class InMemoryGuardrailHandler:
|
|||
)
|
||||
try:
|
||||
self.initialize_guardrail(
|
||||
guardrail=previous_guardrail, config_file_path=config_file_path, source=previous_source
|
||||
guardrail=previous_guardrail,
|
||||
config_file_path=config_file_path,
|
||||
source=previous_source,
|
||||
)
|
||||
except Exception: # noqa: BLE001 # the original failure must propagate even if the restore breaks
|
||||
verbose_proxy_logger.exception("Restoring previous guardrail %s also failed", guardrail_id)
|
||||
|
|
@ -876,7 +890,9 @@ class InMemoryGuardrailHandler:
|
|||
if self._has_guardrail_params_changed(guardrail_id, guardrail):
|
||||
guardrail_name: Final = guardrail.get("guardrail_name", "Unknown")
|
||||
verbose_proxy_logger.info(
|
||||
"Guardrail '%s' (ID: %s) params changed, re-initializing...", guardrail_name, guardrail_id
|
||||
"Guardrail '%s' (ID: %s) params changed, re-initializing...",
|
||||
guardrail_name,
|
||||
guardrail_id,
|
||||
)
|
||||
return self.reinitialize_guardrail(
|
||||
guardrail=guardrail,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue