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 c4e6ed369ed Add deferrable mode to Databricks SQL warehouse operators 
(#71752)
c4e6ed369ed is described below

commit c4e6ed369ed6c5fde5faa85e50581a01c71b474d
Author: deepinsight coder <[email protected]>
AuthorDate: Wed Aug 26 01:00:34 2026 -0700

    Add deferrable mode to Databricks SQL warehouse operators (#71752)
    
    * Add deferrable mode to Databricks SQL warehouse operators
    
    Start and stop waits from #70088 held a worker for the whole warehouse
    lifecycle. This follow-up lands the agreed additive triggerer path so
    those waits no longer occupy a worker slot.
    
    Co-authored-by: Cursor <[email protected]>
    
    * Honor Databricks warehouse deferrable timeout across triggerer restarts
    
    A duration recomputed inside the trigger started a fresh wait after every 
triggerer restart or HA rebalance, so a warehouse that never reached its target 
could stay deferred well past the configured timeout.
    
    * Fix Databricks warehouse trigger docs spellcheck
    
    Sphinx treats the HA term as a misspelling and fails the provider docs job.
    
    Co-authored-by: Cursor <[email protected]>
    
    ---------
    
    Co-authored-by: Cursor <[email protected]>
---
 docs/spelling_wordlist.txt                         |   1 +
 providers/databricks/docs/operators/warehouse.rst  |  15 +-
 .../providers/databricks/hooks/databricks.py       |  23 ++
 .../providers/databricks/operators/warehouse.py    |  65 +++++-
 .../providers/databricks/triggers/databricks.py    | 122 +++++++++-
 .../tests/unit/databricks/hooks/test_databricks.py |  30 +++
 .../unit/databricks/operators/test_warehouse.py    | 252 +++++++++++++++++++++
 .../unit/databricks/triggers/test_databricks.py    | 228 ++++++++++++++++++-
 8 files changed, 724 insertions(+), 12 deletions(-)

diff --git a/docs/spelling_wordlist.txt b/docs/spelling_wordlist.txt
index b60002cf177..a153988c461 100644
--- a/docs/spelling_wordlist.txt
+++ b/docs/spelling_wordlist.txt
@@ -1410,6 +1410,7 @@ ReadOnlyCredentials
 readthedocs
 Realtime
 realtime
+rebalance
 rebase
 Rebasing
 Recency
diff --git a/providers/databricks/docs/operators/warehouse.rst 
b/providers/databricks/docs/operators/warehouse.rst
index 1467bbb7f45..5f6f1157fba 100644
--- a/providers/databricks/docs/operators/warehouse.rst
+++ b/providers/databricks/docs/operators/warehouse.rst
@@ -30,10 +30,17 @@ Both operators require the warehouse ID and use the 
:ref:`Databricks connection
 
 By default, each operator waits for the requested state: ``RUNNING`` when 
starting and ``STOPPED``
 when stopping. Use ``polling_period_seconds`` to control the polling interval, 
``timeout`` to limit
-the wait, or ``wait_for_termination=False`` to return after requesting the 
transition. Repeated task
-attempts are safe: an already running warehouse is not started again, and an 
already stopped warehouse
-is not stopped again. If a start is requested while a warehouse is stopping, 
any transition rejection
-from Databricks is propagated to the task.
+the wait, or ``wait_for_termination=False`` to return after requesting the 
transition. Set
+``deferrable=True`` (or enable ``[operators] default_deferrable``) to wait on 
the triggerer instead
+of holding a worker slot. In deferrable mode the operator records an absolute 
``end_time`` from
+``timeout`` when it defers, so a triggerer restart does not reset the wait. 
``wait_for_termination=False``
+still returns immediately and does not defer. Clearing a deferred warehouse 
wait does not start or
+stop the warehouse: Databricks has no cancel API for these transitions, and a 
cancelled start must
+not stop a warehouse the Dag still needs.
+
+Repeated task attempts are safe: an already running warehouse is not started 
again, and an already
+stopped warehouse is not stopped again. If a start is requested while a 
warehouse is stopping, any
+transition rejection from Databricks is propagated to the task.
 
 Start a SQL warehouse
 ---------------------
diff --git 
a/providers/databricks/src/airflow/providers/databricks/hooks/databricks.py 
b/providers/databricks/src/airflow/providers/databricks/hooks/databricks.py
index 4a0f6454983..2223b5a5a4b 100644
--- a/providers/databricks/src/airflow/providers/databricks/hooks/databricks.py
+++ b/providers/databricks/src/airflow/providers/databricks/hooks/databricks.py
@@ -307,6 +307,13 @@ class WarehouseState:
     def __repr__(self) -> str:
         return str(self.__dict__)
 
+    def to_json(self) -> str:
+        return json.dumps(self.__dict__)
+
+    @classmethod
+    def from_json(cls, data: str) -> WarehouseState:
+        return WarehouseState(**json.loads(data))
+
 
 class DatabricksHook(BaseDatabricksHook):
     """
@@ -791,6 +798,14 @@ class DatabricksHook(BaseDatabricksHook):
         """
         return self._do_api_call(("GET", 
f"{SQL_WAREHOUSES_ENDPOINT}/{warehouse_id}"))
 
+    async def a_get_warehouse(self, warehouse_id: str) -> dict[str, Any]:
+        """
+        Async version of `get_warehouse`.
+
+        :param warehouse_id: ID of the SQL warehouse.
+        """
+        return await self._a_do_api_call(("GET", 
f"{SQL_WAREHOUSES_ENDPOINT}/{warehouse_id}"))
+
     def get_warehouse_state(self, warehouse_id: str) -> WarehouseState:
         """
         Retrieve the state of a Databricks SQL warehouse.
@@ -799,6 +814,14 @@ class DatabricksHook(BaseDatabricksHook):
         """
         return WarehouseState(self.get_warehouse(warehouse_id)["state"])
 
+    async def a_get_warehouse_state(self, warehouse_id: str) -> WarehouseState:
+        """
+        Async version of `get_warehouse_state`.
+
+        :param warehouse_id: ID of the SQL warehouse.
+        """
+        return WarehouseState((await 
self.a_get_warehouse(warehouse_id))["state"])
+
     def start_warehouse(self, warehouse_id: str) -> None:
         """
         Start a Databricks SQL warehouse.
diff --git 
a/providers/databricks/src/airflow/providers/databricks/operators/warehouse.py 
b/providers/databricks/src/airflow/providers/databricks/operators/warehouse.py
index e2d0ce9a69b..374e960e184 100644
--- 
a/providers/databricks/src/airflow/providers/databricks/operators/warehouse.py
+++ 
b/providers/databricks/src/airflow/providers/databricks/operators/warehouse.py
@@ -23,9 +23,11 @@ from collections.abc import Sequence
 from functools import cached_property
 from typing import TYPE_CHECKING, Any
 
-from airflow.providers.common.compat.sdk import BaseOperator
+from airflow.providers.common.compat.sdk import BaseOperator, conf
 from airflow.providers.databricks.exceptions import DatabricksWarehouseError
-from airflow.providers.databricks.hooks.databricks import DatabricksHook
+from airflow.providers.databricks.hooks.databricks import DatabricksHook, 
WarehouseState
+from airflow.providers.databricks.triggers.databricks import 
DatabricksWarehouseStateTrigger
+from airflow.providers.databricks.utils.retry import 
validate_deferrable_databricks_retry_args
 
 if TYPE_CHECKING:
     from airflow.providers.common.compat.sdk import Context
@@ -49,6 +51,7 @@ class _DatabricksWarehouseBaseOperator(BaseOperator):
         databricks_retry_limit: int = 3,
         databricks_retry_delay: int = 1,
         databricks_retry_args: dict[Any, Any] | None = None,
+        deferrable: bool = conf.getboolean("operators", "default_deferrable", 
fallback=False),
         **kwargs,
     ) -> None:
         super().__init__(**kwargs)
@@ -60,6 +63,11 @@ class _DatabricksWarehouseBaseOperator(BaseOperator):
         self.databricks_retry_limit = databricks_retry_limit
         self.databricks_retry_delay = databricks_retry_delay
         self.databricks_retry_args = databricks_retry_args
+        self.deferrable = deferrable
+        if self.deferrable:
+            validate_deferrable_databricks_retry_args(
+                self.databricks_retry_args, owner=self.__class__.__name__
+            )
 
     def _validate_warehouse_id(self) -> None:
         if not self.warehouse_id:
@@ -103,6 +111,51 @@ class _DatabricksWarehouseBaseOperator(BaseOperator):
             f"within {self.timeout}s; last state: {last_state}."
         )
 
+    def _wait_or_defer(self, target: str) -> None:
+        if not self.wait_for_termination:
+            return
+        if self.deferrable:
+            self.defer(
+                trigger=DatabricksWarehouseStateTrigger(
+                    warehouse_id=self.warehouse_id,
+                    target_state=target,
+                    databricks_conn_id=self.databricks_conn_id,
+                    end_time=time.time() + self.timeout,
+                    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.__class__.__name__,
+                ),
+                method_name="execute_complete",
+            )
+            return
+        self._wait_for_state(target)
+
+    def execute_complete(self, context: Context, event: dict[str, Any] | None 
= None) -> None:
+        if not event:
+            raise DatabricksWarehouseError("Databricks SQL warehouse trigger 
completed without an event.")
+        warehouse_id = event.get("warehouse_id", self.warehouse_id)
+        target = event.get("target_state", "unknown")
+        last_state = event.get("last_state", "unknown")
+        status = event.get("status")
+        if status == "success":
+            self.log.info("Databricks SQL warehouse %s reached %s.", 
warehouse_id, target)
+            return
+        if status == "deleted":
+            state = WarehouseState.from_json(event["state"])
+            raise DatabricksWarehouseError(
+                f"Databricks SQL warehouse {warehouse_id} entered 
{state.state} while waiting for {target}."
+            )
+        if status == "timeout":
+            raise DatabricksWarehouseError(
+                f"Databricks SQL warehouse {warehouse_id} did not reach 
{target} "
+                f"within {self.timeout}s; last state: {last_state}."
+            )
+        raise DatabricksWarehouseError(
+            f"Databricks SQL warehouse {warehouse_id} trigger finished with 
unexpected status {status!r}."
+        )
+
 
 class DatabricksStartWarehouseOperator(_DatabricksWarehouseBaseOperator):
     """
@@ -116,6 +169,7 @@ class 
DatabricksStartWarehouseOperator(_DatabricksWarehouseBaseOperator):
     :param databricks_retry_limit: Number of times to retry unavailable 
Databricks requests.
     :param databricks_retry_delay: Number of seconds between Databricks 
request retries.
     :param databricks_retry_args: Additional arguments for 
``tenacity.Retrying``.
+    :param deferrable: Run operator in the deferrable mode. Defaults to 
``[operators] default_deferrable``.
     """
 
     def execute(self, context: Context) -> None:
@@ -126,8 +180,7 @@ class 
DatabricksStartWarehouseOperator(_DatabricksWarehouseBaseOperator):
             return
         if state.state != "STARTING":
             self._hook.start_warehouse(self.warehouse_id)
-        if self.wait_for_termination:
-            self._wait_for_state("RUNNING")
+        self._wait_or_defer("RUNNING")
 
 
 class DatabricksStopWarehouseOperator(_DatabricksWarehouseBaseOperator):
@@ -142,6 +195,7 @@ class 
DatabricksStopWarehouseOperator(_DatabricksWarehouseBaseOperator):
     :param databricks_retry_limit: Number of times to retry unavailable 
Databricks requests.
     :param databricks_retry_delay: Number of seconds between Databricks 
request retries.
     :param databricks_retry_args: Additional arguments for 
``tenacity.Retrying``.
+    :param deferrable: Run operator in the deferrable mode. Defaults to 
``[operators] default_deferrable``.
     """
 
     def execute(self, context: Context) -> None:
@@ -152,5 +206,4 @@ class 
DatabricksStopWarehouseOperator(_DatabricksWarehouseBaseOperator):
             return
         if state.state != "STOPPING":
             self._hook.stop_warehouse(self.warehouse_id)
-        if self.wait_for_termination:
-            self._wait_for_state("STOPPED")
+        self._wait_or_defer("STOPPED")
diff --git 
a/providers/databricks/src/airflow/providers/databricks/triggers/databricks.py 
b/providers/databricks/src/airflow/providers/databricks/triggers/databricks.py
index 9706d2e4adf..3a3bffc2c41 100644
--- 
a/providers/databricks/src/airflow/providers/databricks/triggers/databricks.py
+++ 
b/providers/databricks/src/airflow/providers/databricks/triggers/databricks.py
@@ -21,7 +21,7 @@ import asyncio
 import time
 from typing import Any
 
-from airflow.providers.databricks.hooks.databricks import DatabricksHook
+from airflow.providers.databricks.hooks.databricks import DatabricksHook, 
WarehouseState
 from airflow.providers.databricks.utils.databricks import 
extract_failed_task_errors_async
 from airflow.providers.databricks.utils.retry import 
validate_deferrable_databricks_retry_args
 from airflow.triggers.base import BaseTrigger, TriggerEvent
@@ -336,3 +336,123 @@ class DatabricksSQLStatementExecutionTrigger(BaseTrigger):
                 }
             )
             return
+
+
+class DatabricksWarehouseStateTrigger(BaseTrigger):
+    """
+    Poll a Databricks SQL warehouse until it reaches a target lifecycle state.
+
+    Databricks has no cancel API for warehouse start or stop, so this trigger 
does
+    not override ``on_kill``. Clearing a deferred wait leaves warehouse state 
unchanged.
+
+    :param warehouse_id: ID of the Databricks SQL warehouse.
+    :param target_state: Lifecycle state to wait for (``RUNNING`` or 
``STOPPED``).
+    :param databricks_conn_id: Reference to the :ref:`Databricks connection 
<howto/connection:databricks>`.
+    :param end_time: Absolute Unix timestamp when the wait must stop. The 
operator
+        sets this from ``timeout`` before deferring so a triggerer restart or 
HA
+        rebalance does not reset the deadline.
+    :param polling_period_seconds: Controls the rate of the poll for the 
warehouse state.
+        By default, the trigger will poll every 30 seconds.
+    :param retry_limit: The number of times to retry the connection in case of 
service outages.
+    :param retry_delay: Minimum wait in seconds between retryable attempts 
when using the
+        default retry strategy. The wait uses exponential backoff (doubling 
after each
+        failure, capped at ``2 ** retry_limit`` seconds). May be a floating 
point number.
+    :param retry_args: An optional dictionary with arguments passed to 
``tenacity.Retrying`` class.
+    :param caller: The name of the operator that is calling the hook.
+    """
+
+    def __init__(
+        self,
+        warehouse_id: str,
+        target_state: str,
+        databricks_conn_id: str,
+        end_time: float,
+        polling_period_seconds: int = 30,
+        retry_limit: int = 3,
+        retry_delay: int = 10,
+        retry_args: dict[Any, Any] | None = None,
+        caller: str = "DatabricksWarehouseStateTrigger",
+    ) -> None:
+        super().__init__()
+        # Trigger kwargs cross Airflow's serialization boundary, so fail 
before storing invalid
+        # trigger state or surfacing a generic serializer error without 
Databricks-specific guidance.
+        validate_deferrable_databricks_retry_args(retry_args, owner=caller)
+        self.warehouse_id = warehouse_id
+        self.target_state = target_state
+        self.databricks_conn_id = databricks_conn_id
+        self.end_time = end_time
+        self.polling_period_seconds = polling_period_seconds
+        self.retry_limit = retry_limit
+        self.retry_delay = retry_delay
+        self.retry_args = retry_args
+        self.caller = caller
+        self.hook = DatabricksHook(
+            databricks_conn_id,
+            retry_limit=self.retry_limit,
+            retry_delay=self.retry_delay,
+            retry_args=retry_args,
+            caller=caller,
+        )
+
+    def serialize(self) -> tuple[str, dict[str, Any]]:
+        return (
+            
"airflow.providers.databricks.triggers.databricks.DatabricksWarehouseStateTrigger",
+            {
+                "warehouse_id": self.warehouse_id,
+                "target_state": self.target_state,
+                "databricks_conn_id": self.databricks_conn_id,
+                "end_time": self.end_time,
+                "polling_period_seconds": self.polling_period_seconds,
+                "retry_limit": self.retry_limit,
+                "retry_delay": self.retry_delay,
+                "retry_args": self.retry_args,
+                "caller": self.caller,
+            },
+        )
+
+    def _build_trigger_event(
+        self, *, status: str, last_state: str, state: WarehouseState | None = 
None
+    ) -> TriggerEvent:
+        payload: dict[str, Any] = {
+            "status": status,
+            "warehouse_id": self.warehouse_id,
+            "target_state": self.target_state,
+            "last_state": last_state,
+        }
+        if state is not None:
+            # Same typed JSON payload as RunState / SQLStatementState triggers.
+            payload["state"] = state.to_json()
+        return TriggerEvent(payload)
+
+    async def run(self):
+        async with self.hook:
+            last_state = "unknown"
+            last_warehouse_state: WarehouseState | None = None
+            while time.time() < self.end_time:
+                warehouse_state = await 
self.hook.a_get_warehouse_state(self.warehouse_id)
+                last_warehouse_state = warehouse_state
+                last_state = warehouse_state.state
+                now = time.time()
+                if warehouse_state.state == self.target_state:
+                    yield self._build_trigger_event(
+                        status="success", last_state=last_state, 
state=warehouse_state
+                    )
+                    return
+                if warehouse_state.is_deleted:
+                    yield self._build_trigger_event(
+                        status="deleted", last_state=last_state, 
state=warehouse_state
+                    )
+                    return
+                if now >= self.end_time:
+                    break
+                self.log.info(
+                    "Databricks SQL warehouse %s is %s; waiting for %s.",
+                    self.warehouse_id,
+                    warehouse_state.state,
+                    self.target_state,
+                )
+                await asyncio.sleep(min(self.polling_period_seconds, 
self.end_time - now))
+            yield self._build_trigger_event(
+                status="timeout", last_state=last_state, 
state=last_warehouse_state
+            )
+            return
diff --git 
a/providers/databricks/tests/unit/databricks/hooks/test_databricks.py 
b/providers/databricks/tests/unit/databricks/hooks/test_databricks.py
index be87a82483c..c1959cd2359 100644
--- a/providers/databricks/tests/unit/databricks/hooks/test_databricks.py
+++ b/providers/databricks/tests/unit/databricks/hooks/test_databricks.py
@@ -1564,6 +1564,17 @@ class TestWarehouseLifecycle:
         assert result == {"id": "wh-1", "state": "RUNNING"}
         mock_do_api_call.assert_called_once_with(hook, ("GET", 
"2.0/sql/warehouses/wh-1"))
 
+    @pytest.mark.asyncio
+    @mock.patch.object(DatabricksHook, "_a_do_api_call", autospec=True)
+    async def test_a_get_warehouse_calls_correct_endpoint(self, 
mock_a_do_api_call):
+        mock_a_do_api_call.return_value = {"id": "wh-1", "state": "RUNNING"}
+        hook = DatabricksHook()
+
+        result = await hook.a_get_warehouse("wh-1")
+
+        assert result == {"id": "wh-1", "state": "RUNNING"}
+        mock_a_do_api_call.assert_called_once_with(hook, ("GET", 
"2.0/sql/warehouses/wh-1"))
+
     @mock.patch.object(DatabricksHook, "_do_api_call", autospec=True)
     def test_get_warehouse_state_wraps_state(self, mock_do_api_call):
         mock_do_api_call.return_value = {"state": "RUNNING"}
@@ -1614,6 +1625,25 @@ class TestWarehouseLifecycle:
         with pytest.raises(ValueError, match="Unexpected warehouse state: 
FOO"):
             WarehouseState("FOO")
 
+    def test_warehouse_state_json_round_trip(self):
+        state = WarehouseState("STOPPING")
+
+        restored = WarehouseState.from_json(state.to_json())
+
+        assert restored == state
+        assert json.loads(state.to_json()) == {"state": "STOPPING"}
+
+    @pytest.mark.asyncio
+    @mock.patch.object(DatabricksHook, "a_get_warehouse", autospec=True)
+    async def test_a_get_warehouse_state_wraps_state(self, 
mock_a_get_warehouse):
+        mock_a_get_warehouse.return_value = {"state": "STARTING"}
+        hook = DatabricksHook()
+
+        state = await hook.a_get_warehouse_state("wh-1")
+
+        assert state == WarehouseState("STARTING")
+        mock_a_get_warehouse.assert_called_once_with(hook, "wh-1")
+
 
 class TestRunState:
     def test_is_terminal_true(self):
diff --git 
a/providers/databricks/tests/unit/databricks/operators/test_warehouse.py 
b/providers/databricks/tests/unit/databricks/operators/test_warehouse.py
index aa99e4d54cc..c0aa0de03b1 100644
--- a/providers/databricks/tests/unit/databricks/operators/test_warehouse.py
+++ b/providers/databricks/tests/unit/databricks/operators/test_warehouse.py
@@ -16,16 +16,21 @@
 # under the License.
 from __future__ import annotations
 
+import time
 from unittest import mock
 
 import pytest
+import time_machine
+from tenacity import stop_after_attempt, wait_incrementing
 
+from airflow.providers.common.compat.sdk import TaskDeferred
 from airflow.providers.databricks.exceptions import DatabricksWarehouseError
 from airflow.providers.databricks.hooks.databricks import DatabricksHook, 
WarehouseState
 from airflow.providers.databricks.operators.warehouse import (
     DatabricksStartWarehouseOperator,
     DatabricksStopWarehouseOperator,
 )
+from airflow.providers.databricks.triggers.databricks import 
DatabricksWarehouseStateTrigger
 
 TASK_ID = "warehouse-lifecycle"
 WAREHOUSE_ID = "wh-1"
@@ -404,3 +409,250 @@ class TestDatabricksWarehouseOperatorBase:
             retry_args=retry_args,
             caller="DatabricksStartWarehouseOperator",
         )
+
+
+INVALID_RETRY_ARGS_PATTERN = (
+    "does not support non-serializable retry_args/databricks_retry_args when 
deferrable=True"
+)
+UNSUPPORTED_RETRY_ARGS = [
+    pytest.param({"wait": wait_incrementing(start=1, increment=1, max=3)}, 
id="wait_incrementing"),
+    pytest.param({"stop": stop_after_attempt(3)}, id="stop_after_attempt"),
+]
+
+
+class TestDatabricksWarehouseOperatorDeferrable:
+    @pytest.mark.parametrize(
+        ("operator_class", "initial_state", "transition_method", 
"target_state"),
+        [
+            (DatabricksStartWarehouseOperator, "STOPPED", "start_warehouse", 
"RUNNING"),
+            (DatabricksStopWarehouseOperator, "RUNNING", "stop_warehouse", 
"STOPPED"),
+        ],
+    )
+    @mock.patch.object(DatabricksStartWarehouseOperator, "_hook", 
new_callable=mock.PropertyMock)
+    @mock.patch.object(DatabricksStopWarehouseOperator, "_hook", 
new_callable=mock.PropertyMock)
+    @time_machine.travel("2026-08-19 12:00:00", tick=False)
+    def test_execute_defers_until_target_state(
+        self,
+        mock_stop_hook_property,
+        mock_start_hook_property,
+        operator_class,
+        initial_state,
+        transition_method,
+        target_state,
+    ):
+        hook = mock.MagicMock(spec=DatabricksHook)
+        mock_start_hook_property.return_value = hook
+        mock_stop_hook_property.return_value = hook
+        hook.get_warehouse_state.return_value = WarehouseState(initial_state)
+        operator = operator_class(task_id=TASK_ID, warehouse_id=WAREHOUSE_ID, 
deferrable=True)
+
+        with pytest.raises(TaskDeferred) as exc:
+            operator.execute(None)
+
+        getattr(hook, transition_method).assert_called_once_with(WAREHOUSE_ID)
+        assert isinstance(exc.value.trigger, DatabricksWarehouseStateTrigger)
+        assert exc.value.method_name == "execute_complete"
+        assert exc.value.trigger.warehouse_id == WAREHOUSE_ID
+        assert exc.value.trigger.target_state == target_state
+        assert exc.value.trigger.end_time == pytest.approx(time.time() + 3600)
+        assert exc.value.trigger.caller == operator_class.__name__
+
+    @pytest.mark.parametrize(
+        "first_polled_state",
+        ["STARTING", "STOPPED"],
+        ids=["transitioning", "stale-stopped"],
+    )
+    @mock.patch.object(DatabricksStartWarehouseOperator, "_hook", 
new_callable=mock.PropertyMock)
+    def test_start_defers_after_stale_or_starting_precheck(self, 
mock_hook_property, first_polled_state):
+        hook = mock.MagicMock(spec=DatabricksHook)
+        mock_hook_property.return_value = hook
+        hook.get_warehouse_state.return_value = 
WarehouseState(first_polled_state)
+        operator = DatabricksStartWarehouseOperator(
+            task_id=TASK_ID, warehouse_id=WAREHOUSE_ID, deferrable=True
+        )
+
+        with pytest.raises(TaskDeferred) as exc:
+            operator.execute(None)
+
+        if first_polled_state == "STARTING":
+            hook.start_warehouse.assert_not_called()
+        else:
+            hook.start_warehouse.assert_called_once_with(WAREHOUSE_ID)
+        assert exc.value.trigger.target_state == "RUNNING"
+
+    @pytest.mark.parametrize(
+        ("operator_class", "target_state"),
+        [
+            (DatabricksStartWarehouseOperator, WarehouseState("RUNNING")),
+            (DatabricksStopWarehouseOperator, WarehouseState("STOPPED")),
+        ],
+    )
+    @mock.patch.object(DatabricksStartWarehouseOperator, "_hook", 
new_callable=mock.PropertyMock)
+    @mock.patch.object(DatabricksStopWarehouseOperator, "_hook", 
new_callable=mock.PropertyMock)
+    def test_already_at_target_does_not_defer(
+        self, mock_stop_hook_property, mock_start_hook_property, 
operator_class, target_state
+    ):
+        hook = mock.MagicMock(spec=DatabricksHook)
+        mock_start_hook_property.return_value = hook
+        mock_stop_hook_property.return_value = hook
+        hook.get_warehouse_state.return_value = target_state
+        operator = operator_class(task_id=TASK_ID, warehouse_id=WAREHOUSE_ID, 
deferrable=True)
+
+        operator.execute(None)
+
+        hook.start_warehouse.assert_not_called()
+        hook.stop_warehouse.assert_not_called()
+
+    @pytest.mark.parametrize(
+        ("operator_class", "initial_state", "transition_method"),
+        [
+            (DatabricksStartWarehouseOperator, "STOPPED", "start_warehouse"),
+            (DatabricksStopWarehouseOperator, "RUNNING", "stop_warehouse"),
+        ],
+    )
+    @mock.patch.object(DatabricksStartWarehouseOperator, "_hook", 
new_callable=mock.PropertyMock)
+    @mock.patch.object(DatabricksStopWarehouseOperator, "_hook", 
new_callable=mock.PropertyMock)
+    def test_wait_for_termination_false_does_not_defer(
+        self,
+        mock_stop_hook_property,
+        mock_start_hook_property,
+        operator_class,
+        initial_state,
+        transition_method,
+    ):
+        hook = mock.MagicMock(spec=DatabricksHook)
+        mock_start_hook_property.return_value = hook
+        mock_stop_hook_property.return_value = hook
+        hook.get_warehouse_state.return_value = WarehouseState(initial_state)
+        operator = operator_class(
+            task_id=TASK_ID,
+            warehouse_id=WAREHOUSE_ID,
+            deferrable=True,
+            wait_for_termination=False,
+        )
+
+        operator.execute(None)
+
+        getattr(hook, transition_method).assert_called_once_with(WAREHOUSE_ID)
+
+    @pytest.mark.parametrize(
+        ("operator_class", "initial_state", "target_state", 
"transition_method"),
+        [
+            (DatabricksStartWarehouseOperator, "STARTING", "RUNNING", 
"start_warehouse"),
+            (DatabricksStopWarehouseOperator, "STOPPING", "STOPPED", 
"stop_warehouse"),
+        ],
+    )
+    @mock.patch.object(DatabricksStartWarehouseOperator, "_hook", 
new_callable=mock.PropertyMock)
+    @mock.patch.object(DatabricksStopWarehouseOperator, "_hook", 
new_callable=mock.PropertyMock)
+    def test_in_progress_transition_defers_without_duplicate_request(
+        self,
+        mock_stop_hook_property,
+        mock_start_hook_property,
+        operator_class,
+        initial_state,
+        target_state,
+        transition_method,
+    ):
+        hook = mock.MagicMock(spec=DatabricksHook)
+        mock_start_hook_property.return_value = hook
+        mock_stop_hook_property.return_value = hook
+        hook.get_warehouse_state.return_value = WarehouseState(initial_state)
+        operator = operator_class(task_id=TASK_ID, warehouse_id=WAREHOUSE_ID, 
deferrable=True)
+
+        with pytest.raises(TaskDeferred) as exc:
+            operator.execute(None)
+
+        getattr(hook, transition_method).assert_not_called()
+        assert exc.value.trigger.target_state == target_state
+
+    @mock.patch.object(DatabricksStartWarehouseOperator, "_hook", 
new_callable=mock.PropertyMock)
+    def test_start_while_stopping_defers_after_start_request(self, 
mock_hook_property):
+        hook = mock.MagicMock(spec=DatabricksHook)
+        mock_hook_property.return_value = hook
+        hook.get_warehouse_state.return_value = WarehouseState("STOPPING")
+        operator = DatabricksStartWarehouseOperator(
+            task_id=TASK_ID, warehouse_id=WAREHOUSE_ID, deferrable=True
+        )
+
+        with pytest.raises(TaskDeferred) as exc:
+            operator.execute(None)
+
+        hook.start_warehouse.assert_called_once_with(WAREHOUSE_ID)
+        assert exc.value.trigger.target_state == "RUNNING"
+
+    @pytest.mark.parametrize("retry_args", UNSUPPORTED_RETRY_ARGS)
+    @pytest.mark.parametrize(
+        "operator_class",
+        [DatabricksStartWarehouseOperator, DatabricksStopWarehouseOperator],
+    )
+    def test_deferrable_rejects_non_serializable_retry_args(self, 
operator_class, retry_args):
+        with pytest.raises(ValueError, match=INVALID_RETRY_ARGS_PATTERN):
+            operator_class(
+                task_id=TASK_ID,
+                warehouse_id=WAREHOUSE_ID,
+                deferrable=True,
+                databricks_retry_args=retry_args,
+            )
+
+    def test_execute_complete_success(self):
+        operator = DatabricksStartWarehouseOperator(task_id=TASK_ID, 
warehouse_id=WAREHOUSE_ID)
+        event = {
+            "status": "success",
+            "warehouse_id": WAREHOUSE_ID,
+            "target_state": "RUNNING",
+            "last_state": "RUNNING",
+            "state": WarehouseState("RUNNING").to_json(),
+        }
+
+        assert operator.execute_complete(None, event) is None
+
+    def test_execute_complete_deleted_raises(self):
+        operator = DatabricksStartWarehouseOperator(task_id=TASK_ID, 
warehouse_id=WAREHOUSE_ID)
+        event = {
+            "status": "deleted",
+            "warehouse_id": WAREHOUSE_ID,
+            "target_state": "RUNNING",
+            "last_state": "DELETING",
+            "state": WarehouseState("DELETING").to_json(),
+        }
+
+        with pytest.raises(
+            DatabricksWarehouseError,
+            match="entered DELETING while waiting for RUNNING",
+        ):
+            operator.execute_complete(None, event)
+
+    def test_execute_complete_timeout_raises(self):
+        operator = DatabricksStopWarehouseOperator(task_id=TASK_ID, 
warehouse_id=WAREHOUSE_ID, timeout=12)
+        event = {
+            "status": "timeout",
+            "warehouse_id": WAREHOUSE_ID,
+            "target_state": "STOPPED",
+            "last_state": "STOPPING",
+        }
+
+        with pytest.raises(
+            DatabricksWarehouseError,
+            match="did not reach STOPPED within 12s; last state: STOPPING",
+        ):
+            operator.execute_complete(None, event)
+
+    def test_execute_complete_missing_event_raises(self):
+        operator = DatabricksStartWarehouseOperator(task_id=TASK_ID, 
warehouse_id=WAREHOUSE_ID)
+
+        with pytest.raises(DatabricksWarehouseError, match="completed without 
an event"):
+            operator.execute_complete(None, None)
+
+    def test_execute_complete_unexpected_status_raises(self):
+        operator = DatabricksStartWarehouseOperator(task_id=TASK_ID, 
warehouse_id=WAREHOUSE_ID)
+
+        with pytest.raises(DatabricksWarehouseError, match="unexpected status 
'boom'"):
+            operator.execute_complete(
+                None,
+                {
+                    "status": "boom",
+                    "warehouse_id": WAREHOUSE_ID,
+                    "target_state": "RUNNING",
+                    "last_state": "STARTING",
+                },
+            )
diff --git 
a/providers/databricks/tests/unit/databricks/triggers/test_databricks.py 
b/providers/databricks/tests/unit/databricks/triggers/test_databricks.py
index ad204297dc8..d8495a27aee 100644
--- a/providers/databricks/tests/unit/databricks/triggers/test_databricks.py
+++ b/providers/databricks/tests/unit/databricks/triggers/test_databricks.py
@@ -23,10 +23,11 @@ import pytest
 from tenacity import stop_after_attempt, wait_incrementing
 
 from airflow.models import Connection
-from airflow.providers.databricks.hooks.databricks import RunState, 
SQLStatementState
+from airflow.providers.databricks.hooks.databricks import RunState, 
SQLStatementState, WarehouseState
 from airflow.providers.databricks.triggers.databricks import (
     DatabricksExecutionTrigger,
     DatabricksSQLStatementExecutionTrigger,
+    DatabricksWarehouseStateTrigger,
 )
 from airflow.triggers.base import TriggerEvent
 
@@ -145,6 +146,16 @@ TRIGGER_INIT_CASES = [
         },
         id="sql_statement_trigger",
     ),
+    pytest.param(
+        DatabricksWarehouseStateTrigger,
+        {
+            "warehouse_id": "wh-1",
+            "target_state": "RUNNING",
+            "databricks_conn_id": DEFAULT_CONN_ID,
+            "end_time": 1234567890.0,
+        },
+        id="warehouse_state_trigger",
+    ),
 ]
 
 
@@ -633,3 +644,218 @@ class TestDatabricksSQLStatementExecutionTrigger:
     async def test_on_kill_cancels_statement(self, mock_cancel_sql_statement):
         await self.trigger.on_kill()
         mock_cancel_sql_statement.assert_called_once_with(STATEMENT_ID)
+
+
+WAREHOUSE_ID = "wh-1"
+WAREHOUSE_END_TIME = 9999999999.0
+WAREHOUSE_TIMEOUT_SECONDS = 30.0
+WAREHOUSE_CALLER = "DatabricksStartWarehouseOperator"
+
+
+class TestDatabricksWarehouseStateTrigger:
+    @pytest.fixture(autouse=True)
+    def setup_connections(self, create_connection_without_db):
+        create_connection_without_db(
+            Connection(
+                conn_id=DEFAULT_CONN_ID,
+                conn_type="databricks",
+                host=HOST,
+                login=LOGIN,
+                password=PASSWORD,
+                extra=None,
+            )
+        )
+        self.trigger = DatabricksWarehouseStateTrigger(
+            warehouse_id=WAREHOUSE_ID,
+            target_state="RUNNING",
+            databricks_conn_id=DEFAULT_CONN_ID,
+            polling_period_seconds=POLLING_INTERVAL_SECONDS,
+            end_time=WAREHOUSE_END_TIME,
+        )
+
+    def test_serialize(self):
+        assert self.trigger.serialize() == (
+            
"airflow.providers.databricks.triggers.databricks.DatabricksWarehouseStateTrigger",
+            {
+                "warehouse_id": WAREHOUSE_ID,
+                "target_state": "RUNNING",
+                "databricks_conn_id": DEFAULT_CONN_ID,
+                "end_time": WAREHOUSE_END_TIME,
+                "polling_period_seconds": POLLING_INTERVAL_SECONDS,
+                "retry_delay": 10,
+                "retry_limit": 3,
+                "retry_args": None,
+                "caller": "DatabricksWarehouseStateTrigger",
+            },
+        )
+
+    def test_serialize_round_trip_preserves_end_time(self):
+        trigger = DatabricksWarehouseStateTrigger(
+            warehouse_id=WAREHOUSE_ID,
+            target_state="STOPPED",
+            databricks_conn_id=DEFAULT_CONN_ID,
+            end_time=WAREHOUSE_END_TIME,
+            caller=WAREHOUSE_CALLER,
+        )
+        _, kwargs = trigger.serialize()
+        restored = DatabricksWarehouseStateTrigger(**kwargs)
+        assert restored.caller == WAREHOUSE_CALLER
+        assert restored.target_state == "STOPPED"
+        assert restored.end_time == WAREHOUSE_END_TIME
+
+    @pytest.mark.asyncio
+    
@mock.patch("airflow.providers.databricks.hooks.databricks.DatabricksHook.a_get_warehouse_state")
+    async def test_run_return_success(self, mock_a_get_warehouse_state):
+        mock_a_get_warehouse_state.return_value = WarehouseState("RUNNING")
+
+        events = [event async for event in self.trigger.run()]
+
+        assert events == [
+            TriggerEvent(
+                {
+                    "status": "success",
+                    "warehouse_id": WAREHOUSE_ID,
+                    "target_state": "RUNNING",
+                    "last_state": "RUNNING",
+                    "state": WarehouseState("RUNNING").to_json(),
+                }
+            )
+        ]
+
+    @pytest.mark.asyncio
+    
@mock.patch("airflow.providers.databricks.triggers.databricks.asyncio.sleep")
+    
@mock.patch("airflow.providers.databricks.hooks.databricks.DatabricksHook.a_get_warehouse_state")
+    async def test_run_keeps_polling_stale_stopped(self, 
mock_a_get_warehouse_state, mock_sleep):
+        mock_a_get_warehouse_state.side_effect = [
+            WarehouseState("STOPPED"),
+            WarehouseState("RUNNING"),
+        ]
+
+        events = [event async for event in self.trigger.run()]
+
+        assert events == [
+            TriggerEvent(
+                {
+                    "status": "success",
+                    "warehouse_id": WAREHOUSE_ID,
+                    "target_state": "RUNNING",
+                    "last_state": "RUNNING",
+                    "state": WarehouseState("RUNNING").to_json(),
+                }
+            )
+        ]
+        mock_sleep.assert_called_once_with(POLLING_INTERVAL_SECONDS)
+
+    @pytest.mark.asyncio
+    
@mock.patch("airflow.providers.databricks.hooks.databricks.DatabricksHook.a_get_warehouse_state")
+    async def test_run_return_deleted(self, mock_a_get_warehouse_state):
+        mock_a_get_warehouse_state.return_value = WarehouseState("DELETING")
+
+        events = [event async for event in self.trigger.run()]
+
+        assert events == [
+            TriggerEvent(
+                {
+                    "status": "deleted",
+                    "warehouse_id": WAREHOUSE_ID,
+                    "target_state": "RUNNING",
+                    "last_state": "DELETING",
+                    "state": WarehouseState("DELETING").to_json(),
+                }
+            )
+        ]
+
+    @pytest.mark.asyncio
+    @mock.patch("airflow.providers.databricks.triggers.databricks.time.time", 
return_value=0)
+    
@mock.patch("airflow.providers.databricks.triggers.databricks.asyncio.sleep")
+    
@mock.patch("airflow.providers.databricks.hooks.databricks.DatabricksHook.a_get_warehouse_state")
+    async def test_run_timeout(self, mock_a_get_warehouse_state, mock_sleep, 
mock_time):
+        mock_a_get_warehouse_state.return_value = WarehouseState("STARTING")
+        trigger = DatabricksWarehouseStateTrigger(
+            warehouse_id=WAREHOUSE_ID,
+            target_state="RUNNING",
+            databricks_conn_id=DEFAULT_CONN_ID,
+            polling_period_seconds=POLLING_INTERVAL_SECONDS,
+            end_time=WAREHOUSE_TIMEOUT_SECONDS,
+        )
+
+        def advance_to_deadline(seconds):
+            mock_time.return_value = seconds
+
+        mock_sleep.side_effect = advance_to_deadline
+
+        events = [event async for event in trigger.run()]
+
+        assert events == [
+            TriggerEvent(
+                {
+                    "status": "timeout",
+                    "warehouse_id": WAREHOUSE_ID,
+                    "target_state": "RUNNING",
+                    "last_state": "STARTING",
+                    "state": WarehouseState("STARTING").to_json(),
+                }
+            )
+        ]
+        mock_a_get_warehouse_state.assert_called_once_with(WAREHOUSE_ID)
+        mock_sleep.assert_called_once_with(WAREHOUSE_TIMEOUT_SECONDS)
+
+    @pytest.mark.asyncio
+    @mock.patch("airflow.providers.databricks.triggers.databricks.time.time", 
return_value=0)
+    
@mock.patch("airflow.providers.databricks.triggers.databricks.asyncio.sleep")
+    
@mock.patch("airflow.providers.databricks.hooks.databricks.DatabricksHook.a_get_warehouse_state")
+    async def test_run_accepts_target_received_after_deadline(
+        self, mock_a_get_warehouse_state, mock_sleep, mock_time
+    ):
+        def get_state_after_deadline(_):
+            mock_time.return_value = WAREHOUSE_TIMEOUT_SECONDS + 1
+            return WarehouseState("RUNNING")
+
+        mock_a_get_warehouse_state.side_effect = get_state_after_deadline
+        trigger = DatabricksWarehouseStateTrigger(
+            warehouse_id=WAREHOUSE_ID,
+            target_state="RUNNING",
+            databricks_conn_id=DEFAULT_CONN_ID,
+            polling_period_seconds=POLLING_INTERVAL_SECONDS,
+            end_time=WAREHOUSE_TIMEOUT_SECONDS,
+        )
+
+        events = [event async for event in trigger.run()]
+
+        assert events == [
+            TriggerEvent(
+                {
+                    "status": "success",
+                    "warehouse_id": WAREHOUSE_ID,
+                    "target_state": "RUNNING",
+                    "last_state": "RUNNING",
+                    "state": WarehouseState("RUNNING").to_json(),
+                }
+            )
+        ]
+        mock_sleep.assert_not_called()
+
+    @pytest.mark.asyncio
+    
@mock.patch("airflow.providers.databricks.hooks.databricks.DatabricksHook.a_get_warehouse_state")
+    async def test_run_times_out_when_serialized_end_time_already_passed(self, 
mock_a_get_warehouse_state):
+        trigger = DatabricksWarehouseStateTrigger(
+            warehouse_id=WAREHOUSE_ID,
+            target_state="RUNNING",
+            databricks_conn_id=DEFAULT_CONN_ID,
+            polling_period_seconds=POLLING_INTERVAL_SECONDS,
+            end_time=0,
+        )
+
+        events = [event async for event in trigger.run()]
+
+        assert events == [
+            TriggerEvent(
+                {
+                    "status": "timeout",
+                    "warehouse_id": WAREHOUSE_ID,
+                    "target_state": "RUNNING",
+                    "last_state": "unknown",
+                }
+            )
+        ]
+        mock_a_get_warehouse_state.assert_not_called()

Reply via email to