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]

Reply via email to