diff --git a/providers/snowflake/src/airflow/providers/snowflake/operators/snowflake.py b/providers/snowflake/src/airflow/providers/snowflake/operators/snowflake.py index 8e3cda63ef3d7..288ea6328c507 100644 --- a/providers/snowflake/src/airflow/providers/snowflake/operators/snowflake.py +++ b/providers/snowflake/src/airflow/providers/snowflake/operators/snowflake.py @@ -380,6 +380,9 @@ class SnowflakeSqlApiOperator(ResumableJobMixin, SQLExecuteQueryOperator): To set the timeout to the maximum value (604800 seconds), set timeout to 0. :param deferrable: Run operator in the deferrable mode. :param snowflake_api_retry_args: An optional dictionary with arguments passed to ``tenacity.Retrying`` & ``tenacity.AsyncRetrying`` classes. + :param cancel_on_kill: If True (default), cancel the running Snowflake queries when the task is + killed. This applies both while the operator is running and, for a deferred task, while it + waits in the triggerer. :param durable: When ``True`` (the default), the submitted statement handles are persisted to task state before polling begins. A worker crash on retry reconnects to the existing statements instead of resubmitting the SQL. Set to ``False`` to always submit fresh on @@ -413,6 +416,7 @@ def __init__( timeout: int | None = None, deferrable: bool = conf.getboolean("operators", "default_deferrable", fallback=False), snowflake_api_retry_args: dict[str, Any] | None = None, + cancel_on_kill: bool = True, **kwargs: Any, ) -> None: self.snowflake_conn_id = snowflake_conn_id @@ -425,6 +429,7 @@ def __init__( self.execute_async = False self.snowflake_api_retry_args = snowflake_api_retry_args or {} self.deferrable = deferrable + self.cancel_on_kill = cancel_on_kill self.query_ids: list[str] = [] if any([warehouse, database, role, schema, authenticator, session_parameters]): # pragma: no cover hook_params = kwargs.pop("hook_params", {}) # pragma: no cover @@ -491,6 +496,7 @@ def execute(self, context: Context) -> None: snowflake_conn_id=self.snowflake_conn_id, token_life_time=self.token_life_time, token_renewal_delta=self.token_renewal_delta, + cancel_on_kill=self.cancel_on_kill, ), method_name="execute_complete", ) @@ -617,6 +623,8 @@ def execute_complete(self, context: Context, event: dict[str, str | list[str]] | def on_kill(self) -> None: """Cancel the running query.""" + if not self.cancel_on_kill: + return if self.query_ids: self.log.info("Cancelling the query ids %s", self.query_ids) self._hook.cancel_queries(self.query_ids) diff --git a/providers/snowflake/src/airflow/providers/snowflake/triggers/snowflake_trigger.py b/providers/snowflake/src/airflow/providers/snowflake/triggers/snowflake_trigger.py index e460ea7840ce4..30993a661499d 100644 --- a/providers/snowflake/src/airflow/providers/snowflake/triggers/snowflake_trigger.py +++ b/providers/snowflake/src/airflow/providers/snowflake/triggers/snowflake_trigger.py @@ -20,6 +20,8 @@ from collections.abc import AsyncIterator from typing import TYPE_CHECKING, Any +from asgiref.sync import sync_to_async + from airflow.providers.snowflake.hooks.snowflake_sql_api import SnowflakeSqlApiHook from airflow.triggers.base import BaseTrigger, TriggerEvent @@ -36,6 +38,9 @@ class SnowflakeSqlApiTrigger(BaseTrigger): :param snowflake_conn_id: Reference to Snowflake connection id :param token_life_time: lifetime of the JWT Token in timedelta :param token_renewal_delta: Renewal time of the JWT Token in timedelta + :param cancel_on_kill: If True (default), cancel the running Snowflake queries when the user + kills the deferred task (mark failed, clear, or mark success). Requires a version of + ``apache-airflow`` with ``BaseTrigger.on_kill()`` support; on older versions it is inert. """ def __init__( @@ -45,6 +50,7 @@ def __init__( snowflake_conn_id: str, token_life_time: timedelta, token_renewal_delta: timedelta, + cancel_on_kill: bool = True, ): super().__init__() self.poll_interval = poll_interval @@ -52,6 +58,7 @@ def __init__( self.snowflake_conn_id = snowflake_conn_id self.token_life_time = token_life_time self.token_renewal_delta = token_renewal_delta + self.cancel_on_kill = cancel_on_kill def serialize(self) -> tuple[str, dict[str, Any]]: """Serialize SnowflakeSqlApiTrigger arguments and classpath.""" @@ -63,6 +70,7 @@ def serialize(self) -> tuple[str, dict[str, Any]]: "snowflake_conn_id": self.snowflake_conn_id, "token_life_time": self.token_life_time, "token_renewal_delta": self.token_renewal_delta, + "cancel_on_kill": self.cancel_on_kill, }, ) @@ -93,6 +101,27 @@ async def run(self) -> AsyncIterator[TriggerEvent]: except Exception as e: yield TriggerEvent({"status": "error", "message": str(e)}) + async def on_kill(self) -> None: + """Cancel the running Snowflake queries when the user kills the deferred task.""" + if not self.cancel_on_kill or not self.query_ids: + return + self.log.info("Cancelling Snowflake query ids %s", self.query_ids) + try: + await sync_to_async(self._cancel_queries)() + self.log.info("Snowflake query ids %s cancelled.", self.query_ids) + except Exception: + self.log.exception( + "Failed to cancel Snowflake query ids %s. They may still be running.", self.query_ids + ) + + def _cancel_queries(self) -> None: + hook = SnowflakeSqlApiHook( + self.snowflake_conn_id, + self.token_life_time, + self.token_renewal_delta, + ) + hook.cancel_queries(self.query_ids) + async def get_query_status( self, query_id: str, hook: SnowflakeSqlApiHook | None = None ) -> dict[str, Any]: diff --git a/providers/snowflake/tests/unit/snowflake/operators/test_snowflake.py b/providers/snowflake/tests/unit/snowflake/operators/test_snowflake.py index d37f420f2a66f..f37785ab7c8fe 100644 --- a/providers/snowflake/tests/unit/snowflake/operators/test_snowflake.py +++ b/providers/snowflake/tests/unit/snowflake/operators/test_snowflake.py @@ -456,6 +456,7 @@ def test_snowflake_sql_api_execute_operator_async( assert isinstance(exc.value.trigger, SnowflakeSqlApiTrigger), ( "Trigger is not a SnowflakeSqlApiTrigger" ) + assert exc.value.trigger.cancel_on_kill is True def test_snowflake_sql_api_pushes_query_ids_to_xcom( self, @@ -750,6 +751,22 @@ def test_snowflake_sql_api_on_kill_no_queries(self, mock_cancel_queries): mock_cancel_queries.assert_not_called() + @mock.patch("airflow.providers.snowflake.hooks.snowflake_sql_api.SnowflakeSqlApiHook.cancel_queries") + def test_snowflake_sql_api_on_kill_respects_cancel_on_kill_false(self, mock_cancel_queries): + """on_kill does not cancel queries when cancel_on_kill is disabled.""" + operator = SnowflakeSqlApiOperator( + task_id=TASK_ID, + snowflake_conn_id=CONN_ID, + sql=SQL_MULTIPLE_STMTS, + statement_count=4, + cancel_on_kill=False, + ) + operator.query_ids = ["uuid1", "uuid2"] + + operator.on_kill() + + mock_cancel_queries.assert_not_called() + @pytest.mark.skipif( not AIRFLOW_V_3_3_PLUS, reason="task_state_store (durable execution) requires Airflow 3.3+" diff --git a/providers/snowflake/tests/unit/snowflake/triggers/test_snowflake.py b/providers/snowflake/tests/unit/snowflake/triggers/test_snowflake.py index 42a3a07d224de..5eb941e9e0a77 100644 --- a/providers/snowflake/tests/unit/snowflake/triggers/test_snowflake.py +++ b/providers/snowflake/tests/unit/snowflake/triggers/test_snowflake.py @@ -55,8 +55,60 @@ def test_snowflake_sql_trigger_serialization(self): "snowflake_conn_id": "test_conn", "token_life_time": LIFETIME, "token_renewal_delta": RENEWAL_DELTA, + "cancel_on_kill": True, } + def test_snowflake_sql_trigger_serialization_cancel_on_kill_false(self): + """cancel_on_kill=False round-trips through serialization.""" + trigger = SnowflakeSqlApiTrigger( + poll_interval=POLL_INTERVAL, + query_ids=QUERY_IDS, + snowflake_conn_id="test_conn", + token_life_time=LIFETIME, + token_renewal_delta=RENEWAL_DELTA, + cancel_on_kill=False, + ) + _, kwargs = trigger.serialize() + assert kwargs["cancel_on_kill"] is False + + @pytest.mark.asyncio + @mock.patch(f"{MODULE}.triggers.snowflake_trigger.SnowflakeSqlApiHook") + async def test_on_kill_cancels_the_queries(self, mock_hook): + """on_kill() cancels the running queries when enabled and query_ids are set.""" + await self.TRIGGER.on_kill() + mock_hook.assert_called_once_with("test_conn", LIFETIME, RENEWAL_DELTA) + mock_hook.return_value.cancel_queries.assert_called_once_with(QUERY_IDS) + + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("cancel_on_kill", "query_ids"), + [ + pytest.param(False, QUERY_IDS, id="disabled"), + pytest.param(True, [], id="no-query-ids"), + ], + ) + @mock.patch(f"{MODULE}.triggers.snowflake_trigger.SnowflakeSqlApiHook") + async def test_on_kill_does_not_cancel(self, mock_hook, cancel_on_kill, query_ids): + """on_kill() is a no-op (no hook built) when disabled or without query_ids.""" + trigger = SnowflakeSqlApiTrigger( + poll_interval=POLL_INTERVAL, + query_ids=query_ids, + snowflake_conn_id="test_conn", + token_life_time=LIFETIME, + token_renewal_delta=RENEWAL_DELTA, + cancel_on_kill=cancel_on_kill, + ) + await trigger.on_kill() + mock_hook.assert_not_called() + + @pytest.mark.asyncio + @mock.patch(f"{MODULE}.triggers.snowflake_trigger.SnowflakeSqlApiHook") + async def test_on_kill_swallows_cancel_errors(self, mock_hook): + """on_kill() logs and swallows exceptions raised while cancelling.""" + mock_hook.return_value.cancel_queries.side_effect = Exception("Snowflake API error") + await self.TRIGGER.on_kill() + mock_hook.return_value.cancel_queries.assert_called_once_with(QUERY_IDS) + @pytest.mark.asyncio @mock.patch(f"{MODULE}.triggers.snowflake_trigger.SnowflakeSqlApiTrigger.get_query_status") @mock.patch(f"{MODULE}.hooks.snowflake_sql_api.SnowflakeSqlApiHook.get_sql_api_query_status_async")