diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index db23929e0c3..6a5a8832cc6 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -236,9 +236,11 @@ class CustomStreamWrapper: stream_options=None, make_call: Callable | None = None, _response_headers: dict | httpx.Headers | None = None, + count_prompt_tokens: Callable[[], int] | None = None, ): self.model = model self.make_call = make_call + self.count_prompt_tokens = count_prompt_tokens self.custom_llm_provider = custom_llm_provider self.logging_obj: LiteLLMLoggingObject = logging_obj self.completion_stream = completion_stream @@ -1641,7 +1643,7 @@ class CustomStreamWrapper: except Exception: model_response.choices[0].delta = Delta() else: - if self.stream_options is not None and self.stream_options["include_usage"] is True: + if self.send_stream_usage is True: model_response.choices = [] return model_response self._record_usage_only_chunk(model_response=model_response) @@ -1996,6 +1998,7 @@ class CustomStreamWrapper: chunks=self.chunks, messages=self.messages, logging_obj=self.logging_obj, + count_prompt_tokens=self.count_prompt_tokens, ) except Exception as e: # stream_chunk_builder can re-raise (as APIError) on large agentic @@ -2248,6 +2251,7 @@ class CustomStreamWrapper: chunks=self.chunks, messages=self.messages, logging_obj=self.logging_obj, + count_prompt_tokens=self.count_prompt_tokens, ) except Exception as e: # see sync __next__: a raise from stream_chunk_builder inside this @@ -2371,6 +2375,7 @@ class CustomStreamWrapper: chunks=self.chunks, messages=self.messages if isinstance(self.messages, list) else None, logging_obj=self.logging_obj, + count_prompt_tokens=self.count_prompt_tokens, ) if partial_response is None: return diff --git a/litellm/main.py b/litellm/main.py index 10fb32828f1..17edafcdfca 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -850,6 +850,12 @@ def admission_input_tokens(kwargs: Mapping[str, object]) -> int | None: ) +def admitted_prompt_token_counter(prompt_tokens: int | None) -> Callable[[], int] | None: + if prompt_tokens is None: + return None + return lambda: prompt_tokens + + def mock_completion( model: str, messages: list, @@ -935,6 +941,7 @@ def mock_completion( if stream is True: model_response = ModelResponseStream() + count_prompt_tokens: Final = admitted_prompt_token_counter(prompt_tokens) # don't try to access stream object, if kwargs.get("acompletion", False) is True: return CustomStreamWrapper( @@ -944,6 +951,7 @@ def mock_completion( model=model, custom_llm_provider="openai", logging_obj=logging, + count_prompt_tokens=count_prompt_tokens, ) return CustomStreamWrapper( completion_stream=mock_completion_streaming_obj( @@ -952,6 +960,7 @@ def mock_completion( model=model, custom_llm_provider="openai", logging_obj=logging, + count_prompt_tokens=count_prompt_tokens, ) if isinstance(mock_response, litellm.MockException): raise mock_response diff --git a/tests/test_litellm/test_main.py b/tests/test_litellm/test_main.py index a36ca229981..f71225c6fc5 100644 --- a/tests/test_litellm/test_main.py +++ b/tests/test_litellm/test_main.py @@ -2422,9 +2422,13 @@ def test_mock_completion_usage_falls_back_to_default_without_admission_count(): _ADMISSION_INPUT_TOKENS: Final = 51234 -_ADMISSION_METADATA: Final = { - "user_api_key_budget_reservation": {"reserved_cost": 1.0, "input_tokens": _ADMISSION_INPUT_TOKENS} -} + + +def _admission_metadata(input_tokens: int) -> dict[str, object]: # mutable-ok: logging writes into metadata + return {"user_api_key_budget_reservation": {"reserved_cost": 1.0, "input_tokens": input_tokens}} + + +_ADMISSION_METADATA: Final = _admission_metadata(_ADMISSION_INPUT_TOKENS) _MOCK_STREAM_MESSAGES: Final = [{"role": "user", "content": "hello " * 200}] _STREAM_CHUNK_BUILDER_TOKEN_COUNTER: Final = "litellm.litellm_core_utils.streaming_chunk_builder_utils.token_counter" @@ -2440,7 +2444,7 @@ def _client_usage_chunks(chunks: list[ModelResponseStream]) -> list[Usage]: @pytest.mark.parametrize("n", (None, 2)) def test_mock_completion_stream_usage_reports_admission_input_tokens_without_tokenizer_fallback(n: int | None): with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: - chunks = list( + chunks: Final = list( litellm.completion( model="openai/gpt-5.4-mini", messages=_MOCK_STREAM_MESSAGES, @@ -2453,7 +2457,7 @@ def test_mock_completion_stream_usage_reports_admission_input_tokens_without_tok ) ) - usage_chunks = _client_usage_chunks(chunks) + usage_chunks: Final = _client_usage_chunks(chunks) assert len(usage_chunks) == 1 assert usage_chunks[0].prompt_tokens == _ADMISSION_INPUT_TOKENS assert usage_chunks[0].completion_tokens == litellm_main.DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT @@ -2469,7 +2473,7 @@ async def test_mock_acompletion_stream_usage_reports_admission_input_tokens_with n: int | None, ): with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: - response = await litellm.acompletion( + response: Final = await litellm.acompletion( model="openai/gpt-5.4-mini", messages=_MOCK_STREAM_MESSAGES, mock_response="ok", @@ -2479,9 +2483,9 @@ async def test_mock_acompletion_stream_usage_reports_admission_input_tokens_with stream_options={"include_usage": True}, litellm_metadata=_ADMISSION_METADATA, ) - chunks = [chunk async for chunk in response] + chunks: Final = [chunk async for chunk in response] - usage_chunks = _client_usage_chunks(chunks) + usage_chunks: Final = _client_usage_chunks(chunks) assert len(usage_chunks) == 1 assert usage_chunks[0].prompt_tokens == _ADMISSION_INPUT_TOKENS assert usage_chunks[0].total_tokens == _ADMISSION_INPUT_TOKENS + usage_chunks[0].completion_tokens @@ -2492,7 +2496,7 @@ async def test_mock_acompletion_stream_usage_reports_admission_input_tokens_with def test_mock_completion_stream_without_include_usage_hides_usage_chunk_but_logs_admission_count(): with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: - chunks = list( + chunks: Final = list( litellm.completion( model="openai/gpt-5.4-mini", messages=_MOCK_STREAM_MESSAGES, @@ -2509,10 +2513,48 @@ def test_mock_completion_stream_without_include_usage_hides_usage_chunk_but_logs assert _prompt_token_counter_calls(token_counter) == [] +def test_mock_completion_stream_with_empty_stream_options_completes_and_logs_admission_count(): + with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: + chunks: Final = list( + litellm.completion( + model="openai/gpt-5.4-mini", + messages=_MOCK_STREAM_MESSAGES, + mock_response="ok", + api_key="mock", + stream=True, + stream_options={}, + metadata=_ADMISSION_METADATA, + ) + ) + + assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks) == "ok" + assert _client_usage_chunks(chunks) == [] + assert _prompt_token_counter_calls(token_counter) == [] + + +@pytest.mark.asyncio +async def test_mock_acompletion_stream_with_empty_stream_options_completes_and_logs_admission_count(): + with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: + response: Final = await litellm.acompletion( + model="openai/gpt-5.4-mini", + messages=_MOCK_STREAM_MESSAGES, + mock_response="ok", + api_key="mock", + stream=True, + stream_options={}, + litellm_metadata=_ADMISSION_METADATA, + ) + chunks: Final = [chunk async for chunk in response] + + assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks) == "ok" + assert _client_usage_chunks(chunks) == [] + assert _prompt_token_counter_calls(token_counter) == [] + + def test_mock_completion_stream_without_admission_count_falls_back_to_tokenizer(): expected_prompt_tokens: Final = litellm.token_counter(model="openai/gpt-5.4-mini", messages=_MOCK_STREAM_MESSAGES) with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: - chunks = list( + chunks: Final = list( litellm.completion( model="openai/gpt-5.4-mini", messages=_MOCK_STREAM_MESSAGES, @@ -2524,7 +2566,7 @@ def test_mock_completion_stream_without_admission_count_falls_back_to_tokenizer( ) ) - usage_chunks = _client_usage_chunks(chunks) + usage_chunks: Final = _client_usage_chunks(chunks) assert len(usage_chunks) == 1 assert usage_chunks[0].prompt_tokens == expected_prompt_tokens assert usage_chunks[0].total_tokens == expected_prompt_tokens + usage_chunks[0].completion_tokens @@ -2535,7 +2577,7 @@ def test_mock_completion_stream_without_admission_count_falls_back_to_tokenizer( async def test_mock_acompletion_stream_without_admission_count_falls_back_to_tokenizer(): expected_prompt_tokens: Final = litellm.token_counter(model="openai/gpt-5.4-mini", messages=_MOCK_STREAM_MESSAGES) with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: - response = await litellm.acompletion( + response: Final = await litellm.acompletion( model="openai/gpt-5.4-mini", messages=_MOCK_STREAM_MESSAGES, mock_response="ok", @@ -2543,40 +2585,87 @@ async def test_mock_acompletion_stream_without_admission_count_falls_back_to_tok stream=True, stream_options={"include_usage": True}, ) - chunks = [chunk async for chunk in response] + chunks: Final = [chunk async for chunk in response] - usage_chunks = _client_usage_chunks(chunks) + usage_chunks: Final = _client_usage_chunks(chunks) assert len(usage_chunks) == 1 assert usage_chunks[0].prompt_tokens == expected_prompt_tokens assert len(_prompt_token_counter_calls(token_counter)) >= 1 -def test_mock_completion_stream_and_non_stream_report_the_same_admission_usage(): - non_stream = litellm.completion( +def _usage_triple(usage: Usage) -> tuple[int, int, int]: + return (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) + + +@pytest.mark.parametrize("input_tokens", (_ADMISSION_INPUT_TOKENS, 0)) +def test_mock_completion_stream_and_non_stream_report_the_same_admission_usage(input_tokens: int): + metadata: Final = _admission_metadata(input_tokens) + non_stream: Final = litellm.completion( model="openai/gpt-5.4-mini", messages=_MOCK_STREAM_MESSAGES, mock_response="ok", api_key="mock", - metadata=_ADMISSION_METADATA, + metadata=metadata, ) - chunks = list( - litellm.completion( + with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: + chunks: Final = list( + litellm.completion( + model="openai/gpt-5.4-mini", + messages=_MOCK_STREAM_MESSAGES, + mock_response="ok", + api_key="mock", + stream=True, + stream_options={"include_usage": True}, + metadata=metadata, + ) + ) + + assert _usage_triple(non_stream.usage) == _usage_triple(_client_usage_chunks(chunks)[0]) + assert non_stream.usage.prompt_tokens == input_tokens + assert _prompt_token_counter_calls(token_counter) == [] + + +@pytest.mark.asyncio +async def test_mock_acompletion_stream_reports_zero_admission_input_tokens_without_tokenizer_fallback(): + with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: + response: Final = await litellm.acompletion( model="openai/gpt-5.4-mini", - messages=_MOCK_STREAM_MESSAGES, + messages=[{"role": "user", "content": ""}], mock_response="ok", api_key="mock", stream=True, stream_options={"include_usage": True}, - metadata=_ADMISSION_METADATA, + litellm_metadata=_admission_metadata(0), + ) + chunks: Final = [chunk async for chunk in response] + + usage_chunks: Final = _client_usage_chunks(chunks) + assert len(usage_chunks) == 1 + assert _usage_triple(usage_chunks[0]) == (0, usage_chunks[0].completion_tokens, usage_chunks[0].completion_tokens) + assert _prompt_token_counter_calls(token_counter) == [] + + +def test_mock_text_completion_stream_and_non_stream_report_the_same_zero_admission_usage(): + metadata: Final = _admission_metadata(0) + non_stream: Final = litellm.text_completion( + model="openai/gpt-5.4-mini", prompt="", mock_response="ok", api_key="mock", metadata=metadata + ) + chunks: Final = list( + litellm.text_completion( + model="openai/gpt-5.4-mini", + prompt="", + mock_response="ok", + api_key="mock", + stream=True, + stream_options={"include_usage": True}, + metadata=metadata, ) ) - stream_usage: Final = _client_usage_chunks(chunks)[0] - assert (non_stream.usage.prompt_tokens, non_stream.usage.completion_tokens, non_stream.usage.total_tokens) == ( - stream_usage.prompt_tokens, - stream_usage.completion_tokens, - stream_usage.total_tokens, - ) + stream_usages: Final = tuple(chunk.usage for chunk in chunks if getattr(chunk, "usage", None) is not None) + assert len(stream_usages) == 1 + assert _usage_triple(non_stream.usage) == _usage_triple(stream_usages[0]) + assert non_stream.usage.prompt_tokens == 0 def test_mock_completion_stream_with_model_response():