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"