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):