Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
108 changes: 86 additions & 22 deletions task-sdk/src/airflow/sdk/execution_time/supervisor.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1089,6 +1089,18 @@ class ActivitySubprocess(WatchedSubprocess):

_terminal_state: str | None = attrs.field(default=None, init=False)
_final_state: str | None = attrs.field(default=None, init=False)
# The terminal-state message currently being processed by `_handle_request`,
# captured BEFORE the dedicated API call (succeed / retry / defer /
# reschedule). If the API call raises (network blip, server 5xx, etc.),
# this attribute stays set and the dispatcher in
# `update_task_state_if_needed` re-issues the matching API call on
# subprocess exit — re-attempting the original transition rather than
# falling back to `finish()`, which doesn't accept SUCCESS / DEFERRED /
# SERVER_TERMINATED on the server side. Cleared (and `_terminal_state`
# set) only after the API call returns successfully.
_pending_terminal_state_msg: SucceedTask | RetryTask | DeferTask | RescheduleTask | None = attrs.field(
default=None, init=False
)

_last_successful_heartbeat: float = attrs.field(default=0, init=False)
_last_heartbeat_attempt: float = attrs.field(default=0, init=False)
Expand DownExpand Up@@ -1206,10 +1218,23 @@ def wait(self) -> int:
return self._exit_code

def update_task_state_if_needed(self):
# If the process has finished non-directly patched state (directly means deferred, reschedule, etc.),
# update the state of the TaskInstance to reflect the final state of the process.
# For states like `deferred`, `up_for_reschedule`, the process will exit with 0, but the state will be updated
# by the subprocess in the `handle_requests` method.
# If a direct-state API call (succeed / retry / defer / reschedule)
# was attempted but raised, `_pending_terminal_state_msg` still holds
# the original request. Re-issue the matching dedicated API call so
# the server learns the terminal state we couldn't deliver earlier.
# Without this recovery, a transient API failure during the direct
# call would leave the TI stuck RUNNING on the server — `finish()`
# cannot substitute because the server-side `finish` endpoint does
# not accept SUCCESS / DEFERRED / SERVER_TERMINATED transitions.
if self._pending_terminal_state_msg is not None:
self._replay_pending_terminal_state_msg()
return

# If the process has finished a non-directly-patched state (e.g.
# FAILED, UP_FOR_RETRY without RetryTask), `finish()` is the
# dedicated endpoint for those transitions. For states already in
# STATES_SENT_DIRECTLY whose direct API call succeeded, no further
# action is needed.
if self.final_state not in STATES_SENT_DIRECTLY:
self.client.task_instances.finish(
id=self.id,
Expand All@@ -1218,6 +1243,56 @@ def update_task_state_if_needed(self):
rendered_map_index=self._rendered_map_index,
)

def _send_terminal_state_msg(self, msg: SucceedTask | RetryTask | DeferTask | RescheduleTask) -> None:
# Capture the message BEFORE the API call so the recovery dispatcher
# in `update_task_state_if_needed` can re-issue it if the call raises
# (network blip, transient server 5xx). Clear the pending slot and
# record the resulting state only after the call returns successfully.
self._pending_terminal_state_msg = msg
if isinstance(msg, SucceedTask):
self.client.task_instances.succeed(
id=self.id,
when=msg.end_date,
task_outlets=msg.task_outlets,
outlet_events=msg.outlet_events,
rendered_map_index=self._rendered_map_index,
)
self._terminal_state = msg.state
elif isinstance(msg, RetryTask):
self.client.task_instances.retry(
id=self.id,
end_date=msg.end_date,
rendered_map_index=self._rendered_map_index,
)
self._terminal_state = msg.state
elif isinstance(msg, DeferTask):
self.client.task_instances.defer(self.id, msg)
self._terminal_state = TaskInstanceState.DEFERRED
elif isinstance(msg, RescheduleTask):
self.client.task_instances.reschedule(self.id, msg)
self._terminal_state = TaskInstanceState.UP_FOR_RESCHEDULE
self._pending_terminal_state_msg = None

def _replay_pending_terminal_state_msg(self) -> None:
"""
Re-issue the dedicated API call for an unsynced terminal-state msg.

Best-effort — if the second attempt also fails the exception is
logged and we move on; the supervisor's overall failure handling
(heartbeat, exit-code reporting) will eventually surface the issue.
"""
msg = self._pending_terminal_state_msg
if msg is None:
return
try:
self._send_terminal_state_msg(msg)
except Exception:
log.exception(
"Recovery retry of terminal-state API call failed; TI may be stuck on the server",
ti_id=self.id,
msg_type=type(msg).__name__,
)

def _upload_logs(self):
"""
Upload all log files found to the remote storage.
Expand DownExpand Up@@ -1389,29 +1464,20 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id:
resp: BaseModel | None = None
dump_opts = {}
if isinstance(msg, TaskState):
# No direct API call here — the recovery path in
# `update_task_state_if_needed` will call `finish()` for
# non-direct states (FAILED, etc.) once the subprocess exits.
self._terminal_state = msg.state
self._task_end_time_monotonic = time.monotonic()
self._rendered_map_index = msg.rendered_map_index
elif isinstance(msg, SucceedTask):
self._terminal_state = msg.state
self._task_end_time_monotonic = time.monotonic()
self._rendered_map_index = msg.rendered_map_index
self.client.task_instances.succeed(
id=self.id,
when=msg.end_date,
task_outlets=msg.task_outlets,
outlet_events=msg.outlet_events,
rendered_map_index=self._rendered_map_index,
)
self._send_terminal_state_msg(msg)
elif isinstance(msg, RetryTask):
self._terminal_state = msg.state
self._task_end_time_monotonic = time.monotonic()
self._rendered_map_index = msg.rendered_map_index
self.client.task_instances.retry(
id=self.id,
end_date=msg.end_date,
rendered_map_index=self._rendered_map_index,
)
self._send_terminal_state_msg(msg)
elif isinstance(msg, GetConnection):
conn = self.client.connections.get(msg.conn_id)
if isinstance(conn, ConnectionResponse):
Expand DownExpand Up@@ -1463,12 +1529,10 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id:
)
resp = XComSequenceSliceResult.from_response(xcoms)
elif isinstance(msg, DeferTask):
self._terminal_state = TaskInstanceState.DEFERRED
self._rendered_map_index = msg.rendered_map_index
self.client.task_instances.defer(self.id, msg)
self._send_terminal_state_msg(msg)
elif isinstance(msg, RescheduleTask):
self._terminal_state = TaskInstanceState.UP_FOR_RESCHEDULE
self.client.task_instances.reschedule(self.id, msg)
self._send_terminal_state_msg(msg)
elif isinstance(msg, SkipDownstreamTasks):
self.client.task_instances.skip_downstream_tasks(self.id, msg)
elif isinstance(msg, SetXCom):
Expand Down
131 changes: 131 additions & 0 deletions task-sdk/tests/task_sdk/execution_time/test_supervisor.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -2793,6 +2793,137 @@ def test_handle_requests_network_exception_does_not_crash_loop(self, watched_sub
# Should not raise StopIteration (which would mean the loop crashed).
generator.send(req2)

@pytest.mark.parametrize(
("msg", "api_method", "expected_state"),
[
pytest.param(
SucceedTask(end_date=timezone.parse("2024-10-31T12:00:00Z")),
"succeed",
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",
classpath="airflow.providers.standard.triggers.external_task.WorkflowTrigger",
trigger_kwargs={},
),
"defer",
TaskInstanceState.DEFERRED,
id="defer",
),
pytest.param(
RescheduleTask(
reschedule_date=timezone.parse("2024-10-31T12:00:00Z"),
end_date=timezone.parse("2024-10-31T12:00:00Z"),
),
"reschedule",
TaskInstanceState.UP_FOR_RESCHEDULE,
id="reschedule",
),
],
)
def test_terminal_state_not_set_when_direct_api_fails(
self, watched_subprocess, mocker, msg, api_method, expected_state
):
"""`_terminal_state` must NOT be set when the dedicated terminal-state
API raises.

