fix(guardrails): instrument during-call and post-call guardrail latency (#31414)

litellm_guardrail_latency_seconds was only emitted for pre-call guardrails.
during_call_hook and post_call_success_hook ran guardrails without recording
any latency, so during-call and post-call guardrail time was invisible in the
metric and leaked into litellm_overhead_latency_metric, making the documented
"subtract guardrail latency from overhead" workaround under-report total
guardrail time.

Extract the find-the-PrometheusLogger-and-record step into _emit_guardrail_metrics
and add _run_guardrail_with_metrics, a single wrapper that times a guardrail
coroutine, classifies its outcome (success / intervened / error), enriches any
raised HTTPException, and records the latency under the given hook_type. Route
the pre-call emit, during_call_hook, and post_call_success_hook through it so
every guardrail phase contributes to the metric the same way.

Resolves LIT-3999
This commit is contained in:
Yassin Kortam 2026-06-26 23:07:50 +03:00 • committed by GitHub
parent bd5e046464
commit f2fa23b0ec
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 171 additions and 48 deletions

View file

@ -100,6 +100,7 @@ from litellm.proxy.hooks.sensitive_data_routing import (
_PROXY_SensitiveDataRoutingHandler,
)
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.prometheus import PrometheusLogger
from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
from litellm.integrations.SlackAlerting.utils import _add_langfuse_trace_id_to_alert
from litellm.litellm_core_utils.litellm_logging import Logging
@ -1124,7 +1125,6 @@ class ProxyLogging:
Returns:
Updated data dictionary if guardrail passes, None if guardrail should be skipped
"""
from litellm.integrations.prometheus import PrometheusLogger
from litellm.types.guardrails import GuardrailEventHooks
# Determine the event type based on call type
@ -1197,17 +1197,13 @@ class ProxyLogging:
or "unknown"
)
# Find PrometheusLogger in callbacks and record metrics
for prom_callback in litellm.callbacks:
if isinstance(prom_callback, PrometheusLogger):
prom_callback._record_guardrail_metrics(
guardrail_name=metrics_guardrail_name,
latency_seconds=latency_seconds,
status=status,
error_type=error_type,
hook_type="pre_call",
)
break
self._emit_guardrail_metrics(
guardrail_name=metrics_guardrail_name,
latency_seconds=latency_seconds,
status=status,
error_type=error_type,
hook_type="pre_call",
)
return data
@ -1625,19 +1621,58 @@ class ProxyLogging:
return data
@staticmethod
async def _run_guardrail_task_with_enrichment(
callback: Any, coro: Awaitable[Any]
def _emit_guardrail_metrics(
guardrail_name: str,
latency_seconds: float,
status: str,
error_type: Optional[str],
hook_type: str,
) -> None:
for prom_callback in litellm.callbacks:
if isinstance(prom_callback, PrometheusLogger):
prom_callback._record_guardrail_metrics(
guardrail_name=guardrail_name,
latency_seconds=latency_seconds,
status=status,
error_type=error_type,
hook_type=hook_type,
)
break
@staticmethod
async def _run_guardrail_with_metrics(
callback: Any, coro: Awaitable[Any], hook_type: str
) -> Any:
"""
Await `coro`; if it raises an HTTPException with dict detail,
enrich the detail with the originating callback's `guardrail_name`
and `guardrail_mode` before re-raising.
Await `coro`, recording its latency and status to the
`litellm_guardrail_latency_seconds` metric under `hook_type`, and
enriching any raised HTTPException with the originating callback's
`guardrail_name`/`guardrail_mode` before re-raising.
"""
guardrail_name = (
getattr(callback, "guardrail_name", None) or type(callback).__name__
)
start_time = time.perf_counter()
status = "success"
error_type: Optional[str] = None
try:
return await coro
except SensitiveDataRouteException:
status = "intervened"
raise
except Exception as e:
status = "error"
error_type = type(e).__name__
_enrich_http_exception_with_guardrail_context(e, callback)
raise
finally:
ProxyLogging._emit_guardrail_metrics(
guardrail_name=guardrail_name,
latency_seconds=time.perf_counter() - start_time,
status=status,
error_type=error_type,
hook_type=hook_type,
)
@staticmethod
async def _wrap_streaming_iterator_with_enrichment(
@ -1853,22 +1888,24 @@ class ProxyLogging:
and not getattr(callback, "use_native_during_call_hook", False)
):
data["guardrail_to_apply"] = callback
guardrail_task = self._run_guardrail_task_with_enrichment(
guardrail_task = self._run_guardrail_with_metrics(
callback,
unified_guardrail.async_moderation_hook(
user_api_key_dict=user_api_key_dict,
data=data,
call_type=call_type,
),
"during_call",
)
else:
guardrail_task = self._run_guardrail_task_with_enrichment(
guardrail_task = self._run_guardrail_with_metrics(
callback,
callback.async_moderation_hook(
data=data,
user_api_key_dict=user_api_key_auth_dict, # type: ignore
call_type=call_type, # type: ignore
),
"during_call",
)
guardrail_tasks.append(guardrail_task)
@ -2394,29 +2431,25 @@ class ProxyLogging:
if "apply_guardrail" in type(callback).__dict__:
data["guardrail_to_apply"] = callback
try:
guardrail_response = (
await unified_guardrail.async_post_call_success_hook(
user_api_key_dict=user_api_key_dict,
data=data,
response=response,
)
)
except Exception as e:
_enrich_http_exception_with_guardrail_context(e, callback)
raise
guardrail_response = await self._run_guardrail_with_metrics(
callback,
unified_guardrail.async_post_call_success_hook(
user_api_key_dict=user_api_key_dict,
data=data,
response=response,
),
"post_call",
)
else:
try:
guardrail_response = (
await callback.async_post_call_success_hook(
user_api_key_dict=user_api_key_dict,
data=data,
response=response,
)
)
except Exception as e:
_enrich_http_exception_with_guardrail_context(e, callback)
raise
guardrail_response = await self._run_guardrail_with_metrics(
callback,
callback.async_post_call_success_hook(
user_api_key_dict=user_api_key_dict,
data=data,
response=response,
),
"post_call",
)
if guardrail_response is not None:
response = guardrail_response

View file

@ -4,7 +4,7 @@ Covers ``_should_use_guardrail_load_balancing``, ``_execute_guardrail_hook``,
``_execute_guardrail_with_load_balancing``, ``_process_guardrail_callback``,
``_process_prompt_template``, ``_process_guardrail_metadata``,
``_maybe_execute_pipelines``, ``_handle_pipeline_result``,
``_run_guardrail_task_with_enrichment``.
``_run_guardrail_with_metrics``, ``_emit_guardrail_metrics``.
"""
from __future__ import annotations
@ -21,6 +21,7 @@ from litellm.integrations.custom_guardrail import (
CustomGuardrail,
ModifyResponseException,
)
from litellm.integrations.prometheus import PrometheusLogger
from litellm.proxy.utils import ProxyLogging
from litellm.types.guardrails import GuardrailEventHooks
@ -425,23 +426,42 @@ def test_handle_pipeline_result_unknown_action_returns_data():
# ---------------------------------------------------------------------------
# _run_guardrail_task_with_enrichment
# _run_guardrail_with_metrics
# ---------------------------------------------------------------------------
def _prometheus_callback() -> MagicMock:
"""Stand-in PrometheusLogger that records ``_record_guardrail_metrics`` calls.
``MagicMock(spec=PrometheusLogger)`` passes the ``isinstance`` check inside
``_emit_guardrail_metrics`` while letting us capture the recorded labels.
"""
return MagicMock(spec=PrometheusLogger)
@pytest.mark.asyncio
async def test_run_guardrail_task_with_enrichment_passes_result():
async def test_run_guardrail_with_metrics_passes_result_and_records_success(monkeypatch):
async def task():
return {"a": 1, "b": 2, "c": 3}
out = await ProxyLogging._run_guardrail_task_with_enrichment(
callback=MagicMock(guardrail_name="g"), coro=task()
prom = _prometheus_callback()
monkeypatch.setattr(litellm, "callbacks", [prom])
out = await ProxyLogging._run_guardrail_with_metrics(
callback=MagicMock(guardrail_name="g"), coro=task(), hook_type="during_call"
)
assert out == {"a": 1, "b": 2, "c": 3}
recorded = prom._record_guardrail_metrics.call_args.kwargs
assert recorded["guardrail_name"] == "g"
assert recorded["status"] == "success"
assert recorded["error_type"] is None
assert recorded["hook_type"] == "during_call"
assert recorded["latency_seconds"] >= 0
@pytest.mark.asyncio
async def test_run_guardrail_task_with_enrichment_enriches_http_exception_raises():
async def test_run_guardrail_with_metrics_records_error_and_enriches(monkeypatch):
detail = {"error": "blocked"}
async def task():
@ -450,9 +470,79 @@ async def test_run_guardrail_task_with_enrichment_enriches_http_exception_raises
cb = MagicMock()
cb.guardrail_name = "presidio"
cb.event_hook = "pre_call"
prom = _prometheus_callback()
monkeypatch.setattr(litellm, "callbacks", [prom])
with pytest.raises(HTTPException):
await ProxyLogging._run_guardrail_task_with_enrichment(callback=cb, coro=task())
await ProxyLogging._run_guardrail_with_metrics(
callback=cb, coro=task(), hook_type="post_call"
)
assert detail["guardrail_name"] == "presidio"
recorded = prom._record_guardrail_metrics.call_args.kwargs
assert recorded["status"] == "error"
assert recorded["error_type"] == "HTTPException"
assert recorded["hook_type"] == "post_call"
# ---------------------------------------------------------------------------
# during_call / post_call phases emit the latency metric (LIT-3999 regression)
# ---------------------------------------------------------------------------
def _moderation_guardrail() -> MagicMock:
cb = MagicMock(spec=CustomGuardrail)
cb.__class__ = CustomGuardrail
cb.guardrail_name = "g"
cb.event_hook = GuardrailEventHooks.during_call
cb.use_native_during_call_hook = False
cb.should_run_guardrail = MagicMock(return_value=True)
cb.async_moderation_hook = AsyncMock(return_value=None)
cb.async_post_call_success_hook = AsyncMock(return_value=None)
return cb
@pytest.mark.asyncio
async def test_during_call_hook_records_latency_metric(
proxy_logging, make_user_api_key_auth, monkeypatch
):
cb = _moderation_guardrail()
prom = _prometheus_callback()
monkeypatch.setattr(litellm, "callbacks", [prom, cb])
await proxy_logging.during_call_hook(
data={"model": "m"},
user_api_key_dict=make_user_api_key_auth(),
call_type="completion",
)
cb.async_moderation_hook.assert_awaited_once()
recorded = prom._record_guardrail_metrics.call_args.kwargs
assert recorded["hook_type"] == "during_call"
assert recorded["guardrail_name"] == "g"
assert recorded["status"] == "success"
@pytest.mark.asyncio
async def test_post_call_success_hook_records_latency_metric(
proxy_logging, make_user_api_key_auth, monkeypatch
):
cb = _moderation_guardrail()
prom = _prometheus_callback()
monkeypatch.setattr(litellm, "callbacks", [prom, cb])
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
await proxy_logging.post_call_success_hook(
data={"model": "m"},
response=litellm.ModelResponse(),
user_api_key_dict=make_user_api_key_auth(),
)
cb.async_post_call_success_hook.assert_awaited_once()
recorded = prom._record_guardrail_metrics.call_args.kwargs
assert recorded["hook_type"] == "post_call"
assert recorded["guardrail_name"] == "g"
assert recorded["status"] == "success"
# ---------------------------------------------------------------------------