diff --git a/providers/microsoft/azure/docs/connections/analysis_services.rst b/providers/microsoft/azure/docs/connections/analysis_services.rst index cb750ba444af0..471961a87a484 100644 --- a/providers/microsoft/azure/docs/connections/analysis_services.rst +++ b/providers/microsoft/azure/docs/connections/analysis_services.rst @@ -43,20 +43,35 @@ Region Endpoint Client ID Specify the Microsoft Entra service principal application (client) ID in ``login``. + Leave both Client ID and Client Secret empty to use ``DefaultAzureCredential``. Client Secret - Specify the service principal client secret in ``password``. + Specify the service principal client secret in ``password`` when using client-secret authentication. Tenant ID - Specify the Microsoft Entra tenant ID in the ``tenantId`` extra field. + Specify the Microsoft Entra tenant ID in the ``tenantId`` extra field when using client-secret authentication. + +Managed Identity Client ID + Optionally specify ``managed_identity_client_id`` in extras, together with + ``workload_identity_tenant_id``, to configure ``DefaultAzureCredential``. + +Workload Identity Tenant ID + Optionally specify ``workload_identity_tenant_id`` in extras, together with + ``managed_identity_client_id``, to configure ``DefaultAzureCredential``. Azure Analysis Services currently requires the service principal to be a server administrator for asynchronous refresh REST API calls. Add it to the server administrator role using the format ``app:@``. See `Add a service principal to the server administrator role `__. -This connection supports service principal client-secret authentication. Managed identity authentication is -not supported because Azure Analysis Services does not support managed identities for these operations. +When both Client ID and Client Secret are empty, the hook uses the asynchronous +``DefaultAzureCredential`` chain, including environment, workload identity, and managed identity +credentials. You can configure these through the Azure Identity environment variables, or supply both +identity extra fields above. Setting only one of Client ID or Client Secret is an error; incomplete +client-secret credentials do not fall back to another identity. + +The selected identity must still be authorized to use the Azure Analysis Services refresh API; +obtaining a token does not grant server administrator access. .. spelling:word-list:: diff --git a/providers/microsoft/azure/provider.yaml b/providers/microsoft/azure/provider.yaml index 3bc067c8c38d9..50bc1cd7ef8ce 100644 --- a/providers/microsoft/azure/provider.yaml +++ b/providers/microsoft/azure/provider.yaml @@ -429,6 +429,14 @@ connection-types: label: Tenant ID schema: type: ["string", "null"] + managed_identity_client_id: + label: Managed Identity Client ID + schema: + type: ["string", "null"] + workload_identity_tenant_id: + label: Workload Identity Tenant ID + schema: + type: ["string", "null"] - hook-class-name: airflow.providers.microsoft.azure.hooks.base_azure.AzureBaseHook hook-name: "Azure" connection-type: azure diff --git a/providers/microsoft/azure/src/airflow/providers/microsoft/azure/get_provider_info.py b/providers/microsoft/azure/src/airflow/providers/microsoft/azure/get_provider_info.py index 5a22d603e6380..790206f5c8e1a 100644 --- a/providers/microsoft/azure/src/airflow/providers/microsoft/azure/get_provider_info.py +++ b/providers/microsoft/azure/src/airflow/providers/microsoft/azure/get_provider_info.py @@ -418,7 +418,17 @@ def get_provider_info(): }, "placeholders": {"host": "westus.asazure.windows.net"}, }, - "conn-fields": {"tenantId": {"label": "Tenant ID", "schema": {"type": ["string", "null"]}}}, + "conn-fields": { + "tenantId": {"label": "Tenant ID", "schema": {"type": ["string", "null"]}}, + "managed_identity_client_id": { + "label": "Managed Identity Client ID", + "schema": {"type": ["string", "null"]}, + }, + "workload_identity_tenant_id": { + "label": "Workload Identity Tenant ID", + "schema": {"type": ["string", "null"]}, + }, + }, }, { "hook-class-name": "airflow.providers.microsoft.azure.hooks.base_azure.AzureBaseHook", diff --git a/providers/microsoft/azure/src/airflow/providers/microsoft/azure/hooks/analysis_services.py b/providers/microsoft/azure/src/airflow/providers/microsoft/azure/hooks/analysis_services.py index cdce36a4eb841..e07847f3452ae 100644 --- a/providers/microsoft/azure/src/airflow/providers/microsoft/azure/hooks/analysis_services.py +++ b/providers/microsoft/azure/src/airflow/providers/microsoft/azure/hooks/analysis_services.py @@ -26,6 +26,10 @@ from azure.identity.aio import ClientSecretCredential from airflow.providers.common.compat.sdk import AirflowException, BaseHook +from airflow.providers.microsoft.azure.utils import ( + add_managed_identity_connection_widgets, + get_async_default_azure_credential, +) if TYPE_CHECKING: from azure.core.credentials_async import AsyncTokenCredential @@ -74,9 +78,10 @@ class AzureAnalysisServicesHook(BaseHook): :param azure_analysis_services_conn_id: The Azure Analysis Services connection ID. :param request_timeout: Timeout in seconds for each HTTP request. - The connection must define the region endpoint in ``host``, the service principal client ID in - ``login``, the client secret in ``password``, and the Microsoft Entra tenant ID in the - ``tenantId`` extra field. + The connection must define the region endpoint in ``host``. For service principal authentication, + set the client ID in ``login``, the client secret in ``password``, and the Microsoft Entra tenant ID + in the ``tenantId`` extra field. If both ``login`` and ``password`` are empty, use + ``DefaultAzureCredential`` instead. """ conn_type: str = "azure_analysis_services" @@ -103,6 +108,7 @@ def connection(self) -> Connection: return self.get_connection(self.azure_analysis_services_conn_id) @classmethod + @add_managed_identity_connection_widgets def get_connection_form_widgets(cls) -> dict[str, Any]: """Return connection widgets to add to the connection form.""" from flask_appbuilder.fieldwidgets import BS3TextFieldWidget @@ -135,11 +141,18 @@ def get_conn(self) -> httpx.AsyncClient: return self._client def _get_credential(self) -> AsyncTokenCredential: - """Return and cache the service principal credential.""" + """Return and cache the Azure credential.""" if self._credential is not None: return self._credential connection = self.connection + if not connection.login and not connection.password: + self._credential = get_async_default_azure_credential( + managed_identity_client_id=connection.extra_dejson.get("managed_identity_client_id"), + workload_identity_tenant_id=connection.extra_dejson.get("workload_identity_tenant_id"), + ) + return self._credential + tenant_id = connection.extra_dejson.get("tenantId") if not connection.login: raise ValueError("Client ID is required for Azure Analysis Services authentication") diff --git a/providers/microsoft/azure/tests/unit/microsoft/azure/hooks/test_analysis_services.py b/providers/microsoft/azure/tests/unit/microsoft/azure/hooks/test_analysis_services.py index c1b0d57d5eedf..90db70f3e1ebc 100644 --- a/providers/microsoft/azure/tests/unit/microsoft/azure/hooks/test_analysis_services.py +++ b/providers/microsoft/azure/tests/unit/microsoft/azure/hooks/test_analysis_services.py @@ -23,6 +23,7 @@ import pytest from azure.core.credentials import AccessToken from azure.core.exceptions import ClientAuthenticationError +from azure.identity.aio import DefaultAzureCredential from airflow.models import Connection from airflow.providers.microsoft.azure.hooks.analysis_services import ( @@ -81,7 +82,11 @@ def test_rejects_invalid_request_timeout(self, request_timeout): def test_defines_connection_form_widget(self): pytest.importorskip("flask_appbuilder") - assert set(AzureAnalysisServicesHook.get_connection_form_widgets()) == {"tenantId"} + assert set(AzureAnalysisServicesHook.get_connection_form_widgets()) == { + "tenantId", + "managed_identity_client_id", + "workload_identity_tenant_id", + } def test_defines_connection_ui_field_behaviour(self): assert AzureAnalysisServicesHook.get_ui_field_behaviour() == { @@ -105,8 +110,9 @@ def test_get_conn_creates_and_caches_client(self, client_class): assert first_client is client_class.return_value assert second_client is first_client + @mock.patch(f"{MODULE}.get_async_default_azure_credential", autospec=True) @mock.patch(f"{MODULE}.ClientSecretCredential", autospec=True) - def test_get_credential_creates_and_caches_credential(self, credential_class): + def test_get_credential_creates_and_caches_credential(self, credential_class, default_credential): hook = AzureAnalysisServicesHook(CONN_ID) first_credential = hook._get_credential() @@ -119,6 +125,55 @@ def test_get_credential_creates_and_caches_credential(self, credential_class): ) assert first_credential is credential_class.return_value assert second_credential is first_credential + default_credential.assert_not_called() + + @pytest.mark.asyncio + @pytest.mark.parametrize("empty_value", [None, ""]) + @pytest.mark.parametrize( + "extra", + [ + {}, + { + "managed_identity_client_id": "identity-client-id", + "workload_identity_tenant_id": "identity-tenant-id", + "exclude_environment_credential": True, + }, + ], + ) + @mock.patch(f"{MODULE}.ClientSecretCredential", autospec=True) + @mock.patch(f"{MODULE}.get_async_default_azure_credential", autospec=True) + async def test_default_credential_lifecycle( + self, default_credential, secret_credential, create_mock_connection, empty_value, extra + ): + create_mock_connection( + Connection( + conn_id="default-auth", + conn_type="azure_analysis_services", + host=HOST, + login=empty_value, + password=empty_value, + extra=extra, + ) + ) + credential = mock.create_autospec(DefaultAzureCredential, instance=True) + credential.get_token.return_value = AccessToken("token", 0) + default_credential.return_value = credential + hook = AzureAnalysisServicesHook("default-auth") + + assert hook._get_credential() is credential + assert hook._get_credential() is credential + assert await hook._get_headers() == HEADERS + await hook.aclose() + await hook.aclose() + + default_credential.assert_called_once_with( + managed_identity_client_id=extra.get("managed_identity_client_id"), + workload_identity_tenant_id=extra.get("workload_identity_tenant_id"), + ) + secret_credential.assert_not_called() + credential.get_token.assert_awaited_once_with(TOKEN_SCOPE) + credential.close.assert_awaited_once_with() + assert hook._credential is None @mock.patch(f"{MODULE}.BaseHook.get_connection", autospec=True) def test_caches_connection(self, get_connection):