diff --git a/task-sdk/src/airflow/sdk/execution_time/task_runner.py b/task-sdk/src/airflow/sdk/execution_time/task_runner.py index 22ac90405e027..54016e5dc9390 100644 --- a/task-sdk/src/airflow/sdk/execution_time/task_runner.py +++ b/task-sdk/src/airflow/sdk/execution_time/task_runner.py @@ -1474,9 +1474,10 @@ def _on_term(signum, frame): log.info("::group::Post Execute") if e.args: log.info("Skipping task.", reason=e.args[0]) + ti.end_date = datetime.now(tz=timezone.utc) msg = TaskState( state=TaskInstanceState.SKIPPED, - end_date=datetime.now(tz=timezone.utc), + end_date=ti.end_date, rendered_map_index=ti.rendered_map_index, ) state = TaskInstanceState.SKIPPED @@ -1718,26 +1719,23 @@ def _handle_trigger_dag_run( ) if isinstance(comms_msg, ErrorResponse) and comms_msg.error == ErrorType.DAGRUN_ALREADY_EXISTS: + ti.end_date = datetime.now(tz=timezone.utc) + state: Literal[TaskInstanceState.FAILED, TaskInstanceState.SKIPPED] if drte.skip_when_already_exists: log.info( "Dag Run already exists, skipping task as skip_when_already_exists is set to True.", dag_id=drte.trigger_dag_id, ) - msg = TaskState( - state=TaskInstanceState.SKIPPED, - end_date=datetime.now(tz=timezone.utc), - rendered_map_index=ti.rendered_map_index, - ) state = TaskInstanceState.SKIPPED else: log.error("Dag Run already exists, marking task as failed.", dag_id=drte.trigger_dag_id) - msg = TaskState( - state=TaskInstanceState.FAILED, - end_date=datetime.now(tz=timezone.utc), - rendered_map_index=ti.rendered_map_index, - ) state = TaskInstanceState.FAILED + msg = TaskState( + state=state, + end_date=ti.end_date, + rendered_map_index=ti.rendered_map_index, + ) return msg, state log.info("Dag Run triggered successfully.", trigger_dag_id=drte.trigger_dag_id) diff --git a/task-sdk/tests/task_sdk/execution_time/test_task_runner.py b/task-sdk/tests/task_sdk/execution_time/test_task_runner.py index f8953dc232970..6ba897179baa6 100644 --- a/task-sdk/tests/task_sdk/execution_time/test_task_runner.py +++ b/task-sdk/tests/task_sdk/execution_time/test_task_runner.py @@ -4891,6 +4891,10 @@ def test_handle_trigger_dag_run_conflict( assert state == expected_state assert msg.state == expected_state + # end_date must be set on the local instance (not just the outbound + # message) so finalize() emits the task.duration metric + assert ti.end_date is not None + assert msg.end_date == ti.end_date expected_calls = [ mock.call.send( @@ -5137,6 +5141,44 @@ def test_ti_finish_metric_emitted_for_terminal_states( tags={"dag_id": ti.dag_id, "task_id": ti.task_id, "state": expected_state}, ) + @pytest.mark.parametrize( + ("task_callable", "expected_state"), + [ + pytest.param(lambda: "success", "success", id="success"), + pytest.param(lambda: (_ for _ in ()).throw(AirflowSkipException()), "skipped", id="skipped"), + pytest.param(lambda: (_ for _ in ()).throw(AirflowFailException("fail")), "failed", id="failed"), + ], + ) + def test_task_duration_metric_emitted_for_terminal_states( + self, task_callable, expected_state, create_runtime_ti, mock_supervisor_comms + ): + """task.duration is emitted for every terminal state — success, skipped, failed.""" + task = PythonOperator(task_id="test", python_callable=task_callable) + ti = create_runtime_ti(task=task) + + context = ti.get_template_context() + with mock.patch("airflow.sdk._shared.observability.metrics.stats._get_backend") as mock_get_backend: + backend = mock.MagicMock(spec=StatsLogger) + mock_get_backend.return_value = backend + state, _, error = run(ti, context=context, log=mock.MagicMock()) + finalize(ti, state=state, context=context, log=mock.MagicMock(), error=error) + + # verify task.duration was emitted in tagged format + task_duration_calls = [ + c for c in backend.timing.call_args_list if c.args and c.args[0] == "task.duration" + ] + assert len(task_duration_calls) == 1, ( + f"Expected exactly 1 task.duration emit for state={expected_state}, " + f"got {len(task_duration_calls)}" + ) + assert task_duration_calls[0].kwargs.get("tags") == { + "dag_id": ti.dag_id, + "task_id": ti.task_id, + } + # verify task.duration was also emitted in legacy dotted format via the + # registry's legacy_name derivation in metrics_template.yaml + backend.timing.assert_any_call(f"dag.{ti.dag_id}.{ti.task_id}.duration", mock.ANY) + def test_operator_successes_metrics_emitted(self, create_runtime_ti, mock_supervisor_comms): """Test that operator_successes and ti_successes metrics are emitted on task success.""" task = PythonOperator(task_id="test", python_callable=lambda: "success")