This is an automated email from the ASF dual-hosted git repository.
potiuk 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 39b48040226 Avoid extra_dejson in ADF and Synapse async hooks (#72130)
39b48040226 is described below
commit 39b48040226be44a6a113b4ea1bddeca80dc39b7
Author: deepinsight coder <[email protected]>
AuthorDate: Sat Aug 29 06:24:26 2026 -0700
Avoid extra_dejson in ADF and Synapse async hooks (#72130)
extra_dejson can mask secrets via a sync send on the triggerer event
loop, which raises AsyncToSync. Parse extras with json.loads instead,
matching the MSGraph workaround.
closes: #55728
---
.../microsoft/azure/hooks/data_factory.py | 8 ++-
.../providers/microsoft/azure/hooks/synapse.py | 7 ++-
.../microsoft/azure/hooks/test_data_factory.py | 68 ++++++++++++++++++++--
.../microsoft/azure/hooks/test_synapse_pipeline.py | 43 +++++++++++++-
4 files changed, 117 insertions(+), 9 deletions(-)
diff --git
a/providers/microsoft/azure/src/airflow/providers/microsoft/azure/hooks/data_factory.py
b/providers/microsoft/azure/src/airflow/providers/microsoft/azure/hooks/data_factory.py
index 61ae3c8848b..7df86c7be62 100644
---
a/providers/microsoft/azure/src/airflow/providers/microsoft/azure/hooks/data_factory.py
+++
b/providers/microsoft/azure/src/airflow/providers/microsoft/azure/hooks/data_factory.py
@@ -34,6 +34,7 @@ Spelling exceptions.
from __future__ import annotations
import inspect
+import json
import time
from collections.abc import Callable
from functools import wraps
@@ -1124,7 +1125,8 @@ def provide_targeted_factory_async(func: T) -> T:
if arg not in bound_args.arguments or bound_args.arguments[arg] is
None:
self = args[0]
conn = await get_async_connection(self.conn_id)
- extras = conn.extra_dejson
+ # extra_dejson can call mask_secret -> sync send on the
triggerer loop.
+ extras = json.loads(conn.extra) if conn.extra else {}
default_value = extras.get(default_key) or extras.get(
f"extra__azure_data_factory__{default_key}"
)
@@ -1175,7 +1177,8 @@ class AzureDataFactoryAsyncHook(AzureDataFactoryHook):
return self._async_conn
conn = await get_async_connection(self.conn_id)
- extras = conn.extra_dejson
+ # extra_dejson can call mask_secret -> sync send on the triggerer loop.
+ extras = json.loads(conn.extra) if conn.extra else {}
tenant = get_field(extras, "tenantId")
try:
@@ -1208,6 +1211,7 @@ class AzureDataFactoryAsyncHook(AzureDataFactoryHook):
async def refresh_conn(self) -> AsyncDataFactoryManagementClient: # type:
ignore[override]
self._conn = None
+ await self.close()
return await self.get_async_conn()
@provide_targeted_factory_async
diff --git
a/providers/microsoft/azure/src/airflow/providers/microsoft/azure/hooks/synapse.py
b/providers/microsoft/azure/src/airflow/providers/microsoft/azure/hooks/synapse.py
index e37c6bca9f2..c54f90e6bf7 100644
---
a/providers/microsoft/azure/src/airflow/providers/microsoft/azure/hooks/synapse.py
+++
b/providers/microsoft/azure/src/airflow/providers/microsoft/azure/hooks/synapse.py
@@ -16,6 +16,7 @@
# under the License.
from __future__ import annotations
+import json
import time
from typing import TYPE_CHECKING, Any
@@ -29,6 +30,7 @@ from azure.synapse.artifacts import ArtifactsClient
from azure.synapse.artifacts.aio import ArtifactsClient as AsyncArtifactsClient
from azure.synapse.spark import SparkClient
+from airflow.providers.common.compat.connection import get_async_connection
from airflow.providers.common.compat.sdk import AirflowException,
AirflowTaskTimeout, BaseHook
from airflow.providers.microsoft.azure.utils import (
add_managed_identity_connection_widgets,
@@ -489,8 +491,9 @@ class
AzureSynapsePipelineAsyncHook(AzureSynapsePipelineHook):
if self._async_conn is not None:
return self._async_conn
- conn = self.get_connection(self.conn_id)
- extras = conn.extra_dejson
+ conn = await get_async_connection(self.conn_id)
+ # extra_dejson can call mask_secret -> sync send on the triggerer loop.
+ extras = json.loads(conn.extra) if conn.extra else {}
tenant = self._get_field(extras, "tenantId")
credential: AsyncCredentials
diff --git
a/providers/microsoft/azure/tests/unit/microsoft/azure/hooks/test_data_factory.py
b/providers/microsoft/azure/tests/unit/microsoft/azure/hooks/test_data_factory.py
index 1b64d29c3ce..16e6eb237e8 100644
---
a/providers/microsoft/azure/tests/unit/microsoft/azure/hooks/test_data_factory.py
+++
b/providers/microsoft/azure/tests/unit/microsoft/azure/hooks/test_data_factory.py
@@ -16,6 +16,7 @@
# under the License.
from __future__ import annotations
+import json
import os
from unittest import mock
from unittest.mock import MagicMock, PropertyMock, patch
@@ -716,20 +717,31 @@ class TestAzureDataFactoryAsyncHook:
assert response == mock_status
@pytest.mark.asyncio
- @mock.patch(f"{MODULE}.AzureDataFactoryAsyncHook.get_connection")
+ @mock.patch(f"{MODULE}.get_async_connection", new_callable=mock.AsyncMock)
@mock.patch(f"{MODULE}.AzureDataFactoryAsyncHook.get_async_conn")
- async def test_get_pipeline_run_exception_without_resource(self,
mock_conn, mock_get_connection):
+ async def test_get_pipeline_run_exception_without_resource(self,
mock_conn, mock_get_async_connection):
"""
Test get_pipeline_run function without passing the resource name to
check the decorator function and
raise exception
"""
mock_connection = Connection(extra={"factory_name": DATAFACTORY_NAME})
- mock_get_connection.return_value = mock_connection
+ mock_get_async_connection.return_value = mock_connection
mock_conn.return_value.pipeline_runs.get.return_value = MagicMock()
hook = AzureDataFactoryAsyncHook(AZURE_DATA_FACTORY_CONN_ID)
with pytest.raises(AirflowException):
await hook.get_pipeline_run(RUN_ID, None, DATAFACTORY_NAME)
+ @staticmethod
+ def _conn_with_raising_extra_dejson(extra: dict, login="clientId",
password="clientSecret"):
+ conn = mock.Mock()
+ conn.login = login
+ conn.password = password
+ conn.extra = json.dumps(extra)
+ type(conn).extra_dejson = PropertyMock(
+ side_effect=RuntimeError("You cannot use AsyncToSync in the same
thread as an async event loop")
+ )
+ return conn
+
@pytest.mark.asyncio
@pytest.mark.parametrize(
"mocked_connection",
@@ -786,6 +798,50 @@ class TestAzureDataFactoryAsyncHook:
response = await hook.get_async_conn()
assert isinstance(response, DataFactoryManagementClient)
+ @pytest.mark.asyncio
+ async def test_get_async_conn_does_not_touch_extra_dejson(self):
+ conn = self._conn_with_raising_extra_dejson(
+ {"tenantId": "tenantId", "subscriptionId": "subscriptionId"}
+ )
+ hook = AzureDataFactoryAsyncHook(AZURE_DATA_FACTORY_CONN_ID)
+ with (
+ mock.patch(f"{MODULE}.get_async_connection",
new=mock.AsyncMock(return_value=conn)),
+ mock.patch(f"{MODULE}.AsyncClientSecretCredential"),
+ mock.patch(f"{MODULE}.AsyncDataFactoryManagementClient") as
mock_client,
+ ):
+ response = await hook.get_async_conn()
+ assert response is mock_client.return_value
+
+ @pytest.mark.asyncio
+ async def test_get_async_conn_uses_get_async_connection(self):
+ conn = self._conn_with_raising_extra_dejson(
+ {"tenantId": "tenantId", "subscriptionId": "subscriptionId"}
+ )
+ hook = AzureDataFactoryAsyncHook(AZURE_DATA_FACTORY_CONN_ID)
+ with (
+ mock.patch(
+ f"{MODULE}.get_async_connection",
new=mock.AsyncMock(return_value=conn)
+ ) as mock_get_async_connection,
+ mock.patch(f"{MODULE}.AsyncClientSecretCredential"),
+ mock.patch(f"{MODULE}.AsyncDataFactoryManagementClient"),
+ ):
+ await hook.get_async_conn()
+
mock_get_async_connection.assert_awaited_once_with(AZURE_DATA_FACTORY_CONN_ID)
+
+ @pytest.mark.asyncio
+ @mock.patch(f"{MODULE}.AzureDataFactoryAsyncHook.get_async_conn")
+ async def
test_provide_targeted_factory_async_does_not_touch_extra_dejson(self,
mock_get_async_conn):
+ conn = self._conn_with_raising_extra_dejson(
+ {"resource_group_name": RESOURCE_GROUP_NAME, "factory_name":
DATAFACTORY_NAME}
+ )
+ mock_get_async_conn.return_value.pipeline_runs.get =
mock.AsyncMock(return_value=MagicMock())
+ hook = AzureDataFactoryAsyncHook(AZURE_DATA_FACTORY_CONN_ID)
+ with mock.patch(f"{MODULE}.get_async_connection",
new=mock.AsyncMock(return_value=conn)):
+ await hook.get_pipeline_run(RUN_ID, None, None)
+
mock_get_async_conn.return_value.pipeline_runs.get.assert_awaited_once_with(
+ RESOURCE_GROUP_NAME, DATAFACTORY_NAME, RUN_ID
+ )
+
@pytest.mark.asyncio
@pytest.mark.parametrize(
"mocked_connection",
@@ -904,10 +960,14 @@ class TestAzureDataFactoryAsyncHook:
@pytest.mark.asyncio
@mock.patch(f"{MODULE}.AzureDataFactoryAsyncHook.get_async_conn")
async def test_refresh_conn(self, mock_get_async_conn):
- """Test refresh_conn method _conn is reset and get_async_conn is
called"""
+ """Test refresh_conn closes the async client before recreating it."""
hook = AzureDataFactoryAsyncHook(AZURE_DATA_FACTORY_CONN_ID)
+ mock_async_conn = mock.AsyncMock()
+ hook._async_conn = mock_async_conn
await hook.refresh_conn()
assert not hook._conn
+ mock_async_conn.close.assert_awaited_once()
+ assert hook._async_conn is None
assert mock_get_async_conn.called
@pytest.mark.asyncio
diff --git
a/providers/microsoft/azure/tests/unit/microsoft/azure/hooks/test_synapse_pipeline.py
b/providers/microsoft/azure/tests/unit/microsoft/azure/hooks/test_synapse_pipeline.py
index b35e589d028..436d2aa4474 100644
---
a/providers/microsoft/azure/tests/unit/microsoft/azure/hooks/test_synapse_pipeline.py
+++
b/providers/microsoft/azure/tests/unit/microsoft/azure/hooks/test_synapse_pipeline.py
@@ -16,7 +16,8 @@
# under the License.
from __future__ import annotations
-from unittest.mock import AsyncMock, MagicMock, patch
+import json
+from unittest.mock import AsyncMock, MagicMock, PropertyMock, patch
import pytest
from azure.synapse.artifacts import ArtifactsClient
@@ -247,6 +248,46 @@ class TestAzureSynapsePipelineAsyncHook:
credential=mock_credential.return_value,
)
+ @pytest.mark.asyncio
+ @patch(f"{MODULE}.AsyncArtifactsClient")
+ @patch(f"{MODULE}.AsyncClientSecretCredential")
+ async def test_get_async_conn_does_not_touch_extra_dejson(self,
mock_credential, mock_client):
+ conn = MagicMock()
+ conn.login = "clientId"
+ conn.password = "clientSecret"
+ conn.extra = json.dumps({"tenantId": "tenantId"})
+ type(conn).extra_dejson = PropertyMock(
+ side_effect=RuntimeError("You cannot use AsyncToSync in the same
thread as an async event loop")
+ )
+ hook = AzureSynapsePipelineAsyncHook(
+ azure_synapse_conn_id=DEFAULT_CONNECTION_CLIENT_SECRET,
+
azure_synapse_workspace_dev_endpoint=AZURE_SYNAPSE_WORKSPACE_DEV_ENDPOINT,
+ )
+ with patch(f"{MODULE}.get_async_connection",
new=AsyncMock(return_value=conn)):
+ result = await hook.get_async_conn()
+ assert result is mock_client.return_value
+ mock_credential.assert_called_with(
+ client_id="clientId",
+ client_secret="clientSecret",
+ tenant_id="tenantId",
+ )
+
+ @pytest.mark.asyncio
+ @patch(f"{MODULE}.AsyncArtifactsClient")
+ @patch(f"{MODULE}.AsyncClientSecretCredential")
+ async def test_get_async_conn_uses_get_async_connection(self,
mock_credential, mock_client):
+ conn = MagicMock()
+ conn.login = "clientId"
+ conn.password = "clientSecret"
+ conn.extra = json.dumps({"tenantId": "tenantId"})
+ hook = AzureSynapsePipelineAsyncHook(
+ azure_synapse_conn_id=DEFAULT_CONNECTION_CLIENT_SECRET,
+
azure_synapse_workspace_dev_endpoint=AZURE_SYNAPSE_WORKSPACE_DEV_ENDPOINT,
+ )
+ with patch(f"{MODULE}.get_async_connection",
new=AsyncMock(return_value=conn)) as mock_get:
+ await hook.get_async_conn()
+ mock_get.assert_awaited_once_with(DEFAULT_CONNECTION_CLIENT_SECRET)
+
@pytest.mark.asyncio
@patch(f"{MODULE}.AsyncArtifactsClient")
@patch(f"{MODULE}.AsyncDefaultAzureCredential")