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
1 change: 0 additions & 1 deletion generated/known_airflow_exceptions.txt
Original file line number Diff line number Diff line change
Expand Up @@ -418,5 +418,4 @@ task-sdk/src/airflow/sdk/definitions/connection.py::4
task-sdk/src/airflow/sdk/definitions/decorators/setup_teardown.py::4
task-sdk/src/airflow/sdk/definitions/xcom_arg.py::1
task-sdk/src/airflow/sdk/execution_time/task_runner.py::1
task-sdk/tests/task_sdk/bases/test_sensor.py::1
task-sdk/tests/task_sdk/execution_time/test_task_runner.py::2
30 changes: 19 additions & 11 deletions task-sdk/src/airflow/sdk/bases/sensor.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,13 +35,16 @@
AirflowSensorTimeout,
AirflowSkipException,
AirflowTaskTimeout,
TaskDeferralError,
TaskDeferralTimeout,
)

if TYPE_CHECKING:
from airflow.sdk.definitions.context import Context

# soft_fail turns these into a skip whether they surface from poke() or after resuming from a
# trigger; every other error fails the sensor. Both paths share the tuple so they cannot drift.
_SOFT_FAIL_EXCEPTIONS = (AirflowSensorTimeout, AirflowTaskTimeout, AirflowFailException)


