diff --git a/litellm/integrations/asqav/asqav.py b/litellm/integrations/asqav/asqav.py index 1d9aa98e5b7..5bcce14b34a 100644 --- a/litellm/integrations/asqav/asqav.py +++ b/litellm/integrations/asqav/asqav.py @@ -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 diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 76fa8a9b4c2..4668755986e 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/__init__.py index 6237a81d82a..4bdb12566e8 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/__init__.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/bias_hallucination_estimator.py b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/bias_hallucination_estimator.py index bc28199b8ee..15467d3476b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/bias_hallucination_estimator.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/bias_hallucination_estimator.py @@ -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( diff --git a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/data_sources.py b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/data_sources.py index 73ea2135529..11b2c81e995 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/data_sources.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/data_sources.py @@ -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() ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/estimator_core.py b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/estimator_core.py index a7f96c03a15..86074d6249a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/estimator_core.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/estimator_core.py @@ -1,6 +1,5 @@ from __future__ import annotations - from litellm.types.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator import ( BiasAnalysis, HallucinationAnalysis, diff --git a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/risk_scorer.py b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/risk_scorer.py index 62eb94656f5..981d1907db6 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/risk_scorer.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/risk_scorer.py @@ -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, diff --git a/litellm/proxy/guardrails/usage_endpoints.py b/litellm/proxy/guardrails/usage_endpoints.py index 853ff07ad25..66b335c37cf 100644 --- a/litellm/proxy/guardrails/usage_endpoints.py +++ b/litellm/proxy/guardrails/usage_endpoints.py @@ -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: diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index cabb5c398ee..aa70ef99f1a 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -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