From 48c7c6719038f1cb55e6081003155167402599b7 Mon Sep 17 00:00:00 2001 From: 1fanwang <1fannnw@gmail.com> Date: Wed, 24 Jun 2026 02:02:26 -0700 Subject: [PATCH 1/8] Add durable option to TriggerDagRunOperator to reconnect on retry With wait_for_completion the trigger-and-wait runs in the task runner. A worker crash while polling makes the retry recompute a fresh run_id and trigger a duplicate child run (or fail with DagRunAlreadyExists), even though the run the first attempt started is healthy and still running. The opt-in durable flag persists the triggered run_id to task_state_store before polling, so the retry reconnects to the in-flight run instead of resubmitting. Signed-off-by: 1fanwang <1fannnw@gmail.com> --- .../standard/operators/trigger_dagrun.py | 8 + task-sdk/src/airflow/sdk/exceptions.py | 2 + .../airflow/sdk/execution_time/task_runner.py | 125 ++++++++---- .../execution_time/test_task_runner.py | 179 ++++++++++++++++++ 4 files changed, 275 insertions(+), 39 deletions(-) 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 60d96cbe3f495..e2946acf59875 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 logical_date = _validate_datetime_param("logical_date", logical_date) run_after = _validate_datetime_param("run_after", run_after) self.logical_date = logical_date @@ -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/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 31d0831bbadec..ab23edc3ae522 100644 --- a/task-sdk/src/airflow/sdk/execution_time/task_runner.py +++ b/task-sdk/src/airflow/sdk/execution_time/task_runner.py @@ -1871,51 +1871,100 @@ def _finalize_task_failure( ) +_TRIGGERED_RUN_ID_KEY = "triggered_dag_run_id" + + +def _evaluate_prior_triggered_run( + run_id: str, drte: DagRunTriggerException, log: Logger +) -> Literal["succeeded", "reconnect", "resubmit"]: + """ + Classify a run triggered on a prior attempt so the synchronous wait can resume safely. + + ``"succeeded"`` — already finished in an allowed state; skip resubmission. + ``"reconnect"`` — still running; resume the wait without resubmitting. + ``"resubmit"`` — failed, gone, or state unreadable; trigger a fresh run. + """ + comms_msg = SUPERVISOR_COMMS.send(GetDagRunState(dag_id=drte.trigger_dag_id, run_id=run_id)) + state = comms_msg.state if isinstance(comms_msg, DagRunStateResult) else None + if state in drte.allowed_states: + log.info("Run triggered on a prior attempt already succeeded; not resubmitting.", run_id=run_id) + return "succeeded" + if state is None or state in drte.failed_states: + log.warning( + "Run triggered on a prior attempt is not resumable; resubmitting.", run_id=run_id, state=state + ) + return "resubmit" + log.info("Reconnecting to run triggered on a prior attempt.", run_id=run_id, state=state) + return "reconnect" + + def _handle_trigger_dag_run( drte: DagRunTriggerException, context: Context, ti: RuntimeTaskInstance, log: Logger ) -> tuple[ToSupervisor, TaskInstanceState]: """Handle exception from TriggerDagRunOperator.""" - log.info("Triggering Dag Run.", trigger_dag_id=drte.trigger_dag_id) - 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, - ), + # Crash-safety for the synchronous wait: persist the triggered run id before polling so a worker + # crash mid-wait reconnects to the in-flight run on retry instead of triggering a duplicate. + task_state_store = context.get("task_state_store") + durable = ( + drte.durable and drte.wait_for_completion and not drte.deferrable and task_state_store is not None ) - if isinstance(comms_msg, ErrorResponse) and comms_msg.error == ErrorType.DAGRUN_ALREADY_EXISTS: - if drte.skip_when_already_exists: - log.info( - "Dag Run already exists, skipping task as skip_when_already_exists is set to True.", + run_id = drte.dag_run_id + reconnecting = False + if durable and (stored_run_id := task_state_store.get(_TRIGGERED_RUN_ID_KEY)): + decision = _evaluate_prior_triggered_run(stored_run_id, drte, log) + if decision == "succeeded": + ti.xcom_push(key="trigger_run_id", value=stored_run_id) + return _handle_current_task_success(context, ti) + if decision == "reconnect": + run_id = stored_run_id + reconnecting = True + + if not reconnecting: + log.info("Triggering Dag Run.", trigger_dag_id=drte.trigger_dag_id) + comms_msg = SUPERVISOR_COMMS.send( + TriggerDagRun( dag_id=drte.trigger_dag_id, - ) - msg = TaskState( - state=TaskInstanceState.SKIPPED, - end_date=datetime.now(tz=timezone.utc), - rendered_map_index=ti.rendered_map_index, - ) - state = TaskInstanceState.SKIPPED - else: - log.error("Dag Run already exists, marking task as failed.", dag_id=drte.trigger_dag_id) - msg = TaskState( - state=TaskInstanceState.FAILED, - end_date=datetime.now(tz=timezone.utc), - rendered_map_index=ti.rendered_map_index, - ) - state = TaskInstanceState.FAILED + run_id=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, + ), + ) - return msg, state + if isinstance(comms_msg, ErrorResponse) and comms_msg.error == ErrorType.DAGRUN_ALREADY_EXISTS: + if drte.skip_when_already_exists: + log.info( + "Dag Run already exists, skipping task as skip_when_already_exists is set to True.", + dag_id=drte.trigger_dag_id, + ) + msg = TaskState( + state=TaskInstanceState.SKIPPED, + end_date=datetime.now(tz=timezone.utc), + rendered_map_index=ti.rendered_map_index, + ) + state = TaskInstanceState.SKIPPED + else: + log.error("Dag Run already exists, marking task as failed.", dag_id=drte.trigger_dag_id) + msg = TaskState( + state=TaskInstanceState.FAILED, + end_date=datetime.now(tz=timezone.utc), + rendered_map_index=ti.rendered_map_index, + ) + state = TaskInstanceState.FAILED + + return msg, state - log.info("Dag Run triggered successfully.", trigger_dag_id=drte.trigger_dag_id) + 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. - ti.xcom_push(key="trigger_run_id", value=drte.dag_run_id) + if durable: + # Persist before polling so a crash mid-wait reconnects on retry instead of resubmitting. + task_state_store.set(_TRIGGERED_RUN_ID_KEY, run_id) + + # Store the run id (created above or reconnected to) for the webserver extra link. + ti.xcom_push(key="trigger_run_id", value=run_id) if drte.wait_for_completion: if drte.deferrable: @@ -1940,14 +1989,12 @@ def _handle_trigger_dag_run( log.info( "Waiting for dag run to complete execution in allowed state.", dag_id=drte.trigger_dag_id, - run_id=drte.dag_run_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=drte.dag_run_id) - ) + 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: 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 0031974b86646..444d9899cdb57 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 @@ -5058,6 +5058,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""" @@ -5327,6 +5340,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 ): From b0772fee9efd2e39559a13e172e5a7a11eabc43d Mon Sep 17 00:00:00 2001 From: 1fanwang <1fannnw@gmail.com> Date: Wed, 24 Jun 2026 02:06:10 -0700 Subject: [PATCH 2/8] Add newsfragment for TriggerDagRunOperator durable flag Signed-off-by: 1fanwang <1fannnw@gmail.com> --- airflow-core/newsfragments/68936.feature.rst | 3 +++ 1 file changed, 3 insertions(+) create mode 100644 airflow-core/newsfragments/68936.feature.rst diff --git a/airflow-core/newsfragments/68936.feature.rst b/airflow-core/newsfragments/68936.feature.rst new file mode 100644 index 0000000000000..fed2dd43f93aa --- /dev/null +++ b/airflow-core/newsfragments/68936.feature.rst @@ -0,0 +1,3 @@ +``TriggerDagRunOperator`` gains an opt-in ``durable`` flag. When waiting synchronously for the +triggered run to complete, the run id is persisted before polling so that a worker crash mid-wait +reconnects to the in-flight run on retry instead of triggering a duplicate run. From 10248f8e5a421053a54d0e337422c65d444a5d5a Mon Sep 17 00:00:00 2001 From: 1fanwang <1fannnw@gmail.com> Date: Wed, 24 Jun 2026 10:16:26 -0700 Subject: [PATCH 3/8] Narrow task_state_store and coerce stored run id for mypy Signed-off-by: 1fanwang <1fannnw@gmail.com> --- .../airflow/sdk/execution_time/task_runner.py | 25 ++++++++++--------- 1 file changed, 13 insertions(+), 12 deletions(-) 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 ab23edc3ae522..48deb7dfad011 100644 --- a/task-sdk/src/airflow/sdk/execution_time/task_runner.py +++ b/task-sdk/src/airflow/sdk/execution_time/task_runner.py @@ -1905,20 +1905,21 @@ def _handle_trigger_dag_run( # Crash-safety for the synchronous wait: persist the triggered run id before polling so a worker # crash mid-wait reconnects to the in-flight run on retry instead of triggering a duplicate. task_state_store = context.get("task_state_store") - durable = ( - drte.durable and drte.wait_for_completion and not drte.deferrable and task_state_store is not None - ) + durable = drte.durable and drte.wait_for_completion and not drte.deferrable run_id = drte.dag_run_id reconnecting = False - if durable and (stored_run_id := task_state_store.get(_TRIGGERED_RUN_ID_KEY)): - decision = _evaluate_prior_triggered_run(stored_run_id, drte, log) - if decision == "succeeded": - ti.xcom_push(key="trigger_run_id", value=stored_run_id) - return _handle_current_task_success(context, ti) - if decision == "reconnect": - run_id = stored_run_id - reconnecting = True + if durable and task_state_store is not None: + stored_run_id = task_state_store.get(_TRIGGERED_RUN_ID_KEY) + if stored_run_id is not None: + prior_run_id = str(stored_run_id) + decision = _evaluate_prior_triggered_run(prior_run_id, drte, log) + if decision == "succeeded": + ti.xcom_push(key="trigger_run_id", value=prior_run_id) + return _handle_current_task_success(context, ti) + if decision == "reconnect": + run_id = prior_run_id + reconnecting = True if not reconnecting: log.info("Triggering Dag Run.", trigger_dag_id=drte.trigger_dag_id) @@ -1959,7 +1960,7 @@ def _handle_trigger_dag_run( log.info("Dag Run triggered successfully.", trigger_dag_id=drte.trigger_dag_id) - if durable: + if durable and task_state_store is not None: # Persist before polling so a crash mid-wait reconnects on retry instead of resubmitting. task_state_store.set(_TRIGGERED_RUN_ID_KEY, run_id) From f695b816b79a19b1ff224adeb08f68aac42dad84 Mon Sep 17 00:00:00 2001 From: 1fanwang <1fannnw@gmail.com> Date: Wed, 24 Jun 2026 11:09:10 -0700 Subject: [PATCH 4/8] Make the durable TriggerDagRunOperator newsfragment a single line Signed-off-by: 1fanwang <1fannnw@gmail.com> --- airflow-core/newsfragments/68936.feature.rst | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/airflow-core/newsfragments/68936.feature.rst b/airflow-core/newsfragments/68936.feature.rst index fed2dd43f93aa..c429bf95c571f 100644 --- a/airflow-core/newsfragments/68936.feature.rst +++ b/airflow-core/newsfragments/68936.feature.rst @@ -1,3 +1 @@ -``TriggerDagRunOperator`` gains an opt-in ``durable`` flag. When waiting synchronously for the -triggered run to complete, the run id is persisted before polling so that a worker crash mid-wait -reconnects to the in-flight run on retry instead of triggering a duplicate run. +``TriggerDagRunOperator`` gains an opt-in ``durable`` flag that, on a synchronous ``wait_for_completion``, persists the triggered run id before polling so a worker crash mid-wait reconnects to the in-flight run on retry instead of triggering a duplicate run. From 878942565bb8155e8ada9ac6b7e3ca44e04d6807 Mon Sep 17 00:00:00 2001 From: 1fanwang <1fannnw@gmail.com> Date: Wed, 24 Jun 2026 12:08:14 -0700 Subject: [PATCH 5/8] Make ResumableJobMixin's reconnect core reusable outside execute() TriggerDagRunOperator's durable wait lives in the task runner (it raises DagRunTriggerException and is polled there), not in execute(), so it cannot use ResumableJobMixin and re-implements the persist-and-reconnect logic by hand. Lift the mixin's core into a standalone resume_or_submit() so the same implementation can be driven from the runner now and the triggerer later, instead of duplicated per integration point. Signed-off-by: 1fanwang <1fannnw@gmail.com> --- .../airflow/sdk/bases/resumablejobmixin.py | 229 ++++++++++-------- 1 file changed, 132 insertions(+), 97 deletions(-) diff --git a/task-sdk/src/airflow/sdk/bases/resumablejobmixin.py b/task-sdk/src/airflow/sdk/bases/resumablejobmixin.py index ef1629da97503..92fbb1759c18a 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,26 +105,12 @@ 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 @@ -130,85 +119,20 @@ def execute_resumable(self, context: Context) -> Any: 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"), + 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) From 8d81900237ee8f99cd7373fa1ff51612897070f0 Mon Sep 17 00:00:00 2001 From: 1fanwang <1fannnw@gmail.com> Date: Wed, 24 Jun 2026 13:13:34 -0700 Subject: [PATCH 6/8] Reuse the shared resumable core for TriggerDagRunOperator's durable wait The durable synchronous wait previously hand-rolled the persist-and-reconnect logic in the task runner, duplicating what ResumableJobMixin already implements. Drive the shared resume_or_submit core with runner callbacks instead, so the durability primitive has one implementation across the operator (mixin) and the runner. Signed-off-by: 1fanwang <1fannnw@gmail.com> --- .../airflow/sdk/execution_time/task_runner.py | 239 +++++++++++++----- 1 file changed, 169 insertions(+), 70 deletions(-) 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 48deb7dfad011..7bb74a16b0d20 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 @@ -1874,59 +1875,42 @@ def _finalize_task_failure( _TRIGGERED_RUN_ID_KEY = "triggered_dag_run_id" -def _evaluate_prior_triggered_run( - run_id: str, drte: DagRunTriggerException, log: Logger -) -> Literal["succeeded", "reconnect", "resubmit"]: - """ - Classify a run triggered on a prior attempt so the synchronous wait can resume safely. +_RUN_NOT_FOUND = "__run_not_found__" - ``"succeeded"`` — already finished in an allowed state; skip resubmission. - ``"reconnect"`` — still running; resume the wait without resubmitting. - ``"resubmit"`` — failed, gone, or state unreadable; trigger a fresh run. - """ - comms_msg = SUPERVISOR_COMMS.send(GetDagRunState(dag_id=drte.trigger_dag_id, run_id=run_id)) - state = comms_msg.state if isinstance(comms_msg, DagRunStateResult) else None - if state in drte.allowed_states: - log.info("Run triggered on a prior attempt already succeeded; not resubmitting.", run_id=run_id) - return "succeeded" - if state is None or state in drte.failed_states: - log.warning( - "Run triggered on a prior attempt is not resumable; resubmitting.", run_id=run_id, state=state - ) - return "resubmit" - log.info("Reconnecting to run triggered on a prior attempt.", run_id=run_id, state=state) - return "reconnect" +class _TriggeredRunAlreadyExists(Exception): + """Internal signal from the durable submit callback that the triggered run already exists.""" -def _handle_trigger_dag_run( - drte: DagRunTriggerException, context: Context, ti: RuntimeTaskInstance, log: Logger + 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]: - """Handle exception from TriggerDagRunOperator.""" - # Crash-safety for the synchronous wait: persist the triggered run id before polling so a worker - # crash mid-wait reconnects to the in-flight run on retry instead of triggering a duplicate. - task_state_store = context.get("task_state_store") - durable = drte.durable and drte.wait_for_completion and not drte.deferrable - - run_id = drte.dag_run_id - reconnecting = False - if durable and task_state_store is not None: - stored_run_id = task_state_store.get(_TRIGGERED_RUN_ID_KEY) - if stored_run_id is not None: - prior_run_id = str(stored_run_id) - decision = _evaluate_prior_triggered_run(prior_run_id, drte, log) - if decision == "succeeded": - ti.xcom_push(key="trigger_run_id", value=prior_run_id) - return _handle_current_task_success(context, ti) - if decision == "reconnect": - run_id = prior_run_id - reconnecting = True - - if not reconnecting: - log.info("Triggering Dag Run.", trigger_dag_id=drte.trigger_dag_id) + """ + 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=run_id, + run_id=drte.dag_run_id, logical_date=drte.logical_date, run_after=drte.run_after, conf=drte.conf, @@ -1934,38 +1918,151 @@ def _handle_trigger_dag_run( note=drte.note, ), ) - if isinstance(comms_msg, ErrorResponse) and comms_msg.error == ErrorType.DAGRUN_ALREADY_EXISTS: - if drte.skip_when_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( - "Dag Run already exists, skipping task as skip_when_already_exists is set to True.", - dag_id=drte.trigger_dag_id, + "DagRun finished with allowed state.", dag_id=drte.trigger_dag_id, state=comms_msg.state ) - msg = TaskState( + 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, - ) - state = TaskInstanceState.SKIPPED - else: - log.error("Dag Run already exists, marking task as failed.", dag_id=drte.trigger_dag_id) - msg = TaskState( - state=TaskInstanceState.FAILED, - end_date=datetime.now(tz=timezone.utc), - rendered_map_index=ti.rendered_map_index, - ) - state = TaskInstanceState.FAILED + ), + 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) - return msg, state - log.info("Dag Run triggered successfully.", trigger_dag_id=drte.trigger_dag_id) +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( + 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 durable and task_state_store is not None: - # Persist before polling so a crash mid-wait reconnects on retry instead of resubmitting. - task_state_store.set(_TRIGGERED_RUN_ID_KEY, run_id) + if isinstance(comms_msg, ErrorResponse) and comms_msg.error == ErrorType.DAGRUN_ALREADY_EXISTS: + if drte.skip_when_already_exists: + log.info( + "Dag Run already exists, skipping task as skip_when_already_exists is set to True.", + dag_id=drte.trigger_dag_id, + ) + msg = TaskState( + state=TaskInstanceState.SKIPPED, + end_date=datetime.now(tz=timezone.utc), + rendered_map_index=ti.rendered_map_index, + ) + state = TaskInstanceState.SKIPPED + else: + log.error("Dag Run already exists, marking task as failed.", dag_id=drte.trigger_dag_id) + msg = TaskState( + state=TaskInstanceState.FAILED, + end_date=datetime.now(tz=timezone.utc), + rendered_map_index=ti.rendered_map_index, + ) + state = TaskInstanceState.FAILED - # Store the run id (created above or reconnected to) for the webserver extra link. - ti.xcom_push(key="trigger_run_id", value=run_id) + return msg, state + + log.info("Dag Run triggered successfully.", trigger_dag_id=drte.trigger_dag_id) + + # 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: if drte.deferrable: @@ -1990,12 +2087,14 @@ def _handle_trigger_dag_run( log.info( "Waiting for dag run to complete execution in allowed state.", dag_id=drte.trigger_dag_id, - run_id=run_id, + run_id=drte.dag_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)) + comms_msg = SUPERVISOR_COMMS.send( + GetDagRunState(dag_id=drte.trigger_dag_id, run_id=drte.dag_run_id) + ) if TYPE_CHECKING: assert isinstance(comms_msg, DagRunStateResult) if comms_msg.state in drte.failed_states: From 0472aafab578e624353843694432ef5df0dd7aae Mon Sep 17 00:00:00 2001 From: 1fanwang <1fannnw@gmail.com> Date: Mon, 29 Jun 2026 22:07:57 -0700 Subject: [PATCH 7/8] Remove newsfragment from stacked TriggerDagRunOperator PR The newsfragment number (68936) does not match this PR (68955), so check-newsfragment-pr-number fails. The durable-flag note belongs with the PR that lands the feature, not this stacked change. Signed-off-by: 1fanwang <1fannnw@gmail.com> --- airflow-core/newsfragments/68936.feature.rst | 1 - 1 file changed, 1 deletion(-) delete mode 100644 airflow-core/newsfragments/68936.feature.rst diff --git a/airflow-core/newsfragments/68936.feature.rst b/airflow-core/newsfragments/68936.feature.rst deleted file mode 100644 index c429bf95c571f..0000000000000 --- a/airflow-core/newsfragments/68936.feature.rst +++ /dev/null @@ -1 +0,0 @@ -``TriggerDagRunOperator`` gains an opt-in ``durable`` flag that, on a synchronous ``wait_for_completion``, persists the triggered run id before polling so a worker crash mid-wait reconnects to the in-flight run on retry instead of triggering a duplicate run. From c0f6ebf97fe87cb868a2468a4da6d41218d1caa0 Mon Sep 17 00:00:00 2001 From: 1fanwang <1fannnw@gmail.com> Date: Wed, 5 Aug 2026 04:16:36 -0700 Subject: [PATCH 8/8] Guard against an absent context in execute_resumable Folding the durable/non-durable split into resume_or_submit made the context lookups unconditional. Operators executed with a None context - the idiom provider tests use when asserting a validation error that is raised before any context is needed - then failed with an AttributeError instead of that error. Signed-off-by: 1fanwang <1fannnw@gmail.com> (cherry picked from commit 4e62cb149be58f901e0c34a1642ab13465f694d5) --- task-sdk/src/airflow/sdk/bases/resumablejobmixin.py | 4 ++-- .../tests/task_sdk/bases/test_resumablejobmixin.py | 11 +++++++++++ 2 files changed, 13 insertions(+), 2 deletions(-) diff --git a/task-sdk/src/airflow/sdk/bases/resumablejobmixin.py b/task-sdk/src/airflow/sdk/bases/resumablejobmixin.py index 92fbb1759c18a..2625989932dd6 100644 --- a/task-sdk/src/airflow/sdk/bases/resumablejobmixin.py +++ b/task-sdk/src/airflow/sdk/bases/resumablejobmixin.py @@ -115,14 +115,14 @@ def execute_resumable(self, context: Context) -> Any: # 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 return resume_or_submit( durable=self.durable, external_id_key=self.external_id_key, - task_state_store=context.get("task_state_store"), + 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, 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):