Merge pull request #22 from timrid/bugfix/pickel-enum-by-value

Pickle enums by value instead of name (restores pre-3.11 behavior) to support `dataclasses.asdict`
This commit is contained in:
timrid 2023-05-09 12:13:30 +02:00 committed by GitHub
commit 4d07ed4cc6
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
3 changed files with 69 additions and 2 deletions

View file

@ -27,6 +27,8 @@ class DataclassMixin:
methods exists and every name can be used.
"""
__dataclass_fields__: "t.ClassVar[t.Dict[str, dataclasses.Field[t.Any]]]"
def __getitem__(self, key: str) -> t.Any:
return getattr(self, key)
@ -210,7 +212,7 @@ class DataclassStruct(Adapter[t.Any, t.Any, DataclassType, DataclassType]):
value = obj[field.name]
setattr(dc, field.name, value)
return dc
return dc # type: ignore
def _encode(
self, obj: DataclassType, context: Context, path: PathType
@ -269,4 +271,4 @@ TBitStruct = DataclassBitStruct
TContainerMixin = DataclassMixin
TContainerBase = DataclassMixin
TStructField = csfield
sfield = csfield
sfield = csfield

View file

@ -74,6 +74,13 @@ class EnumBase(enum.IntEnum):
return pseudo_member
return None # will raise the ValueError in Enum.__new__
def __reduce_ex__(self, proto: t.Any) -> t.Tuple[t.Any, ...]:
"""
Pickle enums by value instead of name (restores pre-3.11 behavior).
See https://github.com/python/cpython/pull/26658 for why this exists.
"""
return self.__class__, (self._value_,)
EnumType = t.TypeVar("EnumType", bound=EnumBase)
@ -171,6 +178,13 @@ class FlagsEnumBase(enum.IntFlag):
new_member.__doc__ = "missing value"
return new_member
def __reduce_ex__(self, proto: t.Any) -> t.Tuple[t.Any, ...]:
"""
Pickle enums by value instead of name (restores pre-3.11 behavior).
See https://github.com/python/cpython/pull/26658 for why this exists.
"""
return self.__class__, (self._value_,)
FlagsEnumType = t.TypeVar("FlagsEnumType", bound=FlagsEnumBase)

View file

@ -384,6 +384,32 @@ def test_tenum_no_enumbase() -> None:
assert raises(lambda: cst.TEnum(cs.Byte, cls)) == TypeError
def test_tenum_asdict() -> None:
# see: https://github.com/timrid/construct-typing/issues/21
import construct_typed as cst
import dataclasses
class TestEnum(cst.EnumBase):
one = 1
two = 2
four = 4
eight = 8
@dataclasses.dataclass
class SomeDataclass:
a: TestEnum
dc = SomeDataclass(TestEnum.one)
dc_dict = dataclasses.asdict(dc)
assert dc_dict["a"] == dc.a
assert dc_dict["a"] is dc.a
dc = SomeDataclass(TestEnum(5))
dc_dict = dataclasses.asdict(dc)
assert dc_dict["a"] == dc.a
assert dc_dict["a"] is dc.a
def test_tenum_docstring() -> None:
class TestEnum(cst.EnumBase):
"""
@ -472,6 +498,31 @@ def test_tenum_flags() -> None:
assert raises(d.build, 2) == TypeError
def test_tenum_flags_asdict() -> None:
import construct_typed as cst
import dataclasses
class TestEnum(cst.FlagsEnumBase):
one = 1
two = 2
four = 4
eight = 8
@dataclasses.dataclass
class SomeDataclass:
a: TestEnum
dc = SomeDataclass(TestEnum.one)
dc_dict = dataclasses.asdict(dc)
assert dc_dict["a"] == dc.a
assert dc_dict["a"] is dc.a
dc = SomeDataclass(TestEnum(5))
dc_dict = dataclasses.asdict(dc)
assert dc_dict["a"] == dc.a
assert dc_dict["a"] is dc.a
def test_tenum_flags_docstring() -> None:
class TestEnum(cst.FlagsEnumBase):
"""