diff --git a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/__init__.py new file mode 100644 index 00000000000..6237a81d82a --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/__init__.py @@ -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, +} diff --git a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/bias_hallucination_estimator.py b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/bias_hallucination_estimator.py new file mode 100644 index 00000000000..5c8e6975186 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/bias_hallucination_estimator.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/data_sources.py b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/data_sources.py new file mode 100644 index 00000000000..a0582fc607c --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/data_sources.py @@ -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 [] diff --git a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/estimator_core.py b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/estimator_core.py new file mode 100644 index 00000000000..a7f96c03a15 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/estimator_core.py @@ -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)}." diff --git a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/grounding_checker.py b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/grounding_checker.py new file mode 100644 index 00000000000..9539f002cff --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/grounding_checker.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/patterns.py b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/patterns.py new file mode 100644 index 00000000000..9fb1c9c2b08 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/patterns.py @@ -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, +) diff --git a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/risk_scorer.py b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/risk_scorer.py new file mode 100644 index 00000000000..62eb94656f5 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/risk_scorer.py @@ -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 + ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/utils.py b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/utils.py new file mode 100644 index 00000000000..4d7ad5a8ebe --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/utils.py @@ -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)) diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index c3002713580..1b0140e6b8c 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -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, diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator.py b/litellm/types/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator.py new file mode 100644 index 00000000000..6fc2ccf58b0 --- /dev/null +++ b/litellm/types/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator.py @@ -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" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bias_hallucination_estimator.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bias_hallucination_estimator.py new file mode 100644 index 00000000000..d3da0521267 --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bias_hallucination_estimator.py @@ -0,0 +1,1778 @@ +import json +import tempfile +import unittest.mock +from typing import Any + +import pytest + +from litellm.exceptions import GuardrailRaisedException +from litellm.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator.bias_hallucination_estimator import ( + BiasHallucinationEstimatorGuardrail, +) +from litellm.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator.data_sources import ( + ContextDocumentDataSource, + DataSourceResult, + FactCheckDataSource, + FileDataSource, + KnowledgeGraphDataSource, + URLDataSource, + VectorStoreDataSource, + _get_doc_text, + _keyword_search, + _parse_json_docs, +) +from litellm.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator.estimator_core import ( + BiasDetector, + HallucinationDetector, +) +from litellm.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator.grounding_checker import ( + GroundingChecker, +) +from litellm.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator.risk_scorer import ( + RiskScorer, +) +from litellm.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator.utils import ( + clip_example, + split_sentences, + unique_preserve_order, +) +from litellm.types.guardrails import GuardrailEventHooks, SupportedGuardrailIntegrations +from litellm.types.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator import ( + BiasAnalysis, + BiasHallucinationEstimatorConfigModel, + HallucinationAnalysis, +) + +# --------------------------------------------------------------------------- +# BiasDetector +# --------------------------------------------------------------------------- + + +def test_bias_detector_detects_dogmatic_language() -> None: + analysis = BiasDetector().detect( + "Everyone knows this is obviously the only correct answer. It will definitely work for every user." + ) + + assert analysis.bias_detected is True + assert analysis.score > 0 + assert "dogmatic_language" in analysis.patterns_found + assert "overconfidence" in analysis.patterns_found + + +def test_bias_detector_detects_opinion_as_fact() -> None: + analysis = BiasDetector().detect( + "I believe this should be mandatory for all developers." + ) + + assert analysis.bias_detected is True + assert "opinion_as_fact" in analysis.patterns_found + + +def test_bias_detector_detects_sweeping_generalization() -> None: + analysis = BiasDetector().detect("All engineers are skilled at math.") + + assert analysis.bias_detected is True + assert "sweeping_generalization" in analysis.patterns_found + + +def test_bias_detector_score_accumulates_across_patterns() -> None: + analysis = BiasDetector().detect( + "Obviously all developers must be certified. I believe it is 100% guaranteed." + ) + + assert analysis.score > 0.3 + assert len(analysis.patterns_found) >= 2 + assert len(analysis.examples) >= 2 + + +def test_bias_detector_does_not_flag_neutral_text() -> None: + analysis = BiasDetector().detect( + "The timeout defaults to ten seconds and can be configured per request." + ) + + assert analysis.bias_detected is False + assert analysis.score == 0 + assert analysis.patterns_found == [] + + +def test_bias_detector_score_capped_at_one() -> None: + heavily_biased = " ".join( + [ + "Obviously everyone knows this will definitely always be true.", + "I believe it must be that all users should never question it.", + "The fact is it is 100% guaranteed and certainly cannot be wrong.", + ] + ) + analysis = BiasDetector().detect(heavily_biased) + + assert analysis.score <= 1.0 + + +def test_bias_detector_reasoning_reflects_patterns() -> None: + analysis = BiasDetector().detect("Obviously this is the only solution.") + + assert "dogmatic_language" in analysis.reasoning + + +def test_bias_detector_neutral_text_reasoning() -> None: + analysis = BiasDetector().detect("Configure the retry limit in your settings file.") + + assert analysis.reasoning == "No bias indicators found." + + +# --------------------------------------------------------------------------- +# HallucinationDetector +# --------------------------------------------------------------------------- + + +def test_hallucination_detector_detects_unsourced_statistics_and_citation_gaps() -> ( + None +): + analysis = HallucinationDetector().detect( + "Research shows 73% of users switch products on March 14, 2022." + ) + + assert analysis.hallucination_detected is True + assert analysis.score >= 0.5 + assert "unsourced_statistics" in analysis.patterns_found + assert "missing_citations" in analysis.patterns_found + assert "fabricated_specificity" in analysis.patterns_found + + +def test_hallucination_detector_allows_sourced_statistics() -> None: + analysis = HallucinationDetector().detect( + "According to the 2024 usage report, 73% of requests completed successfully." + ) + + assert analysis.hallucination_detected is False + assert analysis.score == 0 + assert analysis.unsourced_claims == [] + + +def test_hallucination_detector_detects_vague_authority() -> None: + analysis = HallucinationDetector().detect( + "It is widely known that this approach is superior." + ) + + assert "missing_citations" in analysis.patterns_found + assert analysis.hallucination_detected is True + + +def test_hallucination_detector_detects_overly_precise_number() -> None: + analysis = HallucinationDetector().detect( + "The system processed exactly 1,234 requests last month." + ) + + assert "fabricated_specificity" in analysis.patterns_found + + +def test_hallucination_detector_multiple_unsourced_claims_increase_score() -> None: + multi_claim = ( + "Studies show 45% of users prefer option A. " + "Research shows 62% switch after six months. " + "Scientists found 3 out of 4 developers agree." + ) + analysis = HallucinationDetector().detect(multi_claim) + + assert analysis.score > 0.5 + assert len(analysis.unsourced_claims) >= 2 + + +def test_hallucination_detector_score_capped_at_one() -> None: + very_risky = " ".join( + [ + "Research shows 73% of users agree.", + "Studies found 2 out of 3 experts concur.", + "It has been proven that exactly 1,234,567 cases exist.", + "According to experts, on January 15, 2023 the rate was 89%.", + "Data proves that scientists found 99% accuracy.", + ] + ) + analysis = HallucinationDetector().detect(very_risky) + + assert analysis.score <= 1.0 + + +def test_hallucination_detector_citation_indicator_clears_number_in_sentence() -> None: + analysis = HallucinationDetector().detect( + "Published in Nature journal: 85% of trials showed improvement." + ) + + assert analysis.unsourced_claims == [] + + +# --------------------------------------------------------------------------- +# RiskScorer +# --------------------------------------------------------------------------- + + +def test_risk_scorer_blocks_when_hallucination_threshold_is_crossed() -> None: + risk = RiskScorer().compute_risk( + bias_analysis=BiasAnalysis(score=0.1), + hallucination_analysis=HallucinationAnalysis( + hallucination_detected=True, + score=0.6, + patterns_found=["unsourced_statistics"], + ), + ) + + assert risk.overall_risk_percentage == 40 + assert risk.recommendation == "block" + assert risk.detected_issues == ["hallucination:unsourced_statistics"] + + +def test_risk_scorer_flags_medium_weighted_risk() -> None: + risk = RiskScorer( + bias_threshold=0.9, + hallucination_threshold=0.9, + ).compute_risk( + bias_analysis=BiasAnalysis( + bias_detected=True, + score=0.4, + patterns_found=["dogmatic_language"], + ), + hallucination_analysis=HallucinationAnalysis(score=0.2), + ) + + assert risk.overall_risk_percentage == 28 + assert risk.recommendation == "flag" + + +def test_risk_scorer_passes_low_risk() -> None: + risk = RiskScorer().compute_risk( + bias_analysis=BiasAnalysis(score=0.1), + hallucination_analysis=HallucinationAnalysis(score=0.1), + ) + + assert risk.recommendation == "pass" + assert risk.overall_risk_percentage < 25 + + +def test_risk_scorer_blocks_on_high_bias_score_alone() -> None: + risk = RiskScorer(bias_threshold=0.4).compute_risk( + bias_analysis=BiasAnalysis( + bias_detected=True, score=0.5, patterns_found=["overconfidence"] + ), + hallucination_analysis=HallucinationAnalysis(score=0.0), + ) + + assert risk.recommendation == "block" + assert "bias:overconfidence" in risk.detected_issues + + +def test_risk_scorer_detected_issues_prefix_by_type() -> None: + risk = RiskScorer().compute_risk( + bias_analysis=BiasAnalysis( + bias_detected=True, + score=0.6, + patterns_found=["dogmatic_language", "overconfidence"], + ), + hallucination_analysis=HallucinationAnalysis( + hallucination_detected=True, + score=0.6, + patterns_found=["unsourced_statistics"], + ), + ) + + assert "bias:dogmatic_language" in risk.detected_issues + assert "bias:overconfidence" in risk.detected_issues + assert "hallucination:unsourced_statistics" in risk.detected_issues + + +def test_risk_scorer_custom_weights_change_overall_percentage() -> None: + bias_only = RiskScorer(bias_weight=1.0, hallucination_weight=0.0).compute_risk( + bias_analysis=BiasAnalysis(score=0.5), + hallucination_analysis=HallucinationAnalysis(score=0.0), + ) + hallucination_only = RiskScorer( + bias_weight=0.0, hallucination_weight=1.0 + ).compute_risk( + bias_analysis=BiasAnalysis(score=0.0), + hallucination_analysis=HallucinationAnalysis(score=0.5), + ) + + assert ( + bias_only.overall_risk_percentage + == hallucination_only.overall_risk_percentage + == 50 + ) + + +def test_risk_scorer_zero_weight_total_returns_zero_percentage() -> None: + risk = RiskScorer(bias_weight=0.0, hallucination_weight=0.0).compute_risk( + bias_analysis=BiasAnalysis(score=0.1), + hallucination_analysis=HallucinationAnalysis(score=0.1), + ) + + assert risk.overall_risk_percentage == 0 + assert risk.recommendation == "pass" + + +# --------------------------------------------------------------------------- +# BiasHallucinationEstimatorGuardrail +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_guardrail_blocks_high_risk_response_and_logs_metadata() -> None: + guardrail = BiasHallucinationEstimatorGuardrail( + guardrail_name="bias-hallucination", + event_hook=GuardrailEventHooks.post_call, + ) + request_data: dict[str, Any] = {} + + with pytest.raises(GuardrailRaisedException) as exc: + await guardrail.apply_guardrail( + inputs={ + "texts": [ + "Research shows 73% of users switch products on March 14, 2022." + ] + }, + request_data=request_data, + input_type="response", + ) + + assert "High bias/hallucination risk detected" in str(exc.value) + logging_entries = request_data["metadata"]["standard_logging_guardrail_information"] + assert logging_entries[0]["guardrail_status"] == "guardrail_intervened" + assert logging_entries[0]["guardrail_response"]["decision"] == "blocked" + + +@pytest.mark.asyncio +async def test_guardrail_log_payload_excludes_text_snippets() -> None: + guardrail = BiasHallucinationEstimatorGuardrail( + guardrail_name="bias-hallucination", + log_only=True, + ) + request_data: dict[str, Any] = {} + + await guardrail.apply_guardrail( + inputs={ + "texts": ["Research shows 73% of users switch products on March 14, 2022."] + }, + request_data=request_data, + input_type="response", + ) + + response = request_data["metadata"]["standard_logging_guardrail_information"][0][ + "guardrail_response" + ] + snippet_fields = { + "examples", + "unsourced_claims", + "missing_citations", + "fabricated_specificity", + } + for entry in response["bias"]: + assert ( + not snippet_fields & entry.keys() + ), f"bias log entry contains snippet fields: {entry.keys()}" + for entry in response["hallucination"]: + assert ( + not snippet_fields & entry.keys() + ), f"hallucination log entry contains snippet fields: {entry.keys()}" + + +@pytest.mark.asyncio +async def test_guardrail_log_only_flags_without_blocking() -> None: + guardrail = BiasHallucinationEstimatorGuardrail( + guardrail_name="bias-hallucination", + event_hook=GuardrailEventHooks.post_call, + log_only=True, + ) + request_data: dict[str, Any] = {} + + result = await guardrail.apply_guardrail( + inputs={ + "texts": ["Research shows 73% of users switch products on March 14, 2022."] + }, + request_data=request_data, + input_type="response", + ) + + assert result["texts"][0].startswith("Research shows") + logging_entries = request_data["metadata"]["standard_logging_guardrail_information"] + assert logging_entries[0]["guardrail_status"] == "success" + assert logging_entries[0]["guardrail_response"]["decision"] == "flagged" + + +@pytest.mark.asyncio +async def test_guardrail_skips_request_by_default() -> None: + guardrail = BiasHallucinationEstimatorGuardrail( + guardrail_name="bias-hallucination", + event_hook=GuardrailEventHooks.pre_call, + ) + request_data: dict[str, Any] = {} + + result = await guardrail.apply_guardrail( + inputs={"texts": ["Research shows 73% of users switch products."]}, + request_data=request_data, + input_type="request", + ) + + assert result["texts"][0].startswith("Research shows") + assert request_data == {} + + +@pytest.mark.asyncio +async def test_guardrail_passes_low_risk_text() -> None: + guardrail = BiasHallucinationEstimatorGuardrail( + guardrail_name="bias-hallucination", + event_hook=GuardrailEventHooks.post_call, + ) + request_data: dict[str, Any] = {} + + result = await guardrail.apply_guardrail( + inputs={ + "texts": [ + "The timeout defaults to ten seconds and can be changed in configuration." + ] + }, + request_data=request_data, + input_type="response", + ) + + assert result["texts"][0].startswith("The timeout") + logging_entries = request_data["metadata"]["standard_logging_guardrail_information"] + assert logging_entries[0]["guardrail_response"]["decision"] == "passed" + + +@pytest.mark.asyncio +async def test_guardrail_respects_custom_violation_message() -> None: + guardrail = BiasHallucinationEstimatorGuardrail( + guardrail_name="bias-hallucination", + violation_message="Custom block message.", + ) + + with pytest.raises(GuardrailRaisedException) as exc: + await guardrail.apply_guardrail( + inputs={ + "texts": [ + "Research shows 73% of users switch products on March 14, 2022." + ] + }, + request_data={}, + input_type="response", + ) + + assert "Custom block message." in str(exc.value) + + +@pytest.mark.asyncio +async def test_guardrail_tool_calls_are_analyzed() -> None: + guardrail = BiasHallucinationEstimatorGuardrail( + guardrail_name="bias-hallucination", + event_hook=GuardrailEventHooks.post_call, + ) + request_data: dict[str, Any] = {} + + with pytest.raises(GuardrailRaisedException): + await guardrail.apply_guardrail( + inputs={ + "texts": [], + "tool_calls": [ + { + "function": { + "name": "search", + "arguments": '{"query": "Research shows 73% switch on March 14, 2022"}', + } + } + ], + }, + request_data=request_data, + input_type="response", + ) + + +@pytest.mark.asyncio +async def test_guardrail_empty_inputs_returns_unchanged() -> None: + guardrail = BiasHallucinationEstimatorGuardrail(guardrail_name="bias-hallucination") + request_data: dict[str, Any] = {} + + result = await guardrail.apply_guardrail( + inputs={"texts": []}, + request_data=request_data, + input_type="response", + ) + + assert result == {"texts": []} + assert request_data == {} + + +@pytest.mark.asyncio +async def test_guardrail_check_request_enabled_detects_bias() -> None: + guardrail = BiasHallucinationEstimatorGuardrail( + guardrail_name="bias-hallucination", + check_request=True, + check_response=False, + event_hook=GuardrailEventHooks.pre_call, + ) + + with pytest.raises(GuardrailRaisedException): + await guardrail.apply_guardrail( + inputs={ + "texts": [ + "Research shows 73% of users switch products on March 14, 2022." + ] + }, + request_data={}, + input_type="request", + ) + + +def test_guardrail_estimate_returns_structured_result() -> None: + guardrail = BiasHallucinationEstimatorGuardrail(guardrail_name="bias-hallucination") + result = guardrail.estimate_bias_hallucination( + "Everyone knows this is obviously the only correct answer." + ) + + assert "bias" in result + assert "hallucination" in result + assert "risk" in result + assert isinstance(result["risk"], dict) + assert "overall_risk_percentage" in result["risk"] # type: ignore[operator] + + +# --------------------------------------------------------------------------- +# Config model +# --------------------------------------------------------------------------- + + +def test_config_model_defaults_and_ui_name() -> None: + config = BiasHallucinationEstimatorConfigModel() + + assert config.bias_threshold == 0.5 + assert config.hallucination_threshold == 0.5 + assert config.check_response is True + assert config.bias_weight == 0.4 + assert config.hallucination_weight == 0.6 + assert ( + BiasHallucinationEstimatorConfigModel.ui_friendly_name() + == "LiteLLM Bias & Hallucination Estimator" + ) + + +def test_config_model_rejects_out_of_range_threshold() -> None: + import pydantic + + with pytest.raises(pydantic.ValidationError): + BiasHallucinationEstimatorConfigModel(bias_threshold=1.5) + + +# --------------------------------------------------------------------------- +# Registry +# --------------------------------------------------------------------------- + + +def test_guardrail_package_exports_registry_entries() -> None: + from litellm.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator import ( + guardrail_class_registry, + guardrail_initializer_registry, + initialize_guardrail, + ) + + key = SupportedGuardrailIntegrations.BIAS_HALLUCINATION_ESTIMATOR.value + + assert key == "bias_hallucination_estimator" + assert guardrail_initializer_registry[key] == initialize_guardrail + assert guardrail_class_registry[key] == BiasHallucinationEstimatorGuardrail + + +# --------------------------------------------------------------------------- +# Utils +# --------------------------------------------------------------------------- + + +def test_split_sentences_basic() -> None: + result = split_sentences("Hello world. This is a test! Is it working?") + + assert result == ("Hello world.", "This is a test!", "Is it working?") + + +def test_split_sentences_empty_string() -> None: + assert split_sentences("") == () + + +def test_split_sentences_normalizes_whitespace() -> None: + result = split_sentences(" Hello world. Next sentence. ") + + assert result == ("Hello world.", "Next sentence.") + + +def test_clip_example_short_text_unchanged() -> None: + assert clip_example("Short text") == "Short text" + + +def test_clip_example_truncates_long_text() -> None: + long_text = "a" * 200 + result = clip_example(long_text) + + assert len(result) == 160 + assert result.endswith("...") + + +def test_unique_preserve_order_deduplicates() -> None: + result = unique_preserve_order(["b", "a", "b", "c", "a"]) + + assert result == ("b", "a", "c") + + +def test_unique_preserve_order_filters_empty_strings() -> None: + result = unique_preserve_order(["a", "", "b", ""]) + + assert result == ("a", "b") + + +# --------------------------------------------------------------------------- +# FileDataSource +# --------------------------------------------------------------------------- + + +def test_file_data_source_nonexistent_path_returns_empty() -> None: + source = FileDataSource("/nonexistent/path/facts.json") + + assert source._documents == [] + + +@pytest.mark.asyncio +async def test_file_data_source_search_returns_empty_for_empty_source() -> None: + source = FileDataSource("/nonexistent/path/facts.json") + results = await source.search("some query") + + assert results == [] + + +@pytest.mark.asyncio +async def test_context_document_source_finds_matching_content() -> None: + source = ContextDocumentDataSource( + documents=[ + {"text": "The company was founded in 2015 by Jane Smith."}, + {"text": "The product has 3.2 million active users."}, + ] + ) + + results = await source.search("company founded 2015") + + assert len(results) > 0 + assert any("2015" in r.text for r in results) + + +@pytest.mark.asyncio +async def test_context_document_source_no_match_returns_empty() -> None: + source = ContextDocumentDataSource( + documents=[{"text": "The company was founded in 2015."}] + ) + + results = await source.search("quantum physics neutron stars") + + assert results == [] + + +# --------------------------------------------------------------------------- +# GroundingChecker +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_grounding_checker_empty_claim_returns_ungrounded() -> None: + checker = GroundingChecker(data_sources=[]) + result = await checker.check_claim_grounding("") + + assert result.is_grounded is False + assert "empty" in result.reasoning.lower() + + +@pytest.mark.asyncio +async def test_grounding_checker_no_enabled_sources_returns_ungrounded() -> None: + checker = GroundingChecker(data_sources=[]) + result = await checker.check_claim_grounding("Company founded in 2015") + + assert result.is_grounded is False + + +@pytest.mark.asyncio +async def test_grounding_checker_verifies_claim_from_context_docs() -> None: + source = ContextDocumentDataSource( + documents=[{"text": "The company was founded in 2015 by Jane Smith."}] + ) + checker = GroundingChecker(data_sources=[source], confidence_threshold=0.3) + result = await checker.check_claim_grounding("company founded 2015") + + assert result.is_grounded is True + assert result.confidence >= 0.3 + assert len(result.supporting_docs) > 0 + + +@pytest.mark.asyncio +async def test_grounding_checker_unverifiable_claim_not_grounded() -> None: + source = ContextDocumentDataSource( + documents=[{"text": "The company was founded in 2015 by Jane Smith."}] + ) + checker = GroundingChecker(data_sources=[source], confidence_threshold=0.6) + result = await checker.check_claim_grounding( + "quantum entanglement photon polarization" + ) + + assert result.is_grounded is False + + +@pytest.mark.asyncio +async def test_grounding_checker_verify_multiple_claims() -> None: + source = ContextDocumentDataSource( + documents=[{"text": "The company was founded in 2015."}] + ) + checker = GroundingChecker(data_sources=[source], confidence_threshold=0.3) + results = await checker.verify_multiple_claims( + ["company founded 2015", "quantum physics neutron stars"] + ) + + assert len(results) == 2 + grounded_flags = [r.is_grounded for r in results] + assert True in grounded_flags + assert False in grounded_flags + + +# --------------------------------------------------------------------------- +# DataSourceResult +# --------------------------------------------------------------------------- + + +def test_data_source_result_clamps_confidence_below_zero() -> None: + result = DataSourceResult(text="hello", source="s", confidence=-0.5) + + assert result.confidence == 0.0 + + +def test_data_source_result_clamps_confidence_above_one() -> None: + result = DataSourceResult(text="hello", source="s", confidence=1.5) + + assert result.confidence == 1.0 + + +def test_data_source_result_default_metadata_is_empty_dict() -> None: + result = DataSourceResult(text="hello", source="s") + + assert result.metadata == {} + + +def test_data_source_result_repr_contains_source_and_confidence() -> None: + result = DataSourceResult(text="hello", source="my_source", confidence=0.75) + + assert "my_source" in repr(result) + assert "0.75" in repr(result) + + +# --------------------------------------------------------------------------- +# DataSource.verify_fact +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_data_source_verify_fact_returns_true_when_match_found() -> None: + source = ContextDocumentDataSource( + documents=[{"text": "The company was founded in 2015."}] + ) + found, text = await source.verify_fact("company founded 2015") + + assert found is True + assert text is not None + assert "2015" in text + + +@pytest.mark.asyncio +async def test_data_source_verify_fact_returns_false_when_no_match() -> None: + source = ContextDocumentDataSource( + documents=[{"text": "The company was founded in 2015."}] + ) + found, text = await source.verify_fact("quantum entanglement photon") + + assert found is False + assert text is None + + +# --------------------------------------------------------------------------- +# Internal helpers +# --------------------------------------------------------------------------- + + +def test_get_doc_text_returns_string_directly() -> None: + assert _get_doc_text("plain string") == "plain string" + + +def test_get_doc_text_extracts_text_key_from_dict() -> None: + assert _get_doc_text({"text": "from dict"}) == "from dict" + + +def test_get_doc_text_falls_back_to_str_when_no_text_key() -> None: + result = _get_doc_text({"other": "value"}) + + assert "other" in result + + +def test_parse_json_docs_wraps_non_list_in_list() -> None: + result = _parse_json_docs({"text": "single doc"}) + + assert isinstance(result, list) + assert len(result) == 1 + + +def test_parse_json_docs_returns_list_unchanged() -> None: + docs = [{"text": "a"}, {"text": "b"}] + result = _parse_json_docs(docs) + + assert result == docs + + +def test_keyword_search_returns_empty_for_empty_query() -> None: + from litellm.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator.data_sources import ( + _build_keyword_index, + ) + + docs = [{"text": "some content"}] + index = _build_keyword_index(docs) + results = _keyword_search("", docs, index, "src", 5) + + assert results == [] + + +def test_keyword_search_returns_empty_when_no_docs_match() -> None: + from litellm.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator.data_sources import ( + _build_keyword_index, + ) + + docs = [{"text": "apples and oranges"}] + index = _build_keyword_index(docs) + results = _keyword_search("quantum neutron", docs, index, "src", 5) + + assert results == [] + + +def test_keyword_search_respects_limit() -> None: + from litellm.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator.data_sources import ( + _build_keyword_index, + ) + + docs = [{"text": f"document about python {i}"} for i in range(10)] + index = _build_keyword_index(docs) + results = _keyword_search("python document", docs, index, "src", 3) + + assert len(results) <= 3 + + +def test_keyword_search_scores_by_word_overlap() -> None: + from litellm.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator.data_sources import ( + _build_keyword_index, + ) + + docs = [ + {"text": "python programming language"}, + {"text": "python snake reptile"}, + ] + index = _build_keyword_index(docs) + results = _keyword_search("python programming", docs, index, "src", 5) + + assert results[0].text == "python programming language" + assert results[0].confidence > results[1].confidence + + +# --------------------------------------------------------------------------- +# FileDataSource with real files +# --------------------------------------------------------------------------- + + +def test_file_data_source_loads_json_file() -> None: + docs = [ + {"text": "Paris is the capital of France."}, + {"text": "Berlin is the capital of Germany."}, + ] + with tempfile.NamedTemporaryFile(mode="w", suffix=".json", delete=False) as f: + json.dump(docs, f) + path = f.name + + source = FileDataSource(path) + + assert len(source._documents) == 2 + + +@pytest.mark.asyncio +async def test_file_data_source_searches_json_content() -> None: + docs = [ + {"text": "Paris is the capital of France."}, + {"text": "Berlin is the capital of Germany."}, + ] + with tempfile.NamedTemporaryFile(mode="w", suffix=".json", delete=False) as f: + json.dump(docs, f) + path = f.name + + source = FileDataSource(path) + results = await source.search("capital France Paris") + + assert len(results) > 0 + assert any("Paris" in r.text for r in results) + + +def test_file_data_source_loads_txt_file() -> None: + with tempfile.NamedTemporaryFile(mode="w", suffix=".txt", delete=False) as f: + f.write("First fact about astronomy.\nSecond fact about physics.\n") + path = f.name + + source = FileDataSource(path) + + assert len(source._documents) == 2 + + +def test_file_data_source_loads_csv_file() -> None: + with tempfile.NamedTemporaryFile(mode="w", suffix=".csv", delete=False) as f: + f.write("row one data\nrow two data\n") + path = f.name + + source = FileDataSource(path) + + assert len(source._documents) == 2 + + +def test_file_data_source_uses_stem_as_default_name() -> None: + source = FileDataSource("/some/path/my_facts.json") + + assert source.name == "file_my_facts" + + +def test_file_data_source_loads_json_object_not_list() -> None: + doc = {"text": "Single document as object."} + with tempfile.NamedTemporaryFile(mode="w", suffix=".json", delete=False) as f: + json.dump(doc, f) + path = f.name + + source = FileDataSource(path) + + assert len(source._documents) == 1 + + +# --------------------------------------------------------------------------- +# URLDataSource +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_url_data_source_returns_empty_without_aiohttp() -> None: + source = URLDataSource(urls=["http://example.com/data.json"]) + + with unittest.mock.patch( + "litellm.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator.data_sources._AIOHTTP_AVAILABLE", + False, + ): + results = await source.search("any query") + + assert results == [] + + +@pytest.mark.asyncio +async def test_url_data_source_parse_content_handles_json() -> None: + json_content = json.dumps([{"text": "fact one"}, {"text": "fact two"}]) + parsed = URLDataSource._parse_content(json_content) + + assert len(parsed) == 2 + assert parsed[0] == {"text": "fact one"} # type: ignore[comparison-overlap] + + +@pytest.mark.asyncio +async def test_url_data_source_parse_content_handles_plain_text() -> None: + parsed = URLDataSource._parse_content("This is not JSON at all.") + + assert len(parsed) == 1 + assert parsed[0] == {"text": "This is not JSON at all."} # type: ignore[comparison-overlap] + + +@pytest.mark.asyncio +async def test_url_data_source_search_returns_cached_on_second_call() -> None: + source = URLDataSource(urls=[]) + source._fetched = True + source._documents = [{"text": "cached document about python"}] + from litellm.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator.data_sources import ( + _build_keyword_index, + ) + + source._index = _build_keyword_index(source._documents) + + results = await source.search("python") + + assert len(results) == 1 + assert "python" in results[0].text + + +@pytest.mark.asyncio +async def test_url_data_source_concurrent_fetch_only_runs_once() -> None: + fetch_count = 0 + + async def fake_fetch_all() -> list[Any]: + nonlocal fetch_count + fetch_count += 1 + return [{"text": "fetched document"}] + + source = URLDataSource(urls=["http://example.com"]) + + with unittest.mock.patch.object(source, "_fetch_all", side_effect=fake_fetch_all): + await source.search("document") + await source.search("document") + + assert fetch_count == 1 + + +# --------------------------------------------------------------------------- +# VectorStoreDataSource +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_vector_store_data_source_returns_empty_without_client() -> None: + source = VectorStoreDataSource(provider="unknown_provider") + + assert source.client is None + results = await source.search("anything") + + assert results == [] + + +@pytest.mark.asyncio +async def test_vector_store_data_source_returns_empty_without_embedding_model() -> None: + source = VectorStoreDataSource( + provider="pinecone", + client=unittest.mock.MagicMock(), + embedding_model=None, + ) + + results = await source.search("query") + + assert results == [] + + +def test_vector_store_data_source_default_name_uses_provider() -> None: + source = VectorStoreDataSource(provider="pinecone") + + assert source.name == "vectorstore_pinecone" + + +# --------------------------------------------------------------------------- +# KnowledgeGraphDataSource +# --------------------------------------------------------------------------- + + +def test_knowledge_graph_sparql_strips_injection_chars() -> None: + malicious = 'normal text"; DROP TABLE items; --' + query = KnowledgeGraphDataSource._build_sparql_query(malicious) + label_value = query.split('label "')[1].split('"')[0] + + assert ";" not in label_value + assert '"' not in label_value + + +def test_knowledge_graph_sparql_caps_at_128_chars() -> None: + long_input = "a" * 200 + query = KnowledgeGraphDataSource._build_sparql_query(long_input) + label_value = query.split('label "')[1].split('"')[0] + + assert len(label_value) <= 128 + + +def test_knowledge_graph_sparql_preserves_alphanumeric_and_spaces() -> None: + query = KnowledgeGraphDataSource._build_sparql_query("Marie Curie 1867") + label_value = query.split('label "')[1].split('"')[0] + + assert label_value == "Marie Curie 1867" + + +def test_knowledge_graph_extract_text_uses_itemlabel_first() -> None: + binding = {"itemLabel": {"value": "Paris"}, "label": {"value": "Other"}} + + assert KnowledgeGraphDataSource._extract_text(binding) == "Paris" + + +def test_knowledge_graph_extract_text_falls_back_to_result_key() -> None: + binding = {"result": {"value": "Some result"}} + + assert KnowledgeGraphDataSource._extract_text(binding) == "Some result" + + +def test_knowledge_graph_extract_text_falls_back_to_label_key() -> None: + binding = {"label": {"value": "A label"}} + + assert KnowledgeGraphDataSource._extract_text(binding) == "A label" + + +def test_knowledge_graph_extract_text_returns_empty_for_unknown_keys() -> None: + binding = {"unknown_key": {"value": "ignored"}} + + assert KnowledgeGraphDataSource._extract_text(binding) == "" + + +@pytest.mark.asyncio +async def test_knowledge_graph_returns_empty_without_aiohttp() -> None: + source = KnowledgeGraphDataSource() + + with unittest.mock.patch( + "litellm.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator.data_sources._AIOHTTP_AVAILABLE", + False, + ): + results = await source.search("Paris") + + assert results == [] + + +# --------------------------------------------------------------------------- +# FactCheckDataSource +# --------------------------------------------------------------------------- + + +def test_fact_check_data_source_default_name_uses_provider() -> None: + source = FactCheckDataSource() + + assert source.name == "factcheck_snopes" + + +def test_fact_check_data_source_custom_name_overrides_default() -> None: + source = FactCheckDataSource(name="my_checker") + + assert source.name == "my_checker" + + +@pytest.mark.asyncio +async def test_fact_check_data_source_search_returns_empty() -> None: + source = FactCheckDataSource(provider="snopes", api_key="dummy") + results = await source.search("any claim") + + assert results == [] + + +# --------------------------------------------------------------------------- +# URLDataSource._fetch_url via aiohttp mock +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_url_data_source_fetch_url_returns_parsed_json() -> None: + payload = json.dumps([{"text": "fact from url"}]) + mock_response = unittest.mock.AsyncMock() + mock_response.status = 200 + mock_response.text = unittest.mock.AsyncMock(return_value=payload) + mock_response.__aenter__ = unittest.mock.AsyncMock(return_value=mock_response) + mock_response.__aexit__ = unittest.mock.AsyncMock(return_value=False) + + mock_session = unittest.mock.AsyncMock() + mock_session.get = unittest.mock.MagicMock(return_value=mock_response) + mock_session.__aenter__ = unittest.mock.AsyncMock(return_value=mock_session) + mock_session.__aexit__ = unittest.mock.AsyncMock(return_value=False) + + with unittest.mock.patch("aiohttp.ClientSession", return_value=mock_session): + source = URLDataSource(urls=["http://example.com/data.json"]) + results = await source.search("fact url") + + assert any("fact from url" in r.text for r in results) + + +@pytest.mark.asyncio +async def test_url_data_source_fetch_url_handles_non_200_response() -> None: + mock_response = unittest.mock.AsyncMock() + mock_response.status = 404 + mock_response.__aenter__ = unittest.mock.AsyncMock(return_value=mock_response) + mock_response.__aexit__ = unittest.mock.AsyncMock(return_value=False) + + mock_session = unittest.mock.AsyncMock() + mock_session.get = unittest.mock.MagicMock(return_value=mock_response) + mock_session.__aenter__ = unittest.mock.AsyncMock(return_value=mock_session) + mock_session.__aexit__ = unittest.mock.AsyncMock(return_value=False) + + with unittest.mock.patch("aiohttp.ClientSession", return_value=mock_session): + source = URLDataSource(urls=["http://example.com/data.json"]) + results = await source.search("any query") + + assert results == [] + + +# --------------------------------------------------------------------------- +# VectorStoreDataSource.search with mocked client +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_vector_store_data_source_pinecone_returns_results() -> None: + mock_client = unittest.mock.MagicMock() + mock_client.query.return_value = { + "matches": [ + {"metadata": {"text": "pinecone result one"}, "score": 0.9}, + {"metadata": {"text": "pinecone result two"}, "score": 0.7}, + ] + } + mock_model = unittest.mock.MagicMock() + mock_model.encode.return_value = [0.1, 0.2, 0.3] + source = VectorStoreDataSource( + provider="pinecone", client=mock_client, embedding_model=mock_model + ) + + results = await source.search("any query") + + assert len(results) == 2 + assert results[0].text == "pinecone result one" + assert results[0].confidence == 0.9 + + +@pytest.mark.asyncio +async def test_vector_store_data_source_weaviate_returns_results() -> None: + mock_client = unittest.mock.MagicMock() + ( + mock_client.query.get.return_value.with_near_vector.return_value.with_limit.return_value.do.return_value + ) = {"data": {"Get": [{"text": "weaviate result"}]}} + mock_model = unittest.mock.MagicMock() + mock_model.encode.return_value = [0.1, 0.2, 0.3] + source = VectorStoreDataSource( + provider="weaviate", client=mock_client, embedding_model=mock_model + ) + + results = await source.search("any query") + + assert len(results) == 1 + assert results[0].text == "weaviate result" + + +@pytest.mark.asyncio +async def test_vector_store_data_source_embedding_list_passthrough() -> None: + mock_client = unittest.mock.MagicMock() + mock_client.query.return_value = {"matches": []} + mock_model = unittest.mock.MagicMock() + mock_model.encode.return_value = [0.5, 0.6] + source = VectorStoreDataSource( + provider="pinecone", client=mock_client, embedding_model=mock_model + ) + + results = await source.search("query") + + assert results == [] + mock_model.encode.assert_called_once_with("query", convert_to_tensor=False) + + +# --------------------------------------------------------------------------- +# KnowledgeGraphDataSource.search via aiohttp mock +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_knowledge_graph_returns_results_from_sparql_response() -> None: + sparql_response = { + "results": { + "bindings": [ + {"itemLabel": {"value": "Paris"}}, + {"itemLabel": {"value": "Lyon"}}, + ] + } + } + mock_response = unittest.mock.AsyncMock() + mock_response.status = 200 + mock_response.json = unittest.mock.AsyncMock(return_value=sparql_response) + mock_response.__aenter__ = unittest.mock.AsyncMock(return_value=mock_response) + mock_response.__aexit__ = unittest.mock.AsyncMock(return_value=False) + + mock_session = unittest.mock.AsyncMock() + mock_session.get = unittest.mock.MagicMock(return_value=mock_response) + mock_session.__aenter__ = unittest.mock.AsyncMock(return_value=mock_session) + mock_session.__aexit__ = unittest.mock.AsyncMock(return_value=False) + + with unittest.mock.patch("aiohttp.ClientSession", return_value=mock_session): + source = KnowledgeGraphDataSource() + results = await source.search("French cities") + + assert len(results) == 2 + assert results[0].text == "Paris" + assert results[0].confidence == 0.9 + + +@pytest.mark.asyncio +async def test_knowledge_graph_skips_bindings_with_empty_text() -> None: + sparql_response = { + "results": { + "bindings": [ + {"unknown_key": {"value": "ignored"}}, + {"itemLabel": {"value": "Berlin"}}, + ] + } + } + mock_response = unittest.mock.AsyncMock() + mock_response.status = 200 + mock_response.json = unittest.mock.AsyncMock(return_value=sparql_response) + mock_response.__aenter__ = unittest.mock.AsyncMock(return_value=mock_response) + mock_response.__aexit__ = unittest.mock.AsyncMock(return_value=False) + + mock_session = unittest.mock.AsyncMock() + mock_session.get = unittest.mock.MagicMock(return_value=mock_response) + mock_session.__aenter__ = unittest.mock.AsyncMock(return_value=mock_session) + mock_session.__aexit__ = unittest.mock.AsyncMock(return_value=False) + + with unittest.mock.patch("aiohttp.ClientSession", return_value=mock_session): + source = KnowledgeGraphDataSource() + results = await source.search("German cities") + + assert len(results) == 1 + assert results[0].text == "Berlin" + + +@pytest.mark.asyncio +async def test_knowledge_graph_handles_non_200_response() -> None: + mock_response = unittest.mock.AsyncMock() + mock_response.status = 500 + mock_response.__aenter__ = unittest.mock.AsyncMock(return_value=mock_response) + mock_response.__aexit__ = unittest.mock.AsyncMock(return_value=False) + + mock_session = unittest.mock.AsyncMock() + mock_session.get = unittest.mock.MagicMock(return_value=mock_response) + mock_session.__aenter__ = unittest.mock.AsyncMock(return_value=mock_session) + mock_session.__aexit__ = unittest.mock.AsyncMock(return_value=False) + + with unittest.mock.patch("aiohttp.ClientSession", return_value=mock_session): + source = KnowledgeGraphDataSource() + results = await source.search("query") + + assert results == [] + + +# --------------------------------------------------------------------------- +# initialize_guardrail (__init__.py coverage) +# --------------------------------------------------------------------------- + + +def test_initialize_guardrail_wires_up_guardrail_from_litellm_params() -> None: + from litellm.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator import ( + initialize_guardrail, + ) + from litellm.types.guardrails import Guardrail, LitellmParams + + litellm_params = LitellmParams( + guardrail="bias_hallucination_estimator", + mode="post_call", + ) + guardrail: Guardrail = { + "guardrail_name": "test-guardrail", + "litellm_params": litellm_params, + "guardrail_id": "grd-001", + } + + instance = initialize_guardrail(litellm_params=litellm_params, guardrail=guardrail) + + assert isinstance(instance, BiasHallucinationEstimatorGuardrail) + assert instance.guardrail_id == "grd-001" + assert instance.bias_threshold == 0.5 + assert instance.hallucination_threshold == 0.5 + + +def test_initialize_guardrail_uses_custom_thresholds() -> None: + from litellm.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator import ( + initialize_guardrail, + ) + from litellm.types.guardrails import Guardrail, LitellmParams + + litellm_params = LitellmParams( + guardrail="bias_hallucination_estimator", + mode="post_call", + bias_threshold=0.3, + hallucination_threshold=0.4, + block_on_high_risk=False, + log_only=True, + ) + guardrail: Guardrail = { + "guardrail_name": "custom-guardrail", + "litellm_params": litellm_params, + } + + instance = initialize_guardrail(litellm_params=litellm_params, guardrail=guardrail) + + assert instance.bias_threshold == 0.3 + assert instance.hallucination_threshold == 0.4 + assert instance.block_on_high_risk is False + assert instance.log_only is True + + +# --------------------------------------------------------------------------- +# BiasHallucinationEstimatorGuardrail — uncovered static branches +# --------------------------------------------------------------------------- + + +def test_normalize_event_hook_with_mode_returns_mode() -> None: + from litellm.types.guardrails import Mode + + mode = Mode(tags={}, default=None) + result = BiasHallucinationEstimatorGuardrail._normalize_event_hook(mode) + + assert result == mode + + +def test_normalize_event_hook_with_list_returns_list_of_hooks() -> None: + result = BiasHallucinationEstimatorGuardrail._normalize_event_hook( + ["pre_call", "post_call"] + ) + + assert isinstance(result, list) + assert GuardrailEventHooks.pre_call in result + assert GuardrailEventHooks.post_call in result + + +def test_detected_issues_returns_empty_for_non_dict() -> None: + from litellm.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator.bias_hallucination_estimator import ( + BiasHallucinationEstimatorGuardrail, + ) + + assert BiasHallucinationEstimatorGuardrail._detected_issues("not a dict") == () + + +def test_detected_issues_returns_empty_when_detected_issues_not_list() -> None: + assert ( + BiasHallucinationEstimatorGuardrail._detected_issues( + {"detected_issues": "not a list"} + ) + == () + ) + + +def test_violation_categories_returns_empty_when_risk_scores_not_list() -> None: + assert ( + BiasHallucinationEstimatorGuardrail._violation_categories( + {"risk_scores": "bad"} + ) + == [] + ) + + +def test_tool_call_text_returns_none_for_dict_without_function() -> None: + from litellm.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator.bias_hallucination_estimator import ( + BiasHallucinationEstimatorGuardrail, + ) + + assert BiasHallucinationEstimatorGuardrail._tool_call_text({"other": "key"}) is None + + +def test_tool_call_text_returns_none_for_object_without_function_attr() -> None: + assert BiasHallucinationEstimatorGuardrail._tool_call_text(object()) is None + + +def test_tool_call_text_handles_protocol_object() -> None: + class FakeFunction: + name = "my_tool" + arguments = '{"x": 1}' + + class FakeToolCall: + function = FakeFunction() + + result = BiasHallucinationEstimatorGuardrail._tool_call_text(FakeToolCall()) + + assert result == 'my_tool {"x": 1}' + + +def test_string_value_returns_empty_for_none() -> None: + assert BiasHallucinationEstimatorGuardrail._string_value(None) == "" + + +def test_string_value_converts_non_string() -> None: + assert BiasHallucinationEstimatorGuardrail._string_value(42) == "42" + + +def test_get_config_model_returns_config_class() -> None: + result = BiasHallucinationEstimatorGuardrail.get_config_model() + + assert result is not None + + +# --------------------------------------------------------------------------- +# GroundingChecker — uncovered branches +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_grounding_checker_no_verifiable_elements_returns_ungrounded() -> None: + source = ContextDocumentDataSource(documents=[{"text": "some content"}]) + checker = GroundingChecker(data_sources=[source]) + result = await checker.check_claim_grounding("a an the and or") + + assert result.is_grounded is False + assert "verifiable" in result.reasoning.lower() + + +@pytest.mark.asyncio +async def test_grounding_checker_source_failure_returns_ungrounded() -> None: + source = ContextDocumentDataSource(documents=[{"text": "company founded 2015"}]) + source.name = "failing_source" + + async def always_fails(query: str, limit: int = 5) -> list[Any]: + raise RuntimeError("connection refused") + + source.search = always_fails # type: ignore[method-assign] + checker = GroundingChecker(data_sources=[source], confidence_threshold=0.3) + result = await checker.check_claim_grounding("company founded 2015") + + assert result.is_grounded is False + + +@pytest.mark.asyncio +async def test_grounding_checker_partial_match_below_threshold_is_not_grounded() -> ( + None +): + source = ContextDocumentDataSource( + documents=[{"text": "The company was founded in 2015 by Jane Smith."}] + ) + checker = GroundingChecker(data_sources=[source], confidence_threshold=0.99) + # 2 of 5 query words match → base confidence 0.4; boosted still well below 0.99 + result = await checker.check_claim_grounding( + "company founded quantum photon neutron" + ) + + assert result.is_grounded is False + assert "confidence" in result.reasoning.lower() + + +@pytest.mark.asyncio +async def test_grounding_checker_timeout_returns_empty_not_error() -> None: + import asyncio as _asyncio + + source = ContextDocumentDataSource(documents=[{"text": "anything"}]) + + async def slow_search(query: str, limit: int = 5) -> list[Any]: + await _asyncio.sleep(10) + return [] + + source.search = slow_search # type: ignore[method-assign] + checker = GroundingChecker( + data_sources=[source], confidence_threshold=0.3, timeout_per_source=0.01 + ) + result = await checker.check_claim_grounding("company founded 2015") + + assert result.is_grounded is False + + +# --------------------------------------------------------------------------- +# FileDataSource — exception path (corrupt file) +# --------------------------------------------------------------------------- + + +def test_file_data_source_corrupt_json_returns_empty() -> None: + import tempfile + + with tempfile.NamedTemporaryFile(mode="w", suffix=".json", delete=False) as f: + f.write("{ this is not valid json }") + path = f.name + + source = FileDataSource(path) + + assert source._documents == [] + + +# --------------------------------------------------------------------------- +# DataSource.verify_fact — base class helper +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_data_source_verify_fact_returns_true_when_result_found() -> None: + source = ContextDocumentDataSource( + documents=[{"text": "Python was created by Guido van Rossum."}] + ) + found, text = await source.verify_fact("Python created Guido") + assert found is True + assert text is not None + assert "Python" in text + + +@pytest.mark.asyncio +async def test_data_source_verify_fact_returns_false_when_no_result() -> None: + source = ContextDocumentDataSource(documents=[]) + found, text = await source.verify_fact("anything") + assert found is False + assert text is None + + +# --------------------------------------------------------------------------- +# _keyword_search — empty query_words branch +# --------------------------------------------------------------------------- + + +def test_keyword_search_returns_empty_for_non_word_query() -> None: + documents: list[str | dict[str, Any]] = [{"text": "hello world"}] + index = {"hello": [0], "world": [0]} + result = _keyword_search("!!!", documents, index, "src", 5) + assert result == [] + + +# --------------------------------------------------------------------------- +# URLDataSource._fetch_url — exception path +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_url_data_source_fetch_url_handles_exception() -> None: + source = URLDataSource(urls=["http://example.com"]) + + async def boom(*_args: object, **_kwargs: object) -> Any: + raise RuntimeError("network error") + + with unittest.mock.patch("aiohttp.ClientSession.get", side_effect=boom): + import sys + + sys.modules.setdefault("aiohttp", unittest.mock.MagicMock()) + result = await source._fetch_url("http://example.com") + + assert result == [] + + +# --------------------------------------------------------------------------- +# VectorStoreDataSource — _initialize_client provider paths +# --------------------------------------------------------------------------- + + +def test_vector_store_initialize_client_pinecone_returns_index() -> None: + mock_index = object() + mock_pinecone = unittest.mock.MagicMock() + mock_pinecone.Index.return_value = mock_index + + with unittest.mock.patch.dict( + __import__("sys").modules, {"pinecone": mock_pinecone} + ): + client = VectorStoreDataSource._initialize_client( + "pinecone", {"api_key": "k", "index_name": "idx"} + ) + + assert client is mock_index + + +def test_vector_store_initialize_client_pinecone_no_credentials_returns_none() -> None: + mock_pinecone = unittest.mock.MagicMock() + with unittest.mock.patch.dict( + __import__("sys").modules, {"pinecone": mock_pinecone} + ): + client = VectorStoreDataSource._initialize_client("pinecone", {}) + + assert client is None + + +def test_vector_store_initialize_client_weaviate_returns_client() -> None: + mock_client_obj = object() + mock_weaviate = unittest.mock.MagicMock() + mock_weaviate.Client.return_value = mock_client_obj + + with unittest.mock.patch.dict( + __import__("sys").modules, {"weaviate": mock_weaviate} + ): + client = VectorStoreDataSource._initialize_client( + "weaviate", {"url": "http://localhost:8080"} + ) + + assert client is mock_client_obj + + +def test_vector_store_initialize_client_weaviate_no_url_returns_none() -> None: + mock_weaviate = unittest.mock.MagicMock() + with unittest.mock.patch.dict( + __import__("sys").modules, {"weaviate": mock_weaviate} + ): + client = VectorStoreDataSource._initialize_client("weaviate", {}) + + assert client is None + + +def test_vector_store_initialize_client_weaviate_import_error_returns_none() -> None: + import sys + + saved = sys.modules.pop("weaviate", None) + try: + with unittest.mock.patch.dict(sys.modules, {"weaviate": None}): # type: ignore[dict-item] + client = VectorStoreDataSource._initialize_client( + "weaviate", {"url": "http://localhost:8080"} + ) + assert client is None + finally: + if saved is not None: + sys.modules["weaviate"] = saved + + +# --------------------------------------------------------------------------- +# VectorStoreDataSource — _load_embedding_model +# --------------------------------------------------------------------------- + + +def test_vector_store_load_embedding_model_returns_transformer() -> None: + mock_model = object() + mock_st = unittest.mock.MagicMock() + mock_st.SentenceTransformer.return_value = mock_model + + with unittest.mock.patch.dict( + __import__("sys").modules, {"sentence_transformers": mock_st} + ): + model = VectorStoreDataSource._load_embedding_model() + + assert model is mock_model + + +# --------------------------------------------------------------------------- +# VectorStoreDataSource.search — exception path +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_vector_store_search_exception_returns_empty() -> None: + mock_client = unittest.mock.MagicMock() + mock_client.query.side_effect = RuntimeError("pinecone unavailable") + mock_model = unittest.mock.MagicMock() + mock_model.encode.return_value = [0.1, 0.2, 0.3] + + source = VectorStoreDataSource( + provider="pinecone", client=mock_client, embedding_model=mock_model + ) + results = await source.search("test query") + + assert results == [] + + +# --------------------------------------------------------------------------- +# KnowledgeGraphDataSource.search — exception path +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_knowledge_graph_search_exception_returns_empty() -> None: + import sys + + mock_aiohttp = unittest.mock.MagicMock() + mock_session = unittest.mock.AsyncMock() + mock_session.__aenter__ = unittest.mock.AsyncMock( + side_effect=RuntimeError("connection error") + ) + mock_aiohttp.ClientSession.return_value = mock_session + + with unittest.mock.patch.dict(sys.modules, {"aiohttp": mock_aiohttp}): + source = KnowledgeGraphDataSource() + with unittest.mock.patch( + "litellm.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator.data_sources._AIOHTTP_AVAILABLE", + True, + ): + results = await source.search("test") + + assert results == [] + + +# --------------------------------------------------------------------------- +# GroundingChecker._boost_confidence — entity match branch (line 164) +# --------------------------------------------------------------------------- + + +def test_boost_confidence_entity_match_increases_score() -> None: + from litellm.proxy.guardrails.guardrail_hooks.bias_hallucination_estimator.grounding_checker import ( + GroundingChecker, + ) + + result = DataSourceResult( + text="Albert Einstein published the theory of relativity.", + source="test", + confidence=0.5, + ) + claim_elements = { + "numbers": [], + "dates": [], + "entities": ["Albert Einstein"], + "keywords": ["theory", "relativity"], + } + boosted = GroundingChecker._boost_confidence(result, claim_elements) + + assert boosted > 0.5