diff --git a/airflow-ctl/docs/images/command_hashes.txt b/airflow-ctl/docs/images/command_hashes.txt index 84cab0f38ecc4..d689874c2a0e9 100644 --- a/airflow-ctl/docs/images/command_hashes.txt +++ b/airflow-ctl/docs/images/command_hashes.txt @@ -1,4 +1,4 @@ -main:8d768c837899829dfd21d37253d2fb44 +main:32f8c659348c161980e59398a17f1bb0 assets:b3ae2b933e54528bf486ff28e887804d auth:f396d4bce90215599dde6ad0a8f30f29 backfills:725109470cd2613de8cc8af022fb54bc diff --git a/airflow-ctl/docs/images/output_main.svg b/airflow-ctl/docs/images/output_main.svg index 98bf85d2fd28f..a5f8703a8549e 100644 --- a/airflow-ctl/docs/images/output_main.svg +++ b/airflow-ctl/docs/images/output_main.svg @@ -19,135 +19,135 @@ font-weight: 700; } - .terminal-101498644-matrix { + .terminal-2610352951-matrix { font-family: Fira Code, monospace; font-size: 20px; line-height: 24.4px; font-variant-east-asian: full-width; } - .terminal-101498644-title { + .terminal-2610352951-title { font-size: 18px; font-weight: bold; font-family: arial; } - .terminal-101498644-r1 { fill: #c5c8c6 } + .terminal-2610352951-r1 { fill: #c5c8c6 } - + - + - + - + - + - + - + - + - + - + - + - + - + - + - + - + - + - + - + - + - + - + - + - + - + - Command: main + Command: main - + - - Usage: airflowctl [-h] GROUP_OR_COMMAND ... - -Positional Arguments: -  GROUP_OR_COMMAND - -    Groups -      assets        Perform Assets operations -      auth          Manage authentication for CLI. Either pass token from -                    environment variable/parameter or pass username and -                    password. -      backfills     Perform Backfills operations -      config        Perform Config operations -      connections   Perform Connections operations -      dag           Perform Dag operations -      dagrun        Perform DagRun operations -      jobs          Perform Jobs operations -      pools         Perform Pools operations -      providers     Perform Providers operations -      variables     Perform Variables operations - -    Commands: -      version       Show version information - -Options: -  -h, --help        show this help message and exit + + Usage: airflowctl [-h] GROUP_OR_COMMAND ... + +Positional Arguments: +  GROUP_OR_COMMAND + +    Groups +      assets        Perform Assets operations +      auth          Manage authentication for CLI. Either pass token from +                    environment variable/parameter or pass username and +                    password. +      backfills     Perform Backfills operations +      config        Perform Config operations +      connections   Perform Connections operations +      dag           Perform Dag operations +      dagrun        Perform DagRun operations +      jobs          Perform Jobs operations +      pools         Perform Pools operations +      providers     Perform Providers operations +      variables     Perform Variables operations + +    Commands: +      version       Show version information + +Options: +  -h, --help        show this help message and exit diff --git a/airflow-ctl/src/airflowctl/api/operations.py b/airflow-ctl/src/airflowctl/api/operations.py index 0161f566ff790..b247fd787aa65 100644 --- a/airflow-ctl/src/airflowctl/api/operations.py +++ b/airflow-ctl/src/airflowctl/api/operations.py @@ -18,10 +18,11 @@ from __future__ import annotations import datetime -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, TypeVar import httpx import structlog +from pydantic import BaseModel from airflowctl.api.datamodels.auth_generated import LoginBody, LoginResponse from airflowctl.api.datamodels.generated import ( @@ -77,6 +78,8 @@ log = structlog.get_logger(logger_name=__name__) +T = TypeVar("T", bound=BaseModel) + # Generic Server Response Error class ServerResponseError(httpx.HTTPStatusError): @@ -147,6 +150,37 @@ def __init_subclass__(cls, **kwargs): if callable(value): setattr(cls, attr, _check_flag_and_exit_if_server_response_error(value)) + def execute_list( + self, + *, + path: str, + data_model: type[T], + offset: int = 0, + limit: int = 50, + params: dict | None = None, + ) -> T | ServerResponseError: + shared_params = {**(params or {})} + try: + self.response = self.client.get(path, params=shared_params) + first_pass = data_model.model_validate_json(self.response.content) + total_entries = first_pass.total_entries # type: ignore[attr-defined] + if total_entries < limit: + return first_pass + for key, value in first_pass.model_dump().items(): + if key != "total_entries" and isinstance(value, list): + break + entry_list = getattr(first_pass, key) + offset = offset + limit + while offset < total_entries: + self.response = self.client.get(path, params={**shared_params, "offset": offset}) + entry = data_model.model_validate_json(self.response.content) + offset = offset + limit + entry_list.extend(getattr(entry, key)) + obj = data_model(**{key: entry_list, "total_entries": total_entries}) + return data_model.model_validate(obj.model_dump()) + except ServerResponseError as e: + raise e + # Login operations class LoginOperations: @@ -188,19 +222,11 @@ def get_by_alias(self, alias: str) -> AssetAliasResponse | ServerResponseError: def list(self) -> AssetCollectionResponse | ServerResponseError: """List all assets from the API server.""" - try: - self.response = self.client.get("assets") - return AssetCollectionResponse.model_validate_json(self.response.content) - except ServerResponseError as e: - raise e + return super().execute_list(path="assets", data_model=AssetCollectionResponse) def list_by_alias(self) -> AssetAliasCollectionResponse | ServerResponseError: """List all assets by alias from the API server.""" - try: - self.response = self.client.get("/assets/aliases") - return AssetAliasCollectionResponse.model_validate_json(self.response.content) - except ServerResponseError as e: - raise e + return super().execute_list(path="/assets/aliases", data_model=AssetAliasCollectionResponse) def create_event( self, asset_event_body: CreateAssetEventsBody @@ -298,13 +324,10 @@ def get(self, backfill_id: str) -> BackfillResponse | ServerResponseError: except ServerResponseError as e: raise e - def list(self) -> BackfillCollectionResponse | ServerResponseError: + def list(self, dag_id: str) -> BackfillCollectionResponse | ServerResponseError: """List all backfills.""" - try: - self.response = self.client.get("backfills") - return BackfillCollectionResponse.model_validate_json(self.response.content) - except ServerResponseError as e: - raise e + params = {"dag_id": dag_id} + return super().execute_list(path="backfills", data_model=BackfillCollectionResponse, params=params) def pause(self, backfill_id: str) -> BackfillResponse | ServerResponseError: """Pause a backfill.""" @@ -364,11 +387,7 @@ def get(self, conn_id: str) -> ConnectionResponse | ServerResponseError: def list(self) -> ConnectionCollectionResponse | ServerResponseError: """List all connections from the API server.""" - try: - self.response = self.client.get("connections") - return ConnectionCollectionResponse.model_validate_json(self.response.content) - except ServerResponseError as e: - raise e + return super().execute_list(path="connections", data_model=ConnectionCollectionResponse) def create( self, @@ -451,19 +470,11 @@ def get_details(self, dag_id: str) -> DAGDetailsResponse | ServerResponseError: def get_tags(self) -> DAGTagCollectionResponse | ServerResponseError: """Get all DAG tags.""" - try: - self.response = self.client.get("dagTags") - return DAGTagCollectionResponse.model_validate_json(self.response.content) - except ServerResponseError as e: - raise e + return super().execute_list(path="dagTags", data_model=DAGTagCollectionResponse) def list(self) -> DAGCollectionResponse | ServerResponseError: """List DAGs.""" - try: - self.response = self.client.get("dags") - return DAGCollectionResponse.model_validate_json(self.response.content) - except ServerResponseError as e: - raise e + return super().execute_list(path="dags", data_model=DAGCollectionResponse) def patch(self, dag_id: str, dag_body: DAGPatchBody) -> DAGResponse | ServerResponseError: try: @@ -487,11 +498,7 @@ def get_import_error(self, import_error_id: str) -> ImportErrorResponse | Server raise e def list_import_error(self) -> ImportErrorCollectionResponse | ServerResponseError: - try: - self.response = self.client.get("importErrors") - return ImportErrorCollectionResponse.model_validate_json(self.response.content) - except ServerResponseError as e: - raise e + return super().execute_list(path="importErrors", data_model=ImportErrorCollectionResponse) def get_stats(self, dag_ids: list) -> DagStatsCollectionResponse | ServerResponseError: # type: ignore try: @@ -508,18 +515,12 @@ def get_version(self, dag_id: str, version_number: int) -> DagVersionResponse | raise e def list_version(self, dag_id: str) -> DAGVersionCollectionResponse | ServerResponseError: - try: - self.response = self.client.get(f"dags/{dag_id}/dagVersions") - return DAGVersionCollectionResponse.model_validate_json(self.response.content) - except ServerResponseError as e: - raise e + return super().execute_list( + path=f"dags/{dag_id}/dagVersions", data_model=DAGVersionCollectionResponse + ) def list_warning(self) -> DAGWarningCollectionResponse | ServerResponseError: - try: - self.response = self.client.get("dagWarnings") - return DAGWarningCollectionResponse.model_validate_json(self.response.content) - except ServerResponseError as e: - raise e + return super().execute_list(path="dagWarnings", data_model=DAGWarningCollectionResponse) class DagRunOperations(BaseOperations): @@ -542,17 +543,13 @@ def list( limit: int, ) -> DAGRunCollectionResponse | ServerResponseError: """List all dag runs.""" - try: - params = { - "start_date": start_date, - "end_date": end_date, - "state": state, - "limit": limit, - } - self.response = self.client.get("dag_runs", params=params) # type: ignore - return DAGRunCollectionResponse.model_validate_json(self.response.content) - except ServerResponseError as e: - raise e + params = { + "start_date": start_date, + "end_date": end_date, + "state": state, + "limit": limit, + } + return super().execute_list(path="dag_runs", data_model=DAGRunCollectionResponse, params=params) def create( self, dag_id: str, trigger_dag_run: TriggerDAGRunPostBody @@ -573,12 +570,8 @@ def list( self, job_type: str, hostname: str, is_alive: bool ) -> JobCollectionResponse | ServerResponseError: """List all jobs.""" - try: - params = {"job_type": job_type, "hostname": hostname, "is_alive": is_alive} - self.response = self.client.get("jobs", params=params) # type: ignore - return JobCollectionResponse.model_validate_json(self.response.content) - except ServerResponseError as e: - raise e + params = {"job_type": job_type, "hostname": hostname, "is_alive": is_alive} + return super().execute_list(path="jobs", data_model=JobCollectionResponse, params=params) class PoolsOperations(BaseOperations): @@ -594,11 +587,7 @@ def get(self, pool_name: str) -> PoolResponse | ServerResponseError: def list(self) -> PoolCollectionResponse | ServerResponseError: """List all pools.""" - try: - self.response = self.client.get("pools") - return PoolCollectionResponse.model_validate_json(self.response.content) - except ServerResponseError as e: - raise e + return super().execute_list(path="pools", data_model=PoolCollectionResponse) def create(self, pool: PoolBody) -> PoolResponse | ServerResponseError: """Create a pool.""" @@ -638,11 +627,7 @@ class ProvidersOperations(BaseOperations): def list(self) -> ProviderCollectionResponse | ServerResponseError: """List all providers.""" - try: - self.response = self.client.get("providers") - return ProviderCollectionResponse.model_validate_json(self.response.content) - except ServerResponseError as e: - raise e + return super().execute_list(path="providers", data_model=ProviderCollectionResponse) class VariablesOperations(BaseOperations): @@ -658,11 +643,7 @@ def get(self, variable_key: str) -> VariableResponse | ServerResponseError: def list(self) -> VariableCollectionResponse | ServerResponseError: """List all variables.""" - try: - self.response = self.client.get("variables") - return VariableCollectionResponse.model_validate_json(self.response.content) - except ServerResponseError as e: - raise e + return super().execute_list(path="variables", data_model=VariableCollectionResponse) def create(self, variable: VariableBody) -> VariableResponse | ServerResponseError: """Create a variable.""" diff --git a/airflow-ctl/tests/airflow_ctl/api/test_operations.py b/airflow-ctl/tests/airflow_ctl/api/test_operations.py index 1a894990183a0..c8f88f046fba2 100644 --- a/airflow-ctl/tests/airflow_ctl/api/test_operations.py +++ b/airflow-ctl/tests/airflow_ctl/api/test_operations.py @@ -24,6 +24,7 @@ import httpx import pytest +from pydantic import BaseModel from airflowctl.api.client import Client, ClientKind from airflowctl.api.datamodels.auth_generated import LoginBody, LoginResponse @@ -106,6 +107,15 @@ def make_api_client( return Client(base_url=base_url, transport=transport, token=token, kind=kind) +class HelloResponse(BaseModel): + name: str + + +class HelloCollectionResponse(BaseModel): + hellos: list[HelloResponse] + total_entries: int + + class TestBaseOperations: def test_server_connection_refused(self): client = make_api_client(base_url="http://localhost") @@ -114,6 +124,46 @@ def test_server_connection_refused(self): ): client.connections.get("1") + @pytest.mark.parametrize( + "total_entries, limit, expected_response", + [ + (0, 0, (HelloCollectionResponse(hellos=[], total_entries=0))), + (1, 50, (HelloCollectionResponse(hellos=[HelloResponse(name="hello")], total_entries=1))), + ( + 150, + 50, + ( + HelloCollectionResponse( + hellos=[ + HelloResponse(name="hello"), + HelloResponse(name="hello"), + HelloResponse(name="hello"), + ], + total_entries=150, + ) + ), + ), + ( + 90, + 50, + ( + HelloCollectionResponse( + hellos=[HelloResponse(name="hello"), HelloResponse(name="hello")], total_entries=90 + ) + ), + ), + ], + ) + def test_execute_list(self, total_entries, limit, expected_response): + hello_response = [] + if total_entries != 0: + update = (total_entries + limit - 1) // limit + hello_response.extend([HelloResponse(name="hello")] * update) + hello_collection_response = HelloCollectionResponse( + hellos=hello_response, total_entries=total_entries + ) + assert expected_response == hello_collection_response + class TestAssetsOperations: asset_id: int = 1 @@ -367,7 +417,7 @@ def handle_request(request: httpx.Request) -> httpx.Response: return httpx.Response(200, json=json.loads(self.backfills_collection_response.model_dump_json())) client = make_api_client(transport=httpx.MockTransport(handle_request)) - response = client.backfills.list() + response = client.backfills.list(dag_id="dag_id") assert response == self.backfills_collection_response def test_pause(self):