Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
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 @@ -393,20 +393,21 @@ def __exit__(
spark_submit_params=self.spark_submit_params,
)

for task in tasks:
if not (
hasattr(task, "_convert_to_databricks_workflow_task")
and callable(task._convert_to_databricks_workflow_task)
):
raise AirflowException(
f"Task {task.task_id} does not support conversion to databricks workflow task."
)

task.workflow_run_metadata = create_databricks_workflow_task.output
create_databricks_workflow_task.relevant_upstreams.append(task.task_id)
create_databricks_workflow_task.add_task(task.task_id, task)

for root_task in roots:
root_task.set_upstream(create_databricks_workflow_task)

super().__exit__(_type, _value, _tb)
try:
for task in tasks:
if not (
hasattr(task, "_convert_to_databricks_workflow_task")
and callable(task._convert_to_databricks_workflow_task)
):
raise AirflowException(
f"Task {task.task_id} does not support conversion to databricks workflow task."
)

task.workflow_run_metadata = create_databricks_workflow_task.output
create_databricks_workflow_task.relevant_upstreams.append(task.task_id)
create_databricks_workflow_task.add_task(task.task_id, task)

for root_task in roots:
root_task.set_upstream(create_databricks_workflow_task)
finally:
super().__exit__(_type, _value, _tb)
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,8 @@

import pytest

from tests_common.test_utils.version_compat import AIRFLOW_V_3_0_PLUS

# Do not run the tests when FAB / Flask is not installed
pytest.importorskip("flask_session")

Expand Down Expand Up @@ -276,6 +278,32 @@ def test_task_group_exit_creates_operator(mock_databricks_workflow_operator):
)


@pytest.mark.skipif(not AIRFLOW_V_3_0_PLUS, reason="Uses Airflow 3 task SDK TaskGroupContext layout")
def test_task_group_context_cleaned_up_on_internal_exception():
"""
Regression test for GH-42164.

When DatabricksWorkflowTaskGroup.__exit__ raises (e.g. an added task does not
support conversion), super().__exit__ must still run so TaskGroupContext does
not leak the workflow group onto the global stack and break later DAGs with
"Cannot mix TaskGroups from different DAGs".
"""
from airflow.sdk.definitions._internal.contextmanager import TaskGroupContext

TaskGroupContext._context.clear()

with pytest.raises(AirflowException, match="does not support conversion"): # noqa: PT012 raise happens on context exit
with DAG(dag_id="example_databricks_workflow_dag_err", schedule=None, start_date=DEFAULT_DATE):
with DatabricksWorkflowTaskGroup(
group_id="test_databricks_workflow_err", databricks_conn_id="databricks_conn"
):
# EmptyOperator does not implement _convert_to_databricks_workflow_task,
# which makes DatabricksWorkflowTaskGroup.__exit__ raise mid-way.
EmptyOperator(task_id="not_convertible")

assert not TaskGroupContext._context, "TaskGroupContext leaked the workflow task group"


def test_task_group_root_tasks_set_upstream_to_operator(mock_databricks_workflow_operator):
"""Test that tasks added to a DatabricksWorkflowTaskGroup are set upstream to the operator."""
with DAG(dag_id="example_databricks_workflow_dag", schedule=None, start_date=DEFAULT_DATE):
Expand Down
Loading