diff --git a/CHANGELOG.md b/CHANGELOG.md index cdfb8ed..25a2e8f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -48,6 +48,10 @@ Note that 0.3.3 was never released; 0.3.4 follows 0.3.2. awaited instead of failing to encode. ([#100](https://github.com/allen0099/FastAPI-CacheX/issues/100)) +- `CacheManager.get_or_set()` awaits whatever the factory returns when it is + awaitable. The documented `get_or_set(key, lambda: load_user(42))` form used + to store the coroutine itself and fail with `TypeError`. ([#101](https://github.com/allen0099/FastAPI-CacheX/issues/101)) + - The monitoring routes from `add_routes()` now show when Redis entries expire. `AsyncRedisCacheBackend.get_cache_data()` reported every entry as never expiring (`ttl_remaining: null`); it now fetches each key's `PTTL` in diff --git a/fastapi_cachex/manager.py b/fastapi_cachex/manager.py index f92d763..03cc4b8 100644 --- a/fastapi_cachex/manager.py +++ b/fastapi_cachex/manager.py @@ -165,8 +165,10 @@ async def get_or_set( Args: key: Logical cache key (without the manager's prefix). - factory: Zero-argument callable (sync or async) that produces the - JSON-serializable value to cache on a miss. + factory: Zero-argument callable that produces the JSON-serializable + value to cache on a miss. If it returns an awaitable (an async + function, or a lambda or ``functools.partial`` wrapping one), + the result is awaited. ttl: Time-to-live in seconds for a newly created value. If None, uses ``self.default_ttl``. @@ -181,10 +183,11 @@ async def get_or_set( if cached is not sentinel: return cached - if inspect.iscoroutinefunction(factory): - value = await factory() - else: - value = factory() + # Await whatever comes back awaitable, not just from coroutine + # functions: `lambda: load(42)` and `functools.partial` return one too. + value = factory() + if inspect.isawaitable(value): + value = await value await self.set(key, value, ttl=ttl) return value diff --git a/tests/test_cache_manager.py b/tests/test_cache_manager.py index a514e1d..588ca20 100644 --- a/tests/test_cache_manager.py +++ b/tests/test_cache_manager.py @@ -2,6 +2,7 @@ import asyncio from collections.abc import AsyncGenerator +from functools import partial from typing import TYPE_CHECKING from typing import Any @@ -194,6 +195,24 @@ async def factory() -> str: assert await cache_manager.get("nope") == "async_value" +@pytest.mark.asyncio +@pytest.mark.parametrize("wrap", ["lambda", "partial"]) +async def test_get_or_set_awaits_a_sync_callable_returning_an_awaitable( + cache_manager: CacheManager, wrap: str +) -> None: + """The documented `lambda: load_user(42)` form must store the result, not the coroutine.""" + + async def load_user(user_id: int) -> dict[str, int]: + return {"id": user_id} + + factory = (lambda: load_user(42)) if wrap == "lambda" else partial(load_user, 42) + + result = await cache_manager.get_or_set("user", factory) + + assert result == {"id": 42} + assert await cache_manager.get("user") == {"id": 42} + + @pytest.mark.asyncio async def test_get_or_set_honors_ttl_on_created_value( cache_manager: CacheManager,