ramitkataria commented on code in PR #73702:
URL: https://github.com/apache/airflow/pull/73702#discussion_r4128155346


##########
providers/amazon/src/airflow/providers/amazon/aws/operators/glue.py:
##########
@@ -260,6 +267,21 @@ def __init__(
         self.deferrable = deferrable
         self.job_poll_interval = job_poll_interval
         self.stop_job_run_on_kill = stop_job_run_on_kill
+        # In deferrable mode durable reconnects to the still-running job on 
clear, while
+        # stop_job_run_on_kill stops it via the trigger: the two are mutually 
exclusive there. Because
+        # the worker re-executes before the trigger's on_kill runs, a durable 
retry would reconnect to
+        # a run that on_kill is about to stop and then fail, so reject the 
contradiction and keep
+        # durable off when the caller asked to stop on kill. (In synchronous 
mode there is no such
+        # race -- on_kill runs on the worker before the retry -- so the 
combination is allowed.)
+        if self.deferrable and self.stop_job_run_on_kill and self.durable:

Review Comment:
   Could this move to a separate PR? It overrides the 9.35.0 `durable` default, 
and the `ValueError` breaks Dags that parse today (e.g. `durable=True` via 
`default_args`). It also won't give a fresh run: the retry starts before 
`on_kill`, so with `MaxConcurrentRuns=1` it hits 
`ConcurrentRunsExceededException`. Handling it on the worker (stop the old run, 
wait for STOPPED, then submit) would cover sync mode too.



