kaxil commented on code in PR #73986: URL: https://github.com/apache/airflow/pull/73986#discussion_r4150340263
########## providers/http/src/airflow/providers/http/triggers/external_task.py: ########## @@ -0,0 +1,223 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +from collections.abc import Collection, Iterable, Iterator +from functools import cached_property +from typing import TYPE_CHECKING, Any +from urllib.parse import quote + +from asgiref.sync import sync_to_async +from packaging.version import Version + +from airflow.providers.common.compat.sdk import AirflowOptionalProviderFeatureException +from airflow.providers.http.hooks.http import HttpHook +from airflow.providers.http.version_compat import AIRFLOW_V_3_0_PLUS + +if not AIRFLOW_V_3_0_PLUS: + raise AirflowOptionalProviderFeatureException("Waiting for a remote Airflow deployment needs Airflow 3+.") + +from airflow.providers.standard.triggers.external_task import WorkflowTrigger + +if TYPE_CHECKING: + from datetime import datetime + + from requests import Response, Session + +# The ``task_group_id`` filter of the task instances endpoint was added in Airflow 3.2.0. +# Older API servers silently ignore it and would return the task instances of all tasks. +MIN_TASK_GROUP_FILTER_VERSION = Version("3.2.0") + + +class _AirflowApiClient: + """ + Query the REST API (v2) of the remote Airflow deployment of an HTTP connection. + + Login and password of the connection are exchanged for a JWT access token of the remote auth + manager, a password without login is used as bearer token as-is. + """ + + page_limit = 100 + + def __init__(self, http_conn_id: str) -> None: + self.http_conn_id = http_conn_id + self._hook = HttpHook(method="GET", http_conn_id=http_conn_id) + self._session: Session | None = None + self._uses_access_token = False + self._remote_version: Version | None = None + + def _create_session(self) -> Session: + session = self._hook.get_conn() + # Airflow 3 API servers reject the basic auth HttpHook derives from the connection login. + session.auth = None + connection = self._hook.get_connection(self.http_conn_id) + self._uses_access_token = bool(connection.login) + if connection.login: + response = session.post( + self._hook.url_from_endpoint("auth/token"), + json={"username": connection.login, "password": connection.password}, + timeout=self._hook.merged_extra.get("timeout"), + ) + response.raise_for_status() + session.headers["Authorization"] = f"Bearer {response.json()['access_token']}" + elif connection.password: + session.headers["Authorization"] = f"Bearer {connection.password}" + return session + + def _request(self, method: str, endpoint: str, **kwargs: Any) -> Any: + if self._session is None: + self._session = self._create_session() + url = self._hook.url_from_endpoint(f"api/v2/{endpoint}") + timeout = self._hook.merged_extra.get("timeout") + response: Response = self._session.request(method, url, timeout=timeout, **kwargs) + if response.status_code == 401 and self._uses_access_token: + # Access tokens expire, which long-running sensors and triggers outlive. + self._session = self._create_session() + response = self._session.request(method, url, timeout=timeout, **kwargs) + response.raise_for_status() + return response.json() + + def get_dr_count( + self, dag_id: str, logical_dates: Iterable[datetime], states: Collection[str] | None + ) -> int: + return sum( + self._request( + "GET", + f"dags/{quote(dag_id, safe='')}/dagRuns", + params={ + "logical_date_gte": logical_date.isoformat(), + "logical_date_lte": logical_date.isoformat(), + "state": list(states or []), + "limit": 1, + }, + )["total_entries"] + for logical_date in logical_dates + ) + + def get_ti_count( + self, + dag_id: str, + task_ids: Collection[str], + logical_dates: Iterable[datetime], + states: Collection[str] | None, + ) -> int: + return sum( + self._request( + "POST", + "dags/~/dagRuns/~/taskInstances/list", Review Comment: Two side effects of the batch endpoint that the Dag and task group paths don't have. First, `POST .../taskInstances/list` carries `action_logging()`, so every poke writes one audit `Log` row per logical date and state set on the remote. At the default 60s poke interval that's a few thousand rows per sensor per day in their Audit Log. Second, it authorizes on `dag_id="~"` and then filters rows with `readable_ti_filter`, so a user without read access to the awaited Dag gets `total_entries: 0` rather than a 403 and the sensor just waits until timeout. `GET dags/{dag_id}/dagRuns/~/taskInstances` with `task_id`, `state`, `logical_date_gte/lte` and `limit=1` avoids both, at one call per task id. Those filters exist from 3.0.0. ########## providers/http/src/airflow/providers/http/triggers/external_task.py: ########## @@ -0,0 +1,223 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +from collections.abc import Collection, Iterable, Iterator +from functools import cached_property +from typing import TYPE_CHECKING, Any +from urllib.parse import quote + +from asgiref.sync import sync_to_async +from packaging.version import Version + +from airflow.providers.common.compat.sdk import AirflowOptionalProviderFeatureException +from airflow.providers.http.hooks.http import HttpHook +from airflow.providers.http.version_compat import AIRFLOW_V_3_0_PLUS + +if not AIRFLOW_V_3_0_PLUS: + raise AirflowOptionalProviderFeatureException("Waiting for a remote Airflow deployment needs Airflow 3+.") + +from airflow.providers.standard.triggers.external_task import WorkflowTrigger Review Comment: If someone upgrades the http provider without the `standard` extra, they keep whatever standard their Airflow install brought (core only pins `>=1.9.0`), and 1.20.0 has none of these seams: `_poke_af3` calls `ti.get_dr_count` / `get_ti_count` / `get_task_states` directly and `execute` builds `WorkflowTrigger` inline. `HttpExternalTaskSensor` still imports fine, so it silently polls the *local* deployment for `external_dag_id`. That either waits until timeout or succeeds against a local Dag with the same id, which is likely when two deployments share a Dag repo. A triggerer image with new http and old standard gets the same local polling through the old `_get_count_af_3`. Could this module fail fast next to the AF3 guard above, e.g. `if not hasattr(WorkflowTrigger, "_get_dr_count"): raise AirflowOptionalProviderFeatureException(...)` naming `apache-airflow-providers-http[standard]`? The sensor module imports this one first, so a single guard covers both the worker and the triggerer. A test that monkeypatches the attribute away and asserts the import raises would pin it. ########## providers/standard/src/airflow/providers/standard/sensors/external_task.py: ########## @@ -389,6 +376,41 @@ def _get_count(states: list[str]) -> int: count_allowed = self._calculate_count(count, dttm_filter) return count_allowed == len(dttm_filter) + def _get_dr_count( Review Comment: These four methods are now the extension contract `HttpExternalTaskSensor` overrides from a separately released distribution, but the leading `_` marks them private and free to rename, and a rename would bring back the silent local polling with nothing in standard's own tests noticing. Would it be worth making them a documented subclass hook, either dropping the `_` prefix or saying in the docstrings that subclasses override them to change where state comes from? Same question for the `WorkflowTrigger` trio. ########## providers/http/src/airflow/providers/http/triggers/external_task.py: ########## @@ -0,0 +1,223 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +from collections.abc import Collection, Iterable, Iterator +from functools import cached_property +from typing import TYPE_CHECKING, Any +from urllib.parse import quote + +from asgiref.sync import sync_to_async +from packaging.version import Version + +from airflow.providers.common.compat.sdk import AirflowOptionalProviderFeatureException +from airflow.providers.http.hooks.http import HttpHook +from airflow.providers.http.version_compat import AIRFLOW_V_3_0_PLUS + +if not AIRFLOW_V_3_0_PLUS: + raise AirflowOptionalProviderFeatureException("Waiting for a remote Airflow deployment needs Airflow 3+.") + +from airflow.providers.standard.triggers.external_task import WorkflowTrigger + +if TYPE_CHECKING: + from datetime import datetime + + from requests import Response, Session + +# The ``task_group_id`` filter of the task instances endpoint was added in Airflow 3.2.0. +# Older API servers silently ignore it and would return the task instances of all tasks. +MIN_TASK_GROUP_FILTER_VERSION = Version("3.2.0") Review Comment: The `task_group_id` filter on the task instances endpoint was backported to 3.1.6 (#59511); it isn't in 3.1.5 but is in 3.1.6 and later. So a 3.1.6 to 3.1.x remote would answer this correctly and still gets the ValueError. Should the floor be `Version("3.1.6")`, with the comment above, the operators.rst line and the sensor docstring saying the same? `test_get_task_group_states_requires_airflow_3_2` would need a 3.1.5 remote then. ########## providers/http/src/airflow/providers/http/triggers/external_task.py: ########## @@ -0,0 +1,223 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +from collections.abc import Collection, Iterable, Iterator +from functools import cached_property +from typing import TYPE_CHECKING, Any +from urllib.parse import quote + +from asgiref.sync import sync_to_async +from packaging.version import Version + +from airflow.providers.common.compat.sdk import AirflowOptionalProviderFeatureException +from airflow.providers.http.hooks.http import HttpHook +from airflow.providers.http.version_compat import AIRFLOW_V_3_0_PLUS + +if not AIRFLOW_V_3_0_PLUS: + raise AirflowOptionalProviderFeatureException("Waiting for a remote Airflow deployment needs Airflow 3+.") + +from airflow.providers.standard.triggers.external_task import WorkflowTrigger + +if TYPE_CHECKING: + from datetime import datetime + + from requests import Response, Session + +# The ``task_group_id`` filter of the task instances endpoint was added in Airflow 3.2.0. +# Older API servers silently ignore it and would return the task instances of all tasks. +MIN_TASK_GROUP_FILTER_VERSION = Version("3.2.0") + + +class _AirflowApiClient: + """ + Query the REST API (v2) of the remote Airflow deployment of an HTTP connection. + + Login and password of the connection are exchanged for a JWT access token of the remote auth + manager, a password without login is used as bearer token as-is. + """ + + page_limit = 100 + + def __init__(self, http_conn_id: str) -> None: + self.http_conn_id = http_conn_id + self._hook = HttpHook(method="GET", http_conn_id=http_conn_id) + self._session: Session | None = None + self._uses_access_token = False + self._remote_version: Version | None = None + + def _create_session(self) -> Session: + session = self._hook.get_conn() + # Airflow 3 API servers reject the basic auth HttpHook derives from the connection login. + session.auth = None + connection = self._hook.get_connection(self.http_conn_id) + self._uses_access_token = bool(connection.login) + if connection.login: + response = session.post( + self._hook.url_from_endpoint("auth/token"), + json={"username": connection.login, "password": connection.password}, + timeout=self._hook.merged_extra.get("timeout"), + ) + response.raise_for_status() + session.headers["Authorization"] = f"Bearer {response.json()['access_token']}" + elif connection.password: + session.headers["Authorization"] = f"Bearer {connection.password}" + return session + + def _request(self, method: str, endpoint: str, **kwargs: Any) -> Any: + if self._session is None: + self._session = self._create_session() + url = self._hook.url_from_endpoint(f"api/v2/{endpoint}") + timeout = self._hook.merged_extra.get("timeout") + response: Response = self._session.request(method, url, timeout=timeout, **kwargs) + if response.status_code == 401 and self._uses_access_token: + # Access tokens expire, which long-running sensors and triggers outlive. + self._session = self._create_session() + response = self._session.request(method, url, timeout=timeout, **kwargs) + response.raise_for_status() + return response.json() + + def get_dr_count( + self, dag_id: str, logical_dates: Iterable[datetime], states: Collection[str] | None + ) -> int: + return sum( + self._request( + "GET", + f"dags/{quote(dag_id, safe='')}/dagRuns", + params={ + "logical_date_gte": logical_date.isoformat(), + "logical_date_lte": logical_date.isoformat(), + "state": list(states or []), + "limit": 1, + }, + )["total_entries"] + for logical_date in logical_dates + ) + + def get_ti_count( + self, + dag_id: str, + task_ids: Collection[str], + logical_dates: Iterable[datetime], + states: Collection[str] | None, + ) -> int: + return sum( + self._request( + "POST", + "dags/~/dagRuns/~/taskInstances/list", + json={ + "dag_ids": [dag_id], + "task_ids": list(task_ids), + "state": list(states) if states else None, + "logical_date_gte": logical_date.isoformat(), + "logical_date_lte": logical_date.isoformat(), + "page_limit": 1, + }, + )["total_entries"] + for logical_date in logical_dates + ) + + def get_task_group_states( + self, dag_id: str, task_group_id: str, logical_dates: Iterable[datetime] + ) -> dict[str, dict[str, Any]]: + version = self._get_remote_version() + if version.release < MIN_TASK_GROUP_FILTER_VERSION.release: + raise ValueError( + f"Waiting for a task group requires the remote Airflow deployment to run Airflow " + f"{MIN_TASK_GROUP_FILTER_VERSION} or later, but it runs Airflow {version}." + ) + task_states: dict[str, dict[str, Any]] = {} + for logical_date in logical_dates: + for ti in self._iter_task_group_instances(dag_id, task_group_id, logical_date): + # Keyed like the Execution API's task states to reuse the same state matching. + key = ti["task_id"] if ti["map_index"] < 0 else f"{ti['task_id']}_{ti['map_index']}" + task_states.setdefault(ti["dag_run_id"], {})[key] = ti["state"] + return task_states + + def _iter_task_group_instances( + self, dag_id: str, task_group_id: str, logical_date: datetime + ) -> Iterator[dict[str, Any]]: + offset = 0 + while True: + page = self._request( + "GET", + f"dags/{quote(dag_id, safe='')}/dagRuns/~/taskInstances", + params={ + "task_group_id": task_group_id, + "logical_date_gte": logical_date.isoformat(), + "logical_date_lte": logical_date.isoformat(), + "order_by": "id", + "limit": self.page_limit, + "offset": offset, + }, + ) + yield from page["task_instances"] + offset += len(page["task_instances"]) + if not page["task_instances"] or offset >= page["total_entries"]: + return + + def _get_remote_version(self) -> Version: + if self._remote_version is None: + self._remote_version = Version(self._request("GET", "version")["version"]) + return self._remote_version + + +class HttpExternalTaskTrigger(WorkflowTrigger): + """ + Wait for a Dag, task group or task of a remote Airflow deployment to reach a state. + + Behaves like :class:`~airflow.providers.standard.triggers.external_task.WorkflowTrigger`, + but queries the REST API (v2) of the remote Airflow 3 deployment configured in the HTTP connection. + Only ``logical_dates`` are supported to select the remote Dag runs. + + :param http_conn_id: :ref:`http connection<howto/connection:http>` of the remote Airflow deployment. + + All other parameters are the same as those of + :class:`~airflow.providers.standard.triggers.external_task.WorkflowTrigger`. + """ + + def __init__(self, *, http_conn_id: str, **kwargs: Any) -> None: + super().__init__(**kwargs) + self.http_conn_id = http_conn_id + + def serialize(self) -> tuple[str, dict[str, Any]]: + classpath, data = super().serialize() + return classpath, {**data, "http_conn_id": self.http_conn_id} + + @cached_property + def _client(self) -> _AirflowApiClient: + return _AirflowApiClient(self.http_conn_id) + + async def _get_dr_count(self, states: Collection[str] | None) -> int: + return await sync_to_async(self._client.get_dr_count)( Review Comment: `sync_to_async` defaults to `thread_sensitive=True`, and the triggerer runs under `asyncio.run` with no `AsyncToSync` parent, so these three calls land on asgiref's single shared executor thread. Every other trigger in the process that uses `sync_to_async` runs on that thread too: the local `WorkflowTrigger`, other providers' triggers, sync secrets-backend lookups. The requests have no timeout unless the connection extra sets one (lines 74 and 86). So a remote API server that accepts the connection and then stalls parks that thread and stalls all of them, and cancelling this trigger waits on the socket as well. `HttpAsyncHook` with a `ClientTimeout` would match `HttpSensorTrigger` in this provider. Short of that, `thread_sensitive=False` (or `asyncio.to_thread`) on all three calls plus a default timeout in `_request` and `_create_session` would stop one bad remote from blocking the rest. ########## providers/http/src/airflow/providers/http/triggers/external_task.py: ########## @@ -0,0 +1,223 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +from collections.abc import Collection, Iterable, Iterator +from functools import cached_property +from typing import TYPE_CHECKING, Any +from urllib.parse import quote + +from asgiref.sync import sync_to_async +from packaging.version import Version + +from airflow.providers.common.compat.sdk import AirflowOptionalProviderFeatureException +from airflow.providers.http.hooks.http import HttpHook +from airflow.providers.http.version_compat import AIRFLOW_V_3_0_PLUS + +if not AIRFLOW_V_3_0_PLUS: + raise AirflowOptionalProviderFeatureException("Waiting for a remote Airflow deployment needs Airflow 3+.") + +from airflow.providers.standard.triggers.external_task import WorkflowTrigger + +if TYPE_CHECKING: + from datetime import datetime + + from requests import Response, Session + +# The ``task_group_id`` filter of the task instances endpoint was added in Airflow 3.2.0. +# Older API servers silently ignore it and would return the task instances of all tasks. +MIN_TASK_GROUP_FILTER_VERSION = Version("3.2.0") + + +class _AirflowApiClient: + """ + Query the REST API (v2) of the remote Airflow deployment of an HTTP connection. + + Login and password of the connection are exchanged for a JWT access token of the remote auth + manager, a password without login is used as bearer token as-is. + """ + + page_limit = 100 + + def __init__(self, http_conn_id: str) -> None: + self.http_conn_id = http_conn_id + self._hook = HttpHook(method="GET", http_conn_id=http_conn_id) + self._session: Session | None = None + self._uses_access_token = False + self._remote_version: Version | None = None + + def _create_session(self) -> Session: + session = self._hook.get_conn() + # Airflow 3 API servers reject the basic auth HttpHook derives from the connection login. + session.auth = None + connection = self._hook.get_connection(self.http_conn_id) + self._uses_access_token = bool(connection.login) + if connection.login: + response = session.post( + self._hook.url_from_endpoint("auth/token"), + json={"username": connection.login, "password": connection.password}, + timeout=self._hook.merged_extra.get("timeout"), + ) + response.raise_for_status() + session.headers["Authorization"] = f"Bearer {response.json()['access_token']}" + elif connection.password: + session.headers["Authorization"] = f"Bearer {connection.password}" + return session + + def _request(self, method: str, endpoint: str, **kwargs: Any) -> Any: + if self._session is None: + self._session = self._create_session() + url = self._hook.url_from_endpoint(f"api/v2/{endpoint}") + timeout = self._hook.merged_extra.get("timeout") + response: Response = self._session.request(method, url, timeout=timeout, **kwargs) + if response.status_code == 401 and self._uses_access_token: Review Comment: Only an expired token gets a second attempt. A connection reset, a read timeout, or a 429/503 from the remote raises straight out of `run()`, so one blip fails the trigger and the wait restarts on the next try, or fails the sensor outright with `retries=0`. The local path this stands in for goes through the Execution API client, which retries on its own (`API_RETRIES` in task-sdk `api/client.py`). For a sensor that waits across a network for hours, would passing `HttpHook(adapter=HTTPAdapter(max_retries=Retry(total=..., backoff_factor=..., status_forcelist=(429, 502, 503, 504), allowed_methods=None)))` in `__init__` be enough? `allowed_methods=None` matters because the task count is a POST, which urllib3 doesn't retry by default, and `Retry` already honours `Retry-After` on 429/503. ########## providers/http/src/airflow/providers/http/triggers/external_task.py: ########## @@ -0,0 +1,223 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +from collections.abc import Collection, Iterable, Iterator +from functools import cached_property +from typing import TYPE_CHECKING, Any +from urllib.parse import quote + +from asgiref.sync import sync_to_async +from packaging.version import Version + +from airflow.providers.common.compat.sdk import AirflowOptionalProviderFeatureException +from airflow.providers.http.hooks.http import HttpHook +from airflow.providers.http.version_compat import AIRFLOW_V_3_0_PLUS + +if not AIRFLOW_V_3_0_PLUS: + raise AirflowOptionalProviderFeatureException("Waiting for a remote Airflow deployment needs Airflow 3+.") + +from airflow.providers.standard.triggers.external_task import WorkflowTrigger + +if TYPE_CHECKING: + from datetime import datetime + + from requests import Response, Session + +# The ``task_group_id`` filter of the task instances endpoint was added in Airflow 3.2.0. +# Older API servers silently ignore it and would return the task instances of all tasks. +MIN_TASK_GROUP_FILTER_VERSION = Version("3.2.0") + + +class _AirflowApiClient: + """ + Query the REST API (v2) of the remote Airflow deployment of an HTTP connection. + + Login and password of the connection are exchanged for a JWT access token of the remote auth + manager, a password without login is used as bearer token as-is. + """ + + page_limit = 100 + + def __init__(self, http_conn_id: str) -> None: + self.http_conn_id = http_conn_id + self._hook = HttpHook(method="GET", http_conn_id=http_conn_id) + self._session: Session | None = None + self._uses_access_token = False + self._remote_version: Version | None = None + + def _create_session(self) -> Session: + session = self._hook.get_conn() + # Airflow 3 API servers reject the basic auth HttpHook derives from the connection login. + session.auth = None + connection = self._hook.get_connection(self.http_conn_id) + self._uses_access_token = bool(connection.login) + if connection.login: + response = session.post( + self._hook.url_from_endpoint("auth/token"), + json={"username": connection.login, "password": connection.password}, + timeout=self._hook.merged_extra.get("timeout"), + ) + response.raise_for_status() + session.headers["Authorization"] = f"Bearer {response.json()['access_token']}" + elif connection.password: + session.headers["Authorization"] = f"Bearer {connection.password}" + return session + + def _request(self, method: str, endpoint: str, **kwargs: Any) -> Any: + if self._session is None: + self._session = self._create_session() + url = self._hook.url_from_endpoint(f"api/v2/{endpoint}") + timeout = self._hook.merged_extra.get("timeout") + response: Response = self._session.request(method, url, timeout=timeout, **kwargs) + if response.status_code == 401 and self._uses_access_token: + # Access tokens expire, which long-running sensors and triggers outlive. + self._session = self._create_session() + response = self._session.request(method, url, timeout=timeout, **kwargs) + response.raise_for_status() + return response.json() + + def get_dr_count( + self, dag_id: str, logical_dates: Iterable[datetime], states: Collection[str] | None + ) -> int: + return sum( + self._request( + "GET", + f"dags/{quote(dag_id, safe='')}/dagRuns", + params={ + "logical_date_gte": logical_date.isoformat(), + "logical_date_lte": logical_date.isoformat(), + "state": list(states or []), + "limit": 1, + }, + )["total_entries"] + for logical_date in logical_dates + ) + + def get_ti_count( + self, + dag_id: str, + task_ids: Collection[str], + logical_dates: Iterable[datetime], + states: Collection[str] | None, + ) -> int: + return sum( + self._request( + "POST", + "dags/~/dagRuns/~/taskInstances/list", + json={ + "dag_ids": [dag_id], + "task_ids": list(task_ids), + "state": list(states) if states else None, + "logical_date_gte": logical_date.isoformat(), + "logical_date_lte": logical_date.isoformat(), + "page_limit": 1, + }, + )["total_entries"] + for logical_date in logical_dates + ) + + def get_task_group_states( + self, dag_id: str, task_group_id: str, logical_dates: Iterable[datetime] + ) -> dict[str, dict[str, Any]]: + version = self._get_remote_version() + if version.release < MIN_TASK_GROUP_FILTER_VERSION.release: + raise ValueError( + f"Waiting for a task group requires the remote Airflow deployment to run Airflow " + f"{MIN_TASK_GROUP_FILTER_VERSION} or later, but it runs Airflow {version}." + ) + task_states: dict[str, dict[str, Any]] = {} + for logical_date in logical_dates: + for ti in self._iter_task_group_instances(dag_id, task_group_id, logical_date): + # Keyed like the Execution API's task states to reuse the same state matching. + key = ti["task_id"] if ti["map_index"] < 0 else f"{ti['task_id']}_{ti['map_index']}" + task_states.setdefault(ti["dag_run_id"], {})[key] = ti["state"] + return task_states + + def _iter_task_group_instances( + self, dag_id: str, task_group_id: str, logical_date: datetime + ) -> Iterator[dict[str, Any]]: + offset = 0 + while True: + page = self._request( + "GET", + f"dags/{quote(dag_id, safe='')}/dagRuns/~/taskInstances", + params={ + "task_group_id": task_group_id, + "logical_date_gte": logical_date.isoformat(), + "logical_date_lte": logical_date.isoformat(), + "order_by": "id", + "limit": self.page_limit, + "offset": offset, + }, + ) + yield from page["task_instances"] + offset += len(page["task_instances"]) + if not page["task_instances"] or offset >= page["total_entries"]: + return + + def _get_remote_version(self) -> Version: + if self._remote_version is None: + self._remote_version = Version(self._request("GET", "version")["version"]) + return self._remote_version + + +class HttpExternalTaskTrigger(WorkflowTrigger): + """ + Wait for a Dag, task group or task of a remote Airflow deployment to reach a state. + + Behaves like :class:`~airflow.providers.standard.triggers.external_task.WorkflowTrigger`, + but queries the REST API (v2) of the remote Airflow 3 deployment configured in the HTTP connection. + Only ``logical_dates`` are supported to select the remote Dag runs. + + :param http_conn_id: :ref:`http connection<howto/connection:http>` of the remote Airflow deployment. + + All other parameters are the same as those of + :class:`~airflow.providers.standard.triggers.external_task.WorkflowTrigger`. + """ + + def __init__(self, *, http_conn_id: str, **kwargs: Any) -> None: Review Comment: `run_ids` still reaches `WorkflowTrigger.__init__` through `**kwargs`, and the base `run()` compares the count against `len(self.run_ids or self.logical_dates)`, but these overrides only count over `logical_dates`. So constructing this trigger with `run_ids` counts 0 against `len(run_ids)` every time and never fires. The sensor never passes it, but the trigger is public. Should `__init__` raise if `run_ids` is set, matching the docstring's "only `logical_dates` are supported"? ########## providers/standard/tests/unit/standard/triggers/test_external_task.py: ########## @@ -409,6 +409,38 @@ def test_serialization(self): "soft_fail": False, } + def test_serialization_of_subclass(self): + classpath, _ = _WorkflowTriggerSubclass(external_dag_id=self.DAG_ID).serialize() + + assert classpath == f"{__name__}._WorkflowTriggerSubclass" + + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("kwargs", "method", "return_value"), + [ + ({}, "_get_dr_count", 1), + ({"external_task_ids": ["t1", "t2"]}, "_get_ti_count", 2), + ({"external_task_group_id": "g"}, "_get_task_group_states", {"run_id": {"g.t1": "success"}}), + ], + ids=["dag", "tasks", "task_group"], + ) + @mock.patch("airflow.sdk.execution_time.task_runner.RuntimeTaskInstance", autospec=True) + async def test_run_uses_overridable_state_access(self, mock_ti, kwargs, method, return_value): + trigger = WorkflowTrigger( + external_dag_id=self.DAG_ID, run_ids=[self.RUN_ID], allowed_states=["success"], **kwargs + ) + + with mock.patch.object(WorkflowTrigger, method, autospec=True, return_value=return_value) as m: + event = await anext(aiter(trigger.run())) + + assert event == TriggerEvent({"status": "success"}) + m.assert_awaited_once() Review Comment: A bare `assert_awaited_once()` lets the states argument go wrong unnoticed: changing `_get_count_af_3` to call `self._get_dr_count(None)` passes every standard test. The sensor twin of this test checks its arguments; could this one take the expected args from the parametrize (`(trigger, ["success"])` for dag and tasks, `(trigger,)` for the task group) and use `assert_awaited_once_with`? ########## providers/http/tests/unit/http/sensors/test_external_task.py: ########## @@ -0,0 +1,114 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +from datetime import datetime, timedelta, timezone +from unittest import mock + +import pytest + +from tests_common.test_utils.version_compat import AIRFLOW_V_3_0_PLUS + +if not AIRFLOW_V_3_0_PLUS: + pytest.skip("Waiting for a remote Airflow deployment needs Airflow 3+", allow_module_level=True) + +from airflow.providers.common.compat.sdk import TaskDeferred +from airflow.providers.http.sensors.external_task import HttpExternalTaskSensor +from airflow.providers.http.triggers.external_task import HttpExternalTaskTrigger +from airflow.providers.standard.exceptions import ExternalTaskFailedError +from airflow.sdk.execution_time.task_runner import RuntimeTaskInstance + +LOGICAL_DATE = datetime(2026, 1, 2, tzinfo=timezone.utc) + + [email protected] +def client(): + with mock.patch( + "airflow.providers.http.sensors.external_task._AirflowApiClient", autospec=True + ) as client_class: + yield client_class.return_value + + [email protected] +def context(): + return {"logical_date": LOGICAL_DATE, "ti": mock.create_autospec(RuntimeTaskInstance, instance=True)} + + +def create_sensor(**kwargs) -> HttpExternalTaskSensor: + return HttpExternalTaskSensor( + task_id="wait", http_conn_id="remote_airflow", external_dag_id="remote_dag", **kwargs + ) + + +class TestHttpExternalTaskSensor: + def test_attributes(self): + sensor = create_sensor() + + assert sensor._client.http_conn_id == "remote_airflow" + assert "http_conn_id" in sensor.template_fields + assert not sensor.operator_extra_links + + @pytest.mark.parametrize(("count", "expected"), [(1, True), (0, False)]) + def test_poke_dag(self, client, context, count, expected): + client.get_dr_count.return_value = count + + assert create_sensor(execution_delta=timedelta(days=1)).poke(context) is expected + client.get_dr_count.assert_called_once_with( + "remote_dag", [LOGICAL_DATE - timedelta(days=1)], ["success"] + ) + context["ti"].get_dr_count.assert_not_called() + + def test_poke_tasks(self, client, context): + client.get_ti_count.return_value = 1 + + with pytest.raises(ExternalTaskFailedError): + create_sensor(external_task_ids=["t1", "t2"], failed_states=["failed"]).poke(context) + client.get_ti_count.assert_called_once_with("remote_dag", ["t1", "t2"], [LOGICAL_DATE], ["failed"]) + context["ti"].get_ti_count.assert_not_called() + + @pytest.mark.parametrize( + ("task_states", "expected"), + [ + ({"run_1": {"g.a": "success", "g.b_0": "success"}}, True), + ({"run_1": {"g.a": "success", "g.b_0": "running"}}, False), + ], + ids=["all_allowed", "partially_allowed"], + ) + def test_poke_task_group(self, client, context, task_states, expected): + client.get_task_group_states.return_value = task_states + + assert create_sensor(external_task_group_id="g").poke(context) is expected + client.get_task_group_states.assert_called_once_with("remote_dag", "g", [LOGICAL_DATE]) + context["ti"].get_task_states.assert_not_called() + + def test_execute_deferrable(self, context): + sensor = create_sensor( + external_task_ids=["t1"], failed_states=["failed"], deferrable=True, poke_interval=30 + ) + + with pytest.raises(TaskDeferred) as deferred: + sensor.execute(context) + + trigger = deferred.value.trigger + assert isinstance(trigger, HttpExternalTaskTrigger) + assert trigger.http_conn_id == "remote_airflow" + assert trigger.external_dag_id == "remote_dag" + assert trigger.external_task_ids == ["t1"] + assert trigger.failed_states == ["failed"] + assert trigger.logical_dates == [LOGICAL_DATE] + assert trigger.poke_interval == 30 Review Comment: `_get_trigger` forwards ten kwargs and this checks six. Dropping `allowed_states`, `skipped_states`, `soft_fail` or `external_task_group_id` from `_get_trigger` still passes the whole http suite, and a dropped `allowed_states` is the bad one: the remote queries go out with no state filter and the sensor succeeds while the remote run is still queued. Could this pass non-default values for all of them and compare `trigger.serialize()` in one assertion? -- This is an automated message from the Apache Git Service. To respond to the message, please log on to GitHub and use the URL above to go to the specific comment. To unsubscribe, e-mail: [email protected] For queries about this service, please contact Infrastructure at: [email protected]
