Skip to content

Commit e416111

Browse files
fix: preserve non-init state accepted by custom constructors
1 parent 266bd6b commit e416111

3 files changed

Lines changed: 274 additions & 5 deletions

File tree

‎CHANGELOG.md‎

Lines changed: 7 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -10,8 +10,13 @@ adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
1010
FIXED
1111

1212
- Dataclasses with `init=False` fields now reconstruct as their declared type
13-
instead of falling back to a raw dictionary. Derived fields are initialized by
14-
the dataclass constructor rather than passed as unsupported keyword arguments.
13+
instead of falling back to a raw dictionary when their generated constructor
14+
cannot accept those fields. Handwritten constructors continue to receive
15+
serialized non-init fields they accept as keywords, including through
16+
`**kwargs`. Other non-init fields use defaults, default factories, or
17+
`__post_init__` and may reset previously recorded values. Reconstruction must
18+
be deterministic for replay; use an explicit `from_json()` hook to preserve
19+
recorded state that the constructor cannot accept.
1520
- Timer callbacks no longer schedule additional long-timer chunks, retry
1621
activities or sub-orchestrations, or resume orchestrator code after completion,
1722
failure, termination, or continue-as-new. Work scheduled before the terminal

‎durabletask/serialization.py‎

Lines changed: 53 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -144,6 +144,16 @@ class JsonDataConverter(DataConverter):
144144
keeps the core SDK permissive; a stricter, validating converter can be
145145
supplied for callers who want coercion failures to surface as errors.
146146
147+
Dataclass fields marked ``init=False`` are passed to a handwritten
148+
initializer when it accepts them as keywords, including through
149+
``**kwargs``. Otherwise they are omitted: generated initializers use
150+
defaults and default factories, and ``__post_init__`` can recompute derived
151+
values. Previously recorded non-init values may therefore be reset, or
152+
remain unset if the class initializes them externally. Constructors,
153+
default factories, and ``__post_init__`` must behave deterministically for
154+
replay. Define an explicit ``from_json()`` hook to preserve recorded
155+
non-init state that the initializer cannot accept.
156+
147157
> [!NOTE]
148158
> Type-directed reconstruction recurses through dataclass fields,
149159
> ``list``/``Sequence``, ``dict``/``Mapping`` values, ``tuple`` elements,
@@ -494,6 +504,32 @@ def _coerce_generic(value: Any, expected_type: Any, origin: Any,
494504
return value
495505

496506

507+
@functools.lru_cache(maxsize=256)
508+
def _dataclass_init_keywords(initializer: Any) -> tuple[frozenset[str] | None, str | None]:
509+
"""Return keyword names and a potentially reserved receiver name.
510+
511+
A None keyword set represents **kwargs or an unavailable signature.
512+
The caller accounts for receiver binding when using the reserved name.
513+
"""
514+
try:
515+
parameters = list(inspect.signature(initializer).parameters.values())
516+
except (TypeError, ValueError):
517+
# Preserve the legacy keyword-passing behavior when inspection fails.
518+
return None, None
519+
receiver = None
520+
# Positional-only receiver names remain available as keys in **kwargs.
521+
if parameters and parameters[0].kind is inspect.Parameter.POSITIONAL_OR_KEYWORD:
522+
receiver = parameters[0].name
523+
if any(p.kind is inspect.Parameter.VAR_KEYWORD for p in parameters):
524+
return None, receiver
525+
keywords = frozenset(
526+
p.name for p in parameters
527+
if p.kind in (inspect.Parameter.POSITIONAL_OR_KEYWORD,
528+
inspect.Parameter.KEYWORD_ONLY)
529+
)
530+
return keywords, receiver
531+
532+
497533
def _build_dataclass(cls: Any, data: dict[str, Any],
498534
converter: DataConverter | None = None) -> Any:
499535
"""Construct a dataclass from its dict payload, recursing into typed fields."""
@@ -504,10 +540,24 @@ def _build_dataclass(cls: Any, data: dict[str, Any],
504540
globalns = _type_namespace(cls)
505541
kwargs: dict[str, Any] = {}
506542
for field in dataclasses.fields(cls):
507-
# Derived fields are serialized, but are initialized by the dataclass
508-
# itself (for example in __post_init__), not by constructor arguments.
509-
if not field.init or field.name not in data:
543+
if field.name not in data:
510544
continue
545+
if not field.init:
546+
initializer = cls.__init__
547+
try:
548+
init_keywords, receiver = _dataclass_init_keywords(initializer)
549+
except TypeError:
550+
# Unhashable initializer callables cannot use the cache.
551+
init_keywords, receiver = None, None
552+
if init_keywords is not None and field.name not in init_keywords:
553+
continue
554+
if (field.name == receiver
555+
and (inspect.isfunction(initializer) or inspect.ismethoddescriptor(initializer))
556+
and not isinstance(inspect.getattr_static(cls, "__init__"), staticmethod)):
557+
# Normal instance initializers already receive this argument.
558+
# Bound classmethods omit it from their inspected signature;
559+
# staticmethods have no implicit receiver.
560+
continue
511561
# ``get_type_hints`` on Python 3.10 does not deep-resolve forward
512562
# references nested inside container args (e.g. the ``"TreeNode"`` in
513563
# ``list["TreeNode"]`` on a self-referential dataclass), leaving a bare

‎tests/durabletask/test_data_converter.py‎

Lines changed: 214 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,11 +3,16 @@
33

44
"""Tests for the DataConverter abstraction and the default JsonDataConverter."""
55

6+
import inspect
67
import json
78
import logging
89
from dataclasses import dataclass, field
910
from typing import Any
11+
from unittest.mock import patch
1012

13+
import pytest
14+
15+
from durabletask.internal.entity_state_shim import StateShim
1116
from durabletask.serialization import (
1217
DEFAULT_DATA_CONVERTER,
1318
DataConverter,
@@ -57,6 +62,215 @@ def test_round_trip_nested_dataclass_with_derived_field():
5762
assert converter.deserialize(converter.serialize(shipment), Shipment) == shipment
5863

5964

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+
60274
# ----- JsonDataConverter -----
61275

62276

0 commit comments

Comments
 (0)