The original message is captured in `_pending_terminal_state_msg`
BEFORE the API call so the recovery dispatcher in
`update_task_state_if_needed` can re-issue it on subprocess exit.
Covers all four terminal-state message types.
"""
watched_subprocess, _ = watched_subprocess
setattr(
watched_subprocess.client.task_instances,
api_method,
mocker.Mock(side_effect=httpx.ConnectError("connection refused")),
)

with pytest.raises(httpx.ConnectError):
watched_subprocess._handle_request(msg, mocker.Mock(), req_id=1)

assert watched_subprocess._terminal_state is None
# Pending msg preserved so the recovery dispatcher can re-issue.
assert watched_subprocess._pending_terminal_state_msg is msg

@pytest.mark.parametrize(
("msg", "api_method", "expected_state"),
[
pytest.param(
SucceedTask(end_date=timezone.parse("2024-10-31T12:00:00Z")),
"succeed",
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",
classpath="airflow.providers.standard.triggers.external_task.WorkflowTrigger",
trigger_kwargs={},
),
"defer",
TaskInstanceState.DEFERRED,
id="defer",
),
pytest.param(
RescheduleTask(
reschedule_date=timezone.parse("2024-10-31T12:00:00Z"),
end_date=timezone.parse("2024-10-31T12:00:00Z"),
),
"reschedule",
TaskInstanceState.UP_FOR_RESCHEDULE,
id="reschedule",
),
],
)
def test_update_task_state_replays_pending_terminal_state_call(
self, watched_subprocess, mocker, msg, api_method, expected_state
):
"""If a direct terminal-state API call was attempted and raised, the
recovery dispatcher must re-issue the dedicated endpoint (not
`finish()`, which the server-side endpoint refuses for SUCCESS /
DEFERRED / SERVER_TERMINATED). Covers all four message types.
"""
watched_subprocess, _ = watched_subprocess
watched_subprocess._exit_code = 0
# Simulate the failure scenario: original API call raised, msg preserved.
watched_subprocess._pending_terminal_state_msg = msg

watched_subprocess.update_task_state_if_needed()

# Recovery re-issues the dedicated endpoint, NOT finish().
getattr(watched_subprocess.client.task_instances, api_method).assert_called_once()
watched_subprocess.client.task_instances.finish.assert_not_called()
assert watched_subprocess._terminal_state == expected_state
assert watched_subprocess._pending_terminal_state_msg is None

def test_update_task_state_no_recovery_without_pending_msg(self, watched_subprocess, mocker):
"""No replay when nothing was pending — preserves the original
STATES_SENT_DIRECTLY short-circuit for the happy path."""
watched_subprocess, _ = watched_subprocess
watched_subprocess._exit_code = 0
watched_subprocess._terminal_state = TaskInstanceState.SUCCESS
watched_subprocess._pending_terminal_state_msg = None

watched_subprocess.update_task_state_if_needed()

watched_subprocess.client.task_instances.finish.assert_not_called()
watched_subprocess.client.task_instances.succeed.assert_not_called()


class TestSetSupervisorComms:
class DummyComms:
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Add copy buttons to all
 blocks
(function() {
function addCopyButtons() {
document.querySelectorAll('pre code').forEach(function(codeBlock) {
if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;
codeBlock.parentElement.setAttribute('data-copy-added', 'true');
var btn = document.createElement('button');
btn.textContent = 'Copy';
btn.style.cssText = 'position:absolute;top:4px;right:4px;padding:2px 8px;font-size:11px;background:#4ecdc4;border:none;border-radius:4px;color:#1a1a2e;cursor:pointer;opacity:0.7;transition:opacity 0.2s;';
btn.onmouseover = function() { this.style.opacity = '1'; };
btn.onmouseout = function() { this.style.opacity = '0.7'; };
btn.onclick = function() {
navigator.clipboard.writeText(codeBlock.textContent).then(function() {
btn.textContent = 'Copied!';
setTimeout(function() { btn.textContent = 'Copy'; }, 1500);
});
};
codeBlock.parentElement.style.position = 'relative';
codeBlock.parentElement.appendChild(btn);
});
}
addCopyButtons();
// Re-run on dynamic content
var observer = new MutationObserver(addCopyButtons);
observer.observe(document.body, { childList: true, subtree: true });
})();
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
[v3-2-test] Recover stuck TIs when direct terminal-state API call fails (#66574) by vatsrahul1001 · Pull Request #67204 · apache/airflow · GitHub
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
108 changes: 86 additions & 22 deletions task-sdk/src/airflow/sdk/execution_time/supervisor.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1089,6 +1089,18 @@ class ActivitySubprocess(WatchedSubprocess):

_terminal_state: str | None = attrs.field(default=None, init=False)
_final_state: str | None = attrs.field(default=None, init=False)
# The terminal-state message currently being processed by `_handle_request`,
# captured BEFORE the dedicated API call (succeed / retry / defer /
# reschedule). If the API call raises (network blip, server 5xx, etc.),
# this attribute stays set and the dispatcher in
# `update_task_state_if_needed` re-issues the matching API call on
# subprocess exit — re-attempting the original transition rather than
# falling back to `finish()`, which doesn't accept SUCCESS / DEFERRED /
# SERVER_TERMINATED on the server side. Cleared (and `_terminal_state`
# set) only after the API call returns successfully.
_pending_terminal_state_msg: SucceedTask | RetryTask | DeferTask | RescheduleTask | None = attrs.field(
default=None, init=False
)

_last_successful_heartbeat: float = attrs.field(default=0, init=False)
_last_heartbeat_attempt: float = attrs.field(default=0, init=False)
Expand DownExpand Up@@ -1206,10 +1218,23 @@ def wait(self) -> int:
return self._exit_code

def update_task_state_if_needed(self):
# If the process has finished non-directly patched state (directly means deferred, reschedule, etc.),
# update the state of the TaskInstance to reflect the final state of the process.
# For states like `deferred`, `up_for_reschedule`, the process will exit with 0, but the state will be updated
# by the subprocess in the `handle_requests` method.
# If a direct-state API call (succeed / retry / defer / reschedule)
# was attempted but raised, `_pending_terminal_state_msg` still holds
# the original request. Re-issue the matching dedicated API call so
# the server learns the terminal state we couldn't deliver earlier.
# Without this recovery, a transient API failure during the direct
# call would leave the TI stuck RUNNING on the server — `finish()`
# cannot substitute because the server-side `finish` endpoint does
# not accept SUCCESS / DEFERRED / SERVER_TERMINATED transitions.
if self._pending_terminal_state_msg is not None:
self._replay_pending_terminal_state_msg()
return

# If the process has finished a non-directly-patched state (e.g.
# FAILED, UP_FOR_RETRY without RetryTask), `finish()` is the
# dedicated endpoint for those transitions. For states already in
# STATES_SENT_DIRECTLY whose direct API call succeeded, no further
# action is needed.
if self.final_state not in STATES_SENT_DIRECTLY:
self.client.task_instances.finish(
id=self.id,
Expand All@@ -1218,6 +1243,56 @@ def update_task_state_if_needed(self):
rendered_map_index=self._rendered_map_index,
)

def _send_terminal_state_msg(self, msg: SucceedTask | RetryTask | DeferTask | RescheduleTask) -> None:
# Capture the message BEFORE the API call so the recovery dispatcher
# in `update_task_state_if_needed` can re-issue it if the call raises
# (network blip, transient server 5xx). Clear the pending slot and
# record the resulting state only after the call returns successfully.
self._pending_terminal_state_msg = msg
if isinstance(msg, SucceedTask):
self.client.task_instances.succeed(
id=self.id,
when=msg.end_date,
task_outlets=msg.task_outlets,
outlet_events=msg.outlet_events,
rendered_map_index=self._rendered_map_index,
)
self._terminal_state = msg.state
elif isinstance(msg, RetryTask):
self.client.task_instances.retry(
id=self.id,
end_date=msg.end_date,
rendered_map_index=self._rendered_map_index,
)
self._terminal_state = msg.state
elif isinstance(msg, DeferTask):
self.client.task_instances.defer(self.id, msg)
self._terminal_state = TaskInstanceState.DEFERRED
elif isinstance(msg, RescheduleTask):
self.client.task_instances.reschedule(self.id, msg)
self._terminal_state = TaskInstanceState.UP_FOR_RESCHEDULE
self._pending_terminal_state_msg = None

def _replay_pending_terminal_state_msg(self) -> None:
"""
Re-issue the dedicated API call for an unsynced terminal-state msg.

Best-effort — if the second attempt also fails the exception is
logged and we move on; the supervisor's overall failure handling
(heartbeat, exit-code reporting) will eventually surface the issue.
"""
msg = self._pending_terminal_state_msg
if msg is None:
return
try:
self._send_terminal_state_msg(msg)
except Exception:
log.exception(
"Recovery retry of terminal-state API call failed; TI may be stuck on the server",
ti_id=self.id,
msg_type=type(msg).__name__,
)

def _upload_logs(self):
"""
Upload all log files found to the remote storage.
Expand DownExpand Up@@ -1389,29 +1464,20 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id:
resp: BaseModel | None = None
dump_opts = {}
if isinstance(msg, TaskState):
# No direct API call here — the recovery path in
# `update_task_state_if_needed` will call `finish()` for
# non-direct states (FAILED, etc.) once the subprocess exits.
self._terminal_state = msg.state
self._task_end_time_monotonic = time.monotonic()
self._rendered_map_index = msg.rendered_map_index
elif isinstance(msg, SucceedTask):
self._terminal_state = msg.state
self._task_end_time_monotonic = time.monotonic()
self._rendered_map_index = msg.rendered_map_index
self.client.task_instances.succeed(
id=self.id,
when=msg.end_date,
task_outlets=msg.task_outlets,
outlet_events=msg.outlet_events,
rendered_map_index=self._rendered_map_index,
)
self._send_terminal_state_msg(msg)
elif isinstance(msg, RetryTask):
self._terminal_state = msg.state
self._task_end_time_monotonic = time.monotonic()
self._rendered_map_index = msg.rendered_map_index
self.client.task_instances.retry(
id=self.id,
end_date=msg.end_date,
rendered_map_index=self._rendered_map_index,
)
self._send_terminal_state_msg(msg)
elif isinstance(msg, GetConnection):
conn = self.client.connections.get(msg.conn_id)
if isinstance(conn, ConnectionResponse):
Expand DownExpand Up@@ -1463,12 +1529,10 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id:
)
resp = XComSequenceSliceResult.from_response(xcoms)
elif isinstance(msg, DeferTask):
self._terminal_state = TaskInstanceState.DEFERRED
self._rendered_map_index = msg.rendered_map_index
self.client.task_instances.defer(self.id, msg)
self._send_terminal_state_msg(msg)
elif isinstance(msg, RescheduleTask):
self._terminal_state = TaskInstanceState.UP_FOR_RESCHEDULE
self.client.task_instances.reschedule(self.id, msg)
self._send_terminal_state_msg(msg)
elif isinstance(msg, SkipDownstreamTasks):
self.client.task_instances.skip_downstream_tasks(self.id, msg)
elif isinstance(msg, SetXCom):
Expand Down
131 changes: 131 additions & 0 deletions task-sdk/tests/task_sdk/execution_time/test_supervisor.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -2793,6 +2793,137 @@ def test_handle_requests_network_exception_does_not_crash_loop(self, watched_sub
# Should not raise StopIteration (which would mean the loop crashed).
generator.send(req2)

@pytest.mark.parametrize(
("msg", "api_method", "expected_state"),
[
pytest.param(
SucceedTask(end_date=timezone.parse("2024-10-31T12:00:00Z")),
"succeed",
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",
classpath="airflow.providers.standard.triggers.external_task.WorkflowTrigger",
trigger_kwargs={},
),
"defer",
TaskInstanceState.DEFERRED,
id="defer",
),
pytest.param(
RescheduleTask(
reschedule_date=timezone.parse("2024-10-31T12:00:00Z"),
end_date=timezone.parse("2024-10-31T12:00:00Z"),
),
"reschedule",
TaskInstanceState.UP_FOR_RESCHEDULE,
id="reschedule",
),
],
)
def test_terminal_state_not_set_when_direct_api_fails(
self, watched_subprocess, mocker, msg, api_method, expected_state
):
"""`_terminal_state` must NOT be set when the dedicated terminal-state
API raises.

The original message is captured in `_pending_terminal_state_msg`
BEFORE the API call so the recovery dispatcher in
`update_task_state_if_needed` can re-issue it on subprocess exit.
Covers all four terminal-state message types.
"""
watched_subprocess, _ = watched_subprocess
setattr(
watched_subprocess.client.task_instances,
api_method,
mocker.Mock(side_effect=httpx.ConnectError("connection refused")),
)

with pytest.raises(httpx.ConnectError):
watched_subprocess._handle_request(msg, mocker.Mock(), req_id=1)

assert watched_subprocess._terminal_state is None
# Pending msg preserved so the recovery dispatcher can re-issue.
assert watched_subprocess._pending_terminal_state_msg is msg

@pytest.mark.parametrize(
("msg", "api_method", "expected_state"),
[
pytest.param(
SucceedTask(end_date=timezone.parse("2024-10-31T12:00:00Z")),
"succeed",
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",
classpath="airflow.providers.standard.triggers.external_task.WorkflowTrigger",
trigger_kwargs={},
),
"defer",
TaskInstanceState.DEFERRED,
id="defer",
),
pytest.param(
RescheduleTask(
reschedule_date=timezone.parse("2024-10-31T12:00:00Z"),
end_date=timezone.parse("2024-10-31T12:00:00Z"),
),
"reschedule",
TaskInstanceState.UP_FOR_RESCHEDULE,
id="reschedule",
),
],
)
def test_update_task_state_replays_pending_terminal_state_call(
self, watched_subprocess, mocker, msg, api_method, expected_state
):
"""If a direct terminal-state API call was attempted and raised, the
recovery dispatcher must re-issue the dedicated endpoint (not
`finish()`, which the server-side endpoint refuses for SUCCESS /
DEFERRED / SERVER_TERMINATED). Covers all four message types.
"""
watched_subprocess, _ = watched_subprocess
watched_subprocess._exit_code = 0
# Simulate the failure scenario: original API call raised, msg preserved.
watched_subprocess._pending_terminal_state_msg = msg

watched_subprocess.update_task_state_if_needed()

# Recovery re-issues the dedicated endpoint, NOT finish().
getattr(watched_subprocess.client.task_instances, api_method).assert_called_once()
watched_subprocess.client.task_instances.finish.assert_not_called()
assert watched_subprocess._terminal_state == expected_state
assert watched_subprocess._pending_terminal_state_msg is None

def test_update_task_state_no_recovery_without_pending_msg(self, watched_subprocess, mocker):
"""No replay when nothing was pending — preserves the original
STATES_SENT_DIRECTLY short-circuit for the happy path."""
watched_subprocess, _ = watched_subprocess
watched_subprocess._exit_code = 0
watched_subprocess._terminal_state = TaskInstanceState.SUCCESS
watched_subprocess._pending_terminal_state_msg = None

watched_subprocess.update_task_state_if_needed()

watched_subprocess.client.task_instances.finish.assert_not_called()
watched_subprocess.client.task_instances.succeed.assert_not_called()


class TestSetSupervisorComms:
class DummyComms:
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Force GitHub README to respect dark mode (function() { var style = document.createElement('style'); style.textContent = ' .markdown-body { color-scheme: dark light; } .markdown-body pre { background: #161b22 !important; } .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; } .markdown-body table th, .markdown-body table td { border-color: #30363d !important; } .markdown-body img { background: #0d1117; } .markdown-body blockquote { border-left-color: #8b949e; } .markdown-body hr { border-color: #30363d; } '; document.head.appendChild(style); })(); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' [v3-2-test] Recover stuck TIs when direct terminal-state API call fails (#66574) by vatsrahul1001 · Pull Request #67204 · apache/airflow · GitHub
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
108 changes: 86 additions & 22 deletions task-sdk/src/airflow/sdk/execution_time/supervisor.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1089,6 +1089,18 @@ class ActivitySubprocess(WatchedSubprocess):

