diff --git a/providers/snowflake/provider.yaml b/providers/snowflake/provider.yaml index 2189d5a1ef7f8..153bb8aafd340 100644 --- a/providers/snowflake/provider.yaml +++ b/providers/snowflake/provider.yaml @@ -295,6 +295,7 @@ triggers: - integration-name: Snowflake python-modules: - airflow.providers.snowflake.triggers.snowflake_trigger + - airflow.providers.snowflake.triggers.snowpark_containers config: snowflake: diff --git a/providers/snowflake/src/airflow/providers/snowflake/get_provider_info.py b/providers/snowflake/src/airflow/providers/snowflake/get_provider_info.py index 7e20af8004ee2..2fdf08911a572 100644 --- a/providers/snowflake/src/airflow/providers/snowflake/get_provider_info.py +++ b/providers/snowflake/src/airflow/providers/snowflake/get_provider_info.py @@ -164,7 +164,10 @@ def get_provider_info(): "triggers": [ { "integration-name": "Snowflake", - "python-modules": ["airflow.providers.snowflake.triggers.snowflake_trigger"], + "python-modules": [ + "airflow.providers.snowflake.triggers.snowflake_trigger", + "airflow.providers.snowflake.triggers.snowpark_containers", + ], } ], "config": { diff --git a/providers/snowflake/src/airflow/providers/snowflake/operators/snowpark_containers.py b/providers/snowflake/src/airflow/providers/snowflake/operators/snowpark_containers.py index 971be06543ace..842157b7ab895 100644 --- a/providers/snowflake/src/airflow/providers/snowflake/operators/snowpark_containers.py +++ b/providers/snowflake/src/airflow/providers/snowflake/operators/snowpark_containers.py @@ -19,51 +19,25 @@ import time from collections.abc import Sequence -from enum import Enum +from datetime import timedelta from functools import cached_property from typing import TYPE_CHECKING, Any +from airflow.providers.common.compat.sdk import conf from airflow.providers.common.compat.standard.operators import BaseOperator from airflow.providers.common.sql.hooks.handlers import fetch_one_handler from airflow.providers.snowflake.hooks.snowflake import SnowflakeHook +from airflow.providers.snowflake.triggers.snowpark_containers import ( + NON_TERMINAL_STATUSES, + TERMINAL_STATUSES, + SnowparkContainerJobStatus, + SnowparkContainerJobTrigger, +) if TYPE_CHECKING: from airflow.providers.common.compat.sdk import Context -class SnowparkContainerJobStatus(str, Enum): - """Statuses of a Snowpark Container Services.""" - - PENDING = "PENDING" - RUNNING = "RUNNING" - CANCELLING = "CANCELLING" - SUSPENDING = "SUSPENDING" - DELETING = "DELETING" - DONE = "DONE" - FAILED = "FAILED" - CANCELLED = "CANCELLED" - INTERNAL_ERROR = "INTERNAL_ERROR" - - -TERMINAL_STATUSES: frozenset[SnowparkContainerJobStatus] = frozenset( - { - SnowparkContainerJobStatus.DONE, - SnowparkContainerJobStatus.FAILED, - SnowparkContainerJobStatus.CANCELLED, - SnowparkContainerJobStatus.INTERNAL_ERROR, - } -) -NON_TERMINAL_STATUSES: frozenset[SnowparkContainerJobStatus] = frozenset( - { - SnowparkContainerJobStatus.PENDING, - SnowparkContainerJobStatus.RUNNING, - SnowparkContainerJobStatus.CANCELLING, - SnowparkContainerJobStatus.SUSPENDING, - SnowparkContainerJobStatus.DELETING, - } -) - - class SnowparkContainerJobOperator(BaseOperator): """ Execute a job on Snowpark Container Services. @@ -100,6 +74,14 @@ class SnowparkContainerJobOperator(BaseOperator): (default value: 10) :param snowflake_conn_id: Reference to :ref:`Snowflake connection id` + :param deferrable: Run the operator in deferrable mode. Only effective when + ``wait_for_completion`` is True. With ``wait_for_completion=False`` the + operator submits the job and returns immediately without deferring. + (default value: False) + :param timeout: Maximum seconds to wait for the job to reach a terminal + state. When it elapses the job service is dropped and the task fails. + A shorter ``execution_timeout`` preempts this cleanup, so the service + may be left running. (default value: 86400) :param database: name of database (will overwrite database defined in connection) :param schema: name of schema (will overwrite schema defined in @@ -137,6 +119,8 @@ def __init__( drop_on_completion: bool = True, poll_interval: int = 10, snowflake_conn_id: str = "snowflake_default", + deferrable: bool = conf.getboolean("operators", "default_deferrable", fallback=False), + timeout: int = 24 * 60 * 60, database: str | None = None, schema: str | None = None, role: str | None = None, @@ -156,6 +140,8 @@ def __init__( self.drop_on_completion = drop_on_completion self.poll_interval = poll_interval self.snowflake_conn_id = snowflake_conn_id + self.deferrable = deferrable + self.timeout = timeout self.database = database self.schema = schema self.role = role @@ -167,6 +153,8 @@ def __init__( raise ValueError("Cannot specify both 'spec_text' and 'spec'/'spec_stage'") if not self.spec_text and not (self.spec and self.spec_stage): raise ValueError("Must provide either 'spec_text' or both 'spec' and 'spec_stage'") + if self.deferrable and not self.wait_for_completion: + self.log.warning("deferrable has no effect when wait_for_completion is False.") @cached_property def _hook(self) -> SnowflakeHook: @@ -206,6 +194,7 @@ def _submit_job(self) -> str: def _poll_for_status(self) -> str: """Poll until the job reaches a terminal state.""" + end_time = time.time() + self.timeout while True: response = self._run_one(f"DESCRIBE SERVICE {self.job_name}", return_dictionaries=True) status = response.get("status") @@ -213,6 +202,9 @@ def _poll_for_status(self) -> str: return status if status not in NON_TERMINAL_STATUSES: raise RuntimeError(f"Job {self.job_name} returned unexpected status: {status}") + if time.time() > end_time: + self._drop_service() + raise TimeoutError(f"Job {self.job_name} did not reach a terminal status before the timeout.") time.sleep(self.poll_interval) def _log_container_output(self, status: str) -> None: @@ -227,13 +219,25 @@ def _log_container_output(self, status: str) -> None: else: self.log.info("Logs for instance_id %d:\n%s", instance_id, response) + def _drop_service(self) -> None: + """Best-effort drop of the job service.""" + try: + self._hook.run(f"DROP SERVICE IF EXISTS {self.job_name}") + except Exception as e: + self.log.error("Error dropping service %s: %s", self.job_name, e) + def on_kill(self) -> None: """Drop the running service on task kill.""" if self.job_name: - try: - self._hook.run(f"DROP SERVICE IF EXISTS {self.job_name}") - except Exception as e: - self.log.error("Error dropping service %s: %s", self.job_name, e) + self._drop_service() + + def _handle_final_status(self, status: str) -> None: + """Log container output, fail unless the job is DONE, and optionally drop the service on success.""" + self._log_container_output(status) + if status != SnowparkContainerJobStatus.DONE: + raise RuntimeError(f"Job '{self.job_name}' finished with status: {status}") + if self.drop_on_completion: + self._drop_service() def execute(self, context: Context) -> str: """Submit and optionally wait for a Snowpark Container Services job.""" @@ -242,10 +246,35 @@ def execute(self, context: Context) -> str: raise RuntimeError("Job name was not returned") if not self.wait_for_completion: return self.job_name + if self.deferrable: + self.defer( + trigger=SnowparkContainerJobTrigger( + job_name=self.job_name, + snowflake_conn_id=self.snowflake_conn_id, + poll_interval=self.poll_interval, + end_time=time.time() + self.timeout, + database=self.database, + schema=self.schema, + role=self.role, + warehouse=self.warehouse, + ), + # Pad past the trigger's end_time so its timeout event, which drops the service, + # fires before this hard backstop. A user-set execution_timeout takes precedence. + timeout=self.execution_timeout or timedelta(seconds=self.timeout + self.poll_interval + 60), + method_name="execute_complete", + ) status = self._poll_for_status() - self._log_container_output(status) - if status != SnowparkContainerJobStatus.DONE: - raise RuntimeError(f"Job '{self.job_name}' finished with status: {status}") - if self.drop_on_completion: - self._hook.run(f"DROP SERVICE IF EXISTS {self.job_name}") + self._handle_final_status(status) + return self.job_name + + def execute_complete(self, context: Context, event: dict[str, Any]) -> str: + """Resume after the trigger fires.""" + self.job_name = event["job_name"] + status = event["status"] + if status == "timeout": + self._drop_service() + raise TimeoutError(event.get("message", f"Job '{self.job_name}' did not complete: {status}")) + if status == "error": + raise RuntimeError(event.get("message", f"Job '{self.job_name}' did not complete: {status}")) + self._handle_final_status(status) return self.job_name diff --git a/providers/snowflake/src/airflow/providers/snowflake/triggers/snowpark_containers.py b/providers/snowflake/src/airflow/providers/snowflake/triggers/snowpark_containers.py new file mode 100644 index 0000000000000..cd79a06860824 --- /dev/null +++ b/providers/snowflake/src/airflow/providers/snowflake/triggers/snowpark_containers.py @@ -0,0 +1,179 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +import asyncio +import time +from collections.abc import AsyncIterator +from enum import Enum +from typing import Any + +from airflow.providers.common.sql.hooks.handlers import fetch_one_handler +from airflow.providers.snowflake.hooks.snowflake import SnowflakeHook +from airflow.triggers.base import BaseTrigger, TriggerEvent + + +class SnowparkContainerJobStatus(str, Enum): + """Statuses of a Snowpark Container Services job service.""" + + PENDING = "PENDING" + RUNNING = "RUNNING" + CANCELLING = "CANCELLING" + SUSPENDING = "SUSPENDING" + DELETING = "DELETING" + DONE = "DONE" + FAILED = "FAILED" + CANCELLED = "CANCELLED" + INTERNAL_ERROR = "INTERNAL_ERROR" + + +TERMINAL_STATUSES: frozenset[SnowparkContainerJobStatus] = frozenset( + { + SnowparkContainerJobStatus.DONE, + SnowparkContainerJobStatus.FAILED, + SnowparkContainerJobStatus.CANCELLED, + SnowparkContainerJobStatus.INTERNAL_ERROR, + } +) +NON_TERMINAL_STATUSES: frozenset[SnowparkContainerJobStatus] = frozenset( + { + SnowparkContainerJobStatus.PENDING, + SnowparkContainerJobStatus.RUNNING, + SnowparkContainerJobStatus.CANCELLING, + SnowparkContainerJobStatus.SUSPENDING, + SnowparkContainerJobStatus.DELETING, + } +) + + +class SnowparkContainerJobTrigger(BaseTrigger): + """ + Poll a Snowpark Container Services job until it reaches a terminal status. + + :param job_name: name of the submitted job service to poll. + :param snowflake_conn_id: reference to the Snowflake connection id. + :param poll_interval: seconds to sleep between ``DESCRIBE SERVICE`` polls. + :param end_time: epoch deadline (``time.time()`` seconds) after which a ``timeout`` + event is emitted. + :param database: (Optional) name of database. + :param schema: (Optional) name of schema. + :param role: (Optional) name of role. + :param warehouse: (Optional) name of warehouse. + """ + + def __init__( + self, + job_name: str, + snowflake_conn_id: str, + poll_interval: float, + end_time: float, + database: str | None = None, + schema: str | None = None, + role: str | None = None, + warehouse: str | None = None, + ) -> None: + super().__init__() + self.job_name = job_name + self.snowflake_conn_id = snowflake_conn_id + self.poll_interval = poll_interval + self.end_time = end_time + self.database = database + self.schema = schema + self.role = role + self.warehouse = warehouse + + def serialize(self) -> tuple[str, dict[str, Any]]: + """Serialize SnowparkContainerJobTrigger arguments and class path.""" + return ( + "airflow.providers.snowflake.triggers.snowpark_containers.SnowparkContainerJobTrigger", + { + "job_name": self.job_name, + "snowflake_conn_id": self.snowflake_conn_id, + "poll_interval": self.poll_interval, + "end_time": self.end_time, + "database": self.database, + "schema": self.schema, + "role": self.role, + "warehouse": self.warehouse, + }, + ) + + def _get_hook(self) -> SnowflakeHook: + """Build a ``SnowflakeHook`` from the trigger's connection settings.""" + return SnowflakeHook( + snowflake_conn_id=self.snowflake_conn_id, + warehouse=self.warehouse, + database=self.database, + schema=self.schema, + role=self.role, + ) + + async def _describe_status(self, hook: SnowflakeHook) -> str | None: + """Return the job's current status via ``DESCRIBE SERVICE``, or ``None`` if absent.""" + # SnowflakeHook is synchronous. Run the blocking poll off the event loop so a + # single query does not stall every other trigger on this triggerer. + response: Any = await asyncio.to_thread( + hook.run, + f"DESCRIBE SERVICE {self.job_name}", + handler=fetch_one_handler, + return_dictionaries=True, + ) + return response.get("status") if response else None + + async def run(self) -> AsyncIterator[TriggerEvent]: + """Poll the job status and yield exactly one terminal event.""" + hook = self._get_hook() + while True: + try: + status = await self._describe_status(hook=hook) + except Exception as e: + yield TriggerEvent({"status": "error", "job_name": self.job_name, "message": str(e)}) + return + + if status in TERMINAL_STATUSES: + yield TriggerEvent({"status": status, "job_name": self.job_name}) + return + + if status not in NON_TERMINAL_STATUSES: + yield TriggerEvent( + { + "status": "error", + "job_name": self.job_name, + "message": f"Job {self.job_name} returned unexpected status: {status}", + } + ) + return + + if time.time() > self.end_time: + yield TriggerEvent( + { + "status": "timeout", + "job_name": self.job_name, + "message": f"Job {self.job_name} did not reach a terminal status before the timeout.", + } + ) + return + await asyncio.sleep(self.poll_interval) + + async def on_kill(self) -> None: + """Drop the job service when a deferred task is killed.""" + hook = self._get_hook() + try: + await asyncio.to_thread(hook.run, f"DROP SERVICE IF EXISTS {self.job_name}") + self.log.info("on_kill: dropped service %s", self.job_name) + except Exception as e: + self.log.error("on_kill: failed to drop service %s: %s", self.job_name, e) diff --git a/providers/snowflake/tests/unit/snowflake/operators/test_snowpark_containers.py b/providers/snowflake/tests/unit/snowflake/operators/test_snowpark_containers.py index 3b10bfea06d85..dd59d1990f86d 100644 --- a/providers/snowflake/tests/unit/snowflake/operators/test_snowpark_containers.py +++ b/providers/snowflake/tests/unit/snowflake/operators/test_snowpark_containers.py @@ -20,7 +20,9 @@ import pytest +from airflow.providers.common.compat.sdk import TaskDeferred from airflow.providers.snowflake.operators.snowpark_containers import SnowparkContainerJobOperator +from airflow.providers.snowflake.triggers.snowpark_containers import SnowparkContainerJobTrigger TASK_ID = "test_spcs_job" COMPUTE_POOL = "test_pool" @@ -75,6 +77,21 @@ def test_invalid_spec_combinations(self, kwargs, match): with pytest.raises(ValueError, match=match): _make_operator(**kwargs) + @pytest.mark.parametrize( + ("deferrable", "wait_for_completion", "warns"), + ( + pytest.param(True, False, True, id="deferrable_no_wait"), + pytest.param(True, True, False, id="deferrable_wait"), + pytest.param(False, False, False, id="sync_no_wait"), + ), + ) + @mock.patch.object(SnowparkContainerJobOperator, "log") + def test_warns_when_deferrable_without_wait_for_completion( + self, mock_log, deferrable, wait_for_completion, warns + ): + _make_operator(deferrable=deferrable, wait_for_completion=wait_for_completion) + assert mock_log.warning.called is warns + def test_build_sql_with_spec_stage(self): op = _make_operator() sql = op._build_sql() @@ -150,6 +167,21 @@ def test_poll_waits_through_pending_then_done(self, mock_hook_cls, mock_sleep): assert op._poll_for_status() == "DONE" assert mock_sleep.call_count == 2 + @mock.patch("time.sleep") + @mock.patch(MOCK_HOOK_PATH) + def test_poll_raises_on_timeout(self, mock_hook_cls, sleep_mock, time_machine): + mock_hook = mock_hook_cls.return_value + mock_hook.run.return_value = {"status": "RUNNING"} + op = _make_operator(poll_interval=5, timeout=10) + op.job_name = JOB_NAME + + time_machine.move_to(0, tick=False) + sleep_mock.side_effect = lambda seconds: time_machine.shift(seconds + 0.1) + + with pytest.raises(TimeoutError, match="did not reach a terminal status"): + op._poll_for_status() + mock_hook.run.assert_any_call(f"DROP SERVICE IF EXISTS {JOB_NAME}") + @mock.patch(MOCK_HOOK_PATH) def test_log_container_output_uses_info_on_done(self, mock_hook_cls): mock_hook = mock_hook_cls.return_value @@ -280,3 +312,58 @@ def test_execute_skips_drop_when_disabled(self, mock_hook_cls, mock_submit, mock op = _make_operator(drop_on_completion=False) op.execute(context=None) mock_hook.run.assert_not_called() + + @mock.patch.object(SnowparkContainerJobOperator, "_submit_job", return_value=JOB_NAME) + def test_execute_defers_when_deferrable(self, mock_submit): + op = _make_operator(deferrable=True) + with pytest.raises(TaskDeferred) as exc: + op.execute(context=None) + assert isinstance(exc.value.trigger, SnowparkContainerJobTrigger) + assert exc.value.trigger.job_name == JOB_NAME + assert exc.value.method_name == "execute_complete" + + @mock.patch(MOCK_HOOK_PATH) + @mock.patch.object(SnowparkContainerJobOperator, "_log_container_output") + def test_execute_complete_success_drops_and_returns(self, mock_log, mock_hook_cls): + mock_hook = mock_hook_cls.return_value + op = _make_operator(drop_on_completion=True) + result = op.execute_complete(context=None, event={"status": "DONE", "job_name": JOB_NAME}) + assert result == JOB_NAME + mock_log.assert_called_once_with("DONE") + mock_hook.run.assert_called_once_with(f"DROP SERVICE IF EXISTS {JOB_NAME}") + + @mock.patch(MOCK_HOOK_PATH) + @mock.patch.object(SnowparkContainerJobOperator, "_log_container_output") + def test_execute_complete_failure_raises_without_drop(self, mock_log, mock_hook_cls): + mock_hook = mock_hook_cls.return_value + op = _make_operator() + with pytest.raises(RuntimeError, match="FAILED"): + op.execute_complete(context=None, event={"status": "FAILED", "job_name": JOB_NAME}) + mock_log.assert_called_once_with("FAILED") + mock_hook.run.assert_not_called() + + @pytest.mark.parametrize( + ("status", "exc", "drops"), + [("timeout", TimeoutError, True), ("error", RuntimeError, False)], + ) + @mock.patch(MOCK_HOOK_PATH) + def test_execute_complete_raises_and_drops_only_on_timeout(self, mock_hook_cls, status, exc, drops): + mock_hook = mock_hook_cls.return_value + op = _make_operator() + with pytest.raises(exc, match="boom"): + op.execute_complete( + context=None, + event={"status": status, "job_name": JOB_NAME, "message": "boom"}, + ) + if drops: + mock_hook.run.assert_called_once_with(f"DROP SERVICE IF EXISTS {JOB_NAME}") + else: + mock_hook.run.assert_not_called() + + @mock.patch(MOCK_HOOK_PATH) + @mock.patch.object(SnowparkContainerJobOperator, "_log_container_output") + def test_execute_complete_skips_drop_when_disabled(self, mock_log, mock_hook_cls): + mock_hook = mock_hook_cls.return_value + op = _make_operator(drop_on_completion=False) + op.execute_complete(context=None, event={"status": "DONE", "job_name": JOB_NAME}) + mock_hook.run.assert_not_called() diff --git a/providers/snowflake/tests/unit/snowflake/triggers/test_snowpark_containers.py b/providers/snowflake/tests/unit/snowflake/triggers/test_snowpark_containers.py new file mode 100644 index 0000000000000..777112e1bc355 --- /dev/null +++ b/providers/snowflake/tests/unit/snowflake/triggers/test_snowpark_containers.py @@ -0,0 +1,143 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +import time +from unittest import mock + +import pytest + +from airflow.providers.snowflake.triggers.snowpark_containers import SnowparkContainerJobTrigger +from airflow.triggers.base import TriggerEvent + +TRIGGER_PATH = "airflow.providers.snowflake.triggers.snowpark_containers" +CLASSPATH = f"{TRIGGER_PATH}.SnowparkContainerJobTrigger" +HOOK = f"{TRIGGER_PATH}.SnowflakeHook" + +JOB_NAME = "TEST_JOB" +CONN_ID = "snowflake_default" +POLL_INTERVAL = 1.0 + + +class TestSnowparkContainerJobTrigger: + @staticmethod + def _describe(status): + return {"status": status} + + def _trigger(self, end_time=None, **kwargs): + params = { + "job_name": JOB_NAME, + "snowflake_conn_id": CONN_ID, + "poll_interval": POLL_INTERVAL, + "end_time": end_time if end_time is not None else time.time() + 3600, + } + params.update(kwargs) + return SnowparkContainerJobTrigger(**params) + + def test_serialization(self): + end_time = time.time() + 3600 + class_path, kwargs = self._trigger( + end_time, database="db", schema="sc", role="r", warehouse="wh" + ).serialize() + assert class_path == CLASSPATH + assert kwargs == { + "job_name": JOB_NAME, + "snowflake_conn_id": CONN_ID, + "poll_interval": POLL_INTERVAL, + "end_time": end_time, + "database": "db", + "schema": "sc", + "role": "r", + "warehouse": "wh", + } + + @mock.patch(HOOK, autospec=True) + def test_get_hook_builds_hook_from_connection_settings(self, mock_hook_cls): + self._trigger(database="db", schema="sc", role="r", warehouse="wh")._get_hook() + mock_hook_cls.assert_called_once_with( + snowflake_conn_id=CONN_ID, + warehouse="wh", + database="db", + schema="sc", + role="r", + ) + + @pytest.mark.asyncio + @mock.patch(HOOK, autospec=True) + async def test_describe_status_returns_none_when_no_row(self, mock_hook_cls): + mock_hook_cls.return_value.run.return_value = None + trigger = self._trigger() + assert await trigger._describe_status(trigger._get_hook()) is None + + @pytest.mark.asyncio + @pytest.mark.parametrize("status", ("DONE", "FAILED", "CANCELLED", "INTERNAL_ERROR")) + @mock.patch(HOOK, autospec=True) + async def test_terminal_status_yields_status_event(self, mock_hook_cls, status): + mock_hook_cls.return_value.run.return_value = self._describe(status) + event = await self._trigger().run().__anext__() + assert event == TriggerEvent({"status": status, "job_name": JOB_NAME}) + + @pytest.mark.asyncio + @mock.patch(HOOK, autospec=True) + async def test_unexpected_status_yields_error(self, mock_hook_cls): + mock_hook_cls.return_value.run.return_value = self._describe("RANDOM") + event = await self._trigger().run().__anext__() + assert event.payload["status"] == "error" + assert "unexpected status" in event.payload["message"] + + @pytest.mark.asyncio + @mock.patch(HOOK, autospec=True) + async def test_timeout_yields_timeout_event(self, mock_hook_cls): + mock_hook_cls.return_value.run.return_value = self._describe("RUNNING") + event = await self._trigger(end_time=time.time() - 1).run().__anext__() + assert event.payload["status"] == "timeout" + assert event.payload["job_name"] == JOB_NAME + + @pytest.mark.asyncio + @mock.patch(HOOK, autospec=True) + async def test_poll_exception_yields_error(self, mock_hook_cls): + mock_hook_cls.return_value.run.side_effect = RuntimeError("boom") + event = await self._trigger().run().__anext__() + assert event == TriggerEvent({"status": "error", "job_name": JOB_NAME, "message": "boom"}) + + @pytest.mark.asyncio + @mock.patch(f"{TRIGGER_PATH}.asyncio.sleep") + @mock.patch(HOOK, autospec=True) + async def test_polls_through_non_terminal_then_terminal(self, mock_hook_cls, mock_sleep): + mock_hook_cls.return_value.run.side_effect = [ + self._describe("PENDING"), + self._describe("RUNNING"), + self._describe("DONE"), + ] + event = await self._trigger().run().__anext__() + assert event == TriggerEvent({"status": "DONE", "job_name": JOB_NAME}) + assert mock_sleep.await_count == 2 + + @pytest.mark.asyncio + @mock.patch(HOOK, autospec=True) + async def test_on_kill_drops_service(self, mock_hook_cls): + await self._trigger().on_kill() + mock_hook_cls.return_value.run.assert_called_once_with(f"DROP SERVICE IF EXISTS {JOB_NAME}") + + @pytest.mark.asyncio + @mock.patch(HOOK, autospec=True) + async def test_on_kill_logs_error_on_failure(self, mock_hook_cls): + mock_hook_cls.return_value.run.side_effect = RuntimeError("drop failed") + trigger = self._trigger() + with mock.patch.object(trigger.log, "error") as mock_error: + await trigger.on_kill() + mock_error.assert_called_once()