Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
80 changes: 60 additions & 20 deletions providers/http/src/airflow/providers/http/sensors/http.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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:

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Seems like the self.response_check behvaior is being dropped? Is that an issue?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

As I understand it, this condition causes the behavior reported in #40209 because setting response_check forces the sensor onto the synchronous path. The check itself is still evaluated by self.poke(context) before deferral and by execute_complete() after the trigger returns a response. A false result defers the task again.

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)
Comment thread
yuseok89 marked this conversation as resolved.

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)
Comment thread
yuseok89 marked this conversation as resolved.

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)
26 changes: 24 additions & 2 deletions providers/http/src/airflow/providers/http/triggers/http.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is this a pattern that is being used elsewhere (initial_delay, that is)?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I could not find another trigger using the same initial_delay parameter. I added it because a failed response_check creates a new trigger, which would otherwise send its first request immediately instead of respecting poke_interval.
Since this introduces a new trigger parameter, I would be happy to continue discussing in this PR whether keeping it local to HttpSensorTrigger is appropriate or whether a different design would be preferable.

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__(
Expand All @@ -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
Expand All @@ -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."""
Expand All @@ -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"):
Expand Down
122 changes: 112 additions & 10 deletions providers/http/tests/unit/http/sensors/test_http.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
# under the License.
from __future__ import annotations

from types import SimpleNamespace
from unittest import mock
from unittest.mock import patch

Expand Down Expand Up @@ -368,32 +369,133 @@ 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_passes(self):
task = HttpSensor(
task_id="run_now",
endpoint="test-endpoint",
response_check=lambda response: "httpbin" in response.text,
deferrable=True,
)
assert task.execute_complete(context={}, event=self._make_event()) is None

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]
51 changes: 51 additions & 0 deletions providers/http/tests/unit/http/triggers/test_http.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand Down