Skip to content
Merged
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 @@ -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")

Expand All @@ -131,7 +130,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)
Comment thread
amoghrajesh marked this conversation as resolved.
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
Expand Down Expand Up @@ -160,9 +159,18 @@ 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)
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

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Should we re raise this one?

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think this is fine, this raises error happen only when base_xcom_deser_result is object store path, likely this is string that means we should continue to next below to read from ObjectStore with path. WDYT?

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yep seems ok

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:
Expand Down
28 changes: 28 additions & 0 deletions providers/common/io/tests/unit/common/io/xcom/test_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,8 @@
# under the License.
from __future__ import annotations

from unittest.mock import MagicMock

import pytest

import airflow.models.xcom
Expand Down Expand Up @@ -102,6 +104,8 @@ 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"}
)
Expand Down Expand Up @@ -362,3 +366,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, 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