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()