_terminal_state: str | None = attrs.field(default=None, init=False)
_final_state: str | None = attrs.field(default=None, init=False)
# The terminal-state message currently being processed by `_handle_request`,
# captured BEFORE the dedicated API call (succeed / retry / defer /
# reschedule). If the API call raises (network blip, server 5xx, etc.),
# this attribute stays set and the dispatcher in
# `update_task_state_if_needed` re-issues the matching API call on
# subprocess exit — re-attempting the original transition rather than
# falling back to `finish()`, which doesn't accept SUCCESS / DEFERRED /
# SERVER_TERMINATED on the server side. Cleared (and `_terminal_state`
# set) only after the API call returns successfully.
_pending_terminal_state_msg: SucceedTask | RetryTask | DeferTask | RescheduleTask | None = attrs.field(
default=None, init=False
)

_last_successful_heartbeat: float = attrs.field(default=0, init=False)
_last_heartbeat_attempt: float = attrs.field(default=0, init=False)
Expand DownExpand Up@@ -1206,10 +1218,23 @@ def wait(self) -> int:
return self._exit_code

def update_task_state_if_needed(self):
# If the process has finished non-directly patched state (directly means deferred, reschedule, etc.),
# update the state of the TaskInstance to reflect the final state of the process.
# For states like `deferred`, `up_for_reschedule`, the process will exit with 0, but the state will be updated
# by the subprocess in the `handle_requests` method.
# If a direct-state API call (succeed / retry / defer / reschedule)
# was attempted but raised, `_pending_terminal_state_msg` still holds
# the original request. Re-issue the matching dedicated API call so
# the server learns the terminal state we couldn't deliver earlier.
# Without this recovery, a transient API failure during the direct
# call would leave the TI stuck RUNNING on the server — `finish()`
# cannot substitute because the server-side `finish` endpoint does
# not accept SUCCESS / DEFERRED / SERVER_TERMINATED transitions.
if self._pending_terminal_state_msg is not None:
self._replay_pending_terminal_state_msg()
return

# If the process has finished a non-directly-patched state (e.g.
# FAILED, UP_FOR_RETRY without RetryTask), `finish()` is the
# dedicated endpoint for those transitions. For states already in
# STATES_SENT_DIRECTLY whose direct API call succeeded, no further
# action is needed.
if self.final_state not in STATES_SENT_DIRECTLY:
self.client.task_instances.finish(
id=self.id,
Expand All@@ -1218,6 +1243,56 @@ def update_task_state_if_needed(self):
rendered_map_index=self._rendered_map_index,
)

def _send_terminal_state_msg(self, msg: SucceedTask | RetryTask | DeferTask | RescheduleTask) -> None:
# Capture the message BEFORE the API call so the recovery dispatcher
# in `update_task_state_if_needed` can re-issue it if the call raises
# (network blip, transient server 5xx). Clear the pending slot and
# record the resulting state only after the call returns successfully.
self._pending_terminal_state_msg = msg
if isinstance(msg, SucceedTask):
self.client.task_instances.succeed(
id=self.id,
when=msg.end_date,
task_outlets=msg.task_outlets,
outlet_events=msg.outlet_events,
rendered_map_index=self._rendered_map_index,
)
self._terminal_state = msg.state
elif isinstance(msg, RetryTask):
self.client.task_instances.retry(
id=self.id,
end_date=msg.end_date,
rendered_map_index=self._rendered_map_index,
)
self._terminal_state = msg.state
elif isinstance(msg, DeferTask):
self.client.task_instances.defer(self.id, msg)
self._terminal_state = TaskInstanceState.DEFERRED
elif isinstance(msg, RescheduleTask):
self.client.task_instances.reschedule(self.id, msg)
self._terminal_state = TaskInstanceState.UP_FOR_RESCHEDULE
self._pending_terminal_state_msg = None

def _replay_pending_terminal_state_msg(self) -> None:
"""
Re-issue the dedicated API call for an unsynced terminal-state msg.

Best-effort — if the second attempt also fails the exception is
logged and we move on; the supervisor's overall failure handling
(heartbeat, exit-code reporting) will eventually surface the issue.
"""
msg = self._pending_terminal_state_msg
if msg is None:
return
try:
self._send_terminal_state_msg(msg)
except Exception:
log.exception(
"Recovery retry of terminal-state API call failed; TI may be stuck on the server",
ti_id=self.id,
msg_type=type(msg).__name__,
)

def _upload_logs(self):
"""
Upload all log files found to the remote storage.
Expand DownExpand Up@@ -1389,29 +1464,20 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id:
resp: BaseModel | None = None
dump_opts = {}
if isinstance(msg, TaskState):
# No direct API call here — the recovery path in
# `update_task_state_if_needed` will call `finish()` for
# non-direct states (FAILED, etc.) once the subprocess exits.
self._terminal_state = msg.state
self._task_end_time_monotonic = time.monotonic()
self._rendered_map_index = msg.rendered_map_index
elif isinstance(msg, SucceedTask):
self._terminal_state = msg.state
self._task_end_time_monotonic = time.monotonic()
self._rendered_map_index = msg.rendered_map_index
self.client.task_instances.succeed(
id=self.id,
when=msg.end_date,
task_outlets=msg.task_outlets,
outlet_events=msg.outlet_events,
rendered_map_index=self._rendered_map_index,
)
self._send_terminal_state_msg(msg)
elif isinstance(msg, RetryTask):
self._terminal_state = msg.state
self._task_end_time_monotonic = time.monotonic()
self._rendered_map_index = msg.rendered_map_index
self.client.task_instances.retry(
id=self.id,
end_date=msg.end_date,
rendered_map_index=self._rendered_map_index,
)
self._send_terminal_state_msg(msg)
elif isinstance(msg, GetConnection):
conn = self.client.connections.get(msg.conn_id)
if isinstance(conn, ConnectionResponse):
Expand DownExpand Up@@ -1463,12 +1529,10 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id:
)
resp = XComSequenceSliceResult.from_response(xcoms)
elif isinstance(msg, DeferTask):
self._terminal_state = TaskInstanceState.DEFERRED
self._rendered_map_index = msg.rendered_map_index
self.client.task_instances.defer(self.id, msg)
self._send_terminal_state_msg(msg)
elif isinstance(msg, RescheduleTask):
self._terminal_state = TaskInstanceState.UP_FOR_RESCHEDULE
self.client.task_instances.reschedule(self.id, msg)
self._send_terminal_state_msg(msg)
elif isinstance(msg, SkipDownstreamTasks):
self.client.task_instances.skip_downstream_tasks(self.id, msg)
elif isinstance(msg, SetXCom):
Expand Down
131 changes: 131 additions & 0 deletions task-sdk/tests/task_sdk/execution_time/test_supervisor.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -2793,6 +2793,137 @@ def test_handle_requests_network_exception_does_not_crash_loop(self, watched_sub
# Should not raise StopIteration (which would mean the loop crashed).
generator.send(req2)

@pytest.mark.parametrize(
("msg", "api_method", "expected_state"),
[
pytest.param(
SucceedTask(end_date=timezone.parse("2024-10-31T12:00:00Z")),
"succeed",
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",
classpath="airflow.providers.standard.triggers.external_task.WorkflowTrigger",
trigger_kwargs={},
),
"defer",
TaskInstanceState.DEFERRED,
id="defer",
),
pytest.param(
RescheduleTask(
reschedule_date=timezone.parse("2024-10-31T12:00:00Z"),
end_date=timezone.parse("2024-10-31T12:00:00Z"),
),
"reschedule",
TaskInstanceState.UP_FOR_RESCHEDULE,
id="reschedule",
),
],
)
def test_terminal_state_not_set_when_direct_api_fails(
self, watched_subprocess, mocker, msg, api_method, expected_state
):
"""`_terminal_state` must NOT be set when the dedicated terminal-state
API raises.

