diff --git a/lnbits/utils/cache.py b/lnbits/utils/cache.py index 904447e19..6e9dfa149 100644 --- a/lnbits/utils/cache.py +++ b/lnbits/utils/cache.py @@ -25,6 +25,8 @@ class Cache: def __init__(self, interval: float = 10) -> None: self.interval = interval self._values: dict[Any, Cached] = {} + self._refreshing: set[str] = set() + self._tasks: set[asyncio.Task] = set() def value(self, key: str) -> Cached | None: return self._values.get(key) @@ -49,15 +51,35 @@ class Cache: async def save_result(self, coro, key: str, expiry: float = 10): """ - If `key` exists, return its value, otherwise call coro and cache its result + Stale-while-revalidate: return stale value immediately and refresh in + the background. Only blocks on a true cold start (no prior value). """ - cached = self.get(key) - if cached: - return cached - else: + cached = self._values.get(key) + if cached is not None: + if cached.expiry > time(): + return cached.value + # stale: serve old value and refresh in background (one task at a time) + if key not in self._refreshing: + self._refreshing.add(key) + # extend expiry now to prevent a stampede of background tasks + self._values[key] = Cached(cached.value, time() + expiry) + task = asyncio.create_task(self._refresh(coro, key, expiry)) + self._tasks.add(task) + task.add_done_callback(self._tasks.discard) + return cached.value + # cold start: must wait for the first value + value = await coro() + self.set(key, value, expiry=expiry) + return value + + async def _refresh(self, coro, key: str, expiry: float): + try: value = await coro() self.set(key, value, expiry=expiry) - return value + except Exception: + logger.error(f"Error refreshing cache key {key}") + finally: + self._refreshing.discard(key) async def invalidate_forever(self): while settings.lnbits_running: diff --git a/tests/unit/test_cache.py b/tests/unit/test_cache.py index 2f33ffe1d..b7dd76cb4 100644 --- a/tests/unit/test_cache.py +++ b/tests/unit/test_cache.py @@ -90,6 +90,61 @@ async def test_cache_pop_expired_returns_default(cache): assert cache.pop(key, default="fallback") == "fallback" +@pytest.mark.anyio +async def test_cache_coro_stale_returns_immediately(cache): + """Stale entry is served immediately; background refresh updates the value.""" + calls = 0 + + async def test(): + nonlocal calls + calls += 1 + return calls + + # cold start + result = await cache.save_result(test, key="test", expiry=0.01) + assert result == 1 + + # let the entry expire + await asyncio.sleep(0.02) + + # stale-while-revalidate: returns old value immediately + result = await cache.save_result(test, key="test", expiry=0.5) + assert result == 1 # stale value returned, not the new one + + # allow background refresh to complete + await asyncio.sleep(0.05) + assert calls == 2 + # now the cache has the fresh value + result = await cache.save_result(test, key="test", expiry=0.5) + assert result == 2 + + +@pytest.mark.anyio +async def test_cache_coro_no_stampede(cache): + """Multiple concurrent requests on a stale entry spawn only one refresh.""" + calls = 0 + + async def slow_fetch(): + nonlocal calls + calls += 1 + await asyncio.sleep(0.05) + return calls + + await cache.save_result(slow_fetch, key="test", expiry=0.01) + await asyncio.sleep(0.02) + + # fire multiple concurrent requests while stale + results = await asyncio.gather( + cache.save_result(slow_fetch, key="test", expiry=0.5), + cache.save_result(slow_fetch, key="test", expiry=0.5), + cache.save_result(slow_fetch, key="test", expiry=0.5), + ) + + await asyncio.sleep(0.1) # let the single background task finish + assert all(r == 1 for r in results) # all got stale value + assert calls == 2 # cold start + exactly one background refresh + + @pytest.mark.anyio async def test_invalidate_forever_logs_and_recovers_from_errors( settings: Settings, mocker: MockerFixture