From 5a3211c40ee9b03cf4af7ac77b11b944277d7948 Mon Sep 17 00:00:00 2001 From: Wei Lee Date: Fri, 31 Jul 2026 23:22:53 +0800 Subject: [PATCH] Fix airflowctl commands failing against older Airflow servers MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit airflowctl's datamodels are generated from the newest Airflow API spec, so a server on an older Airflow line legitimately omits fields the models declare as required. Collection responses already tolerated that via fill_missing_fields, added in #63388 for exactly this reason, but only inside execute_list. Every single-object response validated strictly, so commands such as dags get-details, dags unpause and dags update failed: ValidationError: 2 validation errors for DAGResponse is_backfillable Field required timetable_periodic Field required against any Airflow older than the line the models were generated from. Lift the validation out of execute_list into a shared validate_response() and use it for all 46 single-object responses. LoginOperations keeps strict validation on purpose — a defaulted token should fail loudly rather than be filled in. --- airflow-ctl/src/airflowctl/api/operations.py | 118 ++++++++++-------- .../tests/airflow_ctl/api/test_operations.py | 35 ++++++ 2 files changed, 98 insertions(+), 55 deletions(-) diff --git a/airflow-ctl/src/airflowctl/api/operations.py b/airflow-ctl/src/airflowctl/api/operations.py index 94b31010abbc2..79f8095cad086 100644 --- a/airflow-ctl/src/airflowctl/api/operations.py +++ b/airflow-ctl/src/airflowctl/api/operations.py @@ -185,6 +185,21 @@ def fill_missing_fields(data: dict, model: type[BaseModel]) -> dict: return data +def validate_response(content: bytes, data_model: type[T]) -> T: + """ + Validate a server response, tolerating fields an older server does not send. + + The datamodels are generated from the newest Airflow API spec, so a server on an + older Airflow line can legitimately omit fields the models declare as required. + Fill those with type defaults rather than failing the whole command. + """ + try: + return data_model.model_validate_json(content) + except ValidationError: + raw = fill_missing_fields(json.loads(content), data_model) + return data_model.model_validate(raw) + + class BaseOperations: """ Base class for operations. @@ -213,15 +228,8 @@ def execute_list(self, *, path, data_model, offset=0, limit=50, params=None): shared_params = {"limit": limit, **(params or {})} - def safe_validate(content: bytes) -> BaseModel: - try: - return data_model.model_validate_json(content) # type: ignore[union-attr] - except ValidationError: - raw = fill_missing_fields(json.loads(content), data_model) - return data_model.model_validate(raw) # type: ignore[union-attr] - self.response = self.client.get(path, params=shared_params) - first_pass = safe_validate(self.response.content) + first_pass = validate_response(self.response.content, data_model) total_entries = first_pass.total_entries # type: ignore[attr-defined] if total_entries < limit: return first_pass @@ -234,7 +242,7 @@ def safe_validate(content: bytes) -> BaseModel: offset = offset + limit while offset < total_entries: self.response = self.client.get(path, params={**shared_params, "offset": offset}) - entry = safe_validate(self.response.content) + entry = validate_response(self.response.content, data_model) offset = offset + limit entry_list.extend(getattr(entry, found_key)) obj = data_model(**{found_key: entry_list, "total_entries": total_entries}) @@ -262,12 +270,12 @@ class AssetsOperations(BaseOperations): def get(self, asset_id: str) -> AssetResponse | ServerResponseError: """Get an asset from the API server.""" self.response = self.client.get(f"assets/{asset_id}") - return AssetResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, AssetResponse) def get_by_alias(self, alias: str) -> AssetAliasResponse | ServerResponseError: """Get an asset by alias from the API server.""" self.response = self.client.get(f"assets/aliases/{alias}") - return AssetAliasResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, AssetAliasResponse) def list(self) -> AssetCollectionResponse | ServerResponseError: """List all assets from the API server.""" @@ -287,29 +295,29 @@ def create_event( self.response = self.client.post( "assets/events", json=asset_event_body.model_dump(mode="json", exclude_none=True) ) - return AssetEventResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, AssetEventResponse) def materialize(self, asset_id: str) -> DAGRunResponse | ServerResponseError: """Materialize an asset.""" self.response = self.client.post(f"assets/{asset_id}/materialize") - return DAGRunResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, DAGRunResponse) def get_queued_events(self, asset_id: str) -> QueuedEventCollectionResponse | ServerResponseError: """Get queued events for an asset.""" self.response = self.client.get(f"assets/{asset_id}/queuedEvents") - return QueuedEventCollectionResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, QueuedEventCollectionResponse) def get_dag_queued_events( self, dag_id: str, before: str ) -> QueuedEventCollectionResponse | ServerResponseError: """Get queued events for a dag.""" self.response = self.client.get(f"dags/{dag_id}/assets/queuedEvents", params={"before": before}) - return QueuedEventCollectionResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, QueuedEventCollectionResponse) def get_dag_queued_event(self, dag_id: str, asset_id: str) -> QueuedEventResponse | ServerResponseError: """Get a queued event for a dag.""" self.response = self.client.get(f"dags/{dag_id}/assets/{asset_id}/queuedEvents") - return QueuedEventResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, QueuedEventResponse) def delete_queued_events(self, asset_id: str) -> str | ServerResponseError: """Delete a queued event for an asset.""" @@ -335,19 +343,19 @@ def create(self, backfill: BackfillPostBody) -> BackfillResponse | ServerRespons self.response = self.client.post( "backfills", json=backfill.model_dump(mode="json", exclude_none=True) ) - return BackfillResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, BackfillResponse) def create_dry_run(self, backfill: BackfillPostBody) -> BackfillResponse | ServerResponseError: """Create a dry run backfill.""" self.response = self.client.post( "backfills/dry_run", json=backfill.model_dump(mode="json", exclude_none=True) ) - return BackfillResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, BackfillResponse) def get(self, backfill_id: str) -> BackfillResponse | ServerResponseError: """Get a backfill.""" self.response = self.client.get(f"backfills/{backfill_id}") - return BackfillResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, BackfillResponse) def list(self, dag_id: str) -> BackfillCollectionResponse | ServerResponseError: """List all backfills.""" @@ -357,17 +365,17 @@ def list(self, dag_id: str) -> BackfillCollectionResponse | ServerResponseError: def pause(self, backfill_id: str) -> BackfillResponse | ServerResponseError: """Pause a backfill.""" self.response = self.client.post(f"backfills/{backfill_id}/pause") - return BackfillResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, BackfillResponse) def unpause(self, backfill_id: str) -> BackfillResponse | ServerResponseError: """Unpause a backfill.""" self.response = self.client.post(f"backfills/{backfill_id}/unpause") - return BackfillResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, BackfillResponse) def cancel(self, backfill_id: str) -> BackfillResponse | ServerResponseError: """Cancel a backfill.""" self.response = self.client.post(f"backfills/{backfill_id}/cancel") - return BackfillResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, BackfillResponse) class ConfigOperations(BaseOperations): @@ -376,12 +384,12 @@ class ConfigOperations(BaseOperations): def get(self, section: str, option: str) -> Config | ServerResponseError: """Get a config from the API server.""" self.response = self.client.get(f"/config/section/{section}/option/{option}") - return Config.model_validate_json(self.response.content) + return validate_response(self.response.content, Config) def list(self) -> Config | ServerResponseError: """List all configs from the API server.""" self.response = self.client.get("/config") - return Config.model_validate_json(self.response.content) + return validate_response(self.response.content, Config) class ConnectionsOperations(BaseOperations): @@ -390,7 +398,7 @@ class ConnectionsOperations(BaseOperations): def get(self, conn_id: str) -> ConnectionResponse | ServerResponseError: """Get a connection from the API server.""" self.response = self.client.get(f"connections/{conn_id}") - return ConnectionResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, ConnectionResponse) def list(self) -> ConnectionCollectionResponse | ServerResponseError: """List all connections from the API server.""" @@ -404,14 +412,14 @@ def create( self.response = self.client.post( "connections", json=connection.model_dump(mode="json", by_alias=True, exclude_none=True) ) - return ConnectionResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, ConnectionResponse) def bulk(self, connections: BulkBodyConnectionBody) -> BulkResponse | ServerResponseError: """CRUD multiple connections.""" self.response = self.client.patch( "connections", json=connections.model_dump(mode="json", by_alias=True) ) - return BulkResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, BulkResponse) def create_defaults(self) -> None | ServerResponseError: """Create default connections.""" @@ -432,7 +440,7 @@ def update( f"connections/{connection.connection_id}", json=connection.model_dump(mode="json", by_alias=True), ) - return ConnectionResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, ConnectionResponse) def test( self, @@ -442,7 +450,7 @@ def test( self.response = self.client.post( "connections/test", json=connection.model_dump(mode="json", by_alias=True) ) - return ConnectionTestResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, ConnectionTestResponse) class DagsOperations(BaseOperations): @@ -451,12 +459,12 @@ class DagsOperations(BaseOperations): def get(self, dag_id: str) -> DAGResponse | ServerResponseError: """Get a Dag.""" self.response = self.client.get(f"dags/{dag_id}") - return DAGResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, DAGResponse) def get_details(self, dag_id: str) -> DAGDetailsResponse | ServerResponseError: """Get a Dag details.""" self.response = self.client.get(f"dags/{dag_id}/details") - return DAGDetailsResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, DAGDetailsResponse) def get_tags(self) -> DAGTagCollectionResponse | ServerResponseError: """Get all Dag tags.""" @@ -468,7 +476,7 @@ def list(self) -> DAGCollectionResponse | ServerResponseError: def update(self, dag_id: str, dag_body: DAGPatchBody) -> DAGResponse | ServerResponseError: self.response = self.client.patch(f"dags/{dag_id}", json=dag_body.model_dump(mode="json")) - return DAGResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, DAGResponse) def delete(self, dag_id: str) -> str | ServerResponseError: self.client.delete(f"dags/{dag_id}") @@ -476,18 +484,18 @@ def delete(self, dag_id: str) -> str | ServerResponseError: def get_import_error(self, import_error_id: str) -> ImportErrorResponse | ServerResponseError: self.response = self.client.get(f"importErrors/{import_error_id}") - return ImportErrorResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, ImportErrorResponse) def list_import_errors(self) -> ImportErrorCollectionResponse | ServerResponseError: return super().execute_list(path="importErrors", data_model=ImportErrorCollectionResponse) def get_stats(self, dag_ids: list) -> DagStatsCollectionResponse | ServerResponseError: # type: ignore self.response = self.client.get("dagStats", params={"dag_ids": dag_ids}) - return DagStatsCollectionResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, DagStatsCollectionResponse) def get_version(self, dag_id: str, version_number: int) -> DagVersionResponse | ServerResponseError: self.response = self.client.get(f"dags/{dag_id}/dagVersions/{version_number}") - return DagVersionResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, DagVersionResponse) def list_version(self, dag_id: str) -> DAGVersionCollectionResponse | ServerResponseError: return super().execute_list( @@ -506,7 +514,7 @@ def trigger( self.response = self.client.post( f"dags/{dag_id}/dagRuns", json=trigger_dag_run.model_dump(mode="json") ) - return DAGRunResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, DAGRunResponse) class DagRunOperations(BaseOperations): @@ -520,7 +528,7 @@ def get( f"/dags/{dag_id}/dagRuns/{dag_run_id}", extensions={"airflowctl_suppress_error_log": suppress_error_log}, ) - return DAGRunResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, DAGRunResponse) def list( self, @@ -582,7 +590,7 @@ def list( params=params, extensions={"airflowctl_suppress_error_log": suppress_error_log}, ) - return DAGRunCollectionResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, DAGRunCollectionResponse) def delete(self, dag_id: str, dag_run_id: str) -> str | ServerResponseError: """Delete a Dag run.""" @@ -618,7 +626,7 @@ def list( if limit is not None or offset is not None: self.response = self.client.get("jobs", params=params) - return JobCollectionResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, JobCollectionResponse) return super().execute_list(path="jobs", data_model=JobCollectionResponse, params=params) @@ -629,7 +637,7 @@ class PoolsOperations(BaseOperations): def get(self, pool_name: str) -> PoolResponse | ServerResponseError: """Get a pool.""" self.response = self.client.get(f"pools/{pool_name}") - return PoolResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, PoolResponse) def list(self) -> PoolCollectionResponse | ServerResponseError: """List all pools.""" @@ -638,12 +646,12 @@ def list(self) -> PoolCollectionResponse | ServerResponseError: def create(self, pool: PoolBody) -> PoolResponse | ServerResponseError: """Create a pool.""" self.response = self.client.post("pools", json=pool.model_dump(mode="json", exclude_none=True)) - return PoolResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, PoolResponse) def bulk(self, pools: BulkBodyPoolBody) -> BulkResponse | ServerResponseError: """CRUD multiple pools.""" self.response = self.client.patch("pools", json=pools.model_dump(mode="json")) - return BulkResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, BulkResponse) def delete(self, pool: str) -> str | ServerResponseError: """Delete a pool.""" @@ -653,7 +661,7 @@ def delete(self, pool: str) -> str | ServerResponseError: def update(self, pool_body: PoolPatchBody) -> PoolResponse | ServerResponseError: """Update a pool.""" self.response = self.client.patch(f"pools/{pool_body.pool}", json=pool_body.model_dump(mode="json")) - return PoolResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, PoolResponse) class ProvidersOperations(BaseOperations): @@ -692,7 +700,7 @@ def get( path, extensions={"airflowctl_suppress_error_log": suppress_error_log}, ) - return TaskInstanceResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, TaskInstanceResponse) def get_dependencies( self, @@ -711,7 +719,7 @@ def get_dependencies( f"{path}/dependencies", extensions={"airflowctl_suppress_error_log": suppress_error_log}, ) - return TaskDependencyCollectionResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, TaskDependencyCollectionResponse) def list(self, dag_id: str, dag_run_id: str) -> TaskInstanceCollectionResponse | ServerResponseError: """List task instances for a Dag run.""" @@ -732,7 +740,7 @@ def clear( f"dags/{dag_id}/clearTaskInstances", json=clear_task_instances.model_dump(mode="json", exclude_none=True), ) - return TaskInstanceCollectionResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, TaskInstanceCollectionResponse) class VariablesOperations(BaseOperations): @@ -741,7 +749,7 @@ class VariablesOperations(BaseOperations): def get(self, variable_key: str) -> VariableResponse | ServerResponseError: """Get a variable.""" self.response = self.client.get(f"variables/{variable_key}") - return VariableResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, VariableResponse) def list(self) -> VariableCollectionResponse | ServerResponseError: """List all variables.""" @@ -752,12 +760,12 @@ def create(self, variable: VariableBody) -> VariableResponse | ServerResponseErr self.response = self.client.post( "variables", json=variable.model_dump(mode="json", exclude_none=True) ) - return VariableResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, VariableResponse) def bulk(self, variables: BulkBodyVariableBody) -> BulkResponse | ServerResponseError: """CRUD multiple variables.""" self.response = self.client.patch("variables", json=variables.model_dump(mode="json")) - return BulkResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, BulkResponse) def delete(self, variable_key: str) -> str | ServerResponseError: """Delete a variable.""" @@ -767,7 +775,7 @@ def delete(self, variable_key: str) -> str | ServerResponseError: def update(self, variable: VariableBody) -> VariableResponse | ServerResponseError: """Update a variable.""" self.response = self.client.patch(f"variables/{variable.key}", json=variable.model_dump(mode="json")) - return VariableResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, VariableResponse) class VersionOperations(BaseOperations): @@ -776,7 +784,7 @@ class VersionOperations(BaseOperations): def get(self) -> VersionInfo | ServerResponseError: """Get the version.""" self.response = self.client.get("version") - return VersionInfo.model_validate_json(self.response.content) + return validate_response(self.response.content, VersionInfo) class XComOperations(BaseOperations): @@ -798,7 +806,7 @@ def get( f"dags/{dag_id}/dagRuns/{dag_run_id}/taskInstances/{task_id}/xcomEntries/{key}", params=params, ) - return XComResponseNative.model_validate_json(self.response.content) + return validate_response(self.response.content, XComResponseNative) def list( self, @@ -843,7 +851,7 @@ def add( f"dags/{dag_id}/dagRuns/{dag_run_id}/taskInstances/{task_id}/xcomEntries", json=body.model_dump(mode="json", exclude_unset=True, exclude_none=True), ) - return XComResponseNative.model_validate_json(self.response.content) + return validate_response(self.response.content, XComResponseNative) def edit( self, @@ -868,7 +876,7 @@ def edit( f"dags/{dag_id}/dagRuns/{dag_run_id}/taskInstances/{task_id}/xcomEntries/{key}", json=body.model_dump(mode="json", exclude_unset=True, exclude_none=True), ) - return XComResponseNative.model_validate_json(self.response.content) + return validate_response(self.response.content, XComResponseNative) def delete( self, @@ -899,4 +907,4 @@ def list(self) -> PluginCollectionResponse | ServerResponseError: def list_import_errors(self) -> PluginImportErrorCollectionResponse | ServerResponseError: """List plugin import errors from the API server.""" self.response = self.client.get("plugins/importErrors") - return PluginImportErrorCollectionResponse.model_validate_json(self.response.content) + return validate_response(self.response.content, PluginImportErrorCollectionResponse) diff --git a/airflow-ctl/tests/airflow_ctl/api/test_operations.py b/airflow-ctl/tests/airflow_ctl/api/test_operations.py index 4868841bc2fd2..003c5472fd63a 100644 --- a/airflow-ctl/tests/airflow_ctl/api/test_operations.py +++ b/airflow-ctl/tests/airflow_ctl/api/test_operations.py @@ -1119,6 +1119,41 @@ def handle_request(request: httpx.Request) -> httpx.Response: response = client.dags.get_details("dag_id") assert response == self.dag_details_response + @pytest.mark.parametrize( + ("operation", "path", "payload_attr"), + [ + (lambda dags, patch_body: dags.get("dag_id"), "/api/v2/dags/dag_id", "dag_response"), + ( + lambda dags, patch_body: dags.get_details("dag_id"), + "/api/v2/dags/dag_id/details", + "dag_details_response", + ), + ( + lambda dags, patch_body: dags.update(dag_id="dag_id", dag_body=patch_body), + "/api/v2/dags/dag_id", + "dag_response", + ), + ], + ids=["get", "get_details", "update"], + ) + @pytest.mark.parametrize("omitted_field", ["is_backfillable", "timetable_periodic"]) + def test_single_object_response_tolerates_field_an_older_server_omits( + self, operation, path, payload_attr, omitted_field + ): + """A server on an older Airflow line omits fields the generated models declare as required.""" + payload = json.loads(getattr(self, payload_attr).model_dump_json()) + del payload[omitted_field] + + def handle_request(request: httpx.Request) -> httpx.Response: + assert request.url.path == path + return httpx.Response(200, json=payload) + + client = make_api_client(transport=httpx.MockTransport(handle_request)) + response = operation(client.dags, self.dag_patch_body) + + assert response.dag_id == self.dag_id + assert getattr(response, omitted_field) is False + def test_get_tags(self): def handle_request(request: httpx.Request) -> httpx.Response: assert request.url.path == "/api/v2/dagTags"