feat: add flag to store post-guardrail proxy response in spend logs

This commit is contained in:
Yuta Saito 2025-12-29 07:06:30 +09:00
parent c6d6ffd567
commit 955de81743
4 changed files with 431 additions and 16 deletions

View file

@ -28,6 +28,49 @@ class _ProxyDBLogger(CustomLogger):
kwargs, response_obj, start_time, end_time
)
async def async_post_call_success_hook(
self,
data: dict,
user_api_key_dict: UserAPIKeyAuth,
response: Any,
):
"""Persist spend logs after proxy guardrails mutate metadata."""
if ProxyUpdateSpend.should_store_proxy_response_in_spend_logs() is not True:
return
logging_obj = data.get("litellm_logging_obj", None)
if logging_obj is None or not hasattr(logging_obj, "model_call_details"):
verbose_proxy_logger.debug(
"async_post_call_success_hook: missing logging_obj or model_call_details"
)
return
kwargs = logging_obj.model_call_details
metadata = data.get("metadata") or data.get("litellm_metadata")
if isinstance(metadata, dict):
kwargs["metadata"] = metadata
litellm_params = kwargs.get("litellm_params", {}) or {}
litellm_params["metadata"] = metadata
kwargs["litellm_params"] = litellm_params
guardrail_info = metadata.get("standard_logging_guardrail_information")
if guardrail_info is not None:
sl_object = kwargs.get("standard_logging_object")
if isinstance(sl_object, dict):
sl_object["guardrail_information"] = guardrail_info
start_time = getattr(logging_obj, "start_time")
end_time = kwargs.get("end_time")
await self._write_proxy_response_spend_log(
kwargs=kwargs,
completion_response=response,
start_time=start_time,
end_time=end_time,
)
async def async_post_call_failure_hook(
self,
request_data: dict,
@ -168,18 +211,22 @@ class _ProxyDBLogger(CustomLogger):
end_user_id=end_user_id,
):
## UPDATE DATABASE
await proxy_logging_obj.db_spend_update_writer.update_database(
token=user_api_key,
response_cost=response_cost,
user_id=user_id,
end_user_id=end_user_id,
team_id=team_id,
kwargs=kwargs,
completion_response=completion_response,
start_time=start_time,
end_time=end_time,
org_id=org_id,
)
if (
ProxyUpdateSpend.should_store_proxy_response_in_spend_logs()
is not True
):
await proxy_logging_obj.db_spend_update_writer.update_database(
token=user_api_key,
response_cost=response_cost,
user_id=user_id,
end_user_id=end_user_id,
team_id=team_id,
kwargs=kwargs,
completion_response=completion_response,
start_time=start_time,
end_time=end_time,
org_id=org_id,
)
# update cache
asyncio.create_task(
@ -254,6 +301,65 @@ class _ProxyDBLogger(CustomLogger):
return False
return
@log_db_metrics
async def _write_proxy_response_spend_log(
self,
*,
kwargs: dict,
completion_response: Any,
start_time,
end_time,
) -> None:
from litellm.proxy.proxy_server import proxy_logging_obj
metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs)
user_id = cast(Optional[str], metadata.get("user_api_key_user_id", None))
team_id = cast(Optional[str], metadata.get("user_api_key_team_id", None))
org_id = cast(Optional[str], metadata.get("user_api_key_org_id", None))
user_api_key = metadata.get("user_api_key", None)
end_user_id = get_end_user_id_for_cost_tracking(
kwargs.get("litellm_params", {}) or {}
)
sl_object: Optional[StandardLoggingPayload] = kwargs.get(
"standard_logging_object", None
)
response_cost = (
sl_object.get("response_cost", None)
if sl_object is not None
else kwargs.get("response_cost", None)
)
if kwargs.get("cache_hit", False) is True:
response_cost = 0.0
if response_cost is None:
verbose_proxy_logger.debug(
"async_post_call_success_hook: missing response_cost, skipping db write"
)
return
if not _should_track_cost_callback(
user_api_key=user_api_key,
user_id=user_id,
team_id=team_id,
end_user_id=end_user_id,
):
return
await proxy_logging_obj.db_spend_update_writer.update_database(
token=user_api_key,
response_cost=response_cost,
user_id=user_id,
end_user_id=end_user_id,
team_id=team_id,
kwargs=kwargs,
completion_response=completion_response,
start_time=start_time,
end_time=end_time,
org_id=org_id,
)
def _should_track_cost_callback(
user_api_key: Optional[str],

View file

@ -1321,12 +1321,29 @@ def load_from_azure_key_vault(use_azure_key_vault: bool = False):
def cost_tracking():
global prisma_client
if prisma_client is not None:
litellm.logging_callback_manager.add_litellm_callback(_ProxyDBLogger())
litellm.logging_callback_manager.add_litellm_async_success_callback(
_ProxyDBLogger()
if prisma_client is None:
return
from litellm.proxy.utils import ProxyUpdateSpend
store_proxy_response = (
ProxyUpdateSpend.should_store_proxy_response_in_spend_logs()
)
disable_spend = ProxyUpdateSpend.disable_spend_updates()
if store_proxy_response is None and not disable_spend:
verbose_proxy_logger.warning(
"general_settings.spend_logs_store_proxy_response is not set. "
"Current default logs the upstream LLM response; this will change in a future release. "
"Set it to True to store the proxy-mutated response or False to keep current behavior."
)
proxy_db_logger = _ProxyDBLogger()
litellm.logging_callback_manager.add_litellm_callback(proxy_db_logger)
litellm.logging_callback_manager.add_litellm_async_success_callback(
proxy_db_logger
)
async def update_cache( # noqa: PLR0915
token: Optional[str],

View file

@ -3647,6 +3647,16 @@ class ProxyUpdateSpend:
return True
return False
@staticmethod
def should_store_proxy_response_in_spend_logs() -> Optional[bool]:
"""
Returns True if spend logs should store the proxy-modified response (after guardrails/post-processing).
False means log the upstream LLM response.
"""
from litellm.proxy.proxy_server import general_settings
return general_settings.get("spend_logs_store_proxy_response", None)
async def update_spend( # noqa: PLR0915
prisma_client: PrismaClient,

View file

@ -17,6 +17,35 @@ from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger
from litellm.types.utils import StandardLoggingPayload
@pytest.fixture(autouse=True)
def mock_proxy_logging_obj(monkeypatch):
class _SlackStub:
def __init__(self):
self.customer_spend_alert = AsyncMock()
class _DBWriterStub:
def __init__(self):
self.update_database = AsyncMock()
class _ProxyLoggingStub:
def __init__(self):
self.db_spend_update_writer = _DBWriterStub()
self.slack_alerting_instance = _SlackStub()
stub = _ProxyLoggingStub()
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_logging_obj",
stub,
raising=False,
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.update_cache",
AsyncMock(),
raising=False,
)
return stub
@pytest.mark.asyncio
async def test_async_post_call_failure_hook():
# Setup
@ -126,3 +155,256 @@ async def test_async_post_call_failure_hook_non_llm_route():
# Assert that update_database was NOT called for non-LLM routes
mock_update_database.assert_not_called()
@pytest.mark.asyncio
async def test_async_log_success_event_writes_db_when_flag_false(monkeypatch, mock_proxy_logging_obj):
logger = _ProxyDBLogger()
kwargs = {
"model": "gpt-4",
"metadata": {
"user_api_key": "sk-test",
"user_api_key_user_id": "user-1",
"user_api_key_team_id": "team-1",
"user_api_key_org_id": "org-1",
},
"litellm_params": {},
"standard_logging_object": StandardLoggingPayload(
id="req-1",
trace_id=None,
call_type="acompletion",
cache_hit=False,
stream=False,
status="success",
status_fields=None,
custom_llm_provider=None,
saved_cache_cost=0,
startTime=0,
endTime=0,
completionStartTime=0,
response_time=None,
model="gpt-4",
metadata={},
cache_key=None,
response_cost=0.123,
cost_breakdown=None,
total_tokens=0,
prompt_tokens=0,
completion_tokens=0,
request_tags=None,
end_user="",
api_base="",
model_group=None,
model_id=None,
requester_ip_address=None,
messages=None,
response=None,
model_parameters=None,
hidden_params={},
model_map_information=None,
error_str=None,
error_information=None,
response_cost_failure_debug_info=None,
guardrail_information=None,
standard_built_in_tools_params=None,
),
}
monkeypatch.setattr(
"litellm.proxy.utils.ProxyUpdateSpend.should_store_proxy_response_in_spend_logs",
lambda: False,
)
await logger.async_log_success_event(
kwargs=kwargs,
response_obj=None,
start_time=datetime.now(),
end_time=datetime.now(),
)
mock_proxy_logging_obj.db_spend_update_writer.update_database.assert_awaited_once()
@pytest.mark.asyncio
async def test_async_log_success_event_skips_db_when_flag_true(monkeypatch, mock_proxy_logging_obj):
logger = _ProxyDBLogger()
kwargs = {
"model": "gpt-4",
"metadata": {
"user_api_key": "sk-test",
"user_api_key_user_id": "user-1",
},
"litellm_params": {},
"standard_logging_object": StandardLoggingPayload(
id="req-1",
trace_id=None,
call_type="acompletion",
cache_hit=False,
stream=False,
status="success",
status_fields=None,
custom_llm_provider=None,
saved_cache_cost=0,
startTime=0,
endTime=0,
completionStartTime=0,
response_time=None,
model="gpt-4",
metadata={},
cache_key=None,
response_cost=0.5,
cost_breakdown=None,
total_tokens=0,
prompt_tokens=0,
completion_tokens=0,
request_tags=None,
end_user="",
api_base="",
model_group=None,
model_id=None,
requester_ip_address=None,
messages=None,
response=None,
model_parameters=None,
hidden_params={},
model_map_information=None,
error_str=None,
error_information=None,
response_cost_failure_debug_info=None,
guardrail_information=None,
standard_built_in_tools_params=None,
),
}
monkeypatch.setattr(
"litellm.proxy.utils.ProxyUpdateSpend.should_store_proxy_response_in_spend_logs",
lambda: True,
)
await logger.async_log_success_event(
kwargs=kwargs,
response_obj=None,
start_time=datetime.now(),
end_time=datetime.now(),
)
mock_proxy_logging_obj.db_spend_update_writer.update_database.assert_not_called()
@pytest.mark.asyncio
async def test_async_log_success_event_with_flag_none(monkeypatch, mock_proxy_logging_obj):
logger = _ProxyDBLogger()
kwargs = {
"model": "gpt-4",
"metadata": {
"user_api_key": "sk-test",
"user_api_key_user_id": "user-1",
"user_api_key_team_id": "team-1",
"user_api_key_org_id": "org-1",
},
"litellm_params": {},
"standard_logging_object": StandardLoggingPayload(
id="req-flag-none",
trace_id=None,
call_type="acompletion",
cache_hit=False,
stream=False,
status="success",
status_fields=None,
custom_llm_provider=None,
saved_cache_cost=0,
startTime=0,
endTime=0,
completionStartTime=0,
response_time=None,
model="gpt-4",
metadata={},
cache_key=None,
response_cost=0.321,
cost_breakdown=None,
total_tokens=0,
prompt_tokens=0,
completion_tokens=0,
request_tags=None,
end_user="",
api_base="",
model_group=None,
model_id=None,
requester_ip_address=None,
messages=None,
response=None,
model_parameters=None,
hidden_params={},
model_map_information=None,
error_str=None,
error_information=None,
response_cost_failure_debug_info=None,
guardrail_information=None,
standard_built_in_tools_params=None,
),
}
monkeypatch.setattr(
"litellm.proxy.utils.ProxyUpdateSpend.should_store_proxy_response_in_spend_logs",
lambda: None,
)
await logger.async_log_success_event(
kwargs=kwargs,
response_obj=None,
start_time=datetime.now(),
end_time=datetime.now(),
)
mock_proxy_logging_obj.db_spend_update_writer.update_database.assert_awaited_once()
@pytest.mark.asyncio
async def test_async_post_call_success_hook_writes_db_with_guardrail_info(
monkeypatch, mock_proxy_logging_obj
):
logger = _ProxyDBLogger()
class _LoggingObj:
def __init__(self):
self.model_call_details = {
"metadata": {
"user_api_key": "sk-test",
"user_api_key_user_id": "user-1",
"user_api_key_team_id": "team-1",
"user_api_key_org_id": "org-1",
},
"litellm_params": {},
"standard_logging_object": {
"response_cost": 0.25,
},
"end_time": datetime.now(),
}
self.start_time = datetime.now()
data = {
"litellm_logging_obj": _LoggingObj(),
"metadata": {
"standard_logging_guardrail_information": [
{"guardrail_name": "noma", "status": "success"}
],
},
}
monkeypatch.setattr(
"litellm.proxy.utils.ProxyUpdateSpend.should_store_proxy_response_in_spend_logs",
lambda: True,
)
await logger.async_post_call_success_hook(
data=data,
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
response=None,
)
mock_proxy_logging_obj.db_spend_update_writer.update_database.assert_awaited_once()
kwargs = mock_proxy_logging_obj.db_spend_update_writer.update_database.call_args.kwargs
assert kwargs["kwargs"]["metadata"]["standard_logging_guardrail_information"]