From 8c39b462915be435525f6bcb538adf4219dc5d13 Mon Sep 17 00:00:00 2001 From: David Blain Date: Wed, 1 Jul 2026 18:33:48 +0200 Subject: [PATCH 1/4] refactor: Added get_async_hook in common.compat provider --- .../providers/common/compat/hook/__init__.py | 46 ++++++++++++++ .../tests/unit/common/compat/hook/__init__.py | 16 +++++ .../unit/common/compat/hook/test_hook.py | 62 +++++++++++++++++++ 3 files changed, 124 insertions(+) create mode 100644 providers/common/compat/src/airflow/providers/common/compat/hook/__init__.py create mode 100644 providers/common/compat/tests/unit/common/compat/hook/__init__.py create mode 100644 providers/common/compat/tests/unit/common/compat/hook/test_hook.py diff --git a/providers/common/compat/src/airflow/providers/common/compat/hook/__init__.py b/providers/common/compat/src/airflow/providers/common/compat/hook/__init__.py new file mode 100644 index 0000000000000..5fe49788a5f43 --- /dev/null +++ b/providers/common/compat/src/airflow/providers/common/compat/hook/__init__.py @@ -0,0 +1,46 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + +from __future__ import annotations + +import logging + +from airflow.providers.common.compat.sdk import BaseHook + +log = logging.getLogger(__name__) + + +async def get_async_hook(conn_id: str, hook_params: dict | None = None) -> BaseHook: + """ + Get an asynchronous Airflow connection that is backwards compatible. + + :param conn_id: The provided connection ID. + :param hook_params: Additional hook params. + :returns: Connection + """ + from asgiref.sync import sync_to_async + + if hasattr(BaseHook, "aget_hook"): + log.debug("Get hook using `BaseHook.aget_hook().") + return await BaseHook.aget_hook(conn_id=conn_id, hook_params=hook_params) + log.debug("Get hook using `BaseHook.get_hook().") + return await sync_to_async(BaseHook.get_hook)(conn_id=conn_id, hook_params=hook_params) + + +__all__ = [ + "get_async_hook", +] diff --git a/providers/common/compat/tests/unit/common/compat/hook/__init__.py b/providers/common/compat/tests/unit/common/compat/hook/__init__.py new file mode 100644 index 0000000000000..13a83393a9124 --- /dev/null +++ b/providers/common/compat/tests/unit/common/compat/hook/__init__.py @@ -0,0 +1,16 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. diff --git a/providers/common/compat/tests/unit/common/compat/hook/test_hook.py b/providers/common/compat/tests/unit/common/compat/hook/test_hook.py new file mode 100644 index 0000000000000..1683e36953d15 --- /dev/null +++ b/providers/common/compat/tests/unit/common/compat/hook/test_hook.py @@ -0,0 +1,62 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +from unittest import mock + +import pytest + +from airflow.providers.common.compat.hook import get_async_hook +from airflow.providers.common.compat.sdk import BaseHook + +_MOCK_HOOK = mock.MagicMock(spec=BaseHook) + + +class MockAgetBaseHook: + def __init__(self, *args, **kwargs): + self.last_call = {} + + async def aget_hook(self, conn_id: str, hook_params: dict | None = None): + self.last_call = {"conn_id": conn_id, "hook_params": hook_params} + return _MOCK_HOOK + + +class MockBaseHook: + def __init__(self, *args, **kwargs): + self.last_call = {} + + def get_hook(self, conn_id: str, hook_params: dict | None = None): + self.last_call = {"conn_id": conn_id, "hook_params": hook_params} + return _MOCK_HOOK + + +class TestGetAsyncHook: + @pytest.mark.parametrize("hook_params", [None, {"key": "value"}]) + @mock.patch("airflow.providers.common.compat.hook.BaseHook", new_callable=MockAgetBaseHook) + @pytest.mark.asyncio + async def test_get_async_hook_uses_aget_hook_when_available(self, mock_hook_class, hook_params): + result = await get_async_hook("test_conn", hook_params=hook_params) + assert result is _MOCK_HOOK + assert mock_hook_class.last_call == {"conn_id": "test_conn", "hook_params": hook_params} + + @pytest.mark.parametrize("hook_params", [None, {"key": "value"}]) + @mock.patch("airflow.providers.common.compat.hook.BaseHook", new_callable=MockBaseHook) + @pytest.mark.asyncio + async def test_get_async_hook_falls_back_to_get_hook(self, mock_hook_class, hook_params): + result = await get_async_hook("test_conn", hook_params=hook_params) + assert result is _MOCK_HOOK + assert mock_hook_class.last_call == {"conn_id": "test_conn", "hook_params": hook_params} From 8edf0560b616ac30c0185a66fa76dbd742c9b24e Mon Sep 17 00:00:00 2001 From: David Blain Date: Fri, 3 Jul 2026 08:30:55 +0200 Subject: [PATCH 2/4] refactor: Fixed docstring and logging statements --- .../src/airflow/providers/common/compat/hook/__init__.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/providers/common/compat/src/airflow/providers/common/compat/hook/__init__.py b/providers/common/compat/src/airflow/providers/common/compat/hook/__init__.py index 5fe49788a5f43..6318594a85ed5 100644 --- a/providers/common/compat/src/airflow/providers/common/compat/hook/__init__.py +++ b/providers/common/compat/src/airflow/providers/common/compat/hook/__init__.py @@ -30,14 +30,14 @@ async def get_async_hook(conn_id: str, hook_params: dict | None = None) -> BaseH :param conn_id: The provided connection ID. :param hook_params: Additional hook params. - :returns: Connection + :returns: BaseHook """ from asgiref.sync import sync_to_async if hasattr(BaseHook, "aget_hook"): - log.debug("Get hook using `BaseHook.aget_hook().") + log.debug("Get hook using `BaseHook.aget_hook()`.") return await BaseHook.aget_hook(conn_id=conn_id, hook_params=hook_params) - log.debug("Get hook using `BaseHook.get_hook().") + log.debug("Get hook using `BaseHook.get_hook()`.") return await sync_to_async(BaseHook.get_hook)(conn_id=conn_id, hook_params=hook_params) From 1615118ee71f2d32a3334dd02a2395f7b4f1f68a Mon Sep 17 00:00:00 2001 From: David Blain Date: Fri, 3 Jul 2026 08:32:29 +0200 Subject: [PATCH 3/4] refactor: Refactored TestGetAsyncHook --- .../tests/unit/common/compat/hook/test_hook.py | 16 ++++------------ 1 file changed, 4 insertions(+), 12 deletions(-) diff --git a/providers/common/compat/tests/unit/common/compat/hook/test_hook.py b/providers/common/compat/tests/unit/common/compat/hook/test_hook.py index 1683e36953d15..da71a55563063 100644 --- a/providers/common/compat/tests/unit/common/compat/hook/test_hook.py +++ b/providers/common/compat/tests/unit/common/compat/hook/test_hook.py @@ -28,20 +28,12 @@ class MockAgetBaseHook: def __init__(self, *args, **kwargs): - self.last_call = {} - - async def aget_hook(self, conn_id: str, hook_params: dict | None = None): - self.last_call = {"conn_id": conn_id, "hook_params": hook_params} - return _MOCK_HOOK + self.aget_hook = mock.AsyncMock(return_value=_MOCK_HOOK) class MockBaseHook: def __init__(self, *args, **kwargs): - self.last_call = {} - - def get_hook(self, conn_id: str, hook_params: dict | None = None): - self.last_call = {"conn_id": conn_id, "hook_params": hook_params} - return _MOCK_HOOK + self.get_hook = mock.MagicMock(return_value=_MOCK_HOOK) class TestGetAsyncHook: @@ -51,7 +43,7 @@ class TestGetAsyncHook: async def test_get_async_hook_uses_aget_hook_when_available(self, mock_hook_class, hook_params): result = await get_async_hook("test_conn", hook_params=hook_params) assert result is _MOCK_HOOK - assert mock_hook_class.last_call == {"conn_id": "test_conn", "hook_params": hook_params} + mock_hook_class.aget_hook.assert_called_once_with(conn_id="test_conn", hook_params=hook_params) @pytest.mark.parametrize("hook_params", [None, {"key": "value"}]) @mock.patch("airflow.providers.common.compat.hook.BaseHook", new_callable=MockBaseHook) @@ -59,4 +51,4 @@ async def test_get_async_hook_uses_aget_hook_when_available(self, mock_hook_clas async def test_get_async_hook_falls_back_to_get_hook(self, mock_hook_class, hook_params): result = await get_async_hook("test_conn", hook_params=hook_params) assert result is _MOCK_HOOK - assert mock_hook_class.last_call == {"conn_id": "test_conn", "hook_params": hook_params} + mock_hook_class.get_hook.assert_called_once_with(conn_id="test_conn", hook_params=hook_params) From e792f07d4dad6f16b20c5b3186376b5608d584ee Mon Sep 17 00:00:00 2001 From: David Blain Date: Mon, 6 Jul 2026 13:35:37 +0200 Subject: [PATCH 4/4] Update providers/common/compat/src/airflow/providers/common/compat/hook/__init__.py Co-authored-by: Tzu-ping Chung --- .../compat/src/airflow/providers/common/compat/hook/__init__.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/providers/common/compat/src/airflow/providers/common/compat/hook/__init__.py b/providers/common/compat/src/airflow/providers/common/compat/hook/__init__.py index 6318594a85ed5..87a5f4d12eae2 100644 --- a/providers/common/compat/src/airflow/providers/common/compat/hook/__init__.py +++ b/providers/common/compat/src/airflow/providers/common/compat/hook/__init__.py @@ -26,7 +26,7 @@ async def get_async_hook(conn_id: str, hook_params: dict | None = None) -> BaseHook: """ - Get an asynchronous Airflow connection that is backwards compatible. + Get an asynchronous Airflow hook that is backwards compatible. :param conn_id: The provided connection ID. :param hook_params: Additional hook params.