Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
20 commits
Select commit Hold shift + click to select a range
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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"
Expand Down Expand Up @@ -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)

Expand All @@ -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

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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 = {
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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()

Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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")

Expand All @@ -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],
Expand Down Expand Up @@ -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}
Expand All @@ -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":
Expand Down
Loading
Loading