This is an automated email from the ASF dual-hosted git repository.
eladkal 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 e8ba9507034 Add Databricks-native retry settings to task operators
(#69182)
e8ba9507034 is described below
commit e8ba9507034c5a17c720b3d6d08662b6d7ecbe0b
Author: Nick <[email protected]>
AuthorDate: Fri Aug 7 09:01:54 2026 -0500
Add Databricks-native retry settings to task operators (#69182)
* Add Databricks-native retry settings to task operators
* Fix Databricks task retry completion handling
* Warn on unparsable Databricks max_retries value
---------
Co-authored-by: Jarek Potiuk <[email protected]>
---
providers/databricks/docs/operators/notebook.rst | 38 ++
providers/databricks/docs/operators/task.rst | 27 +
.../providers/databricks/operators/databricks.py | 275 +++++++--
.../providers/databricks/triggers/databricks.py | 101 +++-
.../unit/databricks/operators/test_databricks.py | 646 +++++++++++++++++++++
.../unit/databricks/triggers/test_databricks.py | 201 +++++++
6 files changed, 1251 insertions(+), 37 deletions(-)
diff --git a/providers/databricks/docs/operators/notebook.rst
b/providers/databricks/docs/operators/notebook.rst
index 83dd899b412..3c8ec27b8d3 100644
--- a/providers/databricks/docs/operators/notebook.rst
+++ b/providers/databricks/docs/operators/notebook.rst
@@ -42,3 +42,41 @@ Running a notebook in Databricks on an existing cluster
:language: python
:start-after: [START howto_operator_databricks_notebook_existing_cluster]
:end-before: [END howto_operator_databricks_notebook_existing_cluster]
+
+Configuring Databricks-native task retries
+-------------------------------------------
+
+Use ``max_retries``, ``min_retry_interval_millis`` and ``retry_on_timeout`` to
configure
+`Databricks-native task retries
<https://docs.databricks.com/api/workspace/jobs/create#tasks-max_retries>`_.
+Databricks reruns failed task attempts within the same job run, so Airflow
sees only the final result.
+Set ``max_retries`` to ``-1`` to retry indefinitely, or ``0`` to disable
retries.
+
+These settings are independent of the Airflow task-level ``retries``
parameter, which retries the
+whole Airflow task:
+
+.. code-block:: python
+
+ DatabricksNotebookOperator(
+ task_id="notebook",
+ notebook_path="/path/to/notebook",
+ source="WORKSPACE",
+ existing_cluster_id="existing_cluster_id",
+ max_retries=3,
+ min_retry_interval_millis=2000,
+ retry_on_timeout=True,
+ )
+
+Airflow ``retries`` behaves differently depending on where the operator runs.
For a standalone
+operator, each retry submits a new Databricks run. Inside a
+:class:`~airflow.providers.databricks.operators.databricks_workflow.DatabricksWorkflowTaskGroup`,
+the Airflow task monitors a sub-run that was already submitted by the workflow
launch task, so a
+retry only re-polls the terminal sub-run. Use ``max_retries`` to retry
Databricks work inside a
+workflow task group.
+
+Inside a
+:class:`~airflow.providers.databricks.operators.databricks_workflow.DatabricksWorkflowTaskGroup`,
+a task that exhausts a finite ``max_retries`` is reported as failed as soon as
its final failed
+attempt is observed, so downstream failure handling is not delayed by
long-running sibling tasks.
+Only unlimited retries (``max_retries=-1``) keep the Airflow task waiting (or
deferring) until the
+parent workflow run reaches a terminal state, because Databricks may still
launch a retry attempt
+under the same ``task_key`` until then. Sibling tasks in the run continue
independently.
diff --git a/providers/databricks/docs/operators/task.rst
b/providers/databricks/docs/operators/task.rst
index 5c446593531..b327aa20ccf 100644
--- a/providers/databricks/docs/operators/task.rst
+++ b/providers/databricks/docs/operators/task.rst
@@ -44,3 +44,30 @@ Running a SQL query in Databricks using
DatabricksTaskOperator
:language: python
:start-after: [START howto_operator_databricks_task_sql]
:end-before: [END howto_operator_databricks_task_sql]
+
+Configuring Databricks-native task retries
+-------------------------------------------
+
+Use ``max_retries``, ``min_retry_interval_millis`` and ``retry_on_timeout`` to
configure
+`Databricks-native task retries
<https://docs.databricks.com/api/workspace/jobs/create#tasks-max_retries>`_.
+Databricks reruns failed task attempts within the same job run, so Airflow
sees only the final result.
+Set ``max_retries`` to ``-1`` to retry indefinitely, or ``0`` to disable
retries.
+
+These settings are independent of the Airflow task-level ``retries``
parameter, which retries the
+whole Airflow task. You can set the same fields directly in ``task_config``.
When both are set, the
+operator parameter takes precedence. If a field is unset, Databricks uses its
default.
+
+Airflow ``retries`` behaves differently depending on where the operator runs.
For a standalone
+operator, each retry submits a new Databricks run. Inside a
+:class:`~airflow.providers.databricks.operators.databricks_workflow.DatabricksWorkflowTaskGroup`,
+the Airflow task monitors a sub-run that was already submitted by the workflow
launch task, so a
+retry only re-polls the terminal sub-run. Use ``max_retries`` to retry
Databricks work inside a
+workflow task group.
+
+Inside a
+:class:`~airflow.providers.databricks.operators.databricks_workflow.DatabricksWorkflowTaskGroup`,
+a task that exhausts a finite ``max_retries`` is reported as failed as soon as
its final failed
+attempt is observed, so downstream failure handling is not delayed by
long-running sibling tasks.
+Only unlimited retries (``max_retries=-1``) keep the Airflow task waiting (or
deferring) until the
+parent workflow run reaches a terminal state, because Databricks may still
launch a retry attempt
+under the same ``task_key`` until then. Sibling tasks in the run continue
independently.
diff --git
a/providers/databricks/src/airflow/providers/databricks/operators/databricks.py
b/providers/databricks/src/airflow/providers/databricks/operators/databricks.py
index a64acb78598..d820b099c8d 100644
---
a/providers/databricks/src/airflow/providers/databricks/operators/databricks.py
+++
b/providers/databricks/src/airflow/providers/databricks/operators/databricks.py
@@ -1673,6 +1673,9 @@ class DatabricksTaskBaseOperator(BaseOperator, ABC):
:param wait_for_termination: if we should wait for termination of the job
run. ``True`` by default.
:param workflow_run_metadata: Metadata for the workflow run. This is used
when the operator is used within
a workflow. It is expected to be a dictionary containing the run_id
and conn_id for the workflow.
+ :param max_retries: Databricks task retry count. Use ``-1`` for unlimited
retries.
+ :param min_retry_interval_millis: Minimum interval between Databricks task
retries.
+ :param retry_on_timeout: Whether Databricks retries timed-out tasks.
"""
def __init__(
@@ -1690,6 +1693,9 @@ class DatabricksTaskBaseOperator(BaseOperator, ABC):
polling_period_seconds: int = 5,
wait_for_termination: bool = True,
workflow_run_metadata: dict[str, Any] | None = None,
+ max_retries: int | str | None = None,
+ min_retry_interval_millis: int | str | None = None,
+ retry_on_timeout: bool | str | None = None,
**kwargs: Any,
):
self.caller = caller
@@ -1705,6 +1711,9 @@ class DatabricksTaskBaseOperator(BaseOperator, ABC):
self.polling_period_seconds = polling_period_seconds
self.wait_for_termination = wait_for_termination
self.workflow_run_metadata = workflow_run_metadata
+ self.max_retries = max_retries
+ self.min_retry_interval_millis = min_retry_interval_millis
+ self.retry_on_timeout = retry_on_timeout
self.databricks_run_id: int | None = None
@@ -1792,21 +1801,75 @@ class DatabricksTaskBaseOperator(BaseOperator, ABC):
"""Get the base json for the task."""
raise NotImplementedError()
+ def _retry_settings(self) -> dict[str, Any]:
+ """Databricks-native task retry settings that were explicitly
provided."""
+ settings: dict[str, Any] = {}
+ if self.max_retries is not None:
+ settings["max_retries"] = int(self.max_retries)
+ if self.min_retry_interval_millis is not None:
+ settings["min_retry_interval_millis"] =
int(self.min_retry_interval_millis)
+ if self.retry_on_timeout is not None:
+ retry_on_timeout = self.retry_on_timeout
+ if isinstance(retry_on_timeout, str):
+ retry_on_timeout = retry_on_timeout.strip().lower() in
("true", "1", "yes")
+ settings["retry_on_timeout"] = bool(retry_on_timeout)
+ return settings
+
+ def _has_retry_settings(self) -> bool:
+ """Whether any Databricks-native retry field is explicitly
configured."""
+ if self._retry_settings():
+ return True
+ task_config = getattr(self, "task_config", {}) or {}
+ return any(
+ task_config.get(key) is not None
+ for key in ("max_retries", "min_retry_interval_millis",
"retry_on_timeout")
+ )
+
+ def _resolved_max_retries(self) -> int | None:
+ """Return the resolved ``max_retries`` value, with operator args
taking precedence."""
+ task_config = getattr(self, "task_config", {}) or {}
+ max_retries = self.max_retries if self.max_retries is not None else
task_config.get("max_retries")
+
+ if isinstance(max_retries, bool) or max_retries is None:
+ return None
+ if isinstance(max_retries, str):
+ try:
+ return int(max_retries)
+ except ValueError:
+ self.log.warning(
+ "Ignoring unparsable max_retries value %r; treating as no
retry limit configured.",
+ max_retries,
+ )
+ return None
+ if not isinstance(max_retries, int):
+ return None
+ return max_retries
+
def _get_run_json(self) -> dict[str, Any]:
"""Get run json to be used for task submissions."""
- run_json = {
- "run_name": self.databricks_task_key,
- **self._get_task_base_json(),
- }
if self.new_cluster and self.existing_cluster_id:
raise ValueError("Both new_cluster and existing_cluster_id are
set. Only one should be set.")
+ cluster: dict[str, Any]
if self.new_cluster:
- run_json["new_cluster"] = self.new_cluster
+ cluster = {"new_cluster": self.new_cluster}
elif self.existing_cluster_id:
- run_json["existing_cluster_id"] = self.existing_cluster_id
+ cluster = {"existing_cluster_id": self.existing_cluster_id}
else:
raise ValueError("Must specify either existing_cluster_id or
new_cluster.")
- return run_json
+
+ if not self._has_retry_settings():
+ # No retry settings: keep the legacy single-task runs/submit shape
unchanged.
+ return {"run_name": self.databricks_task_key,
**self._get_task_base_json(), **cluster}
+
+ # Retry settings are per-task SubmitTask fields, so submit the single
task explicitly.
+ # The explicit task_key also gives monitoring a stable task to look up.
+ task = {
+ **self._get_task_base_json(),
+ "task_key": self.databricks_task_key,
+ **self._retry_settings(),
+ **cluster,
+ }
+ return {"run_name": self.databricks_task_key, "tasks": [task]}
def _launch_job(self, context: Context | None = None) -> int | None:
"""Launch the job on Databricks."""
@@ -1857,13 +1920,14 @@ class DatabricksTaskBaseOperator(BaseOperator, ABC):
base_task_json = self._get_task_base_json()
result = {
- "task_key": self.databricks_task_key,
"depends_on": [
{"task_key": self._generate_databricks_task_key(task_id,
task_dict)}
for task_id in self.upstream_task_ids
if task_id in relevant_upstreams
],
**base_task_json,
+ "task_key": self.databricks_task_key,
+ **self._retry_settings(),
}
trigger_rule_value = (
@@ -1899,49 +1963,170 @@ class DatabricksTaskBaseOperator(BaseOperator, ABC):
def monitor_databricks_job(self) -> None:
"""
- Monitor the Databricks job.
-
- Wait for the job to terminate. If deferrable, defer the task.
+ Monitor the Databricks job until it terminates and surface its result.
+
+ Picks one of three monitoring strategies, depending on whether
Databricks-native retry
+ attempts can occur and whether the operator runs inside a
``DatabricksWorkflowTaskGroup``:
+
+ * standalone with native retries -> :meth:`_monitor_submit_run`,
following the submit run
+ whose own terminal state already accounts for every retry attempt.
+ * workflow task with native retries -> :meth:`_monitor_workflow_task`,
following the task's
+ latest attempt and tolerating in-flight retries until they are
exhausted (or, for unlimited
+ retries, until the shared run is terminal).
+ * otherwise -> :meth:`_monitor_single_attempt`, following one attempt
and reporting as soon
+ as it terminates (the historical behaviour, unchanged).
"""
if self.databricks_run_id is None:
raise ValueError("Databricks job not yet launched. Please run
launch_notebook_job first.")
- current_task_run_id = self._get_current_databricks_task()["run_id"]
- run = self._hook.get_run(current_task_run_id)
- run_page_url = run["run_page_url"]
- self.log.info("Check the task run in Databricks: %s", run_page_url)
- run_state = RunState(**run["state"])
+
+ max_retries = self._resolved_max_retries()
+ if max_retries is None or (max_retries != -1 and max_retries <= 0):
+ self._monitor_single_attempt()
+ elif self._databricks_workflow_task_group is None:
+ self._monitor_submit_run()
+ else:
+ self._monitor_workflow_task()
+
+ def _log_task_state(self, run_state: RunState) -> None:
self.log.info(
"Current state of the databricks task %s is %s",
self.databricks_task_key,
run_state.life_cycle_state,
)
+
+ def _defer_on_run(
+ self,
+ run_id: int,
+ *,
+ workflow_run_id: int | None = None,
+ databricks_task_key: str | None = None,
+ max_retries: int | None = None,
+ ) -> None:
+ """Defer monitoring of ``run_id`` to the trigger, optionally with
workflow-task context."""
+ self.defer(
+ trigger=DatabricksExecutionTrigger(
+ run_id=run_id,
+ databricks_conn_id=self.databricks_conn_id,
+ polling_period_seconds=self.polling_period_seconds,
+ retry_limit=self.databricks_retry_limit,
+ retry_delay=self.databricks_retry_delay,
+ retry_args=self.databricks_retry_args,
+ caller=self.caller,
+ workflow_run_id=workflow_run_id,
+ databricks_task_key=databricks_task_key,
+ max_retries=max_retries,
+ ),
+ method_name=DEFER_METHOD_NAME,
+ )
+
+ def _monitor_single_attempt(self) -> None:
+ """Follow this task's attempt run and report as soon as it terminates
(no native retries)."""
+ current_task_run_id = self._get_current_databricks_task()["run_id"]
+ run = self._hook.get_run(current_task_run_id)
+ self.log.info("Check the task run in Databricks: %s",
run["run_page_url"])
+ run_state = RunState(**run["state"])
+ self._log_task_state(run_state)
+
if self.deferrable and not run_state.is_terminal:
- self.defer(
- trigger=DatabricksExecutionTrigger(
- run_id=current_task_run_id,
- databricks_conn_id=self.databricks_conn_id,
- polling_period_seconds=self.polling_period_seconds,
- retry_limit=self.databricks_retry_limit,
- retry_delay=self.databricks_retry_delay,
- retry_args=self.databricks_retry_args,
- caller=self.caller,
- ),
- method_name=DEFER_METHOD_NAME,
- )
+ self._defer_on_run(current_task_run_id)
+
while not run_state.is_terminal:
time.sleep(self.polling_period_seconds)
run = self._hook.get_run(current_task_run_id)
run_state = RunState(**run["state"])
+ self._log_task_state(run_state)
- self.log.info(
- "Current state of the databricks task %s is %s",
- self.databricks_task_key,
- run_state.life_cycle_state,
+ errors = extract_failed_task_errors(self._hook, run, run_state)
+ self._handle_terminal_run_state(run_state, errors)
+
+ def _monitor_workflow_task(self) -> None:
+ """
+ Follow this task's latest attempt within a shared workflow run,
tolerating native retries.
+
+ Inside a ``DatabricksWorkflowTaskGroup`` the run holds sibling tasks,
so the operator must
+ report when its own task finishes rather than wait for the whole run.
A failed attempt is
+ final once finite retries are exhausted; unlimited retries fall back
to the workflow run's
+ terminal state. Each poll re-resolves the latest attempt for the task
key.
+ """
+ workflow_run_id = self.databricks_run_id
+ if workflow_run_id is None:
+ raise ValueError("Databricks job not yet launched. Please run
launch_notebook_job first.")
+ current_task = self._get_current_databricks_task()
+ current_task_run_id = current_task["run_id"]
+ run = self._hook.get_run(current_task_run_id)
+ self.log.info("Check the task run in Databricks: %s",
run["run_page_url"])
+ run_state = RunState(**run["state"])
+ self._log_task_state(run_state)
+ attempt_number = current_task.get("attempt_number")
+
+ # Defer whenever the outcome is not yet conclusive: a failed attempt
with retries still
+ # available means a retry may follow, and the trigger waits for it
without blocking a worker.
+ if self.deferrable and not self._workflow_task_is_conclusive(
+ run_state, workflow_run_id, attempt_number
+ ):
+ self._defer_on_run(
+ current_task_run_id,
+ workflow_run_id=workflow_run_id,
+ databricks_task_key=self.databricks_task_key,
+ max_retries=self._resolved_max_retries(),
)
- # Extract errors from the run response using utility function
+ while not self._workflow_task_is_conclusive(run_state,
workflow_run_id, attempt_number):
+ time.sleep(self.polling_period_seconds)
+ current_task = self._get_current_databricks_task()
+ current_task_run_id = current_task["run_id"]
+ run = self._hook.get_run(current_task_run_id)
+ run_state = RunState(**run["state"])
+ self._log_task_state(run_state)
+ attempt_number = current_task.get("attempt_number")
+
errors = extract_failed_task_errors(self._hook, run, run_state)
+ self._handle_terminal_run_state(run_state, errors)
+
+ def _workflow_task_is_conclusive(
+ self, run_state: RunState, workflow_run_id: int, attempt_number: int |
None
+ ) -> bool:
+ """Whether the attempt is final: succeeded, retries exhausted, or the
run is terminal."""
+ if not run_state.is_terminal:
+ return False
+ if run_state.is_successful:
+ return True
+ max_retries = self._resolved_max_retries()
+ if (
+ max_retries is not None
+ and max_retries != -1
+ and attempt_number is not None
+ and attempt_number >= max_retries
+ ):
+ return True
+ parent_state = RunState(**self._hook.get_run(workflow_run_id)["state"])
+ return parent_state.is_terminal
+
+ def _monitor_submit_run(self) -> None:
+ """
+ Wait for a standalone submit run to terminate, tolerating
Databricks-native retries.
+
+ The submit run owns exactly this operator's task, so its own terminal
state already
+ accounts for every native retry attempt — we follow the run rather
than any single attempt.
+ """
+ run_id = self.databricks_run_id
+ if run_id is None:
+ raise ValueError("Databricks job not yet launched. Please run
launch_notebook_job first.")
+ run = self._hook.get_run(run_id)
+ self.log.info("Check the job run in Databricks: %s",
run["run_page_url"])
+ run_state = RunState(**run["state"])
+ self.log.info("Current state of the databricks run %s is %s", run_id,
run_state.life_cycle_state)
+ if self.deferrable and not run_state.is_terminal:
+ self._defer_on_run(run_id)
+
+ while not run_state.is_terminal:
+ time.sleep(self.polling_period_seconds)
+ run = self._hook.get_run(run_id)
+ run_state = RunState(**run["state"])
+ self.log.info("Current state of the databricks run %s is %s",
run_id, run_state.life_cycle_state)
+
+ errors = extract_failed_task_errors(self._hook, run, run_state)
self._handle_terminal_run_state(run_state, errors)
def execute(self, context: Context) -> None:
@@ -2033,11 +2218,17 @@ class
DatabricksNotebookOperator(DatabricksTaskBaseOperator):
:param wait_for_termination: if we should wait for termination of the job
run. ``True`` by default.
:param workflow_run_metadata: Metadata for the workflow run. This is used
when the operator is used within
a workflow. It is expected to be a dictionary containing the run_id
and conn_id for the workflow.
+ :param max_retries: Databricks task retry count. Use ``-1`` for unlimited
retries.
+ :param min_retry_interval_millis: Minimum interval between Databricks task
retries.
+ :param retry_on_timeout: Whether Databricks retries timed-out tasks.
"""
template_fields = (
"notebook_params",
"workflow_run_metadata",
+ "max_retries",
+ "min_retry_interval_millis",
+ "retry_on_timeout",
)
CALLER = "DatabricksNotebookOperator"
@@ -2058,6 +2249,9 @@ class
DatabricksNotebookOperator(DatabricksTaskBaseOperator):
polling_period_seconds: int = 5,
wait_for_termination: bool = True,
workflow_run_metadata: dict | None = None,
+ max_retries: int | str | None = None,
+ min_retry_interval_millis: int | str | None = None,
+ retry_on_timeout: bool | str | None = None,
**kwargs: Any,
):
self.notebook_path = notebook_path
@@ -2078,6 +2272,9 @@ class
DatabricksNotebookOperator(DatabricksTaskBaseOperator):
polling_period_seconds=polling_period_seconds,
wait_for_termination=wait_for_termination,
workflow_run_metadata=workflow_run_metadata,
+ max_retries=max_retries,
+ min_retry_interval_millis=min_retry_interval_millis,
+ retry_on_timeout=retry_on_timeout,
**kwargs,
)
@@ -2176,6 +2373,9 @@ class DatabricksTaskOperator(DatabricksTaskBaseOperator):
:param new_cluster: Specs for a new cluster on which this task will be run.
:param polling_period_seconds: Controls the rate which we poll for the
result of this notebook job run.
:param wait_for_termination: if we should wait for termination of the job
run. ``True`` by default.
+ :param max_retries: Databricks task retry count, overriding
``task_config`` when set.
+ :param min_retry_interval_millis: Minimum retry interval, overriding
``task_config`` when set.
+ :param retry_on_timeout: Whether Databricks retries timed-out tasks,
overriding ``task_config`` when set.
"""
CALLER = "DatabricksTaskOperator"
@@ -2183,6 +2383,9 @@ class DatabricksTaskOperator(DatabricksTaskBaseOperator):
"databricks_conn_id",
"task_config",
"workflow_run_metadata",
+ "max_retries",
+ "min_retry_interval_millis",
+ "retry_on_timeout",
)
def __init__(
@@ -2199,6 +2402,9 @@ class DatabricksTaskOperator(DatabricksTaskBaseOperator):
polling_period_seconds: int = 5,
wait_for_termination: bool = True,
workflow_run_metadata: dict | None = None,
+ max_retries: int | str | None = None,
+ min_retry_interval_millis: int | str | None = None,
+ retry_on_timeout: bool | str | None = None,
**kwargs,
):
self.task_config = task_config
@@ -2216,6 +2422,9 @@ class DatabricksTaskOperator(DatabricksTaskBaseOperator):
polling_period_seconds=polling_period_seconds,
wait_for_termination=wait_for_termination,
workflow_run_metadata=workflow_run_metadata,
+ max_retries=max_retries,
+ min_retry_interval_millis=min_retry_interval_millis,
+ retry_on_timeout=retry_on_timeout,
**kwargs,
)
diff --git
a/providers/databricks/src/airflow/providers/databricks/triggers/databricks.py
b/providers/databricks/src/airflow/providers/databricks/triggers/databricks.py
index 57d489550f2..9706d2e4adf 100644
---
a/providers/databricks/src/airflow/providers/databricks/triggers/databricks.py
+++
b/providers/databricks/src/airflow/providers/databricks/triggers/databricks.py
@@ -43,6 +43,9 @@ class DatabricksExecutionTrigger(BaseTrigger):
:param run_page_url: The run page url.
:param repair_run: Repair the databricks run in case of failure.
:param caller: The name of the operator that is calling the hook.
+ :param workflow_run_id: Parent workflow run ID for task-level monitoring.
+ :param databricks_task_key: Task key to monitor within ``workflow_run_id``.
+ :param max_retries: Resolved Databricks-native ``max_retries`` for
task-level monitoring.
"""
def __init__(
@@ -56,6 +59,9 @@ class DatabricksExecutionTrigger(BaseTrigger):
run_page_url: str | None = None,
repair_run: bool = False,
caller: str = "DatabricksExecutionTrigger",
+ workflow_run_id: int | None = None,
+ databricks_task_key: str | None = None,
+ max_retries: int | None = None,
) -> None:
super().__init__()
# Trigger kwargs cross Airflow's serialization boundary, so fail
before storing invalid
@@ -70,6 +76,9 @@ class DatabricksExecutionTrigger(BaseTrigger):
self.run_page_url = run_page_url
self.repair_run = repair_run
self.caller = caller
+ self.workflow_run_id = workflow_run_id
+ self.databricks_task_key = databricks_task_key
+ self.max_retries = max_retries
self.hook = DatabricksHook(
databricks_conn_id,
retry_limit=self.retry_limit,
@@ -91,19 +100,40 @@ class DatabricksExecutionTrigger(BaseTrigger):
"run_page_url": self.run_page_url,
"repair_run": self.repair_run,
"caller": self.caller,
+ "workflow_run_id": self.workflow_run_id,
+ "databricks_task_key": self.databricks_task_key,
+ "max_retries": self.max_retries,
},
)
async def on_kill(self) -> None:
"""Cancel the Databricks run when the trigger is cancelled by a user
action."""
- if self.run_id:
- from asgiref.sync import sync_to_async
+ from asgiref.sync import sync_to_async
+
+ run_id = self.run_id
+ if self.workflow_run_id is not None and self.databricks_task_key is
not None:
+ # self.run_id may be an earlier, now-terminal attempt; cancel the
task's latest attempt
+ # so a retry/repair launched under the same task_key is not left
running.
+ tasks = await
sync_to_async(self.hook.get_run_tasks)(self.workflow_run_id)
+ attempt = {
+ task["task_key"]: task for task in sorted(tasks, key=lambda
task: task["start_time"])
+ }.get(self.databricks_task_key)
+ if attempt:
+ run_id = attempt["run_id"]
+ if run_id:
+ self.log.info("Cancelling Databricks run %s.", run_id)
+ await sync_to_async(self.hook.cancel_run)(run_id)
- self.log.info("Cancelling Databricks run %s.", self.run_id)
- await sync_to_async(self.hook.cancel_run)(self.run_id)
+ def _monitors_workflow_task(self) -> bool:
+ """Whether this trigger follows one task inside a shared workflow run
(see ``workflow_run_id``)."""
+ return bool(self.workflow_run_id and self.databricks_task_key)
async def run(self):
async with self.hook:
+ if self._monitors_workflow_task():
+ async for event in self._run_workflow_task():
+ yield event
+ return
while True:
run_state = await self.hook.a_get_run_state(self.run_id)
if not run_state.is_terminal:
@@ -129,6 +159,69 @@ class DatabricksExecutionTrigger(BaseTrigger):
)
return
+ async def _run_workflow_task(self):
+ """Monitor one task in a workflow run, tolerating in-flight
retries/repairs."""
+ from asgiref.sync import sync_to_async
+
+ while True:
+ tasks = await
sync_to_async(self.hook.get_run_tasks)(self.workflow_run_id)
+ sorted_task_runs = sorted(tasks, key=lambda task:
task["start_time"])
+ attempt = {task["task_key"]: task for task in
sorted_task_runs}.get(self.databricks_task_key)
+
+ if attempt is not None:
+ attempt_run_id = attempt["run_id"]
+ attempt_state = await self.hook.a_get_run_state(attempt_run_id)
+
+ if attempt_state.is_terminal:
+ if attempt_state.is_successful:
+ yield TriggerEvent(
+ {
+ "run_id": attempt_run_id,
+ "run_page_url": self.run_page_url,
+ "run_state": attempt_state.to_json(),
+ "repair_run": self.repair_run,
+ "errors": [],
+ }
+ )
+ return
+ # A failed attempt is final once finite retries are
exhausted; otherwise wait
+ # for the parent run because another attempt may still
appear.
+ attempt_number = attempt.get("attempt_number")
+ retries_exhausted = (
+ self.max_retries is not None
+ and self.max_retries != -1
+ and attempt_number is not None
+ and attempt_number >= self.max_retries
+ )
+ if (
+ retries_exhausted
+ or (await
self.hook.a_get_run_state(self.workflow_run_id)).is_terminal
+ ):
+ run_info = await self.hook.a_get_run(attempt_run_id)
+ failed_tasks = await extract_failed_task_errors_async(
+ self.hook, run_info, attempt_state
+ )
+ yield TriggerEvent(
+ {
+ "run_id": attempt_run_id,
+ "run_page_url": self.run_page_url,
+ "run_state": attempt_state.to_json(),
+ "repair_run": self.repair_run,
+ "errors": failed_tasks,
+ }
+ )
+ return
+
+ # attempt is None when the task has not yet surfaced in the run
(e.g. just after launch);
+ # keep polling rather than crashing on a missing task_key.
+ self.log.info(
+ "databricks task %s not yet conclusive in run %s. sleeping for
%s seconds",
+ self.databricks_task_key,
+ self.workflow_run_id,
+ self.polling_period_seconds,
+ )
+ await asyncio.sleep(self.polling_period_seconds)
+
class DatabricksSQLStatementExecutionTrigger(BaseTrigger):
"""
diff --git
a/providers/databricks/tests/unit/databricks/operators/test_databricks.py
b/providers/databricks/tests/unit/databricks/operators/test_databricks.py
index 65c990a5926..0f190da3661 100644
--- a/providers/databricks/tests/unit/databricks/operators/test_databricks.py
+++ b/providers/databricks/tests/unit/databricks/operators/test_databricks.py
@@ -3687,6 +3687,103 @@ class TestDatabricksNotebookOperator:
"Trigger is not a DatabricksExecutionTrigger"
)
assert exec_info.value.method_name == "execute_complete"
+ # Without native retries configured, the trigger keeps its original
parent-unaware behavior.
+ assert exec_info.value.trigger.workflow_run_id is None
+ assert exec_info.value.trigger.databricks_task_key is None
+
+
@mock.patch("airflow.providers.databricks.operators.databricks.DatabricksHook")
+ def test_execute_standalone_with_retries_defers_on_submit_run(self,
mock_databricks_hook):
+ # A standalone operator with native retries follows its own submit run
to a terminal state,
+ # so the defer targets that run directly (no per-task workflow
context).
+ mock_databricks_hook.return_value.get_run.return_value = {
+ "state": {"life_cycle_state": "PENDING"},
+ "run_page_url": "test_url",
+ }
+ operator = DatabricksNotebookOperator(
+ task_id="test_task",
+ notebook_path="test_path",
+ source="test_source",
+ databricks_conn_id="test_conn_id",
+ wait_for_termination=True,
+ deferrable=True,
+ max_retries=2,
+ )
+ operator.databricks_run_id = 12345
+
+ with pytest.raises(TaskDeferred) as exec_info:
+ operator.monitor_databricks_job()
+ assert exec_info.value.trigger.run_id == 12345
+ assert exec_info.value.trigger.workflow_run_id is None
+ assert exec_info.value.trigger.databricks_task_key is None
+
+ @mock.patch(
+
"airflow.providers.databricks.operators.databricks.DatabricksNotebookOperator._databricks_workflow_task_group",
+ new_callable=mock.PropertyMock,
+ )
+
@mock.patch("airflow.providers.databricks.operators.databricks.DatabricksHook")
+ @mock.patch(
+
"airflow.providers.databricks.operators.databricks.DatabricksNotebookOperator._get_current_databricks_task"
+ )
+ def
test_execute_with_deferrable_passes_workflow_context_when_retries_configured(
+ self, mock_get_current_task, mock_databricks_hook, mock_workflow_tg
+ ):
+ # Inside a workflow task group the run is shared, so the defer carries
the workflow context
+ # and the trigger tracks this task's own attempt within that run.
+ mock_workflow_tg.return_value = MagicMock()
+ mock_get_current_task.return_value = {"run_id": "attempt-1"}
+ mock_databricks_hook.return_value.get_run.return_value = {
+ "state": {"life_cycle_state": "PENDING"},
+ "run_page_url": "test_url",
+ }
+ operator = DatabricksNotebookOperator(
+ task_id="test_task",
+ notebook_path="test_path",
+ source="test_source",
+ databricks_conn_id="test_conn_id",
+ wait_for_termination=True,
+ deferrable=True,
+ max_retries=2,
+ )
+ operator.databricks_run_id = 12345
+
+ with pytest.raises(TaskDeferred) as exec_info:
+ operator.monitor_databricks_job()
+ assert exec_info.value.trigger.workflow_run_id == 12345
+ assert exec_info.value.trigger.databricks_task_key ==
operator.databricks_task_key
+
+ @mock.patch(
+
"airflow.providers.databricks.operators.databricks.DatabricksNotebookOperator._databricks_workflow_task_group",
+ new_callable=mock.PropertyMock,
+ )
+
@mock.patch("airflow.providers.databricks.operators.databricks.DatabricksHook")
+ @mock.patch(
+
"airflow.providers.databricks.operators.databricks.DatabricksNotebookOperator._get_current_databricks_task"
+ )
+ def
test_execute_with_deferrable_and_max_retries_zero_keeps_single_attempt_monitoring(
+ self, mock_get_current_task, mock_databricks_hook, mock_workflow_tg
+ ):
+ mock_workflow_tg.return_value = MagicMock()
+ mock_get_current_task.return_value = {"run_id": "attempt-1"}
+ mock_databricks_hook.return_value.get_run.return_value = {
+ "state": {"life_cycle_state": "PENDING"},
+ "run_page_url": "test_url",
+ }
+ operator = DatabricksNotebookOperator(
+ task_id="test_task",
+ notebook_path="test_path",
+ source="test_source",
+ databricks_conn_id="test_conn_id",
+ wait_for_termination=True,
+ deferrable=True,
+ max_retries=0,
+ )
+ operator.databricks_run_id = 12345
+
+ with pytest.raises(TaskDeferred) as exec_info:
+ operator.monitor_databricks_job()
+ assert exec_info.value.trigger.run_id == "attempt-1"
+ assert exec_info.value.trigger.workflow_run_id is None
+ assert exec_info.value.trigger.databricks_task_key is None
@mock.patch("airflow.providers.databricks.operators.databricks.DatabricksHook")
@mock.patch(
@@ -3762,6 +3859,279 @@ class TestDatabricksNotebookOperator:
exception_message = "Task failed. Final state FAILED. Reason: FAILURE.
Errors: []"
assert exception_message == str(exc_info.value)
+ @mock.patch(
+
"airflow.providers.databricks.operators.databricks.DatabricksNotebookOperator._databricks_workflow_task_group",
+ new_callable=mock.PropertyMock,
+ )
+ @mock.patch("airflow.providers.databricks.operators.databricks.time.sleep")
+
@mock.patch("airflow.providers.databricks.operators.databricks.DatabricksHook")
+ @mock.patch(
+
"airflow.providers.databricks.operators.databricks.DatabricksNotebookOperator._get_current_databricks_task"
+ )
+ def test_monitor_databricks_job_retry_in_flight_succeeds(
+ self, mock_get_current_task, mock_databricks_hook, mock_sleep,
mock_workflow_tg
+ ):
+ # Inside a workflow task group the run is shared, so the task tracks
its own attempt: the
+ # first attempt fails while the workflow run is still active; a
retried attempt then
+ # succeeds and the Airflow task must not fail.
+ mock_workflow_tg.return_value = MagicMock()
+ mock_get_current_task.side_effect = [{"run_id": "attempt-1"},
{"run_id": "attempt-2"}]
+ runs = {
+ "attempt-1": {
+ "state": {
+ "life_cycle_state": "TERMINATED",
+ "result_state": "FAILED",
+ "state_message": "first attempt failed",
+ },
+ "run_page_url": "url-1",
+ },
+ 12345: {"state": {"life_cycle_state": "RUNNING"}, "run_page_url":
"parent"},
+ "attempt-2": {
+ "state": {"life_cycle_state": "TERMINATED", "result_state":
"SUCCESS"},
+ "run_page_url": "url-2",
+ },
+ }
+ mock_databricks_hook.return_value.get_run.side_effect = lambda run_id:
runs[run_id]
+
+ operator = DatabricksNotebookOperator(
+ task_id="test_task",
+ notebook_path="test_path",
+ source="test_source",
+ databricks_conn_id="test_conn_id",
+ max_retries=1,
+ )
+ operator.databricks_run_id = 12345
+
+ operator.monitor_databricks_job()
+ mock_sleep.assert_called_once()
+
+ @mock.patch("airflow.providers.databricks.operators.databricks.time.sleep")
+
@mock.patch("airflow.providers.databricks.operators.databricks.DatabricksHook")
+ def test_monitor_standalone_submit_run_follows_run_to_terminal_state(
+ self, mock_databricks_hook, mock_sleep
+ ):
+ # A standalone submit run stays active while Databricks retries the
task; the operator
+ # follows the run (not an attempt) and reports only once the run
itself terminates.
+ run_states = iter(
+ [
+ {"life_cycle_state": "RUNNING"},
+ {"life_cycle_state": "TERMINATED", "result_state": "SUCCESS"},
+ ]
+ )
+ mock_databricks_hook.return_value.get_run.side_effect = lambda run_id:
{
+ "state": next(run_states),
+ "run_page_url": "url",
+ }
+
+ operator = DatabricksNotebookOperator(
+ task_id="test_task",
+ notebook_path="test_path",
+ source="test_source",
+ databricks_conn_id="test_conn_id",
+ max_retries=1,
+ )
+ operator.databricks_run_id = 12345
+
+ operator.monitor_databricks_job()
+ # Only the submit run is polled; the per-attempt resolver is never
used.
+ mock_databricks_hook.return_value.get_run.assert_called_with(12345)
+ mock_databricks_hook.return_value.get_run_tasks.assert_not_called()
+
+ @mock.patch("airflow.providers.databricks.operators.databricks.time.sleep")
+
@mock.patch("airflow.providers.databricks.operators.databricks.DatabricksHook")
+ @mock.patch(
+
"airflow.providers.databricks.operators.databricks.DatabricksNotebookOperator._get_current_databricks_task"
+ )
+ def test_monitor_databricks_job_without_retries_fails_immediately(
+ self, mock_get_current_task, mock_databricks_hook, mock_sleep
+ ):
+ # Without native retries configured the task fails as soon as its own
attempt terminates,
+ # even if the parent run is still active. The parent run state must
not be polled.
+ mock_get_current_task.return_value = {"run_id": "attempt-1"}
+ mock_databricks_hook.return_value.get_run.return_value = {
+ "state": {
+ "life_cycle_state": "TERMINATED",
+ "result_state": "FAILED",
+ "state_message": "attempt failed",
+ },
+ "run_page_url": "url",
+ }
+
+ operator = DatabricksNotebookOperator(
+ task_id="test_task",
+ notebook_path="test_path",
+ source="test_source",
+ databricks_conn_id="test_conn_id",
+ )
+ operator.databricks_run_id = 12345
+
+ with pytest.raises(AirflowException):
+ operator.monitor_databricks_job()
+ mock_sleep.assert_not_called()
+ # Only the attempt run is polled (once, before the loop); the parent
run is never fetched.
+
mock_databricks_hook.return_value.get_run.assert_called_once_with("attempt-1")
+
+ @mock.patch(
+
"airflow.providers.databricks.operators.databricks.DatabricksNotebookOperator._databricks_workflow_task_group",
+ new_callable=mock.PropertyMock,
+ )
+ @mock.patch("airflow.providers.databricks.operators.databricks.time.sleep")
+
@mock.patch("airflow.providers.databricks.operators.databricks.DatabricksHook")
+ @mock.patch(
+
"airflow.providers.databricks.operators.databricks.DatabricksNotebookOperator._get_current_databricks_task"
+ )
+ def test_monitor_workflow_task_reports_failure_once_retries_exhausted(
+ self, mock_get_current_task, mock_databricks_hook, mock_sleep,
mock_workflow_tg
+ ):
+ mock_workflow_tg.return_value = MagicMock()
+ mock_get_current_task.return_value = {"run_id": "attempt-2",
"attempt_number": 1}
+ mock_databricks_hook.return_value.get_run.return_value = {
+ "state": {
+ "life_cycle_state": "TERMINATED",
+ "result_state": "FAILED",
+ "state_message": "final attempt failed",
+ },
+ "run_page_url": "url",
+ }
+
+ operator = DatabricksNotebookOperator(
+ task_id="test_task",
+ notebook_path="test_path",
+ source="test_source",
+ databricks_conn_id="test_conn_id",
+ max_retries=1,
+ )
+ operator.databricks_run_id = 12345
+
+ with pytest.raises(AirflowException):
+ operator.monitor_databricks_job()
+ mock_sleep.assert_not_called()
+
mock_databricks_hook.return_value.get_run.assert_called_once_with("attempt-2")
+
+ @mock.patch(
+
"airflow.providers.databricks.operators.databricks.DatabricksNotebookOperator._databricks_workflow_task_group",
+ new_callable=mock.PropertyMock,
+ )
+ @mock.patch("airflow.providers.databricks.operators.databricks.time.sleep")
+
@mock.patch("airflow.providers.databricks.operators.databricks.DatabricksHook")
+ @mock.patch(
+
"airflow.providers.databricks.operators.databricks.DatabricksNotebookOperator._get_current_databricks_task"
+ )
+ def test_monitor_workflow_task_unlimited_retries_waits_for_parent(
+ self, mock_get_current_task, mock_databricks_hook, mock_sleep,
mock_workflow_tg
+ ):
+ mock_workflow_tg.return_value = MagicMock()
+ mock_get_current_task.return_value = {"run_id": "attempt-1",
"attempt_number": 5}
+ parent_states = iter(
+ [
+ {"life_cycle_state": "RUNNING"},
+ {"life_cycle_state": "TERMINATED", "result_state": "FAILED",
"state_message": "failed"},
+ ]
+ )
+
+ def fake_get_run(run_id):
+ if run_id == 12345:
+ return {"state": next(parent_states), "run_page_url": "parent"}
+ return {
+ "state": {
+ "life_cycle_state": "TERMINATED",
+ "result_state": "FAILED",
+ "state_message": "failed",
+ },
+ "run_page_url": "url",
+ }
+
+ mock_databricks_hook.return_value.get_run.side_effect = fake_get_run
+
+ operator = DatabricksNotebookOperator(
+ task_id="test_task",
+ notebook_path="test_path",
+ source="test_source",
+ databricks_conn_id="test_conn_id",
+ max_retries=-1,
+ )
+ operator.databricks_run_id = 12345
+
+ with pytest.raises(AirflowException):
+ operator.monitor_databricks_job()
+ mock_sleep.assert_called_once()
+ mock_databricks_hook.return_value.get_run.assert_any_call(12345)
+
+ @mock.patch("airflow.providers.databricks.operators.databricks.time.sleep")
+
@mock.patch("airflow.providers.databricks.operators.databricks.DatabricksHook")
+ def
test_reshaped_submit_run_without_native_retries_resolves_task_for_monitoring(
+ self, mock_databricks_hook, mock_sleep
+ ):
+ operator = DatabricksNotebookOperator(
+ task_id="test_task",
+ notebook_path="test_path",
+ source="test_source",
+ databricks_conn_id="test_conn_id",
+ existing_cluster_id="existing_cluster_id",
+ min_retry_interval_millis=2000,
+ )
+ run_json = operator._get_run_json()
+ assert "tasks" in run_json
+ assert run_json["tasks"][0]["task_key"] == operator.databricks_task_key
+ assert operator._resolved_max_retries() is None
+
+ operator.databricks_run_id = 12345
+ mock_databricks_hook.return_value.get_run_tasks.return_value = [
+ {"task_key": operator.databricks_task_key, "run_id": "attempt-1",
"start_time": 1}
+ ]
+ mock_databricks_hook.return_value.get_run.return_value = {
+ "state": {"life_cycle_state": "TERMINATED", "result_state":
"SUCCESS"},
+ "run_page_url": "url",
+ }
+
+ operator.monitor_databricks_job()
+
mock_databricks_hook.return_value.get_run_tasks.assert_called_once_with(12345)
+
mock_databricks_hook.return_value.get_run.assert_called_once_with("attempt-1")
+
+ @pytest.mark.parametrize(
+ "operator",
+ [
+ DatabricksNotebookOperator(
+ task_id="test_task",
+ notebook_path="test_path",
+ source="WORKSPACE",
+ databricks_conn_id="test_conn_id",
+ existing_cluster_id="existing_cluster_id",
+ max_retries="{{ params.max_retries }}",
+ min_retry_interval_millis="{{ params.min_retry_interval_millis
}}",
+ retry_on_timeout="{{ params.retry_on_timeout }}",
+ ),
+ DatabricksTaskOperator(
+ task_id="test_task",
+ databricks_conn_id="test_conn_id",
+ task_config={},
+ max_retries="{{ params.max_retries }}",
+ min_retry_interval_millis="{{ params.min_retry_interval_millis
}}",
+ retry_on_timeout="{{ params.retry_on_timeout }}",
+ ),
+ ],
+ )
+ def test_retry_params_render_from_templates(self, operator):
+ """Retry fields are templatable: Jinja values render and coerce to
typed retry settings.
+
+ Guards the ``template_fields`` membership on both subclasses (which
keep separate tuples):
+ an un-templated field would keep its literal ``{{ ... }}`` string and
fail coercion.
+ """
+ operator.render_template_fields(
+ context={
+ "params": {
+ "max_retries": 3,
+ "min_retry_interval_millis": 2000,
+ "retry_on_timeout": "true",
+ }
+ }
+ )
+ assert operator._retry_settings() == {
+ "max_retries": 3,
+ "min_retry_interval_millis": 2000,
+ "retry_on_timeout": True,
+ }
+
@mock.patch("airflow.providers.databricks.operators.databricks.DatabricksHook")
def test_launch_notebook_job(self, mock_databricks_hook):
operator = DatabricksNotebookOperator(
@@ -3914,6 +4284,88 @@ class TestDatabricksNotebookOperator:
assert task_json == expected_json
+ @pytest.mark.parametrize(
+ ("retry_kwargs", "expected"),
+ [
+ (
+ {"max_retries": -1, "min_retry_interval_millis": 2000,
"retry_on_timeout": True},
+ {"max_retries": -1, "min_retry_interval_millis": 2000,
"retry_on_timeout": True},
+ ),
+ ({"max_retries": 0}, {"max_retries": 0}),
+ ({"retry_on_timeout": False}, {"retry_on_timeout": False}),
+ ],
+ )
+ def test_get_run_json_retry_settings(self, retry_kwargs, expected):
+ """Retry settings are added to the submitted task only when explicitly
provided.
+
+ They must live inside ``tasks[0]`` (a Databricks ``SubmitTask``); at
the top level of a
+ runs/submit payload Databricks silently ignores them.
+ """
+ operator = DatabricksNotebookOperator(
+ task_id="test_task",
+ notebook_path="test_path",
+ source="test_source",
+ databricks_conn_id="test_conn_id",
+ existing_cluster_id="existing_cluster_id",
+ **retry_kwargs,
+ )
+ run_json = operator._get_run_json()
+ assert "tasks" in run_json
+ task = run_json["tasks"][0]
+ assert task["task_key"] == operator.databricks_task_key
+ assert task["existing_cluster_id"] == "existing_cluster_id"
+ assert task["notebook_task"]["notebook_path"] == "test_path"
+ assert "existing_cluster_id" not in run_json
+ for key in ("max_retries", "min_retry_interval_millis",
"retry_on_timeout"):
+ assert key not in run_json
+ if key in expected:
+ assert task[key] == expected[key]
+ else:
+ assert key not in task
+
+ def test_get_run_json_without_retries_uses_legacy_top_level_shape(self):
+ """Without native retries the payload is the legacy single-task
runs/submit shape (unchanged)."""
+ operator = DatabricksNotebookOperator(
+ task_id="test_task",
+ notebook_path="test_path",
+ source="test_source",
+ databricks_conn_id="test_conn_id",
+ existing_cluster_id="existing_cluster_id",
+ )
+ run_json = operator._get_run_json()
+ assert run_json["run_name"] == operator.databricks_task_key
+ assert run_json["existing_cluster_id"] == "existing_cluster_id"
+ assert run_json["notebook_task"]["notebook_path"] == "test_path"
+ # No reshape into the multi-task form when retries are not configured.
+ assert "tasks" not in run_json
+
+ def test_convert_to_databricks_workflow_task_includes_retry_settings(self):
+ """Retry settings provided to the operator are included in the
workflow task JSON."""
+ dag = DAG(dag_id="example_dag", schedule=None, start_date=DEFAULT_DATE)
+ operator = DatabricksNotebookOperator(
+ notebook_path="/path/to/notebook",
+ source="WORKSPACE",
+ task_id="test_task",
+ max_retries=2,
+ min_retry_interval_millis=1000,
+ retry_on_timeout=True,
+ dag=dag,
+ )
+
+ databricks_workflow_task_group = MagicMock()
+ databricks_workflow_task_group.notebook_packages = []
+ databricks_workflow_task_group.notebook_params = {}
+
+ operator.task_group = databricks_workflow_task_group
+ relevant_upstreams = []
+ task_dict = {}
+
+ task_json =
operator._convert_to_databricks_workflow_task(relevant_upstreams, task_dict)
+
+ assert task_json["max_retries"] == 2
+ assert task_json["min_retry_interval_millis"] == 1000
+ assert task_json["retry_on_timeout"] is True
+
@pytest.mark.parametrize(
("trigger_rule", "expected_run_if"),
[
@@ -4054,6 +4506,69 @@ class TestDatabricksTaskOperator:
assert operator.task_config == task_config
assert task_base_json == task_config
+ def test_get_run_json_operator_task_key_wins_over_task_config(self):
+ """A task_key in task_config must not shadow the operator-managed key
used for monitoring.
+
+ The operator only injects its own task_key in the reshaped
(native-retry) payload, where
+ monitoring re-resolves the task by that key, so the guarantee is
exercised with retries set.
+ """
+ operator = DatabricksTaskOperator(
+ task_id="test_task",
+ databricks_conn_id="test_conn_id",
+ existing_cluster_id="existing_cluster_id",
+ max_retries=2,
+ task_config={"task_key": "user_supplied_key", "notebook_task":
{"notebook_path": "/p"}},
+ )
+ task = operator._get_run_json()["tasks"][0]
+ assert task["task_key"] == operator.databricks_task_key
+
+ def
test_convert_to_databricks_workflow_task_includes_task_config_retry_settings(self):
+ """Retry settings supplied via task_config surface in the workflow
task JSON."""
+ dag = DAG(dag_id="example_dag", schedule=None, start_date=DEFAULT_DATE)
+ operator = DatabricksTaskOperator(
+ task_id="test_task",
+ databricks_conn_id="test_conn_id",
+ task_config={
+ "notebook_task": {"notebook_path": "/path"},
+ "max_retries": 5,
+ "min_retry_interval_millis": 1000,
+ "retry_on_timeout": True,
+ },
+ dag=dag,
+ )
+ operator.task_group = MagicMock()
+
+ task_json = operator._convert_to_databricks_workflow_task([], {})
+
+ assert task_json["max_retries"] == 5
+ assert task_json["min_retry_interval_millis"] == 1000
+ assert task_json["retry_on_timeout"] is True
+
+ def
test_convert_to_databricks_workflow_task_operator_retry_overrides_task_config(self):
+ """An operator-level retry value overrides the task_config value in
the workflow task JSON."""
+ dag = DAG(dag_id="example_dag", schedule=None, start_date=DEFAULT_DATE)
+ operator = DatabricksTaskOperator(
+ task_id="test_task",
+ databricks_conn_id="test_conn_id",
+ task_config={
+ "notebook_task": {"notebook_path": "/path"},
+ "max_retries": 5,
+ "min_retry_interval_millis": 1000,
+ "retry_on_timeout": False,
+ },
+ max_retries=2,
+ min_retry_interval_millis=2000,
+ retry_on_timeout=True,
+ dag=dag,
+ )
+ operator.task_group = MagicMock()
+
+ task_json = operator._convert_to_databricks_workflow_task([], {})
+
+ assert task_json["max_retries"] == 2
+ assert task_json["min_retry_interval_millis"] == 2000
+ assert task_json["retry_on_timeout"] is True
+
def test_generate_databricks_task_key(self):
task_config = {}
operator = DatabricksTaskOperator(
@@ -4154,3 +4669,134 @@ class TestDatabricksTaskOperator:
):
operator.on_kill()
db_mock.cancel_run.assert_not_called()
+
+ @pytest.mark.parametrize(
+ ("field", "task_config_retries", "retry_kwargs", "expected"),
+ [
+ ("max_retries", {"max_retries": 1}, {}, 1),
+ ("max_retries", {"max_retries": 1}, {"max_retries": 0}, 0),
+ (
+ "min_retry_interval_millis",
+ {"min_retry_interval_millis": 1000},
+ {"min_retry_interval_millis": 5000},
+ 5000,
+ ),
+ ("retry_on_timeout", {"retry_on_timeout": True},
{"retry_on_timeout": False}, False),
+ ],
+ )
+ def test_get_run_json_retry_settings(self, field, task_config_retries,
retry_kwargs, expected):
+ """Operator-level retry settings are applied and take precedence over
task_config values."""
+ task_config = {
+ "notebook_task": {"notebook_path": "/path", "source": "WORKSPACE",
"base_parameters": {}},
+ **task_config_retries,
+ }
+ operator = DatabricksTaskOperator(
+ task_id="test_task",
+ databricks_conn_id="test_conn_id",
+ existing_cluster_id="existing_cluster_id",
+ task_config=task_config,
+ **retry_kwargs,
+ )
+ run_json = operator._get_run_json()
+ assert run_json["tasks"][0][field] == expected
+ assert field not in run_json
+
+ @pytest.mark.parametrize("max_retries", [2, "2"])
+ @mock.patch("airflow.providers.databricks.operators.databricks.time.sleep")
+
@mock.patch("airflow.providers.databricks.operators.databricks.DatabricksHook")
+ def test_monitor_uses_retry_strategy_when_retries_set_in_task_config(
+ self, mock_databricks_hook, mock_sleep, max_retries
+ ):
+ # max_retries supplied only through task_config must still select
retry-aware monitoring
+ # (follow the submit run), not the single-attempt path.
+ run_states = iter(
+ [
+ {"life_cycle_state": "RUNNING"},
+ {"life_cycle_state": "TERMINATED", "result_state": "SUCCESS"},
+ ]
+ )
+ mock_databricks_hook.return_value.get_run.side_effect = lambda run_id:
{
+ "state": next(run_states),
+ "run_page_url": "url",
+ }
+
+ operator = DatabricksTaskOperator(
+ task_id="test_task",
+ databricks_conn_id="test_conn_id",
+ existing_cluster_id="existing_cluster_id",
+ task_config={
+ "notebook_task": {"notebook_path": "/path", "source":
"WORKSPACE", "base_parameters": {}},
+ "max_retries": max_retries,
+ },
+ )
+ operator.databricks_run_id = 12345
+
+ operator.monitor_databricks_job()
+ mock_databricks_hook.return_value.get_run.assert_called_with(12345)
+ mock_databricks_hook.return_value.get_run_tasks.assert_not_called()
+
+ @mock.patch("airflow.providers.databricks.operators.databricks.time.sleep")
+
@mock.patch("airflow.providers.databricks.operators.databricks.DatabricksHook")
+ @mock.patch(
+
"airflow.providers.databricks.operators.databricks.DatabricksTaskOperator._get_current_databricks_task"
+ )
+ def
test_monitor_max_retries_zero_operator_arg_overrides_task_config_retries(
+ self, mock_get_current_task, mock_databricks_hook, mock_sleep
+ ):
+ mock_get_current_task.return_value = {"run_id": "attempt-1"}
+ mock_databricks_hook.return_value.get_run.return_value = {
+ "state": {
+ "life_cycle_state": "TERMINATED",
+ "result_state": "FAILED",
+ "state_message": "attempt failed",
+ },
+ "run_page_url": "url",
+ }
+
+ operator = DatabricksTaskOperator(
+ task_id="test_task",
+ databricks_conn_id="test_conn_id",
+ existing_cluster_id="existing_cluster_id",
+ task_config={
+ "notebook_task": {"notebook_path": "/path", "source":
"WORKSPACE", "base_parameters": {}},
+ "max_retries": 2,
+ },
+ max_retries=0,
+ )
+ operator.databricks_run_id = 12345
+
+ with pytest.raises(AirflowException):
+ operator.monitor_databricks_job()
+ mock_sleep.assert_not_called()
+
mock_databricks_hook.return_value.get_run.assert_called_once_with("attempt-1")
+
+ @mock.patch("airflow.providers.databricks.operators.databricks.time.sleep")
+
@mock.patch("airflow.providers.databricks.operators.databricks.DatabricksHook")
+ def test_monitor_waits_out_waiting_for_retry_state(self,
mock_databricks_hook, mock_sleep):
+ # A native retry surfaces WAITING_FOR_RETRY between attempts; the poll
must keep waiting
+ # rather than crash on an unexpected life cycle state.
+ run_states = iter(
+ [
+ {"life_cycle_state": "RUNNING"},
+ {"life_cycle_state": "WAITING_FOR_RETRY"},
+ {"life_cycle_state": "TERMINATED", "result_state": "SUCCESS"},
+ ]
+ )
+ mock_databricks_hook.return_value.get_run.side_effect = lambda run_id:
{
+ "state": next(run_states),
+ "run_page_url": "url",
+ }
+
+ operator = DatabricksTaskOperator(
+ task_id="test_task",
+ databricks_conn_id="test_conn_id",
+ existing_cluster_id="existing_cluster_id",
+ task_config={
+ "notebook_task": {"notebook_path": "/path", "source":
"WORKSPACE", "base_parameters": {}},
+ "max_retries": 2,
+ },
+ )
+ operator.databricks_run_id = 12345
+
+ operator.monitor_databricks_job()
+ assert mock_databricks_hook.return_value.get_run.call_count == 3
diff --git
a/providers/databricks/tests/unit/databricks/triggers/test_databricks.py
b/providers/databricks/tests/unit/databricks/triggers/test_databricks.py
index 8854eb03fb5..ad204297dc8 100644
--- a/providers/databricks/tests/unit/databricks/triggers/test_databricks.py
+++ b/providers/databricks/tests/unit/databricks/triggers/test_databricks.py
@@ -189,6 +189,9 @@ class TestDatabricksExecutionTrigger:
"run_page_url": RUN_PAGE_URL,
"repair_run": False,
"caller": "DatabricksExecutionTrigger",
+ "workflow_run_id": None,
+ "databricks_task_key": None,
+ "max_retries": None,
},
)
@@ -306,12 +309,210 @@ class TestDatabricksExecutionTrigger:
mock_sleep.assert_called_once()
mock_sleep.assert_called_with(POLLING_INTERVAL_SECONDS)
+ @pytest.mark.asyncio
+
@mock.patch("airflow.providers.databricks.hooks.databricks.DatabricksHook.a_get_run_output")
+
@mock.patch("airflow.providers.databricks.hooks.databricks.DatabricksHook.a_get_run")
+
@mock.patch("airflow.providers.databricks.triggers.databricks.asyncio.sleep")
+
@mock.patch("airflow.providers.databricks.hooks.databricks.DatabricksHook.a_get_run_state")
+ async def test_run_waits_out_waiting_for_retry_state(
+ self, mock_get_run_state, mock_sleep, mock_get_run, mock_get_run_output
+ ):
+ # A native retry surfaces WAITING_FOR_RETRY between attempts; the
trigger must keep polling
+ # rather than crash on an unexpected life cycle state.
+ mock_get_run_state.side_effect = [
+ RunState(life_cycle_state="WAITING_FOR_RETRY", state_message="",
result_state=""),
+ RunState(life_cycle_state=LIFE_CYCLE_STATE_TERMINATED,
state_message="", result_state="SUCCESS"),
+ ]
+ mock_get_run.return_value = GET_RUN_RESPONSE_TERMINATED
+ mock_get_run_output.return_value = GET_RUN_OUTPUT_RESPONSE
+
+ async for event in self.trigger.run():
+ assert event == TriggerEvent(
+ {
+ "run_id": RUN_ID,
+ "run_state": RunState(
+ life_cycle_state=LIFE_CYCLE_STATE_TERMINATED,
state_message="", result_state="SUCCESS"
+ ).to_json(),
+ "run_page_url": RUN_PAGE_URL,
+ "repair_run": False,
+ "errors": [],
+ }
+ )
+ mock_sleep.assert_called_once_with(POLLING_INTERVAL_SECONDS)
+
+ @pytest.mark.asyncio
+
@mock.patch("airflow.providers.databricks.triggers.databricks.asyncio.sleep")
+
@mock.patch("airflow.providers.databricks.hooks.databricks.DatabricksHook.get_run_tasks")
+
@mock.patch("airflow.providers.databricks.hooks.databricks.DatabricksHook.a_get_run_state")
+ async def test_run_workflow_task_retry_in_flight_succeeds(
+ self, mock_get_run_state, mock_get_run_tasks, mock_sleep
+ ):
+ # First attempt fails while the parent run is still active; a later
attempt succeeds and
+ # the trigger must emit a success event rather than failing on the
first attempt.
+ mock_get_run_tasks.side_effect = [
+ [{"run_id": TASK_RUN_ID1, "task_key": TASK_RUN_ID1_KEY,
"start_time": 1}],
+ [
+ {"run_id": TASK_RUN_ID1, "task_key": TASK_RUN_ID1_KEY,
"start_time": 1},
+ {"run_id": TASK_RUN_ID2, "task_key": TASK_RUN_ID1_KEY,
"start_time": 2},
+ ],
+ ]
+ mock_get_run_state.side_effect = [
+ RunState(life_cycle_state=LIFE_CYCLE_STATE_TERMINATED,
state_message="", result_state="FAILED"),
+ RunState(life_cycle_state=LIFE_CYCLE_STATE_PENDING,
state_message="", result_state=""),
+ RunState(life_cycle_state=LIFE_CYCLE_STATE_TERMINATED,
state_message="", result_state="SUCCESS"),
+ ]
+
+ trigger = DatabricksExecutionTrigger(
+ run_id=TASK_RUN_ID1,
+ databricks_conn_id=DEFAULT_CONN_ID,
+ polling_period_seconds=POLLING_INTERVAL_SECONDS,
+ run_page_url=RUN_PAGE_URL,
+ workflow_run_id=RUN_ID,
+ databricks_task_key=TASK_RUN_ID1_KEY,
+ )
+
+ events = [event async for event in trigger.run()]
+ assert events == [
+ TriggerEvent(
+ {
+ "run_id": TASK_RUN_ID2,
+ "run_page_url": RUN_PAGE_URL,
+ "run_state": RunState(
+ life_cycle_state=LIFE_CYCLE_STATE_TERMINATED,
state_message="", result_state="SUCCESS"
+ ).to_json(),
+ "repair_run": False,
+ "errors": [],
+ }
+ )
+ ]
+ mock_sleep.assert_called_once_with(POLLING_INTERVAL_SECONDS)
+
+ @pytest.mark.asyncio
+
@mock.patch("airflow.providers.databricks.triggers.databricks.asyncio.sleep")
+
@mock.patch("airflow.providers.databricks.hooks.databricks.DatabricksHook.a_get_run_output")
+
@mock.patch("airflow.providers.databricks.hooks.databricks.DatabricksHook.a_get_run")
+
@mock.patch("airflow.providers.databricks.hooks.databricks.DatabricksHook.get_run_tasks")
+
@mock.patch("airflow.providers.databricks.hooks.databricks.DatabricksHook.a_get_run_state")
+ async def test_run_workflow_task_failed_attempt_waits_for_parent(
+ self, mock_get_run_state, mock_get_run_tasks, mock_get_run,
mock_get_run_output, mock_sleep
+ ):
+ mock_get_run_tasks.return_value = [
+ {"run_id": TASK_RUN_ID1, "task_key": TASK_RUN_ID1_KEY,
"start_time": 1}
+ ]
+ mock_get_run_state.side_effect = [
+ RunState(life_cycle_state=LIFE_CYCLE_STATE_TERMINATED,
state_message="", result_state="FAILED"),
+ RunState(life_cycle_state="RUNNING", state_message="",
result_state=""),
+ RunState(life_cycle_state=LIFE_CYCLE_STATE_TERMINATED,
state_message="", result_state="FAILED"),
+ RunState(life_cycle_state=LIFE_CYCLE_STATE_TERMINATED,
state_message="", result_state="FAILED"),
+ ]
+ mock_get_run.return_value = GET_RUN_RESPONSE_TERMINATED_WITH_FAILED
+ mock_get_run_output.return_value = GET_RUN_OUTPUT_RESPONSE
+
+ trigger = DatabricksExecutionTrigger(
+ run_id=TASK_RUN_ID1,
+ databricks_conn_id=DEFAULT_CONN_ID,
+ polling_period_seconds=POLLING_INTERVAL_SECONDS,
+ run_page_url=RUN_PAGE_URL,
+ workflow_run_id=RUN_ID,
+ databricks_task_key=TASK_RUN_ID1_KEY,
+ )
+
+ events = [event async for event in trigger.run()]
+ assert events == [
+ TriggerEvent(
+ {
+ "run_id": TASK_RUN_ID1,
+ "run_page_url": RUN_PAGE_URL,
+ "run_state": RunState(
+ life_cycle_state=LIFE_CYCLE_STATE_TERMINATED,
state_message="", result_state="FAILED"
+ ).to_json(),
+ "repair_run": False,
+ "errors": [
+ {"task_key": TASK_RUN_ID1_KEY, "run_id": TASK_RUN_ID1,
"error": ERROR_MESSAGE},
+ {"task_key": TASK_RUN_ID3_KEY, "run_id": TASK_RUN_ID3,
"error": ERROR_MESSAGE},
+ ],
+ }
+ )
+ ]
+ mock_sleep.assert_called_once_with(POLLING_INTERVAL_SECONDS)
+
+ @pytest.mark.asyncio
+
@mock.patch("airflow.providers.databricks.hooks.databricks.DatabricksHook.a_get_run_output")
+
@mock.patch("airflow.providers.databricks.hooks.databricks.DatabricksHook.a_get_run")
+
@mock.patch("airflow.providers.databricks.triggers.databricks.asyncio.sleep")
+
@mock.patch("airflow.providers.databricks.hooks.databricks.DatabricksHook.get_run_tasks")
+
@mock.patch("airflow.providers.databricks.hooks.databricks.DatabricksHook.a_get_run_state")
+ async def test_run_workflow_task_reports_failure_once_retries_exhausted(
+ self, mock_get_run_state, mock_get_run_tasks, mock_sleep,
mock_get_run, mock_get_run_output
+ ):
+ mock_get_run_tasks.return_value = [
+ {"run_id": TASK_RUN_ID1, "task_key": TASK_RUN_ID1_KEY,
"start_time": 1, "attempt_number": 1}
+ ]
+ mock_get_run_state.return_value = RunState(
+ life_cycle_state=LIFE_CYCLE_STATE_TERMINATED, state_message="",
result_state="FAILED"
+ )
+ mock_get_run.return_value = GET_RUN_RESPONSE_TERMINATED_WITH_FAILED
+ mock_get_run_output.return_value = GET_RUN_OUTPUT_RESPONSE
+
+ trigger = DatabricksExecutionTrigger(
+ run_id=TASK_RUN_ID1,
+ databricks_conn_id=DEFAULT_CONN_ID,
+ polling_period_seconds=POLLING_INTERVAL_SECONDS,
+ run_page_url=RUN_PAGE_URL,
+ workflow_run_id=RUN_ID,
+ databricks_task_key=TASK_RUN_ID1_KEY,
+ max_retries=1,
+ )
+
+ events = [event async for event in trigger.run()]
+ assert events == [
+ TriggerEvent(
+ {
+ "run_id": TASK_RUN_ID1,
+ "run_page_url": RUN_PAGE_URL,
+ "run_state": RunState(
+ life_cycle_state=LIFE_CYCLE_STATE_TERMINATED,
state_message="", result_state="FAILED"
+ ).to_json(),
+ "repair_run": False,
+ "errors": [
+ {"task_key": TASK_RUN_ID1_KEY, "run_id": TASK_RUN_ID1,
"error": ERROR_MESSAGE},
+ {"task_key": TASK_RUN_ID3_KEY, "run_id": TASK_RUN_ID3,
"error": ERROR_MESSAGE},
+ ],
+ }
+ )
+ ]
+ mock_sleep.assert_not_called()
+ mock_get_run_state.assert_called_once_with(TASK_RUN_ID1)
+
@pytest.mark.asyncio
@mock.patch("airflow.providers.databricks.hooks.databricks.DatabricksHook.cancel_run")
async def test_on_kill_cancels_run(self, mock_cancel_run):
await self.trigger.on_kill()
mock_cancel_run.assert_called_once_with(RUN_ID)
+ @pytest.mark.asyncio
+
@mock.patch("airflow.providers.databricks.hooks.databricks.DatabricksHook.get_run_tasks")
+
@mock.patch("airflow.providers.databricks.hooks.databricks.DatabricksHook.cancel_run")
+ async def test_on_kill_workflow_task_cancels_latest_attempt(self,
mock_cancel_run, mock_get_run_tasks):
+ # The original attempt may be terminal after a retry; on_kill must
cancel the latest attempt
+ # (under the same task_key) instead of the now-stale run_id the
trigger was created with.
+ mock_get_run_tasks.return_value = [
+ {"run_id": TASK_RUN_ID1, "task_key": TASK_RUN_ID1_KEY,
"start_time": 1},
+ {"run_id": TASK_RUN_ID2, "task_key": TASK_RUN_ID1_KEY,
"start_time": 2},
+ ]
+ trigger = DatabricksExecutionTrigger(
+ run_id=TASK_RUN_ID1,
+ databricks_conn_id=DEFAULT_CONN_ID,
+ polling_period_seconds=POLLING_INTERVAL_SECONDS,
+ run_page_url=RUN_PAGE_URL,
+ workflow_run_id=RUN_ID,
+ databricks_task_key=TASK_RUN_ID1_KEY,
+ )
+
+ await trigger.on_kill()
+ mock_get_run_tasks.assert_called_once_with(RUN_ID)
+ mock_cancel_run.assert_called_once_with(TASK_RUN_ID2)
+
class TestDatabricksSQLStatementExecutionTrigger:
@pytest.fixture(autouse=True)