This is an automated email from the ASF dual-hosted git repository.

ashb pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git


The following commit(s) were added to refs/heads/main by this push:
     new e1f00413250 Delay retry reporting until task finalization completes 
(#73253)
e1f00413250 is described below

commit e1f004132507903ff3653db6fdf0f66196b5f936
Author: Ash Berlin-Taylor <[email protected]>
AuthorDate: Thu Sep 24 22:28:48 2026 +0100

    Delay retry reporting until task finalization completes (#73253)
    
    Reporting a retry to the API can replace the task instance UUID before
    on_retry_callback has run. The callback and other finalizers still use
    the old UUID, so their API calls can fail even though they are part of
    finishing that attempt.
    
    This changes things so we hold the RetryTask message until the worker
    exits. Heartbeats continue while finalization runs, and the existing
    overtime limit still applies. Server-directed termination discards the
    pending report.
    
    The report keeps the original end date, retry delay and reason. We don't
    add the retry delay again after finalization, but if that delay expires
    while callbacks are running, the next attempt now waits for them.
    
    Other outcomes are reported at the same point as before. In particular,
    success still reaches the server before callbacks run, so this doesn't
    delay downstream scheduling on success.
    
    | Outcome | Before | After |
    | --- | --- | --- |
    | Success | Before finalization | Unchanged |
    | Skipped | After worker exit | Unchanged |
    | Failed | After worker exit | Unchanged |
    | Retry | Before finalization | After worker exit |
    | Deferred / rescheduled / awaiting input | Immediately | Unchanged |
    
    The tests check that retry callbacks can still use the original attempt,
    including through dag.test(), and that callback errors don't prevent the
    retry report. They also cover heartbeats, the finalization timeout, and
    external state changes while the report is pending.
---
 .../versions/head/test_task_instances.py           |  34 +++-
 airflow-core/tests/unit/models/test_dag.py         |  42 +++++
 .../src/airflow/sdk/execution_time/supervisor.py   |  24 +--
 .../task_sdk/execution_time/test_supervisor.py     | 186 ++++++++++++++++++---
 4 files changed, 251 insertions(+), 35 deletions(-)

diff --git 
a/airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_task_instances.py
 
b/airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_task_instances.py
index 197db78d962..9485c0f874f 100644
--- 
a/airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_task_instances.py
+++ 
b/airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_task_instances.py
@@ -53,7 +53,7 @@ from airflow.models.asset import AssetActive, 
AssetAliasModel, AssetEvent, Asset
 from airflow.models.dag import DagModel
 from airflow.models.log import Log
 from airflow.models.task_state_store import TaskStateStoreModel
-from airflow.models.taskinstance import TaskInstance
+from airflow.models.taskinstance import TaskInstance, clear_task_instances
 from airflow.models.taskinstancehistory import TaskInstanceHistory
 from airflow.providers.standard.operators.empty import EmptyOperator
 from airflow.sdk import Asset, TaskGroup, TriggerRule, task, task_group
@@ -1317,6 +1317,38 @@ class TestTIRunState:
 
 
 class TestTIUpdateState:
+    @pytest.mark.parametrize(
+        ("interruption", "expected_state", "expected_status"),
+        [("clear", State.RESTARTING, 404), ("failed", State.FAILED, 409)],
+    )
+    def test_delayed_retry_report_preserves_external_state(
+        self, client, session, create_task_instance, interruption, 
expected_state, expected_status
+    ):
+        ti = create_task_instance(state=State.RUNNING, 
start_date=DEFAULT_START_DATE, session=session)
+        session.commit()
+        old_id = ti.id
+
+        if interruption == "clear":
+            clear_task_instances([ti], session=session)
+        else:
+            ti.set_state(State.FAILED, session=session)
+        session.commit()
+        expected_id = ti.id
+        expected_end_date = ti.end_date
+
+        response = client.patch(
+            f"/execution/task-instances/{old_id}/state",
+            json={"state": State.UP_FOR_RETRY, "end_date": 
DEFAULT_END_DATE.isoformat()},
+        )
+
+        assert response.status_code == expected_status
+        session.refresh(ti)
+        assert ti.id == expected_id
+        assert ti.state == expected_state
+        assert ti.end_date == expected_end_date
+        if interruption == "clear":
+            assert ti.id != old_id
+
     def setup_method(self):
         clear_db_assets()
         clear_db_logs()
diff --git a/airflow-core/tests/unit/models/test_dag.py 
b/airflow-core/tests/unit/models/test_dag.py
index b22227b695b..00b9a924faa 100644
--- a/airflow-core/tests/unit/models/test_dag.py
+++ b/airflow-core/tests/unit/models/test_dag.py
@@ -65,6 +65,7 @@ from airflow.models.deadline_alert import DeadlineAlert as 
DeadlineAlertModel
 from airflow.models.hitl import HITLDetail
 from airflow.models.serialized_dag import SerializedDagModel
 from airflow.models.taskinstance import TaskInstance as TI
+from airflow.models.taskinstancehistory import TaskInstanceHistory
 from airflow.models.trigger import handle_event_submit
 from airflow.providers.standard.operators.bash import BashOperator
 from airflow.providers.standard.operators.empty import EmptyOperator
@@ -1815,6 +1816,47 @@ class TestDag:
         dag.test()
         mock_object.assert_called_with("output of first task")
 
+    @pytest.mark.parametrize("callback_error", [None, RuntimeError])
+    def test_dag_test_retry_callback_keeps_attempt_live(self, 
testing_dag_bundle, callback_error):
+        observed = []
+
+        def on_retry(context):
+            ti = context["ti"]
+            ti.xcom_push(key="retry_callback", value="written")
+            value = ti.xcom_pull(task_ids=ti.task_id, key="retry_callback")
+            with create_session() as session:
+                stored_ti = session.get(TI, ti.id)
+                observed.append((ti.id, stored_ti.state if stored_ti else 
None, ti.end_date, value))
+            if callback_error:
+                raise callback_error("callback failed")
+
+        with DAG(dag_id="test_retry_callback_live_attempt", schedule=None, 
start_date=DEFAULT_DATE) as dag:
+
+            @task_decorator(retries=1, retry_delay=timedelta(0), 
on_retry_callback=on_retry)
+            def fail_once(**context):
+                if context["ti"].try_number == 1:
+                    raise RuntimeError("retry this attempt")
+
+            fail_once()
+        sync_dag_to_db(dag)
+
+        dr = dag.test()
+
+        assert dr.state == DagRunState.SUCCESS
+        assert len(observed) == 1
+        old_id, state_during_callback, end_date, value = observed[0]
+        assert state_during_callback == TaskInstanceState.RUNNING
+        assert value == "written"
+        with create_session() as session:
+            history = session.scalar(
+                
select(TaskInstanceHistory).where(TaskInstanceHistory.task_instance_id == 
old_id)
+            )
+            assert history is not None
+            assert history.end_date == end_date
+            ti = dr.get_task_instance("fail_once", session=session)
+            assert ti.id != old_id
+            assert ti.try_number == 2
+
     def test_dag_test_with_fail_handler(self, testing_dag_bundle):
         mock_handle_object_1 = mock.MagicMock()
         mock_handle_object_2 = mock.MagicMock()
diff --git a/task-sdk/src/airflow/sdk/execution_time/supervisor.py 
b/task-sdk/src/airflow/sdk/execution_time/supervisor.py
index 34863abb957..8a6d5e199f2 100644
--- a/task-sdk/src/airflow/sdk/execution_time/supervisor.py
+++ b/task-sdk/src/airflow/sdk/execution_time/supervisor.py
@@ -1732,7 +1732,7 @@ class ActivitySubprocess(WatchedSubprocess):
             return
 
         if self._pending_terminal_state_msg is not None:
-            if isinstance(self._pending_terminal_state_msg, TaskState):
+            if isinstance(self._pending_terminal_state_msg, (TaskState, 
RetryTask)):
                 self._send_terminal_state_msg(self._pending_terminal_state_msg)
             else:
                 self._replay_pending_terminal_state_msg()
@@ -1955,7 +1955,7 @@ class ActivitySubprocess(WatchedSubprocess):
 
         Not valid before the process has finished.
         """
-        if self._terminal_state == SERVER_TERMINATED:
+        if self._terminal_state in (SERVER_TERMINATED, 
TaskInstanceState.UP_FOR_RETRY):
             return self._terminal_state
         if self._exit_code == 0:
             return self._terminal_state or TaskInstanceState.SUCCESS
@@ -1977,7 +1977,9 @@ class ActivitySubprocess(WatchedSubprocess):
             log.debug("Received message from task runner", msg=msg)
         super()._handle_request(msg, log, req_id)
 
-    def _handle_task_state(self, msg: TaskState, log: FilteringBoundLogger, 
req_id: int) -> RequestResult:
+    def _handle_task_state(
+        self, msg: TaskState | RetryTask, log: FilteringBoundLogger, req_id: 
int
+    ) -> RequestResult:
         if self._terminal_state != SERVER_TERMINATED:
             self._terminal_state = msg.state
             self._pending_terminal_state_msg = msg
@@ -1986,7 +1988,7 @@ class ActivitySubprocess(WatchedSubprocess):
         return None, {}
 
     def _handle_finished_task(
-        self, msg: SucceedTask | RetryTask, log: FilteringBoundLogger, req_id: 
int
+        self, msg: SucceedTask, log: FilteringBoundLogger, req_id: int
     ) -> RequestResult:
         self._task_end_time_monotonic = time.monotonic()
         self._rendered_map_index = msg.rendered_map_index
@@ -2295,7 +2297,7 @@ class ActivitySubprocess(WatchedSubprocess):
                 register_request_method(GetTaskStateStore, 
_handle_get_task_state_store),
                 register_request_method(RescheduleTask, 
_handle_reschedule_task),
                 register_request_method(ResendLoggingFD, 
_handle_resend_logging_fd),
-                register_request_method(RetryTask, _handle_finished_task),
+                register_request_method(RetryTask, _handle_task_state),
                 register_request_method(SetAssetStateStoreByName, 
_handle_set_asset_state_store_by_name),
                 register_request_method(SetAssetStateStoreByUri, 
_handle_set_asset_state_store_by_uri),
                 register_request_method(SetRenderedFields, 
_handle_set_rendered_fields),
@@ -2492,12 +2494,12 @@ class InProcessTestSupervisor(ActivitySubprocess):
 
                 state, msg, error = run(ti, context, log)
                 context["exception"] = error
-                finalize(ti, state, context, log, error)
-
-                # In the normal subprocess model, the task runner calls this 
before exiting.
-                # Since we're running in-process, we manually notify the API 
server that
-                # the task has finished—unless the terminal state was already 
sent explicitly.
-                supervisor.update_task_state_if_needed()
+                try:
+                    finalize(ti, state, context, log, error)
+                finally:
+                    # In the normal subprocess model, the supervisor reports 
pending outcomes after
+                    # the child exits; in-process execution must do this even 
if finalization raises.
+                    supervisor.update_task_state_if_needed()
 
         return TaskRunResult(ti=ti, state=state, msg=msg, error=error)
 
diff --git a/task-sdk/tests/task_sdk/execution_time/test_supervisor.py 
b/task-sdk/tests/task_sdk/execution_time/test_supervisor.py
index 2767e9f6f06..1fba4c43661 100644
--- a/task-sdk/tests/task_sdk/execution_time/test_supervisor.py
+++ b/task-sdk/tests/task_sdk/execution_time/test_supervisor.py
@@ -2035,17 +2035,6 @@ REQUEST_TEST_CASES = [
         message=RetryTask(
             end_date=timezone.parse("2024-10-31T12:00:00Z"), 
rendered_map_index="test retry task"
         ),
-        client_mock=ClientMock(
-            method_path="task_instances.retry",
-            kwargs={
-                "id": TI_ID,
-                "end_date": timezone.parse("2024-10-31T12:00:00Z"),
-                "rendered_map_index": "test retry task",
-                "retry_delay_seconds": None,
-                "retry_reason": None,
-            },
-            response=OKResponse(ok=True),
-        ),
         test_id="up_for_retry",
     ),
     RequestTestCase(
@@ -3446,6 +3435,84 @@ class TestHandleRequest:
         assert process._terminal_state == TaskInstanceState.FAILED
         assert process._pending_terminal_state_msg is msg
 
+    @pytest.mark.parametrize("exit_code", [0, 1, -signal.SIGTERM])
+    def test_retry_is_reported_after_finalization(self, watched_subprocess, 
exit_code):
+        process, _ = watched_subprocess
+        msg = RetryTask(
+            end_date=timezone.parse("2024-10-31T12:00:00Z"),
+            rendered_map_index="retrying",
+            retry_delay_seconds=37,
+            retry_reason="rate limited",
+        )
+
+        process._handle_request(msg, structlog.get_logger(), req_id=1)
+
+        process.client.task_instances.retry.assert_not_called()
+        process._send_heartbeat_if_needed()
+        process.client.task_instances.heartbeat.assert_called_once()
+        process._exit_code = exit_code
+        assert process.final_state == TaskInstanceState.UP_FOR_RETRY
+
+        process.update_task_state_if_needed()
+
+        process.client.task_instances.retry.assert_called_once_with(
+            id=TI_ID,
+            end_date=msg.end_date,
+            rendered_map_index="retrying",
+            retry_delay_seconds=37,
+            retry_reason="rate limited",
+        )
+        process.client.task_instances.finish.assert_not_called()
+
+    def test_retry_finalization_is_bounded_by_overtime(self, 
watched_subprocess, mocker):
+        process, _ = watched_subprocess
+        kill = mocker.patch.object(ActivitySubprocess, "kill", autospec=True)
+        monotonic = mocker.patch("time.monotonic", autospec=True, 
return_value=1.0)
+        process._handle_request(RetryTask(end_date=timezone.utcnow()), 
structlog.get_logger(), req_id=1)
+
+        monotonic.return_value += supervisor.TASK_OVERTIME_THRESHOLD + 1
+        process._handle_process_overtime_if_needed()
+
+        kill.assert_called_once_with(process, signal.SIGTERM, force=True)
+
+    def test_wait_propagates_delayed_retry_report_failure(self, 
watched_subprocess, mocker):
+        process, _ = watched_subprocess
+        msg = RetryTask(end_date=timezone.utcnow())
+        process.client.task_instances.retry.side_effect = 
httpx.ConnectError("connection refused")
+        process._handle_request(msg, structlog.get_logger(), req_id=1)
+
+        process.client.task_instances.retry.assert_not_called()
+        process._send_heartbeat_if_needed()
+        process.client.task_instances.heartbeat.assert_called_once()
+
+        def child_exited(process):
+            process._exit_code = 0
+
+        mocker.patch.object(
+            ActivitySubprocess, "_monitor_subprocess", autospec=True, 
side_effect=child_exited
+        )
+        upload_logs = mocker.patch.object(ActivitySubprocess, "_upload_logs", 
autospec=True)
+
+        with pytest.raises(httpx.ConnectError, match="connection refused"):
+            process.wait()
+
+        assert process._terminal_state == TaskInstanceState.UP_FOR_RETRY
+        assert process._pending_terminal_state_msg is msg
+        process.client.task_instances.finish.assert_not_called()
+        upload_logs.assert_called_once_with(process)
+
+    def test_server_termination_cancels_pending_retry(self, 
watched_subprocess):
+        process, _ = watched_subprocess
+        process._handle_request(RetryTask(end_date=timezone.utcnow()), 
structlog.get_logger(), req_id=1)
+        process._terminal_state = supervisor.SERVER_TERMINATED
+        process._exit_code = -signal.SIGTERM
+
+        process.update_task_state_if_needed()
+
+        process.client.task_instances.retry.assert_not_called()
+        process.client.task_instances.finish.assert_not_called()
+        assert process.final_state == supervisor.SERVER_TERMINATED
+
     class _OverrideActivitySubprocess(ActivitySubprocess):
         def _handle_set_rendered_map_index(
             self, msg: SetRenderedMapIndex, log: FilteringBoundLogger, req_id: 
int
@@ -3802,12 +3869,6 @@ class TestHandleRequest:
                 TaskInstanceState.SUCCESS,
                 id="succeed",
             ),
-            pytest.param(
-                RetryTask(end_date=timezone.parse("2024-10-31T12:00:00Z")),
-                "retry",
-                TaskInstanceState.UP_FOR_RETRY,
-                id="retry",
-            ),
             pytest.param(
                 DeferTask(
                     next_method="execute_complete",
@@ -3957,17 +4018,17 @@ class TestHandleRequest:
 
 
 @pytest.mark.parametrize(
-    "test_name",
+    ("test_method", "excluded_message_types"),
     [
-        "test_worker_outcome_retained_when_direct_api_fails",
-        "test_update_task_state_replays_pending_terminal_state_call",
+        (TestHandleRequest.test_worker_outcome_retained_when_direct_api_fails, 
{TaskState, RetryTask}),
+        
(TestHandleRequest.test_update_task_state_replays_pending_terminal_state_call, 
{TaskState}),
     ],
 )
-def test_terminal_message_parameter_coverage(test_name):
-    marks = getattr(TestHandleRequest, test_name).pytestmark
+def test_terminal_message_parameter_coverage(test_method, 
excluded_message_types):
+    marks = test_method.pytestmark
     cases = next(mark.args[1] for mark in marks if mark.name == "parametrize")
     message_types = 
set(get_args(get_type_hints(ActivitySubprocess._send_terminal_state_msg)["msg"]))
-    assert {type(case.values[0]) for case in cases} == message_types - 
{TaskState}
+    assert {type(case.values[0]) for case in cases} == message_types - 
excluded_message_types
 
     marks = 
TestHandleRequest.test_task_state_waits_for_exit_and_keeps_heartbeating.pytestmark
     states = next(mark.args[1] for mark in marks if mark.name == "parametrize" 
and mark.args[0] == "state")
@@ -4015,6 +4076,85 @@ class TestSetSupervisorComms:
 
 
 class TestInProcessTestSupervisor:
+    @pytest.mark.parametrize("callback_error", [None, RuntimeError])
+    def test_retry_callback_finishes_before_retry_report(self, 
make_ti_context, mocker, callback_error):
+        client = mocker.Mock(spec=sdk_client.Client)
+        client.task_instances = 
mocker.create_autospec(sdk_client.TaskInstanceOperations, instance=True)
+        client.xcoms = mocker.create_autospec(sdk_client.XComOperations, 
instance=True)
+        client.task_instances.start.return_value = 
make_ti_context(should_retry=True, max_tries=1)
+        observed = []
+
+        def callback(context):
+            ti = context["ti"]
+            observed.append((client.task_instances.retry.call_count, ti.id, 
ti.end_date))
+            ti.xcom_push(key="retry", value="callback")
+            if callback_error:
+                raise callback_error("callback failed")
+
+        class FailingOperator(BaseOperator):
+            def execute(self, context):
+                raise ValueError("task failed")
+
+        with DAG(dag_id="test_dag"):
+            task = FailingOperator(task_id="failing", retries=1, 
on_retry_callback=callback)
+        ti = TaskInstance(
+            id=uuid7(),
+            dag_version_id=uuid7(),
+            dag_id="test_dag",
+            task_id=task.task_id,
+            run_id="test_run",
+            try_number=1,
+            queue="default",
+        )
+
+        result = InProcessTestSupervisor.start(what=ti, task=task, 
client=client)
+
+        assert result.state == TaskInstanceState.UP_FOR_RETRY
+        assert observed == [(0, ti.id, result.msg.end_date)]
+        client.xcoms.set.assert_called_once()
+        client.task_instances.retry.assert_called_once_with(
+            id=ti.id,
+            end_date=result.msg.end_date,
+            rendered_map_index=None,
+            retry_delay_seconds=None,
+            retry_reason=None,
+        )
+        client.task_instances.finish.assert_not_called()
+
+    @pytest.mark.parametrize("finalization_error", [RuntimeError, 
KeyboardInterrupt])
+    def test_pending_retry_is_reported_when_finalization_raises(
+        self, make_ti_context, mocker, finalization_error
+    ):
+        client = mocker.Mock(spec=sdk_client.Client)
+        client.task_instances = 
mocker.create_autospec(sdk_client.TaskInstanceOperations, instance=True)
+        client.task_instances.start.return_value = 
make_ti_context(should_retry=True, max_tries=1)
+        with DAG(dag_id="test_dag"):
+            task = BaseOperator(task_id="failing", retries=1)
+        ti = TaskInstance(
+            id=uuid7(),
+            dag_version_id=uuid7(),
+            dag_id="test_dag",
+            task_id=task.task_id,
+            run_id="test_run",
+            try_number=1,
+            queue="default",
+        )
+        finalize = mocker.patch.object(task_runner, "finalize", autospec=True, 
side_effect=finalization_error)
+
+        with pytest.raises(finalization_error):
+            InProcessTestSupervisor.start(what=ti, task=task, client=client)
+
+        runtime_ti, state, *_ = finalize.call_args.args
+        assert state == TaskInstanceState.UP_FOR_RETRY
+        client.task_instances.retry.assert_called_once_with(
+            id=ti.id,
+            end_date=runtime_ti.end_date,
+            rendered_map_index=None,
+            retry_delay_seconds=None,
+            retry_reason=None,
+        )
+        client.task_instances.finish.assert_not_called()
+
     def test_inprocess_supervisor_comms_roundtrip(self):
         """
         Test that InProcessSupervisorComms correctly sends a message to the 
supervisor,

Reply via email to