mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
d47268939d
commit
ca9284f8da
4 changed files with 185 additions and 10 deletions
|
|
@ -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 = (
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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']}"
|
||||
)
|
||||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue