From db9b35a0c2fd3840ef9f9ac297d70e7e748a1812 Mon Sep 17 00:00:00 2001 From: Tim Rid <6593626+timrid@users.noreply.github.com> Date: Sun, 13 Feb 2022 16:04:27 +0100 Subject: [PATCH 01/24] added `Constructable` Protocol --- construct_typed/generic_wrapper.py | 17 +++++++++++++++++ 1 file changed, 17 insertions(+) diff --git a/construct_typed/generic_wrapper.py b/construct_typed/generic_wrapper.py index aa4af4c..75c0ac2 100644 --- a/construct_typed/generic_wrapper.py +++ b/construct_typed/generic_wrapper.py @@ -39,3 +39,20 @@ else: ConstantOrContextLambda = t.Union[ValueType, t.Callable[[Context], t.Any]] PathType = str + + +@t.runtime_checkable +class Constructable(t.Protocol[ParsedType, BuildTypes]): + def __construct__(self) -> "Construct[ParsedType, BuildTypes]": + raise NotImplementedError + + +def construct( + constr: t.Union[ + Constructable[ParsedType, BuildTypes], "Construct[ParsedType, BuildTypes]" + ], +) -> Construct[ParsedType, BuildTypes]: + """Get construct instance of `Constructable` or `Construct`""" + if isinstance(constr, Constructable): + constr = constr.__construct__() + return constr From 753e4282ee48259d82c0944a8a728abb33412802 Mon Sep 17 00:00:00 2001 From: Tim Rid <6593626+timrid@users.noreply.github.com> Date: Sun, 13 Feb 2022 16:13:11 +0100 Subject: [PATCH 02/24] renamed generic_wrapper.py to generic.py and removed relative imports --- construct_typed/__init__.py | 6 +++--- construct_typed/dataclass_struct.py | 4 ++-- construct_typed/{generic_wrapper.py => generic.py} | 0 construct_typed/tenum.py | 2 +- construct_typed/version.py | 2 +- 5 files changed, 7 insertions(+), 7 deletions(-) rename construct_typed/{generic_wrapper.py => generic.py} (100%) diff --git a/construct_typed/__init__.py b/construct_typed/__init__.py index 00d5093..a11e26d 100644 --- a/construct_typed/__init__.py +++ b/construct_typed/__init__.py @@ -1,4 +1,4 @@ -from .dataclass_struct import ( +from construct_typed.dataclass_struct import ( DataclassBitStruct, DataclassMixin, DataclassStruct, @@ -10,7 +10,7 @@ from .dataclass_struct import ( csfield, sfield, ) -from .generic_wrapper import ( +from construct_typed.generic import ( Adapter, ConstantOrContextLambda, Construct, @@ -18,7 +18,7 @@ from .generic_wrapper import ( ListContainer, PathType, ) -from .tenum import EnumBase, FlagsEnumBase, TEnum, TFlagsEnum +from construct_typed.tenum import EnumBase, FlagsEnumBase, TEnum, TFlagsEnum __all__ = [ "DataclassBitStruct", diff --git a/construct_typed/dataclass_struct.py b/construct_typed/dataclass_struct.py index 78d2e59..5980a82 100644 --- a/construct_typed/dataclass_struct.py +++ b/construct_typed/dataclass_struct.py @@ -12,7 +12,7 @@ from construct.lib.containers import ( ) from construct.lib.py3compat import bytestringtype, reprstring, unicodestringtype -from .generic_wrapper import Adapter, Construct, Context, ParsedType, PathType +from construct_typed.generic import Adapter, Construct, Context, ParsedType, PathType class DataclassMixin: @@ -269,4 +269,4 @@ TBitStruct = DataclassBitStruct TContainerMixin = DataclassMixin TContainerBase = DataclassMixin TStructField = csfield -sfield = csfield \ No newline at end of file +sfield = csfield diff --git a/construct_typed/generic_wrapper.py b/construct_typed/generic.py similarity index 100% rename from construct_typed/generic_wrapper.py rename to construct_typed/generic.py diff --git a/construct_typed/tenum.py b/construct_typed/tenum.py index 7c93e33..91c4173 100644 --- a/construct_typed/tenum.py +++ b/construct_typed/tenum.py @@ -1,7 +1,7 @@ import enum import typing as t -from .generic_wrapper import * +from construct_typed.generic import * # ## TEnum ############################################################################################################ diff --git a/construct_typed/version.py b/construct_typed/version.py index 1a555cf..9498baa 100644 --- a/construct_typed/version.py +++ b/construct_typed/version.py @@ -1,2 +1,2 @@ version = (0, 5, 2) -version_string = "0.5.2" \ No newline at end of file +version_string = "0.5.2" From b896457f907dd4600402e73d938cd37a9eb96fcb Mon Sep 17 00:00:00 2001 From: Tim Rid <6593626+timrid@users.noreply.github.com> Date: Sun, 13 Feb 2022 16:32:30 +0100 Subject: [PATCH 03/24] Changed implementation of DataclassStruct. It is now only nessesary to sublcass DataclassStruct and not to combine it with @dataclasses.dataclass. Also now the DataclassConstruct is included in the DataclassStruct class type itself. --- construct_typed/dataclass_struct.py | 437 ++++++++++++++++------------ 1 file changed, 253 insertions(+), 184 deletions(-) diff --git a/construct_typed/dataclass_struct.py b/construct_typed/dataclass_struct.py index 5980a82..23a00eb 100644 --- a/construct_typed/dataclass_struct.py +++ b/construct_typed/dataclass_struct.py @@ -14,19 +14,247 @@ from construct.lib.py3compat import bytestringtype, reprstring, unicodestringtyp from construct_typed.generic import Adapter, Construct, Context, ParsedType, PathType +T = t.TypeVar("T") -class DataclassMixin: + +# Static type inference support via __dataclass_transform__ implemented as per: +# https://github.com/microsoft/pyright/blob/1.1.135/specs/dataclass_transforms.md +def __dataclass_transform__( + *, + eq_default: bool = True, + order_default: bool = False, + kw_only_default: bool = False, + field_descriptors: t.Tuple[t.Union[type, t.Callable[..., t.Any]], ...] = (()), +) -> t.Callable[[T], T]: + return lambda a: a + + +DATACLASS_METADATA_KEY = "__construct_typed_subcon" + +if t.TYPE_CHECKING: + # specialisation for constructs, that builds from none and dont have to be declared in the __init__ method + @t.overload + def csfield( + subcon: cs.Construct[ParsedType, None], + doc: t.Optional[str] = None, + parsed: t.Optional[t.Callable[[t.Any, Context], None]] = None, + init: t.Literal[False] = False, + ) -> ParsedType: + ... + + @t.overload + def csfield( + subcon: Construct[ParsedType, t.Any], + doc: t.Optional[str] = None, + parsed: t.Optional[t.Callable[[t.Any, Context], None]] = None, + init: bool = True, + ) -> ParsedType: + ... + + +def csfield( + subcon: Construct[ParsedType, t.Any], + doc: t.Optional[str] = None, + parsed: t.Optional[t.Callable[[t.Any, Context], None]] = None, + init: bool = True, +) -> ParsedType: """ - Mixin for the dataclasses which are passed to "DataclassStruct" and "DataclassBitStruct". + Helper method for "DataclassStruct" and "DataclassBitStruct" to create the dataclass fields. - Note: This implementation is different to the 'cs.Container' of the original 'construct' - library. In the original 'cs.Container' some names like "update", "keys", "items", ... can - only accessed via key access (square brackets) and not via attribute access (dot operator), - because they are also method names. This implementation is based on "dataclasses.dataclass" - which only uses modul-level instead of instance-level helper methods.So no instance-level - methods exists and every name can be used. + This method also processes Const and Default, to pass these values als default values to the dataclass. + """ + orig_subcon = 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).strip("\n") + subcon = cs.Renamed(subcon, newdocs=doc, newparsed=parsed) + + if orig_subcon.flagbuildnone is True: + init = False + default = None + else: + init = True + default = dataclasses.MISSING + + # Set default values in case of special sucons + if isinstance(orig_subcon, cs.Const): + const_subcon: "cs.Const[t.Any, t.Any, t.Any, t.Any]" = orig_subcon + default = const_subcon.value + elif isinstance(orig_subcon, cs.Default): + default_subcon: "cs.Default[t.Any, t.Any, t.Any, t.Any]" = orig_subcon + if callable(default_subcon.value): + default = None # context lambda is only defined at parsing/building + else: + default = default_subcon.value + + return t.cast( + ParsedType, + dataclasses.field( + default=default, + init=init, + metadata={DATACLASS_METADATA_KEY: subcon}, + ), + ) + + +class DataclassConstruct(Adapter[t.Any, t.Any, T, T]): + r""" + TODO: Add Documentation + """ + subcon: "cs.Struct[t.Any, t.Any]" + if t.TYPE_CHECKING: + + def __new__( + cls, + dc_type: t.Type[T], + reverse: bool = False, + ) -> "DataclassConstruct[T]": + ... + + def __init__( + self, + dc_type: t.Type[T], + reverse: bool = False, + ) -> None: + if not isinstance(dc_type, DataclassStruct): + raise TypeError(f"'{repr(dc_type)}' has to be a 'DataclassStruct'") + if not dataclasses.is_dataclass(dc_type): + raise TypeError(f"'{repr(dc_type)}' has to be a 'dataclasses.dataclass'") + self.dc_type = dc_type + self.reverse = reverse + + # get all fields from the dataclass + fields = dataclasses.fields(self.dc_type) + if self.reverse: + fields = tuple(reversed(fields)) + + # extract the construct formats from the struct_type + subcon_fields = {} + for field in fields: + subcon_fields[field.name] = field.metadata[DATACLASS_METADATA_KEY] + + # init adatper + super().__init__(cs.Struct(**subcon_fields)) # type: ignore + + def __getattr__(self, name: str) -> t.Any: + return getattr(self.subcon, name) + + def _decode( + self, obj: "cs.Container[t.Any]", context: Context, path: PathType + ) -> T: + # get all fields from the dataclass + fields = dataclasses.fields(self.dc_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 = obj[field.name] + dc_init[field.name] = value + + # create object of dataclass + dc = self.dc_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 = obj[field.name] + setattr(dc, field.name, value) + + return dc + + def _encode(self, obj: T, context: Context, path: PathType) -> t.Dict[str, t.Any]: + if not isinstance(obj, self.dc_type): + raise TypeError(f"'{repr(obj)}' has to be of type {repr(self.dc_type)}") + + # get all fields from the dataclass + fields = dataclasses.fields(self.dc_type) + + # extract all fields from the container, that are used for create the dataclass object + ret_dict: t.Dict[str, t.Any] = {} + for field in fields: + value = getattr(obj, field.name) + ret_dict[field.name] = value + + return ret_dict + + +# Helper object for defining the `constr` of a `struct`. Will be replaced with the proper construct, when class is created. +this_struct: Construct[t.Any, t.Any] = Construct() + + +def _replace_this_struct(constr: "Construct[t.Any, t.Any]", replacement: t.Any): + """Recursive search for `this_struct` in all SubConstructs and replace it with AttrsStruct""" + subcon = getattr(constr, "subcon", None) + if subcon is this_struct: + setattr(constr, "subcon", replacement) + elif subcon is not None: + _replace_this_struct(subcon, replacement) + else: + raise ValueError( + "Could not find `this_struct`. Only SubConstructs are supported" + ) + + +@__dataclass_transform__(field_descriptors=(csfield,)) +class DataclassStruct: + """ + Adapter for a dataclasses for optimised type hints / static autocompletion in comparision to the original Struct. + + + Before this construct can be created a dataclasses.dataclass type must be created, which must also derive from DataclassMixin. In this dataclass all fields must be assigned to a construct type using csfield. + + Internally, all fields are converted to a normal Struct, which does the actual parsing/building. + + Parses to a dataclasses.dataclass instance, and builds from such instance. Size is the sum of all subcon sizes, unless any subcon raises SizeofError. + + :param constr: This can be used if the structure is nested inside a Subconstruct. To represent this struct use the constant `this_struct`. + :param reverse: Flag if the fields of the dataclass should be reversed + + Example:: + + >>> from construct import Bytes, Int8ub, this + >>> from construct_typed import DataclassMixin, DataclassStruct, csfield, construct + ... class Image(DataclassStruct): + ... width: int = csfield(Int8ub) + ... height: int = csfield(Int8ub) + ... pixels: bytes = csfield(Bytes(this.height * this.width)) + >>> d = construct(Image) + >>> d.parse(b"\x01\x0212") + Image(width=1, height=2, pixels=b'12') """ + @classmethod + def __init_subclass__( + cls, + constr: "cs.Construct[t.Any, t.Any]" = this_struct, + reverse_fields: bool = False, + ): + # validate types + if not isinstance(constr, cs.Construct): # type: ignore + raise ValueError("`constr` parameter has to be an `Construct` object") + if not isinstance(reverse_fields, bool): # type: ignore + raise ValueError("`reverse_fields` parameter has to be an `bool` object") + + # create dataclass + cls = dataclasses.dataclass(cls) + + # create construct format + dc_constr = DataclassConstruct(cls, reverse_fields) + if constr is this_struct: + constr = dc_constr + else: + _replace_this_struct(constr, dc_constr) + + # save construct format and make the class compatible to `Constructable` protocol + setattr(cls, "__construct__", lambda: constr) + + return cls + + # the `construct` library is using the [] access internally, so struct objects + # should also make this possible and not only via the dot access. def __getitem__(self, key: str) -> t.Any: return getattr(self, key) @@ -72,201 +300,42 @@ class DataclassMixin: text.append(indentation.join(str(v).split("\n"))) return "".join(text) - -def csfield( - subcon: Construct[ParsedType, t.Any], - doc: t.Optional[str] = None, - parsed: t.Optional[t.Callable[[t.Any, Context], None]] = None, -) -> ParsedType: - """ - Helper method for "DataclassStruct" and "DataclassBitStruct" to create the dataclass fields. - - This method also processes Const and Default, to pass these values als default values to the dataclass. - """ - orig_subcon = 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).strip("\n") - subcon = cs.Renamed(subcon, newdocs=doc, newparsed=parsed) - - if orig_subcon.flagbuildnone is True: - init = False - default = None - else: - init = True - default = dataclasses.MISSING - - # Set default values in case of special sucons - if isinstance(orig_subcon, cs.Const): - const_subcon: "cs.Const[t.Any, t.Any, t.Any, t.Any]" = orig_subcon - default = const_subcon.value - elif isinstance(orig_subcon, cs.Default): - default_subcon: "cs.Default[t.Any, t.Any, t.Any, t.Any]" = orig_subcon - if callable(default_subcon.value): - default = None # context lambda is only defined at parsing/building - else: - default = default_subcon.value - - return t.cast( - ParsedType, - dataclasses.field( - default=default, - init=init, - metadata={"subcon": subcon}, - ), - ) - - -DataclassType = t.TypeVar("DataclassType", bound=DataclassMixin) - - -class DataclassStruct(Adapter[t.Any, t.Any, DataclassType, DataclassType]): - """ - Adapter for a dataclasses for optimised type hints / static autocompletion in comparision to the original Struct. - - Before this construct can be created a dataclasses.dataclass type must be created, which must also derive from DataclassMixin. In this dataclass all fields must be assigned to a construct type using csfield. - - Internally, all fields are converted to a Struct, which does the actual parsing/building. - - Parses to a dataclasses.dataclass instance, and builds from such instance. Size is the sum of all subcon sizes, unless any subcon raises SizeofError. - - :param dc_type: Type of the dataclass, which also inherits from DataclassMixin - :param reverse: Flag if the fields of the dataclass should be reversed - - Example:: - - >>> import dataclasses - >>> from construct import Bytes, Int8ub, this - >>> from construct_typed import DataclassMixin, DataclassStruct, csfield - >>> @dataclasses.dataclass - ... class Image(DataclassMixin): - ... width: int = csfield(Int8ub) - ... height: int = csfield(Int8ub) - ... pixels: bytes = csfield(Bytes(this.height * this.width)) - >>> d = DataclassStruct(Image) - >>> d.parse(b"\x01\x0212") - Image(width=1, height=2, pixels=b'12') - """ - - subcon: "cs.Struct[t.Any, t.Any]" if t.TYPE_CHECKING: - def __new__( - cls, - dc_type: t.Type[DataclassType], - reverse: bool = False, - ) -> "DataclassStruct[DataclassType]": + @classmethod + def __construct__(cls: t.Type[T]) -> "DataclassConstruct[T]": ... - def __init__( - self, - dc_type: t.Type[DataclassType], - reverse: bool = False, - ) -> None: - if not issubclass(dc_type, DataclassMixin): - raise TypeError(f"'{repr(dc_type)}' has to be a '{repr(DataclassMixin)}'") - if not dataclasses.is_dataclass(dc_type): - raise TypeError(f"'{repr(dc_type)}' has to be a 'dataclasses.dataclass'") - self.dc_type = dc_type - self.reverse = reverse - # get all fields from the dataclass - fields = dataclasses.fields(self.dc_type) - if self.reverse: - fields = tuple(reversed(fields)) - - # extract the construct formats from the struct_type - subcon_fields = {} - for field in fields: - subcon_fields[field.name] = field.metadata["subcon"] - - # init adatper - super().__init__(cs.Struct(**subcon_fields)) # type: ignore - - def __getattr__(self, name: str) -> t.Any: - return getattr(self.subcon, name) - - def _decode( - self, obj: "cs.Container[t.Any]", context: Context, path: PathType - ) -> DataclassType: - # get all fields from the dataclass - fields = dataclasses.fields(self.dc_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 = obj[field.name] - dc_init[field.name] = value - - # create object of dataclass - dc = self.dc_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 = obj[field.name] - setattr(dc, field.name, value) - - return dc - - def _encode( - self, obj: DataclassType, context: Context, path: PathType - ) -> t.Dict[str, t.Any]: - if not isinstance(obj, self.dc_type): - raise TypeError(f"'{repr(obj)}' has to be of type {repr(self.dc_type)}") - - # get all fields from the dataclass - fields = dataclasses.fields(self.dc_type) - - # extract all fields from the container, that are used for create the dataclass object - ret_dict: t.Dict[str, t.Any] = {} - for field in fields: - value = getattr(obj, field.name) - ret_dict[field.name] = value - - return ret_dict - - -def DataclassBitStruct( - dc_type: t.Type[DataclassType], reverse: bool = False -) -> t.Union[ - "cs.Transformed[DataclassType, DataclassType]", - "cs.Restreamed[DataclassType, DataclassType]", -]: +class DataclassBitStruct(DataclassStruct): r""" Makes a DataclassStruct inside a Bitwise. See :class:`~construct.core.Bitwise` and :class:`~construct_typed.dataclass_struct.DatclassStruct` for semantics and raisable exceptions. - :param dc_type: Type of the dataclass, which also inherits from DataclassMixin - :param reverse: Flag if the fields of the dataclass should be reversed + :param constr: TODO + :param reverse_fields: Flag if the fields of the dataclass should be reversed Example:: - DataclassBitStruct <--> Bitwise(DataclassStruct(...)) - >>> import dataclasses + TODO: >>> from construct import BitsInteger, Flag, Nibble, Padding - >>> from construct_typed import DataclassBitStruct, DataclassMixin, csfield - >>> @dataclasses.dataclass - ... class TestDataclass(DataclassMixin): + >>> from construct_typed import DataclassBitStruct, csfield, construct + ... class TestDataclass(DataclassBitStruct): ... a: int = csfield(Flag) ... b: int = csfield(Nibble) ... c: int = csfield(BitsInteger(10)) ... d: None = csfield(Padding(1)) - >>> d = DataclassBitStruct(TestDataclass) + >>> d = construct(TestDataclass) >>> d.parse(b"\x01\x02") TestDataclass(a=False, b=0, c=129, d=None) """ - return cs.Bitwise(DataclassStruct(dc_type, reverse)) - -# support legacy names -TStruct = DataclassStruct -TBitStruct = DataclassBitStruct -TContainerMixin = DataclassMixin -TContainerBase = DataclassMixin -TStructField = csfield -sfield = csfield + @classmethod + def __init_subclass__( + cls, + constr: "cs.Construct[t.Any, t.Any]" = this_struct, + reverse_fields: bool = False, + ): + cls = DataclassStruct.__init_subclass__.__func__(cls, cs.Bitwise(constr), reverse_fields) # type: ignore + return cls From 8169f0ed318a497f5d9c884d153ced16e056ed9d Mon Sep 17 00:00:00 2001 From: Tim Rid <6593626+timrid@users.noreply.github.com> Date: Sun, 13 Feb 2022 16:51:20 +0100 Subject: [PATCH 04/24] Changed implementation of TEnum and TEnumFlags. --- construct_typed/__init__.py | 22 ++------ construct_typed/tenum.py | 101 ++++++++++++++++++++++++++++-------- 2 files changed, 84 insertions(+), 39 deletions(-) diff --git a/construct_typed/__init__.py b/construct_typed/__init__.py index a11e26d..10e5b12 100644 --- a/construct_typed/__init__.py +++ b/construct_typed/__init__.py @@ -1,14 +1,7 @@ from construct_typed.dataclass_struct import ( DataclassBitStruct, - DataclassMixin, DataclassStruct, - TBitStruct, - TContainerBase, - TContainerMixin, - TStruct, - TStructField, csfield, - sfield, ) from construct_typed.generic import ( Adapter, @@ -18,23 +11,16 @@ from construct_typed.generic import ( ListContainer, PathType, ) -from construct_typed.tenum import EnumBase, FlagsEnumBase, TEnum, TFlagsEnum +from construct_typed.tenum import TEnum, TFlags, TEnumConstruct, TFlagsConstruct __all__ = [ "DataclassBitStruct", - "DataclassMixin", "DataclassStruct", - "TBitStruct", - "TContainerBase", - "TContainerMixin", - "TStruct", - "TStructField", "csfield", - "sfield", - "EnumBase", - "FlagsEnumBase", "TEnum", - "TFlagsEnum", + "TEnumConstruct", + "TFlags", + "TFlagsConstruct", "Adapter", "ConstantOrContextLambda", "Construct", diff --git a/construct_typed/tenum.py b/construct_typed/tenum.py index 91c4173..4a84584 100644 --- a/construct_typed/tenum.py +++ b/construct_typed/tenum.py @@ -1,28 +1,54 @@ import enum import typing as t -from construct_typed.generic import * +import construct as cs + +from .generic import * + +T = t.TypeVar("T") # ## TEnum ############################################################################################################ -class EnumBase(enum.IntEnum): +class TEnum(enum.IntEnum): """ - Base class for an Enum used in `construct_typed.TEnum`. + Base class for an Enum used in `construct_typed.TEnumConstruct`. This class extends the standard `enum.IntEnum`, so that missing values are automatically generated. """ + @classmethod + def __init_subclass__( + cls, + subcon: "cs.Construct[t.Any, t.Any]", + **kwargs: t.Any, + ): + super().__init_subclass__(**kwargs) + + # validate types + if not isinstance(subcon, cs.Construct): # type: ignore + raise ValueError( + f"`subcon` parameter has to be an `Construct` object but is {type(subcon)}" + ) + + # create construct format + enum_constr = TEnumConstruct(subcon, cls) + + # save construct format and make the class compatible to `Constructable` protocol + setattr(cls, "__construct__", lambda: enum_constr) + + return cls + # Extend 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 @classmethod - def _missing_(cls, value: t.Any) -> t.Optional["EnumBase"]: + def _missing_(cls, value: t.Any) -> t.Optional["TEnum"]: 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, value: int) -> "EnumBase": + def _create_pseudo_member_(cls, value: int) -> "TEnum": pseudo_member = cls._value2member_map_.get(value, None) # type: ignore if pseudo_member is None: new_member = int.__new__(cls, value) @@ -33,11 +59,17 @@ class EnumBase(enum.IntEnum): pseudo_member = cls._value2member_map_.setdefault(value, new_member) # type: ignore return pseudo_member # type: ignore + if t.TYPE_CHECKING: -EnumType = t.TypeVar("EnumType", bound=EnumBase) + @classmethod + def __construct__(cls: "t.Type[EnumType]") -> "TEnumConstruct[EnumType]": + ... -class TEnum(Adapter[int, int, EnumType, EnumType]): +EnumType = t.TypeVar("EnumType", bound=TEnum) + + +class TEnumConstruct(Adapter[int, int, EnumType, EnumType]): """ Typed enum. """ @@ -46,20 +78,20 @@ class TEnum(Adapter[int, int, EnumType, EnumType]): def __new__( cls, subcon: Construct[int, int], enum_type: t.Type[EnumType] - ) -> "TEnum[EnumType]": + ) -> "TEnumConstruct[EnumType]": ... def __init__(self, subcon: Construct[int, int], enum_type: t.Type[EnumType]): - if not issubclass(enum_type, EnumBase): + if not issubclass(enum_type, TEnum): raise TypeError( - "'{}' has to be a '{}'".format(repr(enum_type), repr(EnumBase)) + "'{}' has to be a '{}'".format(repr(enum_type), repr(TEnum)) ) # save enum type self.enum_type = t.cast(t.Type[EnumType], enum_type) # type: ignore # init adatper - super(TEnum, self).__init__(subcon) # type: ignore + super(TEnumConstruct, self).__init__(subcon) # type: ignore def _decode(self, obj: int, context: Context, path: PathType) -> EnumType: return self.enum_type(obj) @@ -76,38 +108,65 @@ class TEnum(Adapter[int, int, EnumType, EnumType]): "'{}' has to be of type {}".format(repr(obj), repr(self.enum_type)) ) +# ## TFlags ####################################################################################################### +class TFlags(enum.IntFlag): + @classmethod + def __init_subclass__( + cls, + subcon: "cs.Construct[t.Any, t.Any]", + **kwargs: t.Any, + ): + super().__init_subclass__(**kwargs) -# ## TFlagsEnum ####################################################################################################### -class FlagsEnumBase(enum.IntFlag): - pass + # validate types + if not isinstance(subcon, cs.Construct): # type: ignore + raise ValueError( + f"`subcon` parameter has to be an `Construct` object but is {type(subcon)}" + ) + + # create construct format + enum_constr = TFlagsConstruct(subcon, cls) + + # save construct format and make the class compatible to `Constructable` protocol + setattr(cls, "__construct__", lambda: enum_constr) + + return cls + + if t.TYPE_CHECKING: + + @classmethod + def __construct__( + cls: "t.Type[FlagsEnumType]", + ) -> "TFlagsConstruct[FlagsEnumType]": + ... -FlagsEnumType = t.TypeVar("FlagsEnumType", bound=FlagsEnumBase) +FlagsEnumType = t.TypeVar("FlagsEnumType", bound=TFlags) -class TFlagsEnum(Adapter[int, int, FlagsEnumType, FlagsEnumType]): +class TFlagsConstruct(Adapter[int, int, FlagsEnumType, FlagsEnumType]): """ - Typed enum. + Typed flags. """ if t.TYPE_CHECKING: def __new__( cls, subcon: Construct[int, int], enum_type: t.Type[FlagsEnumType] - ) -> "TFlagsEnum[FlagsEnumType]": + ) -> "TFlagsConstruct[FlagsEnumType]": ... def __init__(self, subcon: Construct[int, int], enum_type: t.Type[FlagsEnumType]): - if not issubclass(enum_type, FlagsEnumBase): + if not issubclass(enum_type, TFlags): raise TypeError( - "'{}' has to be a '{}'".format(repr(enum_type), repr(FlagsEnumBase)) + "'{}' has to be a '{}'".format(repr(enum_type), repr(TFlags)) ) # save enum type self.enum_type = t.cast(t.Type[FlagsEnumType], enum_type) # type: ignore # init adatper - super(TFlagsEnum, self).__init__(subcon) # type: ignore + super(TFlagsConstruct, self).__init__(subcon) # type: ignore def _decode(self, obj: int, context: Context, path: PathType) -> FlagsEnumType: return self.enum_type(obj) From 9a9b4ab95a2df28683294993bbe7992c97609a2a Mon Sep 17 00:00:00 2001 From: Tim Rid <6593626+timrid@users.noreply.github.com> Date: Sun, 13 Feb 2022 16:59:01 +0100 Subject: [PATCH 05/24] fixed type error --- construct_typed/dataclass_struct.py | 4 ++-- tests/test_typed.py | 2 +- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/construct_typed/dataclass_struct.py b/construct_typed/dataclass_struct.py index 23a00eb..f9f9705 100644 --- a/construct_typed/dataclass_struct.py +++ b/construct_typed/dataclass_struct.py @@ -118,8 +118,8 @@ class DataclassConstruct(Adapter[t.Any, t.Any, T, T]): dc_type: t.Type[T], reverse: bool = False, ) -> None: - if not isinstance(dc_type, DataclassStruct): - raise TypeError(f"'{repr(dc_type)}' has to be a 'DataclassStruct'") + if not issubclass(dc_type, DataclassStruct): + raise TypeError(f"'{repr(dc_type)}' has to be a subclass of 'DataclassStruct'") if not dataclasses.is_dataclass(dc_type): raise TypeError(f"'{repr(dc_type)}' has to be a 'dataclasses.dataclass'") self.dc_type = dc_type diff --git a/tests/test_typed.py b/tests/test_typed.py index d74e025..b76ff91 100644 --- a/tests/test_typed.py +++ b/tests/test_typed.py @@ -6,7 +6,7 @@ import typing as t import construct as cs import construct_typed as cst -from construct_typed import DataclassBitStruct, DataclassMixin, DataclassStruct, csfield +from construct_typed import DataclassBitStruct, DataclassStruct, csfield from .declarativeunittest import common, raises, setattrs From de76415ac58bb022243226f64bccb06c75191dc5 Mon Sep 17 00:00:00 2001 From: Tim Rid <6593626+timrid@users.noreply.github.com> Date: Sun, 13 Feb 2022 17:02:46 +0100 Subject: [PATCH 06/24] changed vscode test configuration to use the native testing api instead of the Test-Explorer --- .vscode/launch.json | 6 ------ .vscode/settings.json | 21 +++++++++++---------- 2 files changed, 11 insertions(+), 16 deletions(-) diff --git a/.vscode/launch.json b/.vscode/launch.json index d94289d..8655026 100644 --- a/.vscode/launch.json +++ b/.vscode/launch.json @@ -9,12 +9,6 @@ "type": "python", "request": "launch", "program": "${file}", - "console": "integratedTerminal" - }, - { - "name": "Debug Tests", - "type": "python", - "request": "test", "console": "integratedTerminal", "justMyCode": false } diff --git a/.vscode/settings.json b/.vscode/settings.json index 76bdfe9..b265945 100644 --- a/.vscode/settings.json +++ b/.vscode/settings.json @@ -1,24 +1,25 @@ { - "python.pythonPath": "python", "python.languageServer": "Pylance", - // "python.testing.unittestEnabled": false, - // "python.testing.nosetestsEnabled": false, - // "python.testing.pytestEnabled": true, - "pythonTestExplorer.testFramework": "pytest", + + // configure code formating "python.formatting.provider": "black", "python.sortImports.path": "isort", "python.sortImports.args": [ "--profile=black", ], - // "[python]": { - // "editor.codeActionsOnSave": { - // "source.organizeImports": true - // } - // } + + // configure pylance "python.analysis.typeCheckingMode": "strict", "python.analysis.autoImportCompletions": false, "python.analysis.diagnosticSeverityOverrides": { "reportPrivateUsage": "information", "reportUntypedNamedTuple": "information", }, + + // configure pytest + "python.testing.pytestArgs": [ + "tests" + ], + "python.testing.unittestEnabled": false, + "python.testing.pytestEnabled": true, } \ No newline at end of file From 7c68aeecd486aaaff5cb189273658eb5f32dd6db Mon Sep 17 00:00:00 2001 From: Tim Rid <6593626+timrid@users.noreply.github.com> Date: Sun, 13 Feb 2022 17:19:13 +0100 Subject: [PATCH 07/24] adapted all tests to the new api --- construct_typed/__init__.py | 2 + tests/test_typed.py | 177 ++++++++++++------------------------ 2 files changed, 61 insertions(+), 118 deletions(-) diff --git a/construct_typed/__init__.py b/construct_typed/__init__.py index 10e5b12..86e7d7a 100644 --- a/construct_typed/__init__.py +++ b/construct_typed/__init__.py @@ -1,3 +1,4 @@ +from construct_typed.generic import construct from construct_typed.dataclass_struct import ( DataclassBitStruct, DataclassStruct, @@ -16,6 +17,7 @@ from construct_typed.tenum import TEnum, TFlags, TEnumConstruct, TFlagsConstruct __all__ = [ "DataclassBitStruct", "DataclassStruct", + "construct", "csfield", "TEnum", "TEnumConstruct", diff --git a/tests/test_typed.py b/tests/test_typed.py index b76ff91..b69e50c 100644 --- a/tests/test_typed.py +++ b/tests/test_typed.py @@ -1,19 +1,22 @@ # -*- coding: utf-8 -*- # pyright: strict -import dataclasses -import enum import typing as t import construct as cs -import construct_typed as cst -from construct_typed import DataclassBitStruct, DataclassStruct, csfield +from construct_typed import ( + DataclassBitStruct, + DataclassStruct, + csfield, + construct, + TEnum, + TFlags, +) from .declarativeunittest import common, raises, setattrs def test_dataclass_const_default() -> None: - @dataclasses.dataclass - class ConstDefaultTest(DataclassMixin): + class ConstDefaultTest(DataclassStruct): const_bytes: bytes = csfield(cs.Const(b"BMP")) const_int: int = csfield(cs.Const(5, cs.Int8ub)) default_int: int = csfield(cs.Default(cs.Int8ub, 28)) @@ -29,8 +32,7 @@ def test_dataclass_const_default() -> None: def test_dataclass_access() -> None: - @dataclasses.dataclass - class TestTContainer(DataclassMixin): + class TestTContainer(DataclassStruct): a: t.Optional[int] = csfield(cs.Const(1, cs.Byte)) b: int = csfield(cs.Int8ub) @@ -50,17 +52,16 @@ def test_dataclass_access() -> None: assert tcontainer["a"] == 6 # wrong creation - assert raises(lambda: TestTContainer(a=0, b=1)) == TypeError + assert raises(lambda: TestTContainer(a=0, b=1)) == TypeError # type: ignore def test_dataclass_str_repr() -> None: - @dataclasses.dataclass - class Image(DataclassMixin): + class Image(DataclassStruct): signature: t.Optional[bytes] = csfield(cs.Const(b"BMP")) width: int = csfield(cs.Int8ub) height: int = csfield(cs.Int8ub) - format = DataclassStruct(Image) + format = construct(Image) obj = Image(width=3, height=2) assert ( str(obj) @@ -74,20 +75,19 @@ def test_dataclass_str_repr() -> None: def test_dataclass_struct() -> None: - @dataclasses.dataclass - class Image(DataclassMixin): + class Image(DataclassStruct): width: int = csfield(cs.Int8ub) height: int = csfield(cs.Int8ub) pixels: bytes = csfield(cs.Bytes(cs.this.height * cs.this.width)) common( - cst.DataclassStruct(Image), + construct(Image), b"\x01\x0212", Image(width=1, height=2, pixels=b"12"), ) # check __getattr__ - c = cst.DataclassStruct(Image) + c = Image.__construct__() assert c.width.name == "width" assert c.height.name == "height" assert c.width.subcon is cs.Int8ub @@ -95,43 +95,36 @@ def test_dataclass_struct() -> None: def test_dataclass_struct_reverse() -> None: - @dataclasses.dataclass - class TestContainer(DataclassMixin): + class TestContainer(DataclassStruct, reverse_fields=True): a: int = csfield(cs.Int16ub) b: int = csfield(cs.Int8ub) common( - DataclassStruct(TestContainer, reverse=True), + construct(TestContainer), b"\x02\x00\x01", TestContainer(a=1, b=2), 3, ) - normal = DataclassStruct(TestContainer) - reverse = DataclassStruct(TestContainer, reverse=True) - assert str(normal.parse(b"\x00\x01\x02")) == str(reverse.parse(b"\x02\x00\x01")) def test_dataclass_struct_nested() -> None: - @dataclasses.dataclass - class TestContainer(DataclassMixin): - @dataclasses.dataclass - class InnerDataclass(DataclassMixin): + class TestContainer(DataclassStruct): + class InnerDataclass(DataclassStruct): b: int = csfield(cs.Byte) c: bytes = csfield(cs.Bytes(cs.this._.length)) length: int = csfield(cs.Byte) - a: InnerDataclass = csfield(DataclassStruct(InnerDataclass)) + a: InnerDataclass = csfield(construct(InnerDataclass)) common( - DataclassStruct(TestContainer), + construct(TestContainer), b"\x02\x01\xF1\xF2", TestContainer(length=2, a=TestContainer.InnerDataclass(b=1, c=b"\xF1\xF2")), ) def test_dataclass_struct_default_field() -> None: - @dataclasses.dataclass - class Image(DataclassMixin): + class Image(DataclassStruct): width: int = csfield(cs.Int8ub) height: int = csfield(cs.Int8ub) pixels: t.Optional[bytes] = csfield( @@ -142,7 +135,7 @@ def test_dataclass_struct_default_field() -> None: ) common( - DataclassStruct(Image), + construct(Image), b"\x02\x03\x00\x00\x00\x00\x00\x00", setattrs(Image(2, 3), pixels=bytes(6)), sample_building=Image(2, 3), @@ -150,12 +143,11 @@ def test_dataclass_struct_default_field() -> None: def test_dataclass_struct_const_field() -> None: - @dataclasses.dataclass - class TestContainer(DataclassMixin): + class TestContainer(DataclassStruct): const_field: t.Optional[bytes] = csfield(cs.Const(b"\x00")) common( - DataclassStruct(TestContainer), + construct(TestContainer), bytes(1), setattrs(TestContainer(), const_field=b"\x00"), 1, @@ -163,7 +155,7 @@ def test_dataclass_struct_const_field() -> None: assert ( raises( - DataclassStruct(TestContainer).build, + construct(TestContainer).build, setattrs(TestContainer(), const_field=b"\x01"), ) == cs.ConstError @@ -171,12 +163,11 @@ def test_dataclass_struct_const_field() -> None: def test_dataclass_struct_array_field() -> None: - @dataclasses.dataclass - class TestContainer(DataclassMixin): + class TestContainer(DataclassStruct): array_field: t.List[int] = csfield(cs.Array(5, cs.Int8ub)) common( - DataclassStruct(TestContainer), + construct(TestContainer), bytes(5), TestContainer(array_field=[0, 0, 0, 0, 0]), 5, @@ -184,15 +175,14 @@ def test_dataclass_struct_array_field() -> None: def test_dataclass_struct_anonymus_fields_1() -> None: - @dataclasses.dataclass - class TestContainer(DataclassMixin): + class TestContainer(DataclassStruct): _1: t.Optional[bytes] = csfield(cs.Const(b"\x00")) _2: None = csfield(cs.Padding(1)) _3: None = csfield(cs.Pass) _4: None = csfield(cs.Terminated) common( - DataclassStruct(TestContainer), + construct(TestContainer), bytes(2), setattrs(TestContainer(), _1=b"\x00"), cs.SizeofError, @@ -200,22 +190,20 @@ def test_dataclass_struct_anonymus_fields_1() -> None: def test_dataclass_struct_anonymus_fields_2() -> None: - @dataclasses.dataclass - class TestContainer(DataclassMixin): + class TestContainer(DataclassStruct): _1: int = csfield(cs.Computed(7)) _2: t.Optional[bytes] = csfield(cs.Const(b"JPEG")) _3: None = csfield(cs.Pass) _4: None = csfield(cs.Terminated) - d = DataclassStruct(TestContainer) + d = construct(TestContainer) assert d.build(TestContainer()) == d.build(TestContainer()) def test_dataclass_struct_overloaded_method() -> None: # Test dot access to some names that are not accessable via dot # in the original 'cs.Container'. - @dataclasses.dataclass - class TestContainer(DataclassMixin): + class TestContainer(DataclassStruct): clear: int = csfield(cs.Int8ul) copy: int = csfield(cs.Int8ul) fromkeys: int = csfield(cs.Int8ul) @@ -231,7 +219,7 @@ def test_dataclass_struct_overloaded_method() -> None: update: int = csfield(cs.Int8ul) values: int = csfield(cs.Int8ul) - d = DataclassStruct(TestContainer) + d = construct(TestContainer) obj = d.parse( d.build( TestContainer( @@ -268,44 +256,22 @@ def test_dataclass_struct_overloaded_method() -> None: assert obj.values == 14 -def test_dataclass_struct_no_dataclass() -> None: - class TestContainer(DataclassMixin): - a: int = csfield(cs.Int16ub) - b: int = csfield(cs.Int8ub) - - assert raises(lambda: DataclassStruct(TestContainer)) == TypeError - - -def test_dataclass_struct_no_DataclassMixin() -> None: - @dataclasses.dataclass - class TestContainer: - a: int = csfield(cs.Int16ub) - b: int = csfield(cs.Int8ub) - - cls = t.cast(t.Type[DataclassMixin], TestContainer) - assert raises(lambda: DataclassStruct(cls)) == TypeError - - def test_dataclass_struct_wrong_container() -> None: - @dataclasses.dataclass - class TestContainer1(DataclassMixin): + class TestContainer1(DataclassStruct): a: int = csfield(cs.Int16ub) b: int = csfield(cs.Int8ub) - @dataclasses.dataclass - class TestContainer2(DataclassMixin): + class TestContainer2(DataclassStruct): a: int = csfield(cs.Int16ub) b: int = csfield(cs.Int8ub) assert ( - raises(DataclassStruct(TestContainer1).build, TestContainer2(a=1, b=2)) - == TypeError + raises(construct(TestContainer1).build, TestContainer2(a=1, b=2)) == TypeError ) def test_dataclass_struct_doc() -> None: - @dataclasses.dataclass - class TestContainer(DataclassMixin): + class TestContainer(DataclassStruct): a: int = csfield(cs.Int16ub, "This is the documentation of a") b: int = csfield( cs.Int8ub, doc="This is the documentation of b\nwhich is multiline" @@ -318,7 +284,7 @@ def test_dataclass_struct_doc() -> None: """, ) - format = DataclassStruct(TestContainer) + format = TestContainer.__construct__() common(format, b"\x00\x01\x02\x03", TestContainer(a=1, b=2, c=3), 4) assert format.subcon.a.docs == "This is the documentation of a" @@ -330,39 +296,36 @@ def test_dataclass_struct_doc() -> None: def test_dataclass_bitstruct() -> None: - @dataclasses.dataclass - class TestContainer(DataclassMixin): + class TestContainer(DataclassBitStruct): a: int = csfield(cs.BitsInteger(7)) b: int = csfield(cs.Bit) c: int = csfield(cs.BitsInteger(8)) - print("") - common( - DataclassBitStruct(TestContainer), + construct(TestContainer), b"\xFD\x12", TestContainer(a=0x7E, b=1, c=0x12), 2, ) # check __getattr__ - c = DataclassStruct(TestContainer) - assert c.a.name == "a" - assert c.b.name == "b" - assert c.c.name == "c" - assert isinstance(c.a.subcon, cs.BitsInteger) - assert c.b.subcon is cs.Bit - assert isinstance(c.c.subcon, cs.BitsInteger) + c = TestContainer.__construct__() + assert c.subcon.a.name == "a" + assert c.subcon.b.name == "b" + assert c.subcon.c.name == "c" + assert isinstance(c.subcon.a.subcon, cs.BitsInteger) + assert c.subcon.b.subcon is cs.Bit + assert isinstance(c.subcon.c.subcon, cs.BitsInteger) def test_tenum() -> None: - class TestEnum(cst.EnumBase): + class TestEnum(TEnum, subcon=cs.Byte): one = 1 two = 2 four = 4 eight = 8 - d = cst.TEnum(cs.Byte, TestEnum) + d = construct(TestEnum) common(d, b"\x01", TestEnum.one, 1) common(d, b"\xff", TestEnum(255), 1) @@ -375,57 +338,35 @@ def test_tenum() -> None: assert raises(d.build, 8) == TypeError -def test_tenum_no_enumbase() -> None: - class E(enum.Enum): +def test_tenum_in_dataclass_struct() -> None: + class TestEnum(TEnum, subcon=cs.Int8ub): a = 1 b = 2 - cls = t.cast(t.Type[cst.EnumBase], E) - assert raises(lambda: cst.TEnum(cs.Byte, cls)) == TypeError - - -def test_dataclass_struct_wrong_enumbase() -> None: - class E1(cst.EnumBase): - a = 1 - b = 2 - - class E2(cst.EnumBase): - a = 1 - b = 2 - - assert raises(cst.TEnum(cs.Byte, E1).build, E2.a) == TypeError - - -def test_tenum_in_tstruct() -> None: - class TestEnum(cst.EnumBase): - a = 1 - b = 2 - - @dataclasses.dataclass - class TestContainer(DataclassMixin): - a: TestEnum = csfield(cst.TEnum(cs.Int8ub, TestEnum)) + class TestContainer(DataclassStruct): + a: TestEnum = csfield(construct(TestEnum)) b: int = csfield(cs.Int8ub) common( - DataclassStruct(TestContainer), + construct(TestContainer), b"\x01\x02", TestContainer(a=TestEnum.a, b=2), 2, ) assert ( - raises(cst.TEnum(cs.Byte, TestEnum).build, TestContainer(a=1, b=2)) == TypeError # type: ignore + raises(construct(TestEnum).build, TestContainer(a=1, b=2)) == TypeError # type: ignore ) def test_tenum_flags() -> None: - class TestEnum(cst.FlagsEnumBase): + class TestEnum(TFlags, subcon=cs.Byte): one = 1 two = 2 four = 4 eight = 8 - d = cst.TFlagsEnum(cs.Byte, TestEnum) + d = 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" From 8525c04165e617b4b0cd326800d858505597728b Mon Sep 17 00:00:00 2001 From: Tim Rid <6593626+timrid@users.noreply.github.com> Date: Sun, 13 Feb 2022 19:29:56 +0100 Subject: [PATCH 08/24] added custom _EnumMeta, because __init_subclass__ is not working correctly together with the standard enum.EnumMeta... --- construct_typed/tenum.py | 128 +++++++++++++++++++++++---------------- 1 file changed, 77 insertions(+), 51 deletions(-) diff --git a/construct_typed/tenum.py b/construct_typed/tenum.py index 4a84584..47b089c 100644 --- a/construct_typed/tenum.py +++ b/construct_typed/tenum.py @@ -8,35 +8,78 @@ from .generic import * T = t.TypeVar("T") -# ## TEnum ############################################################################################################ -class TEnum(enum.IntEnum): +class _EnumMeta(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 `TEnum` 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 TEnum in bases: + enum_constr = TEnumConstruct(subcon, cls) # type: ignore + elif TFlags in bases: + enum_constr = TFlagsConstruct(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 + + +# ## TEnumConstruct ############################################################################################################ +class TEnum(enum.IntEnum, metaclass=_EnumMeta): """ Base class for an Enum used in `construct_typed.TEnumConstruct`. This class extends the standard `enum.IntEnum`, so that missing values are automatically generated. """ - @classmethod - def __init_subclass__( - cls, - subcon: "cs.Construct[t.Any, t.Any]", - **kwargs: t.Any, - ): - super().__init_subclass__(**kwargs) + if t.TYPE_CHECKING: + # unfortunately the metaclass `enum.EnumMeta` does not forward the parameters to __init_subclass__, so that + # we have to make our own metaclass `ConstructEnumMeta`. + # But pylance/pyright is checking the type parameters passed to the class via __init_subclass__, so that we + # have to fake one. + @classmethod + def __init_subclass__( + cls, + subcon: "cs.Construct[t.Any, t.Any]", + ): + ... - # validate types - if not isinstance(subcon, cs.Construct): # type: ignore - raise ValueError( - f"`subcon` parameter has to be an `Construct` object but is {type(subcon)}" - ) - - # create construct format - enum_constr = TEnumConstruct(subcon, cls) - - # save construct format and make the class compatible to `Constructable` protocol - setattr(cls, "__construct__", lambda: enum_constr) - - return cls + @classmethod + def __construct__(cls: "t.Type[EnumType]") -> "TEnumConstruct[EnumType]": + ... # Extend the enum type with __missing__ method. So if a enum value # not found in the enum, a new pseudo member is created. @@ -59,12 +102,6 @@ class TEnum(enum.IntEnum): pseudo_member = cls._value2member_map_.setdefault(value, new_member) # type: ignore return pseudo_member # type: ignore - if t.TYPE_CHECKING: - - @classmethod - def __construct__(cls: "t.Type[EnumType]") -> "TEnumConstruct[EnumType]": - ... - EnumType = t.TypeVar("EnumType", bound=TEnum) @@ -108,31 +145,20 @@ class TEnumConstruct(Adapter[int, int, EnumType, EnumType]): "'{}' has to be of type {}".format(repr(obj), repr(self.enum_type)) ) + # ## TFlags ####################################################################################################### -class TFlags(enum.IntFlag): - @classmethod - def __init_subclass__( - cls, - subcon: "cs.Construct[t.Any, t.Any]", - **kwargs: t.Any, - ): - super().__init_subclass__(**kwargs) - - # validate types - if not isinstance(subcon, cs.Construct): # type: ignore - raise ValueError( - f"`subcon` parameter has to be an `Construct` object but is {type(subcon)}" - ) - - # create construct format - enum_constr = TFlagsConstruct(subcon, cls) - - # save construct format and make the class compatible to `Constructable` protocol - setattr(cls, "__construct__", lambda: enum_constr) - - return cls - +class TFlags(enum.IntFlag, metaclass=_EnumMeta): if t.TYPE_CHECKING: + # unfortunately the metaclass `enum.EnumMeta` does not forward the parameters to __init_subclass__, so that + # we have to make our own metaclass `ConstructEnumMeta`. + # But pylance/pyright is checking the type parameters passed to the class via __init_subclass__, so that we + # have to fake one. + @classmethod + def __init_subclass__( + cls, + subcon: "cs.Construct[t.Any, t.Any]", + ): + ... @classmethod def __construct__( From e99dd5d752e7413a8f58919b620afa4026703fde Mon Sep 17 00:00:00 2001 From: Tim Rid <6593626+timrid@users.noreply.github.com> Date: Sun, 13 Feb 2022 19:30:35 +0100 Subject: [PATCH 09/24] some renaming --- construct_typed/tenum.py | 20 ++++++++++---------- tests/test_typed.py | 20 ++++++++++---------- 2 files changed, 20 insertions(+), 20 deletions(-) diff --git a/construct_typed/tenum.py b/construct_typed/tenum.py index 47b089c..b146549 100644 --- a/construct_typed/tenum.py +++ b/construct_typed/tenum.py @@ -162,15 +162,15 @@ class TFlags(enum.IntFlag, metaclass=_EnumMeta): @classmethod def __construct__( - cls: "t.Type[FlagsEnumType]", - ) -> "TFlagsConstruct[FlagsEnumType]": + cls: "t.Type[FlagsType]", + ) -> "TFlagsConstruct[FlagsType]": ... -FlagsEnumType = t.TypeVar("FlagsEnumType", bound=TFlags) +FlagsType = t.TypeVar("FlagsType", bound=TFlags) -class TFlagsConstruct(Adapter[int, int, FlagsEnumType, FlagsEnumType]): +class TFlagsConstruct(Adapter[int, int, FlagsType, FlagsType]): """ Typed flags. """ @@ -178,28 +178,28 @@ class TFlagsConstruct(Adapter[int, int, FlagsEnumType, FlagsEnumType]): if t.TYPE_CHECKING: def __new__( - cls, subcon: Construct[int, int], enum_type: t.Type[FlagsEnumType] - ) -> "TFlagsConstruct[FlagsEnumType]": + cls, subcon: Construct[int, int], enum_type: t.Type[FlagsType] + ) -> "TFlagsConstruct[FlagsType]": ... - def __init__(self, subcon: Construct[int, int], enum_type: t.Type[FlagsEnumType]): + def __init__(self, subcon: Construct[int, int], enum_type: t.Type[FlagsType]): if not issubclass(enum_type, TFlags): raise TypeError( "'{}' has to be a '{}'".format(repr(enum_type), repr(TFlags)) ) # save enum type - self.enum_type = t.cast(t.Type[FlagsEnumType], enum_type) # type: ignore + self.enum_type = t.cast(t.Type[FlagsType], enum_type) # type: ignore # init adatper super(TFlagsConstruct, self).__init__(subcon) # type: ignore - def _decode(self, obj: int, context: Context, path: PathType) -> FlagsEnumType: + def _decode(self, obj: int, context: Context, path: PathType) -> FlagsType: return self.enum_type(obj) def _encode( self, - obj: FlagsEnumType, + obj: FlagsType, context: Context, path: PathType, ) -> int: diff --git a/tests/test_typed.py b/tests/test_typed.py index b69e50c..8efdbb5 100644 --- a/tests/test_typed.py +++ b/tests/test_typed.py @@ -359,19 +359,19 @@ def test_tenum_in_dataclass_struct() -> None: ) -def test_tenum_flags() -> None: - class TestEnum(TFlags, subcon=cs.Byte): +def test_tflags() -> None: + class TestFlags(TFlags, subcon=cs.Byte): one = 1 two = 2 four = 4 eight = 8 - d = 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" - 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" + d = construct(TestFlags) + common(d, b"\x03", TestFlags.one | TestFlags.two, 1) + assert d.build(TestFlags(0)) == b"\x00" + assert d.build(TestFlags.one | TestFlags.two) == b"\x03" + assert d.build(TestFlags(8)) == b"\x08" + assert d.build(TestFlags(1 | 2)) == b"\x03" + assert d.build(TestFlags(255)) == b"\xff" + assert d.build(TestFlags.eight) == b"\x08" assert raises(d.build, 2) == TypeError From 94a8097bec7212c820ef76e3e23257342cd12ae6 Mon Sep 17 00:00:00 2001 From: Tim Rid <6593626+timrid@users.noreply.github.com> Date: Sun, 13 Feb 2022 19:48:00 +0100 Subject: [PATCH 10/24] added this_struct to public interface --- construct_typed/__init__.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/construct_typed/__init__.py b/construct_typed/__init__.py index 86e7d7a..1c11365 100644 --- a/construct_typed/__init__.py +++ b/construct_typed/__init__.py @@ -3,6 +3,7 @@ from construct_typed.dataclass_struct import ( DataclassBitStruct, DataclassStruct, csfield, + this_struct ) from construct_typed.generic import ( Adapter, @@ -23,6 +24,7 @@ __all__ = [ "TEnumConstruct", "TFlags", "TFlagsConstruct", + "this_struct", "Adapter", "ConstantOrContextLambda", "Construct", From be5ae240fb65c3662331fc1dee3ed0c5aac88469 Mon Sep 17 00:00:00 2001 From: Tim Rid <6593626+timrid@users.noreply.github.com> Date: Sat, 19 Feb 2022 12:46:28 +0100 Subject: [PATCH 11/24] renamed `construct` to `constr` so that there are no naming collisions with the construct package name --- construct_typed/__init__.py | 4 +-- construct_typed/dataclass_struct.py | 6 ++-- construct_typed/generic.py | 6 ++-- construct_typed/tenum.py | 6 ++-- tests/test_typed.py | 46 ++++++++++++++--------------- 5 files changed, 34 insertions(+), 34 deletions(-) diff --git a/construct_typed/__init__.py b/construct_typed/__init__.py index 1c11365..56c0e57 100644 --- a/construct_typed/__init__.py +++ b/construct_typed/__init__.py @@ -1,4 +1,4 @@ -from construct_typed.generic import construct +from construct_typed.generic import constr from construct_typed.dataclass_struct import ( DataclassBitStruct, DataclassStruct, @@ -18,7 +18,7 @@ from construct_typed.tenum import TEnum, TFlags, TEnumConstruct, TFlagsConstruct __all__ = [ "DataclassBitStruct", "DataclassStruct", - "construct", + "constr", "csfield", "TEnum", "TEnumConstruct", diff --git a/construct_typed/dataclass_struct.py b/construct_typed/dataclass_struct.py index f9f9705..d6a8566 100644 --- a/construct_typed/dataclass_struct.py +++ b/construct_typed/dataclass_struct.py @@ -78,7 +78,7 @@ def csfield( init = True default = dataclasses.MISSING - # Set default values in case of special sucons + # Set default values in case of special subcons if isinstance(orig_subcon, cs.Const): const_subcon: "cs.Const[t.Any, t.Any, t.Any, t.Any]" = orig_subcon default = const_subcon.value @@ -249,7 +249,7 @@ class DataclassStruct: _replace_this_struct(constr, dc_constr) # save construct format and make the class compatible to `Constructable` protocol - setattr(cls, "__construct__", lambda: constr) + setattr(cls, "__constr__", lambda: constr) return cls @@ -303,7 +303,7 @@ class DataclassStruct: if t.TYPE_CHECKING: @classmethod - def __construct__(cls: t.Type[T]) -> "DataclassConstruct[T]": + def __constr__(cls: t.Type[T]) -> "DataclassConstruct[T]": ... diff --git a/construct_typed/generic.py b/construct_typed/generic.py index 75c0ac2..b988444 100644 --- a/construct_typed/generic.py +++ b/construct_typed/generic.py @@ -43,16 +43,16 @@ else: @t.runtime_checkable class Constructable(t.Protocol[ParsedType, BuildTypes]): - def __construct__(self) -> "Construct[ParsedType, BuildTypes]": + def __constr__(self) -> "Construct[ParsedType, BuildTypes]": raise NotImplementedError -def construct( +def constr( constr: t.Union[ Constructable[ParsedType, BuildTypes], "Construct[ParsedType, BuildTypes]" ], ) -> Construct[ParsedType, BuildTypes]: """Get construct instance of `Constructable` or `Construct`""" if isinstance(constr, Constructable): - constr = constr.__construct__() + constr = constr.__constr__() return constr diff --git a/construct_typed/tenum.py b/construct_typed/tenum.py index b146549..2cf5316 100644 --- a/construct_typed/tenum.py +++ b/construct_typed/tenum.py @@ -52,7 +52,7 @@ class _EnumMeta(enum.EnumMeta): enum_constr = None # save construct format and make the class compatible to `Constructable` protocol - setattr(cls, "__construct__", lambda: enum_constr) # type: ignore + setattr(cls, "__constr__", lambda: enum_constr) # type: ignore return cls @@ -78,7 +78,7 @@ class TEnum(enum.IntEnum, metaclass=_EnumMeta): ... @classmethod - def __construct__(cls: "t.Type[EnumType]") -> "TEnumConstruct[EnumType]": + def __constr__(cls: "t.Type[EnumType]") -> "TEnumConstruct[EnumType]": ... # Extend the enum type with __missing__ method. So if a enum value @@ -161,7 +161,7 @@ class TFlags(enum.IntFlag, metaclass=_EnumMeta): ... @classmethod - def __construct__( + def __constr__( cls: "t.Type[FlagsType]", ) -> "TFlagsConstruct[FlagsType]": ... diff --git a/tests/test_typed.py b/tests/test_typed.py index 8efdbb5..ef0226b 100644 --- a/tests/test_typed.py +++ b/tests/test_typed.py @@ -7,7 +7,7 @@ from construct_typed import ( DataclassBitStruct, DataclassStruct, csfield, - construct, + constr, TEnum, TFlags, ) @@ -61,7 +61,7 @@ def test_dataclass_str_repr() -> None: width: int = csfield(cs.Int8ub) height: int = csfield(cs.Int8ub) - format = construct(Image) + format = constr(Image) obj = Image(width=3, height=2) assert ( str(obj) @@ -81,13 +81,13 @@ def test_dataclass_struct() -> None: pixels: bytes = csfield(cs.Bytes(cs.this.height * cs.this.width)) common( - construct(Image), + constr(Image), b"\x01\x0212", Image(width=1, height=2, pixels=b"12"), ) # check __getattr__ - c = Image.__construct__() + c = Image.__constr__() assert c.width.name == "width" assert c.height.name == "height" assert c.width.subcon is cs.Int8ub @@ -100,7 +100,7 @@ def test_dataclass_struct_reverse() -> None: b: int = csfield(cs.Int8ub) common( - construct(TestContainer), + constr(TestContainer), b"\x02\x00\x01", TestContainer(a=1, b=2), 3, @@ -114,10 +114,10 @@ def test_dataclass_struct_nested() -> None: c: bytes = csfield(cs.Bytes(cs.this._.length)) length: int = csfield(cs.Byte) - a: InnerDataclass = csfield(construct(InnerDataclass)) + a: InnerDataclass = csfield(constr(InnerDataclass)) common( - construct(TestContainer), + constr(TestContainer), b"\x02\x01\xF1\xF2", TestContainer(length=2, a=TestContainer.InnerDataclass(b=1, c=b"\xF1\xF2")), ) @@ -135,7 +135,7 @@ def test_dataclass_struct_default_field() -> None: ) common( - construct(Image), + constr(Image), b"\x02\x03\x00\x00\x00\x00\x00\x00", setattrs(Image(2, 3), pixels=bytes(6)), sample_building=Image(2, 3), @@ -147,7 +147,7 @@ def test_dataclass_struct_const_field() -> None: const_field: t.Optional[bytes] = csfield(cs.Const(b"\x00")) common( - construct(TestContainer), + constr(TestContainer), bytes(1), setattrs(TestContainer(), const_field=b"\x00"), 1, @@ -155,7 +155,7 @@ def test_dataclass_struct_const_field() -> None: assert ( raises( - construct(TestContainer).build, + constr(TestContainer).build, setattrs(TestContainer(), const_field=b"\x01"), ) == cs.ConstError @@ -167,7 +167,7 @@ def test_dataclass_struct_array_field() -> None: array_field: t.List[int] = csfield(cs.Array(5, cs.Int8ub)) common( - construct(TestContainer), + constr(TestContainer), bytes(5), TestContainer(array_field=[0, 0, 0, 0, 0]), 5, @@ -182,7 +182,7 @@ def test_dataclass_struct_anonymus_fields_1() -> None: _4: None = csfield(cs.Terminated) common( - construct(TestContainer), + constr(TestContainer), bytes(2), setattrs(TestContainer(), _1=b"\x00"), cs.SizeofError, @@ -196,7 +196,7 @@ def test_dataclass_struct_anonymus_fields_2() -> None: _3: None = csfield(cs.Pass) _4: None = csfield(cs.Terminated) - d = construct(TestContainer) + d = constr(TestContainer) assert d.build(TestContainer()) == d.build(TestContainer()) @@ -219,7 +219,7 @@ def test_dataclass_struct_overloaded_method() -> None: update: int = csfield(cs.Int8ul) values: int = csfield(cs.Int8ul) - d = construct(TestContainer) + d = constr(TestContainer) obj = d.parse( d.build( TestContainer( @@ -266,7 +266,7 @@ def test_dataclass_struct_wrong_container() -> None: b: int = csfield(cs.Int8ub) assert ( - raises(construct(TestContainer1).build, TestContainer2(a=1, b=2)) == TypeError + raises(constr(TestContainer1).build, TestContainer2(a=1, b=2)) == TypeError ) @@ -284,7 +284,7 @@ def test_dataclass_struct_doc() -> None: """, ) - format = TestContainer.__construct__() + format = TestContainer.__constr__() common(format, b"\x00\x01\x02\x03", TestContainer(a=1, b=2, c=3), 4) assert format.subcon.a.docs == "This is the documentation of a" @@ -302,14 +302,14 @@ def test_dataclass_bitstruct() -> None: c: int = csfield(cs.BitsInteger(8)) common( - construct(TestContainer), + constr(TestContainer), b"\xFD\x12", TestContainer(a=0x7E, b=1, c=0x12), 2, ) # check __getattr__ - c = TestContainer.__construct__() + c = TestContainer.__constr__() assert c.subcon.a.name == "a" assert c.subcon.b.name == "b" assert c.subcon.c.name == "c" @@ -325,7 +325,7 @@ def test_tenum() -> None: four = 4 eight = 8 - d = construct(TestEnum) + d = constr(TestEnum) common(d, b"\x01", TestEnum.one, 1) common(d, b"\xff", TestEnum(255), 1) @@ -344,18 +344,18 @@ def test_tenum_in_dataclass_struct() -> None: b = 2 class TestContainer(DataclassStruct): - a: TestEnum = csfield(construct(TestEnum)) + a: TestEnum = csfield(constr(TestEnum)) b: int = csfield(cs.Int8ub) common( - construct(TestContainer), + constr(TestContainer), b"\x01\x02", TestContainer(a=TestEnum.a, b=2), 2, ) assert ( - raises(construct(TestEnum).build, TestContainer(a=1, b=2)) == TypeError # type: ignore + raises(constr(TestEnum).build, TestContainer(a=1, b=2)) == TypeError # type: ignore ) @@ -366,7 +366,7 @@ def test_tflags() -> None: four = 4 eight = 8 - d = construct(TestFlags) + d = constr(TestFlags) common(d, b"\x03", TestFlags.one | TestFlags.two, 1) assert d.build(TestFlags(0)) == b"\x00" assert d.build(TestFlags.one | TestFlags.two) == b"\x03" From 75dbd8d8228aa6111fa4c5b22f01230c8359f62c Mon Sep 17 00:00:00 2001 From: Tim Rid <6593626+timrid@users.noreply.github.com> Date: Sat, 19 Feb 2022 14:10:18 +0100 Subject: [PATCH 12/24] fixed some mypy issues --- construct_typed/dataclass_struct.py | 53 ++++++++++++++--------------- construct_typed/tenum.py | 22 ++++++------ tests/declarativeunittest.pyi | 4 +-- 3 files changed, 39 insertions(+), 40 deletions(-) diff --git a/construct_typed/dataclass_struct.py b/construct_typed/dataclass_struct.py index d6a8566..ea89f5a 100644 --- a/construct_typed/dataclass_struct.py +++ b/construct_typed/dataclass_struct.py @@ -31,29 +31,29 @@ def __dataclass_transform__( DATACLASS_METADATA_KEY = "__construct_typed_subcon" -if t.TYPE_CHECKING: - # specialisation for constructs, that builds from none and dont have to be declared in the __init__ method - @t.overload - def csfield( - subcon: cs.Construct[ParsedType, None], - doc: t.Optional[str] = None, - parsed: t.Optional[t.Callable[[t.Any, Context], None]] = None, - init: t.Literal[False] = False, - ) -> ParsedType: - ... +# specialisation for constructs, that builds from none and dont have to be declared in the __init__ method +@t.overload +def csfield( + subcon: "cs.Construct[ParsedType, None]", + doc: t.Optional[str] = None, + parsed: t.Optional[t.Callable[[t.Any, Context], None]] = None, + init: t.Literal[False] = False, +) -> ParsedType: + ... - @t.overload - def csfield( - subcon: Construct[ParsedType, t.Any], - doc: t.Optional[str] = None, - parsed: t.Optional[t.Callable[[t.Any, Context], None]] = None, - init: bool = True, - ) -> ParsedType: - ... + +@t.overload +def csfield( + subcon: "Construct[ParsedType, t.Any]", + doc: t.Optional[str] = None, + parsed: t.Optional[t.Callable[[t.Any, Context], None]] = None, + init: bool = True, +) -> ParsedType: + ... def csfield( - subcon: Construct[ParsedType, t.Any], + subcon: "Construct[ParsedType, t.Any]", doc: t.Optional[str] = None, parsed: t.Optional[t.Callable[[t.Any, Context], None]] = None, init: bool = True, @@ -119,7 +119,9 @@ class DataclassConstruct(Adapter[t.Any, t.Any, T, T]): reverse: bool = False, ) -> None: if not issubclass(dc_type, DataclassStruct): - raise TypeError(f"'{repr(dc_type)}' has to be a subclass of 'DataclassStruct'") + raise TypeError( + f"'{repr(dc_type)}' has to be a subclass of 'DataclassStruct'" + ) if not dataclasses.is_dataclass(dc_type): raise TypeError(f"'{repr(dc_type)}' has to be a 'dataclasses.dataclass'") self.dc_type = dc_type @@ -155,7 +157,7 @@ class DataclassConstruct(Adapter[t.Any, t.Any, T, T]): dc_init[field.name] = value # create object of dataclass - dc = self.dc_type(**dc_init) # type: ignore + dc: T = self.dc_type(**dc_init) # type: ignore # extract all other values from the container, an pass it to the dataclass for field in fields: @@ -185,7 +187,7 @@ class DataclassConstruct(Adapter[t.Any, t.Any, T, T]): this_struct: Construct[t.Any, t.Any] = Construct() -def _replace_this_struct(constr: "Construct[t.Any, t.Any]", replacement: t.Any): +def _replace_this_struct(constr: "Construct[t.Any, t.Any]", replacement: t.Any) -> None: """Recursive search for `this_struct` in all SubConstructs and replace it with AttrsStruct""" subcon = getattr(constr, "subcon", None) if subcon is this_struct: @@ -231,7 +233,7 @@ class DataclassStruct: cls, constr: "cs.Construct[t.Any, t.Any]" = this_struct, reverse_fields: bool = False, - ): + ) -> None: # validate types if not isinstance(constr, cs.Construct): # type: ignore raise ValueError("`constr` parameter has to be an `Construct` object") @@ -251,8 +253,6 @@ class DataclassStruct: # save construct format and make the class compatible to `Constructable` protocol setattr(cls, "__constr__", lambda: constr) - return cls - # the `construct` library is using the [] access internally, so struct objects # should also make this possible and not only via the dot access. def __getitem__(self, key: str) -> t.Any: @@ -336,6 +336,5 @@ class DataclassBitStruct(DataclassStruct): cls, constr: "cs.Construct[t.Any, t.Any]" = this_struct, reverse_fields: bool = False, - ): + ) -> None: cls = DataclassStruct.__init_subclass__.__func__(cls, cs.Bitwise(constr), reverse_fields) # type: ignore - return cls diff --git a/construct_typed/tenum.py b/construct_typed/tenum.py index 2cf5316..83f1d8c 100644 --- a/construct_typed/tenum.py +++ b/construct_typed/tenum.py @@ -12,22 +12,22 @@ class _EnumMeta(enum.EnumMeta): @classmethod def __prepare__( metacls, # type: ignore - name: str, - bases: t.Tuple[type, ...], + __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) + 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], + __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 + cls: T = super().__new__(metacls, __name, __bases, __namespace) # type: ignore # if the `TEnum` class is created, there are no parameters if len(kwargs) == 0: @@ -44,9 +44,9 @@ class _EnumMeta(enum.EnumMeta): raise ValueError(f"unsupported parameter(s) detected: {unsupp_parm}") # create construct format - if TEnum in bases: + if TEnum in __bases: enum_constr = TEnumConstruct(subcon, cls) # type: ignore - elif TFlags in bases: + elif TFlags in __bases: enum_constr = TFlagsConstruct(subcon, cls) # type: ignore else: enum_constr = None @@ -74,7 +74,7 @@ class TEnum(enum.IntEnum, metaclass=_EnumMeta): def __init_subclass__( cls, subcon: "cs.Construct[t.Any, t.Any]", - ): + ) -> None: ... @classmethod @@ -157,7 +157,7 @@ class TFlags(enum.IntFlag, metaclass=_EnumMeta): def __init_subclass__( cls, subcon: "cs.Construct[t.Any, t.Any]", - ): + ) -> None: ... @classmethod diff --git a/tests/declarativeunittest.pyi b/tests/declarativeunittest.pyi index e2f8cab..5afd044 100644 --- a/tests/declarativeunittest.pyi +++ b/tests/declarativeunittest.pyi @@ -6,7 +6,7 @@ import construct_typed as cst Buffer = t.Union[bytes, memoryview, bytearray] ParsedType = t.TypeVar("ParsedType") BuildTypes = t.TypeVar("BuildTypes") -ContainerType = t.TypeVar("ContainerType", bound=cst.TContainerMixin) +ContainerType = t.TypeVar("ContainerType", bound=cst.DataclassStruct) T = t.TypeVar("T") IdentType = t.TypeVar("IdentType") @@ -20,7 +20,7 @@ def raises( ) -> t.Union[t.Any, Exception]: ... @t.overload def common( - format: cst.TStruct[ContainerType], + format: ContainerType, datasample: Buffer, objsample: t.Union[ContainerType, t.Dict[str, t.Any]], sizesample: t.Union[int, t.Type[Exception]] = ..., From b21c43ca65475f3045dc2d92a74c516a9dffb99a Mon Sep 17 00:00:00 2001 From: Tim Rid <6593626+timrid@users.noreply.github.com> Date: Sat, 19 Feb 2022 21:27:52 +0100 Subject: [PATCH 13/24] added dataclass module from python 3.10 --- construct_typed/dataclass_py310.py | 1454 ++++++++++++++++++++++++++++ 1 file changed, 1454 insertions(+) create mode 100644 construct_typed/dataclass_py310.py diff --git a/construct_typed/dataclass_py310.py b/construct_typed/dataclass_py310.py new file mode 100644 index 0000000..fb3c7d5 --- /dev/null +++ b/construct_typed/dataclass_py310.py @@ -0,0 +1,1454 @@ +import re +import sys +import copy +import types +import inspect +import keyword +import builtins +import functools +import abc +import _thread +from types import FunctionType, GenericAlias + + +__all__ = ['dataclass', + 'field', + 'Field', + 'FrozenInstanceError', + 'InitVar', + 'KW_ONLY', + 'MISSING', + + # Helper functions. + 'fields', + 'asdict', + 'astuple', + 'make_dataclass', + 'replace', + 'is_dataclass', + ] + +# Conditions for adding methods. The boxes indicate what action the +# dataclass decorator takes. For all of these tables, when I talk +# about init=, repr=, eq=, order=, unsafe_hash=, or frozen=, I'm +# referring to the arguments to the @dataclass decorator. When +# checking if a dunder method already exists, I mean check for an +# entry in the class's __dict__. I never check to see if an attribute +# is defined in a base class. + +# Key: +# +=========+=========================================+ +# + Value | Meaning | +# +=========+=========================================+ +# | | No action: no method is added. | +# +---------+-----------------------------------------+ +# | add | Generated method is added. | +# +---------+-----------------------------------------+ +# | raise | TypeError is raised. | +# +---------+-----------------------------------------+ +# | None | Attribute is set to None. | +# +=========+=========================================+ + +# __init__ +# +# +--- init= parameter +# | +# v | | | +# | no | yes | <--- class has __init__ in __dict__? +# +=======+=======+=======+ +# | False | | | +# +-------+-------+-------+ +# | True | add | | <- the default +# +=======+=======+=======+ + +# __repr__ +# +# +--- repr= parameter +# | +# v | | | +# | no | yes | <--- class has __repr__ in __dict__? +# +=======+=======+=======+ +# | False | | | +# +-------+-------+-------+ +# | True | add | | <- the default +# +=======+=======+=======+ + + +# __setattr__ +# __delattr__ +# +# +--- frozen= parameter +# | +# v | | | +# | no | yes | <--- class has __setattr__ or __delattr__ in __dict__? +# +=======+=======+=======+ +# | False | | | <- the default +# +-------+-------+-------+ +# | True | add | raise | +# +=======+=======+=======+ +# Raise because not adding these methods would break the "frozen-ness" +# of the class. + +# __eq__ +# +# +--- eq= parameter +# | +# v | | | +# | no | yes | <--- class has __eq__ in __dict__? +# +=======+=======+=======+ +# | False | | | +# +-------+-------+-------+ +# | True | add | | <- the default +# +=======+=======+=======+ + +# __lt__ +# __le__ +# __gt__ +# __ge__ +# +# +--- order= parameter +# | +# v | | | +# | no | yes | <--- class has any comparison method in __dict__? +# +=======+=======+=======+ +# | False | | | <- the default +# +-------+-------+-------+ +# | True | add | raise | +# +=======+=======+=======+ +# Raise because to allow this case would interfere with using +# functools.total_ordering. + +# __hash__ + +# +------------------- unsafe_hash= parameter +# | +----------- eq= parameter +# | | +--- frozen= parameter +# | | | +# v v v | | | +# | no | yes | <--- class has explicitly defined __hash__ +# +=======+=======+=======+========+========+ +# | False | False | False | | | No __eq__, use the base class __hash__ +# +-------+-------+-------+--------+--------+ +# | False | False | True | | | No __eq__, use the base class __hash__ +# +-------+-------+-------+--------+--------+ +# | False | True | False | None | | <-- the default, not hashable +# +-------+-------+-------+--------+--------+ +# | False | True | True | add | | Frozen, so hashable, allows override +# +-------+-------+-------+--------+--------+ +# | True | False | False | add | raise | Has no __eq__, but hashable +# +-------+-------+-------+--------+--------+ +# | True | False | True | add | raise | Has no __eq__, but hashable +# +-------+-------+-------+--------+--------+ +# | True | True | False | add | raise | Not frozen, but hashable +# +-------+-------+-------+--------+--------+ +# | True | True | True | add | raise | Frozen, so hashable +# +=======+=======+=======+========+========+ +# For boxes that are blank, __hash__ is untouched and therefore +# inherited from the base class. If the base is object, then +# id-based hashing is used. +# +# Note that a class may already have __hash__=None if it specified an +# __eq__ method in the class body (not one that was created by +# @dataclass). +# +# See _hash_action (below) for a coded version of this table. + +# __match_args__ +# +# +--- match_args= parameter +# | +# v | | | +# | no | yes | <--- class has __match_args__ in __dict__? +# +=======+=======+=======+ +# | False | | | +# +-------+-------+-------+ +# | True | add | | <- the default +# +=======+=======+=======+ +# __match_args__ is always added unless the class already defines it. It is a +# tuple of __init__ parameter names; non-init fields must be matched by keyword. + + +# Raised when an attempt is made to modify a frozen class. +class FrozenInstanceError(AttributeError): pass + +# A sentinel object for default values to signal that a default +# factory will be used. This is given a nice repr() which will appear +# in the function signature of dataclasses' constructors. +class _HAS_DEFAULT_FACTORY_CLASS: + def __repr__(self): + return '' +_HAS_DEFAULT_FACTORY = _HAS_DEFAULT_FACTORY_CLASS() + +# A sentinel object to detect if a parameter is supplied or not. Use +# a class to give it a better repr. +class _MISSING_TYPE: + pass +MISSING = _MISSING_TYPE() + +# A sentinel object to indicate that following fields are keyword-only by +# default. Use a class to give it a better repr. +class _KW_ONLY_TYPE: + pass +KW_ONLY = _KW_ONLY_TYPE() + +# Since most per-field metadata will be unused, create an empty +# read-only proxy that can be shared among all fields. +_EMPTY_METADATA = types.MappingProxyType({}) + +# Markers for the various kinds of fields and pseudo-fields. +class _FIELD_BASE: + def __init__(self, name): + self.name = name + def __repr__(self): + return self.name +_FIELD = _FIELD_BASE('_FIELD') +_FIELD_CLASSVAR = _FIELD_BASE('_FIELD_CLASSVAR') +_FIELD_INITVAR = _FIELD_BASE('_FIELD_INITVAR') + +# The name of an attribute on the class where we store the Field +# objects. Also used to check if a class is a Data Class. +_FIELDS = '__dataclass_fields__' + +# The name of an attribute on the class that stores the parameters to +# @dataclass. +_PARAMS = '__dataclass_params__' + +# The name of the function, that if it exists, is called at the end of +# __init__. +_POST_INIT_NAME = '__post_init__' + +# String regex that string annotations for ClassVar or InitVar must match. +# Allows "identifier.identifier[" or "identifier[". +# https://bugs.python.org/issue33453 for details. +_MODULE_IDENTIFIER_RE = re.compile(r'^(?:\s*(\w+)\s*\.)?\s*(\w+)') + +class InitVar: + __slots__ = ('type', ) + + def __init__(self, type): + self.type = type + + def __repr__(self): + if isinstance(self.type, type) and not isinstance(self.type, GenericAlias): + type_name = self.type.__name__ + else: + # typing objects, e.g. List[int] + type_name = repr(self.type) + return f'dataclasses.InitVar[{type_name}]' + + def __class_getitem__(cls, type): + return InitVar(type) + +# Instances of Field are only ever created from within this module, +# and only from the field() function, although Field instances are +# exposed externally as (conceptually) read-only objects. +# +# name and type are filled in after the fact, not in __init__. +# They're not known at the time this class is instantiated, but it's +# convenient if they're available later. +# +# When cls._FIELDS is filled in with a list of Field objects, the name +# and type fields will have been populated. +class Field: + __slots__ = ('name', + 'type', + 'default', + 'default_factory', + 'repr', + 'hash', + 'init', + 'compare', + 'metadata', + 'kw_only', + '_field_type', # Private: not to be used by user code. + ) + + def __init__(self, default, default_factory, init, repr, hash, compare, + metadata, kw_only): + self.name = None + self.type = None + self.default = default + self.default_factory = default_factory + self.init = init + self.repr = repr + self.hash = hash + self.compare = compare + self.metadata = (_EMPTY_METADATA + if metadata is None else + types.MappingProxyType(metadata)) + self.kw_only = kw_only + self._field_type = None + + def __repr__(self): + return ('Field(' + f'name={self.name!r},' + f'type={self.type!r},' + f'default={self.default!r},' + f'default_factory={self.default_factory!r},' + f'init={self.init!r},' + f'repr={self.repr!r},' + f'hash={self.hash!r},' + f'compare={self.compare!r},' + f'metadata={self.metadata!r},' + f'kw_only={self.kw_only!r},' + f'_field_type={self._field_type}' + ')') + + # This is used to support the PEP 487 __set_name__ protocol in the + # case where we're using a field that contains a descriptor as a + # default value. For details on __set_name__, see + # https://www.python.org/dev/peps/pep-0487/#implementation-details. + # + # Note that in _process_class, this Field object is overwritten + # with the default value, so the end result is a descriptor that + # had __set_name__ called on it at the right time. + def __set_name__(self, owner, name): + func = getattr(type(self.default), '__set_name__', None) + if func: + # There is a __set_name__ method on the descriptor, call + # it. + func(self.default, owner, name) + + __class_getitem__ = classmethod(GenericAlias) + + +class _DataclassParams: + __slots__ = ('init', + 'repr', + 'eq', + 'order', + 'unsafe_hash', + 'frozen', + ) + + def __init__(self, init, repr, eq, order, unsafe_hash, frozen): + self.init = init + self.repr = repr + self.eq = eq + self.order = order + self.unsafe_hash = unsafe_hash + self.frozen = frozen + + def __repr__(self): + return ('_DataclassParams(' + f'init={self.init!r},' + f'repr={self.repr!r},' + f'eq={self.eq!r},' + f'order={self.order!r},' + f'unsafe_hash={self.unsafe_hash!r},' + f'frozen={self.frozen!r}' + ')') + + +# This function is used instead of exposing Field creation directly, +# so that a type checker can be told (via overloads) that this is a +# function whose type depends on its parameters. +def field(*, default=MISSING, default_factory=MISSING, init=True, repr=True, + hash=None, compare=True, metadata=None, kw_only=MISSING): + """Return an object to identify dataclass fields. + + default is the default value of the field. default_factory is a + 0-argument function called to initialize a field's value. If init + is true, the field will be a parameter to the class's __init__() + function. If repr is true, the field will be included in the + object's repr(). If hash is true, the field will be included in the + object's hash(). If compare is true, the field will be used in + comparison functions. metadata, if specified, must be a mapping + which is stored but not otherwise examined by dataclass. If kw_only + is true, the field will become a keyword-only parameter to + __init__(). + + It is an error to specify both default and default_factory. + """ + + if default is not MISSING and default_factory is not MISSING: + raise ValueError('cannot specify both default and default_factory') + return Field(default, default_factory, init, repr, hash, compare, + metadata, kw_only) + + +def _fields_in_init_order(fields): + # Returns the fields as __init__ will output them. It returns 2 tuples: + # the first for normal args, and the second for keyword args. + + return (tuple(f for f in fields if f.init and not f.kw_only), + tuple(f for f in fields if f.init and f.kw_only) + ) + + +def _tuple_str(obj_name, fields): + # Return a string representing each field of obj_name as a tuple + # member. So, if fields is ['x', 'y'] and obj_name is "self", + # return "(self.x,self.y)". + + # Special case for the 0-tuple. + if not fields: + return '()' + # Note the trailing comma, needed if this turns out to be a 1-tuple. + return f'({",".join([f"{obj_name}.{f.name}" for f in fields])},)' + + +# This function's logic is copied from "recursive_repr" function in +# reprlib module to avoid dependency. +def _recursive_repr(user_function): + # Decorator to make a repr function return "..." for a recursive + # call. + repr_running = set() + + @functools.wraps(user_function) + def wrapper(self): + key = id(self), _thread.get_ident() + if key in repr_running: + return '...' + repr_running.add(key) + try: + result = user_function(self) + finally: + repr_running.discard(key) + return result + return wrapper + + +def _create_fn(name, args, body, *, globals=None, locals=None, + return_type=MISSING): + # Note that we mutate locals when exec() is called. Caller + # beware! The only callers are internal to this module, so no + # worries about external callers. + if locals is None: + locals = {} + if 'BUILTINS' not in locals: + locals['BUILTINS'] = builtins + return_annotation = '' + if return_type is not MISSING: + locals['_return_type'] = return_type + return_annotation = '->_return_type' + args = ','.join(args) + body = '\n'.join(f' {b}' for b in body) + + # Compute the text of the entire function. + txt = f' def {name}({args}){return_annotation}:\n{body}' + + local_vars = ', '.join(locals.keys()) + txt = f"def __create_fn__({local_vars}):\n{txt}\n return {name}" + ns = {} + exec(txt, globals, ns) + return ns['__create_fn__'](**locals) + + +def _field_assign(frozen, name, value, self_name): + # If we're a frozen class, then assign to our fields in __init__ + # via object.__setattr__. Otherwise, just use a simple + # assignment. + # + # self_name is what "self" is called in this function: don't + # hard-code "self", since that might be a field name. + if frozen: + return f'BUILTINS.object.__setattr__({self_name},{name!r},{value})' + return f'{self_name}.{name}={value}' + + +def _field_init(f, frozen, globals, self_name, slots): + # Return the text of the line in the body of __init__ that will + # initialize this field. + + default_name = f'_dflt_{f.name}' + if f.default_factory is not MISSING: + if f.init: + # This field has a default factory. If a parameter is + # given, use it. If not, call the factory. + globals[default_name] = f.default_factory + value = (f'{default_name}() ' + f'if {f.name} is _HAS_DEFAULT_FACTORY ' + f'else {f.name}') + else: + # This is a field that's not in the __init__ params, but + # has a default factory function. It needs to be + # initialized here by calling the factory function, + # because there's no other way to initialize it. + + # For a field initialized with a default=defaultvalue, the + # class dict just has the default value + # (cls.fieldname=defaultvalue). But that won't work for a + # default factory, the factory must be called in __init__ + # and we must assign that to self.fieldname. We can't + # fall back to the class dict's value, both because it's + # not set, and because it might be different per-class + # (which, after all, is why we have a factory function!). + + globals[default_name] = f.default_factory + value = f'{default_name}()' + else: + # No default factory. + if f.init: + if f.default is MISSING: + # There's no default, just do an assignment. + value = f.name + elif f.default is not MISSING: + globals[default_name] = f.default + value = f.name + else: + # If the class has slots, then initialize this field. + if slots and f.default is not MISSING: + globals[default_name] = f.default + value = default_name + else: + # This field does not need initialization: reading from it will + # just use the class attribute that contains the default. + # Signify that to the caller by returning None. + return None + + # Only test this now, so that we can create variables for the + # default. However, return None to signify that we're not going + # to actually do the assignment statement for InitVars. + if f._field_type is _FIELD_INITVAR: + return None + + # Now, actually generate the field assignment. + return _field_assign(frozen, f.name, value, self_name) + + +def _init_param(f): + # Return the __init__ parameter string for this field. For + # example, the equivalent of 'x:int=3' (except instead of 'int', + # reference a variable set to int, and instead of '3', reference a + # variable set to 3). + if f.default is MISSING and f.default_factory is MISSING: + # There's no default, and no default_factory, just output the + # variable name and type. + default = '' + elif f.default is not MISSING: + # There's a default, this will be the name that's used to look + # it up. + default = f'=_dflt_{f.name}' + elif f.default_factory is not MISSING: + # There's a factory function. Set a marker. + default = '=_HAS_DEFAULT_FACTORY' + return f'{f.name}:_type_{f.name}{default}' + + +def _init_fn(fields, std_fields, kw_only_fields, frozen, has_post_init, + self_name, globals, slots): + # fields contains both real fields and InitVar pseudo-fields. + + # Make sure we don't have fields without defaults following fields + # with defaults. This actually would be caught when exec-ing the + # function source code, but catching it here gives a better error + # message, and future-proofs us in case we build up the function + # using ast. + + seen_default = False + for f in std_fields: + # Only consider the non-kw-only fields in the __init__ call. + if f.init: + if not (f.default is MISSING and f.default_factory is MISSING): + seen_default = True + elif seen_default: + raise TypeError(f'non-default argument {f.name!r} ' + 'follows default argument') + + locals = {f'_type_{f.name}': f.type for f in fields} + locals.update({ + 'MISSING': MISSING, + '_HAS_DEFAULT_FACTORY': _HAS_DEFAULT_FACTORY, + }) + + body_lines = [] + for f in fields: + line = _field_init(f, frozen, locals, self_name, slots) + # line is None means that this field doesn't require + # initialization (it's a pseudo-field). Just skip it. + if line: + body_lines.append(line) + + # Does this class have a post-init function? + if has_post_init: + params_str = ','.join(f.name for f in fields + if f._field_type is _FIELD_INITVAR) + body_lines.append(f'{self_name}.{_POST_INIT_NAME}({params_str})') + + # If no body lines, use 'pass'. + if not body_lines: + body_lines = ['pass'] + + _init_params = [_init_param(f) for f in std_fields] + if kw_only_fields: + # Add the keyword-only args. Because the * can only be added if + # there's at least one keyword-only arg, there needs to be a test here + # (instead of just concatenting the lists together). + _init_params += ['*'] + _init_params += [_init_param(f) for f in kw_only_fields] + return _create_fn('__init__', + [self_name] + _init_params, + body_lines, + locals=locals, + globals=globals, + return_type=None) + + +def _repr_fn(fields, globals): + fn = _create_fn('__repr__', + ('self',), + ['return self.__class__.__qualname__ + f"(' + + ', '.join([f"{f.name}={{self.{f.name}!r}}" + for f in fields]) + + ')"'], + globals=globals) + return _recursive_repr(fn) + + +def _frozen_get_del_attr(cls, fields, globals): + locals = {'cls': cls, + 'FrozenInstanceError': FrozenInstanceError} + if fields: + fields_str = '(' + ','.join(repr(f.name) for f in fields) + ',)' + else: + # Special case for the zero-length tuple. + fields_str = '()' + return (_create_fn('__setattr__', + ('self', 'name', 'value'), + (f'if type(self) is cls or name in {fields_str}:', + ' raise FrozenInstanceError(f"cannot assign to field {name!r}")', + f'super(cls, self).__setattr__(name, value)'), + locals=locals, + globals=globals), + _create_fn('__delattr__', + ('self', 'name'), + (f'if type(self) is cls or name in {fields_str}:', + ' raise FrozenInstanceError(f"cannot delete field {name!r}")', + f'super(cls, self).__delattr__(name)'), + locals=locals, + globals=globals), + ) + + +def _cmp_fn(name, op, self_tuple, other_tuple, globals): + # Create a comparison function. If the fields in the object are + # named 'x' and 'y', then self_tuple is the string + # '(self.x,self.y)' and other_tuple is the string + # '(other.x,other.y)'. + + return _create_fn(name, + ('self', 'other'), + [ 'if other.__class__ is self.__class__:', + f' return {self_tuple}{op}{other_tuple}', + 'return NotImplemented'], + globals=globals) + + +def _hash_fn(fields, globals): + self_tuple = _tuple_str('self', fields) + return _create_fn('__hash__', + ('self',), + [f'return hash({self_tuple})'], + globals=globals) + + +def _is_classvar(a_type, typing): + # This test uses a typing internal class, but it's the best way to + # test if this is a ClassVar. + return (a_type is typing.ClassVar + or (type(a_type) is typing._GenericAlias + and a_type.__origin__ is typing.ClassVar)) + + +def _is_initvar(a_type, dataclasses): + # The module we're checking against is the module we're + # currently in (dataclasses.py). + return (a_type is dataclasses.InitVar + or type(a_type) is dataclasses.InitVar) + +def _is_kw_only(a_type, dataclasses): + return a_type is dataclasses.KW_ONLY + + +def _is_type(annotation, cls, a_module, a_type, is_type_predicate): + # Given a type annotation string, does it refer to a_type in + # a_module? For example, when checking that annotation denotes a + # ClassVar, then a_module is typing, and a_type is + # typing.ClassVar. + + # It's possible to look up a_module given a_type, but it involves + # looking in sys.modules (again!), and seems like a waste since + # the caller already knows a_module. + + # - annotation is a string type annotation + # - cls is the class that this annotation was found in + # - a_module is the module we want to match + # - a_type is the type in that module we want to match + # - is_type_predicate is a function called with (obj, a_module) + # that determines if obj is of the desired type. + + # Since this test does not do a local namespace lookup (and + # instead only a module (global) lookup), there are some things it + # gets wrong. + + # With string annotations, cv0 will be detected as a ClassVar: + # CV = ClassVar + # @dataclass + # class C0: + # cv0: CV + + # But in this example cv1 will not be detected as a ClassVar: + # @dataclass + # class C1: + # CV = ClassVar + # cv1: CV + + # In C1, the code in this function (_is_type) will look up "CV" in + # the module and not find it, so it will not consider cv1 as a + # ClassVar. This is a fairly obscure corner case, and the best + # way to fix it would be to eval() the string "CV" with the + # correct global and local namespaces. However that would involve + # a eval() penalty for every single field of every dataclass + # that's defined. It was judged not worth it. + + match = _MODULE_IDENTIFIER_RE.match(annotation) + if match: + ns = None + module_name = match.group(1) + if not module_name: + # No module name, assume the class's module did + # "from dataclasses import InitVar". + ns = sys.modules.get(cls.__module__).__dict__ + else: + # Look up module_name in the class's module. + module = sys.modules.get(cls.__module__) + if module and module.__dict__.get(module_name) is a_module: + ns = sys.modules.get(a_type.__module__).__dict__ + if ns and is_type_predicate(ns.get(match.group(2)), a_module): + return True + return False + + +def _get_field(cls, a_name, a_type, default_kw_only): + # Return a Field object for this field name and type. ClassVars and + # InitVars are also returned, but marked as such (see f._field_type). + # default_kw_only is the value of kw_only to use if there isn't a field() + # that defines it. + + # If the default value isn't derived from Field, then it's only a + # normal default value. Convert it to a Field(). + default = getattr(cls, a_name, MISSING) + if isinstance(default, Field): + f = default + else: + if isinstance(default, types.MemberDescriptorType): + # This is a field in __slots__, so it has no default value. + default = MISSING + f = field(default=default) + + # Only at this point do we know the name and the type. Set them. + f.name = a_name + f.type = a_type + + # Assume it's a normal field until proven otherwise. We're next + # going to decide if it's a ClassVar or InitVar, everything else + # is just a normal field. + f._field_type = _FIELD + + # In addition to checking for actual types here, also check for + # string annotations. get_type_hints() won't always work for us + # (see https://github.com/python/typing/issues/508 for example), + # plus it's expensive and would require an eval for every string + # annotation. So, make a best effort to see if this is a ClassVar + # or InitVar using regex's and checking that the thing referenced + # is actually of the correct type. + + # For the complete discussion, see https://bugs.python.org/issue33453 + + # If typing has not been imported, then it's impossible for any + # annotation to be a ClassVar. So, only look for ClassVar if + # typing has been imported by any module (not necessarily cls's + # module). + typing = sys.modules.get('typing') + if typing: + if (_is_classvar(a_type, typing) + or (isinstance(f.type, str) + and _is_type(f.type, cls, typing, typing.ClassVar, + _is_classvar))): + f._field_type = _FIELD_CLASSVAR + + # If the type is InitVar, or if it's a matching string annotation, + # then it's an InitVar. + if f._field_type is _FIELD: + # The module we're checking against is the module we're + # currently in (dataclasses.py). + dataclasses = sys.modules[__name__] + if (_is_initvar(a_type, dataclasses) + or (isinstance(f.type, str) + and _is_type(f.type, cls, dataclasses, dataclasses.InitVar, + _is_initvar))): + f._field_type = _FIELD_INITVAR + + # Validations for individual fields. This is delayed until now, + # instead of in the Field() constructor, since only here do we + # know the field name, which allows for better error reporting. + + # Special restrictions for ClassVar and InitVar. + if f._field_type in (_FIELD_CLASSVAR, _FIELD_INITVAR): + if f.default_factory is not MISSING: + raise TypeError(f'field {f.name} cannot have a ' + 'default factory') + # Should I check for other field settings? default_factory + # seems the most serious to check for. Maybe add others. For + # example, how about init=False (or really, + # init=)? It makes no sense for + # ClassVar and InitVar to specify init=. + + # kw_only validation and assignment. + if f._field_type in (_FIELD, _FIELD_INITVAR): + # For real and InitVar fields, if kw_only wasn't specified use the + # default value. + if f.kw_only is MISSING: + f.kw_only = default_kw_only + else: + # Make sure kw_only isn't set for ClassVars + assert f._field_type is _FIELD_CLASSVAR + if f.kw_only is not MISSING: + raise TypeError(f'field {f.name} is a ClassVar but specifies ' + 'kw_only') + + # For real fields, disallow mutable defaults for known types. + if f._field_type is _FIELD and isinstance(f.default, (list, dict, set)): + raise ValueError(f'mutable default {type(f.default)} for field ' + f'{f.name} is not allowed: use default_factory') + + return f + +def _set_qualname(cls, value): + # Ensure that the functions returned from _create_fn uses the proper + # __qualname__ (the class they belong to). + if isinstance(value, FunctionType): + value.__qualname__ = f"{cls.__qualname__}.{value.__name__}" + return value + +def _set_new_attribute(cls, name, value): + # Never overwrites an existing attribute. Returns True if the + # attribute already exists. + if name in cls.__dict__: + return True + _set_qualname(cls, value) + setattr(cls, name, value) + return False + + +# Decide if/how we're going to create a hash function. Key is +# (unsafe_hash, eq, frozen, does-hash-exist). Value is the action to +# take. The common case is to do nothing, so instead of providing a +# function that is a no-op, use None to signify that. + +def _hash_set_none(cls, fields, globals): + return None + +def _hash_add(cls, fields, globals): + flds = [f for f in fields if (f.compare if f.hash is None else f.hash)] + return _set_qualname(cls, _hash_fn(flds, globals)) + +def _hash_exception(cls, fields, globals): + # Raise an exception. + raise TypeError(f'Cannot overwrite attribute __hash__ ' + f'in class {cls.__name__}') + +# +# +-------------------------------------- unsafe_hash? +# | +------------------------------- eq? +# | | +------------------------ frozen? +# | | | +---------------- has-explicit-hash? +# | | | | +# | | | | +------- action +# | | | | | +# v v v v v +_hash_action = {(False, False, False, False): None, + (False, False, False, True ): None, + (False, False, True, False): None, + (False, False, True, True ): None, + (False, True, False, False): _hash_set_none, + (False, True, False, True ): None, + (False, True, True, False): _hash_add, + (False, True, True, True ): None, + (True, False, False, False): _hash_add, + (True, False, False, True ): _hash_exception, + (True, False, True, False): _hash_add, + (True, False, True, True ): _hash_exception, + (True, True, False, False): _hash_add, + (True, True, False, True ): _hash_exception, + (True, True, True, False): _hash_add, + (True, True, True, True ): _hash_exception, + } +# See https://bugs.python.org/issue32929#msg312829 for an if-statement +# version of this table. + + +def _process_class(cls, init, repr, eq, order, unsafe_hash, frozen, + match_args, kw_only, slots): + # Now that dicts retain insertion order, there's no reason to use + # an ordered dict. I am leveraging that ordering here, because + # derived class fields overwrite base class fields, but the order + # is defined by the base class, which is found first. + fields = {} + + if cls.__module__ in sys.modules: + globals = sys.modules[cls.__module__].__dict__ + else: + # Theoretically this can happen if someone writes + # a custom string to cls.__module__. In which case + # such dataclass won't be fully introspectable + # (w.r.t. typing.get_type_hints) but will still function + # correctly. + globals = {} + + setattr(cls, _PARAMS, _DataclassParams(init, repr, eq, order, + unsafe_hash, frozen)) + + # Find our base classes in reverse MRO order, and exclude + # ourselves. In reversed order so that more derived classes + # override earlier field definitions in base classes. As long as + # we're iterating over them, see if any are frozen. + any_frozen_base = False + has_dataclass_bases = False + for b in cls.__mro__[-1:0:-1]: + # Only process classes that have been processed by our + # decorator. That is, they have a _FIELDS attribute. + base_fields = getattr(b, _FIELDS, None) + if base_fields is not None: + has_dataclass_bases = True + for f in base_fields.values(): + fields[f.name] = f + if getattr(b, _PARAMS).frozen: + any_frozen_base = True + + # Annotations that are defined in this class (not in base + # classes). If __annotations__ isn't present, then this class + # adds no new annotations. We use this to compute fields that are + # added by this class. + # + # Fields are found from cls_annotations, which is guaranteed to be + # ordered. Default values are from class attributes, if a field + # has a default. If the default value is a Field(), then it + # contains additional info beyond (and possibly including) the + # actual default value. Pseudo-fields ClassVars and InitVars are + # included, despite the fact that they're not real fields. That's + # dealt with later. + cls_annotations = cls.__dict__.get('__annotations__', {}) + + # Now find fields in our class. While doing so, validate some + # things, and set the default values (as class attributes) where + # we can. + cls_fields = [] + # Get a reference to this module for the _is_kw_only() test. + KW_ONLY_seen = False + dataclasses = sys.modules[__name__] + for name, type in cls_annotations.items(): + # See if this is a marker to change the value of kw_only. + if (_is_kw_only(type, dataclasses) + or (isinstance(type, str) + and _is_type(type, cls, dataclasses, dataclasses.KW_ONLY, + _is_kw_only))): + # Switch the default to kw_only=True, and ignore this + # annotation: it's not a real field. + if KW_ONLY_seen: + raise TypeError(f'{name!r} is KW_ONLY, but KW_ONLY ' + 'has already been specified') + KW_ONLY_seen = True + kw_only = True + else: + # Otherwise it's a field of some type. + cls_fields.append(_get_field(cls, name, type, kw_only)) + + for f in cls_fields: + fields[f.name] = f + + # If the class attribute (which is the default value for this + # field) exists and is of type 'Field', replace it with the + # real default. This is so that normal class introspection + # sees a real default value, not a Field. + if isinstance(getattr(cls, f.name, None), Field): + if f.default is MISSING: + # If there's no default, delete the class attribute. + # This happens if we specify field(repr=False), for + # example (that is, we specified a field object, but + # no default value). Also if we're using a default + # factory. The class attribute should not be set at + # all in the post-processed class. + delattr(cls, f.name) + else: + setattr(cls, f.name, f.default) + + # Do we have any Field members that don't also have annotations? + for name, value in cls.__dict__.items(): + if isinstance(value, Field) and not name in cls_annotations: + raise TypeError(f'{name!r} is a field but has no type annotation') + + # Check rules that apply if we are derived from any dataclasses. + if has_dataclass_bases: + # Raise an exception if any of our bases are frozen, but we're not. + if any_frozen_base and not frozen: + raise TypeError('cannot inherit non-frozen dataclass from a ' + 'frozen one') + + # Raise an exception if we're frozen, but none of our bases are. + if not any_frozen_base and frozen: + raise TypeError('cannot inherit frozen dataclass from a ' + 'non-frozen one') + + # Remember all of the fields on our class (including bases). This + # also marks this class as being a dataclass. + setattr(cls, _FIELDS, fields) + + # Was this class defined with an explicit __hash__? Note that if + # __eq__ is defined in this class, then python will automatically + # set __hash__ to None. This is a heuristic, as it's possible + # that such a __hash__ == None was not auto-generated, but it + # close enough. + class_hash = cls.__dict__.get('__hash__', MISSING) + has_explicit_hash = not (class_hash is MISSING or + (class_hash is None and '__eq__' in cls.__dict__)) + + # If we're generating ordering methods, we must be generating the + # eq methods. + if order and not eq: + raise ValueError('eq must be true if order is true') + + # Include InitVars and regular fields (so, not ClassVars). This is + # initialized here, outside of the "if init:" test, because std_init_fields + # is used with match_args, below. + all_init_fields = [f for f in fields.values() + if f._field_type in (_FIELD, _FIELD_INITVAR)] + (std_init_fields, + kw_only_init_fields) = _fields_in_init_order(all_init_fields) + + if init: + # Does this class have a post-init function? + has_post_init = hasattr(cls, _POST_INIT_NAME) + + _set_new_attribute(cls, '__init__', + _init_fn(all_init_fields, + std_init_fields, + kw_only_init_fields, + frozen, + has_post_init, + # The name to use for the "self" + # param in __init__. Use "self" + # if possible. + '__dataclass_self__' if 'self' in fields + else 'self', + globals, + slots, + )) + + # Get the fields as a list, and include only real fields. This is + # used in all of the following methods. + field_list = [f for f in fields.values() if f._field_type is _FIELD] + + if repr: + flds = [f for f in field_list if f.repr] + _set_new_attribute(cls, '__repr__', _repr_fn(flds, globals)) + + if eq: + # Create __eq__ method. There's no need for a __ne__ method, + # since python will call __eq__ and negate it. + flds = [f for f in field_list if f.compare] + self_tuple = _tuple_str('self', flds) + other_tuple = _tuple_str('other', flds) + _set_new_attribute(cls, '__eq__', + _cmp_fn('__eq__', '==', + self_tuple, other_tuple, + globals=globals)) + + if order: + # Create and set the ordering methods. + flds = [f for f in field_list if f.compare] + self_tuple = _tuple_str('self', flds) + other_tuple = _tuple_str('other', flds) + for name, op in [('__lt__', '<'), + ('__le__', '<='), + ('__gt__', '>'), + ('__ge__', '>='), + ]: + if _set_new_attribute(cls, name, + _cmp_fn(name, op, self_tuple, other_tuple, + globals=globals)): + raise TypeError(f'Cannot overwrite attribute {name} ' + f'in class {cls.__name__}. Consider using ' + 'functools.total_ordering') + + if frozen: + for fn in _frozen_get_del_attr(cls, field_list, globals): + if _set_new_attribute(cls, fn.__name__, fn): + raise TypeError(f'Cannot overwrite attribute {fn.__name__} ' + f'in class {cls.__name__}') + + # Decide if/how we're going to create a hash function. + hash_action = _hash_action[bool(unsafe_hash), + bool(eq), + bool(frozen), + has_explicit_hash] + if hash_action: + # No need to call _set_new_attribute here, since by the time + # we're here the overwriting is unconditional. + cls.__hash__ = hash_action(cls, field_list, globals) + + if not getattr(cls, '__doc__'): + # Create a class doc-string. + cls.__doc__ = (cls.__name__ + + str(inspect.signature(cls)).replace(' -> None', '')) + + if match_args: + # I could probably compute this once + _set_new_attribute(cls, '__match_args__', + tuple(f.name for f in std_init_fields)) + + if slots: + cls = _add_slots(cls, frozen) + + abc.update_abstractmethods(cls) + + return cls + + +# _dataclass_getstate and _dataclass_setstate are needed for pickling frozen +# classes with slots. These could be slighly more performant if we generated +# the code instead of iterating over fields. But that can be a project for +# another day, if performance becomes an issue. +def _dataclass_getstate(self): + return [getattr(self, f.name) for f in fields(self)] + + +def _dataclass_setstate(self, state): + for field, value in zip(fields(self), state): + # use setattr because dataclass may be frozen + object.__setattr__(self, field.name, value) + + +def _add_slots(cls, is_frozen): + # Need to create a new class, since we can't set __slots__ + # after a class has been created. + + # Make sure __slots__ isn't already set. + if '__slots__' in cls.__dict__: + raise TypeError(f'{cls.__name__} already specifies __slots__') + + # Create a new dict for our new class. + cls_dict = dict(cls.__dict__) + field_names = tuple(f.name for f in fields(cls)) + cls_dict['__slots__'] = field_names + for field_name in field_names: + # Remove our attributes, if present. They'll still be + # available in _MARKER. + cls_dict.pop(field_name, None) + + # Remove __dict__ itself. + cls_dict.pop('__dict__', None) + + # And finally create the class. + qualname = getattr(cls, '__qualname__', None) + cls = type(cls)(cls.__name__, cls.__bases__, cls_dict) + if qualname is not None: + cls.__qualname__ = qualname + + if is_frozen: + # Need this for pickling frozen classes with slots. + cls.__getstate__ = _dataclass_getstate + cls.__setstate__ = _dataclass_setstate + + return cls + + +def dataclass(cls=None, /, *, init=True, repr=True, eq=True, order=False, + unsafe_hash=False, frozen=False, match_args=True, + kw_only=False, slots=False): + """Returns the same class as was passed in, with dunder methods + added based on the fields defined in the class. + + Examines PEP 526 __annotations__ to determine fields. + + If init is true, an __init__() method is added to the class. If + repr is true, a __repr__() method is added. If order is true, rich + comparison dunder methods are added. If unsafe_hash is true, a + __hash__() method function is added. If frozen is true, fields may + not be assigned to after instance creation. If match_args is true, + the __match_args__ tuple is added. If kw_only is true, then by + default all fields are keyword-only. If slots is true, an + __slots__ attribute is added. + """ + + def wrap(cls): + return _process_class(cls, init, repr, eq, order, unsafe_hash, + frozen, match_args, kw_only, slots) + + # See if we're being called as @dataclass or @dataclass(). + if cls is None: + # We're called with parens. + return wrap + + # We're called as @dataclass without parens. + return wrap(cls) + + +def fields(class_or_instance): + """Return a tuple describing the fields of this dataclass. + + Accepts a dataclass or an instance of one. Tuple elements are of + type Field. + """ + + # Might it be worth caching this, per class? + try: + fields = getattr(class_or_instance, _FIELDS) + except AttributeError: + raise TypeError('must be called with a dataclass type or instance') + + # Exclude pseudo-fields. Note that fields is sorted by insertion + # order, so the order of the tuple is as the fields were defined. + return tuple(f for f in fields.values() if f._field_type is _FIELD) + + +def _is_dataclass_instance(obj): + """Returns True if obj is an instance of a dataclass.""" + return hasattr(type(obj), _FIELDS) + + +def is_dataclass(obj): + """Returns True if obj is a dataclass or an instance of a + dataclass.""" + cls = obj if isinstance(obj, type) and not isinstance(obj, GenericAlias) else type(obj) + return hasattr(cls, _FIELDS) + + +def asdict(obj, *, dict_factory=dict): + """Return the fields of a dataclass instance as a new dictionary mapping + field names to field values. + + Example usage: + + @dataclass + class C: + x: int + y: int + + c = C(1, 2) + assert asdict(c) == {'x': 1, 'y': 2} + + If given, 'dict_factory' will be used instead of built-in dict. + The function applies recursively to field values that are + dataclass instances. This will also look into built-in containers: + tuples, lists, and dicts. + """ + if not _is_dataclass_instance(obj): + raise TypeError("asdict() should be called on dataclass instances") + return _asdict_inner(obj, dict_factory) + + +def _asdict_inner(obj, dict_factory): + if _is_dataclass_instance(obj): + result = [] + for f in fields(obj): + value = _asdict_inner(getattr(obj, f.name), dict_factory) + result.append((f.name, value)) + return dict_factory(result) + elif isinstance(obj, tuple) and hasattr(obj, '_fields'): + # obj is a namedtuple. Recurse into it, but the returned + # object is another namedtuple of the same type. This is + # similar to how other list- or tuple-derived classes are + # treated (see below), but we just need to create them + # differently because a namedtuple's __init__ needs to be + # called differently (see bpo-34363). + + # I'm not using namedtuple's _asdict() + # method, because: + # - it does not recurse in to the namedtuple fields and + # convert them to dicts (using dict_factory). + # - I don't actually want to return a dict here. The main + # use case here is json.dumps, and it handles converting + # namedtuples to lists. Admittedly we're losing some + # information here when we produce a json list instead of a + # dict. Note that if we returned dicts here instead of + # namedtuples, we could no longer call asdict() on a data + # structure where a namedtuple was used as a dict key. + + return type(obj)(*[_asdict_inner(v, dict_factory) for v in obj]) + elif isinstance(obj, (list, tuple)): + # Assume we can create an object of this type by passing in a + # generator (which is not true for namedtuples, handled + # above). + return type(obj)(_asdict_inner(v, dict_factory) for v in obj) + elif isinstance(obj, dict): + return type(obj)((_asdict_inner(k, dict_factory), + _asdict_inner(v, dict_factory)) + for k, v in obj.items()) + else: + return copy.deepcopy(obj) + + +def astuple(obj, *, tuple_factory=tuple): + """Return the fields of a dataclass instance as a new tuple of field values. + + Example usage:: + + @dataclass + class C: + x: int + y: int + + c = C(1, 2) + assert astuple(c) == (1, 2) + + If given, 'tuple_factory' will be used instead of built-in tuple. + The function applies recursively to field values that are + dataclass instances. This will also look into built-in containers: + tuples, lists, and dicts. + """ + + if not _is_dataclass_instance(obj): + raise TypeError("astuple() should be called on dataclass instances") + return _astuple_inner(obj, tuple_factory) + + +def _astuple_inner(obj, tuple_factory): + if _is_dataclass_instance(obj): + result = [] + for f in fields(obj): + value = _astuple_inner(getattr(obj, f.name), tuple_factory) + result.append(value) + return tuple_factory(result) + elif isinstance(obj, tuple) and hasattr(obj, '_fields'): + # obj is a namedtuple. Recurse into it, but the returned + # object is another namedtuple of the same type. This is + # similar to how other list- or tuple-derived classes are + # treated (see below), but we just need to create them + # differently because a namedtuple's __init__ needs to be + # called differently (see bpo-34363). + return type(obj)(*[_astuple_inner(v, tuple_factory) for v in obj]) + elif isinstance(obj, (list, tuple)): + # Assume we can create an object of this type by passing in a + # generator (which is not true for namedtuples, handled + # above). + return type(obj)(_astuple_inner(v, tuple_factory) for v in obj) + elif isinstance(obj, dict): + return type(obj)((_astuple_inner(k, tuple_factory), _astuple_inner(v, tuple_factory)) + for k, v in obj.items()) + else: + return copy.deepcopy(obj) + + +def make_dataclass(cls_name, fields, *, bases=(), namespace=None, init=True, + repr=True, eq=True, order=False, unsafe_hash=False, + frozen=False, match_args=True, kw_only=False, slots=False): + """Return a new dynamically created dataclass. + + The dataclass name will be 'cls_name'. 'fields' is an iterable + of either (name), (name, type) or (name, type, Field) objects. If type is + omitted, use the string 'typing.Any'. Field objects are created by + the equivalent of calling 'field(name, type [, Field-info])'. + + C = make_dataclass('C', ['x', ('y', int), ('z', int, field(init=False))], bases=(Base,)) + + is equivalent to: + + @dataclass + class C(Base): + x: 'typing.Any' + y: int + z: int = field(init=False) + + For the bases and namespace parameters, see the builtin type() function. + + The parameters init, repr, eq, order, unsafe_hash, and frozen are passed to + dataclass(). + """ + + if namespace is None: + namespace = {} + + # While we're looking through the field names, validate that they + # are identifiers, are not keywords, and not duplicates. + seen = set() + annotations = {} + defaults = {} + for item in fields: + if isinstance(item, str): + name = item + tp = 'typing.Any' + elif len(item) == 2: + name, tp, = item + elif len(item) == 3: + name, tp, spec = item + defaults[name] = spec + else: + raise TypeError(f'Invalid field: {item!r}') + + if not isinstance(name, str) or not name.isidentifier(): + raise TypeError(f'Field names must be valid identifiers: {name!r}') + if keyword.iskeyword(name): + raise TypeError(f'Field names must not be keywords: {name!r}') + if name in seen: + raise TypeError(f'Field name duplicated: {name!r}') + + seen.add(name) + annotations[name] = tp + + # Update 'ns' with the user-supplied namespace plus our calculated values. + def exec_body_callback(ns): + ns.update(namespace) + ns.update(defaults) + ns['__annotations__'] = annotations + + # We use `types.new_class()` instead of simply `type()` to allow dynamic creation + # of generic dataclasses. + cls = types.new_class(cls_name, bases, {}, exec_body_callback) + + # Apply the normal decorator. + return dataclass(cls, init=init, repr=repr, eq=eq, order=order, + unsafe_hash=unsafe_hash, frozen=frozen, + match_args=match_args, kw_only=kw_only, slots=slots) + + +def replace(obj, /, **changes): + """Return a new object replacing specified fields with new values. + + This is especially useful for frozen classes. Example usage: + + @dataclass(frozen=True) + class C: + x: int + y: int + + c = C(1, 2) + c1 = replace(c, x=3) + assert c1.x == 3 and c1.y == 2 + """ + + # We're going to mutate 'changes', but that's okay because it's a + # new dict, even if called with 'replace(obj, **my_changes)'. + + if not _is_dataclass_instance(obj): + raise TypeError("replace() should be called on dataclass instances") + + # It's an error to have init=False fields in 'changes'. + # If a field is not in 'changes', read its value from the provided obj. + + for f in getattr(obj, _FIELDS).values(): + # Only consider normal fields or InitVars. + if f._field_type is _FIELD_CLASSVAR: + continue + + if not f.init: + # Error if this field is specified in changes. + if f.name in changes: + raise ValueError(f'field {f.name} is declared with ' + 'init=False, it cannot be specified with ' + 'replace()') + continue + + if f.name not in changes: + if f._field_type is _FIELD_INITVAR and f.default is MISSING: + raise ValueError(f"InitVar {f.name!r} " + 'must be specified with replace()') + changes[f.name] = getattr(obj, f.name) + + # Create the new object, which calls __init__() and + # __post_init__() (if defined), using all of the init fields we've + # added and/or left in 'changes'. If there are values supplied in + # changes that aren't fields, this will correctly raise a + # TypeError. + return obj.__class__(**changes) \ No newline at end of file From a0cadcce854dc1c83ecd3aa45882a183904dedc2 Mon Sep 17 00:00:00 2001 From: Tim Rid <6593626+timrid@users.noreply.github.com> Date: Sat, 19 Feb 2022 21:54:04 +0100 Subject: [PATCH 14/24] removed things that not work in Python 3.8 --- construct_typed/dataclass_py310.py | 10 +++------- 1 file changed, 3 insertions(+), 7 deletions(-) diff --git a/construct_typed/dataclass_py310.py b/construct_typed/dataclass_py310.py index fb3c7d5..9d6d9bb 100644 --- a/construct_typed/dataclass_py310.py +++ b/construct_typed/dataclass_py310.py @@ -8,7 +8,7 @@ import builtins import functools import abc import _thread -from types import FunctionType, GenericAlias +from types import FunctionType __all__ = ['dataclass', @@ -229,7 +229,7 @@ class InitVar: self.type = type def __repr__(self): - if isinstance(self.type, type) and not isinstance(self.type, GenericAlias): + if isinstance(self.type, type): type_name = self.type.__name__ else: # typing objects, e.g. List[int] @@ -309,8 +309,6 @@ class Field: # it. func(self.default, owner, name) - __class_getitem__ = classmethod(GenericAlias) - class _DataclassParams: __slots__ = ('init', @@ -1101,8 +1099,6 @@ def _process_class(cls, init, repr, eq, order, unsafe_hash, frozen, if slots: cls = _add_slots(cls, frozen) - abc.update_abstractmethods(cls) - return cls @@ -1211,7 +1207,7 @@ def _is_dataclass_instance(obj): def is_dataclass(obj): """Returns True if obj is a dataclass or an instance of a dataclass.""" - cls = obj if isinstance(obj, type) and not isinstance(obj, GenericAlias) else type(obj) + cls = obj if isinstance(obj, type) else type(obj) return hasattr(cls, _FIELDS) From f9e2d24d714ef092f398d127cc5ba2dd8f4c5479 Mon Sep 17 00:00:00 2001 From: Tim Rid <6593626+timrid@users.noreply.github.com> Date: Sat, 19 Feb 2022 21:55:16 +0100 Subject: [PATCH 15/24] Use marker instances from the original dataclass module so that hopefully also the original dataclass functions are working --- construct_typed/dataclass_py310.py | 52 ++++++++++++++++++------------ 1 file changed, 32 insertions(+), 20 deletions(-) diff --git a/construct_typed/dataclass_py310.py b/construct_typed/dataclass_py310.py index 9d6d9bb..0a8c4d6 100644 --- a/construct_typed/dataclass_py310.py +++ b/construct_typed/dataclass_py310.py @@ -9,6 +9,18 @@ import functools import abc import _thread from types import FunctionType +from dataclasses import ( + _FIELD, + _FIELD_CLASSVAR, + _FIELD_INITVAR, + _FIELDS, + _PARAMS, + _POST_INIT_NAME, + _EMPTY_METADATA, + MISSING, + _HAS_DEFAULT_FACTORY, + FrozenInstanceError +) __all__ = ['dataclass', @@ -169,21 +181,21 @@ __all__ = ['dataclass', # Raised when an attempt is made to modify a frozen class. -class FrozenInstanceError(AttributeError): pass +# class FrozenInstanceError(AttributeError): pass # A sentinel object for default values to signal that a default # factory will be used. This is given a nice repr() which will appear # in the function signature of dataclasses' constructors. -class _HAS_DEFAULT_FACTORY_CLASS: - def __repr__(self): - return '' -_HAS_DEFAULT_FACTORY = _HAS_DEFAULT_FACTORY_CLASS() +# class _HAS_DEFAULT_FACTORY_CLASS: +# def __repr__(self): +# return '' +# _HAS_DEFAULT_FACTORY = _HAS_DEFAULT_FACTORY_CLASS() # A sentinel object to detect if a parameter is supplied or not. Use # a class to give it a better repr. -class _MISSING_TYPE: - pass -MISSING = _MISSING_TYPE() +# class _MISSING_TYPE: +# pass +# MISSING = _MISSING_TYPE() # A sentinel object to indicate that following fields are keyword-only by # default. Use a class to give it a better repr. @@ -193,29 +205,29 @@ KW_ONLY = _KW_ONLY_TYPE() # Since most per-field metadata will be unused, create an empty # read-only proxy that can be shared among all fields. -_EMPTY_METADATA = types.MappingProxyType({}) +# _EMPTY_METADATA = types.MappingProxyType({}) # Markers for the various kinds of fields and pseudo-fields. -class _FIELD_BASE: - def __init__(self, name): - self.name = name - def __repr__(self): - return self.name -_FIELD = _FIELD_BASE('_FIELD') -_FIELD_CLASSVAR = _FIELD_BASE('_FIELD_CLASSVAR') -_FIELD_INITVAR = _FIELD_BASE('_FIELD_INITVAR') +# class _FIELD_BASE: +# def __init__(self, name): +# self.name = name +# def __repr__(self): +# return self.name +# _FIELD = _FIELD_BASE('_FIELD') +# _FIELD_CLASSVAR = _FIELD_BASE('_FIELD_CLASSVAR') +# _FIELD_INITVAR = _FIELD_BASE('_FIELD_INITVAR') # The name of an attribute on the class where we store the Field # objects. Also used to check if a class is a Data Class. -_FIELDS = '__dataclass_fields__' +# _FIELDS = '__dataclass_fields__' # The name of an attribute on the class that stores the parameters to # @dataclass. -_PARAMS = '__dataclass_params__' +# _PARAMS = '__dataclass_params__' # The name of the function, that if it exists, is called at the end of # __init__. -_POST_INIT_NAME = '__post_init__' +# _POST_INIT_NAME = '__post_init__' # String regex that string annotations for ClassVar or InitVar must match. # Allows "identifier.identifier[" or "identifier[". From 46fc3dd48c161aab8982a0e14fb721561ee8c96a Mon Sep 17 00:00:00 2001 From: Tim Rid <6593626+timrid@users.noreply.github.com> Date: Sat, 19 Feb 2022 21:55:31 +0100 Subject: [PATCH 16/24] renamed file --- construct_typed/{dataclass_py310.py => dataclasses_py310.py} | 0 1 file changed, 0 insertions(+), 0 deletions(-) rename construct_typed/{dataclass_py310.py => dataclasses_py310.py} (100%) diff --git a/construct_typed/dataclass_py310.py b/construct_typed/dataclasses_py310.py similarity index 100% rename from construct_typed/dataclass_py310.py rename to construct_typed/dataclasses_py310.py From 8529bd287897fd61dd6fa9b66e534ee050767482 Mon Sep 17 00:00:00 2001 From: Tim Rid <6593626+timrid@users.noreply.github.com> Date: Sat, 19 Feb 2022 21:56:50 +0100 Subject: [PATCH 17/24] ignore typing issues and removed unnessesary imports --- construct_typed/dataclasses_py310.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/construct_typed/dataclasses_py310.py b/construct_typed/dataclasses_py310.py index 0a8c4d6..a2f76d0 100644 --- a/construct_typed/dataclasses_py310.py +++ b/construct_typed/dataclasses_py310.py @@ -1,3 +1,4 @@ +# type: ignore import re import sys import copy @@ -6,7 +7,6 @@ import inspect import keyword import builtins import functools -import abc import _thread from types import FunctionType from dataclasses import ( From cab9a07c4f204af9fa2df4a9bda7003b75fdbeac Mon Sep 17 00:00:00 2001 From: Tim Rid <6593626+timrid@users.noreply.github.com> Date: Sun, 20 Feb 2022 12:32:49 +0100 Subject: [PATCH 18/24] Use standard dataclasses modul for python >= 3.10 or use the provided dataclasses module from this library fol python < 3.10. So we can use the `kw_only` option. --- construct_typed/dataclass_struct.py | 12 ++++++++++-- 1 file changed, 10 insertions(+), 2 deletions(-) diff --git a/construct_typed/dataclass_struct.py b/construct_typed/dataclass_struct.py index ea89f5a..8b24286 100644 --- a/construct_typed/dataclass_struct.py +++ b/construct_typed/dataclass_struct.py @@ -1,6 +1,7 @@ # -*- coding: utf-8 -*- # pyright: strict import dataclasses +import sys import textwrap import typing as t @@ -14,6 +15,13 @@ from construct.lib.py3compat import bytestringtype, reprstring, unicodestringtyp from construct_typed.generic import Adapter, Construct, Context, ParsedType, PathType +# The `key_only` keyword is possible since python 3.10. To support it in python 3.8 & 3.9 +# the dataclasses module from python 3.10 is copied to this package. +if sys.version_info >= (3, 10) or t.TYPE_CHECKING: + import dataclasses +else: + import construct_typed.dataclasses_py310 as dataclasses + T = t.TypeVar("T") @@ -200,7 +208,7 @@ def _replace_this_struct(constr: "Construct[t.Any, t.Any]", replacement: t.Any) ) -@__dataclass_transform__(field_descriptors=(csfield,)) +@__dataclass_transform__(field_descriptors=(csfield,), kw_only_default=True) class DataclassStruct: """ Adapter for a dataclasses for optimised type hints / static autocompletion in comparision to the original Struct. @@ -241,7 +249,7 @@ class DataclassStruct: raise ValueError("`reverse_fields` parameter has to be an `bool` object") # create dataclass - cls = dataclasses.dataclass(cls) + dataclasses.dataclass(cls, kw_only=True) # type: ignore # create construct format dc_constr = DataclassConstruct(cls, reverse_fields) From a46c180e7a810e878f9d2888a856e3c0a495711a Mon Sep 17 00:00:00 2001 From: Tim Rid <6593626+timrid@users.noreply.github.com> Date: Sun, 20 Feb 2022 12:47:10 +0100 Subject: [PATCH 19/24] Dont parse the provided subcon in csfield. Only check if it builds from none. But now you can use the `const` or `default` parameters instead to provied default or constant values. --- construct_typed/dataclass_struct.py | 87 ++++++++++++++++++++--------- tests/test_typed.py | 38 +++++++------ 2 files changed, 82 insertions(+), 43 deletions(-) diff --git a/construct_typed/dataclass_struct.py b/construct_typed/dataclass_struct.py index 8b24286..4fc3f65 100644 --- a/construct_typed/dataclass_struct.py +++ b/construct_typed/dataclass_struct.py @@ -39,64 +39,97 @@ def __dataclass_transform__( DATACLASS_METADATA_KEY = "__construct_typed_subcon" -# specialisation for constructs, that builds from none and dont have to be declared in the __init__ method +# specialisation for constructs, that builds from none -> this field does not appear in the __init__ method and has a default of None @t.overload def csfield( subcon: "cs.Construct[ParsedType, None]", - doc: t.Optional[str] = None, - parsed: t.Optional[t.Callable[[t.Any, Context], None]] = None, + *, + doc: t.Optional[str] = ..., init: t.Literal[False] = False, -) -> ParsedType: + default: t.Literal[None] = None, + const: t.Literal[dataclasses.MISSING] = ..., +) -> t.Optional[ParsedType]: ... +# normal mode, when neither default nor const is defined -> this field does appear in the __init__ method but has no default value @t.overload def csfield( subcon: "Construct[ParsedType, t.Any]", - doc: t.Optional[str] = None, - parsed: t.Optional[t.Callable[[t.Any, Context], None]] = None, - init: bool = True, + *, + doc: t.Optional[str] = ..., + init: t.Literal[True] = ..., + default: t.Literal[dataclasses.MISSING] = ..., + const: t.Literal[dataclasses.MISSING] = ..., +) -> ParsedType: + ... + + +# specialisation when const parameter is set -> this field does not appear in the __init__ method but has a default value +@t.overload +def csfield( + subcon: "Construct[ParsedType, t.Any]", + *, + doc: t.Optional[str] = ..., + init: t.Literal[False] = ..., + default: t.Literal[dataclasses.MISSING] = ..., + const: t.Optional[ParsedType] = ..., +) -> ParsedType: + ... + + +# specialisation when default parameter is set -> this field does appear in the __init__ method and has a default value +@t.overload +def csfield( + subcon: "Construct[ParsedType, t.Any]", + *, + doc: t.Optional[str] = ..., + init: t.Literal[True] = ..., + default: t.Optional[ParsedType] = ..., + const: t.Literal[dataclasses.MISSING] = ..., ) -> ParsedType: ... def csfield( subcon: "Construct[ParsedType, t.Any]", + *, doc: t.Optional[str] = None, - parsed: t.Optional[t.Callable[[t.Any, Context], None]] = None, - init: bool = True, + init: bool = True, # dont use this parameter, this is only used for `dataclass_transform` + default: t.Optional[t.Any] = dataclasses.MISSING, + const: t.Optional[t.Any] = dataclasses.MISSING, ) -> ParsedType: """ Helper method for "DataclassStruct" and "DataclassBitStruct" to create the dataclass fields. This method also processes Const and Default, to pass these values als default values to the dataclass. + + Only one of the parameters `default` or `const` can be vaild. They are mutually exclusive. """ - orig_subcon = 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).strip("\n") - subcon = cs.Renamed(subcon, newdocs=doc, newparsed=parsed) + if (default is not dataclasses.MISSING) and (const is not dataclasses.MISSING): + raise ValueError("default and const are mutally exclusive") - if orig_subcon.flagbuildnone is True: + # Rename subcon, if doc is available + if doc is not None: + doc = textwrap.dedent(doc).strip("\n") + subcon = cs.Renamed(subcon, newdocs=doc) + + if default is not dataclasses.MISSING: + init = True + default = default + subcon = cs.Default(subcon, default) + elif const is not dataclasses.MISSING: + init = False + default = const + subcon = cs.Const(const, subcon) + elif subcon.flagbuildnone is True: init = False default = None else: init = True default = dataclasses.MISSING - # Set default values in case of special subcons - if isinstance(orig_subcon, cs.Const): - const_subcon: "cs.Const[t.Any, t.Any, t.Any, t.Any]" = orig_subcon - default = const_subcon.value - elif isinstance(orig_subcon, cs.Default): - default_subcon: "cs.Default[t.Any, t.Any, t.Any, t.Any]" = orig_subcon - if callable(default_subcon.value): - default = None # context lambda is only defined at parsing/building - else: - default = default_subcon.value - return t.cast( ParsedType, dataclasses.field( diff --git a/tests/test_typed.py b/tests/test_typed.py index ef0226b..9cf4fc5 100644 --- a/tests/test_typed.py +++ b/tests/test_typed.py @@ -12,28 +12,36 @@ from construct_typed import ( TFlags, ) -from .declarativeunittest import common, raises, setattrs +from tests.declarativeunittest import common, raises, setattrs def test_dataclass_const_default() -> None: class ConstDefaultTest(DataclassStruct): - const_bytes: bytes = csfield(cs.Const(b"BMP")) - const_int: int = csfield(cs.Const(5, cs.Int8ub)) - default_int: int = csfield(cs.Default(cs.Int8ub, 28)) - default_lambda: bytes = csfield( + const_bytes: bytes = csfield(cs.Bytes(3), const=b"BMP") + const_int: int = csfield(cs.Int8ub, const=5) + default_int: int = csfield(cs.Int8ub, default=26) + default_lambda: t.Optional[bytes] = csfield( cs.Default(cs.Bytes(cs.this.const_int), lambda ctx: bytes(ctx.const_int)) ) a = ConstDefaultTest() assert a.const_bytes == b"BMP" assert a.const_int == 5 - assert a.default_int == 28 + assert a.default_int == 26 assert a.default_lambda == None + a = ConstDefaultTest(default_int=1) + assert a.default_int == 1 + + format = ConstDefaultTest.__constr__() + assert isinstance(format.const_bytes.subcon, cs.Const) + assert isinstance(format.const_int.subcon, cs.Const) + assert isinstance(format.default_int.subcon, cs.Default) + assert isinstance(format.default_lambda.subcon, cs.Default) def test_dataclass_access() -> None: class TestTContainer(DataclassStruct): - a: t.Optional[int] = csfield(cs.Const(1, cs.Byte)) + a: int = csfield(cs.Byte, const=1) b: int = csfield(cs.Int8ub) tcontainer = TestTContainer(b=2) @@ -57,7 +65,7 @@ def test_dataclass_access() -> None: def test_dataclass_str_repr() -> None: class Image(DataclassStruct): - signature: t.Optional[bytes] = csfield(cs.Const(b"BMP")) + signature: bytes = csfield(cs.Bytes(3), const=b"BMP") width: int = csfield(cs.Int8ub) height: int = csfield(cs.Int8ub) @@ -137,8 +145,8 @@ def test_dataclass_struct_default_field() -> None: common( constr(Image), b"\x02\x03\x00\x00\x00\x00\x00\x00", - setattrs(Image(2, 3), pixels=bytes(6)), - sample_building=Image(2, 3), + setattrs(Image(width=2, height=3), pixels=bytes(6)), + sample_building=Image(width=2, height=3), ) @@ -191,7 +199,7 @@ def test_dataclass_struct_anonymus_fields_1() -> None: def test_dataclass_struct_anonymus_fields_2() -> None: class TestContainer(DataclassStruct): - _1: int = csfield(cs.Computed(7)) + _1: t.Optional[int] = csfield(cs.Computed(7)) _2: t.Optional[bytes] = csfield(cs.Const(b"JPEG")) _3: None = csfield(cs.Pass) _4: None = csfield(cs.Terminated) @@ -265,20 +273,18 @@ def test_dataclass_struct_wrong_container() -> None: a: int = csfield(cs.Int16ub) b: int = csfield(cs.Int8ub) - assert ( - raises(constr(TestContainer1).build, TestContainer2(a=1, b=2)) == TypeError - ) + assert raises(constr(TestContainer1).build, TestContainer2(a=1, b=2)) == TypeError def test_dataclass_struct_doc() -> None: class TestContainer(DataclassStruct): - a: int = csfield(cs.Int16ub, "This is the documentation of a") + a: int = csfield(cs.Int16ub, doc="This is the documentation of a") b: int = csfield( cs.Int8ub, doc="This is the documentation of b\nwhich is multiline" ) c: int = csfield( cs.Int8ub, - """ + doc=""" This is the documentation of c which is also multiline """, From 935f021a4fb466394634069db75963ef59c34dfe Mon Sep 17 00:00:00 2001 From: Tim Rid <6593626+timrid@users.noreply.github.com> Date: Sun, 20 Feb 2022 16:53:12 +0100 Subject: [PATCH 20/24] added a litte bit of documentation --- construct_typed/dataclass_struct.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/construct_typed/dataclass_struct.py b/construct_typed/dataclass_struct.py index 4fc3f65..1ed95e9 100644 --- a/construct_typed/dataclass_struct.py +++ b/construct_typed/dataclass_struct.py @@ -253,6 +253,11 @@ class DataclassStruct: Parses to a dataclasses.dataclass instance, and builds from such instance. Size is the sum of all subcon sizes, unless any subcon raises SizeofError. + Every construct that builds from None (eg. Const, Default, Index, Rebuild, Check, Checksum, ...) will automatically initialised with None. + If a default or const value should be used in a DataclassStruct the best is to use the `default` or `const` parameters of `csfield`. These + are internally used for creating a `Const` or `Default` construct but also adds the default/const value to the DataclassStruct while creating + it via __init__. + :param constr: This can be used if the structure is nested inside a Subconstruct. To represent this struct use the constant `this_struct`. :param reverse: Flag if the fields of the dataclass should be reversed From acc9f6a9e4a441fe75223d41a753f0ad798188ba Mon Sep 17 00:00:00 2001 From: Tim Rid <6593626+timrid@users.noreply.github.com> Date: Sun, 20 Feb 2022 17:07:41 +0100 Subject: [PATCH 21/24] removed this_struct and added a lambda instead --- construct_typed/__init__.py | 4 +-- construct_typed/dataclass_struct.py | 47 ++++++++++------------------- tests/test_typed.py | 23 ++++++++++++++ 3 files changed, 40 insertions(+), 34 deletions(-) diff --git a/construct_typed/__init__.py b/construct_typed/__init__.py index 56c0e57..537de1a 100644 --- a/construct_typed/__init__.py +++ b/construct_typed/__init__.py @@ -2,8 +2,7 @@ from construct_typed.generic import constr from construct_typed.dataclass_struct import ( DataclassBitStruct, DataclassStruct, - csfield, - this_struct + csfield ) from construct_typed.generic import ( Adapter, @@ -24,7 +23,6 @@ __all__ = [ "TEnumConstruct", "TFlags", "TFlagsConstruct", - "this_struct", "Adapter", "ConstantOrContextLambda", "Construct", diff --git a/construct_typed/dataclass_struct.py b/construct_typed/dataclass_struct.py index 1ed95e9..cee8c89 100644 --- a/construct_typed/dataclass_struct.py +++ b/construct_typed/dataclass_struct.py @@ -224,23 +224,6 @@ class DataclassConstruct(Adapter[t.Any, t.Any, T, T]): return ret_dict -# Helper object for defining the `constr` of a `struct`. Will be replaced with the proper construct, when class is created. -this_struct: Construct[t.Any, t.Any] = Construct() - - -def _replace_this_struct(constr: "Construct[t.Any, t.Any]", replacement: t.Any) -> None: - """Recursive search for `this_struct` in all SubConstructs and replace it with AttrsStruct""" - subcon = getattr(constr, "subcon", None) - if subcon is this_struct: - setattr(constr, "subcon", replacement) - elif subcon is not None: - _replace_this_struct(subcon, replacement) - else: - raise ValueError( - "Could not find `this_struct`. Only SubConstructs are supported" - ) - - @__dataclass_transform__(field_descriptors=(csfield,), kw_only_default=True) class DataclassStruct: """ @@ -258,7 +241,7 @@ class DataclassStruct: are internally used for creating a `Const` or `Default` construct but also adds the default/const value to the DataclassStruct while creating it via __init__. - :param constr: This can be used if the structure is nested inside a Subconstruct. To represent this struct use the constant `this_struct`. + :param constr: Lambda for creating the construct object. This can be used if the DataclassStruct is nested inside a Subconstruct. Eg. `lambda cls: cs.Bitwise(cls)`. :param reverse: Flag if the fields of the dataclass should be reversed Example:: @@ -276,13 +259,15 @@ class DataclassStruct: @classmethod def __init_subclass__( - cls, - constr: "cs.Construct[t.Any, t.Any]" = this_struct, + cls: t.Type[T], + constr: t.Callable[ + [DataclassConstruct[T]], Construct[t.Any, t.Any] + ] = lambda cls: cls, reverse_fields: bool = False, ) -> None: # validate types - if not isinstance(constr, cs.Construct): # type: ignore - raise ValueError("`constr` parameter has to be an `Construct` object") + if not callable(constr): + raise ValueError("`constr` parameter has to be a function or lambda") if not isinstance(reverse_fields, bool): # type: ignore raise ValueError("`reverse_fields` parameter has to be an `bool` object") @@ -290,14 +275,12 @@ class DataclassStruct: dataclasses.dataclass(cls, kw_only=True) # type: ignore # create construct format - dc_constr = DataclassConstruct(cls, reverse_fields) - if constr is this_struct: - constr = dc_constr - else: - _replace_this_struct(constr, dc_constr) + dc_constr = constr(DataclassConstruct(cls, reverse_fields)) + if not isinstance(dc_constr, cs.Construct): # type: ignore + raise ValueError("`constr` sould return a `Construct` object") # save construct format and make the class compatible to `Constructable` protocol - setattr(cls, "__constr__", lambda: constr) + setattr(cls, "__constr__", lambda: dc_constr) # the `construct` library is using the [] access internally, so struct objects # should also make this possible and not only via the dot access. @@ -379,8 +362,10 @@ class DataclassBitStruct(DataclassStruct): @classmethod def __init_subclass__( - cls, - constr: "cs.Construct[t.Any, t.Any]" = this_struct, + cls: t.Type[T], + constr: t.Callable[ + [DataclassConstruct[T]], Construct[t.Any, t.Any] + ] = lambda cls: cls, reverse_fields: bool = False, ) -> None: - cls = DataclassStruct.__init_subclass__.__func__(cls, cs.Bitwise(constr), reverse_fields) # type: ignore + DataclassStruct.__init_subclass__.__func__(cls, lambda cls: cs.Bitwise(constr(cls)), reverse_fields) # type: ignore diff --git a/tests/test_typed.py b/tests/test_typed.py index 9cf4fc5..24e29a0 100644 --- a/tests/test_typed.py +++ b/tests/test_typed.py @@ -301,6 +301,29 @@ def test_dataclass_struct_doc() -> None: ) +def test_dataclass_bitwise() -> None: + class TestContainer(DataclassStruct, constr=lambda cls: cs.Bitwise(cls)): + a: int = csfield(cs.BitsInteger(7)) + b: int = csfield(cs.Bit) + c: int = csfield(cs.BitsInteger(8)) + + common( + constr(TestContainer), + b"\xFD\x12", + TestContainer(a=0x7E, b=1, c=0x12), + 2, + ) + + # check __getattr__ + c = TestContainer.__constr__() + assert c.subcon.a.name == "a" + assert c.subcon.b.name == "b" + assert c.subcon.c.name == "c" + assert isinstance(c.subcon.a.subcon, cs.BitsInteger) + assert c.subcon.b.subcon is cs.Bit + assert isinstance(c.subcon.c.subcon, cs.BitsInteger) + + def test_dataclass_bitstruct() -> None: class TestContainer(DataclassBitStruct): a: int = csfield(cs.BitsInteger(7)) From c5d08f9b87d63c0a5f09d092a68481f283f884e8 Mon Sep 17 00:00:00 2001 From: Tim Rid <6593626+timrid@users.noreply.github.com> Date: Sun, 20 Feb 2022 20:33:28 +0100 Subject: [PATCH 22/24] add docstring to construct `docs` --- construct_typed/dataclass_struct.py | 8 ++++ construct_typed/tenum.py | 13 ++++++- tests/test_typed.py | 57 ++++++++++++++++++++++------- 3 files changed, 62 insertions(+), 16 deletions(-) diff --git a/construct_typed/dataclass_struct.py b/construct_typed/dataclass_struct.py index cee8c89..0b799ee 100644 --- a/construct_typed/dataclass_struct.py +++ b/construct_typed/dataclass_struct.py @@ -271,6 +271,11 @@ class DataclassStruct: if not isinstance(reverse_fields, bool): # type: ignore raise ValueError("`reverse_fields` parameter has to be an `bool` object") + # get documentation before creating the dataclass + docs = "" + if cls.__doc__ is not None: + docs = textwrap.dedent(cls.__doc__).strip("\n") + # create dataclass dataclasses.dataclass(cls, kw_only=True) # type: ignore @@ -279,6 +284,9 @@ class DataclassStruct: if not isinstance(dc_constr, cs.Construct): # type: ignore raise ValueError("`constr` sould return a `Construct` object") + # save docs + dc_constr.docs = docs + # save construct format and make the class compatible to `Constructable` protocol setattr(cls, "__constr__", lambda: dc_constr) diff --git a/construct_typed/tenum.py b/construct_typed/tenum.py index 83f1d8c..5cb9116 100644 --- a/construct_typed/tenum.py +++ b/construct_typed/tenum.py @@ -1,9 +1,10 @@ import enum +import textwrap import typing as t import construct as cs -from .generic import * +from construct_typed.generic import * T = t.TypeVar("T") @@ -26,6 +27,11 @@ class _EnumMeta(enum.EnumMeta): __namespace: t.Dict[str, t.Any], **kwargs: t.Any, ) -> T: + # get documentation before creating the enum + docs = "" + if "__doc__" in __namespace: + docs = textwrap.dedent(__namespace["__doc__"]).strip("\n") + # create new enum object cls: T = super().__new__(metacls, __name, __bases, __namespace) # type: ignore @@ -49,7 +55,10 @@ class _EnumMeta(enum.EnumMeta): elif TFlags in __bases: enum_constr = TFlagsConstruct(subcon, cls) # type: ignore else: - enum_constr = None + raise TypeError("neither `TEnum` nor `TFlags` in bases") + + # save documentation + enum_constr.docs = docs # save construct format and make the class compatible to `Constructable` protocol setattr(cls, "__constr__", lambda: enum_constr) # type: ignore diff --git a/tests/test_typed.py b/tests/test_typed.py index 24e29a0..f3d51d6 100644 --- a/tests/test_typed.py +++ b/tests/test_typed.py @@ -277,28 +277,39 @@ def test_dataclass_struct_wrong_container() -> None: def test_dataclass_struct_doc() -> None: - class TestContainer(DataclassStruct): - a: int = csfield(cs.Int16ub, doc="This is the documentation of a") - b: int = csfield( - cs.Int8ub, doc="This is the documentation of b\nwhich is multiline" - ) + class TestContainer1(DataclassStruct): + """ + Documentation of TestContainer + """ + + a: int = csfield(cs.Int16ub, doc="This is the doc of a") + b: int = csfield(cs.Int8ub, doc="This is the doc of b\nwhich is multiline") c: int = csfield( cs.Int8ub, doc=""" - This is the documentation of c + This is the doc of c which is also multiline """, ) - format = TestContainer.__constr__() - common(format, b"\x00\x01\x02\x03", TestContainer(a=1, b=2, c=3), 4) + format1 = TestContainer1.__constr__() + common(format1, b"\x00\x01\x02\x03", TestContainer1(a=1, b=2, c=3), 4) - assert format.subcon.a.docs == "This is the documentation of a" - assert format.subcon.b.docs == "This is the documentation of b\nwhich is multiline" - assert ( - format.subcon.c.docs - == "This is the documentation of c\nwhich is also multiline" - ) + assert format1.docs == "Documentation of TestContainer" + assert format1.subcon.a.docs == "This is the doc of a" + assert format1.subcon.b.docs == "This is the doc of b\nwhich is multiline" + assert format1.subcon.c.docs == "This is the doc of c\nwhich is also multiline" + + class TestContainer2(DataclassStruct): + a: int = csfield(cs.Int16ub) + b: int = csfield(cs.Int8ub) + c: int = csfield(cs.Int8ub) + + format2 = TestContainer2.__constr__() + assert format2.docs == "" + assert format2.subcon.a.docs == "" + assert format2.subcon.b.docs == "" + assert format2.subcon.c.docs == "" def test_dataclass_bitwise() -> None: @@ -367,6 +378,24 @@ def test_tenum() -> None: assert raises(d.build, 8) == TypeError +def test_tenum_doc() -> None: + class TestEnum1(TEnum, subcon=cs.Byte): + """ + TestEnum documentation + """ + + one = 1 + + d1 = constr(TestEnum1) + assert d1.docs == "TestEnum documentation" + + class TestEnum2(TEnum, subcon=cs.Byte): + two = 2 + + d2 = constr(TestEnum2) + assert d2.docs == "" + + def test_tenum_in_dataclass_struct() -> None: class TestEnum(TEnum, subcon=cs.Int8ub): a = 1 From af1a7799e9701eb2bb2c5a4ed3442dce7029b4cb Mon Sep 17 00:00:00 2001 From: Tim Rid <6593626+timrid@users.noreply.github.com> Date: Sun, 20 Feb 2022 20:42:36 +0100 Subject: [PATCH 23/24] make consistent namings in tests --- tests/test_typed.py | 231 ++++++++++++++++++++++---------------------- 1 file changed, 115 insertions(+), 116 deletions(-) diff --git a/tests/test_typed.py b/tests/test_typed.py index f3d51d6..ed578c8 100644 --- a/tests/test_typed.py +++ b/tests/test_typed.py @@ -16,7 +16,7 @@ from tests.declarativeunittest import common, raises, setattrs def test_dataclass_const_default() -> None: - class ConstDefaultTest(DataclassStruct): + class TestDataclass(DataclassStruct): const_bytes: bytes = csfield(cs.Bytes(3), const=b"BMP") const_int: int = csfield(cs.Int8ub, const=5) default_int: int = csfield(cs.Int8ub, default=26) @@ -24,43 +24,42 @@ def test_dataclass_const_default() -> None: cs.Default(cs.Bytes(cs.this.const_int), lambda ctx: bytes(ctx.const_int)) ) - a = ConstDefaultTest() - assert a.const_bytes == b"BMP" - assert a.const_int == 5 - assert a.default_int == 26 - assert a.default_lambda == None - a = ConstDefaultTest(default_int=1) - assert a.default_int == 1 + obj = TestDataclass() + assert obj.const_bytes == b"BMP" + assert obj.const_int == 5 + assert obj.default_int == 26 + assert obj.default_lambda == None + obj = TestDataclass(default_int=1) + assert obj.default_int == 1 - format = ConstDefaultTest.__constr__() - assert isinstance(format.const_bytes.subcon, cs.Const) - assert isinstance(format.const_int.subcon, cs.Const) - assert isinstance(format.default_int.subcon, cs.Default) - assert isinstance(format.default_lambda.subcon, cs.Default) + fmt = TestDataclass.__constr__() + assert isinstance(fmt.const_bytes.subcon, cs.Const) + assert isinstance(fmt.const_int.subcon, cs.Const) + assert isinstance(fmt.default_int.subcon, cs.Default) + assert isinstance(fmt.default_lambda.subcon, cs.Default) def test_dataclass_access() -> None: - class TestTContainer(DataclassStruct): + class TestDataclass(DataclassStruct): a: int = csfield(cs.Byte, const=1) b: int = csfield(cs.Int8ub) - tcontainer = TestTContainer(b=2) + obj = TestDataclass(b=2) - # tcontainer - assert tcontainer.a == 1 - assert tcontainer["a"] == 1 - assert tcontainer.b == 2 - assert tcontainer["b"] == 2 + assert obj.a == 1 + assert obj["a"] == 1 + assert obj.b == 2 + assert obj["b"] == 2 - tcontainer.a = 5 - assert tcontainer.a == 5 - assert tcontainer["a"] == 5 - tcontainer["a"] = 6 - assert tcontainer.a == 6 - assert tcontainer["a"] == 6 + obj.a = 5 + assert obj.a == 5 + assert obj["a"] == 5 + obj["a"] = 6 + assert obj.a == 6 + assert obj["a"] == 6 # wrong creation - assert raises(lambda: TestTContainer(a=0, b=1)) == TypeError # type: ignore + assert raises(lambda: TestDataclass(a=0, b=1)) == TypeError # type: ignore def test_dataclass_str_repr() -> None: @@ -69,13 +68,13 @@ def test_dataclass_str_repr() -> None: width: int = csfield(cs.Int8ub) height: int = csfield(cs.Int8ub) - format = constr(Image) + fmt = constr(Image) obj = Image(width=3, height=2) assert ( str(obj) == "Image: \n signature = b'BMP' (total 3)\n width = 3\n height = 2" ) - obj = format.parse(format.build(obj)) + obj = fmt.parse(fmt.build(obj)) assert ( str(obj) == "Image: \n signature = b'BMP' (total 3)\n width = 3\n height = 2" @@ -95,28 +94,28 @@ def test_dataclass_struct() -> None: ) # check __getattr__ - c = Image.__constr__() - assert c.width.name == "width" - assert c.height.name == "height" - assert c.width.subcon is cs.Int8ub - assert c.height.subcon is cs.Int8ub + fmt = Image.__constr__() + assert fmt.width.name == "width" + assert fmt.height.name == "height" + assert fmt.width.subcon is cs.Int8ub + assert fmt.height.subcon is cs.Int8ub def test_dataclass_struct_reverse() -> None: - class TestContainer(DataclassStruct, reverse_fields=True): + class TestDataclass(DataclassStruct, reverse_fields=True): a: int = csfield(cs.Int16ub) b: int = csfield(cs.Int8ub) common( - constr(TestContainer), + constr(TestDataclass), b"\x02\x00\x01", - TestContainer(a=1, b=2), + TestDataclass(a=1, b=2), 3, ) def test_dataclass_struct_nested() -> None: - class TestContainer(DataclassStruct): + class TestDataclass(DataclassStruct): class InnerDataclass(DataclassStruct): b: int = csfield(cs.Byte) c: bytes = csfield(cs.Bytes(cs.this._.length)) @@ -125,9 +124,9 @@ def test_dataclass_struct_nested() -> None: a: InnerDataclass = csfield(constr(InnerDataclass)) common( - constr(TestContainer), + constr(TestDataclass), b"\x02\x01\xF1\xF2", - TestContainer(length=2, a=TestContainer.InnerDataclass(b=1, c=b"\xF1\xF2")), + TestDataclass(length=2, a=TestDataclass.InnerDataclass(b=1, c=b"\xF1\xF2")), ) @@ -151,67 +150,67 @@ def test_dataclass_struct_default_field() -> None: def test_dataclass_struct_const_field() -> None: - class TestContainer(DataclassStruct): + class TestDataclass(DataclassStruct): const_field: t.Optional[bytes] = csfield(cs.Const(b"\x00")) common( - constr(TestContainer), + constr(TestDataclass), bytes(1), - setattrs(TestContainer(), const_field=b"\x00"), + setattrs(TestDataclass(), const_field=b"\x00"), 1, ) assert ( raises( - constr(TestContainer).build, - setattrs(TestContainer(), const_field=b"\x01"), + constr(TestDataclass).build, + setattrs(TestDataclass(), const_field=b"\x01"), ) == cs.ConstError ) def test_dataclass_struct_array_field() -> None: - class TestContainer(DataclassStruct): + class TestDataclass(DataclassStruct): array_field: t.List[int] = csfield(cs.Array(5, cs.Int8ub)) common( - constr(TestContainer), + constr(TestDataclass), bytes(5), - TestContainer(array_field=[0, 0, 0, 0, 0]), + TestDataclass(array_field=[0, 0, 0, 0, 0]), 5, ) def test_dataclass_struct_anonymus_fields_1() -> None: - class TestContainer(DataclassStruct): + class TestDataclass(DataclassStruct): _1: t.Optional[bytes] = csfield(cs.Const(b"\x00")) _2: None = csfield(cs.Padding(1)) _3: None = csfield(cs.Pass) _4: None = csfield(cs.Terminated) common( - constr(TestContainer), + constr(TestDataclass), bytes(2), - setattrs(TestContainer(), _1=b"\x00"), + setattrs(TestDataclass(), _1=b"\x00"), cs.SizeofError, ) def test_dataclass_struct_anonymus_fields_2() -> None: - class TestContainer(DataclassStruct): + class TestDataclass(DataclassStruct): _1: t.Optional[int] = csfield(cs.Computed(7)) _2: t.Optional[bytes] = csfield(cs.Const(b"JPEG")) _3: None = csfield(cs.Pass) _4: None = csfield(cs.Terminated) - d = constr(TestContainer) - assert d.build(TestContainer()) == d.build(TestContainer()) + fmt = constr(TestDataclass) + assert fmt.build(TestDataclass()) == fmt.build(TestDataclass()) def test_dataclass_struct_overloaded_method() -> None: # Test dot access to some names that are not accessable via dot # in the original 'cs.Container'. - class TestContainer(DataclassStruct): + class TestDataclass(DataclassStruct): clear: int = csfield(cs.Int8ul) copy: int = csfield(cs.Int8ul) fromkeys: int = csfield(cs.Int8ul) @@ -227,10 +226,10 @@ def test_dataclass_struct_overloaded_method() -> None: update: int = csfield(cs.Int8ul) values: int = csfield(cs.Int8ul) - d = constr(TestContainer) - obj = d.parse( - d.build( - TestContainer( + fmt = constr(TestDataclass) + obj = fmt.parse( + fmt.build( + TestDataclass( clear=1, copy=2, fromkeys=3, @@ -277,9 +276,9 @@ def test_dataclass_struct_wrong_container() -> None: def test_dataclass_struct_doc() -> None: - class TestContainer1(DataclassStruct): + class TestDataclass1(DataclassStruct): """ - Documentation of TestContainer + Documentation of TestDataclass1 """ a: int = csfield(cs.Int16ub, doc="This is the doc of a") @@ -292,70 +291,70 @@ def test_dataclass_struct_doc() -> None: """, ) - format1 = TestContainer1.__constr__() - common(format1, b"\x00\x01\x02\x03", TestContainer1(a=1, b=2, c=3), 4) + fmt1 = TestDataclass1.__constr__() + common(fmt1, b"\x00\x01\x02\x03", TestDataclass1(a=1, b=2, c=3), 4) - assert format1.docs == "Documentation of TestContainer" - assert format1.subcon.a.docs == "This is the doc of a" - assert format1.subcon.b.docs == "This is the doc of b\nwhich is multiline" - assert format1.subcon.c.docs == "This is the doc of c\nwhich is also multiline" + assert fmt1.docs == "Documentation of TestDataclass1" + assert fmt1.subcon.a.docs == "This is the doc of a" + assert fmt1.subcon.b.docs == "This is the doc of b\nwhich is multiline" + assert fmt1.subcon.c.docs == "This is the doc of c\nwhich is also multiline" - class TestContainer2(DataclassStruct): + class TestDataclass2(DataclassStruct): a: int = csfield(cs.Int16ub) b: int = csfield(cs.Int8ub) c: int = csfield(cs.Int8ub) - format2 = TestContainer2.__constr__() - assert format2.docs == "" - assert format2.subcon.a.docs == "" - assert format2.subcon.b.docs == "" - assert format2.subcon.c.docs == "" + fmt2 = TestDataclass2.__constr__() + assert fmt2.docs == "" + assert fmt2.subcon.a.docs == "" + assert fmt2.subcon.b.docs == "" + assert fmt2.subcon.c.docs == "" def test_dataclass_bitwise() -> None: - class TestContainer(DataclassStruct, constr=lambda cls: cs.Bitwise(cls)): + class TestDataclass(DataclassStruct, constr=lambda cls: cs.Bitwise(cls)): a: int = csfield(cs.BitsInteger(7)) b: int = csfield(cs.Bit) c: int = csfield(cs.BitsInteger(8)) common( - constr(TestContainer), + constr(TestDataclass), b"\xFD\x12", - TestContainer(a=0x7E, b=1, c=0x12), + TestDataclass(a=0x7E, b=1, c=0x12), 2, ) # check __getattr__ - c = TestContainer.__constr__() - assert c.subcon.a.name == "a" - assert c.subcon.b.name == "b" - assert c.subcon.c.name == "c" - assert isinstance(c.subcon.a.subcon, cs.BitsInteger) - assert c.subcon.b.subcon is cs.Bit - assert isinstance(c.subcon.c.subcon, cs.BitsInteger) + fmt = TestDataclass.__constr__() + assert fmt.subcon.a.name == "a" + assert fmt.subcon.b.name == "b" + assert fmt.subcon.c.name == "c" + assert isinstance(fmt.subcon.a.subcon, cs.BitsInteger) + assert fmt.subcon.b.subcon is cs.Bit + assert isinstance(fmt.subcon.c.subcon, cs.BitsInteger) def test_dataclass_bitstruct() -> None: - class TestContainer(DataclassBitStruct): + class TestDataclass(DataclassBitStruct): a: int = csfield(cs.BitsInteger(7)) b: int = csfield(cs.Bit) c: int = csfield(cs.BitsInteger(8)) common( - constr(TestContainer), + constr(TestDataclass), b"\xFD\x12", - TestContainer(a=0x7E, b=1, c=0x12), + TestDataclass(a=0x7E, b=1, c=0x12), 2, ) # check __getattr__ - c = TestContainer.__constr__() - assert c.subcon.a.name == "a" - assert c.subcon.b.name == "b" - assert c.subcon.c.name == "c" - assert isinstance(c.subcon.a.subcon, cs.BitsInteger) - assert c.subcon.b.subcon is cs.Bit - assert isinstance(c.subcon.c.subcon, cs.BitsInteger) + fmt = TestDataclass.__constr__() + assert fmt.subcon.a.name == "a" + assert fmt.subcon.b.name == "b" + assert fmt.subcon.c.name == "c" + assert isinstance(fmt.subcon.a.subcon, cs.BitsInteger) + assert fmt.subcon.b.subcon is cs.Bit + assert isinstance(fmt.subcon.c.subcon, cs.BitsInteger) def test_tenum() -> None: @@ -365,17 +364,17 @@ def test_tenum() -> None: four = 4 eight = 8 - d = constr(TestEnum) + fmt = constr(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 + common(fmt, b"\x01", TestEnum.one, 1) + common(fmt, b"\xff", TestEnum(255), 1) + assert fmt.parse(b"\x01") == TestEnum.one + assert fmt.parse(b"\x01") == 1 + assert int(fmt.parse(b"\x01")) == 1 + assert fmt.parse(b"\xff") == TestEnum(255) + assert fmt.parse(b"\xff") == 255 + assert int(fmt.parse(b"\xff")) == 255 + assert raises(fmt.build, 8) == TypeError def test_tenum_doc() -> None: @@ -401,19 +400,19 @@ def test_tenum_in_dataclass_struct() -> None: a = 1 b = 2 - class TestContainer(DataclassStruct): + class TestDataclass(DataclassStruct): a: TestEnum = csfield(constr(TestEnum)) b: int = csfield(cs.Int8ub) common( - constr(TestContainer), + constr(TestDataclass), b"\x01\x02", - TestContainer(a=TestEnum.a, b=2), + TestDataclass(a=TestEnum.a, b=2), 2, ) assert ( - raises(constr(TestEnum).build, TestContainer(a=1, b=2)) == TypeError # type: ignore + raises(constr(TestEnum).build, TestDataclass(a=1, b=2)) == TypeError # type: ignore ) @@ -424,12 +423,12 @@ def test_tflags() -> None: four = 4 eight = 8 - d = constr(TestFlags) - common(d, b"\x03", TestFlags.one | TestFlags.two, 1) - assert d.build(TestFlags(0)) == b"\x00" - assert d.build(TestFlags.one | TestFlags.two) == b"\x03" - assert d.build(TestFlags(8)) == b"\x08" - assert d.build(TestFlags(1 | 2)) == b"\x03" - assert d.build(TestFlags(255)) == b"\xff" - assert d.build(TestFlags.eight) == b"\x08" - assert raises(d.build, 2) == TypeError + fmt = constr(TestFlags) + common(fmt, b"\x03", TestFlags.one | TestFlags.two, 1) + assert fmt.build(TestFlags(0)) == b"\x00" + assert fmt.build(TestFlags.one | TestFlags.two) == b"\x03" + assert fmt.build(TestFlags(8)) == b"\x08" + assert fmt.build(TestFlags(1 | 2)) == b"\x03" + assert fmt.build(TestFlags(255)) == b"\xff" + assert fmt.build(TestFlags.eight) == b"\x08" + assert raises(fmt.build, 2) == TypeError From e730fd8d879466ecc7aecedccb5c11c89b8064a8 Mon Sep 17 00:00:00 2001 From: Tim Rid <6593626+timrid@users.noreply.github.com> Date: Sun, 20 Feb 2022 21:14:05 +0100 Subject: [PATCH 24/24] use a special MISSING enum, because typing.Literal only accepts enums --- construct_typed/dataclass_struct.py | 40 +++++++++++++++++------------ 1 file changed, 23 insertions(+), 17 deletions(-) diff --git a/construct_typed/dataclass_struct.py b/construct_typed/dataclass_struct.py index 0b799ee..2e23dbc 100644 --- a/construct_typed/dataclass_struct.py +++ b/construct_typed/dataclass_struct.py @@ -4,6 +4,7 @@ import dataclasses import sys import textwrap import typing as t +import enum import construct as cs from construct.lib.containers import ( @@ -37,17 +38,22 @@ def __dataclass_transform__( return lambda a: a +# this is nessesary, because typing.Literal type has to be an enum +class Flag(enum.Enum): + MISSING = dataclasses.MISSING + + DATACLASS_METADATA_KEY = "__construct_typed_subcon" # specialisation for constructs, that builds from none -> this field does not appear in the __init__ method and has a default of None @t.overload -def csfield( - subcon: "cs.Construct[ParsedType, None]", +def csfield( # type: ignore + subcon: "Construct[ParsedType, None]", *, doc: t.Optional[str] = ..., - init: t.Literal[False] = False, - default: t.Literal[None] = None, - const: t.Literal[dataclasses.MISSING] = ..., + init: t.Literal[False] = ..., + default: t.Literal[None] = ..., + const: t.Literal[Flag.MISSING] = ..., ) -> t.Optional[ParsedType]: ... @@ -59,8 +65,8 @@ def csfield( *, doc: t.Optional[str] = ..., init: t.Literal[True] = ..., - default: t.Literal[dataclasses.MISSING] = ..., - const: t.Literal[dataclasses.MISSING] = ..., + default: t.Literal[Flag.MISSING] = ..., + const: t.Literal[Flag.MISSING] = ..., ) -> ParsedType: ... @@ -72,7 +78,7 @@ def csfield( *, doc: t.Optional[str] = ..., init: t.Literal[False] = ..., - default: t.Literal[dataclasses.MISSING] = ..., + default: t.Literal[Flag.MISSING] = ..., const: t.Optional[ParsedType] = ..., ) -> ParsedType: ... @@ -86,7 +92,7 @@ def csfield( doc: t.Optional[str] = ..., init: t.Literal[True] = ..., default: t.Optional[ParsedType] = ..., - const: t.Literal[dataclasses.MISSING] = ..., + const: t.Literal[Flag.MISSING] = ..., ) -> ParsedType: ... @@ -95,10 +101,10 @@ def csfield( subcon: "Construct[ParsedType, t.Any]", *, doc: t.Optional[str] = None, - init: bool = True, # dont use this parameter, this is only used for `dataclass_transform` - default: t.Optional[t.Any] = dataclasses.MISSING, - const: t.Optional[t.Any] = dataclasses.MISSING, -) -> ParsedType: + init: bool = True, # dont use `init`, this is only used for `dataclass_transform` + default: t.Optional[t.Any] = Flag.MISSING, + const: t.Optional[t.Any] = Flag.MISSING, +) -> t.Optional[ParsedType]: """ Helper method for "DataclassStruct" and "DataclassBitStruct" to create the dataclass fields. @@ -107,7 +113,7 @@ def csfield( Only one of the parameters `default` or `const` can be vaild. They are mutually exclusive. """ - if (default is not dataclasses.MISSING) and (const is not dataclasses.MISSING): + if (default is not Flag.MISSING) and (const is not Flag.MISSING): raise ValueError("default and const are mutally exclusive") # Rename subcon, if doc is available @@ -115,11 +121,11 @@ def csfield( doc = textwrap.dedent(doc).strip("\n") subcon = cs.Renamed(subcon, newdocs=doc) - if default is not dataclasses.MISSING: + if default is not Flag.MISSING: init = True default = default subcon = cs.Default(subcon, default) - elif const is not dataclasses.MISSING: + elif const is not Flag.MISSING: init = False default = const subcon = cs.Const(const, subcon) @@ -131,7 +137,7 @@ def csfield( default = dataclasses.MISSING return t.cast( - ParsedType, + t.Optional[ParsedType], dataclasses.field( default=default, init=init,