diff --git a/construct_typed/__init__.py b/construct_typed/__init__.py index 374d2cf..e7d171c 100644 --- a/construct_typed/__init__.py +++ b/construct_typed/__init__.py @@ -1,168 +1,232 @@ from enum import IntEnum -from typing import Any, Type, Dict, TYPE_CHECKING, TypeVar, Union -import typing -from construct.core import Construct, Adapter, Struct, BitStruct -from construct.lib.containers import Container +import typing as t +import textwrap +import construct as cs +import dataclasses + +ParsedType = t.TypeVar("ParsedType") +BuildTypes = t.TypeVar("BuildTypes") +SubconParsedType = t.TypeVar("SubconParsedType") +SubconBuildTypes = t.TypeVar("SubconBuildTypes") +EnumType = t.TypeVar("EnumType", bound=IntEnum) +DataclassType = t.TypeVar("DataclassType") -#=============================================================================== -# mappings -#=============================================================================== -if TYPE_CHECKING: - EnumType = TypeVar("EnumType", bound=IntEnum) - - class TypedEnum(Adapter[EnumType, Union[int, str, IntEnum], int, int]): - def __init__(self, subcon: Construct[int, int], enum_type: Type[EnumType]): ... - -# def wrap_enum(enum_type: Type[EnumType]) -> +if t.TYPE_CHECKING: + # while type checking, the original classes are generics, because they are defined in the stubs. + from construct import Construct, Adapter else: - class TypedEnum(Adapter): - def __init__(self, subcon, enum_type): - super(TypedEnum, self).__init__(subcon) + # at runtime, the original classes are no generics, so whe have to make new classes with generics support + class Construct(t.Generic[ParsedType, BuildTypes], cs.Construct): + pass - @classmethod - def _missing_(cls, value): - if isinstance(value, int): - return cls._create_pseudo_member_(value) - return None # will raise the ValueError in Enum.__new__ + class Adapter( + t.Generic[SubconParsedType, SubconBuildTypes, ParsedType, BuildTypes], + cs.Adapter, + ): + pass - @classmethod - def _create_pseudo_member_(cls, value): - pseudo_member = cls._value2member_map_.get(value, None) - if pseudo_member is None: - new_member = int.__new__(cls, value) - # I expect a name attribute to hold a string, hence str(value) - # However, new_member._name_ = value works, too - new_member._name_ = str(value) - new_member._value_ = value - pseudo_member = cls._value2member_map_.setdefault(value, new_member) - return pseudo_member - # Monkey-patch the enum type with __missing__ method. So if a enum value - # not found in the enum a new pseudo member is created. - # The idea is taken from: https://stackoverflow.com/a/57179436 - enum_type._missing_ = _missing_ - enum_type._create_pseudo_member_ = _create_pseudo_member_ - - self.enum_type = enum_type +# =============================================================================== +# mappings +# =============================================================================== +class TEnum(Adapter[int, int, EnumType, t.Union[int, str, EnumType]]): + def __new__( + cls, subcon: Construct[int, int], enum_type: t.Type[EnumType] + ) -> "TEnum[EnumType]": + return super(TEnum, cls).__new__(cls, subcon, enum_type) # type: ignore - def _decode(self, obj, context, path): + def __init__(self, subcon: Construct[int, int], enum_type: t.Type[EnumType]): + @classmethod + def _missing_( + cls: t.Type[EnumType], value: t.Union[int, EnumType] + ) -> t.Optional[EnumType]: + if isinstance(value, int): + return cls._create_pseudo_member_(value) + return None # will raise the ValueError in Enum.__new__ + + @classmethod + def _create_pseudo_member_(cls: t.Type[EnumType], value: int) -> EnumType: + pseudo_member = cls._value2member_map_.get(value, None) + if pseudo_member is None: + new_member = int.__new__(cls, value) + # I expect a name attribute to hold a string, hence str(value) + # However, new_member._name_ = value works, too + new_member._name_ = str(value) + new_member._value_ = value + pseudo_member = cls._value2member_map_.setdefault(value, new_member) + return pseudo_member + + # Monkey-patch the enum type with __missing__ method. So if a enum value + # not found in the enum a new pseudo member is created. + # The idea is taken from: https://stackoverflow.com/a/57179436 + enum_type._missing_ = _missing_ + enum_type._create_pseudo_member_ = _create_pseudo_member_ + + # save enum type + self.enum_type = enum_type + + # init adatper + super(TEnum, self).__init__(subcon) # type: ignore + + def _decode(self, obj: int, context: "cs.Context", path: "cs.PathType") -> EnumType: + return self.enum_type(obj) + + def _encode( + self, + obj: t.Union[int, str, EnumType], + context: "cs.Context", + path: "cs.PathType", + ) -> int: + try: if isinstance(obj, str): return int(self.enum_type[obj]) else: - return self.enum_type(obj) - - def _encode(self, obj, context, path): - try: - if isinstance(obj, str): - return int(self.enum_type[obj]) - else: - return int(self.enum_type(obj)) - except: - raise MappingError("building failed, no mapping for %r" % (obj,), path=path) - - + return int(self.enum_type(obj)) + except: + raise cs.MappingError( + "building failed, no mapping for %r" % (obj,), path=path + ) # =============================================================================== # structures and sequences # =============================================================================== -ParsedType = TypeVar("ParsedType") -BuildTypes = TypeVar("BuildTypes") +def TSubcon( + subcon: Construct[ParsedType, BuildTypes], + doc: t.Optional[str] = None, + parsed: t.Optional[t.Callable[[t.Any, "cs.Context"], None]] = None, +) -> ParsedType: + """ + Create a dataclass field from a subcon. + """ + # Rename subcon, if doc or parsed are available + if (doc is not None) or (parsed is not None): + if doc is not None: + doc = textwrap.dedent(doc) + subcon = cs.Renamed(subcon, newdocs=doc, newparsed=parsed) -# In an optimal way the code look like this (with python 3.9): -# -# def Subcon(subcon: Construct[ParsedType, BuildTypes]) -> Type[ParsedType]: -# return Annotated[ParsedType, subcon] -# -# But this only works if the type annotations are directly in the source -# code. However, in this case the type annotations are stored in a separate -# stub (.pyi) file, which is not available at runtime. -# -# Therefore a small hack is used here. During the type checking, the type -# of the subcon is determined with the help of the stub files and returned. -# At runtime the actual subcon and not its type is returned, so that the -# "TypedStruct" can use this for creating the struct. -# -# This has the disadvantage that it can cause problems if the type is evaluated -# at runtime... -if TYPE_CHECKING: - def Subcon(subcon: Construct[ParsedType, BuildTypes]) -> Type[ParsedType]: ... -else: - def Subcon(subcon: Construct) -> Construct: - return subcon + if subcon.flagbuildnone is True: + # some subcons have a predefined default value. all other have "None" + default: t.Any = None + if isinstance(subcon, (cs.Const, cs.Default)): + if callable(subcon.value): + raise ValueError("lamda as default is not supported") + default = subcon.value -if TYPE_CHECKING: - class TypedContainer(Container[Any]): - ... - ContainerType = TypeVar("ContainerType", bound=TypedContainer) -else: - class TypedContainer(Container): - pass + # if subcon builds from "None", set default to "None" + field = dataclasses.field( + default=default, + init=False, + metadata={"subcon": cs.Renamed(subcon, newdocs=doc)}, + ) + else: + field = dataclasses.field(metadata={"subcon": subcon}) -if TYPE_CHECKING: - class TypedStruct(Adapter[Container[Any], Dict[str, Any], ContainerType, Dict[str, Any]]): - def __init__(self, container_type: Type[ContainerType], swapped: bool = False) -> None: ... -else: - class TypedStruct(Adapter): - def __init__(self, container_type, swapped = False): - if not issubclass(container_type, TypedContainer): - raise TypeError("the subcon has to be a TypedContainer") - self.container_type = container_type - self.swapped = swapped + return field # type: ignore - # extract the construct formats from the struct_type - subcons = {} - subcon_formats = typing.get_type_hints(container_type) - subcon_items = subcon_formats.items() - if swapped: - subcon_items = reversed(subcon_items) - for subcon_name, subcon_format in subcon_items: - subcons[subcon_name] = subcon_format - # init Adatper with a Struct as subcon - super(TypedStruct, self).__init__(Struct(**subcons)) +class _TStruct( + Adapter["cs.Container[t.Any]", t.Dict[str, t.Any], DataclassType, DataclassType] +): + """ + Base class for a typed struct, based on standard dataclasses. + """ - def _decode(self, obj, context, path): - if self.swapped: - return self.container_type({k: v for k, v in reversed(list(obj.items()))}) - else: - return self.container_type(obj) + def __new__( + cls, dataclass_type: t.Type[DataclassType], swapped: bool = False + ) -> "_TStruct[DataclassType]": + return super(_TStruct, cls).__new__(cls, dataclass_type, swapped) # type: ignore - def _encode(self, obj, context, path): - return obj + def __init__( + self, dataclass_type: t.Type[DataclassType], swapped: bool = False + ) -> None: + if not dataclasses.is_dataclass(dataclass_type): + raise TypeError( + 'the "dataclass_type" has to be a dataclass but is "{}"'.format( + type(dataclass_type).__name__ + ) + ) + self.dataclass_type = dataclass_type + self.swapped = swapped -if TYPE_CHECKING: - class TypedBitStruct(Adapter[Container[Any], Dict[str, Any], ContainerType, Dict[str, Any]]): - def __init__(self, container_type: Type[ContainerType], swapped: bool = False) -> None: ... -else: - class TypedBitStruct(Adapter): - def __init__(self, container_type, swapped = False): - if not issubclass(container_type, TypedContainer): - raise TypeError("the subcon has to be a TypedContainer") - self.container_type = container_type - self.swapped = swapped + # get all fields from the dataclass + fields = dataclasses.fields(self.dataclass_type) + if self.swapped: + fields = tuple(reversed(fields)) - # extract the construct formats from the struct_type - subcons = {} - subcon_formats = typing.get_type_hints(container_type) - subcon_items = subcon_formats.items() - if swapped: - subcon_items = reversed(subcon_items) - for subcon_name, subcon_format in subcon_items: - subcons[subcon_name] = subcon_format + # extract the construct formats from the struct_type + subcon_fields = {} + for field in fields: + subcon_fields[field.name] = field.metadata["subcon"] - # init Adatper with a Struct as subcon - super(TypedBitStruct, self).__init__(BitStruct(**subcons)) + # init adatper + super(_TStruct, self).__init__(self._create_subcon(subcon_fields)) # type: ignore - def _decode(self, obj, context, path): - if self.swapped: - return self.container_type({k: v for k, v in reversed(list(obj.items()))}) - else: - return self.container_type(obj) + def _create_subcon( + self, subcon_fields: t.Dict[str, t.Any] + ) -> Construct[t.Any, t.Any]: + raise NotImplementedError - def _encode(self, obj, context, path): - return obj + def _decode( + self, obj: "cs.Container[t.Any]", context: "cs.Context", path: "cs.PathType" + ) -> DataclassType: + # get all fields from the dataclass + fields = dataclasses.fields(self.dataclass_type) + + # extract all fields from the container, that are used for create the dataclass object + dc_init = {} + for field in fields: + if field.init: + value = getattr(obj, field.name) + dc_init[field.name] = value + + # create object of dataclass + dc = self.dataclass_type(**dc_init) # type: ignore + + # extract all other values from the container, an pass it to the dataclass + for field in fields: + if not field.init: + value = getattr(obj, field.name) + setattr(dc, field.name, value) + + return dc + + def _encode( + self, obj: DataclassType, context: "cs.Context", path: "cs.PathType" + ) -> t.Dict[str, t.Any]: + # get all fields from the dataclass + fields = dataclasses.fields(self.dataclass_type) + + # extract all fields from the container, that are used for create the dataclass object + ret_dict = {} + for field in fields: + value = getattr(obj, field.name) + ret_dict[field.name] = value + + return ret_dict + + +class TStruct(_TStruct[DataclassType]): + """ + Typed struct, based on standard dataclasses. + """ + + def _create_subcon( + self, subcon_fields: t.Dict[str, t.Any] + ) -> Construct[t.Any, t.Any]: + return cs.Struct(**subcon_fields) + + +class TBitStruct(_TStruct[DataclassType]): + """ + Typed bit struct, based on standard dataclasses. + """ + + def _create_subcon( + self, subcon_fields: t.Dict[str, t.Any] + ) -> Construct[t.Any, t.Any]: + return cs.BitStruct(**subcon_fields) # TODO: TypedUnion diff --git a/tests/test_typed.py b/tests/test_typed.py index 1ef4e24..72a4725 100644 --- a/tests/test_typed.py +++ b/tests/test_typed.py @@ -1,6 +1,8 @@ # -*- coding: utf-8 -*- import enum +import dataclasses +import typing as t from .declarativeunittest import common, raises from construct import ( Int8ub, @@ -15,68 +17,74 @@ from construct import ( Computed, this, ) -from construct.lib import Container -from construct_typed import TypedContainer, TypedStruct, TypedEnum, Subcon +from construct_typed import TStruct, TBitStruct, TSubcon, TEnum -def test_typed_struct(): - class Container1(TypedContainer): - a: Subcon(Int16ub) - b: Subcon(Int8ub) +def test_typed_struct_1(): + @dataclasses.dataclass + class TestDataclass: + a: int = TSubcon(Int16ub) + b: int = TSubcon(Int8ub) - common(TypedStruct(Container1), b"\x00\x01\x02", Container1(a=1, b=2), 3) + common(TStruct(TestDataclass), b"\x00\x01\x02", TestDataclass(a=1, b=2), 3) common( - TypedStruct(Container1, swapped=True), + TStruct(TestDataclass, swapped=True), b"\x02\x00\x01", - Container(a=1, b=2), + TestDataclass(a=1, b=2), 3, ) - normal = TypedStruct(Container1) - swapped = TypedStruct(Container1, swapped=True) + normal = TStruct(TestDataclass) + swapped = TStruct(TestDataclass, swapped=True) assert str(normal.parse(b"\x00\x01\x02")) == str(swapped.parse(b"\x02\x00\x01")) - class Container2(TypedContainer): - class InnerContainer(TypedContainer): - b: Subcon(Byte) +def test_typed_struct_2(): + @dataclasses.dataclass + class TestDataclass: + @dataclasses.dataclass + class InnerDataclass: + b: int = TSubcon(Byte) - a: Subcon(TypedStruct(InnerContainer)) - - common(TypedStruct(Container2), b"\x01", Container(a=Container(b=1)), 1) - - # TODO: How to get anonymus subcons? - class Container3(TypedContainer): - anonymus1: Subcon(Const(b"\x00")) - anonymus2: Subcon(Padding(1)) - anonymus3: Subcon(Pass) - anonymus4: Subcon(Terminated) + a: InnerDataclass = TSubcon(TStruct(InnerDataclass)) common( - TypedStruct(Container3), + TStruct(TestDataclass), + b"\x01", + TestDataclass(a=TestDataclass.InnerDataclass(b=1)), + 1, + ) + +def test_typed_struct_3(): + @dataclasses.dataclass + class TestDataclass: + _1: t.Optional[bytes] = TSubcon(Const(b"\x00")) + _2: None = TSubcon(Padding(1)) + _3: None = TSubcon(Pass) + _4: None = TSubcon(Terminated) + + common( + TStruct(TestDataclass), bytes(2), - dict(anonymus1=b"\x00", anonymus2=None, anonymus3=None, anonymus4=None), + TestDataclass(), SizeofError, ) - class Container4(TypedContainer): - missingkey: Subcon(Byte) +def test_typed_struct_4(): + @dataclasses.dataclass + class TestDataclass: + _1: bytes = TSubcon(Bytes(this.missing)) - assert raises(TypedStruct(Container4).build, {}) == KeyError + assert raises(TStruct(TestDataclass).sizeof) == SizeofError - # TODO: How to get anonymus subcons? - class Container5(TypedContainer): - anonymus1: Subcon(Bytes(this.missing)) +def test_typed_struct_5(): + @dataclasses.dataclass + class TestDataclass: + _1: int = TSubcon(Computed(7)) + _2: t.Optional[bytes] = TSubcon(Const(b"JPEG")) + _3: None = TSubcon(Pass) + _4: None = TSubcon(Terminated) - assert raises(TypedStruct(Container5).sizeof) == SizeofError - - # TODO: How to get anonymus subcons? - class Container6(TypedContainer): - anonymus1: Subcon(Computed(7)) - anonymus2: Subcon(Const(b"JPEG")) - anonymus3: Subcon(Pass) - anonymus4: Subcon(Terminated) - - d = TypedStruct(Container6) - assert d.build({}) == d.build({}) + d = TStruct(TestDataclass) + assert d.build(TestDataclass()) == d.build(TestDataclass()) def test_typed_bit_struct(): @@ -88,25 +96,40 @@ def test_enum(): a = 1 b = 2 - common(TypedEnum(Byte, E), b"\x01", E.a, 1) - common(TypedEnum(Byte, E), b"\x01", 1, 1) - format = TypedEnum(Byte, E) + a = TEnum(Byte, E) + + common(TEnum(Byte, E), b"\x01", E.a, 1) + common(TEnum(Byte, E), b"\x01", 1, 1) + format = TEnum(Byte, E) obj = format.parse(b"\x01") assert obj == E.a data = format.build("a") assert data == b"\x01" - common(TypedEnum(Byte, E), b"\x02", E.b, 1) - common(TypedEnum(Byte, E), b"\x02", 2, 1) - format = TypedEnum(Byte, E) + common(TEnum(Byte, E), b"\x02", E.b, 1) + common(TEnum(Byte, E), b"\x02", 2, 1) + format = TEnum(Byte, E) obj = format.parse(b"\x02") assert obj == E.b data = format.build("b") assert data == b"\x02" - common(TypedEnum(Byte, E), b"\x03", 3, 1) - format = TypedEnum(Byte, E) + common(TEnum(Byte, E), b"\x03", 3, 1) + format = TEnum(Byte, E) obj = format.parse(b"\x03") assert int(obj) == 3 data = format.build(3) assert data == b"\x03" + + +def test_enum_in_struct(): + class TestEnum(enum.IntEnum): + a = 1 + b = 2 + + @dataclasses.dataclass + class TestDataclass: + a: TestEnum = TSubcon(TEnum(Int8ub, TestEnum)) + b: int = TSubcon(Int8ub) + + common(TStruct(TestDataclass), b"\x00\x01\x02", TestDataclass(a=TestEnum.a, b=TestEnum.b), 3)