class PokeReturnValue:
"""
Expand Down Expand Up @@ -69,8 +72,8 @@ class BaseSensorOperator(BaseOperator):
Sensor operators keep executing at a time interval and succeed when
a criteria is met and fail if and when they time out.

:param soft_fail: Set to true to mark the task as SKIPPED on failure.
Mutually exclusive with never_fail.
:param soft_fail: Set to true to mark the task as SKIPPED when it times out or raises
``AirflowFailException``. Mutually exclusive with never_fail.
:param poke_interval: Time that the job should wait in between each try.
Can be ``timedelta`` or ``float`` seconds.
:param timeout: Time elapsed before the task times out and fails.
Expand Down Expand Up @@ -205,11 +208,7 @@ def run_duration() -> float:
while True:
try:
poke_return = self.poke(context)
except (
AirflowSensorTimeout,
AirflowTaskTimeout,
AirflowFailException,
) as e:
except _SOFT_FAIL_EXCEPTIONS as e:
if self.soft_fail:
raise AirflowSkipException("Skipping due to soft_fail is set to True.") from e
if self.never_fail:
Expand Down Expand Up @@ -251,19 +250,28 @@ def run_duration() -> float:
return xcom_value

def resume_execution(self, next_method: str, next_kwargs: dict[str, Any] | None, context: Context):
# Use nested try/except to convert TaskDeferralTimeout to AirflowSensorTimeout
# while still allowing soft_fail/never_fail to handle both exception types.
# Nested try/except so a trigger timeout becomes AirflowSensorTimeout before soft_fail
# and never_fail are applied to it.
try:
try:
return super().resume_execution(next_method, next_kwargs, context)
except TaskDeferralTimeout as e:
raise AirflowSensorTimeout(*e.args) from e
except (AirflowException, TaskDeferralError) as e:
except _SOFT_FAIL_EXCEPTIONS as e:
if self.soft_fail:
raise AirflowSkipException("Skipping due to soft_fail is set to True.") from e
if self.never_fail:
raise AirflowSkipException("Skipping due to never_fail is set to True.") from e
raise
except AirflowSkipException:
raise
except AirflowException as e:
# execute() only skips the exceptions handled above; anything else here (a crashed
# trigger, an error event raised by execute_complete) is a real failure, and soft_fail
# must not hide it behind a skip.
if self.never_fail:
raise AirflowSkipException("Skipping due to never_fail is set to True.") from e
raise

def _get_next_poke_interval(
self,
Expand Down
50 changes: 30 additions & 20 deletions task-sdk/tests/task_sdk/bases/test_sensor.py
Original file line number Diff line number Diff line change
Expand Up @@ -62,12 +62,13 @@ def poke(self, context: Context):


class DummyAsyncSensor(BaseSensorOperator):
def __init__(self, return_value=False, **kwargs):
def __init__(self, return_value=False, failure_exception: BaseException | None = None, **kwargs):
super().__init__(**kwargs)
self.return_value = return_value
self.failure_exception = failure_exception or AirflowException("Sensor failed")

def execute_complete(self, context, event=None):
raise AirflowException("Should be skipped")
raise self.failure_exception


class DummySensorWithXcomValue(BaseSensorOperator):
Expand Down Expand Up @@ -688,16 +689,25 @@ def test_poke_mode_only_bad_poke(self):


class TestAsyncSensor:
@pytest.mark.parametrize("soft_fail", [True, False])
def test_error_from_execute_complete_fails_regardless_of_soft_fail(self, soft_fail):
async_sensor = DummyAsyncSensor(task_id="dummy_async_sensor", soft_fail=soft_fail)
with pytest.raises(AirflowException, match="Sensor failed"):
async_sensor.resume_execution("execute_complete", None, {})

@pytest.mark.parametrize(
("soft_fail", "expected_exception"),
"failure_exception",
[
(True, AirflowSkipException),
(False, AirflowException),
pytest.param(AirflowSensorTimeout("timed out"), id="sensor-timeout"),
pytest.param(AirflowTaskTimeout("timed out"), id="task-timeout"),
pytest.param(AirflowFailException("failed"), id="fail-exception"),
],
)
def test_fail_after_resuming_deferred_sensor(self, soft_fail, expected_exception):
async_sensor = DummyAsyncSensor(task_id="dummy_async_sensor", soft_fail=soft_fail)
with pytest.raises(expected_exception):
def test_soft_fail_skips_timeout_and_fail_exceptions_after_resuming(self, failure_exception):
async_sensor = DummyAsyncSensor(
task_id="dummy_async_sensor", soft_fail=True, failure_exception=failure_exception
)
with pytest.raises(AirflowSkipException):
async_sensor.resume_execution("execute_complete", None, {})

@pytest.mark.parametrize(
Expand Down Expand Up @@ -727,25 +737,25 @@ def test_timeout_after_resuming_deferred_sensor_with_never_fail(self):
context={},
)

@pytest.mark.parametrize(
("soft_fail", "expected_exception"),
[
(True, AirflowSkipException),
(False, TaskDeferralError),
],
)
def test_trigger_failure_after_resuming_deferred_sensor_with_soft_fail(
self, soft_fail, expected_exception
):
"""Test that deferrable sensors with soft_fail skip on trigger failure instead of failing."""
@pytest.mark.parametrize("soft_fail", [True, False])
def test_trigger_failure_after_resuming_deferred_sensor_fails_regardless_of_soft_fail(self, soft_fail):
async_sensor = DummyAsyncSensor(task_id="dummy_async_sensor", soft_fail=soft_fail)
with pytest.raises(expected_exception):
with pytest.raises(TaskDeferralError):
async_sensor.resume_execution(
next_method="__fail__",
next_kwargs={"error": TriggerFailureReason.TRIGGER_FAILURE},
context={},
)

def test_never_fail_keeps_the_skip_raised_by_execute_complete(self):
async_sensor = DummyAsyncSensor(
task_id="dummy_async_sensor",
never_fail=True,
failure_exception=AirflowSkipException("External job has skipped"),
)
with pytest.raises(AirflowSkipException, match="External job has skipped"):
async_sensor.resume_execution("execute_complete", None, {})

def test_trigger_failure_after_resuming_deferred_sensor_with_never_fail(self):
"""Test that deferrable sensors with never_fail skip on trigger failure."""
async_sensor = DummyAsyncSensor(task_id="dummy_async_sensor", never_fail=True)
Expand Down