diff --git a/task-sdk/src/airflow/sdk/api/client.py b/task-sdk/src/airflow/sdk/api/client.py index 893364d47a9c6..8f2e8dba05aa8 100644 --- a/task-sdk/src/airflow/sdk/api/client.py +++ b/task-sdk/src/airflow/sdk/api/client.py @@ -47,6 +47,7 @@ TIEnterRunningPayload, TIHeartbeatInfo, TIRescheduleStatePayload, + TIRetryStatePayload, TIRunContext, TISkippedDownstreamTasksStatePayload, TISuccessStatePayload, @@ -152,6 +153,11 @@ def finish(self, id: uuid.UUID, state: TerminalStateNonSuccess, when: datetime): body = TITerminalStatePayload(end_date=when, state=TerminalStateNonSuccess(state)) self.client.patch(f"task-instances/{id}/state", content=body.model_dump_json()) + def retry(self, id: uuid.UUID, end_date: datetime): + """Tell the API server that this TI has failed and reached a up_for_retry state.""" + body = TIRetryStatePayload(end_date=end_date) + self.client.patch(f"task-instances/{id}/state", content=body.model_dump_json()) + def succeed(self, id: uuid.UUID, when: datetime, task_outlets, outlet_events): """Tell the API server that this TI has succeeded.""" body = TISuccessStatePayload(end_date=when, task_outlets=task_outlets, outlet_events=outlet_events) diff --git a/task-sdk/src/airflow/sdk/execution_time/supervisor.py b/task-sdk/src/airflow/sdk/execution_time/supervisor.py index fffd7603a1f56..768c25eafa9d0 100644 --- a/task-sdk/src/airflow/sdk/execution_time/supervisor.py +++ b/task-sdk/src/airflow/sdk/execution_time/supervisor.py @@ -87,6 +87,7 @@ PrevSuccessfulDagRunResult, PutVariable, RescheduleTask, + RetryTask, SetRenderedFields, SetXCom, SkipDownstreamTasks, @@ -125,6 +126,7 @@ STATES_SENT_DIRECTLY = [ IntermediateTIState.DEFERRED, IntermediateTIState.UP_FOR_RESCHEDULE, + IntermediateTIState.UP_FOR_RETRY, TerminalTIState.SUCCESS, ] @@ -913,6 +915,13 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger): task_outlets=msg.task_outlets, outlet_events=msg.outlet_events, ) + elif isinstance(msg, RetryTask): + self._terminal_state = msg.state + self._task_end_time_monotonic = time.monotonic() + self.client.task_instances.retry( + id=self.id, + end_date=msg.end_date, + ) elif isinstance(msg, GetConnection): conn = self.client.connections.get(msg.conn_id) if isinstance(conn, ConnectionResponse): diff --git a/task-sdk/tests/task_sdk/api/test_client.py b/task-sdk/tests/task_sdk/api/test_client.py index 541faccc58bf9..8e4ad0cab71b5 100644 --- a/task-sdk/tests/task_sdk/api/test_client.py +++ b/task-sdk/tests/task_sdk/api/test_client.py @@ -375,6 +375,22 @@ def handle_request(request: httpx.Request) -> httpx.Response: ) client.task_instances.reschedule(ti_id, msg) + def test_task_instance_up_for_retry(self): + ti_id = uuid6.uuid7() + + def handle_request(request: httpx.Request) -> httpx.Response: + if request.url.path == f"/task-instances/{ti_id}/state": + actual_body = json.loads(request.read()) + assert actual_body["state"] == "up_for_retry" + assert actual_body["end_date"] == "2024-10-31T12:00:00Z" + return httpx.Response( + status_code=204, + ) + return httpx.Response(status_code=400, json={"detail": "Bad Request"}) + + client = make_client(transport=httpx.MockTransport(handle_request)) + client.task_instances.retry(ti_id, end_date=timezone.parse("2024-10-31T12:00:00Z")) + @pytest.mark.parametrize( "rendered_fields", [ diff --git a/task-sdk/tests/task_sdk/execution_time/test_supervisor.py b/task-sdk/tests/task_sdk/execution_time/test_supervisor.py index 5cae5232e64e3..568d5a38857fa 100644 --- a/task-sdk/tests/task_sdk/execution_time/test_supervisor.py +++ b/task-sdk/tests/task_sdk/execution_time/test_supervisor.py @@ -78,6 +78,7 @@ PrevSuccessfulDagRunResult, PutVariable, RescheduleTask, + RetryTask, SetRenderedFields, SetXCom, SucceedTask, @@ -1165,6 +1166,15 @@ def watched_subprocess(self, mocker): "", id="patch_task_instance_to_skipped", ), + pytest.param( + RetryTask(end_date=timezone.parse("2024-10-31T12:00:00Z")), + b"", + "task_instances.retry", + (), + {"id": TI_ID, "end_date": timezone.parse("2024-10-31T12:00:00Z")}, + "", + id="up_for_retry", + ), pytest.param( SetRenderedFields(rendered_fields={"field1": "rendered_value1", "field2": "rendered_value2"}), b"",