fix(lint): fix UP045, I001, PLR0913 strict-budget violations

- Replace Optional[X] with X | None throughout asqav.py, prometheus.py,
  usage_endpoints.py, ui_sso.py (UP045)
- Fix import ordering in estimator_core.py (I001)
- Consolidate constructor args in bias_hallucination_estimator into
  config dataclasses (RiskThresholds, RiskWeights, GuardrailConfig,
  FetchConfig, VectorStoreClientConfig) to satisfy PLR0913

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
Sameer Kankute 2026-06-24 16:08:43 +05:30
parent 41abaa7aad
commit 9fb239d029
No known key found for this signature in database
9 changed files with 195 additions and 122 deletions

View file

@ -22,7 +22,7 @@ import threading
import time
import traceback
from datetime import datetime, timezone
from typing import Any, BinaryIO, Optional
from typing import Any, BinaryIO
from litellm._logging import verbose_logger
from litellm.integrations.custom_logger import CustomLogger
@ -63,7 +63,7 @@ def _read_tail(fh: BinaryIO, size: int) -> bytes:
chunk_size *= 2
def _content_digest(value: object) -> Optional[str]:
def _content_digest(value: object) -> str | None:
"""Return a SHA-256 hex digest of a content value, or None if empty."""
if value is None:
return None
@ -74,8 +74,8 @@ def _content_digest(value: object) -> Optional[str]:
def _extract_loggable(
kwargs: dict[str, Any],
response_obj: object,
start_time: Optional[datetime],
end_time: Optional[datetime],
start_time: datetime | None,
end_time: datetime | None,
status: str,
) -> dict[str, Any]:
"""Pull metadata + digests out of a callback invocation.
@ -125,7 +125,7 @@ def _extract_loggable(
pass
# Timing
latency_ms: Optional[int] = None
latency_ms: int | None = None
try:
if start_time is not None and end_time is not None:
latency_ms = int((end_time - start_time).total_seconds() * 1000)
@ -133,11 +133,11 @@ def _extract_loggable(
pass
# Usage
prompt_tokens: Optional[int] = None
completion_tokens: Optional[int] = None
total_tokens: Optional[int] = None
finish_reason: Optional[str] = None
provider_request_id: Optional[str] = None
prompt_tokens: int | None = None
completion_tokens: int | None = None
total_tokens: int | None = None
finish_reason: str | None = None
provider_request_id: str | None = None
try:
if hasattr(response_obj, "usage") and response_obj.usage:
prompt_tokens = response_obj.usage.prompt_tokens
@ -153,9 +153,9 @@ def _extract_loggable(
pass
# Content digests (not content itself)
messages_digest: Optional[str] = _content_digest(messages)
messages_digest: str | None = _content_digest(messages)
response_content_digest: Optional[str] = None
response_content_digest: str | None = None
try:
if hasattr(response_obj, "choices") and response_obj.choices:
content = response_obj.choices[0].message.content
@ -164,7 +164,7 @@ def _extract_loggable(
pass
# Standard logging payload may carry call_id / litellm_call_id
call_id: Optional[str] = None
call_id: str | None = None
try:
slp: Any = kwargs.get("standard_logging_object")
if slp and isinstance(slp, dict):
@ -215,7 +215,7 @@ class AsqavLogger(CustomLogger):
def __init__(
self,
log_path: Optional[str] = None,
log_path: str | None = None,
redact_content: bool = True,
) -> None:
super().__init__()
@ -277,8 +277,8 @@ class AsqavLogger(CustomLogger):
self,
kwargs: dict[str, Any],
response_obj: object,
start_time: Optional[datetime],
end_time: Optional[datetime],
start_time: datetime | None,
end_time: datetime | None,
status: str,
) -> None:
"""Build one audit record and append it to the JSONL log.
@ -373,8 +373,8 @@ class AsqavLogger(CustomLogger):
self,
kwargs: dict[str, Any],
response_obj: object,
start_time: Optional[datetime],
end_time: Optional[datetime],
start_time: datetime | None,
end_time: datetime | None,
) -> None:
self._build_and_append(kwargs, response_obj, start_time, end_time, "success")
@ -382,8 +382,8 @@ class AsqavLogger(CustomLogger):
self,
kwargs: dict[str, Any],
response_obj: object,
start_time: Optional[datetime],
end_time: Optional[datetime],
start_time: datetime | None,
end_time: datetime | None,
) -> None:
self._build_and_append(kwargs, response_obj, start_time, end_time, "failure")
@ -391,8 +391,8 @@ class AsqavLogger(CustomLogger):
self,
kwargs: dict[str, Any],
response_obj: object,
start_time: Optional[datetime],
end_time: Optional[datetime],
start_time: datetime | None,
end_time: datetime | None,
) -> None:
await asyncio.to_thread(
self._build_and_append,
@ -407,8 +407,8 @@ class AsqavLogger(CustomLogger):
self,
kwargs: dict[str, Any],
response_obj: object,
start_time: Optional[datetime],
end_time: Optional[datetime],
start_time: datetime | None,
end_time: datetime | None,
) -> None:
await asyncio.to_thread(
self._build_and_append,
@ -423,7 +423,7 @@ class AsqavLogger(CustomLogger):
# Chain verification (utility; not called on the hot path)
# ------------------------------------------------------------------
def verify_chain(self, log_path: Optional[str] = None) -> tuple[bool, str]:
def verify_chain(self, log_path: str | None = None) -> tuple[bool, str]:
"""Verify the integrity of the audit log at log_path.
Returns (True, "ok") when every record's hash matches its content and

View file

@ -2174,7 +2174,7 @@ class PrometheusLogger(CustomLogger):
model_id: Optional[str],
api_base: Optional[str],
llm_provider: Optional[str],
model_group: Optional[str] = None,
model_group: str | None = None,
):
"""
Set the deployment TPM and RPM limits metrics
@ -2743,7 +2743,7 @@ class PrometheusLogger(CustomLogger):
model_id: Optional[str],
api_base: Optional[str],
api_provider: str,
model_group: Optional[str] = None,
model_group: str | None = None,
):
"""
Set the deployment state.
@ -2767,7 +2767,7 @@ class PrometheusLogger(CustomLogger):
model_id: str,
api_base: str,
api_provider: str,
model_group: Optional[str] = None,
model_group: str | None = None,
):
self.set_litellm_deployment_state(
0, litellm_model_name, model_id, api_base, api_provider, model_group
@ -2779,7 +2779,7 @@ class PrometheusLogger(CustomLogger):
model_id: Optional[str],
api_base: Optional[str],
api_provider: str,
model_group: Optional[str] = None,
model_group: str | None = None,
):
self.set_litellm_deployment_state(
1, litellm_model_name, model_id, api_base, api_provider, model_group
@ -2791,7 +2791,7 @@ class PrometheusLogger(CustomLogger):
model_id: Optional[str],
api_base: Optional[str],
api_provider: str,
model_group: Optional[str] = None,
model_group: str | None = None,
):
self.set_litellm_deployment_state(
2, litellm_model_name, model_id, api_base, api_provider, model_group
@ -2804,7 +2804,7 @@ class PrometheusLogger(CustomLogger):
api_base: str,
api_provider: str,
exception_status: str,
model_group: Optional[str] = None,
model_group: str | None = None,
):
"""
increment metric when litellm.Router / load balancing logic places a deployment in cool down