The original message is captured in `_pending_terminal_state_msg`
BEFORE the API call so the recovery dispatcher in
`update_task_state_if_needed` can re-issue it on subprocess exit.
Covers all four terminal-state message types.
"""
watched_subprocess, _ = watched_subprocess
setattr(
watched_subprocess.client.task_instances,
api_method,
mocker.Mock(side_effect=httpx.ConnectError("connection refused")),
)

with pytest.raises(httpx.ConnectError):
watched_subprocess._handle_request(msg, mocker.Mock(), req_id=1)

assert watched_subprocess._terminal_state is None
# Pending msg preserved so the recovery dispatcher can re-issue.
assert watched_subprocess._pending_terminal_state_msg is msg

@pytest.mark.parametrize(
("msg", "api_method", "expected_state"),
[
pytest.param(
SucceedTask(end_date=timezone.parse("2024-10-31T12:00:00Z")),
"succeed",
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",
classpath="airflow.providers.standard.triggers.external_task.WorkflowTrigger",
trigger_kwargs={},
),
"defer",
TaskInstanceState.DEFERRED,
id="defer",
),
pytest.param(
RescheduleTask(
reschedule_date=timezone.parse("2024-10-31T12:00:00Z"),
end_date=timezone.parse("2024-10-31T12:00:00Z"),
),
"reschedule",
TaskInstanceState.UP_FOR_RESCHEDULE,
id="reschedule",
),
],
)
def test_update_task_state_replays_pending_terminal_state_call(
self, watched_subprocess, mocker, msg, api_method, expected_state
):
"""If a direct terminal-state API call was attempted and raised, the
recovery dispatcher must re-issue the dedicated endpoint (not
`finish()`, which the server-side endpoint refuses for SUCCESS /
DEFERRED / SERVER_TERMINATED). Covers all four message types.
"""
watched_subprocess, _ = watched_subprocess
watched_subprocess._exit_code = 0
# Simulate the failure scenario: original API call raised, msg preserved.
watched_subprocess._pending_terminal_state_msg = msg

watched_subprocess.update_task_state_if_needed()

# Recovery re-issues the dedicated endpoint, NOT finish().
getattr(watched_subprocess.client.task_instances, api_method).assert_called_once()
watched_subprocess.client.task_instances.finish.assert_not_called()
assert watched_subprocess._terminal_state == expected_state
assert watched_subprocess._pending_terminal_state_msg is None

def test_update_task_state_no_recovery_without_pending_msg(self, watched_subprocess, mocker):
"""No replay when nothing was pending — preserves the original
STATES_SENT_DIRECTLY short-circuit for the happy path."""
watched_subprocess, _ = watched_subprocess
watched_subprocess._exit_code = 0
watched_subprocess._terminal_state = TaskInstanceState.SUCCESS
watched_subprocess._pending_terminal_state_msg = None

watched_subprocess.update_task_state_if_needed()

watched_subprocess.client.task_instances.finish.assert_not_called()
watched_subprocess.client.task_instances.succeed.assert_not_called()


class TestSetSupervisorComms:
class DummyComms:
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Highlight search terms from Google/DuckDuckGo/Bing referrer (function() { var ref = document.referrer; var terms = []; if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) { var url = new URL(ref); var q = url.searchParams.get('q') || url.searchParams.get('p'); if (q) { terms = q.split(/\s+/).filter(function(t) { return t.length > 2; }); } } if (terms.length === 0) return; var style = document.createElement('style'); style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }'; document.head.appendChild(style); function highlight(node) { if (node.nodeType === 3) { // text node var text = node.textContent; var found = false; terms.forEach(function(term) { var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\]\\]/g, '\\') + ')', 'gi'); if (regex.test(text)) { found = true; var frag = document.createDocumentFragment(); var parts = text.split(regex); parts.forEach(function(part, i) { if (i % 2 === 0) { frag.appendChild(document.createTextNode(part)); } else { var span = document.createElement('span'); span.className = 'userscript-highlight'; span.textContent = part; frag.appendChild(span); } }); node.parentNode.replaceChild(frag, node); } }); } else if (node.nodeType === 1 && node.childNodes) { // element var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT']; if (!skipTags.includes(node.tagName)) { Array.from(node.childNodes).forEach(highlight); } } } highlight(document.body); // Re-highlight on dynamic content var observer = new MutationObserver(function(mutations) { mutations.forEach(function(m) { m.addedNodes.forEach(function(node) { if (node.nodeType === 1 || node.nodeType === 3) highlight(node); }); }); }); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' [v3-2-test] Recover stuck TIs when direct terminal-state API call fails (#66574) by vatsrahul1001 · Pull Request #67204 · apache/airflow · GitHub
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
108 changes: 86 additions & 22 deletions task-sdk/src/airflow/sdk/execution_time/supervisor.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1089,6 +1089,18 @@ class ActivitySubprocess(WatchedSubprocess):

_terminal_state: str | None = attrs.field(default=None, init=False)
_final_state: str | None = attrs.field(default=None, init=False)
# The terminal-state message currently being processed by `_handle_request`,
# captured BEFORE the dedicated API call (succeed / retry / defer /
# reschedule). If the API call raises (network blip, server 5xx, etc.),
# this attribute stays set and the dispatcher in
# `update_task_state_if_needed` re-issues the matching API call on
# subprocess exit — re-attempting the original transition rather than
# falling back to `finish()`, which doesn't accept SUCCESS / DEFERRED /
# SERVER_TERMINATED on the server side. Cleared (and `_terminal_state`
# set) only after the API call returns successfully.
_pending_terminal_state_msg: SucceedTask | RetryTask | DeferTask | RescheduleTask | None = attrs.field(
default=None, init=False
)

_last_successful_heartbeat: float = attrs.field(default=0, init=False)
_last_heartbeat_attempt: float = attrs.field(default=0, init=False)
Expand DownExpand Up@@ -1206,10 +1218,23 @@ def wait(self) -> int:
return self._exit_code

def update_task_state_if_needed(self):
# If the process has finished non-directly patched state (directly means deferred, reschedule, etc.),
# update the state of the TaskInstance to reflect the final state of the process.
# For states like `deferred`, `up_for_reschedule`, the process will exit with 0, but the state will be updated
# by the subprocess in the `handle_requests` method.
# If a direct-state API call (succeed / retry / defer / reschedule)
# was attempted but raised, `_pending_terminal_state_msg` still holds
# the original request. Re-issue the matching dedicated API call so
# the server learns the terminal state we couldn't deliver earlier.
# Without this recovery, a transient API failure during the direct
# call would leave the TI stuck RUNNING on the server — `finish()`
# cannot substitute because the server-side `finish` endpoint does
# not accept SUCCESS / DEFERRED / SERVER_TERMINATED transitions.
if self._pending_terminal_state_msg is not None:
self._replay_pending_terminal_state_msg()
return

# If the process has finished a non-directly-patched state (e.g.
# FAILED, UP_FOR_RETRY without RetryTask), `finish()` is the
# dedicated endpoint for those transitions. For states already in
# STATES_SENT_DIRECTLY whose direct API call succeeded, no further
# action is needed.
if self.final_state not in STATES_SENT_DIRECTLY:
self.client.task_instances.finish(
id=self.id,
Expand All@@ -1218,6 +1243,56 @@ def update_task_state_if_needed(self):
rendered_map_index=self._rendered_map_index,
)

def _send_terminal_state_msg(self, msg: SucceedTask | RetryTask | DeferTask | RescheduleTask) -> None:
# Capture the message BEFORE the API call so the recovery dispatcher
# in `update_task_state_if_needed` can re-issue it if the call raises
# (network blip, transient server 5xx). Clear the pending slot and
# record the resulting state only after the call returns successfully.
self._pending_terminal_state_msg = msg
if isinstance(msg, SucceedTask):
self.client.task_instances.succeed(
id=self.id,
when=msg.end_date,
task_outlets=msg.task_outlets,
outlet_events=msg.outlet_events,
rendered_map_index=self._rendered_map_index,
)
self._terminal_state = msg.state
elif isinstance(msg, RetryTask):
self.client.task_instances.retry(
id=self.id,
end_date=msg.end_date,
rendered_map_index=self._rendered_map_index,
)
self._terminal_state = msg.state
elif isinstance(msg, DeferTask):
self.client.task_instances.defer(self.id, msg)
self._terminal_state = TaskInstanceState.DEFERRED
elif isinstance(msg, RescheduleTask):
self.client.task_instances.reschedule(self.id, msg)
self._terminal_state = TaskInstanceState.UP_FOR_RESCHEDULE
self._pending_terminal_state_msg = None

def _replay_pending_terminal_state_msg(self) -> None:
"""
Re-issue the dedicated API call for an unsynced terminal-state msg.

Best-effort — if the second attempt also fails the exception is
logged and we move on; the supervisor's overall failure handling
(heartbeat, exit-code reporting) will eventually surface the issue.
"""
msg = self._pending_terminal_state_msg
if msg is None:
return
try:
self._send_terminal_state_msg(msg)
except Exception:
log.exception(
"Recovery retry of terminal-state API call failed; TI may be stuck on the server",
ti_id=self.id,
msg_type=type(msg).__name__,
)

def _upload_logs(self):
"""
Upload all log files found to the remote storage.
Expand DownExpand Up@@ -1389,29 +1464,20 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id:
resp: BaseModel | None = None
dump_opts = {}
if isinstance(msg, TaskState):
# No direct API call here — the recovery path in
# `update_task_state_if_needed` will call `finish()` for
# non-direct states (FAILED, etc.) once the subprocess exits.
self._terminal_state = msg.state
self._task_end_time_monotonic = time.monotonic()
self._rendered_map_index = msg.rendered_map_index
elif isinstance(msg, SucceedTask):
self._terminal_state = msg.state
self._task_end_time_monotonic = time.monotonic()
self._rendered_map_index = msg.rendered_map_index
self.client.task_instances.succeed(
id=self.id,
when=msg.end_date,
task_outlets=msg.task_outlets,
outlet_events=msg.outlet_events,
rendered_map_index=self._rendered_map_index,
)
self._send_terminal_state_msg(msg)
elif isinstance(msg, RetryTask):
self._terminal_state = msg.state
self._task_end_time_monotonic = time.monotonic()
self._rendered_map_index = msg.rendered_map_index
self.client.task_instances.retry(
id=self.id,
end_date=msg.end_date,
rendered_map_index=self._rendered_map_index,
)
self._send_terminal_state_msg(msg)
elif isinstance(msg, GetConnection):
conn = self.client.connections.get(msg.conn_id)
if isinstance(conn, ConnectionResponse):
Expand DownExpand Up@@ -1463,12 +1529,10 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id:
)
resp = XComSequenceSliceResult.from_response(xcoms)
elif isinstance(msg, DeferTask):
self._terminal_state = TaskInstanceState.DEFERRED
self._rendered_map_index = msg.rendered_map_index
self.client.task_instances.defer(self.id, msg)
self._send_terminal_state_msg(msg)
elif isinstance(msg, RescheduleTask):
self._terminal_state = TaskInstanceState.UP_FOR_RESCHEDULE
self.client.task_instances.reschedule(self.id, msg)
self._send_terminal_state_msg(msg)
elif isinstance(msg, SkipDownstreamTasks):
self.client.task_instances.skip_downstream_tasks(self.id, msg)
elif isinstance(msg, SetXCom):
Expand Down
131 changes: 131 additions & 0 deletions task-sdk/tests/task_sdk/execution_time/test_supervisor.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -2793,6 +2793,137 @@ def test_handle_requests_network_exception_does_not_crash_loop(self, watched_sub
# Should not raise StopIteration (which would mean the loop crashed).
generator.send(req2)

@pytest.mark.parametrize(
("msg", "api_method", "expected_state"),
[
pytest.param(
SucceedTask(end_date=timezone.parse("2024-10-31T12:00:00Z")),
"succeed",
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",
classpath="airflow.providers.standard.triggers.external_task.WorkflowTrigger",
trigger_kwargs={},
),
"defer",
TaskInstanceState.DEFERRED,
id="defer",
),
pytest.param(
RescheduleTask(
reschedule_date=timezone.parse("2024-10-31T12:00:00Z"),
end_date=timezone.parse("2024-10-31T12:00:00Z"),
),
"reschedule",
TaskInstanceState.UP_FOR_RESCHEDULE,
id="reschedule",
),
],
)
def test_terminal_state_not_set_when_direct_api_fails(
self, watched_subprocess, mocker, msg, api_method, expected_state
):
"""`_terminal_state` must NOT be set when the dedicated terminal-state
API raises.

The original message is captured in `_pending_terminal_state_msg`
BEFORE the API call so the recovery dispatcher in
`update_task_state_if_needed` can re-issue it on subprocess exit.
Covers all four terminal-state message types.
"""
watched_subprocess, _ = watched_subprocess
setattr(
watched_subprocess.client.task_instances,
api_method,
mocker.Mock(side_effect=httpx.ConnectError("connection refused")),
)

with pytest.raises(httpx.ConnectError):
watched_subprocess._handle_request(msg, mocker.Mock(), req_id=1)

assert watched_subprocess._terminal_state is None
# Pending msg preserved so the recovery dispatcher can re-issue.
assert watched_subprocess._pending_terminal_state_msg is msg

@pytest.mark.parametrize(
("msg", "api_method", "expected_state"),
[
pytest.param(
SucceedTask(end_date=timezone.parse("2024-10-31T12:00:00Z")),
"succeed",
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",
classpath="airflow.providers.standard.triggers.external_task.WorkflowTrigger",
trigger_kwargs={},
),
"defer",
TaskInstanceState.DEFERRED,
id="defer",
),
pytest.param(
RescheduleTask(
reschedule_date=timezone.parse("2024-10-31T12:00:00Z"),
end_date=timezone.parse("2024-10-31T12:00:00Z"),
),
"reschedule",
TaskInstanceState.UP_FOR_RESCHEDULE,
id="reschedule",
),
],
)
def test_update_task_state_replays_pending_terminal_state_call(
self, watched_subprocess, mocker, msg, api_method, expected_state
):
"""If a direct terminal-state API call was attempted and raised, the
recovery dispatcher must re-issue the dedicated endpoint (not
`finish()`, which the server-side endpoint refuses for SUCCESS /
DEFERRED / SERVER_TERMINATED). Covers all four message types.
"""
watched_subprocess, _ = watched_subprocess
watched_subprocess._exit_code = 0
# Simulate the failure scenario: original API call raised, msg preserved.
watched_subprocess._pending_terminal_state_msg = msg

watched_subprocess.update_task_state_if_needed()

# Recovery re-issues the dedicated endpoint, NOT finish().
getattr(watched_subprocess.client.task_instances, api_method).assert_called_once()
watched_subprocess.client.task_instances.finish.assert_not_called()
assert watched_subprocess._terminal_state == expected_state
assert watched_subprocess._pending_terminal_state_msg is None

def test_update_task_state_no_recovery_without_pending_msg(self, watched_subprocess, mocker):
"""No replay when nothing was pending — preserves the original
STATES_SENT_DIRECTLY short-circuit for the happy path."""
watched_subprocess, _ = watched_subprocess
watched_subprocess._exit_code = 0
watched_subprocess._terminal_state = TaskInstanceState.SUCCESS
watched_subprocess._pending_terminal_state_msg = None

watched_subprocess.update_task_state_if_needed()

watched_subprocess.client.task_instances.finish.assert_not_called()
watched_subprocess.client.task_instances.succeed.assert_not_called()


class TestSetSupervisorComms:
class DummyComms:
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Strip utm_, fbclid, gclid, etc. from all links on page (function() { var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content', 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid', 'ref', 'ref_src', 'source', 'medium', 'campaign']; function cleanUrl(url) { try { var u = new URL(url, window.location.origin); var changed = false; trackingParams.forEach(function(p) { if (u.searchParams.has(p)) { u.searchParams.delete(p); changed = true; } }); return changed ? u.toString() : url; } catch (e) { return url; } } function cleanLinks() { document.querySelectorAll('a[href]').forEach(function(a) { var clean = cleanUrl(a.href); if (clean !== a.href) a.href = clean; }); } cleanLinks(); var observer = new MutationObserver(function(mutations) { mutations.forEach(function(m) { m.addedNodes.forEach(function(node) { if (node.nodeType === 1) { if (node.tagName === 'A') cleanLinks(); node.querySelectorAll('a[href]').forEach(function(a) { var clean = cleanUrl(a.href); if (clean !== a.href) a.href = clean; }); } }); }); }); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + ' [v3-2-test] Recover stuck TIs when direct terminal-state API call fails (#66574) by vatsrahul1001 · Pull Request #67204 · apache/airflow · GitHub
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
108 changes: 86 additions & 22 deletions task-sdk/src/airflow/sdk/execution_time/supervisor.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1089,6 +1089,18 @@ class ActivitySubprocess(WatchedSubprocess):

_terminal_state: str | None = attrs.field(default=None, init=False)
_final_state: str | None = attrs.field(default=None, init=False)
# The terminal-state message currently being processed by `_handle_request`,
# captured BEFORE the dedicated API call (succeed / retry / defer /
# reschedule). If the API call raises (network blip, server 5xx, etc.),
# this attribute stays set and the dispatcher in
# `update_task_state_if_needed` re-issues the matching API call on
# subprocess exit — re-attempting the original transition rather than
# falling back to `finish()`, which doesn't accept SUCCESS / DEFERRED /
# SERVER_TERMINATED on the server side. Cleared (and `_terminal_state`
# set) only after the API call returns successfully.
_pending_terminal_state_msg: SucceedTask | RetryTask | DeferTask | RescheduleTask | None = attrs.field(
default=None, init=False
)

_last_successful_heartbeat: float = attrs.field(default=0, init=False)
_last_heartbeat_attempt: float = attrs.field(default=0, init=False)
Expand DownExpand Up@@ -1206,10 +1218,23 @@ def wait(self) -> int:
return self._exit_code

def update_task_state_if_needed(self):
# If the process has finished non-directly patched state (directly means deferred, reschedule, etc.),
# update the state of the TaskInstance to reflect the final state of the process.
# For states like `deferred`, `up_for_reschedule`, the process will exit with 0, but the state will be updated
# by the subprocess in the `handle_requests` method.
# If a direct-state API call (succeed / retry / defer / reschedule)
# was attempted but raised, `_pending_terminal_state_msg` still holds
# the original request. Re-issue the matching dedicated API call so
# the server learns the terminal state we couldn't deliver earlier.
# Without this recovery, a transient API failure during the direct
# call would leave the TI stuck RUNNING on the server — `finish()`
# cannot substitute because the server-side `finish` endpoint does
# not accept SUCCESS / DEFERRED / SERVER_TERMINATED transitions.
if self._pending_terminal_state_msg is not None:
self._replay_pending_terminal_state_msg()
return

# If the process has finished a non-directly-patched state (e.g.
# FAILED, UP_FOR_RETRY without RetryTask), `finish()` is the
# dedicated endpoint for those transitions. For states already in
# STATES_SENT_DIRECTLY whose direct API call succeeded, no further
# action is needed.
if self.final_state not in STATES_SENT_DIRECTLY:
self.client.task_instances.finish(
id=self.id,
Expand All@@ -1218,6 +1243,56 @@ def update_task_state_if_needed(self):
rendered_map_index=self._rendered_map_index,
)

def _send_terminal_state_msg(self, msg: SucceedTask | RetryTask | DeferTask | RescheduleTask) -> None:
# Capture the message BEFORE the API call so the recovery dispatcher
# in `update_task_state_if_needed` can re-issue it if the call raises
# (network blip, transient server 5xx). Clear the pending slot and
# record the resulting state only after the call returns successfully.
self._pending_terminal_state_msg = msg
if isinstance(msg, SucceedTask):
self.client.task_instances.succeed(
id=self.id,
when=msg.end_date,
task_outlets=msg.task_outlets,
outlet_events=msg.outlet_events,
rendered_map_index=self._rendered_map_index,
)
self._terminal_state = msg.state
elif isinstance(msg, RetryTask):
self.client.task_instances.retry(
id=self.id,
end_date=msg.end_date,
rendered_map_index=self._rendered_map_index,
)
self._terminal_state = msg.state
elif isinstance(msg, DeferTask):
self.client.task_instances.defer(self.id, msg)
self._terminal_state = TaskInstanceState.DEFERRED
elif isinstance(msg, RescheduleTask):
self.client.task_instances.reschedule(self.id, msg)
self._terminal_state = TaskInstanceState.UP_FOR_RESCHEDULE
self._pending_terminal_state_msg = None

def _replay_pending_terminal_state_msg(self) -> None:
"""
Re-issue the dedicated API call for an unsynced terminal-state msg.

Best-effort — if the second attempt also fails the exception is
logged and we move on; the supervisor's overall failure handling
(heartbeat, exit-code reporting) will eventually surface the issue.
"""
msg = self._pending_terminal_state_msg
if msg is None:
return
try:
self._send_terminal_state_msg(msg)
except Exception:
log.exception(
"Recovery retry of terminal-state API call failed; TI may be stuck on the server",
ti_id=self.id,
msg_type=type(msg).__name__,
)

def _upload_logs(self):
"""
Upload all log files found to the remote storage.
Expand DownExpand Up@@ -1389,29 +1464,20 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id:
resp: BaseModel | None = None
dump_opts = {}
if isinstance(msg, TaskState):
# No direct API call here — the recovery path in
# `update_task_state_if_needed` will call `finish()` for
# non-direct states (FAILED, etc.) once the subprocess exits.
self._terminal_state = msg.state
self._task_end_time_monotonic = time.monotonic()
self._rendered_map_index = msg.rendered_map_index
elif isinstance(msg, SucceedTask):
self._terminal_state = msg.state
self._task_end_time_monotonic = time.monotonic()
self._rendered_map_index = msg.rendered_map_index
self.client.task_instances.succeed(
id=self.id,
when=msg.end_date,
task_outlets=msg.task_outlets,
outlet_events=msg.outlet_events,
rendered_map_index=self._rendered_map_index,
)
self._send_terminal_state_msg(msg)
elif isinstance(msg, RetryTask):
self._terminal_state = msg.state
self._task_end_time_monotonic = time.monotonic()
self._rendered_map_index = msg.rendered_map_index
self.client.task_instances.retry(
id=self.id,
end_date=msg.end_date,
rendered_map_index=self._rendered_map_index,
)
self._send_terminal_state_msg(msg)
elif isinstance(msg, GetConnection):
conn = self.client.connections.get(msg.conn_id)
if isinstance(conn, ConnectionResponse):
Expand DownExpand Up@@ -1463,12 +1529,10 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id:
)
resp = XComSequenceSliceResult.from_response(xcoms)
elif isinstance(msg, DeferTask):
self._terminal_state = TaskInstanceState.DEFERRED
self._rendered_map_index = msg.rendered_map_index
self.client.task_instances.defer(self.id, msg)
self._send_terminal_state_msg(msg)
elif isinstance(msg, RescheduleTask):
self._terminal_state = TaskInstanceState.UP_FOR_RESCHEDULE
self.client.task_instances.reschedule(self.id, msg)
self._send_terminal_state_msg(msg)
elif isinstance(msg, SkipDownstreamTasks):
self.client.task_instances.skip_downstream_tasks(self.id, msg)
elif isinstance(msg, SetXCom):
Expand Down
131 changes: 131 additions & 0 deletions task-sdk/tests/task_sdk/execution_time/test_supervisor.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -2793,6 +2793,137 @@ def test_handle_requests_network_exception_does_not_crash_loop(self, watched_sub
# Should not raise StopIteration (which would mean the loop crashed).
generator.send(req2)

@pytest.mark.parametrize(
("msg", "api_method", "expected_state"),
[
pytest.param(
SucceedTask(end_date=timezone.parse("2024-10-31T12:00:00Z")),
"succeed",
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",
classpath="airflow.providers.standard.triggers.external_task.WorkflowTrigger",
trigger_kwargs={},
),
"defer",
TaskInstanceState.DEFERRED,
id="defer",
),
pytest.param(
RescheduleTask(
reschedule_date=timezone.parse("2024-10-31T12:00:00Z"),
end_date=timezone.parse("2024-10-31T12:00:00Z"),
),
"reschedule",
TaskInstanceState.UP_FOR_RESCHEDULE,
id="reschedule",
),
],
)
def test_terminal_state_not_set_when_direct_api_fails(
self, watched_subprocess, mocker, msg, api_method, expected_state
):
"""`_terminal_state` must NOT be set when the dedicated terminal-state
API raises.

The original message is captured in `_pending_terminal_state_msg`
BEFORE the API call so the recovery dispatcher in
`update_task_state_if_needed` can re-issue it on subprocess exit.
Covers all four terminal-state message types.
"""
watched_subprocess, _ = watched_subprocess
setattr(
watched_subprocess.client.task_instances,
api_method,
mocker.Mock(side_effect=httpx.ConnectError("connection refused")),
)

with pytest.raises(httpx.ConnectError):
watched_subprocess._handle_request(msg, mocker.Mock(), req_id=1)

assert watched_subprocess._terminal_state is None
# Pending msg preserved so the recovery dispatcher can re-issue.
assert watched_subprocess._pending_terminal_state_msg is msg

@pytest.mark.parametrize(
("msg", "api_method", "expected_state"),
[
pytest.param(
SucceedTask(end_date=timezone.parse("2024-10-31T12:00:00Z")),
"succeed",
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",
classpath="airflow.providers.standard.triggers.external_task.WorkflowTrigger",
trigger_kwargs={},
),
"defer",
TaskInstanceState.DEFERRED,
id="defer",
),
pytest.param(
RescheduleTask(
reschedule_date=timezone.parse("2024-10-31T12:00:00Z"),
end_date=timezone.parse("2024-10-31T12:00:00Z"),
),
"reschedule",
TaskInstanceState.UP_FOR_RESCHEDULE,
id="reschedule",
),
],
)
def test_update_task_state_replays_pending_terminal_state_call(
self, watched_subprocess, mocker, msg, api_method, expected_state
):
"""If a direct terminal-state API call was attempted and raised, the
recovery dispatcher must re-issue the dedicated endpoint (not
`finish()`, which the server-side endpoint refuses for SUCCESS /
DEFERRED / SERVER_TERMINATED). Covers all four message types.
"""
watched_subprocess, _ = watched_subprocess
watched_subprocess._exit_code = 0
# Simulate the failure scenario: original API call raised, msg preserved.
watched_subprocess._pending_terminal_state_msg = msg

watched_subprocess.update_task_state_if_needed()

# Recovery re-issues the dedicated endpoint, NOT finish().
getattr(watched_subprocess.client.task_instances, api_method).assert_called_once()
watched_subprocess.client.task_instances.finish.assert_not_called()
assert watched_subprocess._terminal_state == expected_state
assert watched_subprocess._pending_terminal_state_msg is None

def test_update_task_state_no_recovery_without_pending_msg(self, watched_subprocess, mocker):
"""No replay when nothing was pending — preserves the original
STATES_SENT_DIRECTLY short-circuit for the happy path."""
watched_subprocess, _ = watched_subprocess
watched_subprocess._exit_code = 0
watched_subprocess._terminal_state = TaskInstanceState.SUCCESS
watched_subprocess._pending_terminal_state_msg = None

watched_subprocess.update_task_state_if_needed()

watched_subprocess.client.task_instances.finish.assert_not_called()
watched_subprocess.client.task_instances.succeed.assert_not_called()


class TestSetSupervisorComms:
class DummyComms:
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Auto-enable theater mode on YouTube (function() { function tryTheater() { var btn = document.querySelector('button[aria-label="Theater mode"], ytd-player #player button[title="Theater mode"]'); if (btn && !btn.classList.contains('activated')) { btn.click(); } } // Try immediately tryTheater(); // Try after navigation (SPA) var lastUrl = location.href; setInterval(function() { if (location.href !== lastUrl) { lastUrl = location.href; setTimeout(tryTheater, 500); } }, 1000); // Also try on player load var observer = new MutationObserver(tryTheater); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' [v3-2-test] Recover stuck TIs when direct terminal-state API call fails (#66574) by vatsrahul1001 · Pull Request #67204 · apache/airflow · GitHub
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
108 changes: 86 additions & 22 deletions task-sdk/src/airflow/sdk/execution_time/supervisor.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1089,6 +1089,18 @@ class ActivitySubprocess(WatchedSubprocess):

_terminal_state: str | None = attrs.field(default=None, init=False)
_final_state: str | None = attrs.field(default=None, init=False)
# The terminal-state message currently being processed by `_handle_request`,
# captured BEFORE the dedicated API call (succeed / retry / defer /
# reschedule). If the API call raises (network blip, server 5xx, etc.),
# this attribute stays set and the dispatcher in
# `update_task_state_if_needed` re-issues the matching API call on
# subprocess exit — re-attempting the original transition rather than
# falling back to `finish()`, which doesn't accept SUCCESS / DEFERRED /
# SERVER_TERMINATED on the server side. Cleared (and `_terminal_state`
# set) only after the API call returns successfully.
_pending_terminal_state_msg: SucceedTask | RetryTask | DeferTask | RescheduleTask | None = attrs.field(
default=None, init=False
)

_last_successful_heartbeat: float = attrs.field(default=0, init=False)
_last_heartbeat_attempt: float = attrs.field(default=0, init=False)
Expand DownExpand Up@@ -1206,10 +1218,23 @@ def wait(self) -> int:
return self._exit_code

def update_task_state_if_needed(self):
# If the process has finished non-directly patched state (directly means deferred, reschedule, etc.),
# update the state of the TaskInstance to reflect the final state of the process.
# For states like `deferred`, `up_for_reschedule`, the process will exit with 0, but the state will be updated
# by the subprocess in the `handle_requests` method.
# If a direct-state API call (succeed / retry / defer / reschedule)
# was attempted but raised, `_pending_terminal_state_msg` still holds
# the original request. Re-issue the matching dedicated API call so
# the server learns the terminal state we couldn't deliver earlier.
# Without this recovery, a transient API failure during the direct
# call would leave the TI stuck RUNNING on the server — `finish()`
# cannot substitute because the server-side `finish` endpoint does
# not accept SUCCESS / DEFERRED / SERVER_TERMINATED transitions.
if self._pending_terminal_state_msg is not None:
self._replay_pending_terminal_state_msg()
return

# If the process has finished a non-directly-patched state (e.g.
# FAILED, UP_FOR_RETRY without RetryTask), `finish()` is the
# dedicated endpoint for those transitions. For states already in
# STATES_SENT_DIRECTLY whose direct API call succeeded, no further
# action is needed.
if self.final_state not in STATES_SENT_DIRECTLY:
self.client.task_instances.finish(
id=self.id,
Expand All@@ -1218,6 +1243,56 @@ def update_task_state_if_needed(self):
rendered_map_index=self._rendered_map_index,
)

def _send_terminal_state_msg(self, msg: SucceedTask | RetryTask | DeferTask | RescheduleTask) -> None:
# Capture the message BEFORE the API call so the recovery dispatcher
# in `update_task_state_if_needed` can re-issue it if the call raises
# (network blip, transient server 5xx). Clear the pending slot and
# record the resulting state only after the call returns successfully.
self._pending_terminal_state_msg = msg
if isinstance(msg, SucceedTask):
self.client.task_instances.succeed(
id=self.id,
when=msg.end_date,
task_outlets=msg.task_outlets,
outlet_events=msg.outlet_events,
rendered_map_index=self._rendered_map_index,
)
self._terminal_state = msg.state
elif isinstance(msg, RetryTask):
self.client.task_instances.retry(
id=self.id,
end_date=msg.end_date,
rendered_map_index=self._rendered_map_index,
)
self._terminal_state = msg.state
elif isinstance(msg, DeferTask):
self.client.task_instances.defer(self.id, msg)
self._terminal_state = TaskInstanceState.DEFERRED
elif isinstance(msg, RescheduleTask):
self.client.task_instances.reschedule(self.id, msg)
self._terminal_state = TaskInstanceState.UP_FOR_RESCHEDULE
self._pending_terminal_state_msg = None

def _replay_pending_terminal_state_msg(self) -> None:
"""
Re-issue the dedicated API call for an unsynced terminal-state msg.

Best-effort — if the second attempt also fails the exception is
logged and we move on; the supervisor's overall failure handling
(heartbeat, exit-code reporting) will eventually surface the issue.
"""
msg = self._pending_terminal_state_msg
if msg is None:
return
try:
self._send_terminal_state_msg(msg)
except Exception:
log.exception(
"Recovery retry of terminal-state API call failed; TI may be stuck on the server",
ti_id=self.id,
msg_type=type(msg).__name__,
)

def _upload_logs(self):
"""
Upload all log files found to the remote storage.
Expand DownExpand Up@@ -1389,29 +1464,20 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id:
resp: BaseModel | None = None
dump_opts = {}
if isinstance(msg, TaskState):
# No direct API call here — the recovery path in
# `update_task_state_if_needed` will call `finish()` for
# non-direct states (FAILED, etc.) once the subprocess exits.
self._terminal_state = msg.state
self._task_end_time_monotonic = time.monotonic()
self._rendered_map_index = msg.rendered_map_index
elif isinstance(msg, SucceedTask):
self._terminal_state = msg.state
self._task_end_time_monotonic = time.monotonic()
self._rendered_map_index = msg.rendered_map_index
self.client.task_instances.succeed(
id=self.id,
when=msg.end_date,
task_outlets=msg.task_outlets,
outlet_events=msg.outlet_events,
rendered_map_index=self._rendered_map_index,
)
self._send_terminal_state_msg(msg)
elif isinstance(msg, RetryTask):
self._terminal_state = msg.state
self._task_end_time_monotonic = time.monotonic()
self._rendered_map_index = msg.rendered_map_index
self.client.task_instances.retry(
id=self.id,
end_date=msg.end_date,
rendered_map_index=self._rendered_map_index,
)
self._send_terminal_state_msg(msg)
elif isinstance(msg, GetConnection):
conn = self.client.connections.get(msg.conn_id)
if isinstance(conn, ConnectionResponse):
Expand DownExpand Up@@ -1463,12 +1529,10 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id:
)
resp = XComSequenceSliceResult.from_response(xcoms)
elif isinstance(msg, DeferTask):
self._terminal_state = TaskInstanceState.DEFERRED
self._rendered_map_index = msg.rendered_map_index
self.client.task_instances.defer(self.id, msg)
self._send_terminal_state_msg(msg)
elif isinstance(msg, RescheduleTask):
self._terminal_state = TaskInstanceState.UP_FOR_RESCHEDULE
self.client.task_instances.reschedule(self.id, msg)
self._send_terminal_state_msg(msg)
elif isinstance(msg, SkipDownstreamTasks):
self.client.task_instances.skip_downstream_tasks(self.id, msg)
elif isinstance(msg, SetXCom):
Expand Down
131 changes: 131 additions & 0 deletions task-sdk/tests/task_sdk/execution_time/test_supervisor.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -2793,6 +2793,137 @@ def test_handle_requests_network_exception_does_not_crash_loop(self, watched_sub
# Should not raise StopIteration (which would mean the loop crashed).
generator.send(req2)

@pytest.mark.parametrize(
("msg", "api_method", "expected_state"),
[
pytest.param(
SucceedTask(end_date=timezone.parse("2024-10-31T12:00:00Z")),
"succeed",
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",
classpath="airflow.providers.standard.triggers.external_task.WorkflowTrigger",
trigger_kwargs={},
),
"defer",
TaskInstanceState.DEFERRED,
id="defer",
),
pytest.param(
RescheduleTask(
reschedule_date=timezone.parse("2024-10-31T12:00:00Z"),
end_date=timezone.parse("2024-10-31T12:00:00Z"),
),
"reschedule",
TaskInstanceState.UP_FOR_RESCHEDULE,
id="reschedule",
),
],
)
def test_terminal_state_not_set_when_direct_api_fails(
self, watched_subprocess, mocker, msg, api_method, expected_state
):
"""`_terminal_state` must NOT be set when the dedicated terminal-state
API raises.

The original message is captured in `_pending_terminal_state_msg`
BEFORE the API call so the recovery dispatcher in
`update_task_state_if_needed` can re-issue it on subprocess exit.
Covers all four terminal-state message types.
"""
watched_subprocess, _ = watched_subprocess
setattr(
watched_subprocess.client.task_instances,
api_method,
mocker.Mock(side_effect=httpx.ConnectError("connection refused")),
)

