Skip to content

Commit 91dfcf2

Browse files
committed
fix(cache-manager): await an awaitable returned by the get_or_set factory
get_or_set() awaited the factory only when it was a coroutine function, so the documented lambda: load_user(42) form stored the coroutine and failed with TypeError. The factory is now called and its result awaited whenever it is awaitable. Closes #101
1 parent 4ecb77b commit 91dfcf2

3 files changed

Lines changed: 31 additions & 6 deletions

File tree

‎CHANGELOG.md‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -40,6 +40,9 @@ Note that 0.3.3 was never released; 0.3.4 follows 0.3.2.
4040

4141
### Fixed
4242

43+
- `CacheManager.get_or_set()` awaits whatever the factory returns when it is
44+
awaitable. The documented `get_or_set(key, lambda: load_user(42))` form used
45+
to store the coroutine itself and fail with `TypeError`. ([#101](https://github.com/allen0099/FastAPI-CacheX/issues/101))
4346
- The monitoring routes from `add_routes()` now show when Redis entries
4447
expire. `AsyncRedisCacheBackend.get_cache_data()` reported every entry as
4548
never expiring (`ttl_remaining: null`); it now fetches each key's `PTTL` in

‎fastapi_cachex/manager.py‎

Lines changed: 9 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -165,8 +165,10 @@ async def get_or_set(
165165
166166
Args:
167167
key: Logical cache key (without the manager's prefix).
168-
factory: Zero-argument callable (sync or async) that produces the
169-
JSON-serializable value to cache on a miss.
168+
factory: Zero-argument callable that produces the JSON-serializable
169+
value to cache on a miss. If it returns an awaitable (an async
170+
function, or a lambda or ``functools.partial`` wrapping one),
171+
the result is awaited.
170172
ttl: Time-to-live in seconds for a newly created value. If None,
171173
uses ``self.default_ttl``.
172174
@@ -181,10 +183,11 @@ async def get_or_set(
181183
if cached is not sentinel:
182184
return cached
183185

184-
if inspect.iscoroutinefunction(factory):
185-
value = await factory()
186-
else:
187-
value = factory()
186+
# Await whatever comes back awaitable, not just from coroutine
187+
# functions: `lambda: load(42)` and `functools.partial` return one too.
188+
value = factory()
189+
if inspect.isawaitable(value):
190+
value = await value
188191

189192
await self.set(key, value, ttl=ttl)
190193
return value

‎tests/test_cache_manager.py‎

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@
22

33
import asyncio
44
from collections.abc import AsyncGenerator
5+
from functools import partial
56
from typing import TYPE_CHECKING
67
from typing import Any
78

@@ -194,6 +195,24 @@ async def factory() -> str:
194195
assert await cache_manager.get("nope") == "async_value"
195196

196197

198+
@pytest.mark.asyncio
199+
@pytest.mark.parametrize("wrap", ["lambda", "partial"])
200+
async def test_get_or_set_awaits_a_sync_callable_returning_an_awaitable(
201+
cache_manager: CacheManager, wrap: str
202+
) -> None:
203+
"""The documented `lambda: load_user(42)` form must store the result, not the coroutine."""
204+
205+
async def load_user(user_id: int) -> dict[str, int]:
206+
return {"id": user_id}
207+
208+
factory = (lambda: load_user(42)) if wrap == "lambda" else partial(load_user, 42)
209+
210+
result = await cache_manager.get_or_set("user", factory)
211+
212+
assert result == {"id": 42}
213+
assert await cache_manager.get("user") == {"id": 42}
214+
215+
197216
@pytest.mark.asyncio
198217
async def test_get_or_set_honors_ttl_on_created_value(
199218
cache_manager: CacheManager,

0 commit comments

Comments
 (0)