diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 781f0e9f301..f6d9d8975be 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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 diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_guardrail_pipeline.py b/tests/test_litellm/proxy/utils/proxy_logging/test_guardrail_pipeline.py index 1ff9fbf8d83..b9f209e5251 100644 --- a/tests/test_litellm/proxy/utils/proxy_logging/test_guardrail_pipeline.py +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_guardrail_pipeline.py @@ -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" # ---------------------------------------------------------------------------