Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
15 changes: 9 additions & 6 deletions fastapi_cachex/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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``.

Expand All @@ -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
Expand Down
19 changes: 19 additions & 0 deletions tests/test_cache_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@

import asyncio
from collections.abc import AsyncGenerator
from functools import partial
from typing import TYPE_CHECKING
from typing import Any

Expand Down Expand Up @@ -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,
Expand Down
Loading