mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
2405e0d400
commit
7ac26f23d8
5 changed files with 376 additions and 27 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
)
|
||||
Loading…
Add table
Reference in a new issue