diff --git a/airflow/providers/amazon/aws/sensors/lambda_function.py b/airflow/providers/amazon/aws/sensors/lambda_function.py index 772ed0689a3e6..579c49b76d9a0 100644 --- a/airflow/providers/amazon/aws/sensors/lambda_function.py +++ b/airflow/providers/amazon/aws/sensors/lambda_function.py @@ -26,7 +26,7 @@ if TYPE_CHECKING: from airflow.utils.context import Context -from airflow.exceptions import AirflowException +from airflow.exceptions import AirflowException, AirflowSkipException from airflow.sensors.base import BaseSensorOperator @@ -74,9 +74,11 @@ def poke(self, context: Context) -> bool: state = self.hook.conn.get_function(**trim_none_values(get_function_args))["Configuration"]["State"] if state in self.FAILURE_STATES: - raise AirflowException( - "Lambda function state sensor failed because the Lambda is in a failed state" - ) + message = "Lambda function state sensor failed because the Lambda is in a failed state" + # TODO: remove this if block when min_airflow_version is set to higher than 2.7.1 + if self.soft_fail: + raise AirflowSkipException(message) + raise AirflowException(message) return state in self.target_states diff --git a/tests/providers/amazon/aws/sensors/test_lambda_function.py b/tests/providers/amazon/aws/sensors/test_lambda_function.py index f97eb68b34ee9..652d8996fd159 100644 --- a/tests/providers/amazon/aws/sensors/test_lambda_function.py +++ b/tests/providers/amazon/aws/sensors/test_lambda_function.py @@ -20,7 +20,7 @@ import pytest -from airflow.exceptions import AirflowException +from airflow.exceptions import AirflowException, AirflowSkipException from airflow.providers.amazon.aws.hooks.lambda_function import LambdaHook from airflow.providers.amazon.aws.sensors.lambda_function import LambdaFunctionStateSensor @@ -69,3 +69,19 @@ def test_poke(self, get_function_output, expect_failure, expected): mock_conn.get_function.assert_called_once_with( FunctionName=FUNCTION_NAME, ) + + @pytest.mark.parametrize( + "soft_fail, expected_exception", ((False, AirflowException), (True, AirflowSkipException)) + ) + def test_fail_poke(self, soft_fail, expected_exception): + sensor = LambdaFunctionStateSensor( + task_id="test_sensor", + function_name=FUNCTION_NAME, + ) + sensor.soft_fail = soft_fail + message = "Lambda function state sensor failed because the Lambda is in a failed state" + with pytest.raises(expected_exception, match=message), mock.patch( + "airflow.providers.amazon.aws.hooks.lambda_function.LambdaHook.conn" + ) as conn: + conn.get_function.return_value = {"Configuration": {"State": "Failed"}} + sensor.poke(context={})