diff --git a/litellm/proxy/hooks/dynamic_rate_limiter.py b/litellm/proxy/hooks/dynamic_rate_limiter.py index f2e31b77761..83161fca5bb 100644 --- a/litellm/proxy/hooks/dynamic_rate_limiter.py +++ b/litellm/proxy/hooks/dynamic_rate_limiter.py @@ -4,7 +4,8 @@ import asyncio import os -from typing import List, Optional, Tuple, Union +from datetime import datetime +from typing import Callable, List, Optional, Tuple, Union import litellm from litellm import ModelResponse, Router @@ -30,12 +31,13 @@ class DynamicRateLimiterCache: Track number of active projects calling a model. """ - def __init__(self, cache: DualCache) -> None: + def __init__(self, cache: DualCache, time_fn: Callable[[], datetime] = get_utc_datetime) -> None: self.cache = cache self.ttl = 60 # 1 min ttl + self.time_fn = time_fn async def async_get_cache(self, model: str) -> Optional[int]: - dt = get_utc_datetime() + dt = self.time_fn() current_minute = dt.strftime("%H-%M") key_name = "{}:{}".format(current_minute, model) _response = await self.cache.async_get_cache(key=key_name) @@ -59,7 +61,7 @@ class DynamicRateLimiterCache: - Exception, if unable to connect to cache client (if redis caching enabled) """ try: - dt = get_utc_datetime() + dt = self.time_fn() current_minute = dt.strftime("%H-%M") key_name = "{}:{}".format(current_minute, model) @@ -75,8 +77,8 @@ class DynamicRateLimiterCache: class _PROXY_DynamicRateLimitHandler(CustomLogger): # Class variables or attributes - def __init__(self, internal_usage_cache: DualCache): - self.internal_usage_cache = DynamicRateLimiterCache(cache=internal_usage_cache) + def __init__(self, internal_usage_cache: DualCache, time_fn: Callable[[], datetime] = get_utc_datetime): + self.internal_usage_cache = DynamicRateLimiterCache(cache=internal_usage_cache, time_fn=time_fn) def update_variables(self, llm_router: Router): self.llm_router = llm_router diff --git a/tests/local_testing/test_dynamic_rate_limit_handler.py b/tests/local_testing/test_dynamic_rate_limit_handler.py index ff540e22e7a..d288d622cfa 100644 --- a/tests/local_testing/test_dynamic_rate_limit_handler.py +++ b/tests/local_testing/test_dynamic_rate_limit_handler.py @@ -7,7 +7,7 @@ import sys import time import traceback from litellm._uuid import uuid -from datetime import datetime +from datetime import datetime, timezone from typing import Optional, Tuple from dotenv import load_dotenv @@ -38,7 +38,8 @@ Basic test cases: @pytest.fixture def dynamic_rate_limit_handler() -> DynamicRateLimitHandler: internal_cache = DualCache() - return DynamicRateLimitHandler(internal_usage_cache=internal_cache) + frozen_now = datetime(2024, 1, 1, 10, 30, 0, tzinfo=timezone.utc) + return DynamicRateLimitHandler(internal_usage_cache=internal_cache, time_fn=lambda: frozen_now) @pytest.fixture diff --git a/tests/test_litellm/proxy/hooks/test_dynamic_rate_limiter.py b/tests/test_litellm/proxy/hooks/test_dynamic_rate_limiter.py new file mode 100644 index 00000000000..b630ff1605c --- /dev/null +++ b/tests/test_litellm/proxy/hooks/test_dynamic_rate_limiter.py @@ -0,0 +1,44 @@ +from datetime import datetime, timezone + +import pytest + +from litellm.caching.caching import DualCache +from litellm.proxy.hooks.dynamic_rate_limiter import ( + DynamicRateLimiterCache, + _PROXY_DynamicRateLimitHandler, +) + + +@pytest.mark.asyncio +async def test_sadd_and_get_share_injected_clock_window(): + dual_cache = DualCache() + cache = DynamicRateLimiterCache( + cache=dual_cache, + time_fn=lambda: datetime(2024, 1, 1, 10, 30, 0, tzinfo=timezone.utc), + ) + await cache.async_set_cache_sadd(model="my-fake-model", value=["p1", "p2", "p3"]) + assert await cache.async_get_cache(model="my-fake-model") == 3 + assert await dual_cache.async_get_cache(key="10-30:my-fake-model") is not None + + +@pytest.mark.asyncio +async def test_minute_rollover_between_sadd_and_get_reads_empty_window(): + ticks = iter( + ( + datetime(2024, 1, 1, 10, 30, 59, 999999, tzinfo=timezone.utc), + datetime(2024, 1, 1, 10, 31, 0, 0, tzinfo=timezone.utc), + ) + ) + cache = DynamicRateLimiterCache(cache=DualCache(), time_fn=lambda: next(ticks)) + await cache.async_set_cache_sadd(model="my-fake-model", value=["p1"]) + assert await cache.async_get_cache(model="my-fake-model") is None + + +@pytest.mark.asyncio +async def test_handler_threads_time_fn_to_internal_cache(): + handler = _PROXY_DynamicRateLimitHandler( + internal_usage_cache=DualCache(), + time_fn=lambda: datetime(2024, 1, 1, 10, 30, 0, tzinfo=timezone.utc), + ) + await handler.internal_usage_cache.async_set_cache_sadd(model="my-fake-model", value=["p1", "p2"]) + assert await handler.internal_usage_cache.async_get_cache(model="my-fake-model") == 2