@@ -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
20182093class TestDagParamRuntime :
20192094 DEFAULT_ARGS = {
0 commit comments