fix(proxy): count embedding and text completion tokens toward TPM limits (#30105)

* fix(proxy): count embedding and text completion tokens toward TPM limits

The parallel request limiters only read token usage off ModelResponse,
so EmbeddingResponse and TextCompletionResponse objects left
total_tokens at 0 and the per key, user, team, and end user TPM
counters never incremented. Requests to /v1/embeddings and
/v1/completions were effectively free against any tpm_limit. In the v3
limiter this was worse: the post-call reconciliation computed actual
usage as 0 and refunded the pre-call reservation made at request time.

Broaden the isinstance checks to accept EmbeddingResponse and
TextCompletionResponse, which both expose a Usage object, at the four
per-scope sites in parallel_request_limiter.py and at the usage
extraction in parallel_request_limiter_v3.py. ResponsesAPIResponse was
already covered in v3 via BaseLiteLLMOpenAIResponseObject.

Fixes #27738.

* test(proxy): cover v1 limiter TPM counting for embedding and text completion responses

Exercise the broadened isinstance sites in parallel_request_limiter.py
by asserting that async_log_success_event adds total_tokens to the per
key, user, team, and end user TPM counters for EmbeddingResponse and
TextCompletionResponse objects. The counters are pre-seeded at zero so
the assertion is exactly the increment; on the pre-fix code these
responses left total_tokens at 0 and the test fails.
This commit is contained in:
Filippo Menghi 2026-06-10 12:49:09 +02:00 • committed by GitHub
parent d47268939d
commit ca9284f8da
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 185 additions and 10 deletions

View file

@ -7,7 +7,7 @@ from pydantic import BaseModel
from typing_extensions import TypedDict
import litellm
from litellm import DualCache, ModelResponse
from litellm import DualCache, EmbeddingResponse, ModelResponse, TextCompletionResponse
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs
@ -570,7 +570,9 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
total_tokens = 0
if isinstance(response_obj, ModelResponse):
if isinstance(
response_obj, (ModelResponse, EmbeddingResponse, TextCompletionResponse)
):
total_tokens = response_obj.usage.total_tokens # type: ignore
# ------------
@ -659,7 +661,10 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
if user_api_key_user_id is not None:
total_tokens = 0
if isinstance(response_obj, ModelResponse):
if isinstance(
response_obj,
(ModelResponse, EmbeddingResponse, TextCompletionResponse),
):
total_tokens = response_obj.usage.total_tokens # type: ignore
request_count_api_key = (
@ -692,7 +697,10 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
if user_api_key_team_id is not None:
total_tokens = 0
if isinstance(response_obj, ModelResponse):
if isinstance(
response_obj,
(ModelResponse, EmbeddingResponse, TextCompletionResponse),
):
total_tokens = response_obj.usage.total_tokens # type: ignore
request_count_api_key = (
@ -725,7 +733,10 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
if user_api_key_end_user_id is not None:
total_tokens = 0
if isinstance(response_obj, ModelResponse):
if isinstance(
response_obj,
(ModelResponse, EmbeddingResponse, TextCompletionResponse),
):
total_tokens = response_obj.usage.total_tokens # type: ignore
request_count_api_key = (

View file

@ -39,7 +39,13 @@ from litellm.proxy.common_utils.proxy_rate_limit_error import (
from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit
from litellm.types.caching import RedisPipelineIncrementOperation
from litellm.types.llms.openai import BaseLiteLLMOpenAIResponseObject
from litellm.types.utils import CallTypes, ModelResponse, Usage
from litellm.types.utils import (
CallTypes,
EmbeddingResponse,
ModelResponse,
TextCompletionResponse,
Usage,
)
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
@ -2736,9 +2742,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
# Get total tokens from response
total_tokens = 0
# spot fix for /responses api
if isinstance(response_obj, ModelResponse) or isinstance(
response_obj, BaseLiteLLMOpenAIResponseObject
if isinstance(
response_obj,
(
ModelResponse,
EmbeddingResponse,
TextCompletionResponse,
BaseLiteLLMOpenAIResponseObject,
),
):
_usage = getattr(response_obj, "usage", None)
total_tokens = self._get_total_tokens_from_usage(

View file

@ -0,0 +1,86 @@
"""
Unit Tests for the max parallel request limiter v1 for the proxy
"""
from datetime import datetime
import pytest
from litellm.caching.caching import DualCache
from litellm.proxy.hooks.parallel_request_limiter import (
_PROXY_MaxParallelRequestsHandler,
)
from litellm.proxy.utils import InternalUsageCache, hash_token
from litellm.types.utils import EmbeddingResponse, TextCompletionResponse, Usage
@pytest.mark.parametrize(
"response_obj",
[
EmbeddingResponse(
model="text-embedding-3-small",
usage=Usage(prompt_tokens=50, completion_tokens=0, total_tokens=50),
),
TextCompletionResponse(
model="gpt-3.5-turbo-instruct",
usage=Usage(prompt_tokens=20, completion_tokens=30, total_tokens=50),
),
],
)
@pytest.mark.asyncio
async def test_async_log_success_event_counts_non_chat_response_tokens(response_obj):
"""
Embedding and text completion responses must increment the per key, user,
team, and end user TPM counters, not just chat completion ModelResponse
objects.
"""
_api_key = hash_token("sk-12345")
user_id = "ishaan"
team_id = "litellm-team"
end_user_id = "customer-1"
parallel_request_handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache())
)
current_date = datetime.now().strftime("%Y-%m-%d")
current_hour = datetime.now().strftime("%H")
current_minute = datetime.now().strftime("%M")
precise_minute = f"{current_date}-{current_hour}-{current_minute}"
scope_ids = [_api_key, user_id, team_id, end_user_id]
for scope_id in scope_ids:
await parallel_request_handler.internal_usage_cache.async_set_cache(
key=f"{scope_id}::{precise_minute}::request_count",
value={"current_requests": 1, "current_tpm": 0, "current_rpm": 1},
litellm_parent_otel_span=None,
)
kwargs = {
"litellm_params": {
"metadata": {
"user_api_key": _api_key,
"user_api_key_user_id": user_id,
"user_api_key_team_id": team_id,
"user_api_key_model_max_budget": {},
}
},
"user": end_user_id,
}
await parallel_request_handler.async_log_success_event(
kwargs=kwargs,
response_obj=response_obj,
start_time=datetime.now(),
end_time=datetime.now(),
)
for scope_id in scope_ids:
current = await parallel_request_handler.internal_usage_cache.async_get_cache(
key=f"{scope_id}::{precise_minute}::request_count",
litellm_parent_otel_span=None,
)
assert current["current_tpm"] == 50, (
f"expected 50 tokens counted for {scope_id}, "
f"got {current['current_tpm']}"
)

View file

@ -20,7 +20,12 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import (
_PROXY_MaxParallelRequestsHandler_v3 as _PROXY_MaxParallelRequestsHandler,
)
from litellm.proxy.utils import InternalUsageCache, ProxyLogging, hash_token
from litellm.types.utils import ModelResponse, Usage
from litellm.types.utils import (
EmbeddingResponse,
ModelResponse,
TextCompletionResponse,
Usage,
)
class TimeController:
@ -547,6 +552,68 @@ async def test_token_rate_limit_type_respected_v3(monkeypatch, token_rate_limit_
), f"Expected {expected_tokens[token_rate_limit_type]} tokens for type '{token_rate_limit_type}', got {tpm_operation['increment_value']}"
@pytest.mark.parametrize(
"response_obj",
[
EmbeddingResponse(
model="text-embedding-3-small",
usage=Usage(prompt_tokens=50, completion_tokens=0, total_tokens=50),
),
TextCompletionResponse(
model="gpt-3.5-turbo-instruct",
usage=Usage(prompt_tokens=20, completion_tokens=30, total_tokens=50),
),
],
)
@pytest.mark.asyncio
async def test_async_log_success_event_counts_non_chat_response_tokens(
monkeypatch, response_obj
):
"""
Embedding and text completion responses must increment the TPM counter,
not just chat completion ModelResponse objects.
"""
monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60")
_api_key = hash_token("sk-12345")
parallel_request_handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache())
)
monkeypatch.setattr(
parallel_request_handler, "get_rate_limit_type", lambda: "total"
)
mock_kwargs = {
"standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}},
"model": response_obj.model,
}
captured_operations = []
async def mock_increment_pipeline(increment_list, **kwargs):
captured_operations.extend(increment_list)
return True
monkeypatch.setattr(
parallel_request_handler.internal_usage_cache.dual_cache,
"async_increment_cache_pipeline",
mock_increment_pipeline,
)
await parallel_request_handler.async_log_success_event(
kwargs=mock_kwargs,
response_obj=response_obj,
start_time=datetime.now(),
end_time=datetime.now(),
)
tpm_operation = next(
(op for op in captured_operations if op["key"].endswith(":tokens")), None
)
assert tpm_operation is not None, "Should have a TPM increment operation"
assert tpm_operation["increment_value"] == 50
@pytest.mark.asyncio
async def test_async_log_failure_event_v3():
"""