diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index dab5fb1bfd5..b6d661c454a 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -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], diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index f56c0c2b07a..4d812498fe7 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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], diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index d595db4a2e0..2cdbae4df6a 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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, diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index cb6d90103f7..373a604d410 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -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"]