diff --git a/providers/standard/src/airflow/providers/standard/operators/trigger_dagrun.py b/providers/standard/src/airflow/providers/standard/operators/trigger_dagrun.py index 958cc258cf3d8..0ad36d714ee9a 100644 --- a/providers/standard/src/airflow/providers/standard/operators/trigger_dagrun.py +++ b/providers/standard/src/airflow/providers/standard/operators/trigger_dagrun.py @@ -154,6 +154,9 @@ class TriggerDagRunOperator(BaseOperator): Airflow 3.x this requires Airflow 3.2.0+ (it relies on the task-SDK DAG state endpoint added then); on Airflow 3.0/3.1 setting this raises ``NotImplementedError``. :param deferrable: If waiting for completion, whether to defer the task until done, default is ``False``. + :param durable: If ``True`` and waiting for completion synchronously (non-deferrable), persist the + triggered run id before polling so that a worker crash mid-wait reconnects to the in-flight run on + retry instead of triggering a duplicate. Requires Airflow 3.3+ (task_state_store). Default ``False``. :param openlineage_inject_parent_info: whether to include OpenLineage metadata about the parent task in the triggered DAG run's conf, enabling improved lineage tracking. The metadata is only injected if OpenLineage is enabled and running. This option does not modify any other part of the conf, @@ -198,6 +201,7 @@ def __init__( fail_when_dag_is_paused: bool = False, note: str | None = None, deferrable: bool = conf.getboolean("operators", "default_deferrable", fallback=False), + durable: bool = False, openlineage_inject_parent_info: bool = True, **kwargs, ) -> None: @@ -221,6 +225,7 @@ def __init__( self.openlineage_inject_parent_info = openlineage_inject_parent_info self.note = note self.deferrable = deferrable + self.durable = durable run_after = _validate_datetime_param("run_after", run_after) self.logical_date = logical_date self.run_after = run_after @@ -325,6 +330,9 @@ def _trigger_dag_af_3(self, context, run_id, parsed_logical_date, parsed_run_aft if parsed_run_after and "run_after" in parameters: kwargs_accepted["run_after"] = parsed_run_after + if self.durable and "durable" in parameters: + kwargs_accepted["durable"] = self.durable + if isinstance(context, Mapping): from airflow.utils import helpers diff --git a/task-sdk/src/airflow/sdk/bases/resumablejobmixin.py b/task-sdk/src/airflow/sdk/bases/resumablejobmixin.py index ef1629da97503..2625989932dd6 100644 --- a/task-sdk/src/airflow/sdk/bases/resumablejobmixin.py +++ b/task-sdk/src/airflow/sdk/bases/resumablejobmixin.py @@ -25,7 +25,10 @@ from airflow.sdk.bases.operator import BaseOperatorMeta if TYPE_CHECKING: + from collections.abc import Callable + from pydantic import JsonValue + from structlog.typing import FilteringBoundLogger from airflow.sdk.definitions.context import Context from airflow.sdk.types import Logger @@ -102,113 +105,34 @@ def __init__(self, *, durable: bool = True, **kwargs: Any) -> None: def execute_resumable(self, context: Context) -> Any: """ - Core of the resumable execution logic. Call this from execute() when reconnection is supported. - - On initial run: submits the job, persists the external ID to task_state_store, then polls. + Crash-safe submit-and-poll for synchronous operators. Call from ``execute()``. - Behaviour on retry: - - On retry with active job: skips submission, reconnects to the running job. - - On retry with succeeded job: skips submission and polling, returns result immediately. - - On retry with failed job: falls through and resubmits fresh. - - Known limitation: there is a small window between ``submit_job`` returning and - ``task_state_store.set`` completing. If the worker dies in that gap, the next retry still - holds the previous (terminal) ID and will resubmit a fresh job rather than reconnecting. - Closing this window would require atomic "submit + persist", which is not possible across - an external system boundary. + Binds the operator's ``submit_job`` / ``get_job_status`` / … methods to + :func:`resume_or_submit`, which owns the persist + three-state reconnect logic. See that + function for the behaviour and its known submit-vs-persist limitation. """ - if not self.durable: - external_id = self.submit_job(context) - self.poll_until_complete(external_id, context) - return self.get_job_result(external_id, context) - stats_tags = {"operator": type(self).__name__} # The task is team-scoped in multi-team deployments; surface team_name on the # resumable_job metrics via the running task instance's stats tags (omitted when # not multi-team or the task has no team). - ti = context.get("ti") + ti = context.get("ti") if context is not None else None if ti is not None and (team_name := ti.stats_tags.get("team_name")): stats_tags["team_name"] = team_name - reconnect_to: Any = None - already_succeeded_id: Any = None - - with tracer.start_as_current_span("resumable_job.resume_decision") as span: - span.set_attribute("operator", type(self).__name__) - span.set_attribute("resumable.external_id_key", self.external_id_key) - - task_state_store = context.get("task_state_store") - - if task_state_store is None: - span.set_attribute("resumable.decision", "no_task_state_store") - self.log.warning( - "task_state_store not available in context, crash recovery is disabled for this run" - ) - else: - external_id = task_state_store.get(self.external_id_key) - if external_id: - stats.incr("resumable_job.reconnect_attempt", tags=stats_tags) - - status = self.get_job_status(external_id, context) - - span.set_attribute("resumable.external_id", str(external_id)) - span.set_attribute("resumable.prior_status", status) - - if self.is_job_active(status): - # Job is still running, skip submission and reconnect to it. - span.set_attribute("resumable.decision", "reconnect") - stats.incr("resumable_job.reconnect_success", tags=stats_tags) - self.log.info( - "Reconnecting to existing job", - external_id_key=self.external_id_key, - external_id=external_id, - status=status, - ) - reconnect_to = external_id - elif self.is_job_succeeded(status): - # Job already finished successfully, skip polling and return result directly. - span.set_attribute("resumable.decision", "already_succeeded") - stats.incr("resumable_job.already_succeeded", tags=stats_tags) - self.log.info( - "Job already completed successfully, skipping resubmission", - external_id_key=self.external_id_key, - external_id=external_id, - ) - already_succeeded_id = external_id - else: - # Job is in a terminal failed state, fall through and submit a new job. - span.set_attribute("resumable.decision", "terminal_resubmit") - stats.incr("resumable_job.terminal_resubmit", tags=stats_tags) - self.log.warning( - "Prior job in terminal state, resubmitting fresh", - external_id_key=self.external_id_key, - external_id=external_id, - status=status, - ) - else: - span.set_attribute("resumable.decision", "fresh_submit") - stats.incr("resumable_job.fresh_submit", tags=stats_tags) - self.log.debug( - "No stored external ID found; submitting fresh job", - external_id_key=self.external_id_key, - ) - - if reconnect_to is not None: - return self.poll_until_complete(reconnect_to, context) - if already_succeeded_id is not None: - return self.get_job_result(already_succeeded_id, context) - external_id = self.submit_job(context) - - if task_state_store is not None and external_id is not None: - task_state_store.set(self.external_id_key, external_id) - self.log.debug( - "Persisted external ID to task store", - external_id_key=self.external_id_key, - external_id=external_id, - ) - - self.poll_until_complete(external_id, context) - return self.get_job_result(external_id, context) + return resume_or_submit( + durable=self.durable, + external_id_key=self.external_id_key, + task_state_store=context.get("task_state_store") if context is not None else None, + submit=lambda: self.submit_job(context), + get_status=lambda external_id: self.get_job_status(external_id, context), + is_active=self.is_job_active, + is_succeeded=self.is_job_succeeded, + poll=lambda external_id: self.poll_until_complete(external_id, context), + get_result=lambda external_id: self.get_job_result(external_id, context), + log=self.log, + operator_name=type(self).__name__, + stats_tags=stats_tags, + ) @abstractmethod def submit_job(self, context: Context) -> JsonValue: @@ -261,3 +185,114 @@ def poll_until_complete(self, external_id: JsonValue, context: Context) -> None: def get_job_result(self, external_id: JsonValue, context: Context) -> Any: """Return the job result after completion. Return None if not applicable.""" raise NotImplementedError + + +def resume_or_submit( + *, + durable: bool, + external_id_key: str, + task_state_store: Any, + submit: Callable[[], JsonValue], + get_status: Callable[[JsonValue], str], + is_active: Callable[[str], bool], + is_succeeded: Callable[[str], bool], + poll: Callable[[JsonValue], Any], + get_result: Callable[[JsonValue], Any], + log: FilteringBoundLogger, + operator_name: str, + stats_tags: dict[str, str], +) -> Any: + """ + Submit an external job and poll for it, reconnecting to a prior run on retry. + + The reusable core of crash-safe submit-and-poll execution, independent of any operator. + ``ResumableJobMixin`` wraps it for synchronous operators, but it can equally be driven from the + task runner for operators whose wait happens outside ``execute()`` (e.g. ``TriggerDagRunOperator``, + which raises an exception and is polled by the runner). + + The callbacks are the external-system bindings: ``submit`` starts the job and returns its external + id, ``get_status`` reads the raw backend status, ``is_active`` / ``is_succeeded`` classify it, + ``poll`` blocks until terminal (raising on failure), ``get_result`` returns the result. On the + first run the external id is persisted to ``task_state_store`` before polling; on retry it is read + back and the job is reconnected (active), returned (succeeded), or resubmitted (terminal / missing). + + Known limitation: there is a small window between ``submit`` returning and the persist completing. + A crash in that gap resubmits fresh on the next retry rather than reconnecting; closing it would + require an atomic "submit + persist", which is not possible across an external system boundary. + """ + if not durable: + external_id = submit() + poll(external_id) + return get_result(external_id) + + reconnect_to: Any = None + already_succeeded_id: Any = None + + with tracer.start_as_current_span("resumable_job.resume_decision") as span: + span.set_attribute("operator", operator_name) + span.set_attribute("resumable.external_id_key", external_id_key) + + if task_state_store is None: + span.set_attribute("resumable.decision", "no_task_state_store") + log.warning("task_state_store not available in context, crash recovery is disabled for this run") + else: + external_id = task_state_store.get(external_id_key) + if external_id: + stats.incr("resumable_job.reconnect_attempt", tags=stats_tags) + + status = get_status(external_id) + + span.set_attribute("resumable.external_id", str(external_id)) + span.set_attribute("resumable.prior_status", status) + + if is_active(status): + span.set_attribute("resumable.decision", "reconnect") + stats.incr("resumable_job.reconnect_success", tags=stats_tags) + log.info( + "Reconnecting to existing job", + external_id_key=external_id_key, + external_id=external_id, + status=status, + ) + reconnect_to = external_id + elif is_succeeded(status): + span.set_attribute("resumable.decision", "already_succeeded") + stats.incr("resumable_job.already_succeeded", tags=stats_tags) + log.info( + "Job already completed successfully, skipping resubmission", + external_id_key=external_id_key, + external_id=external_id, + ) + already_succeeded_id = external_id + else: + span.set_attribute("resumable.decision", "terminal_resubmit") + stats.incr("resumable_job.terminal_resubmit", tags=stats_tags) + log.warning( + "Prior job in terminal state, resubmitting fresh", + external_id_key=external_id_key, + external_id=external_id, + status=status, + ) + else: + span.set_attribute("resumable.decision", "fresh_submit") + stats.incr("resumable_job.fresh_submit", tags=stats_tags) + log.debug( + "No stored external ID found; submitting fresh job", + external_id_key=external_id_key, + ) + + if reconnect_to is not None: + return poll(reconnect_to) + if already_succeeded_id is not None: + return get_result(already_succeeded_id) + + external_id = submit() + if task_state_store is not None and external_id is not None: + task_state_store.set(external_id_key, external_id) + log.debug( + "Persisted external ID to task store", + external_id_key=external_id_key, + external_id=external_id, + ) + poll(external_id) + return get_result(external_id) diff --git a/task-sdk/src/airflow/sdk/exceptions.py b/task-sdk/src/airflow/sdk/exceptions.py index 6f43d5421ecf2..c2aac7d3a3c54 100644 --- a/task-sdk/src/airflow/sdk/exceptions.py +++ b/task-sdk/src/airflow/sdk/exceptions.py @@ -309,6 +309,7 @@ def __init__( failed_states: list[str], poke_interval: int, deferrable: bool, + durable: bool = False, note: str | None = None, ): super().__init__() @@ -324,6 +325,7 @@ def __init__( self.failed_states = failed_states self.poke_interval = poke_interval self.deferrable = deferrable + self.durable = durable self.note = note diff --git a/task-sdk/src/airflow/sdk/execution_time/task_runner.py b/task-sdk/src/airflow/sdk/execution_time/task_runner.py index 1ce848128d663..2cb76967c78ee 100644 --- a/task-sdk/src/airflow/sdk/execution_time/task_runner.py +++ b/task-sdk/src/airflow/sdk/execution_time/task_runner.py @@ -58,6 +58,7 @@ TIRunContext, ) from airflow.sdk.bases.operator import BaseOperator, ExecutorSafeguard +from airflow.sdk.bases.resumablejobmixin import resume_or_submit from airflow.sdk.bases.xcom import BaseXCom from airflow.sdk.configuration import conf from airflow.sdk.definitions._internal.dag_parsing_context import _airflow_parsing_context_manager @@ -1939,10 +1940,157 @@ def _finalize_task_failure( ) +_TRIGGERED_RUN_ID_KEY = "triggered_dag_run_id" + + +_RUN_NOT_FOUND = "__run_not_found__" + + +class _TriggeredRunAlreadyExists(Exception): + """Internal signal from the durable submit callback that the triggered run already exists.""" + + def __init__(self, *, skip: bool) -> None: + super().__init__() + self.skip = skip + + +class _TriggeredRunFailed(Exception): + """Internal signal that the triggered run reached a failed state during the durable wait.""" + + +def _handle_durable_trigger_dag_run( + drte: DagRunTriggerException, + context: Context, + ti: RuntimeTaskInstance, + log: Logger, + task_state_store: TaskStateStoreAccessor, +) -> tuple[ToSupervisor, TaskInstanceState]: + """ + Crash-safe synchronous wait for the triggered run, reusing the shared resumable core. + + The trigger and wait happen in the runner (the operator only raises ``DagRunTriggerException``), + so ``resume_or_submit`` is driven by runner callbacks rather than subclassing ``ResumableJobMixin``: + the triggered run id is persisted before polling, and on retry the runner reconnects to the + in-flight run instead of triggering a duplicate. + """ + + def submit() -> str: + comms_msg = SUPERVISOR_COMMS.send( + TriggerDagRun( + dag_id=drte.trigger_dag_id, + run_id=drte.dag_run_id, + logical_date=drte.logical_date, + run_after=drte.run_after, + conf=drte.conf, + reset_dag_run=drte.reset_dag_run, + note=drte.note, + ), + ) + if isinstance(comms_msg, ErrorResponse) and comms_msg.error == ErrorType.DAGRUN_ALREADY_EXISTS: + raise _TriggeredRunAlreadyExists(skip=drte.skip_when_already_exists) + log.info("Dag Run triggered successfully.", trigger_dag_id=drte.trigger_dag_id) + return drte.dag_run_id + + def get_status(run_id: JsonValue) -> str: + comms_msg = SUPERVISOR_COMMS.send(GetDagRunState(dag_id=drte.trigger_dag_id, run_id=str(run_id))) + return comms_msg.state if isinstance(comms_msg, DagRunStateResult) else _RUN_NOT_FOUND + + def wait(run_id: JsonValue) -> None: + run_id = str(run_id) + ti.xcom_push(key="trigger_run_id", value=run_id) + while True: + log.info( + "Waiting for dag run to complete execution in allowed state.", + dag_id=drte.trigger_dag_id, + run_id=run_id, + allowed_state=drte.allowed_states, + ) + time.sleep(drte.poke_interval) + comms_msg = SUPERVISOR_COMMS.send(GetDagRunState(dag_id=drte.trigger_dag_id, run_id=run_id)) + if TYPE_CHECKING: + assert isinstance(comms_msg, DagRunStateResult) + if comms_msg.state in drte.failed_states: + log.error( + "DagRun finished with failed state.", dag_id=drte.trigger_dag_id, state=comms_msg.state + ) + raise _TriggeredRunFailed(f"{drte.trigger_dag_id} failed with failed state {comms_msg.state}") + if comms_msg.state in drte.allowed_states: + log.info( + "DagRun finished with allowed state.", dag_id=drte.trigger_dag_id, state=comms_msg.state + ) + return + log.debug( + "DagRun not yet in allowed or failed state.", + dag_id=drte.trigger_dag_id, + state=comms_msg.state, + ) + + def record_link(run_id: JsonValue) -> None: + ti.xcom_push(key="trigger_run_id", value=str(run_id)) + + stats_tags = {"operator": "TriggerDagRunOperator"} + if team_name := ti.stats_tags.get("team_name"): + stats_tags["team_name"] = team_name + + try: + resume_or_submit( + durable=True, + external_id_key=_TRIGGERED_RUN_ID_KEY, + task_state_store=task_state_store, + submit=submit, + get_status=get_status, + is_active=lambda state: ( + state not in drte.allowed_states + and state not in drte.failed_states + and state != _RUN_NOT_FOUND + ), + is_succeeded=lambda state: state in drte.allowed_states, + poll=wait, + get_result=record_link, + log=log, + operator_name="TriggerDagRunOperator", + stats_tags=stats_tags, + ) + except _TriggeredRunAlreadyExists as exc: + if exc.skip: + log.info( + "Dag Run already exists, skipping task as skip_when_already_exists is set to True.", + dag_id=drte.trigger_dag_id, + ) + return ( + TaskState( + state=TaskInstanceState.SKIPPED, + end_date=datetime.now(tz=timezone.utc), + rendered_map_index=ti.rendered_map_index, + ), + TaskInstanceState.SKIPPED, + ) + log.error("Dag Run already exists, marking task as failed.", dag_id=drte.trigger_dag_id) + return ( + TaskState( + state=TaskInstanceState.FAILED, + end_date=datetime.now(tz=timezone.utc), + rendered_map_index=ti.rendered_map_index, + ), + TaskInstanceState.FAILED, + ) + except _TriggeredRunFailed as exc: + # Mirror the deferrable path so a configured retry_policy is honoured: surface the failure + # as an AirflowException through run()'s handler (falling back to the retry-count check). + return _handle_current_task_failed(ti, AirflowException(str(exc)), log, context) + return _handle_current_task_success(context, ti) + + def _handle_trigger_dag_run( drte: DagRunTriggerException, context: Context, ti: RuntimeTaskInstance, log: Logger ) -> tuple[ToSupervisor, TaskInstanceState]: """Handle exception from TriggerDagRunOperator.""" + task_state_store = context.get("task_state_store") + if drte.durable and drte.wait_for_completion and not drte.deferrable and task_state_store is not None: + # Durable synchronous wait: reuse the shared resumable core so a worker crash mid-wait + # reconnects to the in-flight run on retry instead of triggering a duplicate. + return _handle_durable_trigger_dag_run(drte, context, ti, log, task_state_store) + log.info("Triggering Dag Run.", trigger_dag_id=drte.trigger_dag_id) comms_msg = SUPERVISOR_COMMS.send( TriggerDagRun( @@ -1981,8 +2129,7 @@ def _handle_trigger_dag_run( log.info("Dag Run triggered successfully.", trigger_dag_id=drte.trigger_dag_id) - # Store the run id from the dag run (either created or found above) to - # be used when creating the extra link on the webserver. + # Store the run id from the dag run for the webserver extra link. ti.xcom_push(key="trigger_run_id", value=drte.dag_run_id) if drte.wait_for_completion: diff --git a/task-sdk/tests/task_sdk/bases/test_resumablejobmixin.py b/task-sdk/tests/task_sdk/bases/test_resumablejobmixin.py index 1f84baa2d89e4..dd65c71dbaeb8 100644 --- a/task-sdk/tests/task_sdk/bases/test_resumablejobmixin.py +++ b/task-sdk/tests/task_sdk/bases/test_resumablejobmixin.py @@ -235,6 +235,17 @@ def test_default_is_true(self): assert op.durable is True +class TestAbsentContext: + """Operators are executed with a ``None`` context by provider tests asserting validation errors.""" + + @pytest.mark.parametrize("durable", [True, False]) + def test_submits_when_context_is_none(self, durable): + op = ConcreteResumableOperator(task_id="test_task", durable=durable) + + assert op.execute_resumable(None) == "result-of-job-001" + assert op.submitted_ids == ["job-001"] + + class TestExternalIdKey: def test_custom_key_used_for_storage_and_retrieval(self): class CustomKeyOp(ConcreteResumableOperator): diff --git a/task-sdk/tests/task_sdk/execution_time/test_task_runner.py b/task-sdk/tests/task_sdk/execution_time/test_task_runner.py index cc9fb77e08921..edd823ebf40a1 100644 --- a/task-sdk/tests/task_sdk/execution_time/test_task_runner.py +++ b/task-sdk/tests/task_sdk/execution_time/test_task_runner.py @@ -5283,6 +5283,19 @@ class CustomOperator(BaseOperator): assert log.exception.mock_calls == expected_exception_logs +class _FakeTaskStateStore: + """Minimal in-memory task_state_store stub for durable TriggerDagRunOperator tests.""" + + def __init__(self, initial: dict | None = None): + self._store = dict(initial or {}) + + def get(self, key, default=None): + return self._store.get(key, default) + + def set(self, key, value, *, retention=None): + self._store[key] = value + + class TestTriggerDagRunOperator: """Tests to verify various aspects of TriggerDagRunOperator""" @@ -5626,6 +5639,172 @@ def _send_side_effect(*args, **kwargs): assert state == TaskInstanceState.UP_FOR_RETRY + def test_handle_trigger_dag_run_persists_run_id_before_polling( + self, create_runtime_ti, mock_supervisor_comms + ): + """A durable synchronous wait persists the triggered run id before it starts polling.""" + from airflow.sdk.execution_time.task_runner import _TRIGGERED_RUN_ID_KEY + + task = TriggerDagRunOperator( + task_id="test_task", + trigger_dag_id="test_dag", + trigger_run_id="fresh_run_id", + poke_interval=5, + wait_for_completion=True, + deferrable=False, + durable=True, + ) + ti = create_runtime_ti(dag_id="test_persist", run_id="test_run", task=task) + store = _FakeTaskStateStore() + context = ti.get_template_context() + context["task_state_store"] = store + + persisted_at_first_poll = {} + + def _send(*args, **kwargs): + msg = kwargs.get("msg") or (args[0] if args else None) + if isinstance(msg, TriggerDagRun): + return OKResponse(ok=True) + if isinstance(msg, GetDagRunState): + persisted_at_first_poll.setdefault("value", store.get(_TRIGGERED_RUN_ID_KEY)) + return DagRunStateResult(state=DagRunState.SUCCESS) + return None + + mock_supervisor_comms.send.side_effect = _send + log = mock.MagicMock() + with mock.patch("time.sleep", return_value=None): + state, _, _ = run(ti, context, log) + + assert state == TaskInstanceState.SUCCESS + assert store.get(_TRIGGERED_RUN_ID_KEY) == "fresh_run_id" + # the id must already be persisted by the time the first poll runs, so a crash mid-wait reconnects + assert persisted_at_first_poll["value"] == "fresh_run_id" + + def test_handle_trigger_dag_run_reconnects_to_running_prior_run( + self, create_runtime_ti, mock_supervisor_comms + ): + """On retry, a still-running run triggered on a prior attempt is resumed, not re-triggered.""" + from airflow.sdk.execution_time.task_runner import _TRIGGERED_RUN_ID_KEY + + task = TriggerDagRunOperator( + task_id="test_task", + trigger_dag_id="test_dag", + trigger_run_id="new_run_id", + poke_interval=5, + wait_for_completion=True, + deferrable=False, + durable=True, + ) + ti = create_runtime_ti(dag_id="test_reconnect", run_id="test_run", task=task) + store = _FakeTaskStateStore({_TRIGGERED_RUN_ID_KEY: "prior_run_id"}) + context = ti.get_template_context() + context["task_state_store"] = store + + poll_states = iter([DagRunState.RUNNING, DagRunState.SUCCESS]) + triggered, polls = [], [] + + def _send(*args, **kwargs): + msg = kwargs.get("msg") or (args[0] if args else None) + if isinstance(msg, TriggerDagRun): + triggered.append(msg) + return OKResponse(ok=True) + if isinstance(msg, GetDagRunState): + polls.append(msg) + return DagRunStateResult(state=next(poll_states)) + return None + + mock_supervisor_comms.send.side_effect = _send + log = mock.MagicMock() + with mock.patch("time.sleep", return_value=None): + state, _, _ = run(ti, context, log) + + assert state == TaskInstanceState.SUCCESS + assert triggered == [] # reconnected — never resubmitted + assert all(m.run_id == "prior_run_id" for m in polls) # polled the prior run, not the new run_id + + def test_handle_trigger_dag_run_returns_success_for_already_succeeded_prior_run( + self, create_runtime_ti, mock_supervisor_comms + ): + """On retry, a prior run that already succeeded short-circuits to success without re-triggering.""" + from airflow.sdk.execution_time.task_runner import _TRIGGERED_RUN_ID_KEY + + task = TriggerDagRunOperator( + task_id="test_task", + trigger_dag_id="test_dag", + trigger_run_id="new_run_id", + poke_interval=5, + wait_for_completion=True, + deferrable=False, + durable=True, + ) + ti = create_runtime_ti(dag_id="test_already_succeeded", run_id="test_run", task=task) + store = _FakeTaskStateStore({_TRIGGERED_RUN_ID_KEY: "prior_run_id"}) + context = ti.get_template_context() + context["task_state_store"] = store + + triggered, polls = [], [] + + def _send(*args, **kwargs): + msg = kwargs.get("msg") or (args[0] if args else None) + if isinstance(msg, TriggerDagRun): + triggered.append(msg) + return OKResponse(ok=True) + if isinstance(msg, GetDagRunState): + polls.append(msg) + return DagRunStateResult(state=DagRunState.SUCCESS) + return None + + mock_supervisor_comms.send.side_effect = _send + log = mock.MagicMock() + state, _, _ = run(ti, context, log) + + assert state == TaskInstanceState.SUCCESS + assert triggered == [] # never resubmitted + assert len(polls) == 1 # only the resume check, no poll loop + assert polls[0].run_id == "prior_run_id" + + def test_handle_trigger_dag_run_resubmits_after_failed_prior_run( + self, create_runtime_ti, mock_supervisor_comms + ): + """On retry, a prior run in a failed state triggers a fresh run rather than reconnecting.""" + from airflow.sdk.execution_time.task_runner import _TRIGGERED_RUN_ID_KEY + + task = TriggerDagRunOperator( + task_id="test_task", + trigger_dag_id="test_dag", + trigger_run_id="fresh_run_id", + poke_interval=5, + wait_for_completion=True, + deferrable=False, + durable=True, + ) + ti = create_runtime_ti(dag_id="test_resubmit", run_id="test_run", task=task) + store = _FakeTaskStateStore({_TRIGGERED_RUN_ID_KEY: "dead_run_id"}) + context = ti.get_template_context() + context["task_state_store"] = store + + poll_states = iter([DagRunState.FAILED, DagRunState.SUCCESS]) + triggered = [] + + def _send(*args, **kwargs): + msg = kwargs.get("msg") or (args[0] if args else None) + if isinstance(msg, TriggerDagRun): + triggered.append(msg) + return OKResponse(ok=True) + if isinstance(msg, GetDagRunState): + return DagRunStateResult(state=next(poll_states)) + return None + + mock_supervisor_comms.send.side_effect = _send + log = mock.MagicMock() + with mock.patch("time.sleep", return_value=None): + state, _, _ = run(ti, context, log) + + assert state == TaskInstanceState.SUCCESS + assert len(triggered) == 1 # a fresh run was triggered after the prior one failed + assert triggered[0].run_id == "fresh_run_id" + assert store.get(_TRIGGERED_RUN_ID_KEY) == "fresh_run_id" # store overwritten with the new run + def test_handle_trigger_dag_run_wait_for_completion_failed_state_retry_policy_fail( self, create_runtime_ti, mock_supervisor_comms ):