Skip to content
Open
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
20 changes: 9 additions & 11 deletions task-sdk/src/airflow/sdk/execution_time/task_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down
42 changes: 42 additions & 0 deletions task-sdk/tests/task_sdk/execution_time/test_task_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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")
Expand Down