diff --git a/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py b/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py index 9ea730a873f..0d4bb9c8351 100644 --- a/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py +++ b/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py @@ -231,15 +231,7 @@ async def convert_to_streaming_response_async( model_response_object.choices = choice_list if "usage" in response_object and response_object["usage"] is not None: - setattr( - model_response_object, - "usage", - Usage( - completion_tokens=response_object["usage"].get("completion_tokens", 0), - prompt_tokens=response_object["usage"].get("prompt_tokens", 0), - total_tokens=response_object["usage"].get("total_tokens", 0), - ), - ) + setattr(model_response_object, "usage", Usage(**response_object["usage"])) if "id" in response_object: model_response_object.id = response_object["id"] @@ -325,10 +317,7 @@ def convert_to_streaming_response( model_response_object.choices = choice_list if "usage" in response_object and response_object["usage"] is not None: - setattr(model_response_object, "usage", Usage()) - model_response_object.usage.completion_tokens = response_object["usage"].get("completion_tokens", 0) - model_response_object.usage.prompt_tokens = response_object["usage"].get("prompt_tokens", 0) - model_response_object.usage.total_tokens = response_object["usage"].get("total_tokens", 0) + setattr(model_response_object, "usage", Usage(**response_object["usage"])) if "id" in response_object: model_response_object.id = response_object["id"] diff --git a/tests/unit/litellm_core_utils/llm_response_utils/test_convert_dict_to_response.py b/tests/unit/litellm_core_utils/llm_response_utils/test_convert_dict_to_response.py index 8e46ae21de6..557f3233dd1 100644 --- a/tests/unit/litellm_core_utils/llm_response_utils/test_convert_dict_to_response.py +++ b/tests/unit/litellm_core_utils/llm_response_utils/test_convert_dict_to_response.py @@ -1,12 +1,15 @@ from typing import Final import pytest +from pydantic import JsonValue from litellm.constants import RESPONSE_FORMAT_TOOL_NAME from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( _handle_invalid_parallel_tool_calls, _should_convert_tool_call_to_json_mode, convert_to_model_response_object, + convert_to_streaming_response, + convert_to_streaming_response_async, ) from litellm.types.utils import ( ChatCompletionMessageCustomToolCall, @@ -15,6 +18,80 @@ from litellm.types.utils import ( ModelResponse, ) + +@pytest.mark.parametrize("use_async", (False, True), ids=("sync", "async")) +@pytest.mark.parametrize("content", ("done", "cached response text"), ids=("single", "sliced")) +@pytest.mark.asyncio +async def test_cached_streaming_replay_preserves_complete_usage(content: str, use_async: bool) -> None: + usage: Final = { + "prompt_tokens": 100, + "completion_tokens": 40, + "total_tokens": 140, + "prompt_tokens_details": {"cached_tokens": 60, "audio_tokens": 5}, + "completion_tokens_details": { + "reasoning_tokens": 20, + "audio_tokens": 3, + "accepted_prediction_tokens": 2, + "rejected_prediction_tokens": 1, + }, + "server_tool_use": {"web_search_requests": 1}, + "provider_usage": {"custom_units": 2}, + } + response: Final = { + "id": "chatcmpl-cache-regression", + "created": 1, + "model": "cache-test-model", + "choices": [{"message": {"role": "assistant", "content": content}, "finish_reason": "stop"}], + "usage": usage, + } + chunks: Final = ( + [chunk async for chunk in convert_to_streaming_response_async(response)] + if use_async + else list(convert_to_streaming_response(response)) + ) + + assert "".join(chunk.choices[0].delta.content for chunk in chunks) == content + assert chunks[-1].model_dump(exclude_none=True)["usage"] == usage + assert all(getattr(chunk, "usage", None) is None for chunk in chunks[:-1]) + assert chunks[-1].choices[0].finish_reason == "stop" + + +@pytest.mark.parametrize("use_async", (False, True), ids=("sync", "async")) +@pytest.mark.parametrize( + ("usage_fields", "expected_usage"), + ( + ({}, None), + ({"usage": None}, None), + ({"usage": {}}, {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}), + ({"usage": {"prompt_tokens": 5}}, {"prompt_tokens": 5, "completion_tokens": 0, "total_tokens": 0}), + ( + {"usage": {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8}}, + {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8}, + ), + ), + ids=("missing", "null", "empty", "partial", "totals-only"), +) +@pytest.mark.asyncio +async def test_cached_streaming_replay_preserves_legacy_usage( + usage_fields: dict[str, JsonValue], expected_usage: dict[str, int] | None, use_async: bool +) -> None: + response: Final = { + "id": "chatcmpl-cache-regression", + "created": 1, + "model": "cache-test-model", + "choices": [{"message": {"role": "assistant", "content": "done"}, "finish_reason": "stop"}], + **usage_fields, + } + chunks: Final = ( + [chunk async for chunk in convert_to_streaming_response_async(response)] + if use_async + else list(convert_to_streaming_response(response)) + ) + + assert len(chunks) == 1 + assert chunks[0].model_dump(exclude_none=True).get("usage") == expected_usage + + OPENAI_CUSTOM_TOOL_CALL_RESPONSE = { "id": "chatcmpl-abc", "created": 1784657740, @@ -33,7 +110,7 @@ OPENAI_CUSTOM_TOOL_CALL_RESPONSE = { "type": "custom", "custom": { "name": "ApplyPatch", - "input": "*** Begin Patch\n*** Update File: main.py\n@@\n+def hello():\n+ print(\"Hello\")\n*** End Patch\n", + "input": '*** Begin Patch\n*** Update File: main.py\n@@\n+def hello():\n+ print("Hello")\n*** End Patch\n', }, } ],