From ac8800943aa9faa0fa41e4ecd6353182e205c556 Mon Sep 17 00:00:00 2001 From: Tim Rid <6593626+timrid@users.noreply.github.com> Date: Sat, 2 Jan 2021 12:01:20 +0100 Subject: [PATCH] Added "TFlagsEnum" and "FlagsEnumBase" --- construct_typed/__init__.py | 4 ++- construct_typed/tenum.py | 43 +++++++++++++++++++++++++-- tests/test_typed.py | 58 ++++++++++++++++++++++--------------- 3 files changed, 78 insertions(+), 27 deletions(-) diff --git a/construct_typed/__init__.py b/construct_typed/__init__.py index a749902..4de653b 100644 --- a/construct_typed/__init__.py +++ b/construct_typed/__init__.py @@ -7,7 +7,7 @@ from .generic_wrapper import ( PathType, ) from .tarray import TArray -from .tenum import TEnum, EnumBase +from .tenum import TEnum, EnumBase, TFlagsEnum, FlagsEnumBase from .tstruct import TBitStruct, TStruct, TStructField from .tunion import TUnion, TUnionField @@ -26,4 +26,6 @@ __all__ = [ "Context", "ConstantOrContextLambda", "PathType", + "TFlagsEnum", + "FlagsEnumBase" ] diff --git a/construct_typed/tenum.py b/construct_typed/tenum.py index c713b46..5cf2622 100644 --- a/construct_typed/tenum.py +++ b/construct_typed/tenum.py @@ -4,6 +4,7 @@ import typing as t from .generic_wrapper import * +# ## TEnum ############################################################################################################ class EnumBase(enum.IntEnum): """ Base class for an Enum used in `construct_typed.TEnum`. @@ -58,8 +59,46 @@ class TEnum(Adapter[int, int, EnumType, EnumType]): def _encode( self, obj: EnumType, - context: "cs.Context", - path: "cs.PathType", + context: Context, + path: PathType, + ) -> int: + if isinstance(obj, self.enum_type): + return int(obj) + raise TypeError("'{}' has to be of type {}".format(repr(obj), repr(self.enum_type))) + + + +# ## TFlagsEnum ####################################################################################################### +class FlagsEnumBase(enum.IntFlag): + pass + + +FlagsEnumType = t.TypeVar("FlagsEnumType", bound=FlagsEnumBase) + +class TFlagsEnum(Adapter[int, int, FlagsEnumType, FlagsEnumType]): + """ + Typed enum. + """ + def __init__(self, subcon: Construct[int, int], enum_type: t.Type[FlagsEnumType]): + if not issubclass(enum_type, FlagsEnumBase): + raise TypeError( + "'{}' has to be a '{}'".format(repr(enum_type), repr(FlagsEnumBase)) + ) + + # save enum type + self.enum_type = enum_type + + # init adatper + super(TFlagsEnum, self).__init__(subcon) + + def _decode(self, obj: int, context: Context, path: PathType) -> FlagsEnumType: + return self.enum_type(obj) # type: ignore + + def _encode( + self, + obj: FlagsEnumType, + context: Context, + path: PathType, ) -> int: if isinstance(obj, self.enum_type): return int(obj) diff --git a/tests/test_typed.py b/tests/test_typed.py index d242bd7..2adbd52 100644 --- a/tests/test_typed.py +++ b/tests/test_typed.py @@ -151,29 +151,23 @@ def test_tbitstruct() -> None: def test_tenum() -> None: - class E(cst.EnumBase): - a = 1 - b = 2 + class TestEnum(cst.EnumBase): + one = 1 + two = 2 + four = 4 + eight = 8 - common(cst.TEnum(cs.Byte, E), b"\x01", E.a, 1) - common(cst.TEnum(cs.Byte, E), b"\x02", E.b, 1) - common(cst.TEnum(cs.Byte, E), b"\x03", E(3), 1) - common(cst.TEnum(cs.Byte, E), b"\xff", E(255), 1) - - format = cst.TEnum(cs.Byte, E) - obj = format.parse(b"\x01") - assert obj == E.a - assert obj == 1 - obj = format.parse(b"\x02") - assert obj == E.b - assert obj == 2 - obj = format.parse(b"\x03") - assert obj == E(3) - assert obj == 3 - obj = format.parse(b"\xff") - assert obj == E(255) - assert obj == 255 + d = cst.TEnum(cs.Byte, TestEnum) + common(d, b"\x01", TestEnum.one, 1) + common(d, b"\xff", TestEnum(255), 1) + assert d.parse(b"\x01") == TestEnum.one + assert d.parse(b"\x01") == 1 + assert int(d.parse(b"\x01")) == 1 + assert d.parse(b"\xff") == TestEnum(255) + assert d.parse(b"\xff") == 255 + assert int(d.parse(b"\xff")) == 255 + assert raises(d.build, 8) == TypeError def test_tenum_no_enumbase() -> None: @@ -193,9 +187,7 @@ def test_tstruct_wrong_enumbase() -> None: a = 1 b = 2 - assert ( - raises(cst.TEnum(cs.Byte, E1).build, E2.a) == TypeError - ) + assert raises(cst.TEnum(cs.Byte, E1).build, E2.a) == TypeError def test_tenum_in_tstruct() -> None: @@ -218,3 +210,21 @@ def test_tenum_in_tstruct() -> None: assert ( raises(cst.TEnum(cs.Byte, TestEnum).build, TestDataclass(a=1, b=2)) == TypeError # type: ignore ) + + +def test_tenum_flags() -> None: + class TestEnum(cst.FlagsEnumBase): + one = 1 + two = 2 + four = 4 + eight = 8 + + d = cst.TFlagsEnum(cs.Byte, 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" + assert d.build(TestEnum(8)) == b"\x08" + assert d.build(TestEnum(1 | 2)) == b"\x03" + assert d.build(TestEnum(255)) == b"\xff" + assert d.build(TestEnum.eight) == b"\x08" + assert raises(d.build, 2) == TypeError