mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-20 00:11:50 +00:00
feat: add flag to store post-guardrail proxy response in spend logs
This commit is contained in:
parent
c6d6ffd567
commit
955de81743
4 changed files with 431 additions and 16 deletions
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue