diff --git a/airflow-ctl-tests/tests/airflowctl_tests/test_airflowctl_commands.py b/airflow-ctl-tests/tests/airflowctl_tests/test_airflowctl_commands.py index 14a500c4d5bf7..493c3e18b4506 100644 --- a/airflow-ctl-tests/tests/airflowctl_tests/test_airflowctl_commands.py +++ b/airflow-ctl-tests/tests/airflowctl_tests/test_airflowctl_commands.py @@ -106,6 +106,8 @@ def date_param(): 'taskinstances get example_bash_operator "manual__{date_param}" runme_0', 'taskinstances get-dependencies example_bash_operator "manual__{date_param}" runme_0', 'taskinstances list example_bash_operator "manual__{date_param}"', + # Task instance get (auto-generated command, uses positional args) - needs a Dag run with completed tasks + 'taskinstances get example_bash_operator "manual__{date_param}" runme_0', # XCom commands - need a Dag run with completed tasks 'xcom add example_bash_operator "manual__{date_param}" runme_0 {xcom_key} \'{{"test": "value"}}\'', 'xcom get example_bash_operator "manual__{date_param}" runme_0 {xcom_key}', diff --git a/airflow-ctl/src/airflowctl/ctl/commands/task_command.py b/airflow-ctl/src/airflowctl/ctl/commands/task_command.py index abd61b23dde99..76e2b3db81f42 100644 --- a/airflow-ctl/src/airflowctl/ctl/commands/task_command.py +++ b/airflow-ctl/src/airflowctl/ctl/commands/task_command.py @@ -149,3 +149,17 @@ def states_for_dag_run(args, api_client=NEW_API_CLIENT) -> None: data=[_format_task_instance(ti, has_mapped_instances) for ti in task_instances], output=args.output, ) + + +@provide_api_client(kind=ClientKind.CLI) +def task_state(args, api_client=NEW_API_CLIENT) -> None: + """Get the state of a task instance.""" + ti = api_client.task_instances.get( + dag_id=args.dag_id, + dag_run_id=args.dag_run_id, + task_id=args.task_id, + map_index=args.map_index, + ) + # ``state`` is a str-mixin enum; ``str()`` on it yields "TaskInstanceState.SUCCESS". + state = getattr(ti.state, "value", ti.state) + AirflowConsole().print_as(data=[{"state": state}], output=args.output) diff --git a/airflow-ctl/src/airflowctl/ctl/help_texts.yaml b/airflow-ctl/src/airflowctl/ctl/help_texts.yaml index 4944c60dbdf8b..bb428af0fda16 100644 --- a/airflow-ctl/src/airflowctl/ctl/help_texts.yaml +++ b/airflow-ctl/src/airflowctl/ctl/help_texts.yaml @@ -87,6 +87,7 @@ providers: list: "List all installed Airflow providers" taskinstances: + get: "Get a task instance by Dag ID, Dag run ID and task ID, optionally with a map index" list: "List all task instances for a given Dag run" get: "Get a task instance for a given Dag run" get-dependencies: "Get unmet scheduler dependencies for a task instance" diff --git a/airflow-ctl/tests/airflow_ctl/api/test_operations.py b/airflow-ctl/tests/airflow_ctl/api/test_operations.py index 7806fcfb5abae..2ac2f3d4643af 100644 --- a/airflow-ctl/tests/airflow_ctl/api/test_operations.py +++ b/airflow-ctl/tests/airflow_ctl/api/test_operations.py @@ -2369,6 +2369,103 @@ def handle_request(request: httpx.Request) -> httpx.Response: assert response == self.key +class TestTaskInstancesOperations: + """Test suite for task instance operations.""" + + dag_id: str = "test_dag" + dag_run_id: str = "manual__2025-01-24T00:00:00+00:00" + task_id: str = "test_task" + + task_instance_response = TaskInstanceResponse( + id=uuid.uuid4(), + task_id=task_id, + dag_id=dag_id, + dag_run_id=dag_run_id, + map_index=-1, + run_after=datetime.datetime(2025, 1, 24, 0, 0, 0), + try_number=1, + max_tries=1, + task_display_name=task_id, + dag_display_name=dag_id, + pool="default_pool", + pool_slots=1, + executor_config="{}", + state=TaskInstanceState.SUCCESS, + ) + task_instance_collection_response = TaskInstanceCollectionResponse( + task_instances=[task_instance_response], + total_entries=1, + ) + + def test_get(self): + """Test fetching an unmapped task instance hits the standard endpoint.""" + + def handle_request(request: httpx.Request) -> httpx.Response: + assert request.url.path == ( + f"/api/v2/dags/{self.dag_id}/dagRuns/{self.dag_run_id}/taskInstances/{self.task_id}" + ) + return httpx.Response(200, json=json.loads(self.task_instance_response.model_dump_json())) + + client = make_api_client(transport=httpx.MockTransport(handle_request)) + response = client.task_instances.get( + dag_id=self.dag_id, + dag_run_id=self.dag_run_id, + task_id=self.task_id, + ) + assert response == self.task_instance_response + + @pytest.mark.parametrize("map_index", [-1, None]) + def test_get_without_map_index_uses_unmapped_endpoint(self, map_index): + """A negative or omitted ``map_index`` must not append a map index to the path.""" + + def handle_request(request: httpx.Request) -> httpx.Response: + assert request.url.path == ( + f"/api/v2/dags/{self.dag_id}/dagRuns/{self.dag_run_id}/taskInstances/{self.task_id}" + ) + return httpx.Response(200, json=json.loads(self.task_instance_response.model_dump_json())) + + client = make_api_client(transport=httpx.MockTransport(handle_request)) + response = client.task_instances.get( + dag_id=self.dag_id, + dag_run_id=self.dag_run_id, + task_id=self.task_id, + map_index=map_index, + ) + assert response == self.task_instance_response + + @pytest.mark.parametrize("map_index", [0, 1, 7]) + def test_get_with_map_index_uses_mapped_endpoint(self, map_index): + """A non-negative ``map_index`` must hit the mapped task instance endpoint.""" + mapped_response = self.task_instance_response.model_copy(update={"map_index": map_index}) + + def handle_request(request: httpx.Request) -> httpx.Response: + assert request.url.path == ( + f"/api/v2/dags/{self.dag_id}/dagRuns/{self.dag_run_id}/" + f"taskInstances/{self.task_id}/{map_index}" + ) + return httpx.Response(200, json=json.loads(mapped_response.model_dump_json())) + + client = make_api_client(transport=httpx.MockTransport(handle_request)) + response = client.task_instances.get( + dag_id=self.dag_id, + dag_run_id=self.dag_run_id, + task_id=self.task_id, + map_index=map_index, + ) + assert response == mapped_response + + def test_list(self): + def handle_request(request: httpx.Request) -> httpx.Response: + assert request.url.path == f"/api/v2/dags/{self.dag_id}/dagRuns/{self.dag_run_id}/taskInstances" + return httpx.Response( + 200, json=json.loads(self.task_instance_collection_response.model_dump_json()) + ) + + client = make_api_client(transport=httpx.MockTransport(handle_request)) + response = client.task_instances.list(dag_id=self.dag_id, dag_run_id=self.dag_run_id) + assert response == self.task_instance_collection_response + + class TestPluginsOperations: plugin_response = PluginResponse( name="test-plugin", diff --git a/airflow-ctl/tests/airflow_ctl/ctl/commands/test_task_command.py b/airflow-ctl/tests/airflow_ctl/ctl/commands/test_task_command.py index 2f6f23f6120ad..61cfa0b84aceb 100644 --- a/airflow-ctl/tests/airflow_ctl/ctl/commands/test_task_command.py +++ b/airflow-ctl/tests/airflow_ctl/ctl/commands/test_task_command.py @@ -17,12 +17,14 @@ from __future__ import annotations import datetime +import json import uuid from unittest import mock import httpx import pytest +from airflowctl.api.client import ClientKind from airflowctl.api.datamodels.generated import ( TaskDependencyCollectionResponse, TaskDependencyResponse, @@ -631,3 +633,96 @@ def test_states_for_dag_run_propagates_non_404_api_error(self, failing_call): task_command.states_for_dag_run(self.parser.parse_args(argv), api_client=api_client) assert ctx.value is error + + +class TestTaskCommands: + parser = cli_parser.get_parser() + dag_id = "example_dag" + dag_run_id = "manual__2024-01-01T00:00:00+00:00" + task_id = "my_task" + + task_instance_response = TaskInstanceResponse( + id=uuid.uuid4(), + task_id=task_id, + dag_id=dag_id, + dag_run_id=dag_run_id, + map_index=-1, + run_after=datetime.datetime(2024, 1, 1, 0, 0, 0), + try_number=1, + max_tries=1, + task_display_name=task_id, + dag_display_name=dag_id, + pool="default_pool", + pool_slots=1, + executor_config="{}", + state=TaskInstanceState.SUCCESS, + ) + + def test_task_state(self, api_client_maker, capsys): + api_client = api_client_maker( + path=f"/api/v2/dags/{self.dag_id}/dagRuns/{self.dag_run_id}/taskInstances/{self.task_id}", + response_json=self.task_instance_response.model_dump(mode="json"), + expected_http_status_code=200, + kind=ClientKind.CLI, + ) + task_command.task_state( + self.parser.parse_args( + [ + "tasks", + "state", + self.dag_id, + self.dag_run_id, + self.task_id, + ] + ), + api_client=api_client, + ) + assert json.loads(capsys.readouterr().out) == [{"state": "success"}] + + def test_task_state_not_found(self, api_client_maker): + api_client = api_client_maker( + path=f"/api/v2/dags/{self.dag_id}/dagRuns/{self.dag_run_id}/taskInstances/{self.task_id}", + response_json={"detail": "Task instance not found"}, + expected_http_status_code=404, + kind=ClientKind.CLI, + ) + with pytest.raises(ServerResponseError): + task_command.task_state( + self.parser.parse_args( + [ + "tasks", + "state", + self.dag_id, + self.dag_run_id, + self.task_id, + ] + ), + api_client=api_client, + ) + + @pytest.mark.parametrize("map_index", [0, 1, 7]) + def test_task_state_mapped(self, api_client_maker, capsys, map_index): + mapped_response = self.task_instance_response.model_copy(update={"map_index": map_index}) + api_client = api_client_maker( + path=( + f"/api/v2/dags/{self.dag_id}/dagRuns/{self.dag_run_id}" + f"/taskInstances/{self.task_id}/{map_index}" + ), + response_json=mapped_response.model_dump(mode="json"), + expected_http_status_code=200, + kind=ClientKind.CLI, + ) + task_command.task_state( + self.parser.parse_args( + [ + "tasks", + "state", + self.dag_id, + self.dag_run_id, + self.task_id, + f"--map-index={map_index}", + ] + ), + api_client=api_client, + ) + assert json.loads(capsys.readouterr().out) == [{"state": "success"}]