diff --git a/CHANGELOG.md b/CHANGELOG.md index 149ed3d..cdfb8ed 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -40,6 +40,14 @@ Note that 0.3.3 was never released; 0.3.4 follows 0.3.2. ### Fixed +- A sync (`def`) handler under `@cache` runs in the threadpool again. The + cache wrapper is `async`, so FastAPI stopped offloading the handler and + `@cache` called it on the event loop, where blocking I/O stalled every other + request. A handler whose call returns an awaitable (an object with an + `async def __call__`, or a sync callable returning a coroutine) is now + awaited instead of failing to encode. + ([#100](https://github.com/allen0099/FastAPI-CacheX/issues/100)) + - 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/cache.py b/fastapi_cachex/cache.py index 52ed75a..3f86f58 100644 --- a/fastapi_cachex/cache.py +++ b/fastapi_cachex/cache.py @@ -21,6 +21,7 @@ from fastapi import Request from fastapi import Response +from starlette.concurrency import run_in_threadpool from starlette.status import HTTP_200_OK from starlette.status import HTTP_206_PARTIAL_CONTENT from starlette.status import HTTP_300_MULTIPLE_CHOICES @@ -303,6 +304,19 @@ async def _render( return response, body, None if body is None else _etag_for(body) +def _is_coroutine_callable(func: HandlerCallable) -> bool: + """Report whether calling `func` returns a coroutine. + + `inspect.iscoroutinefunction` already sees through `functools.partial`; + an instance with an ``async def __call__`` needs its method checked. The + method is looked up on the type, as the call itself does, so a class + (whose type is ``type``) counts as sync. + """ + return inspect.iscoroutinefunction(func) or inspect.iscoroutinefunction( + type(func).__call__ + ) + + async def get_response( __func: HandlerCallable, __request: Request, @@ -310,11 +324,20 @@ async def get_response( *args: Any, **kwargs: Any, ) -> Response: - """Get the response from the function.""" - if inspect.iscoroutinefunction(__func): - result = await __func(*args, **kwargs) + """Get the response from the function. + + Coroutine handlers are awaited. Sync handlers run in the threadpool, as + FastAPI would run them without the (async) cache wrapper, so blocking I/O + in a ``def`` handler does not stall the event loop. + """ + if _is_coroutine_callable(__func): + result = await cast("Callable[..., Awaitable[object]]", __func)(*args, **kwargs) else: - result = __func(*args, **kwargs) + result = await run_in_threadpool(__func, *args, **kwargs) + # A sync callable can still hand back an awaitable (a lambda wrapping a + # coroutine function, say); await it rather than try to encode it. + if inspect.isawaitable(result): + result = await result # If already a Response object, return it directly if isinstance(result, Response): diff --git a/tests/test_cache.py b/tests/test_cache.py index be27708..2e5a9c5 100644 --- a/tests/test_cache.py +++ b/tests/test_cache.py @@ -1,4 +1,6 @@ +import threading from collections.abc import AsyncGenerator +from functools import partial import pytest from fastapi import FastAPI @@ -210,6 +212,66 @@ def sync_endpoint(): assert response.status_code == 200 +def test_sync_handler_runs_off_the_event_loop_thread(): + threads: dict[str, int] = {} + local_app = FastAPI() + + @local_app.get("/loop") + @cache() + async def on_loop() -> dict[str, str]: + threads["async"] = threading.get_ident() + return {"ok": "async"} + + @local_app.get("/blocking") + @cache() + def blocking() -> dict[str, str]: + threads["sync"] = threading.get_ident() + return {"ok": "sync"} + + with TestClient(local_app) as local_client: + assert local_client.get("/loop").json() == {"ok": "async"} + assert local_client.get("/blocking").json() == {"ok": "sync"} + + # Inline, the sync handler would run on the event loop's thread and block + # every other request while it waits. + assert threads["sync"] != threads["async"] + + +async def _greet(name: str) -> dict[str, str]: + return {"hello": name} + + +class _AsyncCallable: + async def __call__(self) -> dict[str, str]: + return {"hello": "callable"} + + +def _returns_awaitable() -> object: + return _greet("lambda") + + +@pytest.mark.parametrize( + ("handler", "expected"), + [ + (partial(_greet, "partial"), {"hello": "partial"}), + (_AsyncCallable(), {"hello": "callable"}), + (_returns_awaitable, {"hello": "lambda"}), + ], + ids=["partial", "async-call", "sync-returning-awaitable"], +) +def test_handler_returning_a_coroutine_is_awaited( + handler: object, expected: dict[str, str] +) -> None: + local_app = FastAPI() + local_app.get("/greet")(cache(ttl=60)(handler)) # type: ignore[arg-type] + + with TestClient(local_app) as local_client: + response = local_client.get("/greet") + + assert response.status_code == 200 + assert response.json() == expected + + def test_no_cache_with_revalidate(): @app.get("/no-cache-revalidate") @cache(no_cache=True, must_revalidate=True)