From e8d955bcb73be28233e789585c0f3d3c2ff6089c Mon Sep 17 00:00:00 2001 From: ColtenOuO Date: Thu, 27 Aug 2026 06:35:33 +0000 Subject: [PATCH 1/6] Add WebSocketSensor and WebSocketTrigger to the standard provider Deferrable operators that hand a long-lived request off to a remote server have no way to resume once that server replies over a WebSocket connection; a plain HTTP request isn't well suited to holding a long-lived connection open. This adds a generic WebSocketTrigger and WebSocketSensor to the standard provider (not vendor-specific, alongside FileTrigger/FileSensor), gated behind an optional `websocket` extra so the `websockets` dependency isn't pulled in for users who don't need it. --- providers/standard/docs/index.rst | 1 + providers/standard/docs/sensors/websocket.rst | 43 +++++++++ providers/standard/provider.yaml | 3 + providers/standard/pyproject.toml | 4 + .../standard/example_dags/example_sensors.py | 18 ++++ .../providers/standard/get_provider_info.py | 3 + .../providers/standard/sensors/websocket.py | 92 +++++++++++++++++++ .../providers/standard/triggers/websocket.py | 75 +++++++++++++++ .../unit/standard/sensors/test_websocket.py | 67 ++++++++++++++ .../unit/standard/triggers/test_websocket.py | 77 ++++++++++++++++ uv.lock | 12 ++- 11 files changed, 392 insertions(+), 3 deletions(-) create mode 100644 providers/standard/docs/sensors/websocket.rst create mode 100644 providers/standard/src/airflow/providers/standard/sensors/websocket.py create mode 100644 providers/standard/src/airflow/providers/standard/triggers/websocket.py create mode 100644 providers/standard/tests/unit/standard/sensors/test_websocket.py create mode 100644 providers/standard/tests/unit/standard/triggers/test_websocket.py diff --git a/providers/standard/docs/index.rst b/providers/standard/docs/index.rst index a6b8f12bab228..ffd6d9a48dc40 100644 --- a/providers/standard/docs/index.rst +++ b/providers/standard/docs/index.rst @@ -127,6 +127,7 @@ Install them when installing from PyPI. For example: Extra Dependencies =============== ======================================== ``openlineage`` ``apache-airflow-providers-openlineage`` +``websocket`` ``websockets>=14.0`` =============== ======================================== Downloading official packages diff --git a/providers/standard/docs/sensors/websocket.rst b/providers/standard/docs/sensors/websocket.rst new file mode 100644 index 0000000000000..18ab2c16fa91b --- /dev/null +++ b/providers/standard/docs/sensors/websocket.rst @@ -0,0 +1,43 @@ + .. 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. + + + +.. _howto/operator:WebSocketSensor: + +WebSocketSensor +================ + +Use the :class:`~airflow.providers.standard.sensors.websocket.WebSocketSensor` to wait for a message on a +``ws://`` or ``wss://`` WebSocket connection. This is useful when a remote server accepts a long-lived +connection and replies asynchronously once it has handled the requested work, since a plain HTTP request +is not well suited to that kind of long-lived connection. Requires the ``websocket`` extra +(``apache-airflow-providers-standard[websocket]``). + +.. exampleinclude:: /../src/airflow/providers/standard/example_dags/example_sensors.py + :language: python + :dedent: 4 + :start-after: [START example_websocket_sensor] + :end-before: [END example_websocket_sensor] + +Also for this job you can use sensor in the deferrable mode: + +.. exampleinclude:: /../src/airflow/providers/standard/example_dags/example_sensors.py + :language: python + :dedent: 4 + :start-after: [START example_websocket_sensor_async] + :end-before: [END example_websocket_sensor_async] diff --git a/providers/standard/provider.yaml b/providers/standard/provider.yaml index d593430a0ffd2..40c39ad3caea1 100644 --- a/providers/standard/provider.yaml +++ b/providers/standard/provider.yaml @@ -83,6 +83,7 @@ integrations: - /docs/apache-airflow-providers-standard/sensors/datetime.rst - /docs/apache-airflow-providers-standard/sensors/file.rst - /docs/apache-airflow-providers-standard/sensors/external_task_sensor.rst + - /docs/apache-airflow-providers-standard/sensors/websocket.rst operators: - integration-name: Standard @@ -108,6 +109,7 @@ sensors: - airflow.providers.standard.sensors.python - airflow.providers.standard.sensors.filesystem - airflow.providers.standard.sensors.external_task + - airflow.providers.standard.sensors.websocket hooks: - integration-name: Standard python-modules: @@ -122,6 +124,7 @@ triggers: - airflow.providers.standard.triggers.file - airflow.providers.standard.triggers.temporal - airflow.providers.standard.triggers.hitl + - airflow.providers.standard.triggers.websocket extra-links: - airflow.providers.standard.operators.trigger_dagrun.TriggerDagRunLink diff --git a/providers/standard/pyproject.toml b/providers/standard/pyproject.toml index 04e8c68d8c91a..c19b23befef29 100644 --- a/providers/standard/pyproject.toml +++ b/providers/standard/pyproject.toml @@ -69,6 +69,9 @@ dependencies = [ "openlineage" = [ "apache-airflow-providers-openlineage" ] +"websocket" = [ + "websockets>=14.0" +] [dependency-groups] dev = [ @@ -79,6 +82,7 @@ dev = [ "apache-airflow-providers-openlineage", # Additional devel dependencies (do not remove this line and add extra development dependencies) "apache-airflow-providers-mysql", + "websockets>=14.0", ] # To build docs: diff --git a/providers/standard/src/airflow/providers/standard/example_dags/example_sensors.py b/providers/standard/src/airflow/providers/standard/example_dags/example_sensors.py index 73ccaadbe0f02..7e9363e7ce53f 100644 --- a/providers/standard/src/airflow/providers/standard/example_dags/example_sensors.py +++ b/providers/standard/src/airflow/providers/standard/example_dags/example_sensors.py @@ -28,6 +28,7 @@ from airflow.providers.standard.sensors.python import PythonSensor from airflow.providers.standard.sensors.time import TimeSensor from airflow.providers.standard.sensors.time_delta import TimeDeltaSensor +from airflow.providers.standard.sensors.websocket import WebSocketSensor from airflow.providers.standard.sensors.weekday import DayOfWeekSensor from airflow.providers.standard.utils.weekday import WeekDay from airflow.sdk import DAG @@ -125,6 +126,22 @@ def failure_callable(): ) # [END example_day_of_week_sensor] + # [START example_websocket_sensor] + t12 = WebSocketSensor( + task_id="wait_for_websocket_message", url="wss://example.com/socket", timeout=3, soft_fail=True + ) + # [END example_websocket_sensor] + + # [START example_websocket_sensor_async] + t13 = WebSocketSensor( + task_id="wait_for_websocket_message_async", + url="wss://example.com/socket", + deferrable=True, + timeout=3, + soft_fail=True, + ) + # [END example_websocket_sensor_async] + tx = BashOperator(task_id="print_date_in_bash", bash_command="date") tx.trigger_rule = TriggerRule.NONE_FAILED @@ -133,3 +150,4 @@ def failure_callable(): t8 >> tx [t9, t10] >> tx t11 >> tx + [t12, t13] >> tx diff --git a/providers/standard/src/airflow/providers/standard/get_provider_info.py b/providers/standard/src/airflow/providers/standard/get_provider_info.py index 1f7b2049454d1..0be729079a0cb 100644 --- a/providers/standard/src/airflow/providers/standard/get_provider_info.py +++ b/providers/standard/src/airflow/providers/standard/get_provider_info.py @@ -43,6 +43,7 @@ def get_provider_info(): "/docs/apache-airflow-providers-standard/sensors/datetime.rst", "/docs/apache-airflow-providers-standard/sensors/file.rst", "/docs/apache-airflow-providers-standard/sensors/external_task_sensor.rst", + "/docs/apache-airflow-providers-standard/sensors/websocket.rst", ], } ], @@ -75,6 +76,7 @@ def get_provider_info(): "airflow.providers.standard.sensors.python", "airflow.providers.standard.sensors.filesystem", "airflow.providers.standard.sensors.external_task", + "airflow.providers.standard.sensors.websocket", ], } ], @@ -96,6 +98,7 @@ def get_provider_info(): "airflow.providers.standard.triggers.file", "airflow.providers.standard.triggers.temporal", "airflow.providers.standard.triggers.hitl", + "airflow.providers.standard.triggers.websocket", ], } ], diff --git a/providers/standard/src/airflow/providers/standard/sensors/websocket.py b/providers/standard/src/airflow/providers/standard/sensors/websocket.py new file mode 100644 index 0000000000000..57c75ad409109 --- /dev/null +++ b/providers/standard/src/airflow/providers/standard/sensors/websocket.py @@ -0,0 +1,92 @@ +# 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 + +import datetime +from collections.abc import Sequence +from typing import TYPE_CHECKING, Any + +from websockets.sync.client import connect + +from airflow.providers.common.compat.sdk import BaseSensorOperator, conf +from airflow.providers.standard.triggers.websocket import WebSocketTrigger + +if TYPE_CHECKING: + from airflow.sdk import Context + + +class WebSocketSensor(BaseSensorOperator): + """ + Waits for a message on a WebSocket connection. + + :param url: The ``ws://`` or ``wss://`` URL of the WebSocket server to connect to. + :param header: Optional headers sent when opening the connection. + :param message_to_send: Optional message sent right after the connection is established. + :param deferrable: If waiting for completion, whether to defer the task until done, + default is ``False``. + + .. seealso:: + For more information on how to use this sensor, take a look at the guide: + :ref:`howto/operator:WebSocketSensor` + """ + + template_fields: Sequence[str] = ("url",) + + def __init__( + self, + *, + url: str, + header: dict[str, str] | None = None, + message_to_send: str | bytes | None = None, + deferrable: bool = conf.getboolean("operators", "default_deferrable", fallback=False), + **kwargs, + ): + super().__init__(**kwargs) + self.url = url + self.header = header + self.message_to_send = message_to_send + self.deferrable = deferrable + + def poke(self, context: Context) -> bool: + self.log.info("Poking WebSocket %s", self.url) + with connect(self.url, additional_headers=self.header) as websocket: + if self.message_to_send is not None: + websocket.send(self.message_to_send) + try: + websocket.recv(timeout=self.poke_interval) + except TimeoutError: + return False + self.log.info("Received message from %s", self.url) + return True + + def execute(self, context: Context) -> None: + if not self.deferrable: + super().execute(context=context) + if not self.poke(context=context): + self.defer( + timeout=datetime.timedelta(seconds=self.timeout), + trigger=WebSocketTrigger( + url=self.url, + header=self.header, + message_to_send=self.message_to_send, + ), + method_name="execute_complete", + ) + + def execute_complete(self, context: Context, event: Any = None) -> None: + """Handle the event when the trigger fires and return immediately.""" + self.log.info("%s completed successfully with message: %s", self.task_id, event) diff --git a/providers/standard/src/airflow/providers/standard/triggers/websocket.py b/providers/standard/src/airflow/providers/standard/triggers/websocket.py new file mode 100644 index 0000000000000..d3324469a4beb --- /dev/null +++ b/providers/standard/src/airflow/providers/standard/triggers/websocket.py @@ -0,0 +1,75 @@ +# 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 AsyncIterator +from typing import Any + +from websockets.asyncio.client import connect + +from airflow.providers.standard.version_compat import AIRFLOW_V_3_0_PLUS + +if AIRFLOW_V_3_0_PLUS: + from airflow.triggers.base import BaseEventTrigger, TriggerEvent +else: + from airflow.triggers.base import BaseTrigger as BaseEventTrigger, TriggerEvent # type: ignore + + +class WebSocketTrigger(BaseEventTrigger): + """ + A trigger that opens a WebSocket connection and fires once a message is received. + + This is meant for deferrable operators that hand off a long-lived request to a remote + WebSocket server and resume once that server replies, without occupying a worker slot + while waiting. + + :param url: The ``ws://`` or ``wss://`` URL of the WebSocket server to connect to. + :param header: Optional headers sent when opening the connection. + :param message_to_send: Optional message sent right after the connection is established. + """ + + def __init__( + self, + url: str, + header: dict[str, str] | None = None, + message_to_send: str | bytes | None = None, + **kwargs, + ): + super().__init__() + self.url = url + self.header = header + self.message_to_send = message_to_send + + def serialize(self) -> tuple[str, dict[str, Any]]: + """Serialize WebSocketTrigger arguments and classpath.""" + return ( + "airflow.providers.standard.triggers.websocket.WebSocketTrigger", + { + "url": self.url, + "header": self.header, + "message_to_send": self.message_to_send, + }, + ) + + async def run(self) -> AsyncIterator[TriggerEvent]: + """Connect to the WebSocket server and wait for the first message.""" + async with connect(self.url, additional_headers=self.header) as websocket: + if self.message_to_send is not None: + await websocket.send(self.message_to_send) + message = await websocket.recv() + self.log.info("Received message from %s", self.url) + yield TriggerEvent(message) diff --git a/providers/standard/tests/unit/standard/sensors/test_websocket.py b/providers/standard/tests/unit/standard/sensors/test_websocket.py new file mode 100644 index 0000000000000..3e8848b70520d --- /dev/null +++ b/providers/standard/tests/unit/standard/sensors/test_websocket.py @@ -0,0 +1,67 @@ +# 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 unittest import mock + +import pytest + +from airflow.models.dag import DAG +from airflow.providers.common.compat.sdk import TaskDeferred +from airflow.providers.standard.sensors.websocket import WebSocketSensor +from airflow.providers.standard.triggers.websocket import WebSocketTrigger + +from tests_common.test_utils.version_compat import timezone + +URL = "ws://example.com/socket" +DEFAULT_DATE = timezone.datetime(2015, 1, 1) + + +class TestWebSocketSensor: + @classmethod + def setup_class(cls): + args = {"owner": "airflow", "start_date": DEFAULT_DATE} + cls.dag = DAG("test_websocket_sensor", schedule=None, default_args=args) + + @mock.patch("airflow.providers.standard.sensors.websocket.connect") + def test_poke_returns_true_on_message(self, mock_connect): + mock_websocket = mock.MagicMock() + mock_websocket.recv.return_value = "pong" + mock_connect.return_value.__enter__.return_value = mock_websocket + + sensor = WebSocketSensor(task_id="poke_true", url=URL, message_to_send="ping", dag=self.dag) + assert sensor.poke(context={}) is True + mock_websocket.send.assert_called_once_with("ping") + + @mock.patch("airflow.providers.standard.sensors.websocket.connect") + def test_poke_returns_false_on_timeout(self, mock_connect): + mock_websocket = mock.MagicMock() + mock_websocket.recv.side_effect = TimeoutError() + mock_connect.return_value.__enter__.return_value = mock_websocket + + sensor = WebSocketSensor(task_id="poke_false", url=URL, dag=self.dag) + assert sensor.poke(context={}) is False + + def test_task_defer(self): + sensor = WebSocketSensor(task_id="defer", url=URL, deferrable=True, dag=self.dag) + + with mock.patch.object(WebSocketSensor, "poke", return_value=False): + with pytest.raises(TaskDeferred) as exc: + sensor.execute({}) + + assert isinstance(exc.value.trigger, WebSocketTrigger) + assert exc.value.trigger.url == URL diff --git a/providers/standard/tests/unit/standard/triggers/test_websocket.py b/providers/standard/tests/unit/standard/triggers/test_websocket.py new file mode 100644 index 0000000000000..4d5823cff1a76 --- /dev/null +++ b/providers/standard/tests/unit/standard/triggers/test_websocket.py @@ -0,0 +1,77 @@ +# 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 unittest import mock + +import pytest + +from airflow.providers.standard.triggers.websocket import WebSocketTrigger + + +class _FakeConnection: + """Stands in for the object returned by ``websockets.asyncio.client.connect``.""" + + def __init__(self, websocket): + self._websocket = websocket + + async def __aenter__(self): + return self._websocket + + async def __aexit__(self, *exc_info): + return False + + +class TestWebSocketTrigger: + URL = "ws://example.com/socket" + + def test_serialization(self): + """Asserts that the trigger correctly serializes its arguments and classpath.""" + trigger = WebSocketTrigger(url=self.URL, header={"Authorization": "token"}, message_to_send="ping") + classpath, kwargs = trigger.serialize() + assert classpath == "airflow.providers.standard.triggers.websocket.WebSocketTrigger" + assert kwargs == { + "url": self.URL, + "header": {"Authorization": "token"}, + "message_to_send": "ping", + } + + @pytest.mark.asyncio + @mock.patch("airflow.providers.standard.triggers.websocket.connect") + async def test_run_yields_event_with_received_message(self, mock_connect): + mock_websocket = mock.AsyncMock() + mock_websocket.recv.return_value = "pong" + mock_connect.return_value = _FakeConnection(mock_websocket) + + trigger = WebSocketTrigger(url=self.URL, header={"Authorization": "token"}, message_to_send="ping") + event = await trigger.run().__anext__() + + mock_connect.assert_called_once_with(self.URL, additional_headers={"Authorization": "token"}) + mock_websocket.send.assert_awaited_once_with("ping") + assert event.payload == "pong" + + @pytest.mark.asyncio + @mock.patch("airflow.providers.standard.triggers.websocket.connect") + async def test_run_does_not_send_without_message_to_send(self, mock_connect): + mock_websocket = mock.AsyncMock() + mock_websocket.recv.return_value = "pong" + mock_connect.return_value = _FakeConnection(mock_websocket) + + trigger = WebSocketTrigger(url=self.URL) + await trigger.run().__anext__() + + mock_websocket.send.assert_not_awaited() diff --git a/uv.lock b/uv.lock index d5a35a4b581eb..09a2b237cba00 100644 --- a/uv.lock +++ b/uv.lock @@ -3205,7 +3205,7 @@ docs = [{ name = "apache-airflow-devel-common", extras = ["docs"], editable = "d [[package]] name = "apache-airflow-providers-amazon" -version = "9.35.0" +version = "9.35.1" source = { editable = "providers/amazon" } dependencies = [ { name = "apache-airflow" }, @@ -6587,7 +6587,7 @@ docs = [{ name = "apache-airflow-devel-common", extras = ["docs"], editable = "d [[package]] name = "apache-airflow-providers-microsoft-azure" -version = "15.0.0" +version = "15.0.1" source = { editable = "providers/microsoft/azure" } dependencies = [ { name = "adlfs" }, @@ -8272,6 +8272,9 @@ dependencies = [ openlineage = [ { name = "apache-airflow-providers-openlineage" }, ] +websocket = [ + { name = "websockets" }, +] [package.dev-dependencies] dev = [ @@ -8281,6 +8284,7 @@ dev = [ { name = "apache-airflow-providers-mysql" }, { name = "apache-airflow-providers-openlineage" }, { name = "apache-airflow-task-sdk" }, + { name = "websockets" }, ] docs = [ { name = "apache-airflow-devel-common", extra = ["docs"] }, @@ -8291,8 +8295,9 @@ requires-dist = [ { name = "apache-airflow", editable = "." }, { name = "apache-airflow-providers-common-compat", editable = "providers/common/compat" }, { name = "apache-airflow-providers-openlineage", marker = "extra == 'openlineage'", editable = "providers/openlineage" }, + { name = "websockets", marker = "extra == 'websocket'", specifier = ">=14.0" }, ] -provides-extras = ["openlineage"] +provides-extras = ["openlineage", "websocket"] [package.metadata.requires-dev] dev = [ @@ -8302,6 +8307,7 @@ dev = [ { name = "apache-airflow-providers-mysql", editable = "providers/mysql" }, { name = "apache-airflow-providers-openlineage", editable = "providers/openlineage" }, { name = "apache-airflow-task-sdk", editable = "task-sdk" }, + { name = "websockets", specifier = ">=14.0" }, ] docs = [{ name = "apache-airflow-devel-common", extras = ["docs"], editable = "devel-common" }] From 2c17a304cb718154784e02ceedd8e50a99eb5ac4 Mon Sep 17 00:00:00 2001 From: ColtenOuO Date: Thu, 27 Aug 2026 07:43:40 +0000 Subject: [PATCH 2/6] Fix WebSocketSensor connection reuse and base class per review The deferrable path polled once synchronously before deferring, which consumes a WebSocket message and can resend message_to_send before the trigger opens its own connection; the non-deferrable path fell through to a second poke() after the sensor loop already succeeded, which could even call self.defer() while deferrable=False. WebSocket reads are consumptive, so each path must now touch the connection exactly once. WebSocketTrigger also switches from BaseEventTrigger (the event-driven-scheduling marker) to BaseTrigger, matching FileTrigger, since this trigger resumes a deferred task rather than driving event-based scheduling. Also drops the manual edit to the generated get_provider_info.py (provider.yaml is what local/dev discovery reads; the release process regenerates the rest), adds spec/autospec to the WebSocket mocks in both test files, and adds a sample demonstrating message_to_send with header. --- providers/standard/docs/sensors/websocket.rst | 10 +++++++ .../standard/example_dags/example_sensors.py | 14 +++++++++- .../providers/standard/get_provider_info.py | 3 --- .../providers/standard/sensors/websocket.py | 24 ++++++++++------- .../providers/standard/triggers/websocket.py | 9 ++----- .../unit/standard/sensors/test_websocket.py | 27 ++++++++++++++----- .../unit/standard/triggers/test_websocket.py | 9 ++++--- uv.lock | 4 +-- 8 files changed, 67 insertions(+), 33 deletions(-) diff --git a/providers/standard/docs/sensors/websocket.rst b/providers/standard/docs/sensors/websocket.rst index 18ab2c16fa91b..1f9e104c3a070 100644 --- a/providers/standard/docs/sensors/websocket.rst +++ b/providers/standard/docs/sensors/websocket.rst @@ -41,3 +41,13 @@ Also for this job you can use sensor in the deferrable mode: :dedent: 4 :start-after: [START example_websocket_sensor_async] :end-before: [END example_websocket_sensor_async] + +A common use case is to send a request over the connection right after it opens (via +``message_to_send``) and then wait for the remote server's asynchronous reply, optionally +passing connection headers such as an auth token via ``header``: + +.. exampleinclude:: /../src/airflow/providers/standard/example_dags/example_sensors.py + :language: python + :dedent: 4 + :start-after: [START example_websocket_sensor_send_message_async] + :end-before: [END example_websocket_sensor_send_message_async] diff --git a/providers/standard/src/airflow/providers/standard/example_dags/example_sensors.py b/providers/standard/src/airflow/providers/standard/example_dags/example_sensors.py index 7e9363e7ce53f..a1fcdc64538c6 100644 --- a/providers/standard/src/airflow/providers/standard/example_dags/example_sensors.py +++ b/providers/standard/src/airflow/providers/standard/example_dags/example_sensors.py @@ -142,6 +142,18 @@ def failure_callable(): ) # [END example_websocket_sensor_async] + # [START example_websocket_sensor_send_message_async] + t14 = WebSocketSensor( + task_id="request_and_wait_for_websocket_reply", + url="wss://example.com/socket", + message_to_send='{"action": "start_job"}', + header={"Authorization": "Bearer my-token"}, + deferrable=True, + timeout=3, + soft_fail=True, + ) + # [END example_websocket_sensor_send_message_async] + tx = BashOperator(task_id="print_date_in_bash", bash_command="date") tx.trigger_rule = TriggerRule.NONE_FAILED @@ -150,4 +162,4 @@ def failure_callable(): t8 >> tx [t9, t10] >> tx t11 >> tx - [t12, t13] >> tx + [t12, t13, t14] >> tx diff --git a/providers/standard/src/airflow/providers/standard/get_provider_info.py b/providers/standard/src/airflow/providers/standard/get_provider_info.py index 0be729079a0cb..1f7b2049454d1 100644 --- a/providers/standard/src/airflow/providers/standard/get_provider_info.py +++ b/providers/standard/src/airflow/providers/standard/get_provider_info.py @@ -43,7 +43,6 @@ def get_provider_info(): "/docs/apache-airflow-providers-standard/sensors/datetime.rst", "/docs/apache-airflow-providers-standard/sensors/file.rst", "/docs/apache-airflow-providers-standard/sensors/external_task_sensor.rst", - "/docs/apache-airflow-providers-standard/sensors/websocket.rst", ], } ], @@ -76,7 +75,6 @@ def get_provider_info(): "airflow.providers.standard.sensors.python", "airflow.providers.standard.sensors.filesystem", "airflow.providers.standard.sensors.external_task", - "airflow.providers.standard.sensors.websocket", ], } ], @@ -98,7 +96,6 @@ def get_provider_info(): "airflow.providers.standard.triggers.file", "airflow.providers.standard.triggers.temporal", "airflow.providers.standard.triggers.hitl", - "airflow.providers.standard.triggers.websocket", ], } ], diff --git a/providers/standard/src/airflow/providers/standard/sensors/websocket.py b/providers/standard/src/airflow/providers/standard/sensors/websocket.py index 57c75ad409109..aa903a90c9aed 100644 --- a/providers/standard/src/airflow/providers/standard/sensors/websocket.py +++ b/providers/standard/src/airflow/providers/standard/sensors/websocket.py @@ -76,16 +76,20 @@ def poke(self, context: Context) -> bool: def execute(self, context: Context) -> None: if not self.deferrable: super().execute(context=context) - if not self.poke(context=context): - self.defer( - timeout=datetime.timedelta(seconds=self.timeout), - trigger=WebSocketTrigger( - url=self.url, - header=self.header, - message_to_send=self.message_to_send, - ), - method_name="execute_complete", - ) + return + # Each poke opens and consumes a WebSocket connection, so the deferrable path must + # defer immediately: polling here first would send message_to_send and consume the + # reply before handing off, leaving the trigger to open a second connection and + # re-send the request. + self.defer( + timeout=datetime.timedelta(seconds=self.timeout), + trigger=WebSocketTrigger( + url=self.url, + header=self.header, + message_to_send=self.message_to_send, + ), + method_name="execute_complete", + ) def execute_complete(self, context: Context, event: Any = None) -> None: """Handle the event when the trigger fires and return immediately.""" diff --git a/providers/standard/src/airflow/providers/standard/triggers/websocket.py b/providers/standard/src/airflow/providers/standard/triggers/websocket.py index d3324469a4beb..b288014ee7a8b 100644 --- a/providers/standard/src/airflow/providers/standard/triggers/websocket.py +++ b/providers/standard/src/airflow/providers/standard/triggers/websocket.py @@ -21,15 +21,10 @@ from websockets.asyncio.client import connect -from airflow.providers.standard.version_compat import AIRFLOW_V_3_0_PLUS +from airflow.triggers.base import BaseTrigger, TriggerEvent -if AIRFLOW_V_3_0_PLUS: - from airflow.triggers.base import BaseEventTrigger, TriggerEvent -else: - from airflow.triggers.base import BaseTrigger as BaseEventTrigger, TriggerEvent # type: ignore - -class WebSocketTrigger(BaseEventTrigger): +class WebSocketTrigger(BaseTrigger): """ A trigger that opens a WebSocket connection and fires once a message is received. diff --git a/providers/standard/tests/unit/standard/sensors/test_websocket.py b/providers/standard/tests/unit/standard/sensors/test_websocket.py index 3e8848b70520d..443cfb92bdf43 100644 --- a/providers/standard/tests/unit/standard/sensors/test_websocket.py +++ b/providers/standard/tests/unit/standard/sensors/test_websocket.py @@ -19,6 +19,7 @@ from unittest import mock import pytest +from websockets.sync.client import ClientConnection from airflow.models.dag import DAG from airflow.providers.common.compat.sdk import TaskDeferred @@ -37,9 +38,9 @@ def setup_class(cls): args = {"owner": "airflow", "start_date": DEFAULT_DATE} cls.dag = DAG("test_websocket_sensor", schedule=None, default_args=args) - @mock.patch("airflow.providers.standard.sensors.websocket.connect") + @mock.patch("airflow.providers.standard.sensors.websocket.connect", autospec=True) def test_poke_returns_true_on_message(self, mock_connect): - mock_websocket = mock.MagicMock() + mock_websocket = mock.MagicMock(spec=ClientConnection) mock_websocket.recv.return_value = "pong" mock_connect.return_value.__enter__.return_value = mock_websocket @@ -47,21 +48,35 @@ def test_poke_returns_true_on_message(self, mock_connect): assert sensor.poke(context={}) is True mock_websocket.send.assert_called_once_with("ping") - @mock.patch("airflow.providers.standard.sensors.websocket.connect") + @mock.patch("airflow.providers.standard.sensors.websocket.connect", autospec=True) def test_poke_returns_false_on_timeout(self, mock_connect): - mock_websocket = mock.MagicMock() + mock_websocket = mock.MagicMock(spec=ClientConnection) mock_websocket.recv.side_effect = TimeoutError() mock_connect.return_value.__enter__.return_value = mock_websocket sensor = WebSocketSensor(task_id="poke_false", url=URL, dag=self.dag) assert sensor.poke(context={}) is False - def test_task_defer(self): + def test_task_defer_does_not_poke_first(self): + """The deferrable path must defer immediately: poke() consumes the connection, + so polling before deferring would send message_to_send and lose the reply the + trigger is supposed to wait for.""" sensor = WebSocketSensor(task_id="defer", url=URL, deferrable=True, dag=self.dag) - with mock.patch.object(WebSocketSensor, "poke", return_value=False): + with mock.patch.object(WebSocketSensor, "poke") as mock_poke: with pytest.raises(TaskDeferred) as exc: sensor.execute({}) + mock_poke.assert_not_called() assert isinstance(exc.value.trigger, WebSocketTrigger) assert exc.value.trigger.url == URL + + def test_execute_sync_returns_after_one_poke(self): + """The non-deferrable path must return once the sensor loop succeeds, not poke + (and consume) a second WebSocket message.""" + sensor = WebSocketSensor(task_id="sync", url=URL, timeout=0, dag=self.dag) + + with mock.patch.object(WebSocketSensor, "poke", return_value=True) as mock_poke: + sensor.execute({}) + + mock_poke.assert_called_once() diff --git a/providers/standard/tests/unit/standard/triggers/test_websocket.py b/providers/standard/tests/unit/standard/triggers/test_websocket.py index 4d5823cff1a76..81cb683301053 100644 --- a/providers/standard/tests/unit/standard/triggers/test_websocket.py +++ b/providers/standard/tests/unit/standard/triggers/test_websocket.py @@ -19,6 +19,7 @@ from unittest import mock import pytest +from websockets.asyncio.client import ClientConnection from airflow.providers.standard.triggers.websocket import WebSocketTrigger @@ -51,9 +52,9 @@ def test_serialization(self): } @pytest.mark.asyncio - @mock.patch("airflow.providers.standard.triggers.websocket.connect") + @mock.patch("airflow.providers.standard.triggers.websocket.connect", autospec=True) async def test_run_yields_event_with_received_message(self, mock_connect): - mock_websocket = mock.AsyncMock() + mock_websocket = mock.AsyncMock(spec=ClientConnection) mock_websocket.recv.return_value = "pong" mock_connect.return_value = _FakeConnection(mock_websocket) @@ -65,9 +66,9 @@ async def test_run_yields_event_with_received_message(self, mock_connect): assert event.payload == "pong" @pytest.mark.asyncio - @mock.patch("airflow.providers.standard.triggers.websocket.connect") + @mock.patch("airflow.providers.standard.triggers.websocket.connect", autospec=True) async def test_run_does_not_send_without_message_to_send(self, mock_connect): - mock_websocket = mock.AsyncMock() + mock_websocket = mock.AsyncMock(spec=ClientConnection) mock_websocket.recv.return_value = "pong" mock_connect.return_value = _FakeConnection(mock_websocket) diff --git a/uv.lock b/uv.lock index 09a2b237cba00..f8ddf2bac1fe5 100644 --- a/uv.lock +++ b/uv.lock @@ -3205,7 +3205,7 @@ docs = [{ name = "apache-airflow-devel-common", extras = ["docs"], editable = "d [[package]] name = "apache-airflow-providers-amazon" -version = "9.35.1" +version = "9.35.0" source = { editable = "providers/amazon" } dependencies = [ { name = "apache-airflow" }, @@ -6587,7 +6587,7 @@ docs = [{ name = "apache-airflow-devel-common", extras = ["docs"], editable = "d [[package]] name = "apache-airflow-providers-microsoft-azure" -version = "15.0.1" +version = "15.0.0" source = { editable = "providers/microsoft/azure" } dependencies = [ { name = "adlfs" }, From 22dd11b41fc086c67739f86971ebdd6260788e5d Mon Sep 17 00:00:00 2001 From: ColtenOuO Date: Thu, 27 Aug 2026 09:09:47 +0000 Subject: [PATCH 3/6] Fix CI by restoring the amazon/azure uv.lock version bump CI's dependency-sync check regenerates uv.lock and expects it to match each provider's currently declared version; reverting the amazon/azure entries to their stale values (per an earlier review comment) broke that check, since pyproject.toml for both already declares the newer version. Restoring the bump lets uv.lock agree with pyproject.toml again. --- uv.lock | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/uv.lock b/uv.lock index f8ddf2bac1fe5..09a2b237cba00 100644 --- a/uv.lock +++ b/uv.lock @@ -3205,7 +3205,7 @@ docs = [{ name = "apache-airflow-devel-common", extras = ["docs"], editable = "d [[package]] name = "apache-airflow-providers-amazon" -version = "9.35.0" +version = "9.35.1" source = { editable = "providers/amazon" } dependencies = [ { name = "apache-airflow" }, @@ -6587,7 +6587,7 @@ docs = [{ name = "apache-airflow-devel-common", extras = ["docs"], editable = "d [[package]] name = "apache-airflow-providers-microsoft-azure" -version = "15.0.0" +version = "15.0.1" source = { editable = "providers/microsoft/azure" } dependencies = [ { name = "adlfs" }, From 22f1226da7bc4f3950ae8a7ee4acba41a8c71a79 Mon Sep 17 00:00:00 2001 From: ColtenOuO Date: Thu, 27 Aug 2026 09:22:29 +0000 Subject: [PATCH 4/6] Keep one WebSocket connection open across sensor pokes The prior fix only stopped execute() from polling a second time after super().execute() already succeeded; the poke-mode retry loop it drives still reconnected and re-sent message_to_send on every single poke, since poke() opened a fresh connection each call. WebSocket messages are consumptive, so repeatedly reconnecting can duplicate the remote request and drop replies sent while no connection was open. poke() now opens the connection once and reuses it across retries within the same task attempt, and the sensor is marked poke_mode_only since that per-attempt connection state would be lost under reschedule mode. --- .../providers/standard/sensors/websocket.py | 36 +++++++++++---- .../unit/standard/sensors/test_websocket.py | 46 +++++++++++++++++-- 2 files changed, 69 insertions(+), 13 deletions(-) diff --git a/providers/standard/src/airflow/providers/standard/sensors/websocket.py b/providers/standard/src/airflow/providers/standard/sensors/websocket.py index aa903a90c9aed..2f9f3236ddac5 100644 --- a/providers/standard/src/airflow/providers/standard/sensors/websocket.py +++ b/providers/standard/src/airflow/providers/standard/sensors/websocket.py @@ -20,19 +20,26 @@ from collections.abc import Sequence from typing import TYPE_CHECKING, Any -from websockets.sync.client import connect +from websockets.sync.client import ClientConnection, connect -from airflow.providers.common.compat.sdk import BaseSensorOperator, conf +from airflow.providers.common.compat.sdk import BaseSensorOperator, conf, poke_mode_only from airflow.providers.standard.triggers.websocket import WebSocketTrigger if TYPE_CHECKING: from airflow.sdk import Context +@poke_mode_only class WebSocketSensor(BaseSensorOperator): """ Waits for a message on a WebSocket connection. + WebSocket messages are consumptive: once read, a message cannot be read again, and + reconnecting can duplicate ``message_to_send`` against the remote server. The + non-deferrable path therefore keeps a single connection open across pokes instead of + reconnecting and re-sending on every poke, and this sensor will not behave correctly in + reschedule mode, since that state would be lost between rescheduled invocations. + :param url: The ``ws://`` or ``wss://`` URL of the WebSocket server to connect to. :param header: Optional headers sent when opening the connection. :param message_to_send: Optional message sent right after the connection is established. @@ -60,22 +67,31 @@ def __init__( self.header = header self.message_to_send = message_to_send self.deferrable = deferrable + self._connection: ClientConnection | None = None def poke(self, context: Context) -> bool: - self.log.info("Poking WebSocket %s", self.url) - with connect(self.url, additional_headers=self.header) as websocket: + if self._connection is None: + self.log.info("Connecting to WebSocket %s", self.url) + self._connection = connect(self.url, additional_headers=self.header) if self.message_to_send is not None: - websocket.send(self.message_to_send) - try: - websocket.recv(timeout=self.poke_interval) - except TimeoutError: - return False + self._connection.send(self.message_to_send) + try: + self._connection.recv(timeout=self.poke_interval) + except TimeoutError: + return False self.log.info("Received message from %s", self.url) + self._connection.close() + self._connection = None return True def execute(self, context: Context) -> None: if not self.deferrable: - super().execute(context=context) + try: + super().execute(context=context) + finally: + if self._connection is not None: + self._connection.close() + self._connection = None return # Each poke opens and consumes a WebSocket connection, so the deferrable path must # defer immediately: polling here first would send message_to_send and consume the diff --git a/providers/standard/tests/unit/standard/sensors/test_websocket.py b/providers/standard/tests/unit/standard/sensors/test_websocket.py index 443cfb92bdf43..3357596b815a3 100644 --- a/providers/standard/tests/unit/standard/sensors/test_websocket.py +++ b/providers/standard/tests/unit/standard/sensors/test_websocket.py @@ -22,7 +22,7 @@ from websockets.sync.client import ClientConnection from airflow.models.dag import DAG -from airflow.providers.common.compat.sdk import TaskDeferred +from airflow.providers.common.compat.sdk import AirflowSensorTimeout, TaskDeferred from airflow.providers.standard.sensors.websocket import WebSocketSensor from airflow.providers.standard.triggers.websocket import WebSocketTrigger @@ -42,20 +42,60 @@ def setup_class(cls): def test_poke_returns_true_on_message(self, mock_connect): mock_websocket = mock.MagicMock(spec=ClientConnection) mock_websocket.recv.return_value = "pong" - mock_connect.return_value.__enter__.return_value = mock_websocket + mock_connect.return_value = mock_websocket sensor = WebSocketSensor(task_id="poke_true", url=URL, message_to_send="ping", dag=self.dag) assert sensor.poke(context={}) is True mock_websocket.send.assert_called_once_with("ping") + mock_websocket.close.assert_called_once() + assert sensor._connection is None @mock.patch("airflow.providers.standard.sensors.websocket.connect", autospec=True) def test_poke_returns_false_on_timeout(self, mock_connect): mock_websocket = mock.MagicMock(spec=ClientConnection) mock_websocket.recv.side_effect = TimeoutError() - mock_connect.return_value.__enter__.return_value = mock_websocket + mock_connect.return_value = mock_websocket sensor = WebSocketSensor(task_id="poke_false", url=URL, dag=self.dag) assert sensor.poke(context={}) is False + mock_websocket.close.assert_not_called() + + @mock.patch("airflow.providers.standard.sensors.websocket.connect", autospec=True) + def test_poke_reuses_connection_and_sends_message_once(self, mock_connect): + """A second poke() after a timeout must not reopen the connection or re-send + message_to_send — WebSocket reads are consumptive, so resending would restart + the remote job the sensor is waiting on.""" + mock_websocket = mock.MagicMock(spec=ClientConnection) + mock_websocket.recv.side_effect = [TimeoutError(), "pong"] + mock_connect.return_value = mock_websocket + + sensor = WebSocketSensor(task_id="reuse", url=URL, message_to_send="ping", dag=self.dag) + assert sensor.poke(context={}) is False + assert sensor.poke(context={}) is True + + mock_connect.assert_called_once() + mock_websocket.send.assert_called_once_with("ping") + + def test_execute_closes_connection_left_open_on_timeout(self): + """If the sensor loop times out, execute() must close whatever connection poke() + left open rather than leaking it.""" + sensor = WebSocketSensor(task_id="timeout", url=URL, timeout=0, dag=self.dag) + fake_connection = mock.MagicMock(spec=ClientConnection) + + def fake_poke(context): + sensor._connection = fake_connection + return False + + with mock.patch.object(WebSocketSensor, "poke", side_effect=fake_poke): + with pytest.raises(AirflowSensorTimeout): + sensor.execute({}) + + fake_connection.close.assert_called_once() + assert sensor._connection is None + + def test_reschedule_mode_not_allowed(self): + with pytest.raises(ValueError, match="Cannot set mode to 'reschedule'. Only 'poke' is acceptable"): + WebSocketSensor(task_id="reschedule", url=URL, mode="reschedule", dag=self.dag) def test_task_defer_does_not_poke_first(self): """The deferrable path must defer immediately: poke() consumes the connection, From d4ea6bf39450654eb31219b4fc5d75db9461d6f6 Mon Sep 17 00:00:00 2001 From: ColtenOuO Date: Thu, 27 Aug 2026 11:53:41 +0000 Subject: [PATCH 5/6] Fix CI, sensor timeout handling, and document trigger idempotency Regenerate get_provider_info.py via the update-providers-build-files prek hook (the same one CI's static checks run) instead of leaving it stale, which is what failed the "CI image checks / Static checks" job. The prior fix kept a WebSocket connection open across pokes, but each poke still only waited poke_interval before giving up, so BaseSensorOperator's retry loop could still reconnect and re-send message_to_send once poke_interval elapsed, and a poke_interval longer than the sensor's timeout (e.g. the default 60s poke_interval against a 3s example timeout) could block past the declared timeout entirely. poke() now opens exactly one connection and waits up to the sensor's overall timeout, so execute() never needs a retry loop for this sensor. Also documents that, like any Airflow trigger, WebSocketTrigger.run() can execute more than once (triggerer restart or redistribution), so message_to_send may be resent; a test demonstrates this by reconstructing a trigger from its own serialize() output and running it twice. Adds header and message_to_send to template_fields, since real requests commonly need a run id or auth token resolved from the Dag context rather than hard-coded in the Dag file. --- providers/standard/docs/sensors/websocket.rst | 8 ++ .../providers/standard/get_provider_info.py | 3 + .../providers/standard/sensors/websocket.py | 37 ++++----- .../providers/standard/triggers/websocket.py | 6 ++ .../unit/standard/sensors/test_websocket.py | 75 +++++++++---------- .../unit/standard/triggers/test_websocket.py | 21 ++++++ 6 files changed, 89 insertions(+), 61 deletions(-) diff --git a/providers/standard/docs/sensors/websocket.rst b/providers/standard/docs/sensors/websocket.rst index 1f9e104c3a070..72000b387fdda 100644 --- a/providers/standard/docs/sensors/websocket.rst +++ b/providers/standard/docs/sensors/websocket.rst @@ -51,3 +51,11 @@ passing connection headers such as an auth token via ``header``: :dedent: 4 :start-after: [START example_websocket_sensor_send_message_async] :end-before: [END example_websocket_sensor_send_message_async] + +.. warning:: + In deferrable mode, ``message_to_send`` may be sent more than once for a single task + run. Airflow triggers are not guaranteed to execute exactly once — a triggerer + restart or redistribution to another host re-runs the trigger from scratch, opening + a new connection and re-sending ``message_to_send``. If that message has a side + effect on the remote server, such as starting a job, the server must treat a resend + as safe — for example by deduplicating on a request id embedded in the message. diff --git a/providers/standard/src/airflow/providers/standard/get_provider_info.py b/providers/standard/src/airflow/providers/standard/get_provider_info.py index 1f7b2049454d1..0be729079a0cb 100644 --- a/providers/standard/src/airflow/providers/standard/get_provider_info.py +++ b/providers/standard/src/airflow/providers/standard/get_provider_info.py @@ -43,6 +43,7 @@ def get_provider_info(): "/docs/apache-airflow-providers-standard/sensors/datetime.rst", "/docs/apache-airflow-providers-standard/sensors/file.rst", "/docs/apache-airflow-providers-standard/sensors/external_task_sensor.rst", + "/docs/apache-airflow-providers-standard/sensors/websocket.rst", ], } ], @@ -75,6 +76,7 @@ def get_provider_info(): "airflow.providers.standard.sensors.python", "airflow.providers.standard.sensors.filesystem", "airflow.providers.standard.sensors.external_task", + "airflow.providers.standard.sensors.websocket", ], } ], @@ -96,6 +98,7 @@ def get_provider_info(): "airflow.providers.standard.triggers.file", "airflow.providers.standard.triggers.temporal", "airflow.providers.standard.triggers.hitl", + "airflow.providers.standard.triggers.websocket", ], } ], diff --git a/providers/standard/src/airflow/providers/standard/sensors/websocket.py b/providers/standard/src/airflow/providers/standard/sensors/websocket.py index 2f9f3236ddac5..a5d64d0d9f84c 100644 --- a/providers/standard/src/airflow/providers/standard/sensors/websocket.py +++ b/providers/standard/src/airflow/providers/standard/sensors/websocket.py @@ -20,7 +20,7 @@ from collections.abc import Sequence from typing import TYPE_CHECKING, Any -from websockets.sync.client import ClientConnection, connect +from websockets.sync.client import connect from airflow.providers.common.compat.sdk import BaseSensorOperator, conf, poke_mode_only from airflow.providers.standard.triggers.websocket import WebSocketTrigger @@ -36,9 +36,10 @@ class WebSocketSensor(BaseSensorOperator): WebSocket messages are consumptive: once read, a message cannot be read again, and reconnecting can duplicate ``message_to_send`` against the remote server. The - non-deferrable path therefore keeps a single connection open across pokes instead of - reconnecting and re-sending on every poke, and this sensor will not behave correctly in - reschedule mode, since that state would be lost between rescheduled invocations. + non-deferrable path therefore opens exactly one connection and blocks on it for up to + ``timeout`` seconds instead of reconnecting every ``poke_interval`` (which has no + effect on this sensor), and this sensor is marked poke-mode-only since a rescheduled + invocation would need a new connection anyway. :param url: The ``ws://`` or ``wss://`` URL of the WebSocket server to connect to. :param header: Optional headers sent when opening the connection. @@ -51,7 +52,8 @@ class WebSocketSensor(BaseSensorOperator): :ref:`howto/operator:WebSocketSensor` """ - template_fields: Sequence[str] = ("url",) + template_fields: Sequence[str] = ("url", "header", "message_to_send") + template_fields_renderers = {"header": "json"} def __init__( self, @@ -67,31 +69,22 @@ def __init__( self.header = header self.message_to_send = message_to_send self.deferrable = deferrable - self._connection: ClientConnection | None = None def poke(self, context: Context) -> bool: - if self._connection is None: - self.log.info("Connecting to WebSocket %s", self.url) - self._connection = connect(self.url, additional_headers=self.header) + self.log.info("Connecting to WebSocket %s", self.url) + with connect(self.url, additional_headers=self.header) as websocket: if self.message_to_send is not None: - self._connection.send(self.message_to_send) - try: - self._connection.recv(timeout=self.poke_interval) - except TimeoutError: - return False + websocket.send(self.message_to_send) + try: + websocket.recv(timeout=self.timeout) + except TimeoutError: + return False self.log.info("Received message from %s", self.url) - self._connection.close() - self._connection = None return True def execute(self, context: Context) -> None: if not self.deferrable: - try: - super().execute(context=context) - finally: - if self._connection is not None: - self._connection.close() - self._connection = None + super().execute(context=context) return # Each poke opens and consumes a WebSocket connection, so the deferrable path must # defer immediately: polling here first would send message_to_send and consume the diff --git a/providers/standard/src/airflow/providers/standard/triggers/websocket.py b/providers/standard/src/airflow/providers/standard/triggers/websocket.py index b288014ee7a8b..c3436a5b8ac0c 100644 --- a/providers/standard/src/airflow/providers/standard/triggers/websocket.py +++ b/providers/standard/src/airflow/providers/standard/triggers/websocket.py @@ -32,6 +32,12 @@ class WebSocketTrigger(BaseTrigger): WebSocket server and resume once that server replies, without occupying a worker slot while waiting. + Like any Airflow trigger, ``run()`` is not guaranteed to execute only once: a + triggerer restart or redistribution to another host re-runs it from scratch. Each + execution opens a new connection and re-sends ``message_to_send`` if one is set, so + if that message starts a remote job, the remote server must treat a resend as safe — + for example by deduplicating on a request id embedded in the message. + :param url: The ``ws://`` or ``wss://`` URL of the WebSocket server to connect to. :param header: Optional headers sent when opening the connection. :param message_to_send: Optional message sent right after the connection is established. diff --git a/providers/standard/tests/unit/standard/sensors/test_websocket.py b/providers/standard/tests/unit/standard/sensors/test_websocket.py index 3357596b815a3..5dc2158e5ddde 100644 --- a/providers/standard/tests/unit/standard/sensors/test_websocket.py +++ b/providers/standard/tests/unit/standard/sensors/test_websocket.py @@ -42,56 +42,35 @@ def setup_class(cls): def test_poke_returns_true_on_message(self, mock_connect): mock_websocket = mock.MagicMock(spec=ClientConnection) mock_websocket.recv.return_value = "pong" - mock_connect.return_value = mock_websocket + mock_connect.return_value.__enter__.return_value = mock_websocket sensor = WebSocketSensor(task_id="poke_true", url=URL, message_to_send="ping", dag=self.dag) assert sensor.poke(context={}) is True mock_websocket.send.assert_called_once_with("ping") - mock_websocket.close.assert_called_once() - assert sensor._connection is None @mock.patch("airflow.providers.standard.sensors.websocket.connect", autospec=True) def test_poke_returns_false_on_timeout(self, mock_connect): mock_websocket = mock.MagicMock(spec=ClientConnection) mock_websocket.recv.side_effect = TimeoutError() - mock_connect.return_value = mock_websocket + mock_connect.return_value.__enter__.return_value = mock_websocket sensor = WebSocketSensor(task_id="poke_false", url=URL, dag=self.dag) assert sensor.poke(context={}) is False - mock_websocket.close.assert_not_called() @mock.patch("airflow.providers.standard.sensors.websocket.connect", autospec=True) - def test_poke_reuses_connection_and_sends_message_once(self, mock_connect): - """A second poke() after a timeout must not reopen the connection or re-send - message_to_send — WebSocket reads are consumptive, so resending would restart - the remote job the sensor is waiting on.""" + def test_poke_waits_for_the_overall_timeout_not_poke_interval(self, mock_connect): + """recv() must be bounded by the sensor's overall timeout, not poke_interval — + otherwise a single poke could block past the sensor's declared timeout before + that timeout is ever checked.""" mock_websocket = mock.MagicMock(spec=ClientConnection) - mock_websocket.recv.side_effect = [TimeoutError(), "pong"] - mock_connect.return_value = mock_websocket + mock_websocket.recv.return_value = "pong" + mock_connect.return_value.__enter__.return_value = mock_websocket - sensor = WebSocketSensor(task_id="reuse", url=URL, message_to_send="ping", dag=self.dag) - assert sensor.poke(context={}) is False + sensor = WebSocketSensor( + task_id="poke_timeout_arg", url=URL, timeout=45, poke_interval=5, dag=self.dag + ) assert sensor.poke(context={}) is True - - mock_connect.assert_called_once() - mock_websocket.send.assert_called_once_with("ping") - - def test_execute_closes_connection_left_open_on_timeout(self): - """If the sensor loop times out, execute() must close whatever connection poke() - left open rather than leaking it.""" - sensor = WebSocketSensor(task_id="timeout", url=URL, timeout=0, dag=self.dag) - fake_connection = mock.MagicMock(spec=ClientConnection) - - def fake_poke(context): - sensor._connection = fake_connection - return False - - with mock.patch.object(WebSocketSensor, "poke", side_effect=fake_poke): - with pytest.raises(AirflowSensorTimeout): - sensor.execute({}) - - fake_connection.close.assert_called_once() - assert sensor._connection is None + mock_websocket.recv.assert_called_once_with(timeout=45) def test_reschedule_mode_not_allowed(self): with pytest.raises(ValueError, match="Cannot set mode to 'reschedule'. Only 'poke' is acceptable"): @@ -111,12 +90,30 @@ def test_task_defer_does_not_poke_first(self): assert isinstance(exc.value.trigger, WebSocketTrigger) assert exc.value.trigger.url == URL - def test_execute_sync_returns_after_one_poke(self): - """The non-deferrable path must return once the sensor loop succeeds, not poke - (and consume) a second WebSocket message.""" - sensor = WebSocketSensor(task_id="sync", url=URL, timeout=0, dag=self.dag) + def test_execute_sync_calls_poke_exactly_once(self): + """Since poke() already blocks for the full sensor timeout, execute() must never + call it a second time — a second call would open a new connection and re-send + message_to_send.""" + sensor = WebSocketSensor(task_id="sync_timeout", url=URL, timeout=0, dag=self.dag) - with mock.patch.object(WebSocketSensor, "poke", return_value=True) as mock_poke: - sensor.execute({}) + with mock.patch.object(WebSocketSensor, "poke", return_value=False) as mock_poke: + with pytest.raises(AirflowSensorTimeout): + sensor.execute({}) mock_poke.assert_called_once() + + def test_template_fields_are_rendered(self): + """url, header, and message_to_send commonly need runtime values (run_id, an + idempotency key, an auth token), so all three must be templated.""" + sensor = WebSocketSensor( + task_id="templated", + url="wss://example.com/{{ run_id }}", + message_to_send='{"run_id": "{{ run_id }}"}', + header={"Authorization": "Bearer {{ run_id }}"}, + dag=self.dag, + ) + sensor.render_template_fields({"run_id": "manual__2024-01-01"}) + + assert sensor.url == "wss://example.com/manual__2024-01-01" + assert sensor.message_to_send == '{"run_id": "manual__2024-01-01"}' + assert sensor.header == {"Authorization": "Bearer manual__2024-01-01"} diff --git a/providers/standard/tests/unit/standard/triggers/test_websocket.py b/providers/standard/tests/unit/standard/triggers/test_websocket.py index 81cb683301053..15cdb767cfbe7 100644 --- a/providers/standard/tests/unit/standard/triggers/test_websocket.py +++ b/providers/standard/tests/unit/standard/triggers/test_websocket.py @@ -76,3 +76,24 @@ async def test_run_does_not_send_without_message_to_send(self, mock_connect): await trigger.run().__anext__() mock_websocket.send.assert_not_awaited() + + @pytest.mark.asyncio + @mock.patch("airflow.providers.standard.triggers.websocket.connect", autospec=True) + async def test_reconstructed_trigger_resends_message_to_send(self, mock_connect): + """Documents that a triggerer restart or redistribution — which reconstructs the + trigger from its serialize() output and calls run() again — re-sends + message_to_send. Airflow does not guarantee a trigger's run() executes only + once, so callers whose message starts a remote job must make that job + idempotent; this is not something the trigger itself can enforce.""" + mock_websocket = mock.AsyncMock(spec=ClientConnection) + mock_websocket.recv.return_value = "pong" + mock_connect.return_value = _FakeConnection(mock_websocket) + + original = WebSocketTrigger(url=self.URL, message_to_send="start_job") + _, kwargs = original.serialize() + reconstructed = WebSocketTrigger(**kwargs) + + await original.run().__anext__() + await reconstructed.run().__anext__() + + assert mock_websocket.send.await_args_list == [mock.call("start_job"), mock.call("start_job")] From 8905318c7ea84823ee6a5a951d0221a843f6f1e1 Mon Sep 17 00:00:00 2001 From: ColtenOuO Date: Thu, 27 Aug 2026 14:31:08 +0000 Subject: [PATCH 6/6] Bound the WebSocket handshake too, and fix CI spellcheck poke() passed no open_timeout to connect(), so a slow handshake could add websockets' 10s default on top of the recv() wait, letting a single poke run well past the sensor's declared timeout. A handshake timeout also raised outside the existing try block, bypassing soft_fail entirely instead of being treated like any other sensor timeout. Both stages now share one deadline: connect() gets the remaining time as open_timeout, and recv() gets whatever is left after the handshake, with both timeouts handled the same way. Also adds autospec=True to the two remaining unspecced poke() mocks per the testing guidelines, and adds "deduplicating" to the docs spelling wordlist, fixing the CI static-checks docs spellcheck failure (the existing entries only covered deduplicate/deduplicated/deduplication). --- docs/spelling_wordlist.txt | 1 + .../providers/standard/sensors/websocket.py | 16 +++++---- .../unit/standard/sensors/test_websocket.py | 35 +++++++++++++++++-- 3 files changed, 42 insertions(+), 10 deletions(-) diff --git a/docs/spelling_wordlist.txt b/docs/spelling_wordlist.txt index b0cd8ce92ea2c..c91c0cbc45b7b 100644 --- a/docs/spelling_wordlist.txt +++ b/docs/spelling_wordlist.txt @@ -454,6 +454,7 @@ decrypted Decrypts deduplicate deduplicated +deduplicating deduplication deepcopy DefaultAzureCredential diff --git a/providers/standard/src/airflow/providers/standard/sensors/websocket.py b/providers/standard/src/airflow/providers/standard/sensors/websocket.py index a5d64d0d9f84c..bf87bffd36af9 100644 --- a/providers/standard/src/airflow/providers/standard/sensors/websocket.py +++ b/providers/standard/src/airflow/providers/standard/sensors/websocket.py @@ -17,6 +17,7 @@ from __future__ import annotations import datetime +import time from collections.abc import Sequence from typing import TYPE_CHECKING, Any @@ -72,13 +73,14 @@ def __init__( def poke(self, context: Context) -> bool: self.log.info("Connecting to WebSocket %s", self.url) - with connect(self.url, additional_headers=self.header) as websocket: - if self.message_to_send is not None: - websocket.send(self.message_to_send) - try: - websocket.recv(timeout=self.timeout) - except TimeoutError: - return False + deadline = time.monotonic() + self.timeout + try: + with connect(self.url, additional_headers=self.header, open_timeout=self.timeout) as websocket: + if self.message_to_send is not None: + websocket.send(self.message_to_send) + websocket.recv(timeout=max(deadline - time.monotonic(), 0)) + except TimeoutError: + return False self.log.info("Received message from %s", self.url) return True diff --git a/providers/standard/tests/unit/standard/sensors/test_websocket.py b/providers/standard/tests/unit/standard/sensors/test_websocket.py index 5dc2158e5ddde..45ef4e8435036 100644 --- a/providers/standard/tests/unit/standard/sensors/test_websocket.py +++ b/providers/standard/tests/unit/standard/sensors/test_websocket.py @@ -70,7 +70,36 @@ def test_poke_waits_for_the_overall_timeout_not_poke_interval(self, mock_connect task_id="poke_timeout_arg", url=URL, timeout=45, poke_interval=5, dag=self.dag ) assert sensor.poke(context={}) is True - mock_websocket.recv.assert_called_once_with(timeout=45) + + mock_connect.assert_called_once_with(URL, additional_headers=None, open_timeout=45) + (_, kwargs) = mock_websocket.recv.call_args + assert kwargs["timeout"] == pytest.approx(45, abs=1) + + @mock.patch("airflow.providers.standard.sensors.websocket.connect", autospec=True) + def test_poke_returns_false_when_handshake_times_out(self, mock_connect): + """A connect() timeout (slow handshake) must be treated the same as a recv() + timeout — including respecting soft_fail — not propagate as a raw TimeoutError + that bypasses the sensor's normal timeout handling.""" + mock_connect.side_effect = TimeoutError("timed out while waiting for handshake response") + + sensor = WebSocketSensor(task_id="handshake_timeout", url=URL, dag=self.dag) + assert sensor.poke(context={}) is False + + @mock.patch("airflow.providers.standard.sensors.websocket.time.monotonic") + @mock.patch("airflow.providers.standard.sensors.websocket.connect", autospec=True) + def test_poke_recv_gets_remaining_time_after_slow_handshake(self, mock_connect, mock_monotonic): + """If the handshake itself consumes part of the timeout budget, recv() must only + get what's left, not the full timeout again — otherwise total wait time could + exceed the sensor's declared timeout.""" + mock_websocket = mock.MagicMock(spec=ClientConnection) + mock_websocket.recv.return_value = "pong" + mock_connect.return_value.__enter__.return_value = mock_websocket + # deadline computed at t=0 with timeout=10; connect() "takes" 4s, leaving 6s for recv(). + mock_monotonic.side_effect = [0, 4] + + sensor = WebSocketSensor(task_id="slow_handshake", url=URL, timeout=10, dag=self.dag) + assert sensor.poke(context={}) is True + mock_websocket.recv.assert_called_once_with(timeout=6) def test_reschedule_mode_not_allowed(self): with pytest.raises(ValueError, match="Cannot set mode to 'reschedule'. Only 'poke' is acceptable"): @@ -82,7 +111,7 @@ def test_task_defer_does_not_poke_first(self): trigger is supposed to wait for.""" sensor = WebSocketSensor(task_id="defer", url=URL, deferrable=True, dag=self.dag) - with mock.patch.object(WebSocketSensor, "poke") as mock_poke: + with mock.patch.object(WebSocketSensor, "poke", autospec=True) as mock_poke: with pytest.raises(TaskDeferred) as exc: sensor.execute({}) @@ -96,7 +125,7 @@ def test_execute_sync_calls_poke_exactly_once(self): message_to_send.""" sensor = WebSocketSensor(task_id="sync_timeout", url=URL, timeout=0, dag=self.dag) - with mock.patch.object(WebSocketSensor, "poke", return_value=False) as mock_poke: + with mock.patch.object(WebSocketSensor, "poke", autospec=True, return_value=False) as mock_poke: with pytest.raises(AirflowSensorTimeout): sensor.execute({})