mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
feat(guardrails): add native bias and hallucination estimator guardrail (#30931)
* Implement Bias and Hallucination Estimator with Grounding Checker, Risk Scorer, and Utility Functions - Added GroundingChecker for verifying claims against data sources. - Introduced RiskScorer to compute risk scores based on bias and hallucination analyses. - Developed utility functions for sentence splitting, text clipping, and unique value preservation. - Created patterns for detecting bias and hallucination indicators. - Established data models for bias and hallucination analysis results. - Implemented tests for bias detection, hallucination detection, grounding checks, and risk scoring. - Integrated the BiasHallucinationEstimatorGuardrail for managing high-risk responses. * Refactor Bias Hallucination Estimator: Enhance logging, remove unused parameters, and improve concurrency handling - Added logging decorator to `apply_guardrail` method to log guardrail information while excluding sensitive fields. - Removed `use_logprobs` and `uncertainty_weight` parameters from `BiasHallucinationEstimatorGuardrail` and related classes. - Simplified `BiasDetector` and `HallucinationDetector` initialization by removing threshold parameters. - Implemented a lock mechanism in `URLDataSource` to prevent concurrent fetches from causing race conditions. - Updated `RiskScorer` to remove uncertainty handling and adjusted risk calculation logic. - Enhanced test coverage for guardrail logging and data source functionalities, ensuring proper behavior under various conditions. * Refactor bias hallucination estimator code for improved readability and consistency - Updated string formatting for better readability in data_sources.py, estimator_core.py, grounding_checker.py, patterns.py, risk_scorer.py, and utils.py. - Enhanced the clarity of function signatures and method calls across various classes. - Removed unnecessary variables and streamlined logic in grounding_checker.py. - Improved test cases for bias and hallucination detection to enhance coverage and maintainability. - Added tests for initializing guardrails and handling edge cases in grounding checks. * fix(bias-hallucination-estimator): 100% patch coverage, ruff strict gate passing Fix ruff strict gate violations: replace deprecated typing imports (List/Dict/Tuple/Optional) with builtin generics and union syntax (UP006/UP037/UP045); replace Any in public API signatures with object or str (ANN401); remove now-unused imports (F401). Grow test suite from 117 to 130 tests covering all previously uncovered branches: DataSource.verify_fact, _keyword_search empty-word-chars path, URLDataSource._fetch_url exception path, VectorStoreDataSource _initialize_client pinecone/weaviate paths and ImportError fallback, _load_embedding_model via mocked sentence_transformers, VectorStore search exception path, KnowledgeGraph search exception path, and GroundingChecker._boost_confidence entity match branch. Mark the abstract method stub with pragma: no cover. All 8 files now at 100% patch coverage; strict gate clean.
This commit is contained in:
parent
84eac5de06
commit
1d15734fc8
11 changed files with 3316 additions and 1 deletions
|
|
@ -0,0 +1,52 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from litellm import logging_callback_manager
|
||||
from litellm.types.guardrails import SupportedGuardrailIntegrations
|
||||
|
||||
from .bias_hallucination_estimator import BiasHallucinationEstimatorGuardrail
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.guardrails import Guardrail, LitellmParams
|
||||
|
||||
|
||||
def initialize_guardrail(
|
||||
litellm_params: LitellmParams,
|
||||
guardrail: Guardrail,
|
||||
) -> BiasHallucinationEstimatorGuardrail:
|
||||
guardrail_id = guardrail["guardrail_id"] if "guardrail_id" in guardrail else None
|
||||
callback = BiasHallucinationEstimatorGuardrail(
|
||||
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
|
||||
),
|
||||
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
|
||||
) # pyright: ignore[reportUnknownMemberType]
|
||||
return callback
|
||||
|
||||
|
||||
guardrail_initializer_registry = {
|
||||
SupportedGuardrailIntegrations.BIAS_HALLUCINATION_ESTIMATOR.value: initialize_guardrail,
|
||||
}
|
||||
|
||||
|
||||
guardrail_class_registry = {
|
||||
SupportedGuardrailIntegrations.BIAS_HALLUCINATION_ESTIMATOR.value: BiasHallucinationEstimatorGuardrail,
|
||||
}
|
||||
|
|
@ -0,0 +1,398 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Literal,
|
||||
Mapping,
|
||||
Protocol,
|
||||
Sequence,
|
||||
Union,
|
||||
cast,
|
||||
)
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm.exceptions import GuardrailRaisedException
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
log_guardrail_information,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks, Mode
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator import (
|
||||
BiasAnalysis,
|
||||
HallucinationAnalysis,
|
||||
RiskScore,
|
||||
)
|
||||
from litellm.types.utils import (
|
||||
GenericGuardrailAPIInputs,
|
||||
GuardrailStatus,
|
||||
GuardrailTracingDetail,
|
||||
)
|
||||
|
||||
from .estimator_core import BiasDetector, HallucinationDetector
|
||||
from .risk_scorer import RiskScorer
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
GUARDRAIL_PROVIDER = "litellm_native"
|
||||
_LOG_EXCLUDED_FIELDS: frozenset[str] = frozenset(
|
||||
{"examples", "unsourced_claims", "missing_citations", "fabricated_specificity"}
|
||||
)
|
||||
GuardrailEventHookInput = Union[
|
||||
str,
|
||||
Sequence[str],
|
||||
Mode,
|
||||
]
|
||||
NormalizedGuardrailEventHook = Union[
|
||||
GuardrailEventHooks, list[GuardrailEventHooks], Mode
|
||||
]
|
||||
|
||||
|
||||
class FunctionLike(Protocol):
|
||||
name: str | None
|
||||
arguments: str
|
||||
|
||||
|
||||
class ToolCallLike(Protocol):
|
||||
function: FunctionLike
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TextRiskAnalysis:
|
||||
text: str
|
||||
bias: BiasAnalysis
|
||||
hallucination: HallucinationAnalysis
|
||||
risk: RiskScore
|
||||
|
||||
|
||||
class BiasHallucinationEstimatorGuardrail(CustomGuardrail):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
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,
|
||||
) -> None:
|
||||
super().__init__( # pyright: ignore[reportUnknownMemberType]
|
||||
guardrail_name=guardrail_name,
|
||||
supported_event_hooks=[
|
||||
GuardrailEventHooks.pre_call,
|
||||
GuardrailEventHooks.post_call,
|
||||
],
|
||||
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,
|
||||
)
|
||||
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_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,
|
||||
)
|
||||
|
||||
@log_guardrail_information
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict[str, object],
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: LiteLLMLoggingObj | None = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
if not self._should_check(input_type):
|
||||
return inputs
|
||||
|
||||
start_time = datetime.now()
|
||||
texts = self._extract_texts(inputs)
|
||||
if not texts:
|
||||
return inputs
|
||||
|
||||
analyses = tuple(self._analyze_text(text=text) for text in texts)
|
||||
highest_risk = max(
|
||||
analyses, key=lambda analysis: analysis.risk.overall_risk_percentage
|
||||
)
|
||||
decision = self._decision(highest_risk.risk.recommendation)
|
||||
status: GuardrailStatus = (
|
||||
"guardrail_intervened" if decision == "blocked" else "success"
|
||||
)
|
||||
response_payload = self._build_response_payload(
|
||||
analyses=analyses,
|
||||
decision=decision,
|
||||
input_type=input_type,
|
||||
)
|
||||
|
||||
self._log_guardrail_result(
|
||||
request_data=request_data,
|
||||
response_payload=response_payload,
|
||||
status=status,
|
||||
start_time=start_time,
|
||||
highest_risk_percentage=highest_risk.risk.overall_risk_percentage,
|
||||
)
|
||||
|
||||
if decision == "blocked":
|
||||
raise GuardrailRaisedException(
|
||||
guardrail_name=self.guardrail_name,
|
||||
message=self._build_violation_message(highest_risk.risk),
|
||||
should_wrap_with_default_message=False,
|
||||
)
|
||||
|
||||
return inputs
|
||||
|
||||
def estimate_bias_hallucination(self, text: str) -> dict[str, object]:
|
||||
analysis = self._analyze_text(text=text)
|
||||
return {
|
||||
"bias": analysis.bias.model_dump(mode="json"),
|
||||
"hallucination": analysis.hallucination.model_dump(mode="json"),
|
||||
"risk": analysis.risk.model_dump(mode="json"),
|
||||
}
|
||||
|
||||
def _should_check(self, input_type: Literal["request", "response"]) -> bool:
|
||||
if input_type == "request":
|
||||
return self.check_request
|
||||
return self.check_response
|
||||
|
||||
def _analyze_text(self, *, text: str) -> TextRiskAnalysis:
|
||||
bias_analysis = self.bias_detector.detect(text)
|
||||
hallucination_analysis = self.hallucination_detector.detect(text)
|
||||
risk = self.risk_scorer.compute_risk(
|
||||
bias_analysis=bias_analysis,
|
||||
hallucination_analysis=hallucination_analysis,
|
||||
)
|
||||
return TextRiskAnalysis(
|
||||
text=text,
|
||||
bias=bias_analysis,
|
||||
hallucination=hallucination_analysis,
|
||||
risk=risk,
|
||||
)
|
||||
|
||||
def _decision(self, recommendation: Literal["pass", "flag", "block"]) -> str:
|
||||
if recommendation == "block" and self.block_on_high_risk and not self.log_only:
|
||||
return "blocked"
|
||||
if recommendation == "pass":
|
||||
return "passed"
|
||||
return "flagged"
|
||||
|
||||
def _build_response_payload(
|
||||
self,
|
||||
*,
|
||||
analyses: tuple[TextRiskAnalysis, ...],
|
||||
decision: str,
|
||||
input_type: Literal["request", "response"],
|
||||
) -> dict[str, object]:
|
||||
return {
|
||||
"decision": decision,
|
||||
"input_type": input_type,
|
||||
"risk_scores": [
|
||||
analysis.risk.model_dump(mode="json") for analysis in analyses
|
||||
],
|
||||
"bias": [
|
||||
{
|
||||
k: v
|
||||
for k, v in analysis.bias.model_dump(mode="json").items()
|
||||
if k not in _LOG_EXCLUDED_FIELDS
|
||||
}
|
||||
for analysis in analyses
|
||||
],
|
||||
"hallucination": [
|
||||
{
|
||||
k: v
|
||||
for k, v in analysis.hallucination.model_dump(mode="json").items()
|
||||
if k not in _LOG_EXCLUDED_FIELDS
|
||||
}
|
||||
for analysis in analyses
|
||||
],
|
||||
}
|
||||
|
||||
def _log_guardrail_result(
|
||||
self,
|
||||
*,
|
||||
request_data: dict[str, object],
|
||||
response_payload: dict[str, object],
|
||||
status: GuardrailStatus,
|
||||
start_time: datetime,
|
||||
highest_risk_percentage: int,
|
||||
) -> None:
|
||||
detection_methods = self._detection_methods(response_payload)
|
||||
tracing_detail = GuardrailTracingDetail(
|
||||
guardrail_id=self.guardrail_id or self.guardrail_name,
|
||||
detection_method=detection_methods,
|
||||
risk_score=highest_risk_percentage / 100,
|
||||
violation_categories=self._violation_categories(response_payload),
|
||||
guardrail_action=str(response_payload["decision"]),
|
||||
)
|
||||
self.add_standard_logging_guardrail_information_to_request_data( # pyright: ignore[reportUnknownMemberType]
|
||||
guardrail_provider=self.guardrail_provider,
|
||||
guardrail_json_response=response_payload,
|
||||
request_data=request_data,
|
||||
guardrail_status=status,
|
||||
start_time=start_time.timestamp(),
|
||||
end_time=datetime.now().timestamp(),
|
||||
duration=(datetime.now() - start_time).total_seconds(),
|
||||
tracing_detail=tracing_detail,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _detection_methods(response_payload: dict[str, object]) -> str | None:
|
||||
categories = BiasHallucinationEstimatorGuardrail._violation_categories(
|
||||
response_payload
|
||||
)
|
||||
if not categories:
|
||||
return None
|
||||
return "regex,keyword"
|
||||
|
||||
@staticmethod
|
||||
def _violation_categories(response_payload: dict[str, object]) -> list[str]:
|
||||
risk_scores = response_payload.get("risk_scores")
|
||||
if not isinstance(risk_scores, list):
|
||||
return []
|
||||
typed_risk_scores = cast(list[object], risk_scores)
|
||||
return list(
|
||||
dict.fromkeys(
|
||||
issue.split(":", 1)[0]
|
||||
for risk_score in typed_risk_scores
|
||||
for issue in BiasHallucinationEstimatorGuardrail._detected_issues(
|
||||
risk_score
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _detected_issues(risk_score: object) -> tuple[str, ...]:
|
||||
if not isinstance(risk_score, dict):
|
||||
return ()
|
||||
typed_risk_score = cast(dict[str, object], risk_score)
|
||||
detected_issues = typed_risk_score.get("detected_issues")
|
||||
if not isinstance(detected_issues, list):
|
||||
return ()
|
||||
typed_issues = cast(list[object], detected_issues)
|
||||
return tuple(issue for issue in typed_issues if isinstance(issue, str))
|
||||
|
||||
@staticmethod
|
||||
def _normalize_event_hook(
|
||||
event_hook: GuardrailEventHookInput | None,
|
||||
) -> NormalizedGuardrailEventHook | None:
|
||||
if event_hook is None:
|
||||
return None
|
||||
if isinstance(event_hook, Mode):
|
||||
return event_hook
|
||||
if isinstance(event_hook, str):
|
||||
return GuardrailEventHooks(event_hook)
|
||||
return [GuardrailEventHooks(hook) for hook in event_hook]
|
||||
|
||||
def _build_violation_message(self, risk_score: RiskScore) -> str:
|
||||
default_message = f"High bias/hallucination risk detected ({risk_score.overall_risk_percentage}%)."
|
||||
if self.violation_message:
|
||||
return self.violation_message
|
||||
return self.render_violation_message(
|
||||
default=default_message,
|
||||
context={
|
||||
"risk_score": risk_score.overall_risk_percentage,
|
||||
"detected_issues": ", ".join(risk_score.detected_issues),
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _extract_texts(inputs: GenericGuardrailAPIInputs) -> tuple[str, ...]:
|
||||
texts = tuple(text for text in inputs.get("texts", []) if text)
|
||||
tool_call_texts = tuple(
|
||||
text
|
||||
for tool_call in inputs.get("tool_calls", [])
|
||||
for text in (
|
||||
BiasHallucinationEstimatorGuardrail._tool_call_text(tool_call),
|
||||
)
|
||||
if text
|
||||
)
|
||||
return texts + tool_call_texts
|
||||
|
||||
@staticmethod
|
||||
def _tool_call_text(tool_call: object) -> str | None:
|
||||
if isinstance(tool_call, dict):
|
||||
tool_call_dict = cast(Mapping[str, object], tool_call)
|
||||
function = tool_call_dict.get("function")
|
||||
if isinstance(function, dict):
|
||||
function_dict = cast(Mapping[str, object], function)
|
||||
name = function_dict.get("name")
|
||||
arguments = function_dict.get("arguments")
|
||||
return (
|
||||
f"{BiasHallucinationEstimatorGuardrail._string_value(name)} "
|
||||
f"{BiasHallucinationEstimatorGuardrail._string_value(arguments)}"
|
||||
).strip() or None
|
||||
return None
|
||||
|
||||
if not hasattr(tool_call, "function"):
|
||||
return None
|
||||
tool_call_like = cast(ToolCallLike, tool_call)
|
||||
function = tool_call_like.function
|
||||
name = function.name or ""
|
||||
arguments = function.arguments
|
||||
return f"{name} {arguments}".strip() or None
|
||||
|
||||
@staticmethod
|
||||
def _string_value(value: object) -> str:
|
||||
if value is None:
|
||||
return ""
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
return str(value)
|
||||
|
||||
@staticmethod
|
||||
def get_config_model() -> type[GuardrailConfigModel[BaseModel]] | None:
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator import (
|
||||
BiasHallucinationEstimatorConfigModel,
|
||||
)
|
||||
|
||||
return cast(
|
||||
type[GuardrailConfigModel[BaseModel]],
|
||||
BiasHallucinationEstimatorConfigModel,
|
||||
)
|
||||
|
||||
|
||||
BiasHallucinationEstimator = BiasHallucinationEstimatorGuardrail
|
||||
|
|
@ -0,0 +1,402 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import re
|
||||
from abc import ABC, abstractmethod
|
||||
from importlib.util import find_spec as _find_spec
|
||||
from pathlib import Path
|
||||
from typing import Any, cast
|
||||
|
||||
_AIOHTTP_AVAILABLE: bool = _find_spec("aiohttp") is not None
|
||||
|
||||
|
||||
class DataSourceResult:
|
||||
def __init__(
|
||||
self,
|
||||
text: str,
|
||||
source: str,
|
||||
confidence: float = 0.8,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
self.text = text
|
||||
self.source = source
|
||||
self.confidence = min(max(confidence, 0.0), 1.0)
|
||||
self.metadata = metadata or {}
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return (
|
||||
f"DataSourceResult(source={self.source}, confidence={self.confidence:.2f})"
|
||||
)
|
||||
|
||||
|
||||
class DataSource(ABC):
|
||||
def __init__(self, name: str = "", enabled: bool = True, priority: int = 0) -> None:
|
||||
self.name = name
|
||||
self.enabled = enabled
|
||||
self.priority = priority
|
||||
|
||||
@abstractmethod
|
||||
async def search(
|
||||
self, query: str, limit: int = 5
|
||||
) -> list[DataSourceResult]: ... # pragma: no cover
|
||||
|
||||
async def verify_fact(self, claim: str) -> tuple[bool, str | None]:
|
||||
results = await self.search(claim, limit=1)
|
||||
if results:
|
||||
return True, results[0].text
|
||||
return False, None
|
||||
|
||||
|
||||
def _build_keyword_index(documents: list[str | dict[str, Any]]) -> dict[str, list[int]]:
|
||||
index: dict[str, list[int]] = {}
|
||||
for idx, doc in enumerate(documents):
|
||||
text = _get_doc_text(doc).lower()
|
||||
for word in set(re.findall(r"\b\w+\b", text)):
|
||||
index.setdefault(word, []).append(idx)
|
||||
return index
|
||||
|
||||
|
||||
def _get_doc_text(doc: str | dict[str, Any]) -> str:
|
||||
if isinstance(doc, str):
|
||||
return doc
|
||||
return doc.get("text", str(doc))
|
||||
|
||||
|
||||
def _keyword_search(
|
||||
query: str,
|
||||
documents: list[str | dict[str, Any]],
|
||||
index: dict[str, list[int]],
|
||||
source_name: str,
|
||||
limit: int,
|
||||
) -> list[DataSourceResult]:
|
||||
if not documents or not query:
|
||||
return []
|
||||
query_words = set(re.findall(r"\b\w+\b", query.lower()))
|
||||
if not query_words:
|
||||
return []
|
||||
doc_indices: set[int] = set()
|
||||
for word in query_words:
|
||||
if word in index:
|
||||
doc_indices.update(index[word])
|
||||
scored: list[tuple[float, DataSourceResult]] = []
|
||||
for idx in doc_indices:
|
||||
doc = documents[idx]
|
||||
text = _get_doc_text(doc)
|
||||
matching = len(set(re.findall(r"\b\w+\b", text.lower())) & query_words)
|
||||
score = matching / len(query_words)
|
||||
scored.append(
|
||||
(score, DataSourceResult(text=text, source=source_name, confidence=score))
|
||||
)
|
||||
scored.sort(reverse=True, key=lambda x: x[0])
|
||||
return [result for _, result in scored[:limit]]
|
||||
|
||||
|
||||
def _parse_json_docs(raw: object) -> list[dict[str, Any]]:
|
||||
if isinstance(raw, list):
|
||||
return cast(list[dict[str, Any]], raw)
|
||||
return [cast(dict[str, Any], raw)]
|
||||
|
||||
|
||||
class FileDataSource(DataSource):
|
||||
def __init__(
|
||||
self,
|
||||
file_path: str,
|
||||
name: str = "",
|
||||
enabled: bool = True,
|
||||
priority: int = 0,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
name=name or f"file_{Path(file_path).stem}",
|
||||
enabled=enabled,
|
||||
priority=priority,
|
||||
)
|
||||
self.file_path = file_path
|
||||
self._documents: list[str | dict[str, Any]] = self._load_documents(file_path)
|
||||
self._index: dict[str, list[int]] = _build_keyword_index(self._documents)
|
||||
|
||||
@staticmethod
|
||||
def _load_documents(file_path: str) -> list[str | dict[str, Any]]:
|
||||
try:
|
||||
path = Path(file_path)
|
||||
if not path.exists():
|
||||
return []
|
||||
if path.suffix == ".json":
|
||||
with open(path) as f:
|
||||
return cast(
|
||||
list[str | dict[str, Any]], _parse_json_docs(json.load(f))
|
||||
)
|
||||
if path.suffix in {".csv", ".txt"}:
|
||||
with open(path) as f:
|
||||
return cast(
|
||||
list[str | dict[str, Any]],
|
||||
[{"text": line.strip()} for line in f if line.strip()],
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
return []
|
||||
|
||||
async def search(self, query: str, limit: int = 5) -> list[DataSourceResult]:
|
||||
return _keyword_search(query, self._documents, self._index, self.name, limit)
|
||||
|
||||
|
||||
class URLDataSource(DataSource):
|
||||
def __init__(
|
||||
self,
|
||||
urls: list[str],
|
||||
name: str = "url_source",
|
||||
enabled: bool = True,
|
||||
priority: int = 0,
|
||||
cache_ttl: int = 3600,
|
||||
timeout: float = 5.0,
|
||||
) -> None:
|
||||
super().__init__(name=name, enabled=enabled, priority=priority)
|
||||
self.urls = urls
|
||||
self.cache_ttl = cache_ttl
|
||||
self.timeout = timeout
|
||||
self._documents: list[str | dict[str, Any]] = []
|
||||
self._index: dict[str, list[int]] = {}
|
||||
self._fetched = False
|
||||
self._fetch_lock = asyncio.Lock()
|
||||
|
||||
async def search(self, query: str, limit: int = 5) -> list[DataSourceResult]:
|
||||
if not self._fetched:
|
||||
async with self._fetch_lock:
|
||||
if not self._fetched:
|
||||
self._documents = await self._fetch_all()
|
||||
self._index = _build_keyword_index(self._documents)
|
||||
self._fetched = True
|
||||
return _keyword_search(query, self._documents, self._index, self.name, limit)
|
||||
|
||||
async def _fetch_all(self) -> list[str | dict[str, Any]]:
|
||||
if not _AIOHTTP_AVAILABLE:
|
||||
return []
|
||||
results = await asyncio.gather(
|
||||
*[self._fetch_url(u) for u in self.urls], return_exceptions=True
|
||||
)
|
||||
return [doc for batch in results if isinstance(batch, list) for doc in batch]
|
||||
|
||||
async def _fetch_url(self, url: str) -> list[str | dict[str, Any]]:
|
||||
try:
|
||||
import aiohttp # type: ignore[import-untyped]
|
||||
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(
|
||||
url, timeout=aiohttp.ClientTimeout(total=self.timeout)
|
||||
) as response:
|
||||
if response.status == 200:
|
||||
content = await response.text()
|
||||
return self._parse_content(content)
|
||||
except Exception:
|
||||
pass
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
def _parse_content(content: str) -> list[str | dict[str, Any]]:
|
||||
try:
|
||||
return cast(
|
||||
list[str | dict[str, Any]], _parse_json_docs(json.loads(content))
|
||||
)
|
||||
except json.JSONDecodeError:
|
||||
return cast(list[str | dict[str, Any]], [{"text": content}])
|
||||
|
||||
|
||||
class ContextDocumentDataSource(DataSource):
|
||||
def __init__(
|
||||
self,
|
||||
documents: list[dict[str, Any]],
|
||||
name: str = "context_documents",
|
||||
enabled: bool = True,
|
||||
priority: int = 100,
|
||||
) -> None:
|
||||
super().__init__(name=name, enabled=enabled, priority=priority)
|
||||
self._documents: list[str | dict[str, Any]] = cast(
|
||||
list[str | dict[str, Any]], documents
|
||||
)
|
||||
self._index = _build_keyword_index(self._documents)
|
||||
|
||||
async def search(self, query: str, limit: int = 5) -> list[DataSourceResult]:
|
||||
return _keyword_search(query, self._documents, self._index, self.name, limit)
|
||||
|
||||
|
||||
class VectorStoreDataSource(DataSource):
|
||||
def __init__(
|
||||
self,
|
||||
provider: str,
|
||||
name: str = "",
|
||||
enabled: bool = True,
|
||||
priority: int = 0,
|
||||
client: object | None = None,
|
||||
embedding_model: object | None = None,
|
||||
**config: str,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
name=name or f"vectorstore_{provider}", enabled=enabled, priority=priority
|
||||
)
|
||||
self.provider = provider
|
||||
self.config = config
|
||||
self.client: object | None = (
|
||||
client if client is not None else self._initialize_client(provider, config)
|
||||
)
|
||||
self.embedding_model: object | None = (
|
||||
embedding_model
|
||||
if embedding_model is not None
|
||||
else self._load_embedding_model()
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _initialize_client(provider: str, config: dict[str, str]) -> object | None:
|
||||
if provider == "pinecone":
|
||||
try:
|
||||
import pinecone # type: ignore[import-untyped]
|
||||
|
||||
api_key = config.get("api_key")
|
||||
index_name = config.get("index_name")
|
||||
if api_key and index_name:
|
||||
return pinecone.Index(index_name) # type: ignore[no-untyped-call]
|
||||
except ImportError:
|
||||
pass
|
||||
elif provider == "weaviate":
|
||||
try:
|
||||
import weaviate # type: ignore[import-untyped]
|
||||
|
||||
url = config.get("url")
|
||||
if url:
|
||||
return weaviate.Client(url) # type: ignore[no-untyped-call]
|
||||
except ImportError:
|
||||
pass
|
||||
return None
|
||||
|
||||
async def search(self, query: str, limit: int = 5) -> list[DataSourceResult]:
|
||||
if not self.client:
|
||||
return []
|
||||
try:
|
||||
embedding = await self._get_embedding(query)
|
||||
if not embedding:
|
||||
return []
|
||||
if self.provider == "pinecone":
|
||||
results = self.client.query(
|
||||
embedding, top_k=limit, include_metadata=True
|
||||
)
|
||||
return [
|
||||
DataSourceResult(
|
||||
text=match.get("metadata", {}).get("text", ""),
|
||||
source=self.name,
|
||||
confidence=float(match.get("score", 0.5)),
|
||||
)
|
||||
for match in results.get("matches", [])
|
||||
]
|
||||
if self.provider == "weaviate":
|
||||
response = (
|
||||
self.client.query.get(self.config.get("collection", "Document"))
|
||||
.with_near_vector({"vector": embedding})
|
||||
.with_limit(limit)
|
||||
.do()
|
||||
)
|
||||
docs = response.get("data", {}).get("Get", {})
|
||||
return [
|
||||
DataSourceResult(
|
||||
text=doc.get("text", ""), source=self.name, confidence=0.8
|
||||
)
|
||||
for doc in docs
|
||||
]
|
||||
except Exception:
|
||||
pass
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
def _load_embedding_model() -> object | None:
|
||||
try:
|
||||
from sentence_transformers import SentenceTransformer # type: ignore[import-untyped]
|
||||
|
||||
return SentenceTransformer("all-MiniLM-L6-v2")
|
||||
except ImportError:
|
||||
return None
|
||||
|
||||
async def _get_embedding(self, text: str) -> list[float] | None:
|
||||
if self.embedding_model is None:
|
||||
return None
|
||||
embedding = self.embedding_model.encode(text, convert_to_tensor=False)
|
||||
return embedding.tolist() if hasattr(embedding, "tolist") else list(embedding)
|
||||
|
||||
|
||||
class KnowledgeGraphDataSource(DataSource):
|
||||
def __init__(
|
||||
self,
|
||||
endpoint: str = "https://query.wikidata.org/sparql",
|
||||
name: str = "knowledge_graph",
|
||||
enabled: bool = True,
|
||||
priority: int = 0,
|
||||
) -> None:
|
||||
super().__init__(name=name, enabled=enabled, priority=priority)
|
||||
self.endpoint = endpoint
|
||||
|
||||
async def search(self, query: str, limit: int = 5) -> list[DataSourceResult]:
|
||||
if not _AIOHTTP_AVAILABLE:
|
||||
return []
|
||||
sparql_query = self._build_sparql_query(query)
|
||||
try:
|
||||
import aiohttp # type: ignore[import-untyped]
|
||||
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(
|
||||
self.endpoint,
|
||||
params={"query": sparql_query, "format": "json"},
|
||||
timeout=aiohttp.ClientTimeout(total=5),
|
||||
) as response:
|
||||
if response.status == 200:
|
||||
data = await response.json()
|
||||
return [
|
||||
DataSourceResult(
|
||||
text=text, source=self.name, confidence=0.9
|
||||
)
|
||||
for binding in data.get("results", {}).get("bindings", [])[
|
||||
:limit
|
||||
]
|
||||
for text in (self._extract_text(binding),)
|
||||
if text
|
||||
]
|
||||
except Exception:
|
||||
pass
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
def _build_sparql_query(query: str) -> str:
|
||||
safe_query = re.sub(r"[^a-zA-Z0-9 \-]", "", query)[:128]
|
||||
return f"""
|
||||
SELECT ?item ?itemLabel WHERE {{
|
||||
?item rdfs:label "{safe_query}"@en .
|
||||
SERVICE wikibase:label {{ bd:serviceParam wikibase:language "[AUTO_LANGUAGE],en". }}
|
||||
}}
|
||||
LIMIT 10
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def _extract_text(binding: dict[str, Any]) -> str:
|
||||
for key in ("itemLabel", "result", "label"):
|
||||
if key in binding:
|
||||
return binding[key].get("value", "")
|
||||
return ""
|
||||
|
||||
|
||||
class FactCheckDataSource(DataSource):
|
||||
def __init__(
|
||||
self,
|
||||
provider: str = "snopes",
|
||||
name: str = "",
|
||||
enabled: bool = True,
|
||||
priority: int = 0,
|
||||
api_key: str | None = None,
|
||||
**config: str,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
name=name or f"factcheck_{provider}", enabled=enabled, priority=priority
|
||||
)
|
||||
self.provider = provider
|
||||
self.api_key = api_key
|
||||
self.config = config
|
||||
|
||||
async def search(
|
||||
self, query: str, limit: int = 5
|
||||
) -> list[DataSourceResult]: # noqa: ARG002
|
||||
return []
|
||||
|
|
@ -0,0 +1,145 @@
|
|||
from __future__ import annotations
|
||||
|
||||
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator import (
|
||||
BiasAnalysis,
|
||||
HallucinationAnalysis,
|
||||
)
|
||||
|
||||
from .patterns import (
|
||||
BIAS_PATTERNS,
|
||||
CITATION_GAP_PATTERNS,
|
||||
FABRICATED_SPECIFICITY_PATTERNS,
|
||||
SOURCE_INDICATOR_PATTERN,
|
||||
UNSOURCED_STATISTIC_PATTERN,
|
||||
PatternRule,
|
||||
)
|
||||
from .utils import clip_example, split_sentences, unique_preserve_order
|
||||
|
||||
|
||||
class BiasDetector:
|
||||
def detect(self, text: str) -> BiasAnalysis:
|
||||
return self.detect_bias(text)
|
||||
|
||||
def detect_bias(self, text: str) -> BiasAnalysis:
|
||||
matches = self._match_patterns(text, BIAS_PATTERNS)
|
||||
score = min(sum(match[2] for match in matches), 1.0)
|
||||
patterns_found = unique_preserve_order(match[0] for match in matches)
|
||||
examples = unique_preserve_order(match[1] for match in matches)
|
||||
|
||||
return BiasAnalysis(
|
||||
bias_detected=bool(matches),
|
||||
score=round(score, 3),
|
||||
patterns_found=list(patterns_found),
|
||||
examples=list(examples),
|
||||
reasoning=self._build_reasoning(patterns_found),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _match_patterns(
|
||||
text: str, patterns: tuple[PatternRule, ...]
|
||||
) -> tuple[tuple[str, str, float], ...]:
|
||||
return tuple(
|
||||
(rule.name, clip_example(match.group(0)), rule.score)
|
||||
for rule in patterns
|
||||
for match in rule.pattern.finditer(text)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _build_reasoning(patterns_found: tuple[str, ...]) -> str:
|
||||
if not patterns_found:
|
||||
return "No bias indicators found."
|
||||
return f"Detected bias indicators: {', '.join(patterns_found)}."
|
||||
|
||||
|
||||
class HallucinationDetector:
|
||||
def detect(self, text: str) -> HallucinationAnalysis:
|
||||
return self.detect_hallucination(text)
|
||||
|
||||
def detect_hallucination(self, text: str) -> HallucinationAnalysis:
|
||||
sentences = split_sentences(text)
|
||||
unsourced_claims = self._find_unsourced_statistics(sentences)
|
||||
missing_citations = self._find_rule_matches(text, CITATION_GAP_PATTERNS)
|
||||
fabricated_specificity = self._find_rule_matches(
|
||||
text, FABRICATED_SPECIFICITY_PATTERNS
|
||||
)
|
||||
patterns_found = self._patterns_found(
|
||||
has_unsourced_claims=bool(unsourced_claims),
|
||||
has_missing_citations=bool(missing_citations),
|
||||
has_fabricated_specificity=bool(fabricated_specificity),
|
||||
)
|
||||
score = self._score(
|
||||
unsourced_claims=unsourced_claims,
|
||||
missing_citations=missing_citations,
|
||||
fabricated_specificity=fabricated_specificity,
|
||||
)
|
||||
examples = unique_preserve_order(
|
||||
tuple(unsourced_claims)
|
||||
+ tuple(missing_citations)
|
||||
+ tuple(fabricated_specificity)
|
||||
)
|
||||
|
||||
return HallucinationAnalysis(
|
||||
hallucination_detected=score > 0,
|
||||
score=round(score, 3),
|
||||
patterns_found=list(patterns_found),
|
||||
examples=list(examples),
|
||||
unsourced_claims=list(unsourced_claims),
|
||||
fabricated_specificity=list(fabricated_specificity),
|
||||
missing_citations=list(missing_citations),
|
||||
reasoning=self._build_reasoning(patterns_found),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _find_unsourced_statistics(sentences: tuple[str, ...]) -> tuple[str, ...]:
|
||||
return tuple(
|
||||
clip_example(sentence)
|
||||
for sentence in sentences
|
||||
if UNSOURCED_STATISTIC_PATTERN.pattern.search(sentence)
|
||||
and SOURCE_INDICATOR_PATTERN.search(sentence) is None
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _find_rule_matches(
|
||||
text: str, rules: tuple[PatternRule, ...]
|
||||
) -> tuple[str, ...]:
|
||||
return unique_preserve_order(
|
||||
clip_example(match.group(0))
|
||||
for rule in rules
|
||||
for match in rule.pattern.finditer(text)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _patterns_found(
|
||||
*,
|
||||
has_unsourced_claims: bool,
|
||||
has_missing_citations: bool,
|
||||
has_fabricated_specificity: bool,
|
||||
) -> tuple[str, ...]:
|
||||
return tuple(
|
||||
pattern_name
|
||||
for pattern_name, detected in (
|
||||
("unsourced_statistics", has_unsourced_claims),
|
||||
("missing_citations", has_missing_citations),
|
||||
("fabricated_specificity", has_fabricated_specificity),
|
||||
)
|
||||
if detected
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _score(
|
||||
*,
|
||||
unsourced_claims: tuple[str, ...],
|
||||
missing_citations: tuple[str, ...],
|
||||
fabricated_specificity: tuple[str, ...],
|
||||
) -> float:
|
||||
unsourced_score = min(len(unsourced_claims) * 0.32, 0.64)
|
||||
citation_score = min(len(missing_citations) * 0.3, 0.6)
|
||||
specificity_score = min(len(fabricated_specificity) * 0.22, 0.44)
|
||||
return min(unsourced_score + citation_score + specificity_score, 1.0)
|
||||
|
||||
@staticmethod
|
||||
def _build_reasoning(patterns_found: tuple[str, ...]) -> str:
|
||||
if not patterns_found:
|
||||
return "No hallucination risk indicators found."
|
||||
return f"Detected hallucination risk indicators: {', '.join(patterns_found)}."
|
||||
|
|
@ -0,0 +1,171 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
|
||||
from .data_sources import DataSource, DataSourceResult
|
||||
|
||||
|
||||
@dataclass
|
||||
class GroundingResult:
|
||||
claim: str
|
||||
is_grounded: bool = False
|
||||
confidence: float = 0.0
|
||||
reasoning: str = ""
|
||||
supporting_docs: list[dict[str, Any]] = field(default_factory=list)
|
||||
sources_searched: list[str] = field(default_factory=list)
|
||||
metadata: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
class GroundingChecker:
|
||||
def __init__(
|
||||
self,
|
||||
data_sources: list[DataSource],
|
||||
confidence_threshold: float = 0.6,
|
||||
timeout_per_source: float = 5.0,
|
||||
) -> None:
|
||||
self.data_sources = sorted(data_sources, key=lambda x: -x.priority)
|
||||
self.confidence_threshold = confidence_threshold
|
||||
self.timeout_per_source = timeout_per_source
|
||||
|
||||
async def check_claim_grounding(self, claim: str) -> GroundingResult:
|
||||
if not claim or not claim.strip():
|
||||
return GroundingResult(claim=claim, reasoning="Claim is empty")
|
||||
|
||||
claim_elements = self._extract_verifiable_elements(claim)
|
||||
if not any(claim_elements.values()):
|
||||
return GroundingResult(
|
||||
claim=claim, reasoning="No verifiable elements found in claim"
|
||||
)
|
||||
|
||||
enabled_sources = [ds for ds in self.data_sources if ds.enabled]
|
||||
if not enabled_sources:
|
||||
return GroundingResult(
|
||||
claim=claim, reasoning="No enabled data sources available"
|
||||
)
|
||||
|
||||
raw_results = await asyncio.gather(
|
||||
*[self._search_source_safe(source, claim) for source in enabled_sources],
|
||||
)
|
||||
|
||||
supporting_docs: list[dict[str, Any]] = []
|
||||
sources_searched: list[str] = [ds.name for ds in enabled_sources]
|
||||
max_confidence = 0.0
|
||||
|
||||
for items in raw_results:
|
||||
for item in items:
|
||||
boosted = self._boost_confidence(item, claim_elements)
|
||||
supporting_docs.append(
|
||||
{
|
||||
"text": item.text[:200],
|
||||
"source": item.source,
|
||||
"confidence": boosted,
|
||||
"match_score": item.confidence,
|
||||
}
|
||||
)
|
||||
max_confidence = max(max_confidence, boosted)
|
||||
|
||||
is_grounded = max_confidence >= self.confidence_threshold
|
||||
if is_grounded:
|
||||
reasoning = f"Claim verified with {max_confidence:.1%} confidence across {len(supporting_docs)} sources"
|
||||
elif supporting_docs:
|
||||
reasoning = f"Partial match found but confidence ({max_confidence:.1%}) below threshold ({self.confidence_threshold:.1%})"
|
||||
else:
|
||||
reasoning = f"No matching data found in {len(enabled_sources)} sources"
|
||||
|
||||
return GroundingResult(
|
||||
claim=claim,
|
||||
is_grounded=is_grounded,
|
||||
confidence=max_confidence,
|
||||
supporting_docs=supporting_docs,
|
||||
sources_searched=sources_searched,
|
||||
reasoning=reasoning,
|
||||
)
|
||||
|
||||
async def verify_multiple_claims(self, claims: list[str]) -> list[GroundingResult]:
|
||||
return list(
|
||||
await asyncio.gather(*[self.check_claim_grounding(c) for c in claims])
|
||||
)
|
||||
|
||||
async def _search_source_safe(
|
||||
self, source: DataSource, claim: str
|
||||
) -> list[DataSourceResult]:
|
||||
try:
|
||||
return await asyncio.wait_for(
|
||||
source.search(claim, limit=3), timeout=self.timeout_per_source
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
verbose_logger.warning(
|
||||
f"Timeout searching {source.name} for claim: {claim}"
|
||||
)
|
||||
return []
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"Error searching {source.name}: {e}")
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
def _extract_verifiable_elements(claim: str) -> dict[str, list[str]]:
|
||||
stop_words = frozenset(
|
||||
(
|
||||
"the",
|
||||
"a",
|
||||
"an",
|
||||
"and",
|
||||
"or",
|
||||
"but",
|
||||
"in",
|
||||
"is",
|
||||
"are",
|
||||
"was",
|
||||
"be",
|
||||
"been",
|
||||
"being",
|
||||
"have",
|
||||
"has",
|
||||
"had",
|
||||
"do",
|
||||
"does",
|
||||
"did",
|
||||
)
|
||||
)
|
||||
numbers = re.findall(r"\b\d+(?:,\d{3})*(?:\.\d+)?\s?%?|\b\d{4}\b", claim)
|
||||
dates = re.findall(r"\b(?:\d{1,2}/\d{1,2}/\d{4}|\d{4})\b", claim)
|
||||
entities = re.findall(r"\b[A-Z][a-z]+(?:\s+[A-Z][a-z]+)*\b", claim)
|
||||
keywords = [
|
||||
w
|
||||
for w in re.findall(r"\b[a-z]+\b", claim.lower())
|
||||
if w not in stop_words and len(w) > 3
|
||||
]
|
||||
return {
|
||||
"numbers": numbers[:5],
|
||||
"dates": dates[:3],
|
||||
"entities": entities[:5],
|
||||
"keywords": keywords[:8],
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _boost_confidence(
|
||||
result: DataSourceResult, claim_elements: dict[str, list[str]]
|
||||
) -> float:
|
||||
boost = 0.0
|
||||
result_numbers = set(
|
||||
re.findall(r"\b\d+(?:,\d{3})*(?:\.\d+)?\s?%?|\b\d{4}\b", result.text)
|
||||
)
|
||||
if set(claim_elements.get("numbers", [])) & result_numbers:
|
||||
boost += 0.1
|
||||
result_entities = set(
|
||||
re.findall(r"\b[A-Z][a-z]+(?:\s+[A-Z][a-z]+)*\b", result.text)
|
||||
)
|
||||
if set(claim_elements.get("entities", [])) & result_entities:
|
||||
boost += 0.15
|
||||
result_words = set(re.findall(r"\b[a-z]+\b", result.text.lower()))
|
||||
claim_keywords = set(claim_elements.get("keywords", []))
|
||||
if claim_keywords and result_words:
|
||||
boost += min(
|
||||
0.2, len(claim_keywords & result_words) / len(claim_keywords) * 0.25
|
||||
)
|
||||
return round(min(1.0, result.confidence + boost), 2)
|
||||
|
|
@ -0,0 +1,115 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from typing import Pattern
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PatternRule:
|
||||
name: str
|
||||
pattern: Pattern[str]
|
||||
score: float
|
||||
|
||||
|
||||
BIAS_PATTERNS: tuple[PatternRule, ...] = (
|
||||
PatternRule(
|
||||
name="dogmatic_language",
|
||||
pattern=re.compile(
|
||||
r"\b(?:always|never|obviously|clearly|undeniably|without question|"
|
||||
r"everyone knows|common sense|the fact is)\b",
|
||||
re.IGNORECASE,
|
||||
),
|
||||
score=0.18,
|
||||
),
|
||||
PatternRule(
|
||||
name="opinion_as_fact",
|
||||
pattern=re.compile(
|
||||
r"\b(?:i believe|i think|in my opinion|should be|must be|has to be|"
|
||||
r"need to be)\b",
|
||||
re.IGNORECASE,
|
||||
),
|
||||
score=0.16,
|
||||
),
|
||||
PatternRule(
|
||||
name="overconfidence",
|
||||
pattern=re.compile(
|
||||
r"\b(?:100%|guaranteed|certainly|definitely|there is no doubt|"
|
||||
r"cannot be wrong|will definitely)\b",
|
||||
re.IGNORECASE,
|
||||
),
|
||||
score=0.22,
|
||||
),
|
||||
PatternRule(
|
||||
name="sweeping_generalization",
|
||||
pattern=re.compile(
|
||||
r"\b(?:all|no|every)\s+[a-z][a-z\-]{2,}(?:\s+[a-z][a-z\-]{2,}){0,2}"
|
||||
r"\s+(?:are|is|can|cannot|can't|will|won't|should|must)\b",
|
||||
re.IGNORECASE,
|
||||
),
|
||||
score=0.24,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
UNSOURCED_STATISTIC_PATTERN = PatternRule(
|
||||
name="unsourced_statistics",
|
||||
pattern=re.compile(
|
||||
r"(?:\b\d+(?:\.\d+)?\s?%|\b\d+\s+(?:out of|in)\s+\d+\b|"
|
||||
r"\b\d+(?:,\d{3})+(?:\.\d+)?\b)",
|
||||
re.IGNORECASE,
|
||||
),
|
||||
score=0.32,
|
||||
)
|
||||
|
||||
|
||||
CITATION_GAP_PATTERNS: tuple[PatternRule, ...] = (
|
||||
PatternRule(
|
||||
name="unnamed_research",
|
||||
pattern=re.compile(
|
||||
r"\b(?:research shows|studies show|studies found|a study found|"
|
||||
r"scientists found|experts say|experts agree|according to experts)\b",
|
||||
re.IGNORECASE,
|
||||
),
|
||||
score=0.3,
|
||||
),
|
||||
PatternRule(
|
||||
name="vague_authority",
|
||||
pattern=re.compile(
|
||||
r"\b(?:it is widely known|it has been proven|data proves|"
|
||||
r"evidence proves)\b",
|
||||
re.IGNORECASE,
|
||||
),
|
||||
score=0.26,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
FABRICATED_SPECIFICITY_PATTERNS: tuple[PatternRule, ...] = (
|
||||
PatternRule(
|
||||
name="overly_precise_number",
|
||||
pattern=re.compile(
|
||||
r"\b(?:exactly|precisely)\s+\d+(?:,\d{3})*(?:\.\d+)?\b|" r"\b\d+\.\d{3,}\b",
|
||||
re.IGNORECASE,
|
||||
),
|
||||
score=0.24,
|
||||
),
|
||||
PatternRule(
|
||||
name="specific_date_claim",
|
||||
pattern=re.compile(
|
||||
r"\bon\s+(?:jan(?:uary)?|feb(?:ruary)?|mar(?:ch)?|apr(?:il)?|"
|
||||
r"may|jun(?:e)?|jul(?:y)?|aug(?:ust)?|sep(?:tember)?|oct(?:ober)?|"
|
||||
r"nov(?:ember)?|dec(?:ember)?)\s+\d{1,2},\s+\d{4}\b",
|
||||
re.IGNORECASE,
|
||||
),
|
||||
score=0.18,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
SOURCE_INDICATOR_PATTERN = re.compile(
|
||||
r"\b(?:according to|cited by|source:|doi:|isbn|https?://|www\.|"
|
||||
r"published in|published by|journal|report|whitepaper|survey by|study by|"
|
||||
r"dataset|citation|reference)\b",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
|
@ -0,0 +1,108 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from typing import Literal
|
||||
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator import (
|
||||
BiasAnalysis,
|
||||
HallucinationAnalysis,
|
||||
RiskScore,
|
||||
UncertaintyAnalysis,
|
||||
)
|
||||
|
||||
|
||||
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,
|
||||
) -> 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
|
||||
|
||||
def compute_risk(
|
||||
self,
|
||||
bias_analysis: BiasAnalysis,
|
||||
hallucination_analysis: HallucinationAnalysis,
|
||||
uncertainty_analysis: UncertaintyAnalysis | None = None,
|
||||
) -> RiskScore:
|
||||
overall_risk = self._weighted_score(
|
||||
bias_score=bias_analysis.score,
|
||||
hallucination_score=hallucination_analysis.score,
|
||||
)
|
||||
return RiskScore(
|
||||
overall_risk_percentage=round(overall_risk * 100),
|
||||
bias_score=bias_analysis.score,
|
||||
hallucination_score=hallucination_analysis.score,
|
||||
detected_issues=self._detected_issues(
|
||||
bias_analysis=bias_analysis,
|
||||
hallucination_analysis=hallucination_analysis,
|
||||
uncertainty_analysis=uncertainty_analysis,
|
||||
),
|
||||
recommendation=self.determine_recommendation(
|
||||
overall_risk=overall_risk,
|
||||
bias_score=bias_analysis.score,
|
||||
hallucination_score=hallucination_analysis.score,
|
||||
),
|
||||
)
|
||||
|
||||
def _weighted_score(
|
||||
self,
|
||||
*,
|
||||
bias_score: float,
|
||||
hallucination_score: float,
|
||||
) -> float:
|
||||
total_weight = self.bias_weight + self.hallucination_weight
|
||||
if total_weight <= 0:
|
||||
return 0.0
|
||||
return min(
|
||||
(
|
||||
bias_score * self.bias_weight
|
||||
+ hallucination_score * self.hallucination_weight
|
||||
)
|
||||
/ total_weight,
|
||||
1.0,
|
||||
)
|
||||
|
||||
def determine_recommendation(
|
||||
self,
|
||||
*,
|
||||
overall_risk: float,
|
||||
bias_score: float,
|
||||
hallucination_score: float,
|
||||
) -> Literal["pass", "flag", "block"]:
|
||||
if (
|
||||
overall_risk >= self.block_threshold
|
||||
or bias_score >= self.bias_threshold
|
||||
or hallucination_score >= self.hallucination_threshold
|
||||
):
|
||||
return "block"
|
||||
if overall_risk >= self.flag_threshold:
|
||||
return "flag"
|
||||
return "pass"
|
||||
|
||||
@staticmethod
|
||||
def _detected_issues(
|
||||
*,
|
||||
bias_analysis: BiasAnalysis,
|
||||
hallucination_analysis: HallucinationAnalysis,
|
||||
uncertainty_analysis: UncertaintyAnalysis | None,
|
||||
) -> list[str]:
|
||||
uncertainty_issues: tuple[str, ...] = (
|
||||
tuple(f"uncertainty:{p}" for p in uncertainty_analysis.patterns_found)
|
||||
if uncertainty_analysis and uncertainty_analysis.uncertainty_detected
|
||||
else ()
|
||||
)
|
||||
return list(
|
||||
tuple(f"bias:{p}" for p in bias_analysis.patterns_found)
|
||||
+ tuple(f"hallucination:{p}" for p in hallucination_analysis.patterns_found)
|
||||
+ uncertainty_issues
|
||||
)
|
||||
|
|
@ -0,0 +1,28 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from typing import Iterable
|
||||
|
||||
SENTENCE_SPLIT_PATTERN = re.compile(r"(?<=[.!?])\s+")
|
||||
|
||||
|
||||
def split_sentences(text: str) -> tuple[str, ...]:
|
||||
normalized_text = " ".join(text.split())
|
||||
if not normalized_text:
|
||||
return ()
|
||||
return tuple(
|
||||
sentence.strip()
|
||||
for sentence in SENTENCE_SPLIT_PATTERN.split(normalized_text)
|
||||
if sentence.strip()
|
||||
)
|
||||
|
||||
|
||||
def clip_example(text: str, max_length: int = 160) -> str:
|
||||
normalized_text = " ".join(text.split())
|
||||
if len(normalized_text) <= max_length:
|
||||
return normalized_text
|
||||
return f"{normalized_text[: max_length - 3]}..."
|
||||
|
||||
|
||||
def unique_preserve_order(values: Iterable[str]) -> tuple[str, ...]:
|
||||
return tuple(dict.fromkeys(value for value in values if value))
|
||||
|
|
@ -8,6 +8,9 @@ from typing_extensions import Required, TypedDict
|
|||
from litellm.types.proxy.guardrails.guardrail_hooks.akto import (
|
||||
AktoConfigModel,
|
||||
)
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator import (
|
||||
BiasHallucinationEstimatorConfigModel,
|
||||
)
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.block_code_execution import (
|
||||
BlockCodeExecutionGuardrailConfigModel,
|
||||
)
|
||||
|
|
@ -122,7 +125,7 @@ class SupportedGuardrailIntegrations(Enum):
|
|||
RUBRIK = "rubrik"
|
||||
VIGIL_GUARD = "vigil_guard"
|
||||
REPELLOAI = "repelloai"
|
||||
HEADROOM = "headroom"
|
||||
BIAS_HALLUCINATION_ESTIMATOR = "bias_hallucination_estimator"
|
||||
|
||||
|
||||
class Role(Enum):
|
||||
|
|
@ -910,6 +913,7 @@ class LitellmParams(
|
|||
CiscoAIDefenseGuardrailConfigModel,
|
||||
PresidioConfigModel,
|
||||
BedrockGuardrailConfigModel,
|
||||
BiasHallucinationEstimatorConfigModel,
|
||||
LakeraV2GuardrailConfigModel,
|
||||
HeadroomGuardrailConfigModel,
|
||||
RepelloAIGuardrailConfigModel,
|
||||
|
|
|
|||
|
|
@ -0,0 +1,114 @@
|
|||
from typing import Literal
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
|
||||
|
||||
|
||||
class BiasAnalysis(BaseModel):
|
||||
"""Model representing the result of a bias analysis."""
|
||||
|
||||
bias_detected: bool = Field(
|
||||
default=False, description="Indicates if bias was detected in the text"
|
||||
)
|
||||
score: float = Field(
|
||||
default=0.0,
|
||||
ge=0.0,
|
||||
le=1.0,
|
||||
description="Score representing the level of bias detected",
|
||||
)
|
||||
patterns_found: list[str] = Field(
|
||||
default_factory=list,
|
||||
description="List of patterns or phrases that indicate bias in the text",
|
||||
)
|
||||
examples: list[str] = Field(
|
||||
default_factory=list,
|
||||
description="Examples of text segments that were identified as biased",
|
||||
)
|
||||
reasoning: str = Field(
|
||||
default="",
|
||||
description="Explanation of the reasoning behind the bias detection and scoring",
|
||||
)
|
||||
|
||||
|
||||
class HallucinationAnalysis(BaseModel):
|
||||
"""Model representing the result of a hallucination analysis."""
|
||||
|
||||
hallucination_detected: bool = Field(
|
||||
default=False,
|
||||
description="Indicates if hallucination risk was detected in the text",
|
||||
)
|
||||
score: float = Field(
|
||||
default=0.0,
|
||||
ge=0.0,
|
||||
le=1.0,
|
||||
description="Score representing the level of hallucination risk detected",
|
||||
)
|
||||
patterns_found: list[str] = Field(
|
||||
default_factory=list,
|
||||
description="List of patterns or phrases that indicate hallucination risk",
|
||||
)
|
||||
examples: list[str] = Field(
|
||||
default_factory=list,
|
||||
description="Examples of text segments that were identified as risky",
|
||||
)
|
||||
unsourced_claims: list[str] = Field(default_factory=list)
|
||||
fabricated_specificity: list[str] = Field(default_factory=list)
|
||||
missing_citations: list[str] = Field(default_factory=list)
|
||||
reasoning: str = Field(
|
||||
default="",
|
||||
description="Explanation of the reasoning behind the hallucination risk score",
|
||||
)
|
||||
|
||||
|
||||
class UncertaintyAnalysis(BaseModel):
|
||||
"""Model representing the result of an uncertainty analysis based on logprobs."""
|
||||
|
||||
uncertainty_detected: bool = Field(
|
||||
default=False,
|
||||
description="Indicates if high uncertainty was detected in the text based on logprobs",
|
||||
)
|
||||
score: float = Field(
|
||||
default=0.0,
|
||||
ge=0.0,
|
||||
le=1.0,
|
||||
description="Score representing the level of uncertainty detected based on logprobs",
|
||||
)
|
||||
patterns_found: list[str] = Field(default_factory=list)
|
||||
examples: list[str] = Field(default_factory=list)
|
||||
reasoning: str = Field(default="")
|
||||
|
||||
|
||||
class RiskScore(BaseModel):
|
||||
"""Model representing the overall risk score combining bias and hallucination."""
|
||||
|
||||
overall_risk_percentage: int = Field(
|
||||
default=0, ge=0, le=100, description="Overall risk percentage (0-100)"
|
||||
)
|
||||
bias_score: float = Field(default=0.0, ge=0.0, le=1.0)
|
||||
hallucination_score: float = Field(default=0.0, ge=0.0, le=1.0)
|
||||
uncertainty_score: float = Field(default=0.0, ge=0.0, le=1.0)
|
||||
detected_issues: list[str] = Field(default_factory=list)
|
||||
recommendation: Literal["pass", "flag", "block"] = Field(default="pass")
|
||||
|
||||
|
||||
class BiasHallucinationEstimatorConfigModel(
|
||||
GuardrailConfigModel
|
||||
): # pyright: ignore[reportMissingTypeArgument]
|
||||
"""Configuration schema for the native bias and hallucination estimator."""
|
||||
|
||||
bias_threshold: float = Field(default=0.5, ge=0.0, le=1.0)
|
||||
hallucination_threshold: float = Field(default=0.5, ge=0.0, le=1.0)
|
||||
risk_flag_threshold: float = Field(default=0.25, ge=0.0, le=1.0)
|
||||
risk_block_threshold: float = Field(default=0.5, ge=0.0, le=1.0)
|
||||
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 = Field(default=0.4, ge=0.0)
|
||||
hallucination_weight: float = Field(default=0.6, ge=0.0)
|
||||
|
||||
@staticmethod
|
||||
def ui_friendly_name() -> str:
|
||||
return "LiteLLM Bias & Hallucination Estimator"
|
||||
File diff suppressed because it is too large
Load diff
Loading…
Add table
Reference in a new issue