From 7ac26f23d81cee061750f088a1e46363dda9f0d7 Mon Sep 17 00:00:00 2001 From: devarakondasrikanth Date: Fri, 13 Mar 2026 19:36:52 -0700 Subject: [PATCH] feat(proxy): add async_post_guardrail_log_success_event for post-guardrail logging - CustomLogger: new async_post_guardrail_log_success_event (runs after post-call hooks) - ProxyLogging: invoke only when subclass overrides; end_time at invoke, llm_end_time in kwargs - Non-streaming: call after _override_openai_response_model in common_request_processing - Streaming: collect chunks (exclude_none only), fire hook in background; drain on shutdown - Tests: override vs base skip, llm_end_time, exception isolation, guardrail skip Made-with: Cursor --- litellm/integrations/custom_logger.py | 11 + litellm/proxy/common_request_processing.py | 7 + litellm/proxy/proxy_server.py | 107 ++++++--- litellm/proxy/utils.py | 67 +++++- ..._async_post_guardrail_log_success_event.py | 211 ++++++++++++++++++ 5 files changed, 376 insertions(+), 27 deletions(-) create mode 100644 tests/test_litellm/proxy/hooks/test_async_post_guardrail_log_success_event.py diff --git a/litellm/integrations/custom_logger.py b/litellm/integrations/custom_logger.py index 06ba9675ca2..7cc636fa496 100644 --- a/litellm/integrations/custom_logger.py +++ b/litellm/integrations/custom_logger.py @@ -174,6 +174,17 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): pass + async def async_post_guardrail_log_success_event( + self, kwargs, response_obj, start_time, end_time + ): + """ + Called by the proxy after post-call hooks (e.g. guardrails) have run. + Use this to log the final response seen by the client; async_log_success_event + runs before post-call hooks and sees the unmodified response. + Override this method to log post-guardrail responses; the base no-op is not invoked. + """ + pass + async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): pass diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index a9e9d519f6f..b74c1937b0e 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1082,6 +1082,13 @@ class ProxyBaseLLMRequestProcessing: log_context=f"litellm_call_id={logging_obj.litellm_call_id}", ) + await proxy_logging_obj.async_post_guardrail_log_success_event( + data=self.data, + response=response, + user_api_key_dict=user_api_key_dict, + logging_obj=logging_obj, + ) + hidden_params = ( getattr(response, "_hidden_params", {}) or {} ) # get any updated response headers diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 2a9be0a67c9..bc694f49e27 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -98,6 +98,7 @@ from litellm.proxy.common_utils.callback_utils import ( ) from litellm.proxy.common_utils.realtime_utils import _realtime_request_body from litellm.types.utils import ( + LLMResponseTypes, ModelResponse, ModelResponseStream, TextCompletionResponse, @@ -731,6 +732,10 @@ async def proxy_shutdown_event(): # [DO NOT BLOCK shutdown events for this] pass + # Drain in-flight post-guardrail log tasks so they complete before exit + if _post_guardrail_log_tasks: + await asyncio.gather(*_post_guardrail_log_tasks, return_exceptions=True) + ## RESET CUSTOM VARIABLES ## cleanup_router_config_variables() @@ -1574,6 +1579,8 @@ open_telemetry_logger: Optional[OpenTelemetry] = None proxy_logging_obj = ProxyLogging( user_api_key_cache=user_api_key_cache, premium_user=premium_user ) +# Strong refs to post-guardrail log tasks so they complete before shutdown +_post_guardrail_log_tasks: Set[asyncio.Task[None]] = set() ### REDIS QUEUE ### async_result = None celery_app_conn = None @@ -5497,6 +5504,57 @@ def _restamp_streaming_chunk_model( return chunk, model_mismatch_logged +async def _async_data_generator_fire_post_guardrail_log( + request_data: dict, + user_api_key_dict: UserAPIKeyAuth, + chunks_for_log: List[Dict[str, Any]], + logging_obj: Optional[Any], +) -> None: + """Build full response from streaming chunks and run post-guardrail log hook.""" + if not chunks_for_log: + return + try: + complete_response = litellm.stream_chunk_builder(chunks=chunks_for_log) + if complete_response is not None: + await proxy_logging_obj.async_post_guardrail_log_success_event( + data=request_data, + response=cast(LLMResponseTypes, complete_response), + user_api_key_dict=user_api_key_dict, + logging_obj=logging_obj, + ) + except Exception as e: + verbose_proxy_logger.exception("Error in post-guardrail log (streaming): %s", e) + + +async def _async_data_generator_emit_error( + e: Exception, + request_data: dict, + user_api_key_dict: UserAPIKeyAuth, +) -> str: + """Run failure hook and return the SSE error payload to yield. Re-raises HTTPException.""" + await proxy_logging_obj.post_call_failure_hook( + user_api_key_dict=user_api_key_dict, + original_exception=e, + request_data=request_data, + ) + verbose_proxy_logger.debug( + f"\033[1;31mAn error occurred: {e}\n\n Debug this by setting `--debug`, e.g. `litellm --model gpt-3.5-turbo --debug`" + ) + if isinstance(e, HTTPException): + raise e + if isinstance(e, StreamingCallbackError): + error_msg = str(e) + else: + error_msg = str(e) + proxy_exception = ProxyException( + message=getattr(e, "message", error_msg), + type=getattr(e, "type", "None"), + param=getattr(e, "param", "None"), + code=getattr(e, "status_code", 500), + ) + return f"data: {json.dumps({'error': proxy_exception.to_dict()})}\n\n" + + async def async_data_generator( response, user_api_key_dict: UserAPIKeyAuth, request_data: dict ): @@ -5507,6 +5565,8 @@ async def async_data_generator( request_data=request_data ) model_mismatch_logged = False + # Chunks for post-guardrail log: use exclude_none=True only so stream_chunk_builder gets required keys + _streaming_chunks_for_log: List[Dict[str, Any]] = [] # Use a running string instead of list + join to avoid O(n^2) overhead. # Previously "".join(str_so_far_parts) was called every chunk, re-joining # the entire accumulated response. String += is O(n) amortized total. @@ -5536,6 +5596,9 @@ async def async_data_generator( ) if isinstance(chunk, BaseModel): + _streaming_chunks_for_log.append( + chunk.model_dump(mode="json", exclude_none=True) + ) chunk = chunk.model_dump_json(exclude_none=True, exclude_unset=True) elif isinstance(chunk, str) and chunk.startswith("data: "): error_message = chunk @@ -5546,6 +5609,21 @@ async def async_data_generator( except Exception as e: yield f"data: {str(e)}\n\n" + # Post-guardrail log: run in background so we don't block yielding [DONE] + def _discard_task(t: asyncio.Task[None]) -> None: + _post_guardrail_log_tasks.discard(t) + + _task = asyncio.create_task( + _async_data_generator_fire_post_guardrail_log( + request_data=request_data, + user_api_key_dict=user_api_key_dict, + chunks_for_log=_streaming_chunks_for_log, + logging_obj=request_data.get("litellm_logging_obj"), + ) + ) + _post_guardrail_log_tasks.add(_task) + _task.add_done_callback(_discard_task) + # Streaming is done, yield the [DONE] chunk if error_message is not None: yield error_message @@ -5557,33 +5635,10 @@ async def async_data_generator( str(e) ) ) - await proxy_logging_obj.post_call_failure_hook( - user_api_key_dict=user_api_key_dict, - original_exception=e, - request_data=request_data, + error_payload = await _async_data_generator_emit_error( + e, request_data, user_api_key_dict ) - verbose_proxy_logger.debug( - f"\033[1;31mAn error occurred: {e}\n\n Debug this by setting `--debug`, e.g. `litellm --model gpt-3.5-turbo --debug`" - ) - - if isinstance(e, HTTPException): - raise e - elif isinstance(e, StreamingCallbackError): - error_msg = str(e) - else: - # Only include the error message, not the traceback. - # The traceback is already logged above via verbose_proxy_logger.exception(). - # Including it in the SSE response leaks internal details to clients. - error_msg = str(e) - - proxy_exception = ProxyException( - message=getattr(e, "message", error_msg), - type=getattr(e, "type", "None"), - param=getattr(e, "param", "None"), - code=getattr(e, "status_code", 500), - ) - error_returned = json.dumps({"error": proxy_exception.to_dict()}) - yield f"data: {error_returned}\n\n" + yield error_payload finally: # Close the response stream to release the underlying HTTP connection # back to the connection pool. This prevents pool exhaustion when diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index b9a1bfb9062..1b1ee29ee16 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -1976,6 +1976,71 @@ class ProxyLogging: raise e return response + async def async_post_guardrail_log_success_event( + self, + data: dict, + response: LLMResponseTypes, + user_api_key_dict: UserAPIKeyAuth, + logging_obj: Optional[Any] = None, + ) -> None: + """ + Invoke async_post_guardrail_log_success_event on CustomLogger callbacks that + override the method (not the base no-op). Called after post_call_success_hook + so loggers see the post-guardrail response. end_time is when this hook runs; + llm_end_time in kwargs is when the LLM call finished (pre-guardrail). + """ + try: + kwargs = dict(data) + if logging_obj is not None and getattr( + logging_obj, "model_call_details", None + ): + kwargs = {**logging_obj.model_call_details, **kwargs} + kwargs["user_api_key_dict"] = user_api_key_dict + start_time = None + if logging_obj is not None: + start_time = getattr(logging_obj, "completion_start_time", None) or ( + kwargs.get("start_time") + if isinstance(kwargs.get("start_time"), datetime) + else None + ) + if getattr(logging_obj, "model_call_details", {}).get("end_time"): + kwargs["llm_end_time"] = logging_obj.model_call_details["end_time"] + + for callback in litellm.callbacks: + _callback: Optional[CustomLogger] = None + if isinstance(callback, str): + _callback = litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class( + cast(_custom_logger_compatible_callbacks_literal, callback) + ) + else: + _callback = callback # type: ignore + + if _callback is None or isinstance(_callback, CustomGuardrail): + continue + if not isinstance(_callback, CustomLogger): + continue + if ( + type(_callback).async_post_guardrail_log_success_event + is CustomLogger.async_post_guardrail_log_success_event + ): + continue + try: + end_time = datetime.now(timezone.utc) + await _callback.async_post_guardrail_log_success_event( + kwargs=kwargs, + response_obj=response, + start_time=start_time, + end_time=end_time, + ) + except Exception as e: + verbose_proxy_logger.exception( + "Error in async_post_guardrail_log_success_event: %s", e + ) + except Exception as e: + verbose_proxy_logger.exception( + "Error in async_post_guardrail_log_success_event: %s", e + ) + async def post_call_response_headers_hook( self, data: dict, @@ -5233,7 +5298,7 @@ def normalize_route_for_root_path(route: str) -> Optional[str]: root_path = get_server_root_path() if root_path and root_path != "/": if route.startswith(root_path + "/"): - return route[len(root_path):] + return route[len(root_path) :] return None return route diff --git a/tests/test_litellm/proxy/hooks/test_async_post_guardrail_log_success_event.py b/tests/test_litellm/proxy/hooks/test_async_post_guardrail_log_success_event.py new file mode 100644 index 00000000000..88d94592b08 --- /dev/null +++ b/tests/test_litellm/proxy/hooks/test_async_post_guardrail_log_success_event.py @@ -0,0 +1,211 @@ +""" +Tests for async_post_guardrail_log_success_event. + +The hook runs after post-call hooks (e.g. guardrails) so loggers see the final +response. Only CustomLogger subclasses that override the method are invoked; +base no-op is skipped. end_time is when the hook runs; llm_end_time in kwargs +is when the LLM call finished (pre-guardrail). +""" + +import os +import sys +from datetime import datetime, timezone +from typing import Any, Optional +from unittest.mock import AsyncMock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../../..")) + +import litellm +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.integrations.custom_logger import CustomLogger +from litellm.proxy._types import UserAPIKeyAuth +from litellm.types.utils import ModelResponse + + +class PostGuardrailLogger(CustomLogger): + """Logger that overrides async_post_guardrail_log_success_event.""" + + def __init__(self): + self.called = False + self.kwargs: Optional[dict] = None + self.response_obj: Optional[Any] = None + self.start_time: Optional[datetime] = None + self.end_time: Optional[datetime] = None + + async def async_post_guardrail_log_success_event( + self, kwargs, response_obj, start_time, end_time + ): + self.called = True + self.kwargs = kwargs + self.response_obj = response_obj + self.start_time = start_time + self.end_time = end_time + + +class FailingPostGuardrailLogger(CustomLogger): + """Logger that overrides and raises.""" + + async def async_post_guardrail_log_success_event( + self, kwargs, response_obj, start_time, end_time + ): + raise ValueError("callback failed") + + +@pytest.mark.asyncio +async def test_post_guardrail_log_called_with_response_and_kwargs(): + """Override is invoked with correct response and kwargs.""" + logger = PostGuardrailLogger() + response = ModelResponse(id="r1", choices=[], model="gpt-4") + user_api_key_dict = UserAPIKeyAuth(api_key="test-key") + data = {"model": "gpt-4", "messages": []} + + with patch("litellm.callbacks", [logger]): + from litellm.proxy.utils import ProxyLogging + from litellm.caching.caching import DualCache + + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + await proxy_logging.async_post_guardrail_log_success_event( + data=data, + response=response, + user_api_key_dict=user_api_key_dict, + ) + + assert logger.called is True + assert logger.response_obj is response + assert logger.kwargs is not None + assert logger.kwargs.get("model") == "gpt-4" + assert logger.kwargs.get("user_api_key_dict") is user_api_key_dict + assert logger.end_time is not None + + +@pytest.mark.asyncio +async def test_post_guardrail_log_base_custom_logger_not_invoked(): + """Base CustomLogger (no override) is not invoked; only overriders are called.""" + base_logger = CustomLogger() + overriding_logger = PostGuardrailLogger() + # Base first, then override: only overriding should be called + with patch("litellm.callbacks", [base_logger, overriding_logger]): + from litellm.proxy.utils import ProxyLogging + from litellm.caching.caching import DualCache + + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + await proxy_logging.async_post_guardrail_log_success_event( + data={"model": "gpt-4"}, + response=ModelResponse(id="r1", choices=[], model="gpt-4"), + user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), + ) + + assert overriding_logger.called is True + + +@pytest.mark.asyncio +async def test_post_guardrail_log_base_no_op_never_called_when_only_base_in_callbacks(): + """When callbacks contain only base CustomLogger (no override), the hook is never invoked.""" + with patch.object( + CustomLogger, + "async_post_guardrail_log_success_event", + new_callable=AsyncMock, + ) as mock_base: + with patch("litellm.callbacks", [CustomLogger()]): + from litellm.proxy.utils import ProxyLogging + from litellm.caching.caching import DualCache + + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + await proxy_logging.async_post_guardrail_log_success_event( + data={"model": "gpt-4"}, + response=ModelResponse(id="r1", choices=[], model="gpt-4"), + user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), + ) + mock_base.assert_not_called() + + +@pytest.mark.asyncio +async def test_post_guardrail_log_llm_end_time_in_kwargs(): + """When logging_obj has model_call_details['end_time'], kwargs get llm_end_time.""" + logger = PostGuardrailLogger() + llm_end = datetime(2025, 3, 10, 12, 0, 0, tzinfo=timezone.utc) + logging_obj = type("LoggingObj", (), {})() + logging_obj.model_call_details = {"end_time": llm_end} + + with patch("litellm.callbacks", [logger]): + from litellm.proxy.utils import ProxyLogging + from litellm.caching.caching import DualCache + + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + await proxy_logging.async_post_guardrail_log_success_event( + data={"model": "gpt-4"}, + response=ModelResponse(id="r1", choices=[], model="gpt-4"), + user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), + logging_obj=logging_obj, + ) + + assert logger.called is True + assert logger.kwargs.get("llm_end_time") == llm_end + assert logger.end_time is not None + assert logger.end_time != llm_end # end_time is "now", llm_end_time is from details + + +@pytest.mark.asyncio +async def test_post_guardrail_log_exception_in_one_callback_does_not_block_others(): + """One callback raising does not prevent others from being called.""" + failing = FailingPostGuardrailLogger() + ok = PostGuardrailLogger() + + with patch("litellm.callbacks", [failing, ok]): + from litellm.proxy.utils import ProxyLogging + from litellm.caching.caching import DualCache + + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + await proxy_logging.async_post_guardrail_log_success_event( + data={"model": "gpt-4"}, + response=ModelResponse(id="r1", choices=[], model="gpt-4"), + user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), + ) + + assert ok.called is True + + +@pytest.mark.asyncio +async def test_post_guardrail_log_guardrail_callbacks_not_invoked(): + """CustomGuardrail callbacks are not invoked by this hook.""" + logger = PostGuardrailLogger() + + class FakeGuardrail(CustomGuardrail): + async def async_post_guardrail_log_success_event( + self, kwargs, response_obj, start_time, end_time + ): + self.post_guardrail_log_called = True # would be set if we were called + + guardrail = FakeGuardrail() + guardrail.post_guardrail_log_called = False + + with patch("litellm.callbacks", [guardrail, logger]): + from litellm.proxy.utils import ProxyLogging + from litellm.caching.caching import DualCache + + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + await proxy_logging.async_post_guardrail_log_success_event( + data={"model": "gpt-4"}, + response=ModelResponse(id="r1", choices=[], model="gpt-4"), + user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), + ) + + assert logger.called is True + assert getattr(guardrail, "post_guardrail_log_called", False) is False + + +@pytest.mark.asyncio +async def test_post_guardrail_log_no_callbacks(): + """No callbacks does not raise.""" + with patch("litellm.callbacks", []): + from litellm.proxy.utils import ProxyLogging + from litellm.caching.caching import DualCache + + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + await proxy_logging.async_post_guardrail_log_success_event( + data={"model": "gpt-4"}, + response=ModelResponse(id="r1", choices=[], model="gpt-4"), + user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), + )