This is an automated email from the ASF dual-hosted git repository.

ashb pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git


The following commit(s) were added to refs/heads/main by this push:
     new bc9ca5b1227 Correctly shutdown async sessions on exit. (#73838)
bc9ca5b1227 is described below

commit bc9ca5b1227651a3ec6a84b1b9ae482867695568
Author: Ash Berlin-Taylor <[email protected]>
AuthorDate: Thu Oct 1 13:02:37 2026 +0100

    Correctly shutdown async sessions on exit. (#73838)
    
    This isn't "a problem" per se, as the process is about to exit anyway, but 
this does end up with a confusing/scary looking message in the API server logs 
of:
    
            Traceback (most recent call last):
            File 
"/usr/python/lib/python3.12/site-packages/sqlalchemy/pool/base.py", line 375, 
in _close_connection
                self._dialect.do_close(connection)
            File 
"/usr/python/lib/python3.12/site-packages/sqlalchemy/engine/default.py", line 
721, in do_close
                dbapi_connection.close()
            File 
"/usr/python/lib/python3.12/site-packages/sqlalchemy/dialects/sqlite/aiosqlite.py",
 line 362, in close
                self._handle_exception(error)
            File 
"/usr/python/lib/python3.12/site-packages/sqlalchemy/dialects/sqlite/aiosqlite.py",
 line 373, in _handle_exception
                raise error
            File 
"/usr/python/lib/python3.12/site-packages/sqlalchemy/dialects/sqlite/aiosqlite.py",
 line 350, in close
                self.await_(self._connection.close())
            File 
"/usr/python/lib/python3.12/site-packages/sqlalchemy/util/_concurrency_py3k.py",
 line 123, in await_only
                raise exc.MissingGreenlet(
            sqlalchemy.exc.MissingGreenlet: greenlet_spawn has not been called; 
can't call await_only() here. Was IO attempted in an unexpected place? 
(Background on this error at: https://sqlalche.me/e/20/xd2s)
    
    This also appears in some unit tests (though only visible when a test 
fails, as otherwise logs don't get shown), hence the unit test fixture changes
---
 airflow-core/src/airflow/api_fastapi/app.py        |   2 +
 .../src/airflow/api_fastapi/execution_api/app.py   |   2 +
 airflow-core/src/airflow/settings.py               |  14 +--
 .../api_fastapi/auth/managers/simple/conftest.py   |   3 +-
 airflow-core/tests/unit/api_fastapi/conftest.py    |  27 +++--
 .../core_api/routes/public/test_auth.py            |  28 +++--
 .../core_api/routes/public/test_backfills.py       |  20 ++--
 .../core_api/routes/public/test_connections.py     |   9 +-
 .../core_api/routes/public/test_dag_bundles.py     |  43 ++++----
 .../core_api/routes/public/test_dag_parsing.py     |  18 +---
 .../core_api/routes/public/test_dag_run.py         |  78 +++++---------
 .../core_api/routes/public/test_task_instances.py  | 111 +++++++++----------
 .../unit/api_fastapi/execution_api/test_app.py     |  84 ++++++++++++++-
 airflow-core/tests/unit/api_fastapi/test_app.py    | 108 +++++++++++++++++--
 airflow-core/tests/unit/core/test_settings.py      | 120 +++++++++++++++++++++
 airflow-core/tests/unit/state/test_metastore.py    |   9 ++
 airflow-core/tests/unit/utils/test_session.py      |  18 ++--
 17 files changed, 480 insertions(+), 214 deletions(-)

diff --git a/airflow-core/src/airflow/api_fastapi/app.py 
b/airflow-core/src/airflow/api_fastapi/app.py
index 524cb32cac2..65121ab53d0 100644
--- a/airflow-core/src/airflow/api_fastapi/app.py
+++ b/airflow-core/src/airflow/api_fastapi/app.py
@@ -27,6 +27,7 @@ from fastapi import FastAPI
 from fastapi.routing import Mount
 from starlette.middleware import Middleware
 
+from airflow import settings
 from airflow.api_fastapi.common.dagbag import create_dag_bag
 from airflow.api_fastapi.common.exceptions import init_error_handlers
 from airflow.api_fastapi.common.http_access_log import HttpAccessLogMiddleware
@@ -100,6 +101,7 @@ def _initialize_api_server_stats() -> None:
 async def lifespan(app: FastAPI):
     _initialize_api_server_stats()
     async with AsyncExitStack() as stack:
+        stack.push_async_callback(settings.dispose_async_engine)
         for route in app.routes:
             if isinstance(route, Mount) and isinstance(route.app, FastAPI):
                 await stack.enter_async_context(
diff --git a/airflow-core/src/airflow/api_fastapi/execution_api/app.py 
b/airflow-core/src/airflow/api_fastapi/execution_api/app.py
index 53be7e5af67..4b881cda03c 100644
--- a/airflow-core/src/airflow/api_fastapi/execution_api/app.py
+++ b/airflow-core/src/airflow/api_fastapi/execution_api/app.py
@@ -38,6 +38,7 @@ from fastapi.routing import APIRoute
 from opentelemetry import context as otel_context, propagate as otel_propagate
 from starlette.middleware.base import BaseHTTPMiddleware
 
+from airflow import settings
 from airflow.api_fastapi.auth.tokens import (
     JWTGenerator,
     JWTValidator,
@@ -437,6 +438,7 @@ class InProcessExecutionAPI:
 
         # https://github.com/abersheeran/a2wsgi/discussions/64
         async def start_lifespan(cm: AsyncExitStack, app: FastAPI):
+            cm.push_async_callback(settings.dispose_async_engine)
             await cm.enter_async_context(app.router.lifespan_context(app))
 
         cm = AsyncExitStack()
diff --git a/airflow-core/src/airflow/settings.py 
b/airflow-core/src/airflow/settings.py
index 08d4fbbfe72..0ba2ec16b59 100644
--- a/airflow-core/src/airflow/settings.py
+++ b/airflow-core/src/airflow/settings.py
@@ -427,13 +427,7 @@ def create_async_metadata_engine(
 
 
 def _configure_async_session() -> None:
-    """
-    Configure async SQLAlchemy session.
-
-    This exists so tests can reconfigure the session. How SQLAlchemy configures
-    this does not work well with Pytest and you can end up with issues when the
-    session and runs in a different event loop from the test itself.
-    """
+    """Configure the async engine and session factory."""
     global AsyncSession, async_engine
 
     if not SQL_ALCHEMY_CONN_ASYNC:
@@ -670,6 +664,12 @@ def dispose_orm(do_log: bool = True):
         AsyncSession = None
 
 
+async def dispose_async_engine() -> None:
+    """Close checked-in connections on their event loop, retaining the engine 
and session factory."""
+    if async_engine is not None:
+        await async_engine.dispose()
+
+
 def reconfigure_orm(disable_connection_pool=False, pool_class=None):
     """Properly close database connections and re-configure ORM."""
     dispose_orm()
diff --git 
a/airflow-core/tests/unit/api_fastapi/auth/managers/simple/conftest.py 
b/airflow-core/tests/unit/api_fastapi/auth/managers/simple/conftest.py
index 122e8a35cbe..2ed7643d583 100644
--- a/airflow-core/tests/unit/api_fastapi/auth/managers/simple/conftest.py
+++ b/airflow-core/tests/unit/api_fastapi/auth/managers/simple/conftest.py
@@ -75,4 +75,5 @@ def test_client():
             ): 
"airflow.api_fastapi.auth.managers.simple.simple_auth_manager.SimpleAuthManager"
         }
     ):
-        return TestClient(create_app("core"))
+        with TestClient(create_app("core")) as client:
+            yield client
diff --git a/airflow-core/tests/unit/api_fastapi/conftest.py 
b/airflow-core/tests/unit/api_fastapi/conftest.py
index f275d48f967..931178faf12 100644
--- a/airflow-core/tests/unit/api_fastapi/conftest.py
+++ b/airflow-core/tests/unit/api_fastapi/conftest.py
@@ -133,11 +133,12 @@ def _authed_test_client(app: FastAPI, request):
             ),
         )
     with mock.patch("airflow.models.revoked_token.RevokedToken.is_revoked", 
return_value=False):
-        yield TestClient(
+        with TestClient(
             app,
             headers={"Authorization": f"Bearer {token}"},
             base_url=f"{BASE_URL}{get_api_path(request)}",
-        )
+        ) as test_client:
+            yield test_client
 
 
 @pytest.fixture
@@ -167,22 +168,28 @@ def fresh_test_client(request):
 
 @pytest.fixture
 def unauthenticated_test_client(request, _isolated_shared_app):
-    return TestClient(_isolated_shared_app, 
base_url=f"{BASE_URL}{get_api_path(request)}")
+    with TestClient(_isolated_shared_app, 
base_url=f"{BASE_URL}{get_api_path(request)}") as test_client:
+        yield test_client
 
 
 @pytest.fixture
-def unauthorized_test_client(request, _isolated_shared_app):
-    app = _isolated_shared_app
-    auth_manager: SimpleAuthManager = app.state.auth_manager
+def unauthorized_headers(_isolated_shared_app):
+    auth_manager: SimpleAuthManager = _isolated_shared_app.state.auth_manager
     token = auth_manager._get_token_signer().generate(
         auth_manager.serialize_user(SimpleAuthManagerUser(username="dummy", 
role=None))
     )
+    return {"Authorization": f"Bearer {token}"}
+
+
[email protected]
+def unauthorized_test_client(request, _isolated_shared_app, 
unauthorized_headers):
     with mock.patch("airflow.models.revoked_token.RevokedToken.is_revoked", 
return_value=False):
-        yield TestClient(
-            app,
-            headers={"Authorization": f"Bearer {token}"},
+        with TestClient(
+            _isolated_shared_app,
+            headers=unauthorized_headers,
             base_url=f"{BASE_URL}{get_api_path(request)}",
-        )
+        ) as test_client:
+            yield test_client
 
 
 @pytest.fixture
diff --git 
a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_auth.py 
b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_auth.py
index b2d616129e4..a36ab808701 100644
--- a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_auth.py
+++ b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_auth.py
@@ -22,6 +22,7 @@ from urllib.parse import parse_qs, urlencode
 
 import jwt
 import pytest
+from fastapi.testclient import TestClient
 
 from airflow.api_fastapi.auth.managers.base_auth_manager import 
COOKIE_NAME_JWT_TOKEN
 from airflow.models.revoked_token import RevokedToken
@@ -169,10 +170,7 @@ class TestLogoutTokenRevocation:
         clear_db_revoked_tokens()
 
     @pytest.fixture
-    def logout_client(self):
-        """A test client without the is_revoked mock so revocation tests hit 
the real DB."""
-        from fastapi.testclient import TestClient
-
+    def logout_app(self):
         from airflow.api_fastapi.app import create_app
 
         with conf_vars(
@@ -183,8 +181,13 @@ class TestLogoutTokenRevocation:
                 ): 
"airflow.api_fastapi.auth.managers.simple.simple_auth_manager.SimpleAuthManager"
             }
         ):
-            app = create_app()
-            yield TestClient(app, base_url="http://testserver/api/v2";)
+            yield create_app()
+
+    @pytest.fixture
+    def logout_client(self, logout_app):
+        """A test client without the is_revoked mock so revocation tests hit 
the real DB."""
+        with TestClient(logout_app, base_url="http://testserver/api/v2";) as 
client:
+            yield client
 
     def test_logout_revokes_token(self, logout_client):
         """Test that logout revokes the JWT token and persists it in the 
database."""
@@ -282,7 +285,7 @@ class TestLogoutTokenRevocation:
         assert RevokedToken.is_revoked("test-jti-both-bearer") is True
         assert RevokedToken.is_revoked("test-jti-both-cookie") is True
 
-    def test_logout_revokes_both_even_when_a_trusted_user_is_cached(self, 
logout_client):
+    def test_logout_revokes_both_even_when_a_trusted_user_is_cached(self, 
logout_app):
         """The trusted-middleware shortcut must not change what logout revokes.
 
         On protected routes `get_user()` can return a user cached by 
JWTRefreshMiddleware
@@ -291,7 +294,7 @@ class TestLogoutTokenRevocation:
         """
         from airflow.api_fastapi.core_api.security import 
USER_INJECTED_BY_TRUSTED_MIDDLEWARE
 
-        auth_manager = logout_client.app.state.auth_manager
+        auth_manager = logout_app.state.auth_manager
         bearer_token = self._mint(auth_manager, "test-jti-trusted-bearer")
         cookie_token = self._mint(auth_manager, "test-jti-trusted-cookie")
 
@@ -300,9 +303,12 @@ class TestLogoutTokenRevocation:
             request.state.user_authenticated_via = 
USER_INJECTED_BY_TRUSTED_MIDDLEWARE
             return await call_next(request)
 
-        logout_client.app.middleware("http")(_inject)
-        logout_client.cookies.set(COOKIE_NAME_JWT_TOKEN, cookie_token)
-        with patch.object(auth_manager, "get_url_logout", return_value=None):
+        logout_app.middleware("http")(_inject)
+        with (
+            TestClient(logout_app, base_url="http://testserver/api/v2";) as 
logout_client,
+            patch.object(auth_manager, "get_url_logout", return_value=None),
+        ):
+            logout_client.cookies.set(COOKIE_NAME_JWT_TOKEN, cookie_token)
             response = logout_client.get(
                 "/auth/logout",
                 headers={"Authorization": f"Bearer {bearer_token}"},
diff --git 
a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_backfills.py 
b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_backfills.py
index eeebcd6a220..d67fc96be0b 100644
--- 
a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_backfills.py
+++ 
b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_backfills.py
@@ -23,7 +23,6 @@ from unittest import mock
 
 import pendulum
 import pytest
-from fastapi.testclient import TestClient
 from sqlalchemy import and_, func, select
 from sqlalchemy.exc import OperationalError, ProgrammingError
 
@@ -83,18 +82,13 @@ def clean_db():
 
 
 @pytest.fixture
-def dag_reader_test_client(test_client):
+def dag_reader_headers(test_client):
     """A caller who may read the Dags but not write them: viewer is below the 
role edits require."""
     auth_manager = test_client.app.state.auth_manager
     token = auth_manager._get_token_signer().generate(
         auth_manager.serialize_user(SimpleAuthManagerUser(username="reader", 
role="viewer"))
     )
-    with mock.patch("airflow.models.revoked_token.RevokedToken.is_revoked", 
return_value=False):
-        yield TestClient(
-            test_client.app,
-            headers={"Authorization": f"Bearer {token}"},
-            base_url=str(test_client.base_url),
-        )
+    return {"Authorization": f"Bearer {token}"}
 
 
 def make_dags():
@@ -1599,22 +1593,24 @@ class TestPauseBackfill(TestBackfillEndpoint):
         response = 
unauthorized_test_client.put(f"/backfills/{backfill.id}/pause")
         assert response.status_code == 404
 
-    def test_pause_backfill_403(self, session, dag_reader_test_client):
+    def test_pause_backfill_403(self, session, dag_reader_headers, 
test_client):
         (dag,) = self._create_dag_models()
         from_date = timezone.utcnow()
         to_date = timezone.utcnow()
         backfill = Backfill(dag_id=dag.dag_id, from_date=from_date, 
to_date=to_date)
         session.add(backfill)
         session.commit()
-        response = 
dag_reader_test_client.put(f"/backfills/{backfill.id}/pause")
+        response = test_client.put(f"/backfills/{backfill.id}/pause", 
headers=dag_reader_headers)
         assert response.status_code == 403
 
     def test_pause_backfill_unknown_id_is_not_authorized_by_a_body_dag_id(
-        self, session, dag_reader_test_client
+        self, session, dag_reader_headers, test_client
     ):
         (dag,) = self._create_dag_models()
         session.commit()
-        response = dag_reader_test_client.put(f"/backfills/{231984098}/pause", 
json={"dag_id": dag.dag_id})
+        response = test_client.put(
+            f"/backfills/{231984098}/pause", json={"dag_id": dag.dag_id}, 
headers=dag_reader_headers
+        )
         assert response.status_code == 404
         assert response.json().get("detail") == "Backfill not found"
 
diff --git 
a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_connections.py
 
b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_connections.py
index 3f94a3a763e..cb94ac41c1e 100644
--- 
a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_connections.py
+++ 
b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_connections.py
@@ -1732,16 +1732,15 @@ class TestAsyncConnectionTest(TestConnectionEndpoint):
         assert response.status_code == 422
 
     @mock.patch.dict(os.environ, {"AIRFLOW__CORE__TEST_CONNECTION": "Enabled"})
-    def test_get_status_unauthorized_user_does_not_leak_row(
-        self, test_client, unauthorized_test_client, session
-    ):
+    def test_get_status_unauthorized_user_does_not_leak_row(self, test_client, 
unauthorized_headers, session):
         """A user without rights on the conn_id never sees the row payload via 
GET-by-token."""
         post_response = test_client.post("/connections/enqueue-test", 
json=self.TEST_REQUEST_BODY)
         assert post_response.status_code == 202
         token = post_response.json()["token"]
 
-        response = unauthorized_test_client.get(
-            "/connections/enqueue-test", 
headers={"Airflow-Connection-Test-Token": token}
+        response = test_client.get(
+            "/connections/enqueue-test",
+            headers={**unauthorized_headers, "Airflow-Connection-Test-Token": 
token},
         )
         assert response.status_code in (401, 403, 404)
         body = (
diff --git 
a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_dag_bundles.py
 
b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_dag_bundles.py
index d28bbcb63e8..b01506a0a6d 100644
--- 
a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_dag_bundles.py
+++ 
b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_dag_bundles.py
@@ -21,7 +21,6 @@ from typing import TYPE_CHECKING
 from unittest import mock
 
 import pytest
-from fastapi.testclient import TestClient
 from itsdangerous import URLSafeSerializer
 from sqlalchemy import insert, update
 
@@ -275,7 +274,7 @@ def dag_scoped_client(test_client, readable_dag_ids):
 
 
 @pytest.fixture
-def viewer_client(test_client, readable_dag_ids):
+def viewer_headers(test_client, readable_dag_ids):
     """
     A viewer with the same readable Dags: may read import errors, but not the 
admin-gated view.
 
@@ -286,12 +285,7 @@ def viewer_client(test_client, readable_dag_ids):
     token = auth_manager._get_token_signer().generate(
         auth_manager.serialize_user(SimpleAuthManagerUser(username="viewer", 
role="viewer"))
     )
-    with mock.patch("airflow.models.revoked_token.RevokedToken.is_revoked", 
return_value=False):
-        yield TestClient(
-            test_client.app,
-            headers={"Authorization": f"Bearer {token}"},
-            base_url=str(test_client.base_url),
-        )
+    return {"Authorization": f"Bearer {token}"}
 
 
 class TestGetDagBundles:
@@ -325,7 +319,7 @@ class TestGetDagBundles:
         assert DAGLESS_BUNDLE not in [bundle["name"] for bundle in 
body["dag_bundles"]]
         assert body["total_entries"] == 3
 
-    def test_hides_a_bundle_with_no_dags_from_a_viewer(self, viewer_client):
+    def test_hides_a_bundle_with_no_dags_from_a_viewer(self, viewer_headers, 
test_client):
         """
         The bundle name and its version are the disclosure, so a viewer must 
not get them.
 
@@ -333,7 +327,7 @@ class TestGetDagBundles:
         with no Dag to authorize against, there is nothing weaker than the 
admin view to fall back
         on.
         """
-        body = viewer_client.get("/dagBundles").json()
+        body = test_client.get("/dagBundles", headers=viewer_headers).json()
 
         assert DAGLESS_BUNDLE not in [bundle["name"] for bundle in 
body["dag_bundles"]]
         # Absent from the count too, so its existence does not leak through 
pagination.
@@ -381,12 +375,12 @@ class TestGetDagBundles:
 
         assert GIT_BUNDLE in [bundle["name"] for bundle in body["dag_bundles"]]
 
-    def test_returns_nothing_when_no_dag_is_readable(self, viewer_client):
+    def test_returns_nothing_when_no_dag_is_readable(self, viewer_headers, 
test_client):
         # The fixture already patched this attribute, so retarget its mock 
rather than nesting a
         # second autospec patch over it.
-        
viewer_client.app.state.auth_manager.get_authorized_dag_ids.return_value = set()
+        test_client.app.state.auth_manager.get_authorized_dag_ids.return_value 
= set()
 
-        body = viewer_client.get("/dagBundles").json()
+        body = test_client.get("/dagBundles", headers=viewer_headers).json()
 
         assert body == {"dag_bundles": [], "total_entries": 0}
 
@@ -483,7 +477,9 @@ class TestGetDagBundles:
         assert bundle["active"] is False
         assert bundle["version"] == "deadbeef"
 
-    def 
test_import_error_count_authorizes_on_the_same_terms_as_import_errors(self, 
viewer_client):
+    def test_import_error_count_authorizes_on_the_same_terms_as_import_errors(
+        self, viewer_headers, test_client
+    ):
         """
         Count on the same terms as ``GET /importErrors``, not "every row for 
this bundle".
 
@@ -498,7 +494,7 @@ class TestGetDagBundles:
 
         Dropping either restriction takes the count to 2, so one assertion 
pins both halves.
         """
-        body = viewer_client.get("/dagBundles").json()
+        body = test_client.get("/dagBundles", headers=viewer_headers).json()
         bundle = next(b for b in body["dag_bundles"] if b["name"] == 
GIT_BUNDLE)
 
         assert bundle["import_error_count"] == 1
@@ -755,7 +751,7 @@ class TestGetDagBundle:
     def test_404_for_an_unknown_bundle(self, dag_scoped_client):
         assert dag_scoped_client.get("/dagBundles/no_such_bundle").status_code 
== 404
 
-    def test_import_error_count_is_gated_like_the_collection(self, 
admin_client, viewer_client):
+    def test_import_error_count_is_gated_like_the_collection(self, 
admin_client, viewer_headers):
         """
         The admin sees the unregistered-file error as well; the viewer sees 
only the registered one.
 
@@ -763,7 +759,10 @@ class TestGetDagBundle:
         cannot drift from the collection route it shares a helper with.
         """
         assert 
admin_client.get(f"/dagBundles/{GIT_BUNDLE}").json()["import_error_count"] == 2
-        assert 
viewer_client.get(f"/dagBundles/{GIT_BUNDLE}").json()["import_error_count"] == 1
+        assert (
+            admin_client.get(f"/dagBundles/{GIT_BUNDLE}", 
headers=viewer_headers).json()["import_error_count"]
+            == 1
+        )
 
     def test_import_error_count_is_withheld_without_permission(self, 
admin_client):
         auth_manager = admin_client.app.state.auth_manager
@@ -856,14 +855,14 @@ class TestGetDagBundleFiles:
         assert by_path[UNREGISTERED_FILE]["last_parsed_time"] is None
         assert by_path[UNREGISTERED_FILE]["last_parse_duration"] is None
 
-    def test_excludes_a_file_whose_dag_is_not_readable(self, admin_client, 
viewer_client):
+    def test_excludes_a_file_whose_dag_is_not_readable(self, admin_client, 
viewer_headers):
         """``UNREADABLE_FILE`` is registered, so only the readable-Dag filter 
keeps it out."""
-        for client in (admin_client, viewer_client):
-            body = client.get(f"/dagBundles/{GIT_BUNDLE}/files").json()
+        for headers in ({}, viewer_headers):
+            body = admin_client.get(f"/dagBundles/{GIT_BUNDLE}/files", 
headers=headers).json()
             assert UNREADABLE_FILE not in {file["relative_fileloc"] for file 
in body["dag_bundle_files"]}
 
-    def test_viewer_does_not_see_the_unregistered_file(self, viewer_client):
-        body = viewer_client.get(f"/dagBundles/{GIT_BUNDLE}/files").json()
+    def test_viewer_does_not_see_the_unregistered_file(self, viewer_headers, 
test_client):
+        body = test_client.get(f"/dagBundles/{GIT_BUNDLE}/files", 
headers=viewer_headers).json()
 
         assert [file["relative_fileloc"] for file in body["dag_bundle_files"]] 
== [REGISTERED_FILE]
         assert body["total_entries"] == 1
diff --git 
a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_dag_parsing.py
 
b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_dag_parsing.py
index 7f645f81cc3..91bc6c75828 100644
--- 
a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_dag_parsing.py
+++ 
b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_dag_parsing.py
@@ -16,10 +16,7 @@
 # under the License.
 from __future__ import annotations
 
-from unittest import mock
-
 import pytest
-from fastapi.testclient import TestClient
 from sqlalchemy import select
 
 from airflow.api_fastapi.auth.managers.simple.user import SimpleAuthManagerUser
@@ -44,18 +41,13 @@ TEST_MULTIPLE_DAGS_ID = "asset_produces_1"
 
 
 @pytest.fixture
-def dag_reader_test_client(test_client):
+def dag_reader_headers(test_client):
     """A caller who may read the Dags (and import errors) but not edit them: 
viewer is below the role edits require."""
     auth_manager = test_client.app.state.auth_manager
     token = auth_manager._get_token_signer().generate(
         auth_manager.serialize_user(SimpleAuthManagerUser(username="reader", 
role="viewer"))
     )
-    with mock.patch("airflow.models.revoked_token.RevokedToken.is_revoked", 
return_value=False):
-        yield TestClient(
-            test_client.app,
-            headers={"Authorization": f"Bearer {token}"},
-            base_url=str(test_client.base_url),
-        )
+    return {"Authorization": f"Bearer {token}"}
 
 
 class TestDagParsingEndpoint:
@@ -155,7 +147,7 @@ class TestDagParsingEndpoint:
         assert session.scalars(select(DagPriorityParsingRequest)).all() == []
 
     def 
test_reparse_import_error_file_forbidden_for_basic_import_errors_viewer(
-        self, url_safe_serializer, session, dag_reader_test_client
+        self, url_safe_serializer, session, dag_reader_headers, test_client
     ):
         # Reparsing a file with no registered Dag requires the dedicated 
REPARSE_ALL permission
         # (admin-by-default), so a caller who can view the import-errors list 
(basic IMPORT_ERRORS)
@@ -166,8 +158,8 @@ class TestDagParsingEndpoint:
             {"bundle_name": "some_bundle", "relative_fileloc": 
"dags/broken.py"}
         )
 
-        response = dag_reader_test_client.put(
-            f"/parseDagFile/{token}", headers={"Accept": "application/json"}
+        response = test_client.put(
+            f"/parseDagFile/{token}", headers={"Accept": "application/json", 
**dag_reader_headers}
         )
 
         assert response.status_code == 403
diff --git 
a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_dag_run.py 
b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_dag_run.py
index 075bfb0ede2..eeee6f196ef 100644
--- a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_dag_run.py
+++ b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_dag_run.py
@@ -24,7 +24,6 @@ from unittest import mock
 
 import pytest
 import time_machine
-from fastapi.testclient import TestClient
 from sqlalchemy import delete, func, select, update
 
 from airflow import plugins_manager
@@ -43,7 +42,6 @@ from airflow.models.team import Team
 from airflow.models.xcom import XComModel
 from airflow.providers.standard.operators.empty import EmptyOperator
 from airflow.sdk import Asset, Param, result, task
-from airflow.settings import _configure_async_session
 from airflow.timetables.interval import CronDataIntervalTimetable
 from airflow.timetables.simple import PartitionedAssetTimetable, 
PartitionedAtRuntime
 from airflow.timetables.trigger import CronPartitionTimetable
@@ -2565,24 +2563,17 @@ class TestBulkClearDagRuns:
                 SimpleAuthManagerUser(username="limited-user", role="user", 
teams=[]),
             )
         )
-        with (
-            mock.patch("airflow.models.revoked_token.RevokedToken.is_revoked", 
return_value=False),
-            TestClient(
-                test_client.app,
-                headers={"Authorization": f"Bearer {token}"},
-                base_url=str(test_client.base_url),
-            ) as limited_test_client,
-        ):
-            response = limited_test_client.post(
-                "/dags/~/clearDagRuns",
-                json={
-                    "dry_run": False,
-                    "dag_runs": [
-                        {"dag_id": DAG1_ID, "dag_run_id": DAG1_RUN1_ID},
-                        {"dag_id": DAG2_ID, "dag_run_id": DAG2_RUN1_ID},
-                    ],
-                },
-            )
+        response = test_client.post(
+            "/dags/~/clearDagRuns",
+            json={
+                "dry_run": False,
+                "dag_runs": [
+                    {"dag_id": DAG1_ID, "dag_run_id": DAG1_RUN1_ID},
+                    {"dag_id": DAG2_ID, "dag_run_id": DAG2_RUN1_ID},
+                ],
+            },
+            headers={"Authorization": f"Bearer {token}"},
+        )
 
         assert response.status_code == 403
         # The batched auth check rejects the whole request, so the authorized 
Dag's run is not cleared either.
@@ -4414,16 +4405,6 @@ class TestResolveRunOnLatestVersion:
 
 
 class TestWaitDagRun:
-    # The way we init async engine does not work well with FastAPI app init.
-    # Creating the engine implicitly creates an event loop, which Airflow does
-    # once for the entire process; creating the FastAPI app also does, but our
-    # test setup does it once for each test. I don't know how to properly fix
-    # this without rewriting how Airflow does db; re-configuring the db for 
each
-    # test at least makes the tests run correctly.
-    @pytest.fixture(autouse=True)
-    def reconfigure_async_db_engine(self):
-        _configure_async_session()
-
     def test_should_respond_401(self, unauthenticated_test_client):
         response = unauthenticated_test_client.get(
             f"/dags/{DAG1_ID}/dagRuns/{DAG1_RUN1_ID}/wait",
@@ -4913,28 +4894,21 @@ class TestBulkDagRuns:
                 SimpleAuthManagerUser(username="limited-user", role="user", 
teams=[]),
             )
         )
-        with (
-            mock.patch("airflow.models.revoked_token.RevokedToken.is_revoked", 
return_value=False),
-            TestClient(
-                test_client.app,
-                headers={"Authorization": f"Bearer {token}"},
-                base_url=str(test_client.base_url),
-            ) as limited_test_client,
-        ):
-            response = limited_test_client.patch(
-                self.WILDCARD_ENDPOINT,
-                json={
-                    "actions": [
-                        {
-                            "action": "delete",
-                            "entities": [
-                                {"dag_id": DAG1_ID, "dag_run_id": 
DAG1_RUN1_ID},
-                                {"dag_id": DAG2_ID, "dag_run_id": 
DAG2_RUN1_ID},
-                            ],
-                        }
-                    ]
-                },
-            )
+        response = test_client.patch(
+            self.WILDCARD_ENDPOINT,
+            json={
+                "actions": [
+                    {
+                        "action": "delete",
+                        "entities": [
+                            {"dag_id": DAG1_ID, "dag_run_id": DAG1_RUN1_ID},
+                            {"dag_id": DAG2_ID, "dag_run_id": DAG2_RUN1_ID},
+                        ],
+                    }
+                ]
+            },
+            headers={"Authorization": f"Bearer {token}"},
+        )
 
         assert response.status_code == 403
         session.expire_all()
diff --git 
a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_task_instances.py
 
b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_task_instances.py
index 44d8a65f2f7..265bab04ca7 100644
--- 
a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_task_instances.py
+++ 
b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_task_instances.py
@@ -27,7 +27,6 @@ from unittest import mock
 
 import pendulum
 import pytest
-from fastapi.testclient import TestClient
 from sqlalchemy import delete, func, select, update
 from sqlalchemy.orm import joinedload
 
@@ -7085,38 +7084,31 @@ class TestBulkTaskInstances(TestTaskInstanceEndpoint):
                 SimpleAuthManagerUser(username="limited-user", role="user", 
teams=[]),
             )
         )
-        with (
-            mock.patch("airflow.models.revoked_token.RevokedToken.is_revoked", 
return_value=False),
-            TestClient(
-                test_client.app,
-                headers={"Authorization": f"Bearer {token}"},
-                base_url=str(test_client.base_url),
-            ) as limited_test_client,
-        ):
-            response = limited_test_client.patch(
-                self.WILDCARD_ENDPOINT,
-                json={
-                    "actions": [
-                        {
-                            "action": "update",
-                            "entities": [
-                                {
-                                    "dag_id": self.BASH_DAG_ID,
-                                    "dag_run_id": self.RUN_ID,
-                                    "task_id": self.BASH_TASK_ID,
-                                    "new_state": "success",
-                                },
-                                {
-                                    "dag_id": self.DAG_ID,
-                                    "dag_run_id": self.RUN_ID,
-                                    "task_id": self.TASK_ID,
-                                    "new_state": "success",
-                                },
-                            ],
-                        }
-                    ]
-                },
-            )
+        response = test_client.patch(
+            self.WILDCARD_ENDPOINT,
+            json={
+                "actions": [
+                    {
+                        "action": "update",
+                        "entities": [
+                            {
+                                "dag_id": self.BASH_DAG_ID,
+                                "dag_run_id": self.RUN_ID,
+                                "task_id": self.BASH_TASK_ID,
+                                "new_state": "success",
+                            },
+                            {
+                                "dag_id": self.DAG_ID,
+                                "dag_run_id": self.RUN_ID,
+                                "task_id": self.TASK_ID,
+                                "new_state": "success",
+                            },
+                        ],
+                    }
+                ]
+            },
+            headers={"Authorization": f"Bearer {token}"},
+        )
 
         assert response.status_code == 200
         assert response.json()["update"]["success"] == 
[f"{self.DAG_ID}.{self.RUN_ID}.{self.TASK_ID}[-1]"]
@@ -7160,36 +7152,29 @@ class TestBulkTaskInstances(TestTaskInstanceEndpoint):
                 SimpleAuthManagerUser(username="limited-user", role="user", 
teams=[]),
             )
         )
-        with (
-            mock.patch("airflow.models.revoked_token.RevokedToken.is_revoked", 
return_value=False),
-            TestClient(
-                test_client.app,
-                headers={"Authorization": f"Bearer {token}"},
-                base_url=str(test_client.base_url),
-            ) as limited_test_client,
-        ):
-            response = limited_test_client.patch(
-                self.WILDCARD_ENDPOINT,
-                json={
-                    "actions": [
-                        {
-                            "action": "delete",
-                            "entities": [
-                                {
-                                    "dag_id": self.BASH_DAG_ID,
-                                    "dag_run_id": self.RUN_ID,
-                                    "task_id": self.BASH_TASK_ID,
-                                },
-                                {
-                                    "dag_id": self.DAG_ID,
-                                    "dag_run_id": self.RUN_ID,
-                                    "task_id": self.TASK_ID,
-                                },
-                            ],
-                        }
-                    ]
-                },
-            )
+        response = test_client.patch(
+            self.WILDCARD_ENDPOINT,
+            json={
+                "actions": [
+                    {
+                        "action": "delete",
+                        "entities": [
+                            {
+                                "dag_id": self.BASH_DAG_ID,
+                                "dag_run_id": self.RUN_ID,
+                                "task_id": self.BASH_TASK_ID,
+                            },
+                            {
+                                "dag_id": self.DAG_ID,
+                                "dag_run_id": self.RUN_ID,
+                                "task_id": self.TASK_ID,
+                            },
+                        ],
+                    }
+                ]
+            },
+            headers={"Authorization": f"Bearer {token}"},
+        )
 
         assert response.status_code == 200
         assert response.json()["delete"]["success"] == 
[f"{self.DAG_ID}.{self.RUN_ID}.{self.TASK_ID}[-1]"]
diff --git a/airflow-core/tests/unit/api_fastapi/execution_api/test_app.py 
b/airflow-core/tests/unit/api_fastapi/execution_api/test_app.py
index bb2d2d557dc..3938b490fec 100644
--- a/airflow-core/tests/unit/api_fastapi/execution_api/test_app.py
+++ b/airflow-core/tests/unit/api_fastapi/execution_api/test_app.py
@@ -19,18 +19,21 @@ from __future__ import annotations
 import asyncio
 import gc
 import threading
+from contextlib import asynccontextmanager
 from unittest import mock
 from uuid import UUID
 
 import httpx
 import pytest
-from fastapi import Request, status
+from fastapi import FastAPI, Request, status
 from fastapi.params import Security as SecurityParam
 from fastapi.routing import APIRoute
 from fastapi.testclient import TestClient
 from opentelemetry import context as otel_context, propagate as otel_propagate
+from sqlalchemy import event, text
 from sqlalchemy.exc import SQLAlchemyError
 
+from airflow import settings
 from airflow.api_fastapi.execution_api.app import (
     InProcessExecutionAPI,
     _extract_w3c_trace_context,
@@ -40,6 +43,7 @@ from 
airflow.api_fastapi.execution_api.datamodels.taskinstance import TaskInstan
 from airflow.api_fastapi.execution_api.datamodels.token import TIClaims, 
TIToken
 from airflow.api_fastapi.execution_api.security import require_auth
 from airflow.api_fastapi.execution_api.versions import bundle
+from airflow.utils.session import create_session_async
 
 from tests_common.test_utils.config import conf_vars
 
@@ -193,6 +197,84 @@ def test_in_process_execution_api_transport_lifecycle():
     assert not thread.is_alive()
 
 
[email protected]
+def in_process_db_app():
+    engine = settings.async_engine
+    opened, closed = [], []
+    app = FastAPI()
+
+    def record_connect(connection, record):
+        opened.append((connection, asyncio.get_running_loop()))
+
+    def record_close(connection, record):
+        closed.append((connection, asyncio.get_running_loop()))
+
+    @app.get("/")
+    async def query():
+        async with create_session_async() as session:
+            return (await session.execute(text("SELECT 1"))).scalar_one()
+
+    event.listen(engine.sync_engine, "connect", record_connect)
+    event.listen(engine.sync_engine, "close", record_close)
+    try:
+        yield app, opened, closed
+    finally:
+        event.remove(engine.sync_engine, "connect", record_connect)
+        event.remove(engine.sync_engine, "close", record_close)
+
+
[email protected]("fail_shutdown", [False, True])
+def 
test_in_process_shutdown_closes_connections_after_lifespan(in_process_db_app, 
fail_shutdown):
+    app, opened, closed = in_process_db_app
+    shutdown_loops = []
+
+    @asynccontextmanager
+    async def lifespan(app):
+        yield
+        async with create_session_async() as session:
+            assert (await session.execute(text("SELECT 1"))).scalar_one() == 1
+        assert closed == []
+        shutdown_loops.append(asyncio.get_running_loop())
+        if fail_shutdown:
+            raise RuntimeError("shutdown failed")
+
+    app.router.lifespan_context = lifespan
+    api = InProcessExecutionAPI(app)
+    with httpx.Client(transport=api.transport) as client:
+        assert client.get("http://localhost/";).json() == 1
+    del client, api
+    gc.collect()
+
+    assert len(opened) == 1
+    assert closed == opened
+    assert shutdown_loops == [opened[0][1]]
+
+
+def 
test_session_factory_remains_usable_after_in_process_shutdown(in_process_db_app):
+    app, opened, closed = in_process_db_app
+    engine, factory = settings.async_engine, settings.AsyncSession
+    api = InProcessExecutionAPI(app)
+    with httpx.Client(transport=api.transport) as client:
+        assert client.get("http://localhost/";).json() == 1
+    del client, api
+    gc.collect()
+
+    assert settings.async_engine is engine
+    assert settings.AsyncSession is factory
+
+    async def query_after_shutdown():
+        try:
+            async with create_session_async() as session:
+                assert (await session.execute(text("SELECT 1"))).scalar_one() 
== 1
+        finally:
+            await settings.dispose_async_engine()
+
+    asyncio.run(query_after_shutdown())
+    assert len(opened) == 2
+    assert opened[0][1] is not opened[1][1]
+    assert closed == opened
+
+
 class TestCorrelationIdMiddleware:
     def test_correlation_id_echoed_in_response_headers(self, client):
         """Test that correlation-id from request is echoed back in response 
headers."""
diff --git a/airflow-core/tests/unit/api_fastapi/test_app.py 
b/airflow-core/tests/unit/api_fastapi/test_app.py
index 19174592e65..de74dbc8d7c 100644
--- a/airflow-core/tests/unit/api_fastapi/test_app.py
+++ b/airflow-core/tests/unit/api_fastapi/test_app.py
@@ -16,21 +16,109 @@
 # under the License.
 from __future__ import annotations
 
+import asyncio
 import threading
+from contextlib import asynccontextmanager
 from unittest import mock
 
 import pytest
 from fastapi import FastAPI
+from fastapi.testclient import TestClient
+from sqlalchemy import event, text
+from sqlalchemy.engine import Engine
 
 import airflow.api_fastapi.app as app_module
 import airflow.plugins_manager as plugins_manager
+from airflow import settings
 from airflow.api_fastapi.common.http_access_log import HttpAccessLogMiddleware
+from airflow.utils.session import create_session_async
 
 from tests_common.test_utils.config import conf_vars
 
 pytestmark = pytest.mark.db_test
 
 
[email protected]
+def async_db_app():
+    app = FastAPI(lifespan=app_module.lifespan)
+    opened = []
+    closed = []
+
+    def record_connect(connection, record):
+        opened.append((connection, asyncio.get_running_loop()))
+
+    def record_close(connection, record):
+        closed.append((connection, asyncio.get_running_loop()))
+
+    event.listen(Engine, "connect", record_connect)
+    event.listen(Engine, "close", record_close)
+
+    @app.get("/")
+    async def query(fail: bool = False):
+        async with create_session_async() as session:
+            value = (await session.execute(text("SELECT 1"))).scalar_one()
+            if fail:
+                raise RuntimeError("request failed")
+            return value
+
+    try:
+        yield app, opened, closed
+    finally:
+        event.remove(Engine, "connect", record_connect)
+        event.remove(Engine, "close", record_close)
+
+
+def 
test_async_connections_are_reused_and_disposed_on_the_client_loop(async_db_app):
+    app, opened, closed = async_db_app
+    configured_engine = settings.async_engine
+    configured_factory = settings.AsyncSession
+    for count in (1, 2):
+        with TestClient(app) as client:
+            assert settings.async_engine is configured_engine
+            assert settings.AsyncSession is configured_factory
+            assert client.get("/").json() == 1
+            assert client.get("/").json() == 1
+            assert len(opened) == count
+            assert len(closed) == count - 1
+        assert closed == opened
+        assert settings.async_engine is configured_engine
+        assert settings.AsyncSession is configured_factory
+    configured_engine.sync_engine.dispose()
+    assert opened[0][1] is not opened[1][1]
+
+
+def test_async_pool_is_disposed_after_a_request_error(async_db_app):
+    app, opened, closed = async_db_app
+    with pytest.raises(RuntimeError, match="request failed"):
+        with TestClient(app) as client:
+            client.get("/?fail=true")
+    assert len(opened) == 1
+    assert closed == opened
+
+
[email protected]("fail_at", ["startup", "shutdown"])
+def test_async_pool_is_disposed_after_a_lifespan_error(async_db_app, fail_at):
+    app, opened, closed = async_db_app
+
+    @asynccontextmanager
+    async def lifespan(app):
+        async with create_session_async() as session:
+            await session.execute(text("SELECT 1"))
+        if fail_at == "startup":
+            raise RuntimeError("startup failed")
+        yield
+        async with create_session_async() as session:
+            await session.execute(text("SELECT 1"))
+        raise RuntimeError("shutdown failed")
+
+    app.mount("/child", FastAPI(lifespan=lifespan))
+    with pytest.raises(RuntimeError, match=f"{fail_at} failed"):
+        with TestClient(app):
+            pass
+    assert len(opened) == 1
+    assert closed == opened
+
+
 def test_main_app_lifespan(client):
     with client() as test_client:
         test_app = test_client.app
@@ -43,8 +131,8 @@ def test_main_app_lifespan(client):
 @mock.patch("airflow.api_fastapi.app.init_views")
 @mock.patch("airflow.api_fastapi.app.init_plugins")
 @mock.patch("airflow.api_fastapi.app.create_task_execution_api_app")
-def test_core_api_app(mock_create_task_exec_api, mock_init_plugins, 
mock_init_views, client):
-    test_app = client(apps="core").app
+def test_core_api_app(mock_create_task_exec_api, mock_init_plugins, 
mock_init_views):
+    test_app = app_module.create_app(apps="core")
 
     # Assert that core-related functions were called
     mock_init_views.assert_called_once_with(test_app)
@@ -57,8 +145,8 @@ def test_core_api_app(mock_create_task_exec_api, 
mock_init_plugins, mock_init_vi
 @mock.patch("airflow.api_fastapi.app.init_views")
 @mock.patch("airflow.api_fastapi.app.init_plugins")
 @mock.patch("airflow.api_fastapi.app.create_task_execution_api_app")
-def test_execution_api_app(mock_create_task_exec_api, mock_init_plugins, 
mock_init_views, client):
-    client(apps="execution")
+def test_execution_api_app(mock_create_task_exec_api, mock_init_plugins, 
mock_init_views):
+    app_module.create_app(apps="execution")
 
     # Assert that execution-related functions were called
     mock_create_task_exec_api.assert_called_once()
@@ -78,8 +166,8 @@ def test_execution_api_app_lifespan(client, 
get_execution_app):
 @mock.patch("airflow.api_fastapi.app.init_views")
 @mock.patch("airflow.api_fastapi.app.init_plugins")
 @mock.patch("airflow.api_fastapi.app.create_task_execution_api_app")
-def test_all_apps(mock_create_task_exec_api, mock_init_plugins, 
mock_init_views, client):
-    test_app = client(apps="all").app
+def test_all_apps(mock_create_task_exec_api, mock_init_plugins, 
mock_init_views):
+    test_app = app_module.create_app(apps="all")
 
     # Assert that core-related functions were called
     mock_init_views.assert_called_once_with(test_app)
@@ -90,25 +178,25 @@ def test_all_apps(mock_create_task_exec_api, 
mock_init_plugins, mock_init_views,
 
 
 @pytest.mark.parametrize("apps", ["all", "core", "execution"])
-def 
test_access_log_middleware_installed_outermost_for_every_apps_selection(apps, 
client):
+def 
test_access_log_middleware_installed_outermost_for_every_apps_selection(apps):
     """Both server backends disable their own access logger, so a selection 
that skips this
     middleware has no access logging at all. It must also stay outermost so it 
times the full
     request including inner middlewares (GZip compression in particular — see 
#60165); the
     test default config has no CORS so index 0 is HttpAccessLogMiddleware."""
-    installed = [m.cls for m in client(apps=apps).app.user_middleware]
+    installed = [m.cls for m in 
app_module.create_app(apps=apps).user_middleware]
 
     assert installed.count(HttpAccessLogMiddleware) == 1
     assert installed[0] is HttpAccessLogMiddleware
 
 
-def test_catch_all_route_last(client):
+def test_catch_all_route_last():
     """
     Ensure the catch all route that returns the initial html is the last route 
in the fastapi app.
 
     If it's not, it results in any routes/apps added afterwards to not be 
reachable, as the catch all
     route responds instead.
     """
-    test_app = client(apps="all").app
+    test_app = app_module.create_app(apps="all")
     assert test_app.routes[-1].path == "/{rest_of_path:path}"
 
 
diff --git a/airflow-core/tests/unit/core/test_settings.py 
b/airflow-core/tests/unit/core/test_settings.py
index be9e34f611a..c9ecb61309d 100644
--- a/airflow-core/tests/unit/core/test_settings.py
+++ b/airflow-core/tests/unit/core/test_settings.py
@@ -17,10 +17,13 @@
 # under the License.
 from __future__ import annotations
 
+import asyncio
 import contextlib
 import os
+import subprocess
 import sys
 import tempfile
+import textwrap
 from unittest import mock
 from unittest.mock import MagicMock, call, patch
 
@@ -218,6 +221,11 @@ class TestLocalSettings:
 class TestMetadataEngineHooks:
     """Tests for the overridable create_metadata_engine / 
create_async_metadata_engine hooks."""
 
+    @pytest.fixture(autouse=True)
+    def isolate_orm(self, monkeypatch):
+        for attr in ("engine", "Session", "NonScopedSession", "async_engine", 
"AsyncSession"):
+            monkeypatch.setattr(settings, attr, getattr(settings, attr))
+
     def setup_method(self):
         self.old_modules = dict(sys.modules)
         from airflow import settings
@@ -588,3 +596,115 @@ class TestDisposeOrm:
             settings.dispose_orm(do_log=False)
 
         mock_close.assert_not_called()
+
+
+class TestDisposeAsyncEngine:
+    @pytest.fixture(autouse=True)
+    def isolate_async_orm(self, monkeypatch):
+        monkeypatch.setattr(settings, "async_engine", None)
+        monkeypatch.setattr(settings, "AsyncSession", None)
+
+    def test_disposal_without_an_async_engine_is_a_noop(self):
+        asyncio.run(settings.dispose_async_engine())
+        assert settings.async_engine is None
+        assert settings.AsyncSession is None
+
+    def test_disposes_async_pool_without_changing_sync_resources(self, 
monkeypatch):
+        engine = mock.create_autospec(AsyncEngine, instance=True)
+        factory = mock.create_autospec(settings.async_sessionmaker, 
instance=True)
+        monkeypatch.setattr(settings, "async_engine", engine)
+        monkeypatch.setattr(settings, "AsyncSession", factory)
+        sync_engine, sync_factory = settings.engine, settings.Session
+
+        async def dispose():
+            await settings.dispose_async_engine()
+            await settings.dispose_async_engine()
+
+        asyncio.run(dispose())
+
+        assert engine.dispose.await_count == 2
+        engine.dispose.assert_awaited_with()
+        assert settings.async_engine is engine
+        assert settings.AsyncSession is factory
+        assert settings.engine is sync_engine
+        assert settings.Session is sync_factory
+
+    @pytest.mark.parametrize("error", [RuntimeError("disposal failed"), 
asyncio.CancelledError()])
+    def test_failed_disposal_preserves_resources_for_retry(self, monkeypatch, 
error):
+        engine = mock.create_autospec(AsyncEngine, instance=True)
+        factory = mock.create_autospec(settings.async_sessionmaker, 
instance=True)
+        monkeypatch.setattr(settings, "async_engine", engine)
+        monkeypatch.setattr(settings, "AsyncSession", factory)
+        engine.dispose.side_effect = error
+
+        with pytest.raises(type(error)):
+            asyncio.run(settings.dispose_async_engine())
+
+        assert settings.async_engine is engine
+        assert settings.AsyncSession is factory
+        engine.dispose.side_effect = None
+        asyncio.run(settings.dispose_async_engine())
+        assert settings.async_engine is engine
+        assert settings.AsyncSession is factory
+
+
[email protected]_test
[email protected](
+    "driver",
+    [
+        pytest.param("postgresql+psycopg_async", 
marks=pytest.mark.backend("postgres")),
+        pytest.param("postgresql+asyncpg", 
marks=pytest.mark.backend("postgres")),
+        pytest.param("mysql+aiomysql", marks=pytest.mark.backend("mysql")),
+        pytest.param("sqlite+aiosqlite", marks=pytest.mark.backend("sqlite")),
+    ],
+)
+def test_async_pool_is_closed_before_process_shutdown(driver):
+    script = textwrap.dedent(
+        """
+        import asyncio
+        import sys
+        from sqlalchemy import text
+        from sqlalchemy.engine import make_url
+        from airflow import settings
+
+        driver = sys.argv[1]
+        url = make_url(settings.SQL_ALCHEMY_CONN_ASYNC).set(drivername=driver)
+        settings.SQL_ALCHEMY_CONN_ASYNC = 
url.render_as_string(hide_password=False)
+        settings._configure_async_session()
+
+        async def run():
+            async with settings.AsyncSession() as session:
+                assert (await session.execute(text("SELECT 1"))).scalar_one() 
== 1
+                connection = await session.connection()
+                raw = (await connection.get_raw_connection()).driver_connection
+            await settings.dispose_async_engine()
+            if driver == "sqlite+aiosqlite":
+                try:
+                    await raw.execute("SELECT 1")
+                except ValueError:
+                    pass
+                else:
+                    raise AssertionError("connection remained open")
+            else:
+                assert raw.is_closed() if driver == "postgresql+asyncpg" else 
raw.closed
+
+        asyncio.run(run())
+        """
+    )
+    result = subprocess.run(
+        [sys.executable, "-W", "error::RuntimeWarning", "-c", script, driver],
+        check=False,
+        capture_output=True,
+        text=True,
+        timeout=30,
+    )
+    output = result.stdout + result.stderr
+    assert result.returncode == 0, output
+    for diagnostic in (
+        "MissingGreenlet",
+        "Event loop is closed",
+        "Exception closing connection",
+        "Exception ignored",
+        "was never awaited",
+    ):
+        assert diagnostic not in output, output
diff --git a/airflow-core/tests/unit/state/test_metastore.py 
b/airflow-core/tests/unit/state/test_metastore.py
index 732cece1458..ee0c578c21b 100644
--- a/airflow-core/tests/unit/state/test_metastore.py
+++ b/airflow-core/tests/unit/state/test_metastore.py
@@ -23,8 +23,10 @@ from typing import TYPE_CHECKING
 from unittest.mock import patch
 
 import pytest
+import pytest_asyncio
 from sqlalchemy import Delete, select
 
+from airflow import settings
 from airflow._shared.state import AssetStateStoreWriterKind
 from airflow._shared.timezones import timezone
 from airflow.configuration import conf
@@ -575,6 +577,13 @@ class TestMetastoreBackendAssetScope:
         )
 
 
+@pytest_asyncio.fixture(scope="class", loop_scope="class")
+async def dispose_async_engine():
+    yield
+    await settings.dispose_async_engine()
+
+
[email protected]("dispose_async_engine")
 @pytest.mark.asyncio(loop_scope="class")
 class TestMetastoreBackendAsync:
     async def test_aset_and_aget_task_roundtrip(self, backend: 
MetastoreBackend, dag_run_committed: DagRun):
diff --git a/airflow-core/tests/unit/utils/test_session.py 
b/airflow-core/tests/unit/utils/test_session.py
index 32b2a568ae5..11d680b47f4 100644
--- a/airflow-core/tests/unit/utils/test_session.py
+++ b/airflow-core/tests/unit/utils/test_session.py
@@ -20,6 +20,7 @@ from __future__ import annotations
 import pytest
 from sqlalchemy import select
 
+from airflow import settings
 from airflow.models import Log
 from airflow.utils.session import provide_session
 
@@ -58,10 +59,13 @@ class TestSession:
 
     @pytest.mark.asyncio
     async def test_async_session(self):
-        from airflow.settings import AsyncSession
-
-        session = AsyncSession()
-        session.add(Log(event="hihi1234"))
-        await session.commit()
-        my_special_log_event = await 
session.scalar(select(Log).where(Log.event == "hihi1234").limit(1))
-        assert my_special_log_event.event == "hihi1234"
+        try:
+            async with settings.AsyncSession() as session:
+                session.add(Log(event="hihi1234"))
+                await session.commit()
+                my_special_log_event = await session.scalar(
+                    select(Log).where(Log.event == "hihi1234").limit(1)
+                )
+                assert my_special_log_event.event == "hihi1234"
+        finally:
+            await settings.dispose_async_engine()

Reply via email to