From b5a6e87040712e299eda650b64c7e02343343fd0 Mon Sep 17 00:00:00 2001 From: Pavan Kumar Date: Sun, 13 Apr 2025 04:29:18 +0100 Subject: [PATCH 1/7] Fix XComObjectStorageBackend deserialize_value to use json loads --- .../io/src/airflow/providers/common/io/xcom/backend.py | 8 ++++++-- .../common/io/tests/unit/common/io/xcom/test_backend.py | 6 +++++- 2 files changed, 11 insertions(+), 3 deletions(-) diff --git a/providers/common/io/src/airflow/providers/common/io/xcom/backend.py b/providers/common/io/src/airflow/providers/common/io/xcom/backend.py index 5fc86b8f182ee..3c7a621e1367e 100644 --- a/providers/common/io/src/airflow/providers/common/io/xcom/backend.py +++ b/providers/common/io/src/airflow/providers/common/io/xcom/backend.py @@ -160,9 +160,13 @@ def deserialize_value(result) -> Any: Compression is inferred from the file extension. """ - data = BaseXCom.deserialize_value(result) + base_xcom_deser_result = BaseXCom.deserialize_value(result) + + # When XComObjectStorageBackend is used, xcom value will be serialized using json.dumps + # likely, we need to deserialize it using json.loads + data = json.loads(base_xcom_deser_result, cls=XComDecoder) try: - path = XComObjectStorageBackend._get_full_path(data) + path = XComObjectStorageBackend._get_full_path(base_xcom_deser_result) except (TypeError, ValueError): # Likely value stored directly in the database. return data try: diff --git a/providers/common/io/tests/unit/common/io/xcom/test_backend.py b/providers/common/io/tests/unit/common/io/xcom/test_backend.py index 802106024b887..67869edcd5242 100644 --- a/providers/common/io/tests/unit/common/io/xcom/test_backend.py +++ b/providers/common/io/tests/unit/common/io/xcom/test_backend.py @@ -17,6 +17,8 @@ # under the License. from __future__ import annotations +import json + import pytest import airflow.models.xcom @@ -102,8 +104,10 @@ def test_value_db(self, task_instance, mock_supervisor_comms, session): ) if AIRFLOW_V_3_0_PLUS: + # When using XComObjectStorageBackend, the value is stored in the db is serialized with json dumps + # so we need to mimic that same behavior below. mock_supervisor_comms.get_message.return_value = XComResult( - key="return_value", value={"key": "value"} + key="return_value", value=json.dumps({"key": "value"}) ) value = XCom.get_value( From af98bd7a801f9b8b708b5fd30d1a612acf2eedad Mon Sep 17 00:00:00 2001 From: Pavan Kumar Date: Sun, 13 Apr 2025 09:31:02 +0100 Subject: [PATCH 2/7] Fix tests --- .../io/src/airflow/providers/common/io/xcom/backend.py | 5 ++++- .../common/io/tests/unit/common/io/xcom/test_backend.py | 7 ++++--- 2 files changed, 8 insertions(+), 4 deletions(-) diff --git a/providers/common/io/src/airflow/providers/common/io/xcom/backend.py b/providers/common/io/src/airflow/providers/common/io/xcom/backend.py index 3c7a621e1367e..7079758ce7844 100644 --- a/providers/common/io/src/airflow/providers/common/io/xcom/backend.py +++ b/providers/common/io/src/airflow/providers/common/io/xcom/backend.py @@ -164,7 +164,10 @@ def deserialize_value(result) -> Any: # When XComObjectStorageBackend is used, xcom value will be serialized using json.dumps # likely, we need to deserialize it using json.loads - data = json.loads(base_xcom_deser_result, cls=XComDecoder) + try: + data = json.loads(base_xcom_deser_result, cls=XComDecoder) + except (TypeError, ValueError): + data = base_xcom_deser_result try: path = XComObjectStorageBackend._get_full_path(base_xcom_deser_result) except (TypeError, ValueError): # Likely value stored directly in the database. diff --git a/providers/common/io/tests/unit/common/io/xcom/test_backend.py b/providers/common/io/tests/unit/common/io/xcom/test_backend.py index 67869edcd5242..cc3cae589e850 100644 --- a/providers/common/io/tests/unit/common/io/xcom/test_backend.py +++ b/providers/common/io/tests/unit/common/io/xcom/test_backend.py @@ -170,7 +170,7 @@ def test_value_storage(self, task_instance, mock_supervisor_comms, session): if AIRFLOW_V_3_0_PLUS: mock_supervisor_comms.get_message.return_value = XComResult( - key=XCOM_RETURN_KEY, value={"key": "bigvaluebigvaluebigvalue" * 100} + key=XCOM_RETURN_KEY, value=json.dumps({"key": "bigvaluebigvaluebigvalue" * 100}) ) value = XCom.get_value( @@ -197,6 +197,7 @@ def test_value_storage(self, task_instance, mock_supervisor_comms, session): session=session, ) assert str(p) == qry.first().value + raise def test_clear(self, task_instance, session, mock_supervisor_comms): session.add(task_instance) @@ -252,7 +253,7 @@ def test_clear(self, task_instance, session, mock_supervisor_comms): if AIRFLOW_V_3_0_PLUS: mock_supervisor_comms.get_message.return_value = XComResult( - key=XCOM_RETURN_KEY, value={"key": "superlargevalue" * 100} + key=XCOM_RETURN_KEY, value=json.dumps({"key": "superlargevalue" * 100}) ) value = XCom.get_value( key=XCOM_RETURN_KEY, @@ -357,7 +358,7 @@ def test_compression(self, task_instance, session, mock_supervisor_comms): if AIRFLOW_V_3_0_PLUS: mock_supervisor_comms.get_message.return_value = XComResult( - key=XCOM_RETURN_KEY, value={"key": "superlargevalue" * 100} + key=XCOM_RETURN_KEY, value=json.dumps({"key": "superlargevalue" * 100}) ) value = XCom.get_value( From 63d717977e52295674936404f51197d11bdf7134 Mon Sep 17 00:00:00 2001 From: Pavan Kumar Date: Sun, 13 Apr 2025 09:47:36 +0100 Subject: [PATCH 3/7] remove raise --- providers/common/io/tests/unit/common/io/xcom/test_backend.py | 1 - 1 file changed, 1 deletion(-) diff --git a/providers/common/io/tests/unit/common/io/xcom/test_backend.py b/providers/common/io/tests/unit/common/io/xcom/test_backend.py index cc3cae589e850..09716c0db370c 100644 --- a/providers/common/io/tests/unit/common/io/xcom/test_backend.py +++ b/providers/common/io/tests/unit/common/io/xcom/test_backend.py @@ -197,7 +197,6 @@ def test_value_storage(self, task_instance, mock_supervisor_comms, session): session=session, ) assert str(p) == qry.first().value - raise def test_clear(self, task_instance, session, mock_supervisor_comms): session.add(task_instance) From f18de860e991aa87d609114fea777508e7e7520a Mon Sep 17 00:00:00 2001 From: Pavan Kumar Date: Sun, 13 Apr 2025 12:36:05 +0100 Subject: [PATCH 4/7] Add basic serialization_deserialization tests --- .../tests/unit/common/io/xcom/test_backend.py | 25 +++++++++++++++++++ 1 file changed, 25 insertions(+) diff --git a/providers/common/io/tests/unit/common/io/xcom/test_backend.py b/providers/common/io/tests/unit/common/io/xcom/test_backend.py index 09716c0db370c..18834fd8e4efe 100644 --- a/providers/common/io/tests/unit/common/io/xcom/test_backend.py +++ b/providers/common/io/tests/unit/common/io/xcom/test_backend.py @@ -18,6 +18,7 @@ from __future__ import annotations import json +from unittest.mock import MagicMock import pytest @@ -366,3 +367,27 @@ def test_compression(self, task_instance, session, mock_supervisor_comms): ) assert value == {"key": "superlargevalue" * 100} + + @pytest.mark.parametrize( + "value, expected_value", + [ + pytest.param(1, 1, id="int"), + pytest.param(1.0, 1.0, id="float"), + pytest.param("string", "string", id="str"), + pytest.param(True, True, id="bool"), + pytest.param({"key": "value"}, {"key": "value"}, id="dict"), + pytest.param({"key": {"key": "value"}}, {"key": {"key": "value"}}, id="nested_dict"), + pytest.param([1, 2, 3], [1, 2, 3], id="list"), + pytest.param((1, 2, 3), (1, 2, 3), id="tuple"), + pytest.param(None, None, id="none"), + ], + ) + def test_serialization_deserialization_basic(self, task_instance, value, expected_value): + XCom = resolve_xcom_backend() + airflow.models.xcom.XCom = XCom + + serialized_data = XCom.serialize_value(value) + mock_xcom_ser = MagicMock(value=serialized_data) + deserialized_data = XCom.deserialize_value(mock_xcom_ser) + + assert deserialized_data == expected_value From f04fa852c7ae2fd89dec873506f9bda7c0d3ab6b Mon Sep 17 00:00:00 2001 From: Pavan Kumar Date: Sun, 13 Apr 2025 13:21:29 +0100 Subject: [PATCH 5/7] Add basic serialization_deserialization tests --- providers/common/io/tests/unit/common/io/xcom/test_backend.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/providers/common/io/tests/unit/common/io/xcom/test_backend.py b/providers/common/io/tests/unit/common/io/xcom/test_backend.py index 18834fd8e4efe..9a386692593a5 100644 --- a/providers/common/io/tests/unit/common/io/xcom/test_backend.py +++ b/providers/common/io/tests/unit/common/io/xcom/test_backend.py @@ -382,7 +382,7 @@ def test_compression(self, task_instance, session, mock_supervisor_comms): pytest.param(None, None, id="none"), ], ) - def test_serialization_deserialization_basic(self, task_instance, value, expected_value): + def test_serialization_deserialization_basic(self, value, expected_value): XCom = resolve_xcom_backend() airflow.models.xcom.XCom = XCom From 1ecfa3f071bb227bb1a24ad4c0e3fdb1368e1649 Mon Sep 17 00:00:00 2001 From: Pavan Kumar Date: Mon, 14 Apr 2025 08:02:52 +0100 Subject: [PATCH 6/7] Use BaseXCom serialize_value when objectstorage_threshold is lessthan given input --- .../providers/common/io/xcom/backend.py | 18 ++++++++++-------- .../tests/unit/common/io/xcom/test_backend.py | 9 ++++----- 2 files changed, 14 insertions(+), 13 deletions(-) diff --git a/providers/common/io/src/airflow/providers/common/io/xcom/backend.py b/providers/common/io/src/airflow/providers/common/io/xcom/backend.py index 7079758ce7844..c59098adf8405 100644 --- a/providers/common/io/src/airflow/providers/common/io/xcom/backend.py +++ b/providers/common/io/src/airflow/providers/common/io/xcom/backend.py @@ -131,7 +131,7 @@ def serialize_value( # type: ignore[override] threshold = _get_threshold() if threshold < 0 or len(s_val_encoded) < threshold: # Either no threshold or value is small enough. if AIRFLOW_V_3_0_PLUS: - return s_val + return BaseXCom.serialize_value(value) else: # TODO: Remove this branch once we drop support for Airflow 2 # This is for Airflow 2.10 where the value is expected to be bytes @@ -161,13 +161,15 @@ def deserialize_value(result) -> Any: Compression is inferred from the file extension. """ base_xcom_deser_result = BaseXCom.deserialize_value(result) - - # When XComObjectStorageBackend is used, xcom value will be serialized using json.dumps - # likely, we need to deserialize it using json.loads - try: - data = json.loads(base_xcom_deser_result, cls=XComDecoder) - except (TypeError, ValueError): - data = base_xcom_deser_result + data = base_xcom_deser_result + + if not AIRFLOW_V_3_0_PLUS: + try: + # When XComObjectStorageBackend is used, xcom value will be serialized using json.dumps + # likely, we need to deserialize it using json.loads + data = json.loads(base_xcom_deser_result, cls=XComDecoder) + except (TypeError, ValueError): + pass try: path = XComObjectStorageBackend._get_full_path(base_xcom_deser_result) except (TypeError, ValueError): # Likely value stored directly in the database. diff --git a/providers/common/io/tests/unit/common/io/xcom/test_backend.py b/providers/common/io/tests/unit/common/io/xcom/test_backend.py index 9a386692593a5..99fb46a66c7e5 100644 --- a/providers/common/io/tests/unit/common/io/xcom/test_backend.py +++ b/providers/common/io/tests/unit/common/io/xcom/test_backend.py @@ -17,7 +17,6 @@ # under the License. from __future__ import annotations -import json from unittest.mock import MagicMock import pytest @@ -108,7 +107,7 @@ def test_value_db(self, task_instance, mock_supervisor_comms, session): # When using XComObjectStorageBackend, the value is stored in the db is serialized with json dumps # so we need to mimic that same behavior below. mock_supervisor_comms.get_message.return_value = XComResult( - key="return_value", value=json.dumps({"key": "value"}) + key="return_value", value={"key": "value"} ) value = XCom.get_value( @@ -171,7 +170,7 @@ def test_value_storage(self, task_instance, mock_supervisor_comms, session): if AIRFLOW_V_3_0_PLUS: mock_supervisor_comms.get_message.return_value = XComResult( - key=XCOM_RETURN_KEY, value=json.dumps({"key": "bigvaluebigvaluebigvalue" * 100}) + key=XCOM_RETURN_KEY, value={"key": "bigvaluebigvaluebigvalue" * 100} ) value = XCom.get_value( @@ -253,7 +252,7 @@ def test_clear(self, task_instance, session, mock_supervisor_comms): if AIRFLOW_V_3_0_PLUS: mock_supervisor_comms.get_message.return_value = XComResult( - key=XCOM_RETURN_KEY, value=json.dumps({"key": "superlargevalue" * 100}) + key=XCOM_RETURN_KEY, value={"key": "superlargevalue" * 100} ) value = XCom.get_value( key=XCOM_RETURN_KEY, @@ -358,7 +357,7 @@ def test_compression(self, task_instance, session, mock_supervisor_comms): if AIRFLOW_V_3_0_PLUS: mock_supervisor_comms.get_message.return_value = XComResult( - key=XCOM_RETURN_KEY, value=json.dumps({"key": "superlargevalue" * 100}) + key=XCOM_RETURN_KEY, value={"key": "superlargevalue" * 100} ) value = XCom.get_value( From a7ebc36ebe57f6de965219934c4dc0a6079b271f Mon Sep 17 00:00:00 2001 From: Pavan Kumar Date: Mon, 14 Apr 2025 08:05:40 +0100 Subject: [PATCH 7/7] Use BaseXCom serialize_value when objectstorage_threshold is lessthan given input --- .../common/io/src/airflow/providers/common/io/xcom/backend.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/providers/common/io/src/airflow/providers/common/io/xcom/backend.py b/providers/common/io/src/airflow/providers/common/io/xcom/backend.py index c59098adf8405..c5144222e5cb3 100644 --- a/providers/common/io/src/airflow/providers/common/io/xcom/backend.py +++ b/providers/common/io/src/airflow/providers/common/io/xcom/backend.py @@ -118,8 +118,7 @@ def serialize_value( # type: ignore[override] run_id: str | None = None, map_index: int | None = None, ) -> bytes | str: - # we will always serialize ourselves and not by BaseXCom as the deserialize method - # from BaseXCom accepts only XCom objects and not the value directly + # We will use this serialized value to write to the object store. s_val = json.dumps(value, cls=XComEncoder) s_val_encoded = s_val.encode("utf-8")