diff --git a/CHANGELOG.md b/CHANGELOG.md index 2fb24b979c..8c147d0a5c 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -25,6 +25,7 @@ This project adheres to [Semantic Versioning](https://semver.org/). - [#3986](https://github.com/plotly/dash/pull/3986) Adjust `_run_before_hooks` in the `fastapi` backend to honor a response returned by a `before_request` function, matching the `flask` backend's behavior. ### Fixed +- [#4027](https://github.com/plotly/dash/pull/4027) Support `async def` hook routes (`dash.hooks.route`) on all backends, stop hook routes from leaking the app context to the caller, and let hook routes and MCP read the request on the `fastapi` backend. - [#3980](https://github.com/plotly/dash/pull/3980) Fix the three `before_request` hooks (`Dash._setup_server` and the pages `router_sync` / `router_async`) publishing their "already done" guard flag before the setup work behind it had run. Under a multi-threaded WSGI worker such as `gunicorn -k gthread` (or under an ASGI worker for the async router), a second request arriving mid-setup could observe the flag already set, skip setup, then read `registered_paths` / `callback_map` / the pages router callback while they were still being registered - causing the first burst of component bundle requests after a restart to 500 with `Error loading dependency. "" is not a registered library`, or the pages router to hit `DuplicateCallback` when two workers raced past the guard. Each hook body now runs under a lock (`threading.Lock` for the two sync hooks, an `asyncio.Lock` bound to the running loop for the async router) and only publishes the flag after all work completes. Fixes [#3971](https://github.com/plotly/dash/issues/3971). - [#3944](https://github.com/plotly/dash/pull/3944) Fix `dash.testing` runner backend detection for wrapped FastAPI/Quart servers so threaded Flask-only options are not passed to ASGI runners. - [#3955](https://github.com/plotly/dash/pull/3955) Unpin `selenium` in the testing requirements (was capped at `<=4.2.0` from 2022) and require `>=4.11.0`, so it can drive current stable Chrome via Selenium Manager and stop the widespread CI flakiness. diff --git a/dash/_get_app.py b/dash/_get_app.py index ab0b897f81..3c6f8c2310 100644 --- a/dash/_get_app.py +++ b/dash/_get_app.py @@ -1,4 +1,5 @@ import functools +import inspect from contextvars import ContextVar, copy_context from textwrap import dedent @@ -32,10 +33,25 @@ async def wrap(self, *args, **kwargs): def with_app_context_factory(func, app): + if inspect.iscoroutinefunction(func): + + @functools.wraps(func) + async def async_wrap(*args, **kwargs): + # The coroutine runs in the awaiting task's context, so set the app + # there and restore it after, rather than in a copied context. + token = app_context.set(app) + try: + return await func(*args, **kwargs) + finally: + app_context.reset(token) + + return async_wrap + @functools.wraps(func) def wrap(*args, **kwargs): - app_context.set(app) + # Set the app in the copy only, so it doesn't leak to the caller. ctx = copy_context() + ctx.run(app_context.set, app) return ctx.run(func, *args, **kwargs) return wrap diff --git a/dash/_hooks.py b/dash/_hooks.py index f260b1fcb0..8a1504f3d1 100644 --- a/dash/_hooks.py +++ b/dash/_hooks.py @@ -126,6 +126,10 @@ def route( ): """ Add a route to the Dash server. + + The route function can be `async def`; with the Flask backend this + requires `flask[async]`. Read the request with + `dash.get_app().backend.request_adapter()`. """ def wrap(func: _t.Callable[[], _t.Any]): diff --git a/dash/backends/_fastapi.py b/dash/backends/_fastapi.py index 210f942a39..5e72f65570 100644 --- a/dash/backends/_fastapi.py +++ b/dash/backends/_fastapi.py @@ -148,6 +148,20 @@ def get_current_request() -> Request: _ENV_CONFIG = "_DASH_FASTAPI_CONFIG" +def _replay_body(body: bytes, receive: Receive) -> Receive: + """ASGI receive that sends an already read request body once.""" + sent = False + + async def replay(): + nonlocal sent + if not sent: + sent = True + return {"type": "http.request", "body": body, "more_body": False} + return await receive() + + return replay + + class DashMiddleware: # pylint: disable=too-few-public-methods """Consolidated middleware for all Dash/FastAPI integration needs.""" @@ -176,22 +190,24 @@ async def _initialize_dev_tools(self) -> None: self.dash_app.enable_dev_tools(**config, first_run=False) self._dev_tools_initialized = True - async def _setup_timing(self, request: Request) -> None: - """Set up timing information for the request.""" + async def _setup_timing(self, request: Request) -> bytes | None: + """Set up timing information for the request. + + Returns the request body when it had to be read to parse the JSON. + """ + body = None + request.state.json_body = None try: - request.state.json_body = ( - await request.json() - if request.headers.get("content-type", "").startswith( - "application/json" - ) - else None - ) + if request.headers.get("content-type", "").startswith("application/json"): + body = await request.body() + request.state.json_body = json.loads(body) except Exception: # pylint: disable=broad-exception-caught request.state.json_body = None if self.enable_timing: request.state.timing_information = { "__dash_server": {"dur": time.time(), "desc": None} } + return body async def _run_before_hooks(self) -> None: """Run all before-request hooks.""" @@ -270,7 +286,8 @@ async def _receive_with_shutdown(): await self.app(scope, receive, send) return - # Non-Dash routes pass through to avoid consuming body stream + # Non-Dash routes pass through to avoid consuming body stream. + # Routes registered through Dash (hook routes, MCP) are Dash routes too. path = scope["path"] prefix = self.dash_app.config.routes_pathname_prefix dash_prefix = prefix.rstrip("/") + "/_dash-" @@ -278,6 +295,7 @@ async def _receive_with_shutdown(): not path.startswith(dash_prefix) and path != prefix and path != prefix.rstrip("/") + and path not in self.dash_app.routes ): await self.app(scope, receive, send) return @@ -287,9 +305,12 @@ async def _receive_with_shutdown(): token = set_current_request(request) try: - await self._setup_timing(request) + body = await self._setup_timing(request) await self._run_before_hooks() + if body is not None: + # The body stream is consumed, replay it for handlers reading it. + receive = _replay_body(body, receive) await self.app(scope, receive, send) await self._run_after_hooks() diff --git a/tests/backend_tests/test_hook_routes.py b/tests/backend_tests/test_hook_routes.py new file mode 100644 index 0000000000..813af55468 --- /dev/null +++ b/tests/backend_tests/test_hook_routes.py @@ -0,0 +1,99 @@ +"""Hook routes on every backend: sync and async views reading the request +through the backend's request adapter, and FastAPI handlers reading the body +themselves after the Dash middleware parsed it.""" +import asyncio +import inspect + +import pytest + +from dash import Dash, get_app, hooks, html + + +@pytest.fixture(autouse=True) +def routes_cleanup(): + yield + hooks._ns["routes"] = [] + hooks._ns["setup"] = [] + + +@pytest.fixture(params=["flask", "quart", "fastapi"]) +def backend(request): + if request.param != "flask": + pytest.importorskip(request.param) + return request.param + + +def post(app, path, body): + """POST JSON with the test client of the app's backend.""" + server_type = app.backend.server_type + if server_type == "fastapi": + from starlette.testclient import TestClient + + with TestClient(app.server) as client: + response = client.post(path, json=body) + return response.status_code, response.json() + + if server_type == "quart": + + async def run(): + response = await app.server.test_client().post(path, json=body) + return response.status_code, await response.get_json() + + return asyncio.run(run()) + + response = app.server.test_client().post(path, json=body) + return response.status_code, response.get_json() + + +def make_app(backend): + app = Dash(__name__, backend=backend) + app.layout = html.Div() + return app + + +def test_hook_route_sync(backend): + if backend == "quart": + pytest.skip("Quart's request adapter get_json is async") + + @hooks.route("sync_echo", methods=("POST",)) + def sync_echo(): + adapter = get_app().backend.request_adapter() + return get_app().backend.jsonify({"echo": adapter.get_json()}) + + assert post(make_app(backend), "/sync_echo", {"a": 1}) == (200, {"echo": {"a": 1}}) + + +def test_hook_route_async(backend): + if backend == "flask": + pytest.importorskip("asgiref") + + @hooks.route("async_echo", methods=("POST",)) + async def async_echo(): + app = get_app() + data = app.backend.request_adapter().get_json() + if inspect.isawaitable(data): + data = await data + return app.backend.jsonify({"echo": data, "title": app.title}) + + app = make_app(backend) + app.title = "hook app" + # get_app() must return the app serving the request, not the last created. + make_app(backend).title = "other app" + assert post(app, "/async_echo", {"a": 1}) == ( + 200, + {"echo": {"a": 1}, "title": "hook app"}, + ) + + +def test_fastapi_route_reads_own_body(): + pytest.importorskip("fastapi") + from starlette.requests import Request + + @hooks.setup() + def add_route(app): + async def own_body(request: Request): + return app.backend.jsonify({"own": await request.json()}) + + app._add_url("own_body", own_body, ["POST"]) + + assert post(make_app("fastapi"), "/own_body", {"a": 1}) == (200, {"own": {"a": 1}}) diff --git a/tests/integration/test_hooks.py b/tests/integration/test_hooks.py index 518439d819..c244c04f3d 100644 --- a/tests/integration/test_hooks.py +++ b/tests/integration/test_hooks.py @@ -10,7 +10,7 @@ def hook_cleanup(): yield hooks._ns["layout"] = [] hooks._ns["setup"] = [] - hooks._ns["route"] = [] + hooks._ns["routes"] = [] hooks._ns["error"] = [] hooks._ns["callback"] = [] hooks._ns["index"] = []