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]