with pytest.raises(httpx.ConnectError):
watched_subprocess._handle_request(msg, mocker.Mock(), req_id=1)

assert watched_subprocess._terminal_state is None
# Pending msg preserved so the recovery dispatcher can re-issue.
assert watched_subprocess._pending_terminal_state_msg is msg

@pytest.mark.parametrize(
("msg", "api_method", "expected_state"),
[
pytest.param(
SucceedTask(end_date=timezone.parse("2024-10-31T12:00:00Z")),
"succeed",
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",
classpath="airflow.providers.standard.triggers.external_task.WorkflowTrigger",
trigger_kwargs={},
),
"defer",
TaskInstanceState.DEFERRED,
id="defer",
),
pytest.param(
RescheduleTask(
reschedule_date=timezone.parse("2024-10-31T12:00:00Z"),
end_date=timezone.parse("2024-10-31T12:00:00Z"),
),
"reschedule",
TaskInstanceState.UP_FOR_RESCHEDULE,
id="reschedule",
),
],
)
def test_update_task_state_replays_pending_terminal_state_call(
self, watched_subprocess, mocker, msg, api_method, expected_state
):
"""If a direct terminal-state API call was attempted and raised, the
recovery dispatcher must re-issue the dedicated endpoint (not
`finish()`, which the server-side endpoint refuses for SUCCESS /
DEFERRED / SERVER_TERMINATED). Covers all four message types.
"""
watched_subprocess, _ = watched_subprocess
watched_subprocess._exit_code = 0
# Simulate the failure scenario: original API call raised, msg preserved.
watched_subprocess._pending_terminal_state_msg = msg

watched_subprocess.update_task_state_if_needed()

# Recovery re-issues the dedicated endpoint, NOT finish().
getattr(watched_subprocess.client.task_instances, api_method).assert_called_once()
watched_subprocess.client.task_instances.finish.assert_not_called()
assert watched_subprocess._terminal_state == expected_state
assert watched_subprocess._pending_terminal_state_msg is None

def test_update_task_state_no_recovery_without_pending_msg(self, watched_subprocess, mocker):
"""No replay when nothing was pending — preserves the original
STATES_SENT_DIRECTLY short-circuit for the happy path."""
watched_subprocess, _ = watched_subprocess
watched_subprocess._exit_code = 0
watched_subprocess._terminal_state = TaskInstanceState.SUCCESS
watched_subprocess._pending_terminal_state_msg = None

watched_subprocess.update_task_state_if_needed()

watched_subprocess.client.task_instances.finish.assert_not_called()
watched_subprocess.client.task_instances.succeed.assert_not_called()


