From 242e6c5b4fcf4cb72925e12dc906c739093472cb Mon Sep 17 00:00:00 2001 From: yuseok89 Date: Thu, 8 Oct 2026 09:08:44 +0900 Subject: [PATCH] Make HttpSensor defer when response_check is provided --- .../airflow/providers/http/sensors/http.py | 80 +++++++++---- .../airflow/providers/http/triggers/http.py | 26 +++- .../http/tests/unit/http/sensors/test_http.py | 113 ++++++++++++++++-- .../tests/unit/http/triggers/test_http.py | 51 ++++++++ 4 files changed, 238 insertions(+), 32 deletions(-) diff --git a/providers/http/src/airflow/providers/http/sensors/http.py b/providers/http/src/airflow/providers/http/sensors/http.py index 2d8c4ac85c62e..80f1a5f82182b 100644 --- a/providers/http/src/airflow/providers/http/sensors/http.py +++ b/providers/http/src/airflow/providers/http/sensors/http.py @@ -21,12 +21,18 @@ from datetime import timedelta from typing import TYPE_CHECKING, Any -from airflow.providers.common.compat.sdk import AirflowException, BaseSensorOperator, conf +from airflow.providers.common.compat.sdk import ( + AirflowException, + BaseSensorOperator, + PokeReturnValue, + conf, + timezone, +) from airflow.providers.http.hooks.http import HttpHook -from airflow.providers.http.triggers.http import HttpSensorTrigger +from airflow.providers.http.triggers.http import HttpResponseSerializer, HttpSensorTrigger if TYPE_CHECKING: - from airflow.sdk import Context, PokeReturnValue + from airflow.sdk import Context class HttpSensor(BaseSensorOperator): @@ -72,7 +78,9 @@ def response_check(response, task_instance): :param response_check: A check against the 'requests' response object. The callable takes the response object as the first positional argument and optionally any number of keyword arguments available in the context dictionary. - It should return True for 'pass' and False otherwise. + It should return True for 'pass' and False otherwise. In deferrable mode, the + triggerer waits for the endpoint to respond without error and the check runs on + the worker; if it returns False, the task defers again until the check passes. :param extra_options: Extra options for the 'requests' library, see the 'requests' documentation (options to modify timeout, ssl, etc.) :param tcp_keep_alive: Enable TCP Keep Alive for the connection. @@ -158,22 +166,54 @@ def poke(self, context: Context) -> bool | PokeReturnValue: return True def execute(self, context: Context) -> Any: - if not self.deferrable or self.response_check: + if not self.deferrable: return super().execute(context=context) - if not self.poke(context): - self.defer( - timeout=timedelta(seconds=self.timeout), - trigger=HttpSensorTrigger( - endpoint=self.endpoint, - http_conn_id=self.http_conn_id, - data=self.request_params, - headers=self.headers, - method=self.method, - extra_options=self.extra_options, - poke_interval=self.poke_interval, - ), - method_name="execute_complete", - ) + result = self.poke(context) + + if not result: + self._defer(context=context, initial_delay=self.poke_interval if self.response_check else 0) + # Keep sync mode's contract of returning the xcom value from a truthy PokeReturnValue. + if isinstance(result, PokeReturnValue): + return result.xcom_value + + def _defer(self, context: Context, initial_delay: float = 0.0) -> None: + remaining_timeout = max( + self.timeout - (timezone.utcnow() - context["ti"].start_date).total_seconds(), 0 + ) + self.defer( + timeout=timedelta(seconds=remaining_timeout), + trigger=HttpSensorTrigger( + endpoint=self.endpoint, + http_conn_id=self.http_conn_id, + data=self.request_params, + headers=self.headers, + method=self.method, + extra_options=self.extra_options, + poke_interval=self.poke_interval, + initial_delay=initial_delay, + should_return_response=bool(self.response_check), + ), + method_name="execute_complete", + ) - def execute_complete(self, context: Context, event: dict[str, Any] | None = None) -> None: + def execute_complete(self, context: Context, event: dict[str, Any] | None = None) -> Any: + if self.response_check: + from airflow.utils.operator_helpers import determine_kwargs + + if not isinstance(event, dict) or "response" not in event: + raise ValueError( + "The trigger event does not contain the HTTP response required to " + "evaluate response_check. The deferred task was most likely resumed by a " + "trigger serialized with an older version of the http provider." + ) + response = HttpResponseSerializer.deserialize(event["response"]) + kwargs = determine_kwargs(self.response_check, [response], context) + result = self.response_check(response, **kwargs) + + if not result: + # The check did not pass yet; hand polling back to the triggerer. + self._defer(context=context, initial_delay=self.poke_interval) + if isinstance(result, PokeReturnValue): + self.log.info("%s completed successfully.", self.task_id) + return result.xcom_value self.log.info("%s completed successfully.", self.task_id) diff --git a/providers/http/src/airflow/providers/http/triggers/http.py b/providers/http/src/airflow/providers/http/triggers/http.py index d91559bd689c9..f6e5d5b2a28e8 100644 --- a/providers/http/src/airflow/providers/http/triggers/http.py +++ b/providers/http/src/airflow/providers/http/triggers/http.py @@ -220,6 +220,10 @@ class HttpSensorTrigger(BaseTrigger): :param extra_options: Additional kwargs to pass when creating a request. For example, ``run(json=obj)`` is passed as ``aiohttp.ClientSession().get(json=obj)`` :param poke_interval: Time to sleep using asyncio + :param initial_delay: Time to sleep before the first request. Used when the + sensor re-defers after evaluating ``response_check`` on the worker, so + consecutive attempts keep the ``poke_interval`` pacing. + :param should_return_response: Whether to include the serialized HTTP response in the trigger event. """ def __init__( @@ -231,6 +235,8 @@ def __init__( headers: dict[str, str] | None = None, extra_options: dict[str, Any] | None = None, poke_interval: float = 5.0, + initial_delay: float = 0.0, + should_return_response: bool = False, ): super().__init__() self.endpoint = endpoint @@ -240,6 +246,8 @@ def __init__( self.extra_options = extra_options or {} self.http_conn_id = http_conn_id self.poke_interval = poke_interval + self.initial_delay = initial_delay + self.should_return_response = should_return_response def serialize(self) -> tuple[str, dict[str, Any]]: """Serialize HttpTrigger arguments and classpath.""" @@ -253,23 +261,37 @@ def serialize(self) -> tuple[str, dict[str, Any]]: "extra_options": self.extra_options, "http_conn_id": self.http_conn_id, "poke_interval": self.poke_interval, + "initial_delay": self.initial_delay, + "should_return_response": self.should_return_response, }, ) async def run(self) -> AsyncIterator[TriggerEvent]: """Make a series of asynchronous http calls via an http hook.""" hook = self._get_async_hook() + if self.initial_delay > 0: + await asyncio.sleep(self.initial_delay) while True: try: async with aiohttp.ClientSession() as session: - await hook.run( + client_response = await hook.run( session=session, endpoint=self.endpoint, data=self.data, headers=self.headers, extra_options=self.extra_options, ) - yield TriggerEvent(True) + if self.should_return_response: + response = await HttpTrigger._convert_response(client_response) + if self.should_return_response: + yield TriggerEvent( + { + "status": "success", + "response": HttpResponseSerializer.serialize(response), + } + ) + else: + yield TriggerEvent(True) return except AirflowException as exc: if str(exc).startswith("404"): diff --git a/providers/http/tests/unit/http/sensors/test_http.py b/providers/http/tests/unit/http/sensors/test_http.py index d048adc88020a..00cd02db6eebb 100644 --- a/providers/http/tests/unit/http/sensors/test_http.py +++ b/providers/http/tests/unit/http/sensors/test_http.py @@ -17,6 +17,7 @@ # under the License. from __future__ import annotations +from types import SimpleNamespace from unittest import mock from unittest.mock import patch @@ -378,32 +379,124 @@ def test_execute_finished_before_deferred( "airflow.providers.http.sensors.http.HttpSensor.poke", return_value=False, ) - def test_execute_is_deferred(self, mock_poke): + def test_execute_is_deferred(self, mock_poke, time_machine): """ Asserts that a task is deferred and a HttpTrigger will be fired when the HttpSensor is executed in deferrable mode. """ + time_machine.move_to(DEFAULT_DATE, tick=False) task = HttpSensor(task_id="run_now", endpoint="test-endpoint", deferrable=True) + context = {"ti": SimpleNamespace(start_date=DEFAULT_DATE)} with pytest.raises(TaskDeferred) as exc: - task.execute({}) + task.execute(context) assert isinstance(exc.value.trigger, HttpSensorTrigger), "Trigger is not a HttpTrigger" + assert exc.value.trigger.should_return_response is False - @mock.patch("airflow.providers.http.sensors.http.HttpSensor.defer") @mock.patch( - "airflow.sdk.bases.sensor.BaseSensorOperator.execute" - if AIRFLOW_V_3_0_PLUS - else "airflow.sensors.base.BaseSensorOperator.execute" + "airflow.providers.http.sensors.http.HttpSensor.poke", + return_value=False, ) - def test_execute_not_defer_when_response_check_is_not_none(self, mock_execute, mock_defer): + def test_execute_defers_when_response_check_is_not_none(self, mock_poke, time_machine): + """A response_check must not force the sensor back onto the synchronous path.""" + time_machine.move_to(DEFAULT_DATE, tick=False) task = HttpSensor( task_id="run_now", endpoint="test-endpoint", response_check=lambda response: "httpbin" in response.text, + poke_interval=42, deferrable=True, ) - task.execute({}) - mock_execute.assert_called_once() - mock_defer.assert_not_called() + context = {"ti": SimpleNamespace(start_date=DEFAULT_DATE)} + with pytest.raises(TaskDeferred) as exc: + task.execute(context) + assert isinstance(exc.value.trigger, HttpSensorTrigger) + assert exc.value.trigger.initial_delay == 42 + assert exc.value.trigger.should_return_response is True + assert exc.value.timeout.total_seconds() == task.timeout + + @mock.patch( + "airflow.providers.http.sensors.http.HttpSensor.poke", + return_value=PokeReturnValue(is_done=True, xcom_value="payload"), + ) + def test_execute_returns_xcom_when_first_poke_succeeds(self, mock_poke): + """Immediate success must keep the sync-mode contract of returning the xcom value.""" + task = HttpSensor( + task_id="run_now", + endpoint="test-endpoint", + response_check=lambda response: PokeReturnValue(is_done=True, xcom_value=response.text), + deferrable=True, + ) + assert task.execute({}) == "payload" + + @staticmethod + def _make_event(text: str = "httpbin rocks") -> dict: + from airflow.providers.http.triggers.http import HttpResponseSerializer + + response = requests.Response() + response.status_code = 200 + response._content = text.encode() + response.url = "http://test-endpoint" + return {"status": "success", "response": HttpResponseSerializer.serialize(response)} + + def test_execute_complete_response_check_fails_defers_again(self, time_machine): + time_machine.move_to(DEFAULT_DATE, tick=False) + task = HttpSensor( + task_id="run_now", + endpoint="test-endpoint", + response_check=lambda response: "other" in response.text, + poke_interval=42, + timeout=300, + deferrable=True, + ) + context = {"ti": SimpleNamespace(start_date=DEFAULT_DATE)} + time_machine.shift(120) + with pytest.raises(TaskDeferred) as exc: + task.execute_complete(context=context, event=self._make_event()) + assert isinstance(exc.value.trigger, HttpSensorTrigger) + # Re-deferring must keep the poke_interval pacing instead of refiring immediately. + assert exc.value.trigger.initial_delay == 42 + assert exc.value.timeout.total_seconds() == 180 + + def test_execute_complete_response_check_receives_response(self): + seen = {} + + def response_check(response): + seen["text"] = response.text + seen["status_code"] = response.status_code + return True + + task = HttpSensor( + task_id="run_now", + endpoint="test-endpoint", + response_check=response_check, + deferrable=True, + ) + task.execute_complete(context={}, event=self._make_event("payload")) + assert seen == {"text": "payload", "status_code": 200} + + def test_execute_complete_response_check_poke_return_value(self): + task = HttpSensor( + task_id="run_now", + endpoint="test-endpoint", + response_check=lambda response: PokeReturnValue(is_done=True, xcom_value=response.text), + deferrable=True, + ) + assert task.execute_complete(context={}, event=self._make_event("payload")) == "payload" + + def test_execute_complete_legacy_event_with_response_check_raises(self): + """An in-flight trigger from an older provider version cannot satisfy response_check.""" + task = HttpSensor( + task_id="run_now", + endpoint="test-endpoint", + response_check=lambda response: True, + deferrable=True, + ) + with pytest.raises(ValueError, match="does not contain the HTTP response"): + task.execute_complete(context={}, event=True) # type: ignore[arg-type] + + def test_execute_complete_legacy_event_without_response_check(self): + task = HttpSensor(task_id="run_now", endpoint="test-endpoint", deferrable=True) + assert task.execute_complete(context={}, event=True) is None # type: ignore[arg-type] diff --git a/providers/http/tests/unit/http/triggers/test_http.py b/providers/http/tests/unit/http/triggers/test_http.py index a1542a41508b0..ed74bd783081d 100644 --- a/providers/http/tests/unit/http/triggers/test_http.py +++ b/providers/http/tests/unit/http/triggers/test_http.py @@ -194,6 +194,12 @@ async def test_trigger_on_post_with_data(self, mock_http_post, trigger): class TestHttpSensorTrigger: + @staticmethod + def _mock_run_result(result_to_mock): + f = Future() + f.set_result(result_to_mock) + return f + def test_serialization(self, sensor_trigger): """ Asserts that the HttpSensorTrigger correctly serializes its arguments @@ -209,8 +215,53 @@ def test_serialization(self, sensor_trigger): "data": TEST_DATA, "extra_options": TEST_EXTRA_OPTIONS, "poke_interval": 5.0, + "initial_delay": 0.0, + "should_return_response": False, } + @pytest.mark.asyncio + @mock.patch(HTTP_PATH.format("HttpAsyncHook")) + async def test_yields_serialized_response_on_success(self, mock_hook, sensor_trigger, client_response): + """The event carries the response so the sensor can evaluate response_check on the worker.""" + mock_hook.return_value.run.return_value = self._mock_run_result(client_response) + sensor_trigger.should_return_response = True + response = await HttpTrigger._convert_response(client_response) + + generator = sensor_trigger.run() + actual = await generator.asend(None) + assert actual == TriggerEvent( + {"status": "success", "response": HttpResponseSerializer.serialize(response)} + ) + + @pytest.mark.asyncio + @mock.patch(HTTP_PATH.format("HttpAsyncHook")) + @mock.patch.object(HttpTrigger, "_convert_response", autospec=True) + async def test_does_not_return_response_when_not_requested( + self, mock_convert_response, mock_hook, sensor_trigger, client_response + ): + mock_hook.return_value.run.return_value = self._mock_run_result(client_response) + + generator = sensor_trigger.run() + actual = await generator.asend(None) + assert actual == TriggerEvent(True) + mock_convert_response.assert_not_awaited() + + @pytest.mark.asyncio + @mock.patch(HTTP_PATH.format("asyncio.sleep")) + @mock.patch(HTTP_PATH.format("HttpAsyncHook")) + async def test_initial_delay_sleeps_before_first_request(self, mock_hook, mock_sleep, client_response): + mock_hook.return_value.run.return_value = self._mock_run_result(client_response) + trigger = HttpSensorTrigger( + http_conn_id=TEST_CONN_ID, + endpoint=TEST_ENDPOINT, + method=TEST_METHOD, + initial_delay=42.0, + ) + + generator = trigger.run() + await generator.asend(None) + mock_sleep.assert_awaited_once_with(42.0) + class TestHttpEventTrigger: @staticmethod