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:
A Emmanuel 2026-06-23 18:38:54 +05:30 • committed by Sameer Kankute
parent 84eac5de06
commit 1d15734fc8
No known key found for this signature in database
11 changed files with 3316 additions and 1 deletions

View file

@ -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,
}

View file

@ -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

View file

@ -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 []

View file

@ -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)}."

View file

@ -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)

View file

@ -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,
)

View file

@ -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
)

View file

@ -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))

View file

@ -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,

View file

@ -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"