Skip to content
Draft
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
Original file line number Diff line number Diff line change
Expand Up @@ -803,7 +803,7 @@ def execute(self, context: Context):
if not self.deferrable:
return self.execute_sync(context)

self.execute_async(context)
return self.execute_async(context)

def execute_sync(self, context: Context):
result = None
Expand Down Expand Up @@ -955,7 +955,7 @@ def _refresh_cached_properties(self):
del self.client
del self.pod_manager

def execute_async(self, context: Context) -> None:
def execute_async(self, context: Context) -> Any:
if self.pod_request_obj is None:
self.pod_request_obj = self.build_pod_request_obj(context)
for callback in self.callbacks:
Expand Down Expand Up @@ -991,9 +991,8 @@ def execute_async(self, context: Context) -> None:
# provider where invoke_defer_method does not accept context parameter
sig = inspect.signature(self.invoke_defer_method)
if "context" in sig.parameters:
self.invoke_defer_method(context=context)
else:
self.invoke_defer_method()
return self.invoke_defer_method(context=context)
return self.invoke_defer_method()

def convert_config_file_to_dict(self):
"""Convert passed config_file to dict representation."""
Expand All @@ -1006,7 +1005,7 @@ def convert_config_file_to_dict(self):

def invoke_defer_method(
self, last_log_time: DateTime | None = None, context: Context | None = None
) -> None:
) -> Any:
"""Redefine triggers which are being used in child classes."""
self.convert_config_file_to_dict()

Expand Down Expand Up @@ -1071,7 +1070,7 @@ def invoke_defer_method(
pod_container_state == ContainerState.TERMINATED or pod_container_state == ContainerState.FAILED
):
self.log.info("Skipping deferral as pod is already in a terminal state")
self.trigger_reentry(
return self.trigger_reentry(
context=context,
event={
"status": "failed" if pod_container_state == ContainerState.FAILED else "success",
Expand All @@ -1084,8 +1083,7 @@ def invoke_defer_method(
**(self.trigger_kwargs or {}),
},
)
else:
self.defer(trigger=trigger, method_name="trigger_reentry", timeout=defer_timeout)
self.defer(trigger=trigger, method_name="trigger_reentry", timeout=defer_timeout)

def trigger_reentry(self, context: Context, event: dict[str, Any]) -> Any:
"""
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -355,8 +355,7 @@ def execute(self, context: Context):
self._setup_spark_configuration(context)

if self.deferrable:
self.execute_async(context)
return
return self.execute_async(context)

return super().execute(context)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -3155,6 +3155,73 @@ def test_invoke_defer_method_passes_execution_deadline_when_execution_timeout_se
assert trigger.trigger_kwargs["_execution_deadline"] == expected_deadline
assert exc.value.timeout == datetime.timedelta(seconds=270 + 60)

@patch(KUB_OP_PATH.format("trigger_reentry"))
@patch(KUB_OP_PATH.format("convert_config_file_to_dict"))
@patch("airflow.providers.cncf.kubernetes.operators.pod.BaseHook.get_connection")
def test_invoke_defer_method_returns_trigger_reentry_result_when_pod_already_terminal(
self, mocked_get_connection, mocked_convert_config, mocked_trigger_reentry
):
"""
When the pod is already terminal, ``invoke_defer_method`` calls
``trigger_reentry`` inline instead of deferring. Its return value carries the
XCom sidecar output, so it has to be propagated to the caller -- in the regular
deferral path Airflow takes it from the resume method and stores it as
``return_value``.

Dropping it makes ``do_xcom_push`` silently produce no ``return_value`` while
the task still succeeds, and downstream tasks pulling that XCom get ``None``.
"""
mocked_get_connection.side_effect = AirflowNotFoundException("connection not found")
mocked_trigger_reentry.return_value = {"key": "value"}

k = KubernetesPodOperator(
task_id=TEST_TASK_ID,
namespace=TEST_NAMESPACE,
image=TEST_IMAGE,
name=TEST_NAME,
on_finish_action="keep_pod",
in_cluster=True,
deferrable=True,
do_xcom_push=True,
)
k.pod = MagicMock()
k.pod.metadata.name = TEST_NAME
k.pod.metadata.namespace = TEST_NAMESPACE

context = {"ti": MagicMock()}
with patch(f"{TRIGGER_CLASS}.define_pod_container_state", return_value=ContainerState.TERMINATED):
result = k.invoke_defer_method(context=context)

mocked_trigger_reentry.assert_called_once()
assert result == {"key": "value"}

@patch(KUB_OP_PATH.format("invoke_defer_method"))
@patch(KUB_OP_PATH.format("build_pod_request_obj"))
@patch(KUB_OP_PATH.format("get_or_create_pod"))
def test_execute_returns_deferrable_result(
self, mocked_get_or_create_pod, mocked_build_pod_request_obj, mocked_invoke_defer_method
):
"""
``execute`` has to hand the deferrable result back to Airflow the same way the
synchronous branch does, otherwise the value produced by the inline
``trigger_reentry`` call never becomes the task's ``return_value`` XCom.
"""
mocked_invoke_defer_method.return_value = {"key": "value"}
mocked_get_or_create_pod.return_value = MagicMock()

k = KubernetesPodOperator(
task_id=TEST_TASK_ID,
namespace=TEST_NAMESPACE,
image=TEST_IMAGE,
name=TEST_NAME,
on_finish_action="keep_pod",
in_cluster=True,
deferrable=True,
do_xcom_push=True,
)

assert k.execute(context={"ti": MagicMock()}) == {"key": "value"}

@patch(KUB_OP_PATH.format("convert_config_file_to_dict"))
@patch("airflow.providers.cncf.kubernetes.operators.pod.BaseHook.get_connection")
def test_invoke_defer_method_pads_defer_timeout_for_slow_poll_interval(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1272,7 +1272,8 @@ def test_execute_deferrable_does_not_call_super(

mock_execute_async.assert_called_once_with(context)
mock_parent_execute.assert_not_called()
assert result is None
# The deferrable result carries the XCom sidecar output, so it has to reach the caller.
assert result is mock_execute_async.return_value

def test_execute_non_deferrable_calls_super(
self,
Expand Down