mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
41abaa7aad
commit
9fb239d029
9 changed files with 195 additions and 122 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
from __future__ import annotations
|
||||
|
||||
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator import (
|
||||
BiasAnalysis,
|
||||
HallucinationAnalysis,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue