diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_as_a_judge/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/llm_as_a_judge/__init__.py new file mode 100644 index 00000000000..30095f2a511 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_as_a_judge/__init__.py @@ -0,0 +1,274 @@ +"""LLM-as-a-Judge guardrail: uses an LLM to score responses against weighted criteria.""" + +import json +from datetime import datetime +from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, cast + +from fastapi import HTTPException + +from litellm._logging import verbose_logger +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + log_guardrail_information, +) +from litellm.types.guardrails import GuardrailEventHooks, SupportedGuardrailIntegrations +from litellm.types.utils import GenericGuardrailAPIInputs, GuardrailStatus + +if TYPE_CHECKING: + from litellm.types.guardrails import Guardrail, LitellmParams + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.types.utils import StandardLoggingEvalInformation + +JUDGE_SYSTEM_PROMPT = """You are a quality judge. Evaluate the assistant's response against the criteria provided. +For each criterion, assign a score from 0 to 100 and provide concise reasoning. +Return ONLY valid JSON in this exact format: +{ + "verdicts": [ + {"criterion_name": "", "score": <0-100>, "reasoning": "", "passed": , "weight": } + ], + "overall_score": +}""" + + +def _build_judge_prompt( + criteria: List[Dict[str, Any]], + messages: List[Dict[str, Any]], + response_text: str, +) -> str: + criteria_block = "\n".join( + f'- {c["name"]} (weight {c["weight"]}%): {c.get("description", "")}' + for c in criteria + ) + conversation = "\n".join( + f'{m.get("role", "user").upper()}: {m.get("content", "")}' + for m in messages + if isinstance(m.get("content"), str) + ) + return ( + f"Criteria to evaluate:\n{criteria_block}\n\n" + f"Conversation:\n{conversation}\n\n" + f"Assistant response to evaluate:\n{response_text}" + ) + + +class LLMAsAJudgeGuardrail(CustomGuardrail): + """Post-call (and optionally pre-call) guardrail that judges response quality via an LLM.""" + + def __init__( + self, + guardrail_name: str, + judge_model: str, + criteria: List[Dict[str, Any]], + overall_threshold: float = 80.0, + on_failure: Literal["block", "log"] = "block", + event_hook: Optional[ + Union[GuardrailEventHooks, List[GuardrailEventHooks]] + ] = None, + default_on: bool = False, + **kwargs: Any, + ) -> None: + # Normalize event_hook strings to enum values (matches block_code_execution pattern) + _event_hook: Optional[Union[GuardrailEventHooks, List[GuardrailEventHooks]]] = ( + None + ) + if event_hook is not None: + if isinstance(event_hook, list): + _event_hook = [ + GuardrailEventHooks(h) if isinstance(h, str) else h + for h in event_hook + ] + else: + _event_hook = ( + GuardrailEventHooks(event_hook) + if isinstance(event_hook, str) + else event_hook + ) + + super().__init__( + guardrail_name=guardrail_name, + supported_event_hooks=[ + GuardrailEventHooks.pre_call, + GuardrailEventHooks.post_call, + ], + event_hook=_event_hook or GuardrailEventHooks.post_call, + default_on=default_on, + **kwargs, + ) + self.judge_model = judge_model + self.criteria = criteria + self.overall_threshold = overall_threshold + self.on_failure = on_failure + + async def _run_judge( + self, + messages: List[Dict[str, Any]], + response_text: str, + ) -> Dict[str, Any]: + import litellm + + judge_messages = [ + {"role": "system", "content": JUDGE_SYSTEM_PROMPT}, + { + "role": "user", + "content": _build_judge_prompt(self.criteria, messages, response_text), + }, + ] + response = await litellm.acompletion( + model=self.judge_model, + messages=judge_messages, + response_format={"type": "json_object"}, + temperature=0, + ) + raw = response.choices[0].message.content or "{}" # type: ignore[union-attr] + return json.loads(raw) + + @log_guardrail_information + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional["LiteLLMLoggingObj"] = None, + ) -> GenericGuardrailAPIInputs: + start_time = datetime.now() + status: GuardrailStatus = "success" + judge_result: Dict[str, Any] = {} + + try: + texts = inputs.get("texts") or [] + response_text = " ".join(texts) + messages: List[Dict[str, Any]] = request_data.get("messages") or [] + + if not response_text: + return inputs + + try: + judge_result = await self._run_judge(messages, response_text) + except Exception as judge_err: + verbose_logger.warning( + f"llm_as_a_judge guardrail: judge call failed, failing open. Error: {judge_err}" + ) + return inputs + + overall_score = float(judge_result.get("overall_score", 100)) + passed = overall_score >= self.overall_threshold + + # Write judge result to eval_information so EvalViewer on the logs page can render it + _metadata = request_data.setdefault("metadata", {}) + _eval_info = cast( + "StandardLoggingEvalInformation", + { + "eval_name": self.guardrail_name, + "overall_score": overall_score, + "passed": passed, + "judge_model": self.judge_model, + "threshold": self.overall_threshold, + "verdicts": judge_result.get("verdicts", []), + }, + ) + existing = _metadata.get("eval_information") + if isinstance(existing, list): + existing.append(_eval_info) + elif existing is not None: + _metadata["eval_information"] = [existing, _eval_info] + else: + _metadata["eval_information"] = _eval_info + + if not passed and self.on_failure == "block": + status = "guardrail_intervened" + raise HTTPException( + status_code=422, + detail={ + "error": "LLM judge rejected response: score below threshold", + "overall_score": overall_score, + "threshold": self.overall_threshold, + "verdicts": judge_result.get("verdicts", []), + }, + ) + + return inputs + + except HTTPException: + status = "guardrail_intervened" + raise + except Exception as e: + verbose_logger.warning(f"llm_as_a_judge guardrail unexpected error: {e}") + return inputs + finally: + event_type = ( + GuardrailEventHooks.post_call + if input_type == "response" + else GuardrailEventHooks.pre_call + ) + self.add_standard_logging_guardrail_information_to_request_data( + guardrail_provider="llm_as_a_judge", + guardrail_json_response=judge_result, + request_data=request_data, + guardrail_status=status, + start_time=start_time.timestamp(), + end_time=datetime.now().timestamp(), + event_type=event_type, + ) + + +def initialize_guardrail( + litellm_params: "LitellmParams", + guardrail: "Guardrail", +) -> LLMAsAJudgeGuardrail: + import litellm + + def _get(key: str, default: Any = None) -> Any: + val = getattr(litellm_params, key, None) + if val is not None: + return val + raw = guardrail.get("litellm_params") + if isinstance(raw, dict) and key in raw: + return raw[key] + return default + + guardrail_name = guardrail.get("guardrail_name") + if not guardrail_name: + raise ValueError("llm_as_a_judge guardrail requires a guardrail_name") + + judge_model = _get("judge_model") + if not judge_model: + raise ValueError( + "llm_as_a_judge guardrail requires judge_model in litellm_params" + ) + + criteria = _get("criteria") or [] + if not criteria: + raise ValueError("llm_as_a_judge guardrail requires at least one criterion") + + overall_threshold = float(_get("overall_threshold", 80.0)) + on_failure = _get("on_failure", "block") + mode = _get("mode") + event_hook: Optional[GuardrailEventHooks] = None + if isinstance(mode, str) and mode in {e.value for e in GuardrailEventHooks}: + event_hook = GuardrailEventHooks(mode) + + instance = LLMAsAJudgeGuardrail( + guardrail_name=guardrail_name, + judge_model=judge_model, + criteria=criteria, + overall_threshold=overall_threshold, + on_failure=on_failure, + event_hook=event_hook, + default_on=bool(_get("default_on", False)), + ) + litellm.logging_callback_manager.add_litellm_callback(instance) + return instance + + +guardrail_initializer_registry = { + SupportedGuardrailIntegrations.LLM_AS_A_JUDGE.value: initialize_guardrail, +} + +guardrail_class_registry = { + SupportedGuardrailIntegrations.LLM_AS_A_JUDGE.value: LLMAsAJudgeGuardrail, +} + +__all__ = [ + "LLMAsAJudgeGuardrail", + "initialize_guardrail", +]