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


##########
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:
   Agreed, dropped the `CancelledError` handler and the 
`safe_to_cancel`/`get_task_state` helpers; the trigger now relies on 
`on_kill()` only, like `EmrContainerTrigger`, and the hook is inert below 3.3 
(noted in the docstring). Thanks for catching the mapped-task key. The same 
`[run_id][task_id]` lookup is in `EmrServerlessStartJobTrigger` and a few 
Google triggers, so I'll look at that separately.
   
   ---
   Drafted-by: Claude Code (Opus 5.5); reviewed by @Leondon9 before posting



##########
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:
   Removed. The operator change is now just passing `stop_job_run_on_kill` to 
the trigger, so the `durable` default and parsing behaviour are untouched. 
You're right that it didn't give a fresh run either: the retry submits while 
the old run is still RUNNING/STOPPING. I've noted that as out of scope in the 
description; stopping the prior run on the worker and waiting for STOPPED 
before submitting looks like the right fix, and I'll take that up separately.
   
   ---
   Drafted-by: Claude Code (Opus 5.5); reviewed by @Leondon9 before posting



##########
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:
   Done: `on_kill` now uses `get_async_conn()` and logs the `Errors` from 
`batch_stop_job_run`. Added a docstring note that released 3.3.x cores also 
call it on triggerer failover (fixed in #73454).
   
   ---
   Drafted-by: Claude Code (Opus 5.5); reviewed by @Leondon9 before posting



-- 
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