mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Merge pull request #41915 from BerriAI/litellm_distributed_tpm_shared_redis_check
fix(router): enforce model tpm limits against shared redis usage across replicas
This commit is contained in:
commit
4487a9915b
2 changed files with 136 additions and 4 deletions
|
|
@ -18,6 +18,7 @@ import httpx
|
|||
import litellm
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.caching.redis_cache import RedisCircuitBreakerOpenError
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.router_utils.pre_call_checks.io_token_rate_limit_check import (
|
||||
ITPM_RESERVED_KEY,
|
||||
|
|
@ -136,6 +137,26 @@ class ModelRateLimitingCheck(CustomLogger):
|
|||
|
||||
return tpm_key, rpm_key
|
||||
|
||||
def _get_current_tpm(self, tpm_key: str, tpm_limit: int) -> int | None:
|
||||
local_tpm: Final = self.dual_cache.get_cache(key=tpm_key, local_only=True)
|
||||
redis_cache: Final = self.dual_cache.redis_cache
|
||||
if redis_cache is None or (local_tpm is not None and local_tpm >= tpm_limit):
|
||||
return local_tpm
|
||||
try:
|
||||
return redis_cache.get_cache(key=tpm_key)
|
||||
except RedisCircuitBreakerOpenError:
|
||||
return local_tpm
|
||||
|
||||
async def _async_get_current_tpm(self, tpm_key: str, tpm_limit: int, parent_otel_span: Span | None) -> int | None:
|
||||
local_tpm: Final = await self.dual_cache.async_get_cache(key=tpm_key, local_only=True)
|
||||
redis_cache: Final = self.dual_cache.redis_cache
|
||||
if redis_cache is None or (local_tpm is not None and local_tpm >= tpm_limit):
|
||||
return local_tpm
|
||||
try:
|
||||
return await redis_cache.async_get_cache(key=tpm_key, parent_otel_span=parent_otel_span)
|
||||
except RedisCircuitBreakerOpenError:
|
||||
return local_tpm
|
||||
|
||||
def pre_call_check(self, deployment: dict) -> dict | None:
|
||||
"""
|
||||
Synchronous pre-call check for model rate limits.
|
||||
|
|
@ -168,8 +189,7 @@ class ModelRateLimitingCheck(CustomLogger):
|
|||
|
||||
# Check TPM limit
|
||||
if tpm_limit is not None:
|
||||
# First check local cache
|
||||
current_tpm: Final = self.dual_cache.get_cache(key=tpm_key, local_only=True)
|
||||
current_tpm: Final = self._get_current_tpm(tpm_key, tpm_limit)
|
||||
if current_tpm is not None and current_tpm >= tpm_limit:
|
||||
raise litellm.RateLimitError(
|
||||
message=f"Model rate limit exceeded. TPM limit={tpm_limit}, current usage={current_tpm}",
|
||||
|
|
@ -249,8 +269,7 @@ class ModelRateLimitingCheck(CustomLogger):
|
|||
|
||||
# Check TPM limit
|
||||
if tpm_limit is not None:
|
||||
# First check local cache
|
||||
current_tpm: Final = await self.dual_cache.async_get_cache(key=tpm_key, local_only=True)
|
||||
current_tpm: Final = await self._async_get_current_tpm(tpm_key, tpm_limit, parent_otel_span)
|
||||
if current_tpm is not None and current_tpm >= tpm_limit:
|
||||
raise litellm.RateLimitError(
|
||||
message=f"Model rate limit exceeded. TPM limit={tpm_limit}, current usage={current_tpm}",
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ regardless of the routing strategy being used.
|
|||
"""
|
||||
|
||||
import asyncio
|
||||
from datetime import timedelta
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
|
@ -13,10 +14,28 @@ import pytest
|
|||
import litellm
|
||||
from litellm import Router
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.caching.redis_cache import RedisCircuitBreakerOpenError
|
||||
from litellm.router_utils.pre_call_checks.model_rate_limit_check import (
|
||||
ModelRateLimitingCheck,
|
||||
)
|
||||
|
||||
TPM_DEPLOYMENT = {
|
||||
"tpm": 1000,
|
||||
"litellm_params": {"model": "gpt-4"},
|
||||
"model_info": {"id": "replica-test-id"},
|
||||
"model_name": "test-model",
|
||||
}
|
||||
|
||||
|
||||
def _dual_cache_with_local_tpm(local_tpm: int, redis_cache: MagicMock | None) -> DualCache:
|
||||
dual_cache = DualCache(redis_cache=redis_cache)
|
||||
check = ModelRateLimitingCheck(dual_cache=dual_cache)
|
||||
now = litellm.utils.get_utc_datetime()
|
||||
for minute in (now, now + timedelta(minutes=1)):
|
||||
tpm_key, _ = check._get_cache_keys(TPM_DEPLOYMENT, minute.strftime("%H-%M"))
|
||||
dual_cache.set_cache(key=tpm_key, value=local_tpm, local_only=True)
|
||||
return dual_cache
|
||||
|
||||
|
||||
class TestModelRateLimitingCheck:
|
||||
"""Test the ModelRateLimitingCheck class directly."""
|
||||
|
|
@ -144,6 +163,50 @@ class TestModelRateLimitingCheck:
|
|||
assert "TPM limit=1000" in str(exc_info.value)
|
||||
assert "current usage=1000" in str(exc_info.value)
|
||||
|
||||
def test_pre_call_check_rejects_when_shared_tpm_is_over_limit_but_local_is_under(self):
|
||||
redis_cache = MagicMock()
|
||||
redis_cache.get_cache.return_value = 1000
|
||||
check = ModelRateLimitingCheck(dual_cache=_dual_cache_with_local_tpm(5, redis_cache))
|
||||
|
||||
with pytest.raises(litellm.RateLimitError) as exc_info:
|
||||
check.pre_call_check(TPM_DEPLOYMENT)
|
||||
|
||||
assert "current usage=1000" in str(exc_info.value)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"redis_get", [MagicMock(return_value=None), MagicMock(side_effect=RedisCircuitBreakerOpenError())]
|
||||
)
|
||||
def test_pre_call_check_keeps_rejecting_on_local_usage_when_redis_read_fails(self, redis_get):
|
||||
redis_cache = MagicMock()
|
||||
redis_cache.get_cache = redis_get
|
||||
check = ModelRateLimitingCheck(dual_cache=_dual_cache_with_local_tpm(1000, redis_cache))
|
||||
|
||||
with pytest.raises(litellm.RateLimitError) as exc_info:
|
||||
check.pre_call_check(TPM_DEPLOYMENT)
|
||||
|
||||
assert "current usage=1000" in str(exc_info.value)
|
||||
|
||||
def test_pre_call_check_falls_back_to_local_tpm_and_still_checks_rpm_when_redis_circuit_is_open(self):
|
||||
redis_cache = MagicMock()
|
||||
redis_cache.get_cache.side_effect = RedisCircuitBreakerOpenError()
|
||||
redis_cache.increment_cache.return_value = 2
|
||||
check = ModelRateLimitingCheck(dual_cache=_dual_cache_with_local_tpm(5, redis_cache))
|
||||
|
||||
with pytest.raises(litellm.RateLimitError) as exc_info:
|
||||
check.pre_call_check({**TPM_DEPLOYMENT, "rpm": 1})
|
||||
|
||||
assert "RPM limit=1" in str(exc_info.value)
|
||||
|
||||
def test_pre_call_check_without_redis_enforces_local_tpm_and_rpm(self):
|
||||
check = ModelRateLimitingCheck(dual_cache=_dual_cache_with_local_tpm(5, None))
|
||||
deployment = {**TPM_DEPLOYMENT, "rpm": 1}
|
||||
|
||||
assert check.pre_call_check(deployment) == deployment
|
||||
with pytest.raises(litellm.RateLimitError) as exc_info:
|
||||
check.pre_call_check(deployment)
|
||||
|
||||
assert "RPM limit=1" in str(exc_info.value)
|
||||
|
||||
def test_log_success_event_increments_cache(self):
|
||||
"""Test that log_success_event correctly increments the cache."""
|
||||
mock_cache = MagicMock()
|
||||
|
|
@ -245,6 +308,56 @@ class TestModelRateLimitingCheckAsync:
|
|||
|
||||
assert "TPM limit=1000" in str(exc_info.value)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_pre_call_check_rejects_when_shared_tpm_is_over_limit_but_local_is_under(self):
|
||||
redis_cache = MagicMock()
|
||||
redis_cache.async_get_cache = AsyncMock(return_value=1000)
|
||||
check = ModelRateLimitingCheck(dual_cache=_dual_cache_with_local_tpm(5, redis_cache))
|
||||
|
||||
with pytest.raises(litellm.RateLimitError) as exc_info:
|
||||
await check.async_pre_call_check(TPM_DEPLOYMENT)
|
||||
|
||||
assert "current usage=1000" in str(exc_info.value)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"redis_get", [AsyncMock(return_value=None), AsyncMock(side_effect=RedisCircuitBreakerOpenError())]
|
||||
)
|
||||
async def test_async_pre_call_check_keeps_rejecting_on_local_usage_when_redis_read_fails(self, redis_get):
|
||||
redis_cache = MagicMock()
|
||||
redis_cache.async_get_cache = redis_get
|
||||
check = ModelRateLimitingCheck(dual_cache=_dual_cache_with_local_tpm(1000, redis_cache))
|
||||
|
||||
with pytest.raises(litellm.RateLimitError) as exc_info:
|
||||
await check.async_pre_call_check(TPM_DEPLOYMENT)
|
||||
|
||||
assert "current usage=1000" in str(exc_info.value)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_pre_call_check_falls_back_to_local_tpm_and_still_checks_rpm_when_redis_circuit_is_open(
|
||||
self,
|
||||
):
|
||||
redis_cache = MagicMock()
|
||||
redis_cache.async_get_cache = AsyncMock(side_effect=RedisCircuitBreakerOpenError())
|
||||
redis_cache.async_increment = AsyncMock(return_value=2)
|
||||
check = ModelRateLimitingCheck(dual_cache=_dual_cache_with_local_tpm(5, redis_cache))
|
||||
|
||||
with pytest.raises(litellm.RateLimitError) as exc_info:
|
||||
await check.async_pre_call_check({**TPM_DEPLOYMENT, "rpm": 1})
|
||||
|
||||
assert "RPM limit=1" in str(exc_info.value)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_pre_call_check_without_redis_enforces_local_tpm_and_rpm(self):
|
||||
check = ModelRateLimitingCheck(dual_cache=_dual_cache_with_local_tpm(5, None))
|
||||
deployment = {**TPM_DEPLOYMENT, "rpm": 1}
|
||||
|
||||
assert await check.async_pre_call_check(deployment) == deployment
|
||||
with pytest.raises(litellm.RateLimitError) as exc_info:
|
||||
await check.async_pre_call_check(deployment)
|
||||
|
||||
assert "RPM limit=1" in str(exc_info.value)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_log_success_event_increments_cache(self):
|
||||
"""Test that async_log_success_event correctly increments the cache."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue