From 63f10eef274d1e00dd6e850b76d02dd2b0ac0b60 Mon Sep 17 00:00:00 2001 From: Nitish Agarwal <1592163+nitishagar@users.noreply.github.com> Date: Thu, 18 Jun 2026 13:38:21 +0530 Subject: [PATCH] fix(spend-counter): dedup Redis increment to prevent cluster-mode double-apply inflation RedisCache runs with cluster_error_retry_attempts=3, so a blind INCRBYFLOAT can be re-applied when the client retries after a TimeoutError. That inflates the cross-pod spend counter and can trip false 429s on keys nowhere near their budget. async_increment_idempotent gates the INCRBYFLOAT behind a Lua SET NX on a per-increment dedup key, so a cluster re-issue of the same eval finds the marker already set and returns the current value without applying the increment a second time. It falls back to plain async_increment when scripting is unavailable or the script has already failed once, keeping the degrade path safe. _increment_spend_counter_cache calls async_increment_idempotent with a fresh uuid4 dedup id per logical increment when the method is available, leaving the in-memory and cold-seed paths untouched --- litellm/caching/redis_cache.py | 79 ++++++ litellm/proxy/proxy_server.py | 19 +- .../test_litellm/caching/test_redis_cache.py | 128 ++++++++- .../proxy/proxy_server/test_spend_counters.py | 260 ++++++++++++++++-- tests/test_litellm/proxy/test_proxy_server.py | 7 + 5 files changed, 470 insertions(+), 23 deletions(-) diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index ba07511448a..0bf8af9eacc 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -14,6 +14,7 @@ import functools import hashlib import inspect import json +import re import time from datetime import timedelta from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union, cast @@ -258,6 +259,8 @@ class RedisCache(BaseCache): recovery_timeout=REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT, enabled=REDIS_CIRCUIT_BREAKER_ENABLED, ) + self._dedup_increment_script: Optional[Any] = None + self._dedup_increment_script_disabled: bool = False self._setup_health_pings() @@ -940,6 +943,82 @@ class RedisCache(BaseCache): result = result.decode() return float(result) + _SPEND_DEDUP_INCREMENT_SCRIPT = """ +if redis.call('SET', KEYS[2], '1', 'NX', 'EX', tonumber(ARGV[2])) then + local v = redis.call('INCRBYFLOAT', KEYS[1], tonumber(ARGV[1])) + if ARGV[3] ~= '' then + redis.call('EXPIRE', KEYS[1], tonumber(ARGV[3])) + end + return v +else + local v = redis.call('GET', KEYS[1]) + if v == false then return '0' else return v end +end +""" + + async def async_increment_idempotent( + self, + key: str, + value: float, + dedup_id: str, + ttl: Optional[int] = None, + refresh_ttl: bool = False, + ) -> float: + """ + Atomically increment a spend counter at most once per (key, dedup_id) pair. + + Uses a Lua script with SET NX as a dedup gate so a re-issued INCRBYFLOAT + (e.g. cluster-mode retry after a client-side timeout) is a no-op on the + second run rather than double-applying the increment. Falls back to the + plain async_increment when scripting is unavailable. The dedup_id should + be a server-generated UUID (uuid4) to avoid client-supplied values + suppressing spend tracking. + """ + if self._dedup_increment_script_disabled: + return await self.async_increment( + key=key, value=value, ttl=ttl, refresh_ttl=refresh_ttl + ) + + key = self.check_and_fix_namespace(key=key) + _used_ttl = self.get_ttl(ttl=ttl) + _m = re.search(r"\{([^}]+)\}", key) + _slot_key = _m.group(1) if _m else key + dedup_key = f"{{{_slot_key}}}:dedup:{dedup_id}" + counter_ttl_arg = ( + str(_used_ttl) if refresh_ttl and _used_ttl is not None else "" + ) + + script_register = getattr(self, "async_register_script", None) + if not callable(script_register): + return await self.async_increment( + key=key, value=value, ttl=ttl, refresh_ttl=refresh_ttl + ) + + try: + if self._dedup_increment_script is None: + self._dedup_increment_script = script_register( + self._SPEND_DEDUP_INCREMENT_SCRIPT + ) + + raw = await self._dedup_increment_script( + keys=[key, dedup_key], + args=[str(value), "300", counter_ttl_arg], + ) + + if inspect.isawaitable(raw): + raw = await raw + + result = float(raw) + return result + except Exception: + self._dedup_increment_script_disabled = True + verbose_logger.warning( + "LiteLLM Redis: idempotent increment Lua script failed, falling back to plain increment", + ) + return await self.async_increment( + key=key, value=value, ttl=ttl, refresh_ttl=refresh_ttl + ) + async def flush_cache_buffer(self): print_verbose( f"flushing to redis....reached size of buffer {len(self.redis_batch_writing_buffer)}" diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 3e056c21614..5aa580b70d3 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -2676,11 +2676,20 @@ async def _is_spend_counter_cache_warm(counter_key: str) -> bool: async def _increment_spend_counter_cache(counter_key: str, increment: float): if spend_counter_cache.redis_cache is not None: try: - current_value = await spend_counter_cache.redis_cache.async_increment( - key=counter_key, - value=increment, - refresh_ttl=True, - ) + redis_cache = spend_counter_cache.redis_cache + if callable(getattr(redis_cache, "async_increment_idempotent", None)): + current_value = await redis_cache.async_increment_idempotent( + key=counter_key, + value=increment, + dedup_id=str(uuid.uuid4()), + refresh_ttl=True, + ) + else: + current_value = await redis_cache.async_increment( + key=counter_key, + value=increment, + refresh_ttl=True, + ) except Exception: await _invalidate_spend_counter(counter_key=counter_key) raise diff --git a/tests/test_litellm/caching/test_redis_cache.py b/tests/test_litellm/caching/test_redis_cache.py index 78192400fb0..babd6783b41 100644 --- a/tests/test_litellm/caching/test_redis_cache.py +++ b/tests/test_litellm/caching/test_redis_cache.py @@ -3,7 +3,6 @@ import sys from unittest.mock import MagicMock, patch import pytest -from fastapi.testclient import TestClient sys.path.insert( 0, os.path.abspath("../../..") @@ -517,3 +516,130 @@ async def test_async_lpop_with_float_redis_version( # Verify the method completed without error assert result is not None + + +# --------------------------------------------------------------------------- +# async_increment_idempotent +# --------------------------------------------------------------------------- + + +class _FakeRedisSpendStore: + """ + Minimal stand-in for RedisCache.async_register_script that faithfully models + the Lua SET NX-gated INCRBYFLOAT without a real Redis server or Lua runtime. + + The registered callable uses an in-process set as the dedup store, so a + re-issue with the same dedup_key is a no-op (exactly-once semantics). + """ + + def __init__(self): + self._counters: dict = {} + self._dedup: set = set() + + def async_register_script(self, script: str): + async def _run(keys, args): + counter_key, dedup_key = keys[0], keys[1] + amount = float(args[0]) + if dedup_key not in self._dedup: + self._dedup.add(dedup_key) + self._counters[counter_key] = ( + self._counters.get(counter_key, 0.0) + amount + ) + return str(self._counters.get(counter_key, 0.0)) + + return _run + + +@pytest.mark.asyncio +async def test_async_increment_idempotent_same_dedup_id_increments_once( + monkeypatch, redis_no_ping +): + """ + Re-issuing async_increment_idempotent with the same dedup_id must advance + the counter only once — this is the core fix for the cluster-mode double-apply bug. + """ + monkeypatch.setenv("REDIS_HOST", "https://my-test-host") + redis_cache = RedisCache() + fake_store = _FakeRedisSpendStore() + redis_cache.async_register_script = fake_store.async_register_script # type: ignore[method-assign] + + await redis_cache.async_increment_idempotent( + key="spend:key:abc", value=0.01, dedup_id="req-1", refresh_ttl=True + ) + await redis_cache.async_increment_idempotent( + key="spend:key:abc", value=0.01, dedup_id="req-1", refresh_ttl=True + ) + + assert fake_store._counters.get("spend:key:abc") == pytest.approx( + 0.01 + ), "re-issued increment must not double-apply" + + +@pytest.mark.asyncio +async def test_async_increment_idempotent_different_dedup_ids_are_independent( + monkeypatch, redis_no_ping +): + """Two distinct requests (different dedup_id) must each apply once — INV-2.""" + monkeypatch.setenv("REDIS_HOST", "https://my-test-host") + redis_cache = RedisCache() + fake_store = _FakeRedisSpendStore() + redis_cache.async_register_script = fake_store.async_register_script # type: ignore[method-assign] + + await redis_cache.async_increment_idempotent( + key="spend:key:abc", value=0.01, dedup_id="req-1", refresh_ttl=True + ) + await redis_cache.async_increment_idempotent( + key="spend:key:abc", value=0.02, dedup_id="req-2", refresh_ttl=True + ) + + assert fake_store._counters.get("spend:key:abc") == pytest.approx( + 0.03 + ), "distinct dedup_ids must each contribute independently" + + +@pytest.mark.asyncio +async def test_async_increment_idempotent_script_error_falls_back_to_plain_increment( + monkeypatch, redis_no_ping +): + """On Lua script failure (NOSCRIPT / scripting disabled), fall back to + plain async_increment — INV-4.""" + monkeypatch.setenv("REDIS_HOST", "https://my-test-host") + redis_cache = RedisCache() + + async def _bad_script(keys, args): + raise RuntimeError("NOSCRIPT") + + redis_cache._dedup_increment_script = _bad_script + + plain_increment = AsyncMock(return_value=0.01) + redis_cache.async_increment = plain_increment # type: ignore[method-assign] + + result = await redis_cache.async_increment_idempotent( + key="spend:key:abc", value=0.01, dedup_id="req-1", refresh_ttl=True + ) + + assert plain_increment.called, "must fall back to async_increment on script error" + assert result == pytest.approx(0.01) + assert ( + redis_cache._dedup_increment_script_disabled + ), "script must be permanently disabled after failure to avoid re-registration loop" + + +@pytest.mark.asyncio +async def test_async_increment_idempotent_no_register_script_falls_back_to_plain( + monkeypatch, redis_no_ping +): + """When async_register_script is unavailable, fall back to plain async_increment — INV-4.""" + monkeypatch.setenv("REDIS_HOST", "https://my-test-host") + redis_cache = RedisCache() + redis_cache.async_register_script = None # type: ignore[assignment] + + plain_increment = AsyncMock(return_value=0.05) + redis_cache.async_increment = plain_increment # type: ignore[method-assign] + + result = await redis_cache.async_increment_idempotent( + key="spend:key:abc", value=0.05, dedup_id="req-1" + ) + + assert plain_increment.called + assert result == pytest.approx(0.05) diff --git a/tests/test_litellm/proxy/proxy_server/test_spend_counters.py b/tests/test_litellm/proxy/proxy_server/test_spend_counters.py index a839d82984c..5405017ae9d 100644 --- a/tests/test_litellm/proxy/proxy_server/test_spend_counters.py +++ b/tests/test_litellm/proxy/proxy_server/test_spend_counters.py @@ -53,6 +53,9 @@ def _make_spend_counter_cache( return_value=redis_increment_value, side_effect=redis_increment_side_effect, ) + cache.redis_cache.async_increment_idempotent = AsyncMock( + return_value=redis_increment_value, + ) cache.redis_cache.async_delete_cache = AsyncMock() cache.redis_cache.async_set_cache = AsyncMock() cache.redis_cache.async_set_max = AsyncMock() @@ -409,12 +412,12 @@ async def test_increment_spend_counters_increments_all_buckets(monkeypatch): ) observed = { - "redis_increment_called": fake_cache.redis_cache.async_increment.called, - "increment_calls": fake_cache.redis_cache.async_increment.call_count, + "redis_idempotent_increment_called": fake_cache.redis_cache.async_increment_idempotent.called, + "increment_calls": fake_cache.redis_cache.async_increment_idempotent.call_count, "user_cache_used": fake_user_cache.async_get_cache.called, } assert normalize(observed) == { - "redis_increment_called": True, + "redis_idempotent_increment_called": True, "increment_calls": 4, "user_cache_used": True, } @@ -441,6 +444,7 @@ async def test_increment_spend_counters_zero_cost_is_noop_finalizes_reservation( assert reservation == {"finalized": True} assert fake_cache.redis_cache.async_increment.called is False + assert fake_cache.redis_cache.async_increment_idempotent.called is False # --------------------------------------------------------------------------- @@ -514,9 +518,9 @@ async def test_increment_end_user_and_tag_spend_counters_increments_each_unique_ ) observed = { - "increment_calls": fake_cache.redis_cache.async_increment.call_count, + "increment_calls": fake_cache.redis_cache.async_increment_idempotent.call_count, "in_memory_set_calls": fake_cache.in_memory_cache.set_cache.call_count, - "called": fake_cache.redis_cache.async_increment.called, + "called": fake_cache.redis_cache.async_increment_idempotent.called, } assert normalize(observed) == { "increment_calls": 3, @@ -567,9 +571,9 @@ async def test_increment_org_spend_counter_increments_when_org_present(monkeypat ) observed = { - "increment_called": fake_cache.redis_cache.async_increment.called, - "increment_calls": fake_cache.redis_cache.async_increment.call_count, - "counter_key_arg": fake_cache.redis_cache.async_increment.call_args.kwargs[ + "increment_called": fake_cache.redis_cache.async_increment_idempotent.called, + "increment_calls": fake_cache.redis_cache.async_increment_idempotent.call_count, + "counter_key_arg": fake_cache.redis_cache.async_increment_idempotent.call_args.kwargs[ "key" ], } @@ -639,7 +643,7 @@ async def test_init_and_increment_unreserved_spend_counter_proceeds_when_not_res ) observed = { - "increment_called": fake_cache.redis_cache.async_increment.called, + "increment_called": fake_cache.redis_cache.async_increment_idempotent.called, "redis_get_called": fake_cache.redis_cache.async_get_cache.called, "reseed_consulted": True, } @@ -675,7 +679,7 @@ async def test_init_and_increment_spend_counter_warm_cache_skips_reseed(monkeypa observed = { "reseed_called": reseed.called, - "increment_called": fake_cache.redis_cache.async_increment.called, + "increment_called": fake_cache.redis_cache.async_increment_idempotent.called, "in_memory_seeded_from_redis": fake_cache.in_memory_cache.set_cache.called, } assert normalize(observed) == { @@ -714,8 +718,8 @@ async def test_init_and_increment_window_spend_counter_increments_when_initializ ) observed = { - "redis_increment_called": fake_cache.redis_cache.async_increment.called, - "increment_calls": fake_cache.redis_cache.async_increment.call_count, + "redis_increment_called": fake_cache.redis_cache.async_increment_idempotent.called, + "increment_calls": fake_cache.redis_cache.async_increment_idempotent.call_count, "in_memory_set_calls": fake_cache.in_memory_cache.set_cache.call_count, } assert normalize(observed) == { @@ -799,7 +803,7 @@ async def test_ensure_spend_counter_initialized_cold_seeds_from_source_cache( observed = { "source_cache_called": fake_user_cache.async_get_cache.called, - "seed_increment_called": fake_cache.redis_cache.async_increment.called, + "seed_increment_called": fake_cache.redis_cache.async_increment_idempotent.called, "warm_check_done": fake_cache.redis_cache.async_get_cache.called, } assert normalize(observed) == { @@ -965,12 +969,12 @@ async def test_increment_spend_counter_cache_redis_path_returns_new_value(monkey observed = { "result": result, - "redis_increment_called": fake_cache.redis_cache.async_increment.called, + "redis_idempotent_increment_called": fake_cache.redis_cache.async_increment_idempotent.called, "in_memory_set_called": fake_cache.in_memory_cache.set_cache.called, } assert normalize(observed) == { "result": 44.0, - "redis_increment_called": True, + "redis_idempotent_increment_called": True, "in_memory_set_called": True, } @@ -979,8 +983,9 @@ async def test_increment_spend_counter_cache_redis_path_returns_new_value(monkey async def test_increment_spend_counter_cache_redis_error_raises_and_invalidates( monkeypatch, ): - fake_cache = _make_spend_counter_cache( - redis_increment_side_effect=RuntimeError("incr fail") + fake_cache = _make_spend_counter_cache() + fake_cache.redis_cache.async_increment_idempotent = AsyncMock( + side_effect=RuntimeError("incr fail") ) monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) @@ -1084,3 +1089,224 @@ async def test_update_cache_user_cache_failure_invalid_state_is_swallowed(monkey ) assert result is None + + +# --------------------------------------------------------------------------- +# _increment_spend_counter_cache — idempotent routing (INV-3, INV-4, INV-5) +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_increment_spend_counter_cache_always_uses_idempotent( + monkeypatch, +): + """async_increment_idempotent is always taken when available — plain async_increment + is never used for regular counter increments.""" + fake_cache = _make_spend_counter_cache(redis_increment_value=0.01) + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + + await ps._increment_spend_counter_cache(counter_key="spend:key:abc", increment=0.01) + + assert ( + fake_cache.redis_cache.async_increment_idempotent.called + ), "idempotent method must always be called when available" + assert ( + not fake_cache.redis_cache.async_increment.called + ), "plain increment must not be called when idempotent path is available" + call_kwargs = fake_cache.redis_cache.async_increment_idempotent.call_args.kwargs + assert "dedup_id" in call_kwargs, "a server-generated dedup_id must be passed" + + +@pytest.mark.asyncio +async def test_increment_spend_counter_cache_uses_plain_when_idempotent_unavailable( + monkeypatch, +): + """When async_increment_idempotent is absent from the redis_cache, fall back + to plain async_increment.""" + fake_cache = _make_spend_counter_cache(redis_increment_value=0.01) + del fake_cache.redis_cache.async_increment_idempotent + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + + await ps._increment_spend_counter_cache(counter_key="spend:key:abc", increment=0.01) + + assert ( + fake_cache.redis_cache.async_increment.called + ), "plain increment must be used when idempotent method is absent" + + +@pytest.mark.asyncio +async def test_increment_spend_counters_uses_idempotent_for_all_buckets( + monkeypatch, +): + """increment_spend_counters uses async_increment_idempotent with a + server-generated UUID for each of the four main counter buckets (key, team, + team_member, user). Each call gets its own unique dedup_id.""" + fake_cache = _make_spend_counter_cache( + redis_get_value=None, redis_increment_value=5.0 + ) + fake_user_cache = _make_user_api_key_cache(get_value=None) + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache) + monkeypatch.setattr(ps, "prisma_client", None) + monkeypatch.setattr( + ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None) + ) + + await ps.increment_spend_counters( + token="hashed-tok", + team_id="t1", + user_id="u1", + response_cost=5.0, + ) + + assert fake_cache.redis_cache.async_increment_idempotent.called + assert not fake_cache.redis_cache.async_increment.called + dedup_ids = [ + call.kwargs["dedup_id"] + for call in fake_cache.redis_cache.async_increment_idempotent.call_args_list + ] + assert len(dedup_ids) == 4, "expected one idempotent increment per counter bucket" + assert len(set(dedup_ids)) == 4, "each counter increment must use a unique dedup_id" + + +@pytest.mark.asyncio +async def test_seed_fallback_uses_idempotent_increment(monkeypatch): + """The cold-seed fallback in _ensure_spend_counter_initialized calls + _increment_spend_counter_cache which always uses async_increment_idempotent + with a server-generated UUID.""" + fake_cache = _make_spend_counter_cache( + redis_get_value=None, redis_increment_value=10.0 + ) + fake_user_cache = _make_user_api_key_cache(get_value={"spend": 10.0}) + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache) + monkeypatch.setattr(ps, "prisma_client", None) + monkeypatch.setattr( + ps.SpendCounterReseed, + "coalesced", + AsyncMock(return_value=None), + ) + + await ps._ensure_spend_counter_initialized( + counter_key="spend:key:tok", source_cache_key="tok" + ) + + assert ( + fake_cache.redis_cache.async_increment_idempotent.called + ), "cold-seed fallback must use async_increment_idempotent" + assert not fake_cache.redis_cache.async_increment.called + + +# --------------------------------------------------------------------------- +# Proxy-level regression: reissued increment does not inflate the counter +# --------------------------------------------------------------------------- + + +class _FakeSpendRedisCache: + """ + Minimal Redis stand-in that models exactly-once dedup for the proxy-level + regression test. Implements the same SET NX-gated INCRBYFLOAT as the Lua + script but in plain Python. + """ + + def __init__(self): + self._counters: dict = {} + self._dedup: set = set() + + async def async_get_cache(self, key, **kwargs): + return self._counters.get(key) + + async def async_increment(self, key, value, refresh_ttl=False, **kwargs): + self._counters[key] = self._counters.get(key, 0.0) + value + return self._counters[key] + + async def async_increment_idempotent( + self, key, value, dedup_id, refresh_ttl=False, **kwargs + ): + dedup_key = f"{key}:dedup:{dedup_id}" + if dedup_key not in self._dedup: + self._dedup.add(dedup_key) + self._counters[key] = self._counters.get(key, 0.0) + value + return self._counters.get(key, 0.0) + + async def async_delete_cache(self, key, **kwargs): + self._counters.pop(key, None) + + +@pytest.mark.asyncio +async def test_two_independent_increments_each_apply_once(monkeypatch): + """ + Two separate calls to _increment_spend_counter_cache each generate a distinct + server-side UUID, so both increments are applied independently. The dedup + mechanism protects against redis-py internal retries (same UUID within one + call), not against two distinct spend events. + """ + fake_redis = _FakeSpendRedisCache() + fake_cache = MagicMock() + fake_cache.in_memory_cache = MagicMock() + fake_cache.in_memory_cache.set_cache = MagicMock() + fake_cache.redis_cache = fake_redis + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + + await ps._increment_spend_counter_cache(counter_key="spend:key:abc", increment=0.01) + await ps._increment_spend_counter_cache(counter_key="spend:key:abc", increment=0.01) + + spend = await ps.get_current_spend(counter_key="spend:key:abc", fallback_spend=0.0) + assert spend == pytest.approx( + 0.02 + ), f"two independent increments must each apply; got {spend}" + + +@pytest.mark.asyncio +async def test_window_spend_counter_uses_idempotent_path(monkeypatch): + """_init_and_increment_window_spend_counter uses async_increment_idempotent + with a server-generated UUID.""" + fake_cache = _make_spend_counter_cache( + redis_get_value=0.0, redis_increment_value=5.0 + ) + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + monkeypatch.setattr( + ps.SpendCounterReseed, + "coalesced_window", + AsyncMock(return_value=0.0), + ) + monkeypatch.setattr(ps, "prisma_client", None) + + await ps._init_and_increment_window_spend_counter( + counter_key="spend:key:tok:window:1d", + entity_type="Key", + entity_id="tok", + window_start=datetime(2024, 1, 1), + increment=0.05, + ) + + assert fake_cache.redis_cache.async_increment_idempotent.called + assert not fake_cache.redis_cache.async_increment.called + call_kwargs = fake_cache.redis_cache.async_increment_idempotent.call_args.kwargs + assert "dedup_id" in call_kwargs, "a server-generated dedup_id must be passed" + + +@pytest.mark.asyncio +async def test_end_user_and_tag_counters_use_idempotent_path(monkeypatch): + """_increment_end_user_and_tag_spend_counters uses async_increment_idempotent + with server-generated UUIDs, same as all other spend counter paths.""" + fake_cache = _make_spend_counter_cache( + redis_get_value=None, redis_increment_value=3.0 + ) + fake_user_cache = _make_user_api_key_cache() + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache) + monkeypatch.setattr(ps, "prisma_client", None) + monkeypatch.setattr( + ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None) + ) + + await ps._increment_end_user_and_tag_spend_counters( + end_user_id="eu-1", + tags=["tag-a"], + response_cost=0.03, + reserved_counter_keys=set(), + ) + + assert fake_cache.redis_cache.async_increment_idempotent.called + assert not fake_cache.redis_cache.async_increment.called diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 6017b9555e9..5987ccf5b7e 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -6273,6 +6273,7 @@ async def test_init_and_increment_spend_counter_reseeds_from_db_on_counter_miss( fake_redis = AsyncMock() fake_redis.async_increment = AsyncMock(side_effect=record_increment) + fake_redis.async_increment_idempotent = AsyncMock(side_effect=record_increment) fake_redis.async_get_cache = AsyncMock(return_value=None) # counter missing fake_redis.async_set_cache = AsyncMock(return_value=True) # SET NX wins counter_cache.redis_cache = fake_redis @@ -6569,6 +6570,7 @@ async def test_init_spend_counter_redis_clean_miss_skips_stale_in_memory(): fake_redis = AsyncMock() fake_redis.async_get_cache = AsyncMock(return_value=None) fake_redis.async_increment = AsyncMock(side_effect=redis_increment) + fake_redis.async_increment_idempotent = AsyncMock(side_effect=redis_increment) fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache) counter_cache.redis_cache = fake_redis @@ -6634,6 +6636,7 @@ async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory(): fake_redis.async_get_cache = AsyncMock(return_value=None) fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache) fake_redis.async_increment = AsyncMock(side_effect=redis_increment) + fake_redis.async_increment_idempotent = AsyncMock(side_effect=redis_increment) counter_cache.redis_cache = fake_redis fake_prisma = MagicMock() @@ -6698,6 +6701,7 @@ async def test_window_spend_counter_redis_concurrent_seed_does_not_double_seed() fake_redis.async_get_cache = AsyncMock(side_effect=redis_get_cache) fake_redis.async_set_cache = AsyncMock(return_value=False) fake_redis.async_increment = AsyncMock(side_effect=redis_increment) + fake_redis.async_increment_idempotent = AsyncMock(side_effect=redis_increment) counter_cache.redis_cache = fake_redis fake_prisma = MagicMock() @@ -6957,6 +6961,9 @@ async def test_increment_spend_counter_invalidates_stale_cache_on_redis_failure( counter_cache.in_memory_cache.set_cache(key="spend:team:redis-fail", value=4.0) fake_redis = AsyncMock() fake_redis.async_increment = AsyncMock(side_effect=RuntimeError("redis down")) + fake_redis.async_increment_idempotent = AsyncMock( + side_effect=RuntimeError("redis down") + ) fake_redis.async_delete_cache = AsyncMock() counter_cache.redis_cache = fake_redis