##########
providers/amazon/src/airflow/providers/amazon/aws/triggers/glue.py:
##########
@@ -85,8 +103,114 @@ def __init__(
         self.job_name = job_name
         self.run_id = run_id
         self.verbose = verbose
+        self.stop_job_run_on_kill = stop_job_run_on_kill
+
+    if not AIRFLOW_V_3_0_PLUS:
+
+        @provide_session
+        def get_task_instance(self, *, session: Session) -> TaskInstance:
+            """Get the task instance for the current trigger (Airflow 2.x 
compatibility)."""
+            from sqlalchemy import select
+
+            ti = self.task_instance
+            if ti is None:
+                raise RuntimeError("task_instance is not set on the trigger")
+            query = select(TaskInstance).where(
+                TaskInstance.dag_id == ti.dag_id,
+                TaskInstance.task_id == ti.task_id,
+                TaskInstance.run_id == ti.run_id,
+                TaskInstance.map_index == ti.map_index,
+            )
+            task_instance = session.scalars(query).one_or_none()
+            if task_instance is None:
+                raise ValueError(
+                    f"TaskInstance with dag_id: {ti.dag_id}, "
+                    f"task_id: {ti.task_id}, "
+                    f"run_id: {ti.run_id} and "
+                    f"map_index: {ti.map_index} is not found"
+                )
+            return task_instance
+
+    async def get_task_state(self):
+        """Get the current state of the task instance (Airflow 3.x)."""
+        from airflow.sdk.execution_time.task_runner import RuntimeTaskInstance
+
+        task_states_response = await 
sync_to_async(RuntimeTaskInstance.get_task_states)(
+            dag_id=self.task_instance.dag_id,
+            task_ids=[self.task_instance.task_id],
+            run_ids=[self.task_instance.run_id],
+            map_index=self.task_instance.map_index,
+        )
+        try:
+            task_state = 
task_states_response[self.task_instance.run_id][self.task_instance.task_id]
+        except Exception:
+            raise ValueError(
+                f"TaskInstance with dag_id: {self.task_instance.dag_id}, "
+                f"task_id: {self.task_instance.task_id}, "
+                f"run_id: {self.task_instance.run_id} and "
+                f"map_index: {self.task_instance.map_index} is not found"
+            )
+        return task_state
+
+    async def safe_to_cancel(self) -> bool:
+        """
+        Whether it is safe to stop the Glue job run.
+
+        Returns True if the task is NOT DEFERRED (a user-initiated 
clear/kill). Returns False if the
+        task is still DEFERRED, which means the triggerer is merely restarting 
and the job must keep
+        running.
+        """
+        if AIRFLOW_V_3_0_PLUS:
+            task_state = await self.get_task_state()
+        else:
+            task_instance = self.get_task_instance()  # type: ignore[call-arg]
+            task_state = task_instance.state
+        return task_state != TaskInstanceState.DEFERRED
 
     async def run(self) -> AsyncIterator[TriggerEvent]:
+        """
+        Watch the Glue job run to completion.
+
+        If the task is killed while waiting, stop the Glue job run when 
``stop_job_run_on_kill`` is
+        enabled and it is safe to do so.
+        """
+        try:
+            async for event in self._watch():
+                yield event
+        except asyncio.CancelledError as e:

Review Comment:
   I'd drop this handler and the helpers above, and keep only `on_kill` like 
`EmrContainerTrigger` does. The copied logic breaks for mapped tasks 
(`get_task_states` keys by `f"{task_id}_{map_index}"`, so the lookup raises and 
fails the task on failover). It also still runs on 3.3+ reassignment cancels 
and calls `self.hook().conn` synchronously on the event loop.



##########
providers/amazon/src/airflow/providers/amazon/aws/triggers/glue.py:
##########
@@ -85,8 +103,114 @@ def __init__(
         self.job_name = job_name
         self.run_id = run_id
         self.verbose = verbose
+        self.stop_job_run_on_kill = stop_job_run_on_kill
+
+    if not AIRFLOW_V_3_0_PLUS:
+
+        @provide_session
+        def get_task_instance(self, *, session: Session) -> TaskInstance:
+            """Get the task instance for the current trigger (Airflow 2.x 
compatibility)."""
+            from sqlalchemy import select
+
+            ti = self.task_instance
+            if ti is None:
+                raise RuntimeError("task_instance is not set on the trigger")
+            query = select(TaskInstance).where(
+                TaskInstance.dag_id == ti.dag_id,
+                TaskInstance.task_id == ti.task_id,
+                TaskInstance.run_id == ti.run_id,
+                TaskInstance.map_index == ti.map_index,
+            )
+            task_instance = session.scalars(query).one_or_none()
+            if task_instance is None:
+                raise ValueError(
+                    f"TaskInstance with dag_id: {ti.dag_id}, "
+                    f"task_id: {ti.task_id}, "
+                    f"run_id: {ti.run_id} and "
+                    f"map_index: {ti.map_index} is not found"
+                )
+            return task_instance
+
+    async def get_task_state(self):
+        """Get the current state of the task instance (Airflow 3.x)."""
+        from airflow.sdk.execution_time.task_runner import RuntimeTaskInstance
+
+        task_states_response = await 
sync_to_async(RuntimeTaskInstance.get_task_states)(
+            dag_id=self.task_instance.dag_id,
+            task_ids=[self.task_instance.task_id],
+            run_ids=[self.task_instance.run_id],
+            map_index=self.task_instance.map_index,
+        )
+        try:
+            task_state = 
task_states_response[self.task_instance.run_id][self.task_instance.task_id]
+        except Exception:
+            raise ValueError(
+                f"TaskInstance with dag_id: {self.task_instance.dag_id}, "
+                f"task_id: {self.task_instance.task_id}, "
+                f"run_id: {self.task_instance.run_id} and "
+                f"map_index: {self.task_instance.map_index} is not found"
+            )
+        return task_state
+
+    async def safe_to_cancel(self) -> bool:
+        """
+        Whether it is safe to stop the Glue job run.
+
+        Returns True if the task is NOT DEFERRED (a user-initiated 
clear/kill). Returns False if the
+        task is still DEFERRED, which means the triggerer is merely restarting 
and the job must keep
+        running.
+        """
+        if AIRFLOW_V_3_0_PLUS:
+            task_state = await self.get_task_state()
+        else:
+            task_instance = self.get_task_instance()  # type: ignore[call-arg]
+            task_state = task_instance.state
+        return task_state != TaskInstanceState.DEFERRED
 
     async def run(self) -> AsyncIterator[TriggerEvent]:
+        """
+        Watch the Glue job run to completion.
+
+        If the task is killed while waiting, stop the Glue job run when 
``stop_job_run_on_kill`` is
+        enabled and it is safe to do so.
+        """
+        try:
+            async for event in self._watch():
+                yield event
+        except asyncio.CancelledError as e:
+            # TODO: Remove this handler once the minimum supported Airflow 
version is 3.3+.
+            # On Airflow 3.3+ the triggerer passes a sentinel via 
task.cancel(msg) for
+            # user-initiated kills and calls on_kill() separately -- skip here 
to avoid stopping
+            # the job twice. On older Airflow there is no sentinel, so we 
handle it here.
+            if not (e.args and e.args[0] == "__airflow_user_action__"):
+                if self.run_id and self.stop_job_run_on_kill and await 
self.safe_to_cancel():
+                    self.log.info(
+                        "Task was cancelled. Stopping AWS Glue job %s run 
%s.", self.job_name, self.run_id
+                    )
+                    self.hook().conn.batch_stop_job_run(JobName=self.job_name, 
JobRunIds=[self.run_id])
+                else:
+                    self.log.info(
+                        "Trigger may have shut down or stop_job_run_on_kill is 
disabled. "
+                        "Skipping stop of AWS Glue job %s run %s.",
+                        self.job_name,
+                        self.run_id,
+                    )
+            raise
+
+    async def on_kill(self) -> None:
+        """
+        Stop the Glue job run when the trigger is cancelled by a user action.
+
+        Available on Airflow 3.3+ via ``BaseTrigger.on_kill()``. On older 
Airflow the
+        ``CancelledError`` handler in ``run()`` provides the same behaviour.
+        """
+        if self.run_id and self.stop_job_run_on_kill:
+            self.log.info("Stopping AWS Glue job %s run %s.", self.job_name, 
self.run_id)
+            await sync_to_async(self.hook().conn.batch_stop_job_run)(

Review Comment:
   Could this use `get_async_conn()` like `_watch`, and check `Errors` in the 
response? `sync_to_async(self.hook().conn...)` still builds the client on the 
loop thread. A docstring note would help too: released 3.3.x cores also send 
the user-action cancel on triggerer failover (fixed in #73454, main only).



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]

Reply via email to