feat(guardrails): add self-contained llm_as_a_judge guardrail hook

This commit is contained in:
Ishaan Jaffer 2026-04-23 15:18:20 -07:00
parent 96a2b3e42f
commit ffdf63df76
No known key found for this signature in database

View file

@ -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": "<name>", "score": <0-100>, "reasoning": "<one sentence>", "passed": <true|false>, "weight": <weight>}
],
"overall_score": <weighted average 0-100>
}"""
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",
]