mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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
This commit is contained in:
parent
31eca17007
commit
63f10eef27
5 changed files with 470 additions and 23 deletions
|
|
@ -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)}"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue