- Added new metaclass for enum, so that the construct instance is saved inside the Enum instance.

- renamed TEnum to EnumConstruct
- renamed TFlagsEnum to FlagsEnumConstruct
This commit is contained in:
Tim Rid 2022-01-15 23:10:56 +01:00
parent f5b45792c2
commit c943a92fad
3 changed files with 88 additions and 25 deletions

View file

@ -25,7 +25,7 @@ from .generics import (
Constructable,
construct
)
from .tenum import EnumBase, FlagsEnumBase, TEnum, TFlagsEnum
from .tenum import EnumBase, FlagsEnumBase, EnumConstruct, FlagsEnumConstruct
__all__ = [
"AttrsStruct",
@ -45,8 +45,8 @@ __all__ = [
"sfield",
"EnumBase",
"FlagsEnumBase",
"TEnum",
"TFlagsEnum",
"EnumConstruct",
"FlagsEnumConstruct",
"Adapter",
"ConstantOrContextLambda",
"Construct",

View file

@ -1,13 +1,66 @@
import enum
import typing as t
import construct as cs
from .generics import *
T = t.TypeVar("T")
# ## TEnum ############################################################################################################
class EnumBase(enum.IntEnum):
class ConstructEnumMeta(enum.EnumMeta):
@classmethod
def __prepare__(
metacls, # type: ignore
name: str,
bases: t.Tuple[type, ...],
**kwargs: t.Any,
) -> t.Mapping[str, object]:
# This method is needed, because the original __prepare__ method does not accept kwargs.
return super().__prepare__(name, bases)
def __new__(
metacls: t.Type[T], # type: ignore
name: str,
bases: t.Tuple[type, ...],
namespace: t.Dict[str, t.Any],
**kwargs: t.Any,
) -> T:
# create new enum object
cls = super().__new__(metacls, name, bases, namespace) # type: ignore
# if the `EnumBase` class is created, there are no parameters
if len(kwargs) == 0:
return cls
# extract parameters from kwargs
subcon: "cs.Construct[t.Any, t.Any]" = kwargs.pop("subcon", None)
if not isinstance(subcon, cs.Construct): # type: ignore
raise ValueError(
f"`subcon` parameter has to be an `Construct` object but is {type(subcon)}"
)
if len(kwargs) > 0: # check remaining parameters
unsupp_parm = ", ".join([f"'{k}'" for k in kwargs.keys()])
raise ValueError(f"unsupported parameter(s) detected: {unsupp_parm}")
# create construct format
if EnumBase in bases:
enum_constr = EnumConstruct(subcon, cls) # type: ignore
elif FlagsEnumBase in bases:
enum_constr = FlagsEnumConstruct(subcon, cls) # type: ignore
else:
enum_constr = None
# save construct format and make the class compatible to `Constructable` protocol
setattr(cls, "__construct__", lambda: enum_constr) # type: ignore
return cls
# ## EnumConstruct ############################################################################################################
class EnumBase(enum.IntEnum, metaclass=ConstructEnumMeta):
"""
Base class for an Enum used in `construct_typed.TEnum`.
Base class for an Enum used in `construct_typed.EnumConstruct`.
This class extends the standard `enum.IntEnum`, so that missing values are automatically generated.
"""
@ -33,13 +86,17 @@ class EnumBase(enum.IntEnum):
pseudo_member = cls._value2member_map_.setdefault(value, new_member) # type: ignore
return pseudo_member # type: ignore
# TODO: Add `__construct__` method to support `Constructable` protocol
if t.TYPE_CHECKING:
@classmethod
def __construct__(cls: "t.Type[EnumType]") -> "EnumConstruct[EnumType]":
...
EnumType = t.TypeVar("EnumType", bound=EnumBase)
class TEnum(Adapter[int, int, EnumType, EnumType]):
class EnumConstruct(Adapter[int, int, EnumType, EnumType]):
"""
Typed enum.
"""
@ -48,7 +105,7 @@ class TEnum(Adapter[int, int, EnumType, EnumType]):
def __new__(
cls, subcon: Construct[int, int], enum_type: t.Type[EnumType]
) -> "TEnum[EnumType]":
) -> "EnumConstruct[EnumType]":
...
def __init__(self, subcon: Construct[int, int], enum_type: t.Type[EnumType]):
@ -61,7 +118,7 @@ class TEnum(Adapter[int, int, EnumType, EnumType]):
self.enum_type = t.cast(t.Type[EnumType], enum_type) # type: ignore
# init adatper
super(TEnum, self).__init__(subcon) # type: ignore
super(EnumConstruct, self).__init__(subcon) # type: ignore
def _decode(self, obj: int, context: Context, path: PathType) -> EnumType:
return self.enum_type(obj)
@ -79,15 +136,21 @@ class TEnum(Adapter[int, int, EnumType, EnumType]):
)
# ## TFlagsEnum #######################################################################################################
class FlagsEnumBase(enum.IntFlag):
pass
# ## FlagsEnumConstruct #######################################################################################################
class FlagsEnumBase(enum.IntFlag, metaclass=ConstructEnumMeta):
if t.TYPE_CHECKING:
@classmethod
def __construct__(
cls: "t.Type[FlagsEnumType]",
) -> "FlagsEnumConstruct[FlagsEnumType]":
...
FlagsEnumType = t.TypeVar("FlagsEnumType", bound=FlagsEnumBase)
class TFlagsEnum(Adapter[int, int, FlagsEnumType, FlagsEnumType]):
class FlagsEnumConstruct(Adapter[int, int, FlagsEnumType, FlagsEnumType]):
"""
Typed enum.
"""
@ -96,7 +159,7 @@ class TFlagsEnum(Adapter[int, int, FlagsEnumType, FlagsEnumType]):
def __new__(
cls, subcon: Construct[int, int], enum_type: t.Type[FlagsEnumType]
) -> "TFlagsEnum[FlagsEnumType]":
) -> "FlagsEnumConstruct[FlagsEnumType]":
...
def __init__(self, subcon: Construct[int, int], enum_type: t.Type[FlagsEnumType]):
@ -109,7 +172,7 @@ class TFlagsEnum(Adapter[int, int, FlagsEnumType, FlagsEnumType]):
self.enum_type = t.cast(t.Type[FlagsEnumType], enum_type) # type: ignore
# init adatper
super(TFlagsEnum, self).__init__(subcon) # type: ignore
super(FlagsEnumConstruct, self).__init__(subcon) # type: ignore
def _decode(self, obj: int, context: Context, path: PathType) -> FlagsEnumType:
return self.enum_type(obj)

View file

@ -491,13 +491,13 @@ def test_dataclass_bitstruct() -> None:
def test_tenum() -> None:
class TestEnum(cst.EnumBase):
class TestEnum(cst.EnumBase, subcon=cs.Byte):
one = 1
two = 2
four = 4
eight = 8
d = cst.TEnum(cs.Byte, TestEnum)
d = cst.construct(TestEnum)
common(d, b"\x01", TestEnum.one, 1)
common(d, b"\xff", TestEnum(255), 1)
@ -516,7 +516,7 @@ def test_tenum_no_enumbase() -> None:
b = 2
cls = t.cast(t.Type[cst.EnumBase], E)
assert raises(lambda: cst.TEnum(cs.Byte, cls)) == TypeError
assert raises(lambda: cst.EnumConstruct(cs.Byte, cls)) == TypeError
def test_dataclass_struct_wrong_enumbase() -> None:
@ -528,17 +528,17 @@ def test_dataclass_struct_wrong_enumbase() -> None:
a = 1
b = 2
assert raises(cst.TEnum(cs.Byte, E1).build, E2.a) == TypeError
assert raises(cst.EnumConstruct(cs.Byte, E1).build, E2.a) == TypeError
def test_tenum_in_tstruct() -> None:
class TestEnum(cst.EnumBase):
class TestEnum(cst.EnumBase, subcon=cs.Int8ub):
a = 1
b = 2
@dataclasses.dataclass
class TestContainer(DataclassMixin):
a: TestEnum = csfield(cst.TEnum(cs.Int8ub, TestEnum))
a: TestEnum = csfield(cst.construct(TestEnum))
b: int = csfield(cs.Int8ub)
common(
@ -549,18 +549,18 @@ def test_tenum_in_tstruct() -> None:
)
assert (
raises(cst.TEnum(cs.Byte, TestEnum).build, TestContainer(a=1, b=2)) == TypeError # type: ignore
raises(cst.construct(TestEnum).build, TestContainer(a=1, b=2)) == TypeError # type: ignore
)
def test_tenum_flags() -> None:
class TestEnum(cst.FlagsEnumBase):
class TestEnum(cst.FlagsEnumBase, subcon=cs.Byte):
one = 1
two = 2
four = 4
eight = 8
d = cst.TFlagsEnum(cs.Byte, TestEnum)
d = cst.construct(TestEnum)
common(d, b"\x03", TestEnum.one | TestEnum.two, 1)
assert d.build(TestEnum(0)) == b"\x00"
assert d.build(TestEnum.one | TestEnum.two) == b"\x03"