diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 1b6036fef8..b943323927 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -1747,6 +1747,11 @@ class CustomStreamWrapper: if is_empty: continue print_verbose(f"final returned processed chunk: {processed_chunk}") + + # add usage as hidden param + if self.sent_last_chunk is True and self.stream_options is None: + usage = calculate_total_usage(chunks=self.chunks) + processed_chunk._hidden_params["usage"] = usage return processed_chunk raise StopAsyncIteration else: # temporary patch for non-aiohttp async calls @@ -1790,6 +1795,7 @@ class CustomStreamWrapper: messages=self.messages, logging_obj=self.logging_obj, ) + response = self.model_response_creator() if complete_streaming_response is not None: setattr( diff --git a/tests/local_testing/test_custom_callback_input.py b/tests/local_testing/test_custom_callback_input.py index 9cc41869fe..5da87d2be6 100644 --- a/tests/local_testing/test_custom_callback_input.py +++ b/tests/local_testing/test_custom_callback_input.py @@ -1599,7 +1599,7 @@ def test_logging_key_masking_gemini(): assert "PART" == trimmed_key -@pytest.mark.parametrize("sync_mode", [True]) +@pytest.mark.parametrize("sync_mode", [True, False]) @pytest.mark.asyncio async def test_standard_logging_payload_stream_usage(sync_mode): """