diff --git a/providers/databricks/src/airflow/providers/databricks/hooks/databricks_base.py b/providers/databricks/src/airflow/providers/databricks/hooks/databricks_base.py index fe9d2245705e7..9fd5d17e57854 100644 --- a/providers/databricks/src/airflow/providers/databricks/hooks/databricks_base.py +++ b/providers/databricks/src/airflow/providers/databricks/hooks/databricks_base.py @@ -37,6 +37,7 @@ import aiohttp import requests from aiohttp.client_exceptions import ClientConnectorError +from asgiref.sync import sync_to_async from requests import PreparedRequest, exceptions as requests_exceptions from requests.auth import AuthBase, HTTPBasicAuth from requests.exceptions import JSONDecodeError @@ -59,7 +60,7 @@ from airflow.hooks.base import BaseHook as BaseHook # type: ignore if TYPE_CHECKING: - from airflow.models import Connection + from airflow.providers.common.compat.sdk import Connection # https://docs.microsoft.com/en-us/azure/active-directory/managed-identities-azure-resources/how-to-use-vm-token AZURE_METADATA_SERVICE_TOKEN_URL = "http://169.254.169.254/metadata/identity/oauth2/token" @@ -144,6 +145,10 @@ def __init__( self._metadata_expiry: float = 0 self._metadata_ttl: int = 300 + # Cache for lack of an async @cached_property + self._a_databricks_conn: Connection | None = None + self._a_host: str | None = None + def my_after_func(retry_state): self._log_request_error(retry_state.attempt_number, retry_state.outcome) @@ -163,6 +168,14 @@ def my_after_func(retry_state): def databricks_conn(self) -> Connection: return self.get_connection(self.databricks_conn_id) # type: ignore[return-value] + async def a_databricks_conn(self) -> Connection: + if self._a_databricks_conn is None: + if hasattr(self, "aget_connection"): + self._a_databricks_conn = await self.aget_connection(self.databricks_conn_id) + else: + self._a_databricks_conn = await sync_to_async(self.get_connection)(self.databricks_conn_id) + return self._a_databricks_conn # type: ignore[return-value] + def get_conn(self) -> Connection: return self.databricks_conn @@ -193,6 +206,16 @@ def host(self) -> str | None: host = self._parse_host(self.databricks_conn.host) return host + async def a_host(self) -> str | None: + """Async version of `host` property.""" + if self._a_host is None: + conn = await self.a_databricks_conn() + if "host" in conn.extra_dejson: + self._a_host = self._parse_host(conn.extra_dejson["host"]) + elif conn.host: + self._a_host = self._parse_host(conn.host) + return self._a_host + async def __aenter__(self): self._session = aiohttp.ClientSession() return self @@ -232,6 +255,13 @@ def _get_connection_attr(self, attr_name: str) -> str: raise ValueError(f"`{attr_name}` must be present in Connection") return attr + async def _a_get_connection_attr(self, attr_name: str) -> str: + """Async version of `_get_connection_attr`.""" + conn = await self.a_databricks_conn() + if not (attr := getattr(conn, attr_name)): + raise ValueError(f"`{attr_name}` must be present in Connection") + return attr + def _get_retry_object(self) -> Retrying: """ Instantiate a retry object. @@ -292,13 +322,12 @@ async def _a_get_sp_token(self, resource: str) -> str: self.log.info("Existing Service Principal token is expired, or going to expire soon. Refreshing...") try: + conn = await self.a_databricks_conn() async for attempt in self._a_get_retry_object(): with attempt: async with self._session.post( resource, - auth=aiohttp.BasicAuth( - self._get_connection_attr("login"), self.databricks_conn.password - ), + auth=aiohttp.BasicAuth(await self._a_get_connection_attr("login"), conn.password), data="grant_type=client_credentials&scope=all-apis", headers={ **self.user_agent_header, @@ -384,16 +413,17 @@ async def _a_get_aad_token(self, resource: str) -> str: ManagedIdentityCredential as AsyncManagedIdentityCredential, ) + conn = await self.a_databricks_conn() async for attempt in self._a_get_retry_object(): with attempt: - if self.databricks_conn.extra_dejson.get("use_azure_managed_identity", False): + if conn.extra_dejson.get("use_azure_managed_identity", False): async with AsyncManagedIdentityCredential() as credential: token = await credential.get_token(f"{resource}/.default") else: async with AsyncClientSecretCredential( - client_id=self._get_connection_attr("login"), - client_secret=self.databricks_conn.password, - tenant_id=self.databricks_conn.extra_dejson["azure_tenant_id"], + client_id=await self._a_get_connection_attr("login"), + client_secret=conn.password, + tenant_id=conn.extra_dejson["azure_tenant_id"], ) as credential: token = await credential.get_token(f"{resource}/.default") jsn = { @@ -523,11 +553,10 @@ async def _a_get_aad_headers(self) -> dict: :return: dictionary with filled AAD headers """ headers = {} - if "azure_resource_id" in self.databricks_conn.extra_dejson: + conn = await self.a_databricks_conn() + if "azure_resource_id" in conn.extra_dejson: mgmt_token = await self._a_get_aad_token(AZURE_MANAGEMENT_ENDPOINT) - headers["X-Databricks-Azure-Workspace-Resource-Id"] = self.databricks_conn.extra_dejson[ - "azure_resource_id" - ] + headers["X-Databricks-Azure-Workspace-Resource-Id"] = conn.extra_dejson["azure_resource_id"] headers["X-Databricks-Azure-SP-Management-Token"] = mgmt_token return headers @@ -564,7 +593,8 @@ def _get_k8s_jwt_token(self) -> str: async def _a_get_k8s_jwt_token(self) -> str: """Async version of _get_k8s_jwt_token().""" - if "k8s_projected_volume_token_path" in self.databricks_conn.extra_dejson: + conn = await self.a_databricks_conn() + if "k8s_projected_volume_token_path" in conn.extra_dejson: self.log.info("Using Kubernetes projected volume token") return await self._a_get_k8s_projected_volume_token() @@ -621,8 +651,9 @@ def _get_aiofiles(): async def _a_get_k8s_projected_volume_token(self) -> str: """Async version of _get_k8s_projected_volume_token().""" aiofiles = self._get_aiofiles() + conn = await self.a_databricks_conn() - projected_token_path: str = self.databricks_conn.extra_dejson["k8s_projected_volume_token_path"] + projected_token_path: str = conn.extra_dejson["k8s_projected_volume_token_path"] try: async with aiofiles.open(projected_token_path) as f: @@ -732,15 +763,12 @@ def _get_k8s_token_request_api(self) -> str: async def _a_get_k8s_token_request_api(self) -> str: """Async version of _get_k8s_token_request_api().""" aiofiles = self._get_aiofiles() + conn = await self.a_databricks_conn() - audience = self.databricks_conn.extra_dejson.get("audience", DEFAULT_K8S_AUDIENCE) - expiration_seconds = self.databricks_conn.extra_dejson.get("expiration_seconds", 3600) - token_path = self.databricks_conn.extra_dejson.get( - "k8s_token_path", DEFAULT_K8S_SERVICE_ACCOUNT_TOKEN_PATH - ) - namespace_path = self.databricks_conn.extra_dejson.get( - "k8s_namespace_path", DEFAULT_K8S_NAMESPACE_PATH - ) + audience = conn.extra_dejson.get("audience", DEFAULT_K8S_AUDIENCE) + expiration_seconds = conn.extra_dejson.get("expiration_seconds", 3600) + token_path = conn.extra_dejson.get("k8s_token_path", DEFAULT_K8S_SERVICE_ACCOUNT_TOKEN_PATH) + namespace_path = conn.extra_dejson.get("k8s_namespace_path", DEFAULT_K8S_NAMESPACE_PATH) try: async with aiofiles.open(token_path) as f: @@ -819,6 +847,20 @@ def _get_required_client_id(self) -> str: ) return client_id + async def _a_get_required_client_id(self) -> str: + """Async version of `_get_required_client_id()`.""" + conn = await self.a_databricks_conn() + client_id = conn.extra_dejson.get("client_id") + if not client_id: + # see: https://github.com/kubernetes/kubernetes/issues/116638 + raise AirflowException( + "client_id is required for Kubernetes OIDC token federation. " + "Kubernetes service account tokens do not support custom claims, " + "so service principal-level federation must be used. " + "Please provide client_id in the connection extra parameters." + ) + return client_id + def _get_federated_databricks_token(self, resource: str) -> str: """ Get Databricks OAuth token by exchanging Kubernetes JWT token. @@ -882,7 +924,7 @@ async def _a_get_federated_databricks_token(self, resource: str) -> str: self.log.info("Existing federated token is expired or missing. Fetching new token...") - client_id = self._get_required_client_id() + client_id = await self._a_get_required_client_id() # Get JWT from Kubernetes jwt_token = await self._a_get_k8s_jwt_token() @@ -1018,37 +1060,36 @@ def _get_token(self, raise_error: bool = False) -> str | None: return None async def _a_get_token(self, raise_error: bool = False) -> str | None: - if "token" in self.databricks_conn.extra_dejson: + conn = await self.a_databricks_conn() + if "token" in conn.extra_dejson: self.log.info( "Using token auth. For security reasons, please set token in Password field instead of extra" ) - return self.databricks_conn.extra_dejson["token"] - if not self.databricks_conn.login and self.databricks_conn.password: + return conn.extra_dejson["token"] + if not conn.login and conn.password: self.log.debug("Using token auth.") - return self.databricks_conn.password - if "azure_tenant_id" in self.databricks_conn.extra_dejson: - if self.databricks_conn.login == "" or self.databricks_conn.password == "": + return conn.password + if "azure_tenant_id" in conn.extra_dejson: + if conn.login == "" or conn.password == "": raise AirflowException("Azure SPN credentials aren't provided") self.log.debug("Using AAD Token for SPN.") return await self._a_get_aad_token(DEFAULT_DATABRICKS_SCOPE) - if self.databricks_conn.extra_dejson.get("use_azure_managed_identity", False): + if conn.extra_dejson.get("use_azure_managed_identity", False): self.log.debug("Using AAD Token for managed identity.") await self._a_check_azure_metadata_service() return await self._a_get_aad_token(DEFAULT_DATABRICKS_SCOPE) - if self.databricks_conn.extra_dejson.get(DEFAULT_AZURE_CREDENTIAL_SETTING_KEY, False): + if conn.extra_dejson.get(DEFAULT_AZURE_CREDENTIAL_SETTING_KEY, False): self.log.debug("Using AzureDefaultCredential for authentication.") return await self._a_get_aad_token_for_default_az_credential(DEFAULT_DATABRICKS_SCOPE) - if self.databricks_conn.extra_dejson.get("service_principal_oauth", False): - if self.databricks_conn.login == "" or self.databricks_conn.password == "": + if conn.extra_dejson.get("service_principal_oauth", False): + if conn.login == "" or conn.password == "": raise AirflowException("Service Principal credentials aren't provided") self.log.debug("Using Service Principal Token.") - return await self._a_get_sp_token(self._get_oidc_token_service_url()) - if self.databricks_conn.login == "federated_k8s" or self.databricks_conn.extra_dejson.get( - "federated_k8s", False - ): + return await self._a_get_sp_token(await self._a_get_oidc_token_service_url()) + if conn.login == "federated_k8s" or conn.extra_dejson.get("federated_k8s", False): self.log.debug("Using Kubernetes OIDC token federation.") - return await self._a_get_federated_databricks_token(self._get_oidc_token_service_url()) + return await self._a_get_federated_databricks_token(await self._a_get_oidc_token_service_url()) if raise_error: raise AirflowException("Token authentication isn't configured") @@ -1065,11 +1106,25 @@ def _get_oidc_token_service_url(self) -> str: """ return OIDC_TOKEN_SERVICE_URL.format(f"https://{self.host}") + async def _a_get_oidc_token_service_url(self) -> str: + """ + Async version of `_get_oidc_token_service_url()`. + + :return: Full URL to the OIDC token service endpoint + """ + return OIDC_TOKEN_SERVICE_URL.format(f"https://{await self.a_host()}") + def _endpoint_url(self, endpoint): port = f":{self.databricks_conn.port}" if self.databricks_conn.port else "" schema = self.databricks_conn.schema or "https" return f"{schema}://{self.host}{port}/{endpoint}" + async def _a_endpoint_url(self, endpoint): + conn = await self.a_databricks_conn() + port = f":{conn.port}" if conn.port else "" + schema = conn.schema or "https" + return f"{schema}://{await self.a_host()}{port}/{endpoint}" + def _do_api_call( self, endpoint_info: tuple[str, str], @@ -1156,7 +1211,7 @@ async def _a_do_api_call(self, endpoint_info: tuple[str, str], json: dict[str, A method, endpoint = endpoint_info full_endpoint = f"api/{endpoint}" - url = self._endpoint_url(full_endpoint) + url = await self._a_endpoint_url(full_endpoint) aad_headers = await self._a_get_aad_headers() headers = {**self.user_agent_header, **aad_headers} @@ -1167,7 +1222,8 @@ async def _a_do_api_call(self, endpoint_info: tuple[str, str], json: dict[str, A auth = BearerAuth(token) else: self.log.info("Using basic auth.") - auth = aiohttp.BasicAuth(self._get_connection_attr("login"), self.databricks_conn.password) + conn = await self.a_databricks_conn() + auth = aiohttp.BasicAuth(await self._a_get_connection_attr("login"), conn.password) request_func: Any if method == "GET": diff --git a/providers/databricks/tests/unit/databricks/hooks/test_databricks_base.py b/providers/databricks/tests/unit/databricks/hooks/test_databricks_base.py index 8c8794ebda946..c01b843e0e372 100644 --- a/providers/databricks/tests/unit/databricks/hooks/test_databricks_base.py +++ b/providers/databricks/tests/unit/databricks/hooks/test_databricks_base.py @@ -188,7 +188,8 @@ def test_get_sp_token_retry_error(self, mock_post): @pytest.mark.asyncio @time_machine.travel("2025-07-12 12:00:00") @mock.patch("aiohttp.ClientSession.post") - async def test_a_get_sp_token(self, mock_post): + @mock.patch("airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.a_databricks_conn") + async def test_a_get_sp_token(self, mock_conn, mock_post): expiry_date = int((datetime(2025, 7, 12, 12, 0, 0) + timedelta(minutes=60)).timestamp()) mock_response = mock.AsyncMock() mock_response.__aenter__.return_value = mock_response @@ -200,12 +201,10 @@ async def test_a_get_sp_token(self, mock_post): "token_type": "Bearer", } mock_post.return_value = mock_response - mock_conn = mock.Mock() - mock_conn.login = "client_id" - mock_conn.password = "client_secret" + + mock_conn.return_value = Connection(login="client_id", password="client_secret") hook = BaseDatabricksHook() - hook.databricks_conn = mock_conn hook.user_agent_header = {"User-Agent": "test-agent"} hook.token_timeout_seconds = 10 async with aiohttp.ClientSession() as session: @@ -505,10 +504,7 @@ def test_get_token_not_configured_raises(self, mock_conn): hook._get_token(raise_error=True) @pytest.mark.asyncio - @mock.patch( - "airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.databricks_conn", - new_callable=mock.PropertyMock, - ) + @mock.patch("airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.a_databricks_conn") async def test_a_get_token_from_extra_dejson(self, mock_conn): extra = {"token": "test_token"} mock_conn.return_value = Connection(extra=extra) @@ -521,10 +517,7 @@ async def test_a_get_token_from_extra_dejson(self, mock_conn): ) @pytest.mark.asyncio - @mock.patch( - "airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.databricks_conn", - new_callable=mock.PropertyMock, - ) + @mock.patch("airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.a_databricks_conn") async def test_a_token_from_password_when_login_missing(self, mock_conn): mock_conn.return_value = Connection(login=None, password="pw-token") hook = BaseDatabricksHook() @@ -538,10 +531,7 @@ async def test_a_token_from_password_when_login_missing(self, mock_conn): "airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook._a_get_sp_token", new_callable=mock.AsyncMock, ) - @mock.patch( - "airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.databricks_conn", - new_callable=mock.PropertyMock, - ) + @mock.patch("airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.a_databricks_conn") async def test_a_get_token_service_principal_oauth_success(self, mock_conn, mock_get_sp_token): mock_conn.return_value = Connection( host="example.databricks.com", @@ -562,10 +552,7 @@ async def test_a_get_token_service_principal_oauth_success(self, mock_conn, mock "airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook._a_get_aad_token", new_callable=mock.AsyncMock, ) - @mock.patch( - "airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.databricks_conn", - new_callable=mock.PropertyMock, - ) + @mock.patch("airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.a_databricks_conn") async def test_a_get_token_azure_spn_success(self, mock_conn, mock_get_aad_token): extra = {"azure_tenant_id": "tenant_id"} mock_conn.return_value = Connection(login="spn_client_id", password="spn_client_secret", extra=extra) @@ -578,10 +565,7 @@ async def test_a_get_token_azure_spn_success(self, mock_conn, mock_get_aad_token mock_log_debug.assert_called_once_with("Using AAD Token for SPN.") @pytest.mark.asyncio - @mock.patch( - "airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.databricks_conn", - new_callable=mock.PropertyMock, - ) + @mock.patch("airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.a_databricks_conn") async def test_a_get_token_azure_spn_missing_credentials_raises(self, mock_conn): mock_conn.return_value = Connection(login="", password="", extra={"azure_tenant_id": "tenant_id"}) hook = BaseDatabricksHook() @@ -597,10 +581,7 @@ async def test_a_get_token_azure_spn_missing_credentials_raises(self, mock_conn) "airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook._a_get_aad_token", new_callable=mock.AsyncMock, ) - @mock.patch( - "airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.databricks_conn", - new_callable=mock.PropertyMock, - ) + @mock.patch("airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.a_databricks_conn") async def test_a_get_token_managed_identity(self, mock_conn, mock_get_aad_token, mock_check_metadata): mock_conn.return_value = Connection(extra={"use_azure_managed_identity": True}) mock_get_aad_token.return_value = "mi_token" @@ -617,10 +598,7 @@ async def test_a_get_token_managed_identity(self, mock_conn, mock_get_aad_token, "airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook._a_get_aad_token_for_default_az_credential", new_callable=mock.AsyncMock, ) - @mock.patch( - "airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.databricks_conn", - new_callable=mock.PropertyMock, - ) + @mock.patch("airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.a_databricks_conn") async def test_a_get_token_default_azure_credential(self, mock_conn, mock_get_default_cred_token): extra = {DEFAULT_AZURE_CREDENTIAL_SETTING_KEY: True} mock_conn.return_value = Connection(extra=extra) @@ -633,10 +611,7 @@ async def test_a_get_token_default_azure_credential(self, mock_conn, mock_get_de mock_log_debug.assert_called_once_with("Using AzureDefaultCredential for authentication.") @pytest.mark.asyncio - @mock.patch( - "airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.databricks_conn", - new_callable=mock.PropertyMock, - ) + @mock.patch("airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.a_databricks_conn") async def test_a_get_token_service_principal_oauth_missing_credentials(self, mock_conn): mock_conn.return_value = Connection( host="host", login="", password="", extra={"service_principal_oauth": True} @@ -646,10 +621,7 @@ async def test_a_get_token_service_principal_oauth_missing_credentials(self, moc await hook._a_get_token() @pytest.mark.asyncio - @mock.patch( - "airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.databricks_conn", - new_callable=mock.PropertyMock, - ) + @mock.patch("airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.a_databricks_conn") async def test_a_get_token_not_configured_raises(self, mock_conn): mock_conn.return_value = Connection( host="host", @@ -770,6 +742,107 @@ def test_get_error_code_with_http_error_and_valid_error_code(self): hook = BaseDatabricksHook() assert hook._get_error_code(exception) == "INVALID_REQUEST" + @pytest.mark.asyncio + @mock.patch( + "airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.aget_connection", + create=True, + new_callable=mock.AsyncMock, + ) + async def test_cached_a_databricks_conn(self, mock_aget_connection): + """Verify aget_connection caching.""" + mock_aget_connection.return_value = Connection(login="foo", password="bar") + hook = BaseDatabricksHook() + await hook.a_databricks_conn() + await hook.a_databricks_conn() + mock_aget_connection.assert_called_once() + + @pytest.mark.asyncio + @mock.patch( + "airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.aget_connection", + create=True, + new_callable=mock.AsyncMock, + ) + async def test_no_sync_get_connection(self, mock_aget_connection): + """Ensure sync databricks_conn isn't referenced during async methods.""" + with mock.patch.object( + BaseDatabricksHook, "databricks_conn", new_callable=mock.PropertyMock + ) as mock_databricks_conn: + mock_databricks_conn.side_effect = AssertionError( + "databricks_conn should not be accessed running async" + ) + + mock_aget_connection.return_value = Connection(login="foo", password="bar") + hook = BaseDatabricksHook() + await hook._a_get_token() + await hook._a_get_aad_headers() + await hook._a_endpoint_url(endpoint="foobar") + + @pytest.mark.asyncio + @mock.patch("aiohttp.ClientSession.post") + @mock.patch("ssl.create_default_context") + @mock.patch( + "airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.aget_connection", + create=True, + new_callable=mock.AsyncMock, + ) + async def test_no_sync_get_connection_federated_k8s(self, mock_aget_connection, _mock_ssl_ctx, mock_post): + """Ensure sync databricks_conn isn't referenced during async federated K8s auth.""" + with mock.patch.object( + BaseDatabricksHook, "databricks_conn", new_callable=mock.PropertyMock + ) as mock_databricks_conn: + mock_databricks_conn.side_effect = AssertionError( + "databricks_conn should not be accessed running async" + ) + + mock_aget_connection.return_value = Connection( + host="my-workspace.cloud.databricks.com", + login="federated_k8s", + extra={"client_id": "test-client-id"}, + ) + + # Mock aiofiles.open for reading K8s token and namespace + mock_token_file = mock.AsyncMock() + mock_token_file.read = mock.AsyncMock(return_value="in_cluster_token") + mock_token_file.__aenter__ = mock.AsyncMock(return_value=mock_token_file) + mock_token_file.__aexit__ = mock.AsyncMock(return_value=None) + + mock_namespace_file = mock.AsyncMock() + mock_namespace_file.read = mock.AsyncMock(return_value="default") + mock_namespace_file.__aenter__ = mock.AsyncMock(return_value=mock_namespace_file) + mock_namespace_file.__aexit__ = mock.AsyncMock(return_value=None) + + # Mock K8s TokenRequest API response + k8s_response = mock.AsyncMock() + k8s_response.__aenter__.return_value = k8s_response + k8s_response.__aexit__.return_value = None + k8s_response.raise_for_status.return_value = None + k8s_response.json.return_value = {"status": {"token": "k8s_jwt_token"}} + + # Mock Databricks token exchange response + db_response = mock.AsyncMock() + db_response.__aenter__.return_value = db_response + db_response.__aexit__.return_value = None + db_response.raise_for_status.return_value = None + db_response.json.return_value = { + "access_token": "async_databricks_token", + "expires_in": 3600, + "token_type": "Bearer", + } + + mock_post.side_effect = [k8s_response, db_response] + + hook = BaseDatabricksHook() + hook.user_agent_header = {"User-Agent": "test-agent"} + hook.token_timeout_seconds = 10 + + with mock.patch("aiofiles.open", side_effect=[mock_token_file, mock_namespace_file]): + async with aiohttp.ClientSession() as session: + hook._session = session + # This should NOT access self.databricks_conn (sync property) + token = await hook._a_get_token() + + assert token == "async_databricks_token" + @mock.patch("requests.get") @time_machine.travel("2025-07-12 12:00:00") def test_check_azure_metadata_service_normal(self, mock_get): @@ -1381,11 +1454,18 @@ def test_get_token_with_federated_k8s_extra(self, mock_file): @pytest.mark.asyncio @mock.patch("aiohttp.ClientSession.post") + @mock.patch("airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.a_databricks_conn") @time_machine.travel("2025-07-12 12:00:00") @mock.patch("ssl.create_default_context") - async def test_a_get_federated_token(self, _mock_ssl_ctx, mock_post): + async def test_a_get_federated_token(self, _mock_ssl_ctx, mock_conn, mock_post): """Test async version of federated token exchange.""" expiry_date = int((datetime(2025, 7, 12, 12, 0, 0) + timedelta(minutes=60)).timestamp()) + conn = Connection( + host="my-workspace.cloud.databricks.com", + login="federated_k8s", + extra={"client_id": "test-client-id"}, + ) + mock_conn.return_value = conn # Mock aiofiles.open for reading K8s token and namespace mock_token_file = mock.AsyncMock() @@ -1418,20 +1498,14 @@ async def test_a_get_federated_token(self, _mock_ssl_ctx, mock_post): mock_post.side_effect = [k8s_response, db_response] - mock_conn = mock.Mock() - mock_conn.host = "my-workspace.cloud.databricks.com" - mock_conn.login = "federated_k8s" - mock_conn.extra_dejson = {"client_id": "test-client-id"} - hook = BaseDatabricksHook() - hook.databricks_conn = mock_conn hook.user_agent_header = {"User-Agent": "test-agent"} hook.token_timeout_seconds = 10 with mock.patch("aiofiles.open", side_effect=[mock_token_file, mock_namespace_file]): async with aiohttp.ClientSession() as session: hook._session = session - resource = f"https://{mock_conn.host}/oidc/v1/token" + resource = f"https://{conn.host}/oidc/v1/token" token = await hook._a_get_federated_databricks_token(resource) assert token == "async_databricks_token" @@ -1446,15 +1520,9 @@ async def test_a_get_federated_token(self, _mock_ssl_ctx, mock_post): @time_machine.travel("2025-07-12 12:00:00") async def test_a_get_federated_token_cached_valid(self): """Test that async version returns cached valid token without fetching new one.""" - mock_conn = mock.Mock() - mock_conn.host = "my-workspace.cloud.databricks.com" - mock_conn.login = "federated_k8s" - mock_conn.extra_dejson = {"client_id": "test-client-id"} - hook = BaseDatabricksHook() - hook.databricks_conn = mock_conn - resource = f"https://{mock_conn.host}/oidc/v1/token" + resource = "https://my-workspace.cloud.databricks.com/oidc/v1/token" # Set expiration far in the future future_expiry = int(datetime(2025, 7, 12, 12, 0, 0).timestamp()) + 10000 hook.oauth_tokens[resource] = { @@ -1469,33 +1537,37 @@ async def test_a_get_federated_token_cached_valid(self): assert token == "cached_async_token" @pytest.mark.asyncio - async def test_a_get_federated_token_k8s_not_available(self): + @mock.patch("airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.a_databricks_conn") + async def test_a_get_federated_token_k8s_not_available(self, mock_conn): """Test async error when Kubernetes service account token is not available.""" - mock_conn = mock.Mock() - mock_conn.host = "my-workspace.cloud.databricks.com" - mock_conn.login = "federated_k8s" - mock_conn.extra_dejson = {"client_id": "test-client-id"} + conn = Connection( + host="my-workspace.cloud.databricks.com", + login="federated_k8s", + extra={"client_id": "test-client-id"}, + ) + mock_conn.return_value = conn hook = BaseDatabricksHook() - hook.databricks_conn = mock_conn - resource = f"https://{mock_conn.host}/oidc/v1/token" + resource = f"https://{conn.host}/oidc/v1/token" with mock.patch("aiofiles.open", side_effect=FileNotFoundError()): with pytest.raises(AirflowException, match="Kubernetes service account token not found"): await hook._a_get_federated_databricks_token(resource) @pytest.mark.asyncio - async def test_a_get_federated_token_missing_client_id(self): + @mock.patch("airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.a_databricks_conn") + async def test_a_get_federated_token_missing_client_id(self, mock_conn): """Test async error when client_id is missing from connection extra.""" - mock_conn = mock.Mock() - mock_conn.host = "my-workspace.cloud.databricks.com" - mock_conn.login = "federated_k8s" - mock_conn.extra_dejson = {} # Missing client_id + conn = Connection( + host="my-workspace.cloud.databricks.com", + login="federated_k8s", + extra={}, # Missing client_id + ) + mock_conn.return_value = conn hook = BaseDatabricksHook() - hook.databricks_conn = mock_conn - resource = f"https://{mock_conn.host}/oidc/v1/token" + resource = f"https://{conn.host}/oidc/v1/token" with pytest.raises( AirflowException, match="client_id is required for Kubernetes OIDC token federation" ): @@ -1504,8 +1576,16 @@ async def test_a_get_federated_token_missing_client_id(self): @pytest.mark.asyncio @mock.patch("aiohttp.ClientSession.post") @mock.patch("ssl.create_default_context") - async def test_a_get_federated_token_databricks_error(self, _mock_ssl_ctx, mock_post): + @mock.patch("airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.a_databricks_conn") + async def test_a_get_federated_token_databricks_error(self, _mock_ssl_ctx, mock_conn, mock_post): """Test async error handling when Databricks token exchange fails.""" + conn = Connection( + host="my-workspace.cloud.databricks.com", + login="federated_k8s", + extra={"client_id": "test-client-id"}, + ) + mock_conn.return_value = conn + # Mock aiofiles.open for reading K8s token and namespace mock_token_file = mock.AsyncMock() mock_token_file.read = mock.AsyncMock(return_value="in_cluster_token") @@ -1539,33 +1619,28 @@ async def test_a_get_federated_token_databricks_error(self, _mock_ssl_ctx, mock_ mock_post.side_effect = [k8s_response, db_response] - mock_conn = mock.Mock() - mock_conn.host = "my-workspace.cloud.databricks.com" - mock_conn.login = "federated_k8s" - mock_conn.extra_dejson = {"client_id": "test-client-id"} - hook = BaseDatabricksHook() - hook.databricks_conn = mock_conn hook.user_agent_header = {"User-Agent": "test-agent"} hook.token_timeout_seconds = 10 with mock.patch("aiofiles.open", side_effect=[mock_token_file, mock_namespace_file]): async with aiohttp.ClientSession() as session: hook._session = session - resource = f"https://{mock_conn.host}/oidc/v1/token" + resource = f"https://{conn.host}/oidc/v1/token" with pytest.raises( AirflowException, match="Failed to exchange Kubernetes JWT for Databricks token" ): await hook._a_get_federated_databricks_token(resource) @pytest.mark.asyncio - async def test_a_get_k8s_projected_volume_token_success(self): + @mock.patch("airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.a_databricks_conn") + async def test_a_get_k8s_projected_volume_token_success(self, mock_conn): """Test async successfully reading token from Kubernetes projected volume.""" - mock_conn = mock.Mock() - mock_conn.extra_dejson = {"k8s_projected_volume_token_path": "/var/run/secrets/databricks/token"} + mock_conn.return_value = Connection( + extra={"k8s_projected_volume_token_path": "/var/run/secrets/databricks/token"} + ) hook = BaseDatabricksHook() - hook.databricks_conn = mock_conn # Mock aiofiles.open mock_file = mock.AsyncMock() @@ -1579,13 +1654,14 @@ async def test_a_get_k8s_projected_volume_token_success(self): assert token == "projected_token_content" @pytest.mark.asyncio - async def test_a_get_k8s_projected_volume_token_file_not_found(self): + @mock.patch("airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.a_databricks_conn") + async def test_a_get_k8s_projected_volume_token_file_not_found(self, mock_conn): """Test async error when projected volume token file is not found.""" - mock_conn = mock.Mock() - mock_conn.extra_dejson = {"k8s_projected_volume_token_path": "/var/run/secrets/databricks/token"} + mock_conn.return_value = Connection( + extra={"k8s_projected_volume_token_path": "/var/run/secrets/databricks/token"} + ) hook = BaseDatabricksHook() - hook.databricks_conn = mock_conn with mock.patch("aiofiles.open", side_effect=FileNotFoundError()): with pytest.raises( @@ -1595,13 +1671,14 @@ async def test_a_get_k8s_projected_volume_token_file_not_found(self): await hook._a_get_k8s_projected_volume_token() @pytest.mark.asyncio - async def test_a_get_k8s_projected_volume_token_permission_denied(self): + @mock.patch("airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.a_databricks_conn") + async def test_a_get_k8s_projected_volume_token_permission_denied(self, mock_conn): """Test async error when permission denied reading projected volume token.""" - mock_conn = mock.Mock() - mock_conn.extra_dejson = {"k8s_projected_volume_token_path": "/var/run/secrets/databricks/token"} + mock_conn.return_value = Connection( + extra={"k8s_projected_volume_token_path": "/var/run/secrets/databricks/token"} + ) hook = BaseDatabricksHook() - hook.databricks_conn = mock_conn with mock.patch("aiofiles.open", side_effect=PermissionError()): with pytest.raises( @@ -1611,13 +1688,14 @@ async def test_a_get_k8s_projected_volume_token_permission_denied(self): await hook._a_get_k8s_projected_volume_token() @pytest.mark.asyncio - async def test_a_get_k8s_projected_volume_token_empty_file(self): + @mock.patch("airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.a_databricks_conn") + async def test_a_get_k8s_projected_volume_token_empty_file(self, mock_conn): """Test async error when projected volume token file is empty.""" - mock_conn = mock.Mock() - mock_conn.extra_dejson = {"k8s_projected_volume_token_path": "/var/run/secrets/databricks/token"} + mock_conn.return_value = Connection( + extra={"k8s_projected_volume_token_path": "/var/run/secrets/databricks/token"} + ) hook = BaseDatabricksHook() - hook.databricks_conn = mock_conn # Mock aiofiles.open with empty content mock_file = mock.AsyncMock() @@ -1632,13 +1710,14 @@ async def test_a_get_k8s_projected_volume_token_empty_file(self): await hook._a_get_k8s_projected_volume_token() @pytest.mark.asyncio - async def test_a_get_k8s_jwt_token_uses_projected_volume_when_configured(self): + @mock.patch("airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.a_databricks_conn") + async def test_a_get_k8s_jwt_token_uses_projected_volume_when_configured(self, mock_conn): """Test that async _a_get_k8s_jwt_token delegates to projected volume method when configured.""" - mock_conn = mock.Mock() - mock_conn.extra_dejson = {"k8s_projected_volume_token_path": "/var/run/secrets/databricks/token"} + mock_conn.return_value = Connection( + extra={"k8s_projected_volume_token_path": "/var/run/secrets/databricks/token"} + ) hook = BaseDatabricksHook() - hook.databricks_conn = mock_conn with mock.patch.object( hook, @@ -1656,13 +1735,12 @@ async def test_a_get_k8s_jwt_token_uses_projected_volume_when_configured(self): mock_token_request.assert_not_called() @pytest.mark.asyncio - async def test_a_get_k8s_jwt_token_uses_token_request_api_when_no_projected_path(self): + @mock.patch("airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.a_databricks_conn") + async def test_a_get_k8s_jwt_token_uses_token_request_api_when_no_projected_path(self, mock_conn): """Test that async _a_get_k8s_jwt_token delegates to TokenRequest API when no projected path.""" - mock_conn = mock.Mock() - mock_conn.extra_dejson = {} # No projected volume path + mock_conn.return_value = Connection(extra={}) # No projected volume path hook = BaseDatabricksHook() - hook.databricks_conn = mock_conn with mock.patch.object( hook, "_a_get_k8s_projected_volume_token", new_callable=mock.AsyncMock @@ -1681,17 +1759,17 @@ async def test_a_get_k8s_jwt_token_uses_token_request_api_when_no_projected_path @pytest.mark.asyncio @mock.patch("ssl.create_default_context") - async def test_a_get_k8s_token_request_api_uses_ca_cert_for_tls(self, mock_ssl_ctx): + @mock.patch("airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.a_databricks_conn") + async def test_a_get_k8s_token_request_api_uses_ca_cert_for_tls(self, mock_conn, mock_ssl_ctx): """Verify async TokenRequest API call uses the in-cluster CA bundle for TLS verification.""" import ssl fake_ctx = mock.MagicMock(spec=ssl.SSLContext) mock_ssl_ctx.return_value = fake_ctx - mock_conn = mock.Mock() - mock_conn.extra_dejson = {} + mock_conn.return_value = Connection(extra={}) hook = BaseDatabricksHook() - hook.databricks_conn = mock_conn + hook._a_databricks_conn = mock_conn mock_response = mock.AsyncMock() mock_response.json = mock.AsyncMock(return_value={"status": {"token": "jwt_token"}}) @@ -1720,8 +1798,9 @@ async def test_a_get_k8s_token_request_api_uses_ca_cert_for_tls(self, mock_ssl_c assert call_kwargs["ssl"] is fake_ctx @pytest.mark.asyncio + @mock.patch("airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.a_databricks_conn") @time_machine.travel("2025-07-12 12:00:00") - async def test_a_get_federated_token_with_projected_volume(self): + async def test_a_get_federated_token_with_projected_volume(self, mock_conn): """Test async end-to-end federated token flow using projected volume.""" # Mock Databricks token exchange response expiry_date = int((datetime(2025, 7, 12, 12, 0, 0) + timedelta(minutes=60)).timestamp()) @@ -1731,16 +1810,17 @@ async def test_a_get_federated_token_with_projected_volume(self): "token_type": "Bearer", } - mock_conn = mock.Mock() - mock_conn.host = "my-workspace.cloud.databricks.com" - mock_conn.login = "federated_k8s" - mock_conn.extra_dejson = { - "k8s_projected_volume_token_path": "/var/run/secrets/databricks/token", - "client_id": "test-client-id", - } + conn = Connection( + host="my-workspace.cloud.databricks.com", + login="federated_k8s", + extra={ + "k8s_projected_volume_token_path": "/var/run/secrets/databricks/token", + "client_id": "test-client-id", + }, + ) + mock_conn.return_value = conn hook = BaseDatabricksHook() - hook.databricks_conn = mock_conn hook.user_agent_header = {"User-Agent": "test-agent"} # Mock aiofiles.open for projected volume @@ -1764,7 +1844,7 @@ async def test_a_get_federated_token_with_projected_volume(self): "aiofiles.open", return_value=mock.MagicMock(__aenter__=mock.AsyncMock(return_value=mock_file)), ): - resource = f"https://{mock_conn.host}/oidc/v1/token" + resource = f"https://{conn.host}/oidc/v1/token" token = await hook._a_get_federated_databricks_token(resource) assert token == "databricks_token" @@ -1783,8 +1863,16 @@ async def test_a_get_federated_token_with_projected_volume(self): @pytest.mark.asyncio @mock.patch("aiohttp.ClientSession.post") @mock.patch("ssl.create_default_context") - async def test_a_get_token_with_federated_k8s_login(self, _mock_ssl_ctx, mock_post): + @mock.patch("airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.a_databricks_conn") + async def test_a_get_token_with_federated_k8s_login(self, mock_conn, _mock_ssl_ctx, mock_post): """Test async _a_get_token with login='federated_k8s'.""" + mock_conn.return_value = Connection( + host="my-workspace.cloud.databricks.com", + login="federated_k8s", + password=None, + extra={"client_id": "test-client-id"}, + ) + # Mock aiofiles.open for reading K8s token and namespace mock_token_file = mock.AsyncMock() mock_token_file.read = mock.AsyncMock(return_value="in_cluster_token") @@ -1816,14 +1904,7 @@ async def test_a_get_token_with_federated_k8s_login(self, _mock_ssl_ctx, mock_po mock_post.side_effect = [k8s_response, db_response] - mock_conn = mock.Mock() - mock_conn.host = "my-workspace.cloud.databricks.com" - mock_conn.login = "federated_k8s" - mock_conn.password = None - mock_conn.extra_dejson = {"client_id": "test-client-id"} - hook = BaseDatabricksHook() - hook.databricks_conn = mock_conn hook.user_agent_header = {"User-Agent": "test-agent"} hook.token_timeout_seconds = 10 @@ -1837,8 +1918,16 @@ async def test_a_get_token_with_federated_k8s_login(self, _mock_ssl_ctx, mock_po @pytest.mark.asyncio @mock.patch("aiohttp.ClientSession.post") @mock.patch("ssl.create_default_context") - async def test_a_get_token_with_federated_k8s_extra(self, _mock_ssl_ctx, mock_post): + @mock.patch("airflow.providers.databricks.hooks.databricks_base.BaseDatabricksHook.a_databricks_conn") + async def test_a_get_token_with_federated_k8s_extra(self, mock_conn, _mock_ssl_ctx, mock_post): """Test async _a_get_token with federated_k8s in extras.""" + mock_conn.return_value = Connection( + host="my-workspace.cloud.databricks.com", + login=None, + password=None, + extra={"federated_k8s": True, "client_id": "test-client-id"}, + ) + # Mock aiofiles.open for reading K8s token and namespace mock_token_file = mock.AsyncMock() mock_token_file.read = mock.AsyncMock(return_value="in_cluster_token") @@ -1870,14 +1959,7 @@ async def test_a_get_token_with_federated_k8s_extra(self, _mock_ssl_ctx, mock_po mock_post.side_effect = [k8s_response, db_response] - mock_conn = mock.Mock() - mock_conn.host = "my-workspace.cloud.databricks.com" - mock_conn.login = None - mock_conn.password = None - mock_conn.extra_dejson = {"federated_k8s": True, "client_id": "test-client-id"} - hook = BaseDatabricksHook() - hook.databricks_conn = mock_conn hook.user_agent_header = {"User-Agent": "test-agent"} hook.token_timeout_seconds = 10