This commit is contained in:
yyobject 2026-10-03 20:08:32 +08:00 • committed by GitHub
commit 9639e90cc8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 80 additions and 14 deletions

View file

@ -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"]

View file

@ -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',
},
}
],