View file

@ -5,7 +5,13 @@ from typing import TYPE_CHECKING
from litellm import logging_callback_manager
from litellm.types.guardrails import SupportedGuardrailIntegrations
from .bias_hallucination_estimator import BiasHallucinationEstimatorGuardrail
from .bias_hallucination_estimator import (
BiasHallucinationEstimatorGuardrail,
GuardrailBehaviorConfig,
GuardrailConfig,
GuardrailSessionConfig,
)
from .risk_scorer import RiskThresholds, RiskWeights
if TYPE_CHECKING:
from litellm.types.guardrails import Guardrail, LitellmParams
@ -20,21 +26,34 @@ def initialize_guardrail(
guardrail_name=guardrail["guardrail_name"],
guardrail_id=guardrail_id,
event_hook=litellm_params.mode,
default_on=litellm_params.default_on or False,
bias_threshold=getattr(litellm_params, "bias_threshold", 0.5),
hallucination_threshold=getattr(litellm_params, "hallucination_threshold", 0.5),
risk_flag_threshold=getattr(litellm_params, "risk_flag_threshold", 0.25),
risk_block_threshold=getattr(litellm_params, "risk_block_threshold", 0.5),
block_on_high_risk=getattr(litellm_params, "block_on_high_risk", True),
log_only=getattr(litellm_params, "log_only", False),
check_request=getattr(litellm_params, "check_request", False),
check_response=getattr(litellm_params, "check_response", True),
violation_message=getattr(litellm_params, "violation_message", None),
violation_message_template=getattr(
litellm_params, "violation_message_template", None
config=GuardrailConfig(
thresholds=RiskThresholds(
bias_threshold=getattr(litellm_params, "bias_threshold", 0.5),
hallucination_threshold=getattr(
litellm_params, "hallucination_threshold", 0.5
),
flag_threshold=getattr(litellm_params, "risk_flag_threshold", 0.25),
block_threshold=getattr(litellm_params, "risk_block_threshold", 0.5),
),
weights=RiskWeights(
bias_weight=getattr(litellm_params, "bias_weight", 0.4),
hallucination_weight=getattr(
litellm_params, "hallucination_weight", 0.6
),
),
behavior=GuardrailBehaviorConfig(
block_on_high_risk=getattr(litellm_params, "block_on_high_risk", True),
log_only=getattr(litellm_params, "log_only", False),
check_request=getattr(litellm_params, "check_request", False),
check_response=getattr(litellm_params, "check_response", True),
violation_message=getattr(litellm_params, "violation_message", None),
),
session=GuardrailSessionConfig(
violation_message_template=getattr(
litellm_params, "violation_message_template", None
),
),
),
bias_weight=getattr(litellm_params, "bias_weight", 0.4),
hallucination_weight=getattr(litellm_params, "hallucination_weight", 0.6),
)
logging_callback_manager.add_litellm_callback(
callback

View file

@ -1,6 +1,7 @@
from __future__ import annotations
from dataclasses import dataclass
from dataclasses import field as dataclasses_field
from datetime import datetime, timezone
from typing import (
TYPE_CHECKING,
@ -33,7 +34,7 @@ from litellm.types.utils import (
)
from .estimator_core import BiasDetector, HallucinationDetector
from .risk_scorer import RiskScorer
from .risk_scorer import RiskScorer, RiskThresholds, RiskWeights
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -69,6 +70,46 @@ class TextRiskAnalysis:
risk: RiskScore
@dataclass
class GuardrailBehaviorConfig:
"""Controls what the guardrail checks and how it reacts to violations."""
block_on_high_risk: bool = True
log_only: bool = False
check_request: bool = False
check_response: bool = True
violation_message: str | None = None
@dataclass
class GuardrailSessionConfig:
"""Session-level routing and messaging settings."""
mask_request_content: bool = False
mask_response_content: bool = False
violation_message_template: str | None = None
end_session_after_n_fails: int | None = None
on_violation: str | None = None
realtime_violation_message: str | None = None
on_sensitive_data: str | None = None
sensitive_data_route_to_model: str | None = None
sticky_session_routing: bool = True
@dataclass
class GuardrailConfig:
"""Top-level configuration bundle for the bias/hallucination guardrail."""
thresholds: RiskThresholds = dataclasses_field(default_factory=RiskThresholds)
weights: RiskWeights = dataclasses_field(default_factory=RiskWeights)
behavior: GuardrailBehaviorConfig = dataclasses_field(
default_factory=GuardrailBehaviorConfig
)
session: GuardrailSessionConfig = dataclasses_field(
default_factory=GuardrailSessionConfig
)
class BiasHallucinationEstimatorGuardrail(CustomGuardrail):
def __init__(
self,
@ -76,28 +117,13 @@ class BiasHallucinationEstimatorGuardrail(CustomGuardrail):
guardrail_name: str | None = None,
guardrail_id: str | None = None,
event_hook: GuardrailEventHookInput | None = None,
default_on: bool = False,
bias_threshold: float = 0.5,
hallucination_threshold: float = 0.5,
risk_flag_threshold: float = 0.25,
risk_block_threshold: float = 0.5,
block_on_high_risk: bool = True,
log_only: bool = False,
check_request: bool = False,
check_response: bool = True,
violation_message: str | None = None,
bias_weight: float = 0.4,
hallucination_weight: float = 0.6,
mask_request_content: bool = False,
mask_response_content: bool = False,
violation_message_template: str | None = None,
end_session_after_n_fails: int | None = None,
on_violation: str | None = None,
realtime_violation_message: str | None = None,
on_sensitive_data: str | None = None,
sensitive_data_route_to_model: str | None = None,
sticky_session_routing: bool = True,
config: GuardrailConfig | None = None,
) -> None:
_config = config or GuardrailConfig()
_thresholds = _config.thresholds
_weights = _config.weights
_behavior = _config.behavior
_session = _config.session
super().__init__( # pyright: ignore[reportUnknownMemberType]
guardrail_name=guardrail_name,
supported_event_hooks=[
@ -106,38 +132,30 @@ class BiasHallucinationEstimatorGuardrail(CustomGuardrail):
],
event_hook=self._normalize_event_hook(event_hook)
or GuardrailEventHooks.post_call,
default_on=default_on,
mask_request_content=mask_request_content,
mask_response_content=mask_response_content,
violation_message_template=violation_message_template,
end_session_after_n_fails=end_session_after_n_fails,
on_violation=on_violation,
realtime_violation_message=realtime_violation_message,
on_sensitive_data=on_sensitive_data,
sensitive_data_route_to_model=sensitive_data_route_to_model,
sticky_session_routing=sticky_session_routing,
mask_request_content=_session.mask_request_content,
mask_response_content=_session.mask_response_content,
violation_message_template=_session.violation_message_template,
end_session_after_n_fails=_session.end_session_after_n_fails,
on_violation=_session.on_violation,
realtime_violation_message=_session.realtime_violation_message,
on_sensitive_data=_session.on_sensitive_data,
sensitive_data_route_to_model=_session.sensitive_data_route_to_model,
sticky_session_routing=_session.sticky_session_routing,
)
self.guardrail_provider = GUARDRAIL_PROVIDER
self.guardrail_id = guardrail_id
self.bias_threshold = bias_threshold
self.hallucination_threshold = hallucination_threshold
self.risk_flag_threshold = risk_flag_threshold
self.risk_block_threshold = risk_block_threshold
self.block_on_high_risk = block_on_high_risk
self.log_only = log_only
self.check_request = check_request
self.check_response = check_response
self.violation_message = violation_message
self.bias_threshold = _thresholds.bias_threshold
self.hallucination_threshold = _thresholds.hallucination_threshold
self.risk_flag_threshold = _thresholds.flag_threshold
self.risk_block_threshold = _thresholds.block_threshold
self.block_on_high_risk = _behavior.block_on_high_risk
self.log_only = _behavior.log_only
self.check_request = _behavior.check_request
self.check_response = _behavior.check_response
self.violation_message = _behavior.violation_message
self.bias_detector = BiasDetector()
self.hallucination_detector = HallucinationDetector()
self.risk_scorer = RiskScorer(
bias_weight=bias_weight,
hallucination_weight=hallucination_weight,
bias_threshold=bias_threshold,
hallucination_threshold=hallucination_threshold,
flag_threshold=risk_flag_threshold,
block_threshold=risk_block_threshold,
)
self.risk_scorer = RiskScorer(thresholds=_thresholds, weights=_weights)
@log_guardrail_information
async def apply_guardrail(

View file

@ -4,10 +4,28 @@ import asyncio
import json
import re
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from importlib.util import find_spec as _find_spec
from pathlib import Path
from typing import Any, cast
@dataclass
class FetchConfig:
"""HTTP fetch settings for URL-backed data sources."""
cache_ttl: int = 3600
timeout: float = 5.0
@dataclass
class VectorStoreClientConfig:
"""Optional pre-built client and embedding model for vector store sources."""
client: object | None = None
embedding_model: object | None = None
_AIOHTTP_AVAILABLE: bool = _find_spec("aiohttp") is not None
@ -147,13 +165,13 @@ class URLDataSource(DataSource):
name: str = "url_source",
enabled: bool = True,
priority: int = 0,
cache_ttl: int = 3600,
timeout: float = 5.0,
fetch_config: FetchConfig | None = None,
) -> None:
super().__init__(name=name, enabled=enabled, priority=priority)
self.urls = urls
self.cache_ttl = cache_ttl
self.timeout = timeout
_fetch = fetch_config or FetchConfig()
self.cache_ttl = _fetch.cache_ttl
self.timeout = _fetch.timeout
self._documents: list[str | dict[str, Any]] = []
self._index: dict[str, list[int]] = {}
self._fetched = False
@ -226,8 +244,7 @@ class VectorStoreDataSource(DataSource):
name: str = "",
enabled: bool = True,
priority: int = 0,
client: object | None = None,
embedding_model: object | None = None,
client_config: VectorStoreClientConfig | None = None,
**config: str,
) -> None:
super().__init__(
@ -235,12 +252,15 @@ class VectorStoreDataSource(DataSource):
)
self.provider = provider
self.config = config
_client_config = client_config or VectorStoreClientConfig()
self.client: object | None = (
client if client is not None else self._initialize_client(provider, config)
_client_config.client
if _client_config.client is not None
else self._initialize_client(provider, config)
)
self.embedding_model: object | None = (
embedding_model
if embedding_model is not None
_client_config.embedding_model
if _client_config.embedding_model is not None
else self._load_embedding_model()
)

View file

@ -1,6 +1,5 @@
from __future__ import annotations
from litellm.types.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator import (
BiasAnalysis,
HallucinationAnalysis,

View file

@ -1,5 +1,6 @@
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Literal
from litellm.types.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator import (
@ -10,23 +11,39 @@ from litellm.types.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator
)
@dataclass
class RiskThresholds:
"""Thresholds used to classify risk levels."""
bias_threshold: float = 0.5
hallucination_threshold: float = 0.5
flag_threshold: float = 0.25
block_threshold: float = 0.5
@dataclass
class RiskWeights:
"""Weights applied when computing the combined risk score."""
bias_weight: float = 0.4
hallucination_weight: float = 0.6
class RiskScorer:
def __init__(
self,
*,
bias_weight: float = 0.4,
hallucination_weight: float = 0.6,
bias_threshold: float = 0.5,
hallucination_threshold: float = 0.5,
flag_threshold: float = 0.25,
block_threshold: float = 0.5,
thresholds: RiskThresholds | None = None,
weights: RiskWeights | None = None,
) -> None:
self.bias_weight = bias_weight
self.hallucination_weight = hallucination_weight
self.bias_threshold = bias_threshold
self.hallucination_threshold = hallucination_threshold
self.flag_threshold = flag_threshold
self.block_threshold = block_threshold
_thresholds = thresholds or RiskThresholds()
_weights = weights or RiskWeights()
self.bias_weight = _weights.bias_weight
self.hallucination_weight = _weights.hallucination_weight
self.bias_threshold = _thresholds.bias_threshold
self.hallucination_threshold = _thresholds.hallucination_threshold
self.flag_threshold = _thresholds.flag_threshold
self.block_threshold = _thresholds.block_threshold
def compute_risk(
self,

View file

@ -179,7 +179,7 @@ def _get_config_loaded_guardrails() -> list[Any]:
def _find_config_loaded_guardrail(
guardrail_id_or_name: str,
) -> Optional[object]:
) -> object | None:
for guardrail in _get_config_loaded_guardrails():
gid, display_name = _get_guardrail_attrs(guardrail)
if guardrail_id_or_name in (gid, display_name):
@ -479,7 +479,7 @@ def _build_usage_logs_where(
def _usage_log_entry_from_row(
r: object, sl: object, action_filter: Optional[str]
r: object, sl: object, action_filter: str | None
) -> Optional[UsageLogEntry]:
meta = sl.metadata
if isinstance(meta, str):
@ -530,7 +530,7 @@ def _usage_log_entry_from_row(
)
def _snippet(text: object, max_len: int = 200) -> Optional[str]:
def _snippet(text: object, max_len: int = 200) -> str | None:
if text is None:
return None
if isinstance(text, str):
@ -552,7 +552,7 @@ def _snippet(text: object, max_len: int = 200) -> Optional[str]:
return result
def _input_snippet_for_log(sl: object) -> Optional[str]:
def _input_snippet_for_log(sl: object) -> str | None:
"""Snippet for request input: prefer messages, fall back to proxy_server_request (same as drawer)."""
out = _snippet(sl.messages)
if out:

View file

@ -2192,7 +2192,7 @@ async def cli_poll_key(
async def insert_sso_user(
result_openid: Optional[Union[OpenID, dict]],
user_defined_values: Optional[SSOUserDefinedValues] = None,
prisma_client: Optional[PrismaClient] = None,
prisma_client: PrismaClient | None = None,
) -> NewUserResponse:
"""
Helper function to create a New User in LiteLLM DB after a successful SSO login