diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 6abfca1d3a0..418c3900533 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -334,22 +334,52 @@ class _ProxyDBLogger(CustomLogger): "that no poll task will settle" ) return - await _release_budget_reservation(budget_reservation=budget_reservation) # Non-model call types (health checks, afile_delete) have no model or standard_logging_object. # Use .get() for "stream" to avoid KeyError on health checks. # WS session wrappers (_aresponses_websocket, _arealtime) also reach here with # result=None; their per-turn costs are tracked on the inner aresponses/realtime calls. - if sl_object is None and ( + skippable_non_model_call = sl_object is None and ( not kwargs.get("model") or kwargs.get("call_type") in ("_aresponses_websocket", "_arealtime") - ): + ) + completed_call = kwargs.get("stream") is not True or ( + kwargs.get("stream") is True + and ("complete_streaming_response" in kwargs or "async_complete_streaming_response" in kwargs) + ) + if skippable_non_model_call: + await _release_budget_reservation(budget_reservation=budget_reservation) verbose_proxy_logger.warning( "Cost tracking - skipping, no standard_logging_object for call_type=%s", kwargs.get("call_type", "unknown"), ) return - if kwargs.get("stream") is not True or ( - kwargs.get("stream") is True and "complete_streaming_response" in kwargs - ): + if completed_call: + # Releasing to $0 treats the call as free. Leaving the hold + # open is also wrong: the next priced request only + # reconciles its own reservation, so this one would keep + # blocking shared counters until TTL. Settle at the + # admission estimate instead. No spend-log row — there is + # no real cost to write. + reserved_cost = float(budget_reservation.get("reserved_cost") or 0.0) if budget_reservation else 0.0 + try: + await _reconcile_budget_reservation( + budget_reservation=budget_reservation, + actual_cost=reserved_cost, + ) + except Exception: # noqa: BLE001 # settle can fail on cache/redis; still raise cost-tracking after invalidating + verbose_proxy_logger.exception( + "Failed to settle budget reservation after unpriced successful call" + ) + try: + await _invalidate_budget_reservation_counters( + budget_reservation=budget_reservation, + ) + except Exception: # noqa: BLE001 # invalidate is best-effort so the outer cost-tracking error still surfaces + verbose_proxy_logger.exception( + "Failed to invalidate budget reservation counters after settle failed" + ) + finally: + if budget_reservation is not None: + budget_reservation["finalized"] = True if sl_object is not None: cost_tracking_failure_debug_info: dict | str = ( sl_object["response_cost_failure_debug_info"] @@ -361,6 +391,7 @@ class _ProxyDBLogger(CustomLogger): raise Exception( f"Cost tracking failed for model={model}.\nDebug info - {cost_tracking_failure_debug_info}\nAdd custom pricing - https://docs.litellm.ai/docs/proxy/custom_pricing" ) + await _release_budget_reservation(budget_reservation=budget_reservation) except Exception as e: error_msg = f"Error in tracking cost callback - {e}\n Traceback:{traceback.format_exc()}" model = kwargs.get("model", "") @@ -634,6 +665,23 @@ async def _release_budget_reservation(budget_reservation: dict | None) -> None: ) +async def _reconcile_budget_reservation( + budget_reservation: dict | None, # mutable-ok: same reservation payload _release_budget_reservation takes + actual_cost: float, +) -> None: + if budget_reservation is None: + return + + from litellm.proxy.spend_tracking.budget_reservation import ( + reconcile_budget_reservation, + ) + + await reconcile_budget_reservation( + budget_reservation=budget_reservation, + actual_cost=actual_cost, + ) + + async def _invalidate_budget_reservation_counters( budget_reservation: dict | None, ) -> None: diff --git a/test-quality-budget.json b/test-quality-budget.json index 4a7bc7edff2..0350d813743 100644 --- a/test-quality-budget.json +++ b/test-quality-budget.json @@ -3,7 +3,7 @@ "limit": 744 }, "TQ002": { - "limit": 742 + "limit": 741 }, "TQ003": { "limit": 62 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 ca517474a5c..96b2bf15bb8 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 @@ -1,4 +1,3 @@ - import pytest @@ -70,9 +69,7 @@ async def test_async_post_call_failure_hook(): # Check that metadata was properly updated assert "litellm_params" in call_args["kwargs"] - assert call_args["kwargs"]["litellm_params"]["proxy_server_request"] == { - "request_id": "test_request_id" - } + assert call_args["kwargs"]["litellm_params"]["proxy_server_request"] == {"request_id": "test_request_id"} metadata = call_args["kwargs"]["litellm_params"]["metadata"] assert metadata["user_api_key"] == "test_api_key" assert metadata["status"] == "failure" @@ -336,9 +333,7 @@ async def test_should_continue_failure_tracking_when_budget_release_fails(): ) assert mock_invalidate_budget_reservation_counters.await_count == 1 assert ( - mock_invalidate_budget_reservation_counters.await_args.kwargs[ - "budget_reservation" - ] + mock_invalidate_budget_reservation_counters.await_args.kwargs["budget_reservation"] is user_api_key_dict.budget_reservation ) assert user_api_key_dict.budget_reservation["finalized"] is True @@ -383,7 +378,13 @@ async def test_track_cost_callback_releases_budget_reservation_when_spend_tracki @pytest.mark.asyncio -async def test_track_cost_callback_releases_budget_reservation_when_response_cost_missing(): +async def test_track_cost_callback_settles_budget_reservation_when_response_cost_missing(): + """A successful unpriced model call must not be refunded to $0. + + Settling at the admission estimate converts the hold into budget spend + without inventing a spend-log row. Health checks still release, covered + separately. + """ logger = _ProxyDBLogger() budget_reservation = {"reserved_cost": 0.5, "entries": []} user_api_key_auth = UserAPIKeyAuth(budget_reservation=budget_reservation) @@ -405,13 +406,21 @@ async def test_track_cost_callback_releases_budget_reservation_when_response_cos } with ( - patch( + patch( # test-quality-ok: callback reads proxy_logging_obj from the module; no injection seam "litellm.proxy.proxy_server.proxy_logging_obj", ) as mock_proxy_logging, - patch( + patch( # test-quality-ok: assert the hold is settled, not refunded through release_budget_reservation "litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation", new_callable=AsyncMock, ) as mock_release_budget_reservation, + patch( # test-quality-ok: settle is a proxy-internal reservation call, not an HTTP boundary + "litellm.proxy.spend_tracking.budget_reservation.reconcile_budget_reservation", + new_callable=AsyncMock, + ) as mock_reconcile_budget_reservation, + patch( # test-quality-ok: unpriced settle must not write a spend-log row + "litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database", + new_callable=AsyncMock, + ) as mock_update_database, ): mock_proxy_logging.failed_tracking_alert = AsyncMock() @@ -422,9 +431,175 @@ async def test_track_cost_callback_releases_budget_reservation_when_response_cos end_time=datetime.now(), ) + mock_release_budget_reservation.assert_not_awaited() + mock_reconcile_budget_reservation.assert_awaited_once_with( + budget_reservation=budget_reservation, + actual_cost=0.5, + ) + mock_update_database.assert_not_called() + mock_proxy_logging.failed_tracking_alert.assert_called() + # The hold is converted, not dropped: later traffic still sees the reserved cost. + assert budget_reservation["reserved_cost"] == 0.5 + assert budget_reservation.get("finalized") is not True + + +@pytest.mark.asyncio +async def test_track_cost_callback_settles_async_stream_when_response_cost_missing(): + """Async streams record completion on async_complete_streaming_response.""" + logger = _ProxyDBLogger() + budget_reservation = {"reserved_cost": 0.5, "entries": []} + user_api_key_auth = UserAPIKeyAuth(budget_reservation=budget_reservation) + + kwargs = { + "model": "gpt-4", + "call_type": "acompletion", + "stream": True, + "async_complete_streaming_response": {"usage": {"total_tokens": 10}}, + "litellm_params": { + "metadata": { + "user_api_key_auth": user_api_key_auth, + }, + }, + "standard_logging_object": { + "response_cost": None, + "response_cost_failure_debug_info": "missing custom price", + "request_tags": None, + }, + } + + with ( + patch( # test-quality-ok: callback reads proxy_logging_obj from the module; no injection seam + "litellm.proxy.proxy_server.proxy_logging_obj", + ) as mock_proxy_logging, + patch( # test-quality-ok: async-complete streams must settle, not release + "litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation", + new_callable=AsyncMock, + ) as mock_release_budget_reservation, + patch( # test-quality-ok: settle is a proxy-internal reservation call, not an HTTP boundary + "litellm.proxy.spend_tracking.budget_reservation.reconcile_budget_reservation", + new_callable=AsyncMock, + ) as mock_reconcile_budget_reservation, + patch( # test-quality-ok: unpriced settle must not write a spend-log row + "litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database", + new_callable=AsyncMock, + ) as mock_update_database, + ): + mock_proxy_logging.failed_tracking_alert = AsyncMock() + + await logger._PROXY_track_cost_callback( + kwargs=kwargs, + completion_response=None, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + mock_release_budget_reservation.assert_not_awaited() + mock_reconcile_budget_reservation.assert_awaited_once_with( + budget_reservation=budget_reservation, + actual_cost=0.5, + ) + mock_update_database.assert_not_called() + mock_proxy_logging.failed_tracking_alert.assert_called() + assert budget_reservation["reserved_cost"] == 0.5 + assert budget_reservation.get("finalized") is not True + + +@pytest.mark.asyncio +async def test_track_cost_callback_invalidates_reservation_when_settle_fails(): + """A failed settle must not leave the hold pinning later traffic until TTL.""" + logger = _ProxyDBLogger() + budget_reservation = {"reserved_cost": 0.5, "entries": []} + user_api_key_auth = UserAPIKeyAuth(budget_reservation=budget_reservation) + + kwargs = { + "model": "gpt-4", + "call_type": "acompletion", + "litellm_params": { + "metadata": { + "user_api_key_auth": user_api_key_auth, + }, + }, + "standard_logging_object": { + "response_cost": None, + "response_cost_failure_debug_info": "missing custom price", + "request_tags": None, + }, + "stream": False, + } + + with ( + patch( # test-quality-ok: callback reads proxy_logging_obj from the module; no injection seam + "litellm.proxy.proxy_server.proxy_logging_obj", + ) as mock_proxy_logging, + patch( # test-quality-ok: a failed settle must not refund through release_budget_reservation + "litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation", + new_callable=AsyncMock, + ) as mock_release_budget_reservation, + patch( # test-quality-ok: force reconcile to fail so the invalidate path can be observed + "litellm.proxy.spend_tracking.budget_reservation.reconcile_budget_reservation", + new_callable=AsyncMock, + side_effect=RuntimeError("redis down"), + ), + patch( # test-quality-ok: invalidate is the only way to unpin counters after settle fails + "litellm.proxy.spend_tracking.budget_reservation.invalidate_budget_reservation_counters", + new_callable=AsyncMock, + ) as mock_invalidate_budget_reservation_counters, + patch( # test-quality-ok: failed settle must not write a spend-log row + "litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database", + new_callable=AsyncMock, + ) as mock_update_database, + ): + mock_proxy_logging.failed_tracking_alert = AsyncMock() + + await logger._PROXY_track_cost_callback( + kwargs=kwargs, + completion_response=None, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + mock_release_budget_reservation.assert_not_awaited() + mock_invalidate_budget_reservation_counters.assert_awaited_once() + settled = mock_invalidate_budget_reservation_counters.await_args.kwargs["budget_reservation"] + assert settled["reserved_cost"] == 0.5 + assert settled["finalized"] is True + mock_update_database.assert_not_called() + mock_proxy_logging.failed_tracking_alert.assert_called() + + +@pytest.mark.asyncio +async def test_track_cost_callback_releases_budget_reservation_for_non_model_calls(): + logger = _ProxyDBLogger() + budget_reservation = {"reserved_cost": 0.5, "entries": []} + user_api_key_auth = UserAPIKeyAuth(budget_reservation=budget_reservation) + + kwargs = { + "call_type": "health", + "litellm_params": { + "metadata": { + "user_api_key_auth": user_api_key_auth, + }, + }, + "stream": False, + } + + with patch( # test-quality-ok: health checks have no cost row; release is the observable contract + "litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation", + new_callable=AsyncMock, + ) as mock_release_budget_reservation: + await logger._PROXY_track_cost_callback( + kwargs=kwargs, + completion_response=None, + start_time=datetime.now(), + end_time=datetime.now(), + ) + mock_release_budget_reservation.assert_awaited_once_with( budget_reservation=budget_reservation, ) + # Health checks refund the hold; they must not stamp it settled. + assert budget_reservation.get("finalized") is not True + assert budget_reservation["reserved_cost"] == 0.5 def test_get_budget_reservation_from_metadata_handles_dict_auth_object(): @@ -433,36 +608,21 @@ def test_get_budget_reservation_from_metadata_handles_dict_auth_object(): "entries": [{"counter_key": "spend:key:test_api_key"}], } + assert _get_budget_reservation_from_metadata(metadata={"user_api_key_auth": dict(UserAPIKeyAuth())}) is None assert ( _get_budget_reservation_from_metadata( - metadata={"user_api_key_auth": dict(UserAPIKeyAuth())} - ) - is None - ) - assert ( - _get_budget_reservation_from_metadata( - metadata={ - "user_api_key_auth": UserAPIKeyAuth( - budget_reservation=budget_reservation - ) - } + metadata={"user_api_key_auth": UserAPIKeyAuth(budget_reservation=budget_reservation)} ) == budget_reservation ) assert ( _get_budget_reservation_from_metadata( - metadata={ - "user_api_key_auth": dict( - UserAPIKeyAuth(budget_reservation=budget_reservation) - ) - } + metadata={"user_api_key_auth": dict(UserAPIKeyAuth(budget_reservation=budget_reservation))} ) == budget_reservation ) assert ( - _get_budget_reservation_from_metadata( - metadata={"user_api_key_budget_reservation": budget_reservation} - ) + _get_budget_reservation_from_metadata(metadata={"user_api_key_budget_reservation": budget_reservation}) is budget_reservation ) @@ -470,9 +630,7 @@ def test_get_budget_reservation_from_metadata_handles_dict_auth_object(): @pytest.mark.asyncio async def test_update_database_and_spend_counters_releases_reservation_when_db_update_fails(): proxy_logging_obj = MagicMock() - proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock( - side_effect=Exception("db unavailable") - ) + proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock(side_effect=Exception("db unavailable")) increment_spend_counters = AsyncMock() budget_reservation = {"reserved_cost": 0.5, "entries": []} @@ -508,9 +666,7 @@ async def test_update_database_and_spend_counters_releases_reservation_when_db_u async def test_update_database_and_spend_counters_preserves_db_exception_when_release_fails(): proxy_logging_obj = MagicMock() db_exception = RuntimeError("db unavailable") - proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock( - side_effect=db_exception - ) + proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock(side_effect=db_exception) increment_spend_counters = AsyncMock() budget_reservation = {"reserved_cost": 0.5, "entries": []} @@ -554,12 +710,8 @@ async def test_update_database_and_spend_counters_preserves_db_exception_when_re budget_reservation=budget_reservation, ) assert mock_log_exception.call_count == 2 - mock_log_exception.assert_any_call( - "Failed to release budget reservation after database update failed" - ) - mock_log_exception.assert_any_call( - "Failed to invalidate budget reservation counters after release failed" - ) + mock_log_exception.assert_any_call("Failed to release budget reservation after database update failed") + mock_log_exception.assert_any_call("Failed to invalidate budget reservation counters after release failed") increment_spend_counters.assert_not_awaited() @@ -1097,10 +1249,7 @@ async def test_async_post_call_failure_hook_propagates_trace_id_from_logging_obj # standard_logging_object should have been propagated from logging obj assert call_kwargs.get("standard_logging_object") is not None - assert ( - call_kwargs["standard_logging_object"]["trace_id"] - == "trace-id-from-logging-obj" - ) + assert call_kwargs["standard_logging_object"]["trace_id"] == "trace-id-from-logging-obj" # litellm_trace_id should also be propagated as a fallback assert call_kwargs.get("litellm_trace_id") == "trace-id-from-logging-obj" @@ -1687,9 +1836,7 @@ async def test_async_post_call_failure_hook_records_recovered_partial_spend(): "metadata": {}, "proxy_server_request": {"request_id": "rid"}, "response_cost": 3.5e-05, - "combined_usage_object": Usage( - prompt_tokens=30, completion_tokens=1, total_tokens=31 - ), + "combined_usage_object": Usage(prompt_tokens=30, completion_tokens=1, total_tokens=31), } with patch( @@ -1768,15 +1915,10 @@ async def test_track_cost_callback_enriches_user_id_for_mcp_style_metadata(): assert mock_increment.call_args.kwargs["team_id"] == "team-123" assert mock_increment.call_args.kwargs["org_id"] == "org-456" - update_kwargs = ( - mock_proxy_logging.db_spend_update_writer.update_database.await_args.kwargs - ) + update_kwargs = mock_proxy_logging.db_spend_update_writer.update_database.await_args.kwargs assert update_kwargs["user_id"] == "mcp-user@example.com" assert update_kwargs["team_id"] == "team-123" - assert ( - kwargs["litellm_params"]["metadata"]["user_api_key_user_id"] - == "mcp-user@example.com" - ) + assert kwargs["litellm_params"]["metadata"]["user_api_key_user_id"] == "mcp-user@example.com" @pytest.mark.parametrize( @@ -1824,9 +1966,7 @@ def test_should_track_cost_callback_pass_through_without_owner(call_type, expect ], ) @pytest.mark.asyncio -async def test_track_cost_callback_logs_unauthenticated_pass_through_request( - call_type, expect_spend_log -): +async def test_track_cost_callback_logs_unauthenticated_pass_through_request(call_type, expect_spend_log): """Regression for LIT-3782: a pass-through request with auth=false reaches the cost callback with no key/user/team/end-user. Before the fix the spend-log write was skipped and the request never appeared in request/usage logs. It @@ -1872,6 +2012,4 @@ async def test_track_cost_callback_logs_unauthenticated_pass_through_request( end_time=datetime.now(), ) - assert mock_proxy_logging.db_spend_update_writer.update_database.await_count == ( - 1 if expect_spend_log else 0 - ) + assert mock_proxy_logging.db_spend_update_writer.update_database.await_count == (1 if expect_spend_log else 0)