mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(caching): preserve usage details on streaming cache hits
This commit is contained in:
parent
0d45883312
commit
4a2682c202
2 changed files with 80 additions and 14 deletions
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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',
|
||||
},
|
||||
}
|
||||
],
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue