mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
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:
parent
0d89873daf
commit
dcd38ab9f0
7 changed files with 139 additions and 45 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
39
tests/test_litellm/litellm_core_utils/event_loop_lag.py
Normal file
39
tests/test_litellm/litellm_core_utils/event_loop_lag.py
Normal 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"
|
||||
|
|
@ -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])
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue