diff --git a/providers/apache/kafka/src/airflow/providers/apache/kafka/sensors/kafka.py b/providers/apache/kafka/src/airflow/providers/apache/kafka/sensors/kafka.py index 35b35454a6733..5ee539616e291 100644 --- a/providers/apache/kafka/src/airflow/providers/apache/kafka/sensors/kafka.py +++ b/providers/apache/kafka/src/airflow/providers/apache/kafka/sensors/kafka.py @@ -204,6 +204,10 @@ def __init__( ) def execute(self, context, event=None) -> Any: + if isinstance(self.timeout, (int, float)): + timeout = timedelta(seconds=self.timeout) + else: + timeout = self.timeout self.defer( trigger=AwaitMessageTrigger( topics=self.topics, @@ -215,6 +219,7 @@ def execute(self, context, event=None) -> Any: poll_interval=self.poll_interval, ), method_name="execute_complete", + timeout=timeout, ) return event @@ -222,6 +227,10 @@ def execute(self, context, event=None) -> Any: def execute_complete(self, context, event=None): self.event_triggered_function(event, **context) + if isinstance(self.timeout, (int, float)): + timeout = timedelta(seconds=self.timeout) + else: + timeout = self.timeout self.defer( trigger=AwaitMessageTrigger( topics=self.topics, @@ -233,4 +242,5 @@ def execute_complete(self, context, event=None): poll_interval=self.poll_interval, ), method_name="execute_complete", + timeout=timeout, ) diff --git a/providers/apache/kafka/tests/unit/apache/kafka/sensors/test_kafka.py b/providers/apache/kafka/tests/unit/apache/kafka/sensors/test_kafka.py index 0934d86ae1efb..c409d60891227 100644 --- a/providers/apache/kafka/tests/unit/apache/kafka/sensors/test_kafka.py +++ b/providers/apache/kafka/tests/unit/apache/kafka/sensors/test_kafka.py @@ -19,6 +19,7 @@ import json import logging +from datetime import timedelta import pytest @@ -129,6 +130,27 @@ def test_await_message_trigger_function_with_timeout_parameter(self): assert sensor.timeout == 600 + def test_await_message_trigger_function_forwards_timeout_to_deferral(self): + """The timeout must be forwarded to every deferral, not silently ignored.""" + sensor = AwaitMessageTriggerFunctionSensor( + kafka_config_id="kafka_d", + topics=["test"], + task_id="test", + apply_function=_return_true, + event_triggered_function=_return_true, + timeout=600, + ) + + with pytest.raises(TaskDeferred) as exc_info: + sensor.execute(context={}) + assert exc_info.value.timeout == timedelta(seconds=600) + + # The sensor re-defers after every processed event, so the timeout must be + # applied to that deferral as well. + with pytest.raises(TaskDeferred) as exc_info: + sensor.execute_complete(context={}) + assert exc_info.value.timeout == timedelta(seconds=600) + def test_await_message_trigger_function_with_soft_fail_parameter(self): """Test that AwaitMessageTriggerFunctionSensor accepts soft_fail parameter.""" sensor = AwaitMessageTriggerFunctionSensor(