diff --git a/construct_typed/__init__.py b/construct_typed/__init__.py index d74e136..374d2cf 100644 --- a/construct_typed/__init__.py +++ b/construct_typed/__init__.py @@ -1,7 +1,7 @@ from enum import IntEnum from typing import Any, Type, Dict, TYPE_CHECKING, TypeVar, Union import typing -from construct.core import Construct, Adapter, Struct +from construct.core import Construct, Adapter, Struct, BitStruct from construct.lib.containers import Container @@ -95,12 +95,12 @@ else: if TYPE_CHECKING: class TypedContainer(Container[Any]): ... + ContainerType = TypeVar("ContainerType", bound=TypedContainer) else: class TypedContainer(Container): pass if TYPE_CHECKING: - ContainerType = TypeVar("ContainerType", bound=TypedContainer) class TypedStruct(Adapter[Container[Any], Dict[str, Any], ContainerType, Dict[str, Any]]): def __init__(self, container_type: Type[ContainerType], swapped: bool = False) -> None: ... else: @@ -109,6 +109,7 @@ else: if not issubclass(container_type, TypedContainer): raise TypeError("the subcon has to be a TypedContainer") self.container_type = container_type + self.swapped = swapped # extract the construct formats from the struct_type subcons = {} @@ -123,12 +124,48 @@ else: super(TypedStruct, self).__init__(Struct(**subcons)) def _decode(self, obj, context, path): - return self.container_type(obj) + if self.swapped: + return self.container_type({k: v for k, v in reversed(list(obj.items()))}) + else: + return self.container_type(obj) + + def _encode(self, obj, context, path): + return obj + +if TYPE_CHECKING: + class TypedBitStruct(Adapter[Container[Any], Dict[str, Any], ContainerType, Dict[str, Any]]): + def __init__(self, container_type: Type[ContainerType], swapped: bool = False) -> None: ... +else: + class TypedBitStruct(Adapter): + def __init__(self, container_type, swapped = False): + if not issubclass(container_type, TypedContainer): + raise TypeError("the subcon has to be a TypedContainer") + self.container_type = container_type + self.swapped = swapped + + # extract the construct formats from the struct_type + subcons = {} + subcon_formats = typing.get_type_hints(container_type) + subcon_items = subcon_formats.items() + if swapped: + subcon_items = reversed(subcon_items) + for subcon_name, subcon_format in subcon_items: + subcons[subcon_name] = subcon_format + + # init Adatper with a Struct as subcon + super(TypedBitStruct, self).__init__(BitStruct(**subcons)) + + def _decode(self, obj, context, path): + if self.swapped: + return self.container_type({k: v for k, v in reversed(list(obj.items()))}) + else: + return self.container_type(obj) def _encode(self, obj, context, path): return obj - # TODO: TypedUnion -# TODO: AnonymSubcon: \ No newline at end of file +# TODO: TypedLazyStruct +# TODO: TypedSequence: Based on typing.namedtuple +# TODO: FocusedSeq: Based on typing.namedtuple \ No newline at end of file diff --git a/tests/test_typed.py b/tests/test_typed.py index f5011cc..1ef4e24 100644 --- a/tests/test_typed.py +++ b/tests/test_typed.py @@ -25,17 +25,15 @@ def test_typed_struct(): b: Subcon(Int8ub) common(TypedStruct(Container1), b"\x00\x01\x02", Container1(a=1, b=2), 3) - - class Container1Swapped(TypedContainer): - a: Subcon(Int16ub) - b: Subcon(Int8ub) - common( - TypedStruct(Container1Swapped, swapped=True), + TypedStruct(Container1, swapped=True), b"\x02\x00\x01", Container(a=1, b=2), 3, ) + normal = TypedStruct(Container1) + swapped = TypedStruct(Container1, swapped=True) + assert str(normal.parse(b"\x00\x01\x02")) == str(swapped.parse(b"\x02\x00\x01")) class Container2(TypedContainer): class InnerContainer(TypedContainer): @@ -81,6 +79,10 @@ def test_typed_struct(): assert d.build({}) == d.build({}) +def test_typed_bit_struct(): + assert False + + def test_enum(): class E(enum.IntEnum): a = 1