|
3 | 3 |
|
4 | 4 | """Tests for the DataConverter abstraction and the default JsonDataConverter.""" |
5 | 5 |
|
| 6 | +import inspect |
6 | 7 | import json |
7 | 8 | import logging |
8 | 9 | from dataclasses import dataclass, field |
9 | 10 | from typing import Any |
| 11 | +from unittest.mock import patch |
10 | 12 |
|
| 13 | +import pytest |
| 14 | + |
| 15 | +from durabletask.internal.entity_state_shim import StateShim |
11 | 16 | from durabletask.serialization import ( |
12 | 17 | DEFAULT_DATA_CONVERTER, |
13 | 18 | DataConverter, |
@@ -57,6 +62,215 @@ def test_round_trip_nested_dataclass_with_derived_field(): |
57 | 62 | assert converter.deserialize(converter.serialize(shipment), Shipment) == shipment |
58 | 63 |
|
59 | 64 |
|
| 65 | +@dataclass |
| 66 | +class CounterState: |
| 67 | + name: str |
| 68 | + counter: int = field(init=False, default=0) |
| 69 | + |
| 70 | + def __init__(self, name: str, counter: int = 0): |
| 71 | + self.name = name |
| 72 | + self.counter = counter |
| 73 | + |
| 74 | + |
| 75 | +@dataclass |
| 76 | +class KeywordCounterState: |
| 77 | + name: str |
| 78 | + counter: int = field(init=False, default=0) |
| 79 | + |
| 80 | + def __init__(self, name: str, *, counter: int = 0): |
| 81 | + self.name = name |
| 82 | + self.counter = counter |
| 83 | + |
| 84 | + |
| 85 | +@dataclass |
| 86 | +class KwargsCounterState: |
| 87 | + name: str |
| 88 | + counter: int = field(init=False, default=0) |
| 89 | + |
| 90 | + def __init__(self, name: str, **kwargs): |
| 91 | + self.name = name |
| 92 | + self.counter = kwargs.get("counter", 0) |
| 93 | + |
| 94 | + |
| 95 | +@pytest.mark.parametrize("state_type", [CounterState, KeywordCounterState, KwargsCounterState]) |
| 96 | +def test_non_init_field_preserved_by_custom_constructor(state_type): |
| 97 | + converter = JsonDataConverter() |
| 98 | + original = state_type("persisted", counter=42) |
| 99 | + encoded = converter.serialize(original) |
| 100 | + restored = converter.deserialize(encoded, state_type) |
| 101 | + assert restored == original |
| 102 | + assert restored.counter == 42 |
| 103 | + |
| 104 | + state = StateShim(encoded, converter, is_serialized=True) |
| 105 | + restored_state = state.get_state(state_type) |
| 106 | + assert restored_state.counter == 42 |
| 107 | + state.set_state(restored_state) |
| 108 | + assert json.loads(state.encode_state()) == {"name": "persisted", "counter": 42} |
| 109 | + |
| 110 | + |
| 111 | +@pytest.mark.parametrize("accept_kwargs", [False, True]) |
| 112 | +def test_non_init_custom_constructor_field_is_recursively_coerced(accept_kwargs): |
| 113 | + @dataclass |
| 114 | + class CustomShipment: |
| 115 | + order: Order = field(init=False) |
| 116 | + |
| 117 | + if accept_kwargs: |
| 118 | + def initializer(self, **kwargs): |
| 119 | + self.order = kwargs["order"] |
| 120 | + else: |
| 121 | + def initializer(self, *, order): |
| 122 | + self.order = order |
| 123 | + CustomShipment.__init__ = initializer |
| 124 | + |
| 125 | + converter = JsonDataConverter() |
| 126 | + restored = converter.deserialize('{"order": {"item": "book", "quantity": 42}}', CustomShipment) |
| 127 | + assert isinstance(restored, CustomShipment) |
| 128 | + assert restored.order == Order("book", 42) |
| 129 | + |
| 130 | + |
| 131 | +def test_non_init_positional_only_constructor_parameter_is_omitted(): |
| 132 | + @dataclass |
| 133 | + class PositionalCounter: |
| 134 | + counter: int = field(init=False) |
| 135 | + |
| 136 | + def __init__(self, counter=0, /): |
| 137 | + self.counter = counter |
| 138 | + |
| 139 | + restored = JsonDataConverter().deserialize('{"counter": 42}', PositionalCounter) |
| 140 | + assert isinstance(restored, PositionalCounter) |
| 141 | + assert restored.counter == 0 |
| 142 | + |
| 143 | + |
| 144 | +@pytest.mark.parametrize("method_type", [classmethod, staticmethod]) |
| 145 | +def test_non_init_field_preserved_by_descriptor_initializer(method_type): |
| 146 | + @dataclass |
| 147 | + class DescriptorCounter: |
| 148 | + counter: int = field(init=False, default=0) |
| 149 | + |
| 150 | + if method_type is classmethod: |
| 151 | + def initializer(cls, counter=0): |
| 152 | + cls.counter = counter |
| 153 | + else: |
| 154 | + def initializer(counter=0): |
| 155 | + DescriptorCounter.counter = counter |
| 156 | + DescriptorCounter.__init__ = method_type(initializer) |
| 157 | + |
| 158 | + converter = JsonDataConverter() |
| 159 | + encoded = converter.serialize(DescriptorCounter(counter=42)) |
| 160 | + DescriptorCounter.counter = 0 |
| 161 | + restored = converter.deserialize(encoded, DescriptorCounter) |
| 162 | + assert isinstance(restored, DescriptorCounter) |
| 163 | + assert restored.counter == 42 |
| 164 | + |
| 165 | + |
| 166 | +def test_non_init_field_matching_receiver_name_is_omitted(): |
| 167 | + @dataclass |
| 168 | + class ReceiverCounter: |
| 169 | + counter: int = field(init=False) |
| 170 | + |
| 171 | + def __init__(counter): |
| 172 | + counter.counter = 7 |
| 173 | + |
| 174 | + restored = JsonDataConverter().deserialize('{"counter": 42}', ReceiverCounter) |
| 175 | + assert isinstance(restored, ReceiverCounter) |
| 176 | + assert restored.counter == 7 |
| 177 | + |
| 178 | + |
| 179 | +def test_positional_only_receiver_name_can_be_passed_through_kwargs(): |
| 180 | + @dataclass |
| 181 | + class PositionalReceiver: |
| 182 | + self: int = field(init=False, default=0) |
| 183 | + |
| 184 | + def __init__(self, /, **kwargs): |
| 185 | + self.self = kwargs.get("self", 0) |
| 186 | + |
| 187 | + converter = JsonDataConverter() |
| 188 | + encoded = converter.serialize(PositionalReceiver(**{"self": 42})) |
| 189 | + restored = converter.deserialize(encoded, PositionalReceiver) |
| 190 | + assert isinstance(restored, PositionalReceiver) |
| 191 | + assert restored.self == 42 |
| 192 | + |
| 193 | + |
| 194 | +@pytest.mark.parametrize("hashable", [True, False]) |
| 195 | +def test_non_init_field_preserved_by_callable_initializer(hashable): |
| 196 | + @dataclass |
| 197 | + class CallableCounter: |
| 198 | + counter: int = field(init=False, default=0) |
| 199 | + |
| 200 | + class Initializer: |
| 201 | + def __call__(self, counter=0): |
| 202 | + CallableCounter.counter = counter |
| 203 | + |
| 204 | + if not hashable: |
| 205 | + Initializer.__hash__ = None |
| 206 | + CallableCounter.__init__ = Initializer() |
| 207 | + |
| 208 | + converter = JsonDataConverter() |
| 209 | + encoded = converter.serialize(CallableCounter(counter=42)) |
| 210 | + CallableCounter.counter = 0 |
| 211 | + restored = converter.deserialize(encoded, CallableCounter) |
| 212 | + assert isinstance(restored, CallableCounter) |
| 213 | + assert restored.counter == 42 |
| 214 | + |
| 215 | + |
| 216 | +@pytest.mark.parametrize("error_type", [TypeError, ValueError]) |
| 217 | +def test_non_init_fields_retained_when_constructor_signature_unavailable(error_type): |
| 218 | + @dataclass |
| 219 | + class UninspectableCounter: |
| 220 | + counter: int = field(init=False) |
| 221 | + |
| 222 | + def __init__(self, counter=0): |
| 223 | + self.counter = counter |
| 224 | + |
| 225 | + with patch("durabletask.serialization.inspect.signature", side_effect=error_type): |
| 226 | + restored = JsonDataConverter().deserialize('{"counter": 42}', UninspectableCounter) |
| 227 | + assert isinstance(restored, UninspectableCounter) |
| 228 | + assert restored.counter == 42 |
| 229 | + |
| 230 | + |
| 231 | +def test_dataclass_constructor_typeerror_is_not_retried(): |
| 232 | + calls = [] |
| 233 | + |
| 234 | + @dataclass |
| 235 | + class FailingCounter: |
| 236 | + counter: int = field(init=False) |
| 237 | + |
| 238 | + def __init__(self, counter=0): |
| 239 | + calls.append(counter) |
| 240 | + raise TypeError("failure inside user constructor") |
| 241 | + |
| 242 | + restored = JsonDataConverter().deserialize('{"counter": 42}', FailingCounter) |
| 243 | + assert restored == {"counter": 42} |
| 244 | + assert calls == [42] |
| 245 | + |
| 246 | + |
| 247 | +def test_ordinary_dataclass_skips_constructor_signature_inspection(): |
| 248 | + with patch("durabletask.serialization.inspect.signature", side_effect=AssertionError("unexpected inspection")): |
| 249 | + restored = JsonDataConverter().deserialize('{"item": "book", "quantity": 42}', Order) |
| 250 | + assert restored == Order("book", 42) |
| 251 | + |
| 252 | + |
| 253 | +def test_constructor_signature_cached_by_initializer(): |
| 254 | + @dataclass |
| 255 | + class SharedInitializer: |
| 256 | + counter: int = field(init=False) |
| 257 | + |
| 258 | + def __init__(self, counter=0): |
| 259 | + self.counter = counter |
| 260 | + |
| 261 | + @dataclass(init=False) |
| 262 | + class InheritedInitializer(SharedInitializer): |
| 263 | + pass |
| 264 | + |
| 265 | + converter = JsonDataConverter() |
| 266 | + with patch("durabletask.serialization.inspect.signature", wraps=inspect.signature) as signature: |
| 267 | + for cls in (SharedInitializer, InheritedInitializer, SharedInitializer): |
| 268 | + restored = converter.deserialize('{"counter": 42}', cls) |
| 269 | + assert isinstance(restored, cls) |
| 270 | + assert restored.counter == 42 |
| 271 | + signature.assert_called_once_with(SharedInitializer.__init__) |
| 272 | + |
| 273 | + |
60 | 274 | # ----- JsonDataConverter ----- |
61 | 275 |
|
62 | 276 |
|
|
0 commit comments