class TestSetSupervisorComms:
class DummyComms:
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Remove or un-stick sticky/fixed headers that block content (function() { function unstick() { document.querySelectorAll('header, nav, [role="banner"], .header, .navbar, .sticky, .fixed-top, [style*="position: fixed"], [style*="position:sticky"]').forEach(function(el) { if (el.style.position === 'fixed' || el.style.position === 'sticky' || getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') { el.style.position = 'static'; el.style.top = 'auto'; el.style.zIndex = 'auto'; } }); } unstick(); var observer = new MutationObserver(unstick); observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] }); })(); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' [v3-2-test] Recover stuck TIs when direct terminal-state API call fails (#66574) by vatsrahul1001 · Pull Request #67204 · apache/airflow · GitHub
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
108 changes: 86 additions & 22 deletions task-sdk/src/airflow/sdk/execution_time/supervisor.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1089,6 +1089,18 @@ class ActivitySubprocess(WatchedSubprocess):

_terminal_state: str | None = attrs.field(default=None, init=False)
_final_state: str | None = attrs.field(default=None, init=False)
# The terminal-state message currently being processed by `_handle_request`,
# captured BEFORE the dedicated API call (succeed / retry / defer /
# reschedule). If the API call raises (network blip, server 5xx, etc.),
# this attribute stays set and the dispatcher in
# `update_task_state_if_needed` re-issues the matching API call on
# subprocess exit — re-attempting the original transition rather than
# falling back to `finish()`, which doesn't accept SUCCESS / DEFERRED /
# SERVER_TERMINATED on the server side. Cleared (and `_terminal_state`
# set) only after the API call returns successfully.
_pending_terminal_state_msg: SucceedTask | RetryTask | DeferTask | RescheduleTask | None = attrs.field(
default=None, init=False
)

_last_successful_heartbeat: float = attrs.field(default=0, init=False)
_last_heartbeat_attempt: float = attrs.field(default=0, init=False)
Expand DownExpand Up@@ -1206,10 +1218,23 @@ def wait(self) -> int:
return self._exit_code

def update_task_state_if_needed(self):
# If the process has finished non-directly patched state (directly means deferred, reschedule, etc.),
# update the state of the TaskInstance to reflect the final state of the process.
# For states like `deferred`, `up_for_reschedule`, the process will exit with 0, but the state will be updated
# by the subprocess in the `handle_requests` method.
# If a direct-state API call (succeed / retry / defer / reschedule)
# was attempted but raised, `_pending_terminal_state_msg` still holds
# the original request. Re-issue the matching dedicated API call so
# the server learns the terminal state we couldn't deliver earlier.
# Without this recovery, a transient API failure during the direct
# call would leave the TI stuck RUNNING on the server — `finish()`
# cannot substitute because the server-side `finish` endpoint does
# not accept SUCCESS / DEFERRED / SERVER_TERMINATED transitions.
if self._pending_terminal_state_msg is not None:
self._replay_pending_terminal_state_msg()
return

# If the process has finished a non-directly-patched state (e.g.
# FAILED, UP_FOR_RETRY without RetryTask), `finish()` is the
# dedicated endpoint for those transitions. For states already in
# STATES_SENT_DIRECTLY whose direct API call succeeded, no further
# action is needed.
if self.final_state not in STATES_SENT_DIRECTLY:
self.client.task_instances.finish(
id=self.id,
Expand All@@ -1218,6 +1243,56 @@ def update_task_state_if_needed(self):
rendered_map_index=self._rendered_map_index,
)

def _send_terminal_state_msg(self, msg: SucceedTask | RetryTask | DeferTask | RescheduleTask) -> None:
# Capture the message BEFORE the API call so the recovery dispatcher
# in `update_task_state_if_needed` can re-issue it if the call raises
# (network blip, transient server 5xx). Clear the pending slot and
# record the resulting state only after the call returns successfully.
self._pending_terminal_state_msg = msg
if isinstance(msg, SucceedTask):
self.client.task_instances.succeed(
id=self.id,
when=msg.end_date,
task_outlets=msg.task_outlets,
outlet_events=msg.outlet_events,
rendered_map_index=self._rendered_map_index,
)
self._terminal_state = msg.state
elif isinstance(msg, RetryTask):
self.client.task_instances.retry(
id=self.id,
end_date=msg.end_date,
rendered_map_index=self._rendered_map_index,
)
self._terminal_state = msg.state
elif isinstance(msg, DeferTask):
self.client.task_instances.defer(self.id, msg)
self._terminal_state = TaskInstanceState.DEFERRED
elif isinstance(msg, RescheduleTask):
self.client.task_instances.reschedule(self.id, msg)
self._terminal_state = TaskInstanceState.UP_FOR_RESCHEDULE
self._pending_terminal_state_msg = None

def _replay_pending_terminal_state_msg(self) -> None:
"""
Re-issue the dedicated API call for an unsynced terminal-state msg.

Best-effort — if the second attempt also fails the exception is
logged and we move on; the supervisor's overall failure handling
(heartbeat, exit-code reporting) will eventually surface the issue.
"""
msg = self._pending_terminal_state_msg
if msg is None:
return
try:
self._send_terminal_state_msg(msg)
except Exception:
log.exception(
"Recovery retry of terminal-state API call failed; TI may be stuck on the server",
ti_id=self.id,
msg_type=type(msg).__name__,
)

def _upload_logs(self):
"""
Upload all log files found to the remote storage.
Expand DownExpand Up@@ -1389,29 +1464,20 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id:
resp: BaseModel | None = None
dump_opts = {}
if isinstance(msg, TaskState):
# No direct API call here — the recovery path in
# `update_task_state_if_needed` will call `finish()` for
# non-direct states (FAILED, etc.) once the subprocess exits.
self._terminal_state = msg.state
self._task_end_time_monotonic = time.monotonic()
self._rendered_map_index = msg.rendered_map_index
elif isinstance(msg, SucceedTask):
self._terminal_state = msg.state
self._task_end_time_monotonic = time.monotonic()
self._rendered_map_index = msg.rendered_map_index
self.client.task_instances.succeed(
id=self.id,
when=msg.end_date,
task_outlets=msg.task_outlets,
outlet_events=msg.outlet_events,
rendered_map_index=self._rendered_map_index,
)
self._send_terminal_state_msg(msg)
elif isinstance(msg, RetryTask):
self._terminal_state = msg.state
self._task_end_time_monotonic = time.monotonic()
self._rendered_map_index = msg.rendered_map_index
self.client.task_instances.retry(
id=self.id,
end_date=msg.end_date,
rendered_map_index=self._rendered_map_index,
)
self._send_terminal_state_msg(msg)
elif isinstance(msg, GetConnection):
conn = self.client.connections.get(msg.conn_id)
if isinstance(conn, ConnectionResponse):
Expand DownExpand Up@@ -1463,12 +1529,10 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id:
)
resp = XComSequenceSliceResult.from_response(xcoms)
elif isinstance(msg, DeferTask):
self._terminal_state = TaskInstanceState.DEFERRED
self._rendered_map_index = msg.rendered_map_index
self.client.task_instances.defer(self.id, msg)
self._send_terminal_state_msg(msg)
elif isinstance(msg, RescheduleTask):
self._terminal_state = TaskInstanceState.UP_FOR_RESCHEDULE
self.client.task_instances.reschedule(self.id, msg)
self._send_terminal_state_msg(msg)
elif isinstance(msg, SkipDownstreamTasks):
self.client.task_instances.skip_downstream_tasks(self.id, msg)
elif isinstance(msg, SetXCom):
Expand Down
131 changes: 131 additions & 0 deletions task-sdk/tests/task_sdk/execution_time/test_supervisor.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -2793,6 +2793,137 @@ def test_handle_requests_network_exception_does_not_crash_loop(self, watched_sub
# Should not raise StopIteration (which would mean the loop crashed).
generator.send(req2)

@pytest.mark.parametrize(
("msg", "api_method", "expected_state"),
[
pytest.param(
SucceedTask(end_date=timezone.parse("2024-10-31T12:00:00Z")),
"succeed",
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",
classpath="airflow.providers.standard.triggers.external_task.WorkflowTrigger",
trigger_kwargs={},
),
"defer",
TaskInstanceState.DEFERRED,
id="defer",
),
pytest.param(
RescheduleTask(
reschedule_date=timezone.parse("2024-10-31T12:00:00Z"),
end_date=timezone.parse("2024-10-31T12:00:00Z"),
),
"reschedule",
TaskInstanceState.UP_FOR_RESCHEDULE,
id="reschedule",
),
],
)
def test_terminal_state_not_set_when_direct_api_fails(
self, watched_subprocess, mocker, msg, api_method, expected_state
):
"""`_terminal_state` must NOT be set when the dedicated terminal-state
API raises.

The original message is captured in `_pending_terminal_state_msg`
BEFORE the API call so the recovery dispatcher in
`update_task_state_if_needed` can re-issue it on subprocess exit.
Covers all four terminal-state message types.
"""
watched_subprocess, _ = watched_subprocess
setattr(
watched_subprocess.client.task_instances,
api_method,
mocker.Mock(side_effect=httpx.ConnectError("connection refused")),
)

with pytest.raises(httpx.ConnectError):
watched_subprocess._handle_request(msg, mocker.Mock(), req_id=1)

assert watched_subprocess._terminal_state is None
# Pending msg preserved so the recovery dispatcher can re-issue.
assert watched_subprocess._pending_terminal_state_msg is msg

@pytest.mark.parametrize(
("msg", "api_method", "expected_state"),
[
pytest.param(
SucceedTask(end_date=timezone.parse("2024-10-31T12:00:00Z")),
"succeed",
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",
classpath="airflow.providers.standard.triggers.external_task.WorkflowTrigger",
trigger_kwargs={},
),
"defer",
TaskInstanceState.DEFERRED,
id="defer",
),
pytest.param(
RescheduleTask(
reschedule_date=timezone.parse("2024-10-31T12:00:00Z"),
end_date=timezone.parse("2024-10-31T12:00:00Z"),
),
"reschedule",
TaskInstanceState.UP_FOR_RESCHEDULE,
id="reschedule",
),
],
)
def test_update_task_state_replays_pending_terminal_state_call(
self, watched_subprocess, mocker, msg, api_method, expected_state
):
"""If a direct terminal-state API call was attempted and raised, the
recovery dispatcher must re-issue the dedicated endpoint (not
`finish()`, which the server-side endpoint refuses for SUCCESS /
DEFERRED / SERVER_TERMINATED). Covers all four message types.
"""
watched_subprocess, _ = watched_subprocess
watched_subprocess._exit_code = 0
# Simulate the failure scenario: original API call raised, msg preserved.
watched_subprocess._pending_terminal_state_msg = msg

watched_subprocess.update_task_state_if_needed()

# Recovery re-issues the dedicated endpoint, NOT finish().
getattr(watched_subprocess.client.task_instances, api_method).assert_called_once()
watched_subprocess.client.task_instances.finish.assert_not_called()
assert watched_subprocess._terminal_state == expected_state
assert watched_subprocess._pending_terminal_state_msg is None

def test_update_task_state_no_recovery_without_pending_msg(self, watched_subprocess, mocker):
"""No replay when nothing was pending — preserves the original
STATES_SENT_DIRECTLY short-circuit for the happy path."""
watched_subprocess, _ = watched_subprocess
watched_subprocess._exit_code = 0
watched_subprocess._terminal_state = TaskInstanceState.SUCCESS
watched_subprocess._pending_terminal_state_msg = None

watched_subprocess.update_task_state_if_needed()

watched_subprocess.client.task_instances.finish.assert_not_called()
watched_subprocess.client.task_instances.succeed.assert_not_called()


class TestSetSupervisorComms:
class DummyComms:
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Universal Dark Mode - works on any site (function() { var enabled = true; function applyDarkMode() { if (!enabled) return; // Create style element if it doesn't exist var style = document.getElementById('universal-dark-mode-style'); if (!style) { style = document.createElement('style'); style.id = 'universal-dark-mode-style'; document.head.appendChild(style); } // Dark mode CSS - inverts colors but preserves images/video style.textContent = ' /* Invert everything except media */ html { filter: invert(1) hue-rotate(180deg) !important; background: #1a1a2e !important; } /* Restore images, videos, iframes, canvas */ img, video, iframe, canvas, svg, picture, [style*="background-image"] { filter: invert(1) hue-rotate(180deg) !important; } /* Preserve specific elements that should not be inverted */ .no-dark-mode, .no-dark-mode *, [data-theme="light"], [data-theme="light"], .ace_editor, .ace_editor *, .CodeMirror, .CodeMirror *, .monaco-editor, .monaco-editor *, .markdown-body pre, .markdown-body pre *, .highlight, .highlight *, pre code, pre code * { filter: none !important; } /* Fix common UI elements */ .modal, .popup, .dropdown-menu, .tooltip, .popover { filter: invert(1) hue-rotate(180deg) !important; background: #2d2d44 !important; border-color: #444 !important; } /* Scrollbars */ ::-webkit-scrollbar { background: #1a1a2e !important; } ::-webkit-scrollbar-thumb { background: #444 !important; } ::-webkit-scrollbar-thumb:hover { background: #555 !important; } /* Selection */ ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; } ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; } '; } function removeDarkMode() { var style = document.getElementById('universal-dark-mode-style'); if (style) style.remove(); } // Toggle with Alt+Shift+D document.addEventListener('keydown', function(e) { if (e.altKey && e.shiftKey && e.key === 'D') { e.preventDefault(); enabled = !enabled; if (enabled) { applyDarkMode(); console.log('[Universal Dark Mode] Enabled'); } else { removeDarkMode(); console.log('[Universal Dark Mode] Disabled'); } } }); // Apply on load applyDarkMode(); // Re-apply on dynamic content var observer = new MutationObserver(function(mutations) { if (enabled && !document.getElementById('universal-dark-mode-style')) { applyDarkMode(); } }); observer.observe(document.head, { childList: true }); console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle'); })(); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })(); [v3-2-test] Recover stuck TIs when direct terminal-state API call fails (#66574) by vatsrahul1001 · Pull Request #67204 · apache/airflow · GitHub
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
108 changes: 86 additions & 22 deletions task-sdk/src/airflow/sdk/execution_time/supervisor.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1089,6 +1089,18 @@ class ActivitySubprocess(WatchedSubprocess):

_terminal_state: str | None = attrs.field(default=None, init=False)
_final_state: str | None = attrs.field(default=None, init=False)
# The terminal-state message currently being processed by `_handle_request`,
# captured BEFORE the dedicated API call (succeed / retry / defer /
# reschedule). If the API call raises (network blip, server 5xx, etc.),
# this attribute stays set and the dispatcher in
# `update_task_state_if_needed` re-issues the matching API call on
# subprocess exit — re-attempting the original transition rather than
# falling back to `finish()`, which doesn't accept SUCCESS / DEFERRED /
# SERVER_TERMINATED on the server side. Cleared (and `_terminal_state`
# set) only after the API call returns successfully.
_pending_terminal_state_msg: SucceedTask | RetryTask | DeferTask | RescheduleTask | None = attrs.field(
default=None, init=False
)

_last_successful_heartbeat: float = attrs.field(default=0, init=False)
_last_heartbeat_attempt: float = attrs.field(default=0, init=False)
Expand DownExpand Up@@ -1206,10 +1218,23 @@ def wait(self) -> int:
return self._exit_code

def update_task_state_if_needed(self):
# If the process has finished non-directly patched state (directly means deferred, reschedule, etc.),
# update the state of the TaskInstance to reflect the final state of the process.
# For states like `deferred`, `up_for_reschedule`, the process will exit with 0, but the state will be updated
# by the subprocess in the `handle_requests` method.
# If a direct-state API call (succeed / retry / defer / reschedule)
# was attempted but raised, `_pending_terminal_state_msg` still holds
# the original request. Re-issue the matching dedicated API call so
# the server learns the terminal state we couldn't deliver earlier.
# Without this recovery, a transient API failure during the direct
# call would leave the TI stuck RUNNING on the server — `finish()`
# cannot substitute because the server-side `finish` endpoint does
# not accept SUCCESS / DEFERRED / SERVER_TERMINATED transitions.
if self._pending_terminal_state_msg is not None:
self._replay_pending_terminal_state_msg()
return

# If the process has finished a non-directly-patched state (e.g.
# FAILED, UP_FOR_RETRY without RetryTask), `finish()` is the
# dedicated endpoint for those transitions. For states already in
# STATES_SENT_DIRECTLY whose direct API call succeeded, no further
# action is needed.
if self.final_state not in STATES_SENT_DIRECTLY:
self.client.task_instances.finish(
id=self.id,
Expand All@@ -1218,6 +1243,56 @@ def update_task_state_if_needed(self):
rendered_map_index=self._rendered_map_index,
)

def _send_terminal_state_msg(self, msg: SucceedTask | RetryTask | DeferTask | RescheduleTask) -> None:
# Capture the message BEFORE the API call so the recovery dispatcher
# in `update_task_state_if_needed` can re-issue it if the call raises
# (network blip, transient server 5xx). Clear the pending slot and
# record the resulting state only after the call returns successfully.
self._pending_terminal_state_msg = msg
if isinstance(msg, SucceedTask):
self.client.task_instances.succeed(
id=self.id,
when=msg.end_date,
task_outlets=msg.task_outlets,
outlet_events=msg.outlet_events,
rendered_map_index=self._rendered_map_index,
)
self._terminal_state = msg.state
elif isinstance(msg, RetryTask):
self.client.task_instances.retry(
id=self.id,
end_date=msg.end_date,
rendered_map_index=self._rendered_map_index,
)
self._terminal_state = msg.state
elif isinstance(msg, DeferTask):
self.client.task_instances.defer(self.id, msg)
self._terminal_state = TaskInstanceState.DEFERRED
elif isinstance(msg, RescheduleTask):
self.client.task_instances.reschedule(self.id, msg)
self._terminal_state = TaskInstanceState.UP_FOR_RESCHEDULE
self._pending_terminal_state_msg = None

def _replay_pending_terminal_state_msg(self) -> None:
"""
Re-issue the dedicated API call for an unsynced terminal-state msg.

Best-effort — if the second attempt also fails the exception is
logged and we move on; the supervisor's overall failure handling
(heartbeat, exit-code reporting) will eventually surface the issue.
"""
msg = self._pending_terminal_state_msg
if msg is None:
return
try:
self._send_terminal_state_msg(msg)
except Exception:
log.exception(
"Recovery retry of terminal-state API call failed; TI may be stuck on the server",
ti_id=self.id,
msg_type=type(msg).__name__,
)

def _upload_logs(self):
"""
Upload all log files found to the remote storage.
Expand DownExpand Up@@ -1389,29 +1464,20 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id:
resp: BaseModel | None = None
dump_opts = {}
if isinstance(msg, TaskState):
# No direct API call here — the recovery path in
# `update_task_state_if_needed` will call `finish()` for
# non-direct states (FAILED, etc.) once the subprocess exits.
self._terminal_state = msg.state
self._task_end_time_monotonic = time.monotonic()
self._rendered_map_index = msg.rendered_map_index
elif isinstance(msg, SucceedTask):
self._terminal_state = msg.state
self._task_end_time_monotonic = time.monotonic()
self._rendered_map_index = msg.rendered_map_index
self.client.task_instances.succeed(
id=self.id,
when=msg.end_date,
task_outlets=msg.task_outlets,
outlet_events=msg.outlet_events,
rendered_map_index=self._rendered_map_index,
)
self._send_terminal_state_msg(msg)
elif isinstance(msg, RetryTask):
self._terminal_state = msg.state
self._task_end_time_monotonic = time.monotonic()
self._rendered_map_index = msg.rendered_map_index
self.client.task_instances.retry(
id=self.id,
end_date=msg.end_date,
rendered_map_index=self._rendered_map_index,
)
self._send_terminal_state_msg(msg)
elif isinstance(msg, GetConnection):
conn = self.client.connections.get(msg.conn_id)
if isinstance(conn, ConnectionResponse):
Expand DownExpand Up@@ -1463,12 +1529,10 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id:
)
resp = XComSequenceSliceResult.from_response(xcoms)
elif isinstance(msg, DeferTask):
self._terminal_state = TaskInstanceState.DEFERRED
self._rendered_map_index = msg.rendered_map_index
self.client.task_instances.defer(self.id, msg)
self._send_terminal_state_msg(msg)
elif isinstance(msg, RescheduleTask):
self._terminal_state = TaskInstanceState.UP_FOR_RESCHEDULE
self.client.task_instances.reschedule(self.id, msg)
self._send_terminal_state_msg(msg)
elif isinstance(msg, SkipDownstreamTasks):
self.client.task_instances.skip_downstream_tasks(self.id, msg)
elif isinstance(msg, SetXCom):
Expand Down
131 changes: 131 additions & 0 deletions task-sdk/tests/task_sdk/execution_time/test_supervisor.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -2793,6 +2793,137 @@ def test_handle_requests_network_exception_does_not_crash_loop(self, watched_sub
# Should not raise StopIteration (which would mean the loop crashed).
generator.send(req2)

@pytest.mark.parametrize(
("msg", "api_method", "expected_state"),
[
pytest.param(
SucceedTask(end_date=timezone.parse("2024-10-31T12:00:00Z")),
"succeed",
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",
classpath="airflow.providers.standard.triggers.external_task.WorkflowTrigger",
trigger_kwargs={},
),
"defer",
TaskInstanceState.DEFERRED,
id="defer",
),
pytest.param(
RescheduleTask(
reschedule_date=timezone.parse("2024-10-31T12:00:00Z"),
end_date=timezone.parse("2024-10-31T12:00:00Z"),
),
"reschedule",
TaskInstanceState.UP_FOR_RESCHEDULE,
id="reschedule",
),
],
)
def test_terminal_state_not_set_when_direct_api_fails(
self, watched_subprocess, mocker, msg, api_method, expected_state
):
"""`_terminal_state` must NOT be set when the dedicated terminal-state
API raises.

The original message is captured in `_pending_terminal_state_msg`
BEFORE the API call so the recovery dispatcher in
`update_task_state_if_needed` can re-issue it on subprocess exit.
Covers all four terminal-state message types.
"""
watched_subprocess, _ = watched_subprocess
setattr(
watched_subprocess.client.task_instances,
api_method,
mocker.Mock(side_effect=httpx.ConnectError("connection refused")),
)

with pytest.raises(httpx.ConnectError):
watched_subprocess._handle_request(msg, mocker.Mock(), req_id=1)

assert watched_subprocess._terminal_state is None
# Pending msg preserved so the recovery dispatcher can re-issue.
assert watched_subprocess._pending_terminal_state_msg is msg

@pytest.mark.parametrize(
("msg", "api_method", "expected_state"),
[
pytest.param(
SucceedTask(end_date=timezone.parse("2024-10-31T12:00:00Z")),
"succeed",
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",
classpath="airflow.providers.standard.triggers.external_task.WorkflowTrigger",
trigger_kwargs={},
),
"defer",
TaskInstanceState.DEFERRED,
id="defer",
),
pytest.param(
RescheduleTask(
reschedule_date=timezone.parse("2024-10-31T12:00:00Z"),
end_date=timezone.parse("2024-10-31T12:00:00Z"),
),
"reschedule",
TaskInstanceState.UP_FOR_RESCHEDULE,
id="reschedule",
),
],
)
def test_update_task_state_replays_pending_terminal_state_call(
self, watched_subprocess, mocker, msg, api_method, expected_state
):
"""If a direct terminal-state API call was attempted and raised, the
recovery dispatcher must re-issue the dedicated endpoint (not
`finish()`, which the server-side endpoint refuses for SUCCESS /
DEFERRED / SERVER_TERMINATED). Covers all four message types.
"""
watched_subprocess, _ = watched_subprocess
watched_subprocess._exit_code = 0
# Simulate the failure scenario: original API call raised, msg preserved.
watched_subprocess._pending_terminal_state_msg = msg

watched_subprocess.update_task_state_if_needed()

# Recovery re-issues the dedicated endpoint, NOT finish().
getattr(watched_subprocess.client.task_instances, api_method).assert_called_once()
watched_subprocess.client.task_instances.finish.assert_not_called()
assert watched_subprocess._terminal_state == expected_state
assert watched_subprocess._pending_terminal_state_msg is None

def test_update_task_state_no_recovery_without_pending_msg(self, watched_subprocess, mocker):
"""No replay when nothing was pending — preserves the original
STATES_SENT_DIRECTLY short-circuit for the happy path."""
watched_subprocess, _ = watched_subprocess
watched_subprocess._exit_code = 0
watched_subprocess._terminal_state = TaskInstanceState.SUCCESS
watched_subprocess._pending_terminal_state_msg = None

watched_subprocess.update_task_state_if_needed()

watched_subprocess.client.task_instances.finish.assert_not_called()
watched_subprocess.client.task_instances.succeed.assert_not_called()


class TestSetSupervisorComms:
class DummyComms:
Expand Down
Loading