fix(proxy): count failure and rate-limit input tokens off the event loop

The failure hook's usage estimate and the project ITPM reservation both
called litellm.token_counter inline on the event loop, so a large request
that failed or hit the limiter stalled the gateway the same way the
count_tokens endpoints did. Both now run through asyncify. The loop-lag
probe the existing tests used moves into a shared helper that warms the
tokenizer first, and two new tests fail when either count runs inline
This commit is contained in:
mateo-berri 2026-09-07 19:15:36 -07:00
parent 0d89873daf
commit dcd38ab9f0
7 changed files with 139 additions and 45 deletions

View file

@ -28,6 +28,7 @@ from litellm import DualCache
from litellm._logging import verbose_proxy_logger
from litellm.constants import DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE, INTERNAL_CALL_ORIGIN_METADATA_KEY
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.asyncify import asyncify
from litellm.litellm_core_utils.prompt_templates.common_utils import (
get_str_from_messages,
)
@ -3303,7 +3304,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
min_configured_tpm_limit=min_configured_otpm_limit,
call_type=call_type,
)
raw_estimated_input_tokens: Final = self._estimate_precise_input_tokens(
raw_estimated_input_tokens: Final = await asyncify(self._estimate_precise_input_tokens)(
data=data, model=requested_model, call_type=call_type
)
estimated_input_tokens: Final = max(raw_estimated_input_tokens, 1)

View file

@ -91,6 +91,7 @@ from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.prometheus import PrometheusLogger
from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
from litellm.integrations.SlackAlerting.utils import _add_langfuse_trace_id_to_alert
from litellm.litellm_core_utils.asyncify import asyncify
from litellm.litellm_core_utils.core_helpers import (
coerce_token_limit,
independent_snapshot,
@ -2569,7 +2570,7 @@ class ProxyLogging:
original_exception=original_exception,
)
request_data.update(_failure_fields_to_lift(request_data))
request_data.update(await asyncify(_failure_fields_to_lift)(request_data))
# Remove before callbacks iterate — not serialisable
request_data.pop("litellm_logging_obj", None)

View file

@ -0,0 +1,39 @@
import asyncio
import time
from collections.abc import Awaitable, Callable
from typing import Final, TypeVar
import litellm
T = TypeVar("T")
def warm_tokenizer(model: str) -> None:
litellm.token_counter(model=model, text="load the tokenizer before anything is timed")
async def loop_wake_lags(until: asyncio.Event) -> tuple[float, ...]:
async def wake_lag() -> float:
started: Final = time.perf_counter()
await asyncio.sleep(0.001)
return time.perf_counter() - started - 0.001
return tuple([await wake_lag() for _ in iter(until.is_set, True)])
async def timed_with_loop_lags(run: Callable[[], Awaitable[T]]) -> tuple[T, float, tuple[float, ...]]:
finished: Final = asyncio.Event()
async def timed() -> tuple[T, float]:
started: Final = time.perf_counter()
try:
return await run(), time.perf_counter() - started
finally:
finished.set()
(result, took), lags = await asyncio.gather(timed(), loop_wake_lags(finished))
return result, took, lags
def assert_loop_stayed_free(took: float, lags: tuple[float, ...]) -> None:
assert max(lags) < took / 4, f"the event loop stalled {max(lags):.3f}s during a {took:.3f}s count"

View file

@ -24,6 +24,11 @@ from litellm.litellm_core_utils.token_counter import (
)
from litellm.litellm_core_utils.token_counter import token_counter as token_counter_new
from tests.large_text import text
from tests.test_litellm.litellm_core_utils.event_loop_lag import (
assert_loop_stayed_free,
timed_with_loop_lags,
warm_tokenizer,
)
from tests.test_litellm.litellm_core_utils.messages_with_counts import (
MESSAGES_TEXT,
MESSAGES_WITH_IMAGES,
@ -127,30 +132,15 @@ def test_valid_chunk_size_config_is_honoured(monkeypatch):
importlib.reload(litellm.constants)
async def _loop_wake_lags(until: asyncio.Event) -> tuple[float, ...]:
async def wake_lag() -> float:
started: Final = time.perf_counter()
await asyncio.sleep(0.001)
return time.perf_counter() - started - 0.001
return tuple([await wake_lag() for _ in iter(until.is_set, True)])
async def test_huggingface_count_in_a_worker_thread_leaves_the_event_loop_free():
counted: Final = asyncio.Event()
warm_tokenizer("claude-fable-5")
async def count_off_the_loop() -> tuple[int, float]:
started: Final = time.perf_counter()
try:
tokens: Final = await asyncify(token_counter_new)(model="claude-fable-5", text=text * 100)
return tokens, time.perf_counter() - started
finally:
counted.set()
(tokens, took), lags = await asyncio.gather(count_off_the_loop(), _loop_wake_lags(counted))
tokens, took, lags = await timed_with_loop_lags(
lambda: asyncify(token_counter_new)(model="claude-fable-5", text=text * 100)
)
assert tokens > 0
assert max(lags) < took / 4, f"the event loop stalled {max(lags):.3f}s during a {took:.3f}s count"
assert_loop_stayed_free(took, lags)
@pytest.mark.parametrize("max_exact_chars", [64, 1_000, 2_500])

View file

@ -3670,5 +3670,42 @@ async def test_post_call_success_hook_contains_header_merge_failures(
)
@pytest.mark.asyncio
async def test_the_project_itpm_reservation_counts_the_request_off_the_event_loop(rate_limiter):
from tests.large_text import text
from tests.test_litellm.litellm_core_utils.event_loop_lag import (
assert_loop_stayed_free,
timed_with_loop_lags,
warm_tokenizer,
)
handler, _cache = rate_limiter
stash = get_or_create_request_stash()
warm_tokenizer("claude-fable-5")
data: Dict[str, Any] = {
"model": "claude-fable-5",
"messages": [{"role": "user", "content": text * 100}],
}
itpm_descriptor = {
"key": PROJECT_ITPM_DESCRIPTOR_KEY,
"value": "proj-loop:claude-fable-5",
"rate_limit": {"tokens_per_unit": 10_000_000, "window_size": 60},
}
_, took, lags = await timed_with_loop_lags(
lambda: handler._reserve_project_io_tokens_or_raise(
descriptors=[itpm_descriptor],
data=data,
requested_model="claude-fable-5",
user_api_key_dict=UserAPIKeyAuth(api_key=hash_token("sk-itpm-loop"), project_id="proj-loop"),
tpm_reservation_scopes=[],
tpm_reservation_amount=0,
)
)
assert stash.rate_limit_response is not None
assert_loop_stayed_free(took, lags)
if __name__ == "__main__":
pytest.main([__file__, "-v", "-s"])

View file

@ -12860,32 +12860,20 @@ async def test_update_general_settings_keeps_yaml_openai_websocket_passthrough()
assert ps.general_settings["enable_openai_websocket_passthrough"] is False
async def _loop_wake_lags(until: asyncio.Event) -> tuple[float, ...]:
async def wake_lag() -> float:
started: Final = time.perf_counter()
await asyncio.sleep(0.001)
return time.perf_counter() - started - 0.001
return tuple([await wake_lag() for _ in iter(until.is_set, True)])
async def test_token_counter_keeps_the_event_loop_free_during_a_huggingface_count(monkeypatch):
from tests.large_text import text
from tests.test_litellm.litellm_core_utils.event_loop_lag import (
assert_loop_stayed_free,
timed_with_loop_lags,
warm_tokenizer,
)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None)
counted: Final = asyncio.Event()
warm_tokenizer("claude-fable-5")
async def count_off_the_loop() -> tuple[int, float]:
started: Final = time.perf_counter()
try:
response: Final = await proxy_server_module.token_counter(
TokenCountRequest(model="claude-fable-5", prompt=text * 100)
)
return response.total_tokens, time.perf_counter() - started
finally:
counted.set()
response, took, lags = await timed_with_loop_lags(
lambda: proxy_server_module.token_counter(TokenCountRequest(model="claude-fable-5", prompt=text * 100))
)
(tokens, took), lags = await asyncio.gather(count_off_the_loop(), _loop_wake_lags(counted))
assert tokens > 0
assert max(lags) < took / 4, f"the event loop stalled {max(lags):.3f}s during a {took:.3f}s count"
assert response.total_tokens > 0
assert_loop_stayed_free(took, lags)

View file

@ -1823,6 +1823,44 @@ def test_a_dispatched_failure_lifts_the_four_fields_the_spend_log_needs():
assert lifted["standard_logging_object"] == {"id": "log-1"}
@pytest.mark.asyncio
async def test_a_dispatched_failure_is_counted_off_the_event_loop():
from unittest.mock import AsyncMock, patch
from tests.large_text import text
from tests.test_litellm.litellm_core_utils.event_loop_lag import (
assert_loop_stayed_free,
timed_with_loop_lags,
warm_tokenizer,
)
warm_tokenizer("claude-fable-5")
request_data = {
"litellm_logging_obj": _LoggingObj(
{
"first_api_call_start_time": 1700000000.0,
"call_type": "acompletion",
"model": "claude-fable-5",
"messages": [{"role": "user", "content": text * 100}],
}
),
"metadata": {},
}
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging_obj.alert_types = []
with patch.object(proxy_logging_obj, "update_request_status", new=AsyncMock()):
_, took, lags = await timed_with_loop_lags(
lambda: proxy_logging_obj.post_call_failure_hook(
request_data=request_data,
original_exception=Exception("boom"),
user_api_key_dict=UserAPIKeyAuth(),
)
)
assert request_data["combined_usage_object"].prompt_tokens > 0
assert_loop_stayed_free(took, lags)
@pytest.mark.asyncio
async def test_proxy_only_error_expected_4xx_skips_traceback_for_both_handlers(monkeypatch):
"""Regression for LIT-6043: an expected 4xx must not format a traceback for