Skip to content

Commit 2b993be

Browse files
amoghrajeshkaxil
authored andcommitted
Fix custom xcom backend serialize when BaseXCom.get_all is used (#53814)
(cherry picked from commit a8c4ba3)
1 parent 18e17f2 commit 2b993be

2 files changed

Lines changed: 82 additions & 4 deletions

File tree

task-sdk/src/airflow/sdk/bases/xcom.py

Lines changed: 7 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,7 @@
1717

1818
from __future__ import annotations
1919

20+
import collections
2021
from typing import Any, Protocol
2122

2223
import structlog
@@ -30,6 +31,9 @@
3031
XComSequenceSliceResult,
3132
)
3233

34+
# Lightweight wrapper for XCom values
35+
_XComValueWrapper = collections.namedtuple("_XComValueWrapper", "value")
36+
3337
log = structlog.get_logger(logger_name="task")
3438

3539

@@ -290,7 +294,6 @@ def get_all(
290294
:return: List of all XCom values if found.
291295
"""
292296
from airflow.sdk.execution_time.task_runner import SUPERVISOR_COMMS
293-
from airflow.serialization.serde import deserialize
294297

295298
msg = SUPERVISOR_COMMS.send(
296299
msg=GetXComSequenceSlice(
@@ -307,10 +310,10 @@ def get_all(
307310
if not isinstance(msg, XComSequenceSliceResult):
308311
raise TypeError(f"Expected XComSequenceSliceResult, received: {type(msg)} {msg}")
309312

310-
result = deserialize(msg.root)
311-
if not result:
313+
if not msg.root:
312314
return None
313-
return result
315+
316+
return [cls.deserialize_value(_XComValueWrapper(value)) for value in msg.root]
314317

315318
@staticmethod
316319
def serialize_value(

task-sdk/tests/task_sdk/execution_time/test_task_runner.py

Lines changed: 75 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2014,6 +2014,81 @@ def execute(self, context):
20142014
for x in mock_supervisor_comms.send.call_args_list
20152015
)
20162016

2017+
def test_get_all_uses_custom_deserialize_value(self, mock_supervisor_comms):
2018+
"""
2019+
Tests that XCom.get_all() calls the custom deserialize_value method.
2020+
"""
2021+
2022+
class CustomXCom(BaseXCom):
2023+
@classmethod
2024+
def deserialize_value(cls, result):
2025+
"""Custom deserialization that adds a prefix to show it was called."""
2026+
original_value = super().deserialize_value(result)
2027+
return f"from custom xcom deserialize:{original_value}"
2028+
2029+
serialized_values = ["value1", "value2", "value3"]
2030+
mock_supervisor_comms.send.return_value = XComSequenceSliceResult(root=serialized_values)
2031+
2032+
result = CustomXCom.get_all(key="test_key", dag_id="test_dag", task_id="test_task", run_id="test_run")
2033+
2034+
expected = [
2035+
"from custom xcom deserialize:value1",
2036+
"from custom xcom deserialize:value2",
2037+
"from custom xcom deserialize:value3",
2038+
]
2039+
assert result == expected
2040+
2041+
@pytest.mark.parametrize(
2042+
("include_prior_dates", "expected_value"),
2043+
[
2044+
pytest.param(True, True, id="include_prior_dates_true"),
2045+
pytest.param(False, False, id="include_prior_dates_false"),
2046+
pytest.param(None, False, id="include_prior_dates_default"),
2047+
],
2048+
)
2049+
def test_xcom_pull_with_include_prior_dates(
2050+
self,
2051+
create_runtime_ti,
2052+
mock_supervisor_comms,
2053+
include_prior_dates,
2054+
expected_value,
2055+
):
2056+
"""Test that xcom_pull with include_prior_dates parameter correctly behaves as we expect."""
2057+
task = BaseOperator(task_id="pull_task")
2058+
runtime_ti = create_runtime_ti(task=task)
2059+
2060+
value = {"previous_run_data": "test_value"}
2061+
ser_value = BaseXCom.serialize_value(value)
2062+
2063+
def mock_send_side_effect(*args, **kwargs):
2064+
msg = kwargs.get("msg") or args[0]
2065+
if isinstance(msg, GetXComSequenceSlice):
2066+
assert msg.include_prior_dates is expected_value, (
2067+
f"include_prior_dates should be {expected_value} in GetXComSequenceSlice"
2068+
)
2069+
return XComSequenceSliceResult(root=[ser_value])
2070+
return XComResult(key="test_key", value=None)
2071+
2072+
mock_supervisor_comms.send.side_effect = mock_send_side_effect
2073+
kwargs = {"key": "test_key", "task_ids": "previous_task"}
2074+
if include_prior_dates is not None:
2075+
kwargs["include_prior_dates"] = include_prior_dates
2076+
result = runtime_ti.xcom_pull(**kwargs)
2077+
assert result == value
2078+
2079+
mock_supervisor_comms.send.assert_called_once_with(
2080+
msg=GetXComSequenceSlice(
2081+
key="test_key",
2082+
dag_id=runtime_ti.dag_id,
2083+
run_id=runtime_ti.run_id,
2084+
task_id="previous_task",
2085+
start=None,
2086+
stop=None,
2087+
step=None,
2088+
include_prior_dates=expected_value,
2089+
),
2090+
)
2091+
20172092

20182093
class TestDagParamRuntime:
20192094
DEFAULT_ARGS = {

0 commit comments

Comments
 (0)