Skip to content
Open
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
8 changes: 8 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,14 @@ adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).

FIXED

- Dataclasses with `init=False` fields now reconstruct as their declared type
instead of falling back to a raw dictionary when their generated constructor
cannot accept those fields. Handwritten constructors continue to receive
serialized non-init fields they accept as keywords, including through
`**kwargs`. Other non-init fields use defaults, default factories, or
`__post_init__` and may reset previously recorded values. Reconstruction must
be deterministic for replay; use an explicit `from_json()` hook to preserve
recorded state that the constructor cannot accept.
- Single-instance `purge_orchestration()` requests now explicitly target
orchestrations rather than entities in both synchronous and asynchronous clients.
- Timer callbacks no longer schedule additional long-timer chunks, retry
Expand Down
52 changes: 52 additions & 0 deletions durabletask/serialization.py
Original file line number Diff line number Diff line change
Expand Up @@ -144,6 +144,16 @@ class JsonDataConverter(DataConverter):
keeps the core SDK permissive; a stricter, validating converter can be
supplied for callers who want coercion failures to surface as errors.

Dataclass fields marked ``init=False`` are passed to a handwritten
initializer when it accepts them as keywords, including through
``**kwargs``. Otherwise they are omitted: generated initializers use
defaults and default factories, and ``__post_init__`` can recompute derived
values. Previously recorded non-init values may therefore be reset, or
remain unset if the class initializes them externally. Constructors,
default factories, and ``__post_init__`` must behave deterministically for
replay. Define an explicit ``from_json()`` hook to preserve recorded
non-init state that the initializer cannot accept.

