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
This commit is contained in:
devarakondasrikanth 2026-03-13 19:36:52 -07:00
parent 2405e0d400
commit 7ac26f23d8
5 changed files with 376 additions and 27 deletions

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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"),
)