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))