> [!NOTE]
> Type-directed reconstruction recurses through dataclass fields,
> ``list``/``Sequence``, ``dict``/``Mapping`` values, ``tuple`` elements,
Expand Down Expand Up @@ -494,6 +504,32 @@ def _coerce_generic(value: Any, expected_type: Any, origin: Any,
return value


@functools.lru_cache(maxsize=256)
def _dataclass_init_keywords(initializer: Any) -> tuple[frozenset[str] | None, str | None]:
"""Return keyword names and a potentially reserved receiver name.

A None keyword set represents **kwargs or an unavailable signature.
The caller accounts for receiver binding when using the reserved name.
"""
try:
parameters = list(inspect.signature(initializer, follow_wrapped=False).parameters.values())
except (TypeError, ValueError):
# Preserve the legacy keyword-passing behavior when inspection fails.
return None, None
receiver = None
# Positional-only receiver names remain available as keys in **kwargs.
if parameters and parameters[0].kind is inspect.Parameter.POSITIONAL_OR_KEYWORD:
receiver = parameters[0].name
if any(p.kind is inspect.Parameter.VAR_KEYWORD for p in parameters):
return None, receiver
keywords = frozenset(
p.name for p in parameters
if p.kind in (inspect.Parameter.POSITIONAL_OR_KEYWORD,
inspect.Parameter.KEYWORD_ONLY)
)
return keywords, receiver


def _build_dataclass(cls: Any, data: dict[str, Any],
converter: DataConverter | None = None) -> Any:
"""Construct a dataclass from its dict payload, recursing into typed fields."""
Expand All @@ -506,6 +542,22 @@ def _build_dataclass(cls: Any, data: dict[str, Any],
for field in dataclasses.fields(cls):
if field.name not in data:
continue
if not field.init:
initializer = cls.__init__
try:
init_keywords, receiver = _dataclass_init_keywords(initializer)
except TypeError:
# Unhashable initializer callables cannot use the cache.
init_keywords, receiver = None, None
if init_keywords is not None and field.name not in init_keywords:
continue
if (field.name == receiver
and (inspect.isfunction(initializer) or inspect.ismethoddescriptor(initializer))
and not isinstance(inspect.getattr_static(cls, "__init__"), staticmethod)):
# Normal instance initializers already receive this argument.
# Bound classmethods omit it from their inspected signature;
# staticmethods have no implicit receiver.
continue
# ``get_type_hints`` on Python 3.10 does not deep-resolve forward
# references nested inside container args (e.g. the ``"TreeNode"`` in
# ``list["TreeNode"]`` on a self-referential dataclass), leaving a bare
Expand Down
286 changes: 285 additions & 1 deletion tests/durabletask/test_data_converter.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,11 +3,17 @@

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

import inspect
import json
import logging
from dataclasses import dataclass
from dataclasses import dataclass, field
from functools import wraps
from typing import Any
from unittest.mock import patch

import pytest

from durabletask.internal.entity_state_shim import StateShim
from durabletask.serialization import (
DEFAULT_DATA_CONVERTER,
DataConverter,
Expand All @@ -21,6 +27,284 @@ class Order:
quantity: int


@dataclass(frozen=True)
class PricedOrder:
quantity: int
unit_price: int
total: int = field(init=False)

def __post_init__(self):
object.__setattr__(self, "total", self.quantity * self.unit_price)


@dataclass
class Shipment:
order: PricedOrder


def test_round_trip_dataclass_with_derived_field():
converter = JsonDataConverter()
order = PricedOrder(3, 10)
encoded = converter.serialize(order)
assert json.loads(encoded) == {"quantity": 3, "unit_price": 10, "total": 30}
assert converter.deserialize(encoded, PricedOrder) == order


def test_coerce_dataclass_recomputes_derived_field():
converter = JsonDataConverter()
result = converter.coerce({"quantity": 3, "unit_price": 10, "total": 999}, PricedOrder)
assert result == PricedOrder(3, 10)
assert result.total == 30


def test_round_trip_nested_dataclass_with_derived_field():
converter = JsonDataConverter()
shipment = Shipment(PricedOrder(3, 10))
assert converter.deserialize(converter.serialize(shipment), Shipment) == shipment


@dataclass
class CounterState:
name: str
counter: int = field(init=False, default=0)

def __init__(self, name: str, counter: int = 0):
self.name = name
self.counter = counter


@dataclass
class KeywordCounterState:
name: str
counter: int = field(init=False, default=0)

def __init__(self, name: str, *, counter: int = 0):
self.name = name
self.counter = counter


@dataclass
class KwargsCounterState:
name: str
counter: int = field(init=False, default=0)

def __init__(self, name: str, **kwargs):
self.name = name
self.counter = kwargs.get("counter", 0)


@pytest.mark.parametrize("state_type", [CounterState, KeywordCounterState, KwargsCounterState])
def test_non_init_field_preserved_by_custom_constructor(state_type):
converter = JsonDataConverter()
original = state_type("persisted", counter=42)
encoded = converter.serialize(original)
restored = converter.deserialize(encoded, state_type)
assert restored == original
assert restored.counter == 42

state = StateShim(encoded, converter, is_serialized=True)
restored_state = state.get_state(state_type)
assert restored_state.counter == 42
state.set_state(restored_state)
assert json.loads(state.encode_state()) == {"name": "persisted", "counter": 42}


@pytest.mark.parametrize("accept_kwargs", [False, True])
def test_non_init_field_preserved_by_decorated_initializer(accept_kwargs):
@dataclass
class DecoratedCounterState:
name: str
counter: int = field(init=False, default=0)

generated_init = DecoratedCounterState.__init__
if accept_kwargs:
@wraps(generated_init)
def restore_counter(self, *args, **kwargs):
counter = kwargs.pop("counter", 0)
generated_init(self, *args, **kwargs)
self.counter = counter
else:
@wraps(generated_init)
def restore_counter(self, *args, counter=0, **kwargs):
generated_init(self, *args, **kwargs)
self.counter = counter
DecoratedCounterState.__init__ = restore_counter

converter = JsonDataConverter()
original = DecoratedCounterState("persisted", counter=42)
encoded = converter.serialize(original)
assert converter.deserialize(encoded, DecoratedCounterState).counter == 42

state = StateShim(encoded, converter, is_serialized=True)
restored = state.get_state(DecoratedCounterState)
assert restored.counter == 42
state.set_state(restored)
assert json.loads(state.encode_state()) == {"name": "persisted", "counter": 42}


@pytest.mark.parametrize("accept_kwargs", [False, True])
def test_non_init_custom_constructor_field_is_recursively_coerced(accept_kwargs):
@dataclass
class CustomShipment:
order: Order = field(init=False)

if accept_kwargs:
def initializer(self, **kwargs):
self.order = kwargs["order"]
else:
def initializer(self, *, order):
self.order = order
CustomShipment.__init__ = initializer

converter = JsonDataConverter()
restored = converter.deserialize('{"order": {"item": "book", "quantity": 42}}', CustomShipment)
assert isinstance(restored, CustomShipment)
assert restored.order == Order("book", 42)


def test_non_init_positional_only_constructor_parameter_is_omitted():
@dataclass
class PositionalCounter:
counter: int = field(init=False)

def __init__(self, counter=0, /):
self.counter = counter

restored = JsonDataConverter().deserialize('{"counter": 42}', PositionalCounter)
assert isinstance(restored, PositionalCounter)
assert restored.counter == 0


@pytest.mark.parametrize("method_type", [classmethod, staticmethod])
def test_non_init_field_preserved_by_descriptor_initializer(method_type):
@dataclass
class DescriptorCounter:
counter: int = field(init=False, default=0)

if method_type is classmethod:
def initializer(cls, counter=0):
cls.counter = counter
else:
def initializer(counter=0):
DescriptorCounter.counter = counter
DescriptorCounter.__init__ = method_type(initializer)

converter = JsonDataConverter()
encoded = converter.serialize(DescriptorCounter(counter=42))
DescriptorCounter.counter = 0
restored = converter.deserialize(encoded, DescriptorCounter)
assert isinstance(restored, DescriptorCounter)
assert restored.counter == 42


def test_non_init_field_matching_receiver_name_is_omitted():
@dataclass
class ReceiverCounter:
counter: int = field(init=False)

def __init__(counter):
counter.counter = 7

restored = JsonDataConverter().deserialize('{"counter": 42}', ReceiverCounter)
assert isinstance(restored, ReceiverCounter)
assert restored.counter == 7


def test_positional_only_receiver_name_can_be_passed_through_kwargs():
@dataclass
class PositionalReceiver:
self: int = field(init=False, default=0)

def __init__(self, /, **kwargs):
self.self = kwargs.get("self", 0)

converter = JsonDataConverter()
encoded = converter.serialize(PositionalReceiver(**{"self": 42}))
restored = converter.deserialize(encoded, PositionalReceiver)
assert isinstance(restored, PositionalReceiver)
assert restored.self == 42


@pytest.mark.parametrize("hashable", [True, False])
def test_non_init_field_preserved_by_callable_initializer(hashable):
@dataclass
class CallableCounter:
counter: int = field(init=False, default=0)

class Initializer:
def __call__(self, counter=0):
CallableCounter.counter = counter

if not hashable:
Initializer.__hash__ = None
CallableCounter.__init__ = Initializer()

converter = JsonDataConverter()
encoded = converter.serialize(CallableCounter(counter=42))
CallableCounter.counter = 0
restored = converter.deserialize(encoded, CallableCounter)
assert isinstance(restored, CallableCounter)
assert restored.counter == 42


@pytest.mark.parametrize("error_type", [TypeError, ValueError])
def test_non_init_fields_retained_when_constructor_signature_unavailable(error_type):
@dataclass
class UninspectableCounter:
counter: int = field(init=False)

def __init__(self, counter=0):
self.counter = counter

with patch("durabletask.serialization.inspect.signature", side_effect=error_type):
restored = JsonDataConverter().deserialize('{"counter": 42}', UninspectableCounter)
assert isinstance(restored, UninspectableCounter)
assert restored.counter == 42


def test_dataclass_constructor_typeerror_is_not_retried():
calls = []

@dataclass
class FailingCounter:
counter: int = field(init=False)

def __init__(self, counter=0):
calls.append(counter)
raise TypeError("failure inside user constructor")

restored = JsonDataConverter().deserialize('{"counter": 42}', FailingCounter)
assert restored == {"counter": 42}
assert calls == [42]


def test_ordinary_dataclass_skips_constructor_signature_inspection():
with patch("durabletask.serialization.inspect.signature", side_effect=AssertionError("unexpected inspection")):
restored = JsonDataConverter().deserialize('{"item": "book", "quantity": 42}', Order)
assert restored == Order("book", 42)


def test_constructor_signature_cached_by_initializer():
@dataclass
class SharedInitializer:
counter: int = field(init=False)

def __init__(self, counter=0):
self.counter = counter

@dataclass(init=False)
class InheritedInitializer(SharedInitializer):
pass

converter = JsonDataConverter()
with patch("durabletask.serialization.inspect.signature", wraps=inspect.signature) as signature:
for cls in (SharedInitializer, InheritedInitializer, SharedInitializer):
restored = converter.deserialize('{"counter": 42}', cls)
assert isinstance(restored, cls)
assert restored.counter == 42
signature.assert_called_once_with(SharedInitializer.__init__, follow_wrapped=False)


# ----- JsonDataConverter -----


Expand Down