Compare commits

..

25 commits

Author SHA1 Message Date
Tim Rid
e730fd8d87 use a special MISSING enum, because typing.Literal only accepts enums 2022-02-20 21:14:05 +01:00
Tim Rid
af1a7799e9 make consistent namings in tests 2022-02-20 20:42:36 +01:00
Tim Rid
c5d08f9b87 add docstring to construct docs 2022-02-20 20:33:28 +01:00
Tim Rid
acc9f6a9e4 removed this_struct and added a lambda instead 2022-02-20 17:07:41 +01:00
Tim Rid
935f021a4f added a litte bit of documentation 2022-02-20 16:53:12 +01:00
Tim Rid
a46c180e7a 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. 2022-02-20 12:47:10 +01:00
Tim Rid
cab9a07c4f 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. 2022-02-20 12:32:49 +01:00
Tim Rid
8529bd2878 ignore typing issues and removed unnessesary imports 2022-02-19 21:56:50 +01:00
Tim Rid
46fc3dd48c renamed file 2022-02-19 21:55:31 +01:00
Tim Rid
f9e2d24d71 Use marker instances from the original dataclass module so that hopefully also the original dataclass functions are working 2022-02-19 21:55:16 +01:00
Tim Rid
a0cadcce85 removed things that not work in Python 3.8 2022-02-19 21:54:04 +01:00
Tim Rid
b21c43ca65 added dataclass module from python 3.10 2022-02-19 21:27:52 +01:00
Tim Rid
75dbd8d822 fixed some mypy issues 2022-02-19 14:10:18 +01:00
Tim Rid
be5ae240fb renamed construct to constr so that there are no naming collisions with the construct package name 2022-02-19 12:46:28 +01:00
Tim Rid
94a8097bec added this_struct to public interface 2022-02-13 19:48:00 +01:00
Tim Rid
e99dd5d752 some renaming 2022-02-13 19:30:35 +01:00
Tim Rid
8525c04165 added custom _EnumMeta, because __init_subclass__ is not working correctly together with the standard enum.EnumMeta... 2022-02-13 19:29:56 +01:00
Tim Rid
7c68aeecd4 adapted all tests to the new api 2022-02-13 19:01:10 +01:00
Tim Rid
96ef565044 Merge branch 'main' into feature/new-structure 2022-02-13 17:03:38 +01:00
Tim Rid
de76415ac5 changed vscode test configuration to use the native testing api instead of the Test-Explorer 2022-02-13 17:02:46 +01:00
Tim Rid
9a9b4ab95a fixed type error 2022-02-13 17:00:28 +01:00
Tim Rid
8169f0ed31 Changed implementation of TEnum and TEnumFlags. 2022-02-13 16:58:54 +01:00
Tim Rid
b896457f90 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. 2022-02-13 16:32:30 +01:00
Tim Rid
753e4282ee renamed generic_wrapper.py to generic.py and removed relative imports 2022-02-13 16:13:11 +01:00
Tim Rid
db9b35a0c2 added Constructable Protocol 2022-02-13 16:04:27 +01:00
26 changed files with 2709 additions and 1637 deletions

View file

@ -1,10 +1,6 @@
name: CI name: CI
on: on: [push, pull_request]
push:
pull_request:
workflow_dispatch:
workflow_call:
jobs: jobs:
build: build:
@ -12,7 +8,7 @@ jobs:
strategy: strategy:
matrix: matrix:
os: ['ubuntu-latest', 'windows-latest'] os: ['ubuntu-latest', 'windows-latest']
python-version: [ '3.9', '3.10', '3.11', '3.12', '3.13' ] python-version: [ '3.7', '3.8', '3.9' ]
runs-on: ${{ matrix.os }} runs-on: ${{ matrix.os }}
name: OS ${{ matrix.os }}, Python ${{ matrix.python-version }} name: OS ${{ matrix.os }}, Python ${{ matrix.python-version }}
@ -20,26 +16,25 @@ jobs:
steps: steps:
# Checks out a copy of your repository on the machine # Checks out a copy of your repository on the machine
- name: Checkout code - name: Checkout code
uses: actions/checkout@v3 uses: actions/checkout@v1
# Setup python # Setup python
- name: Setup python - name: Setup python
uses: actions/setup-python@v4 uses: actions/setup-python@v1
with: with:
python-version: ${{ matrix.python-version }} python-version: ${{ matrix.python-version }}
architecture: x64 architecture: x64
# Setup node.js (for pyright) # Setup node.js (for pyright)
- name: Setup node.js (for pyright) - name: Setup node.js (for pyright)
uses: actions/setup-node@v3 uses: actions/setup-node@v2
with: with:
node-version: 16 node-version: '14'
# Install pyright # Install pyright
- name: Install pyright - name: Install pyright
run: | run: |
npm install -g pyright npm install -g pyright
pyright --version
# Install this package # Install this package
- name: Install this package - name: Install this package
@ -66,30 +61,3 @@ jobs:
- name: Run pyright - name: Run pyright
run: | run: |
pyright pyright
create_wheel_and_sdist:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v3
- name: Set up Python
uses: actions/setup-python@v4
with:
python-version: '3.13'
architecture: x64
- name: Install dependencies
run: |
python -m pip install --upgrade pip
pip install wheel build
- name: Build wheel and sdist
run: |
python -m build
- name: Upload wheel and sdist as artifact
uses: actions/upload-artifact@v4
with:
name: Package-Distributions-construct-typing
path: dist/

View file

@ -1,3 +1,6 @@
# This workflows will upload a Python Package using Twine when a release is created
# For more information see: https://help.github.com/en/actions/language-and-framework-guides/using-python-with-github-actions#publishing-to-package-registries
name: Upload Python Package name: Upload Python Package
on: on:
@ -5,26 +8,24 @@ on:
types: [created] types: [created]
jobs: jobs:
create_wheel_and_sdist:
name: create_wheel_and_sdist
uses: ./.github/workflows/main.yml
deploy: deploy:
needs: [ create_wheel_and_sdist ]
runs-on: ubuntu-latest runs-on: ubuntu-latest
environment: pypi
permissions:
id-token: write # IMPORTANT: this permission is mandatory for Trusted Publishing
steps: steps:
- uses: actions/checkout@v3 - uses: actions/checkout@v2
- name: Set up Python
- name: Download artifacts uses: actions/setup-python@v2
uses: actions/download-artifact@v4
with: with:
name: Package-Distributions-construct-typing python-version: '3.x'
path: ./dist - name: Install dependencies
run: |
- name: Publish package distributions to PyPI python -m pip install --upgrade pip
uses: pypa/gh-action-pypi-publish@release/v1 pip install setuptools wheel twine
- name: Build and publish
env:
TWINE_USERNAME: ${{ secrets.PYPI_USERNAME }}
TWINE_PASSWORD: ${{ secrets.PYPI_PASSWORD }}
run: |
python setup.py sdist bdist_wheel
twine upload dist/*

3
.gitignore vendored
View file

@ -129,6 +129,3 @@ dmypy.json
example_737 example_737
example_888 example_888
example_ksy.ksy example_ksy.ksy
# Test stuff
devtest/

6
.vscode/launch.json vendored
View file

@ -9,12 +9,6 @@
"type": "python", "type": "python",
"request": "launch", "request": "launch",
"program": "${file}", "program": "${file}",
"console": "integratedTerminal"
},
{
"name": "Debug Tests",
"type": "python",
"request": "test",
"console": "integratedTerminal", "console": "integratedTerminal",
"justMyCode": false "justMyCode": false
} }

21
.vscode/settings.json vendored
View file

@ -1,6 +1,14 @@
{ {
// static analysis
"python.languageServer": "Pylance", "python.languageServer": "Pylance",
// configure code formating
"python.formatting.provider": "black",
"python.sortImports.path": "isort",
"python.sortImports.args": [
"--profile=black",
],
// configure pylance
"python.analysis.typeCheckingMode": "strict", "python.analysis.typeCheckingMode": "strict",
"python.analysis.autoImportCompletions": false, "python.analysis.autoImportCompletions": false,
"python.analysis.diagnosticSeverityOverrides": { "python.analysis.diagnosticSeverityOverrides": {
@ -8,16 +16,7 @@
"reportUntypedNamedTuple": "information", "reportUntypedNamedTuple": "information",
}, },
// formating // configure pytest
"python.formatting.provider": "black",
// sorting
"python.sortImports.path": "isort",
"python.sortImports.args": [
"--profile=black",
],
// tests
"python.testing.pytestArgs": [ "python.testing.pytestArgs": [
"tests" "tests"
], ],

View file

@ -1,12 +1,3 @@
## Modified version of "construct-typing" module used in my projects.
This modification features:
- **[EnhancedDataclassMixin](https://github.com/waszil/construct-typing/commit/479b51344bfd95149596a75ee574ac2e63c032df) with additional features**
- **ConstantOrContextLambda2 type**
- **Typing for Subconstruct**
- **Type hint for Computed**
- **Switch typing fixes**
The original README.md file was described down below:
# construct-typing # construct-typing
[![PyPI](https://img.shields.io/pypi/v/construct-typing)](https://pypi.org/project/construct-typing/) [![PyPI](https://img.shields.io/pypi/v/construct-typing)](https://pypi.org/project/construct-typing/)
![PyPI - Implementation](https://img.shields.io/pypi/implementation/construct-typing) ![PyPI - Implementation](https://img.shields.io/pypi/implementation/construct-typing)

View file

@ -3,10 +3,6 @@ from construct.debug import *
from construct.expr import * from construct.expr import *
from construct.lib import * from construct.lib import *
from construct.version import * from construct.version import *
from construct import lib
__author__: str
__version__: str
#=============================================================================== #===============================================================================
# exposed names # exposed names

File diff suppressed because it is too large Load diff

View file

@ -1,3 +1,4 @@
import operator
import typing as t import typing as t
from construct.core import * from construct.core import *
@ -469,7 +470,7 @@ class ExprMixin(t.Generic[ReturnType], object):
@t.overload @t.overload
def __eq__(self: ExprMixin[float], other: ConstOrCallable[float]) -> BinExpr[bool]: ... def __eq__(self: ExprMixin[float], other: ConstOrCallable[float]) -> BinExpr[bool]: ...
@t.overload @t.overload
def __eq__(self, other: ConstOrCallable[t.Any]) -> BinExpr[t.Any]: ... # type: ignore def __eq__(self, other: t.Any) -> BinExpr[t.Any]: ...
# __ne__ ########################################################################################################### # __ne__ ###########################################################################################################
@t.overload @t.overload
@ -487,7 +488,7 @@ class ExprMixin(t.Generic[ReturnType], object):
@t.overload @t.overload
def __ne__(self: ExprMixin[float], other: ConstOrCallable[float]) -> BinExpr[bool]: ... def __ne__(self: ExprMixin[float], other: ConstOrCallable[float]) -> BinExpr[bool]: ...
@t.overload @t.overload
def __ne__(self, other: t.Any) -> BinExpr[t.Any]: ... # type: ignore def __ne__(self, other: t.Any) -> BinExpr[t.Any]: ...
# __neg__ ########################################################################################################## # __neg__ ##########################################################################################################
@t.overload @t.overload
@ -497,7 +498,7 @@ class ExprMixin(t.Generic[ReturnType], object):
@t.overload @t.overload
def __neg__(self: ExprMixin[float]) -> BinExpr[float]: ... def __neg__(self: ExprMixin[float]) -> BinExpr[float]: ...
@t.overload @t.overload
def __neg__(self) -> BinExpr[t.Any]: ... def __neg__(self) -> UniExpr[t.Any]: ...
# __pos__ ########################################################################################################## # __pos__ ##########################################################################################################
@t.overload @t.overload
@ -507,7 +508,7 @@ class ExprMixin(t.Generic[ReturnType], object):
@t.overload @t.overload
def __pos__(self: ExprMixin[float]) -> BinExpr[float]: ... def __pos__(self: ExprMixin[float]) -> BinExpr[float]: ...
@t.overload @t.overload
def __pos__(self) -> BinExpr[t.Any]: ... def __pos__(self) -> UniExpr[t.Any]: ...
# __invert__ ####################################################################################################### # __invert__ #######################################################################################################
@t.overload @t.overload
@ -515,7 +516,7 @@ class ExprMixin(t.Generic[ReturnType], object):
@t.overload @t.overload
def __invert__(self: ExprMixin[bool]) -> BinExpr[int]: ... def __invert__(self: ExprMixin[bool]) -> BinExpr[int]: ...
@t.overload @t.overload
def __invert__(self) -> BinExpr[t.Any]: ... def __invert__(self) -> UniExpr[t.Any]: ...
# __inv__ ########################################################################################################## # __inv__ ##########################################################################################################
def __inv__(self) -> UniExpr[t.Any]: ... def __inv__(self) -> UniExpr[t.Any]: ...
@ -542,7 +543,7 @@ class Path2(ExprMixin[ReturnType]):
class FuncPath(ExprMixin[ReturnType]): class FuncPath(ExprMixin[ReturnType]):
def __init__(self, func: t.Callable[[t.Any], ReturnType], operand: t.Optional[t.Any] = ...) -> None: ... def __init__(self, func: t.Callable[[t.Any], t.Any], operand: t.Optional[t.Any] = ...) -> None: ...
def __call__(self, operand: t.Any, *args: t.Any) -> ReturnType: ... def __call__(self, operand: t.Any, *args: t.Any) -> ReturnType: ...

View file

@ -19,7 +19,7 @@ def recursion_lock(
class Container(t.Generic[ContainerType], t.Dict[str, ContainerType]): class Container(t.Generic[ContainerType], t.Dict[str, ContainerType]):
def __getattr__(self, name: str) -> ContainerType: ... def __getattr__(self, name: str) -> ContainerType: ...
def update( # type: ignore def update(
self, self,
seqordict: t.Union[t.Dict[str, ContainerType], t.Tuple[str, ContainerType]], seqordict: t.Union[t.Dict[str, ContainerType], t.Tuple[str, ContainerType]],
) -> None: ... ) -> None: ...

View file

@ -1,5 +1,6 @@
import typing as t import typing as t
class HexDisplayedInteger(int): ... class HexDisplayedInteger(int): ...
class HexDisplayedBytes(bytes): ... class HexDisplayedBytes(bytes): ...
@ -9,6 +10,3 @@ V = t.TypeVar("V")
class HexDisplayedDict(t.Dict[K, V]): ... class HexDisplayedDict(t.Dict[K, V]): ...
class HexDumpDisplayedBytes(bytes): ... class HexDumpDisplayedBytes(bytes): ...
class HexDumpDisplayedDict(t.Dict[K, V]): ... class HexDumpDisplayedDict(t.Dict[K, V]): ...
def hexdump(data: bytes, linesize: int) -> str: ...
def hexundump(data: str, linesize: int) -> bytes: ...

View file

@ -1,6 +1,5 @@
import typing as t import typing as t
PY: t.Tuple[int, int]
PY2: bool PY2: bool
PY3: bool PY3: bool
PYPY: bool PYPY: bool

View file

@ -1,55 +1,32 @@
from .dataclass_struct import ( from construct_typed.generic import constr
from construct_typed.dataclass_struct import (
DataclassBitStruct, DataclassBitStruct,
DataclassMixin,
DataclassStruct, DataclassStruct,
TBitStruct, csfield
TContainerBase,
TContainerMixin,
TStruct,
TStructField,
csfield,
sfield,
EnhancedDataclassMixin
) )
from .generic_wrapper import ( from construct_typed.generic import (
Adapter, Adapter,
ConstantOrContextLambda, ConstantOrContextLambda,
ConstantOrContextLambda2,
Construct, Construct,
Context, Context,
ListContainer, ListContainer,
PathType, PathType,
Array,
Subconstruct,
Computed,
) )
from .tenum import EnumBase, EnumValue, FlagsEnumBase, TEnum, TFlagsEnum from construct_typed.tenum import TEnum, TFlags, TEnumConstruct, TFlagsConstruct
__all__ = [ __all__ = [
"DataclassBitStruct", "DataclassBitStruct",
"DataclassMixin",
"DataclassStruct", "DataclassStruct",
"TBitStruct", "constr",
"TContainerBase",
"TContainerMixin",
"TStruct",
"TStructField",
"csfield", "csfield",
"sfield",
"EnhancedDataclassMixin",
"EnumBase",
"EnumValue",
"FlagsEnumBase",
"TEnum", "TEnum",
"TFlagsEnum", "TEnumConstruct",
"TFlags",
"TFlagsConstruct",
"Adapter", "Adapter",
"ConstantOrContextLambda", "ConstantOrContextLambda",
"ConstantOrContextLambda2",
"Construct", "Construct",
"Context", "Context",
"ListContainer", "ListContainer",
"PathType", "PathType",
"Array",
"Subconstruct",
"Computed"
] ]

View file

@ -1,9 +1,10 @@
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
# pyright: strict # pyright: strict
# pyright: reportIncompatibleVariableOverride=false, reportAny=false
import dataclasses import dataclasses
import sys
import textwrap import textwrap
import typing as t import typing as t
import enum
import construct as cs import construct as cs
from construct.lib.containers import ( from construct.lib.containers import (
@ -12,25 +13,291 @@ from construct.lib.containers import (
recursion_lock, recursion_lock,
) )
from construct.lib.py3compat import bytestringtype, reprstring, unicodestringtype from construct.lib.py3compat import bytestringtype, reprstring, unicodestringtype
from typing_extensions import override
from .generic_wrapper import Adapter, Construct, Context, ParsedType, PathType 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")
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
# 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( # type: ignore
subcon: "Construct[ParsedType, None]",
*,
doc: t.Optional[str] = ...,
init: t.Literal[False] = ...,
default: t.Literal[None] = ...,
const: t.Literal[Flag.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] = ...,
init: t.Literal[True] = ...,
default: t.Literal[Flag.MISSING] = ...,
const: t.Literal[Flag.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[Flag.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[Flag.MISSING] = ...,
) -> ParsedType:
...
def csfield(
subcon: "Construct[ParsedType, t.Any]",
*,
doc: t.Optional[str] = None,
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]:
""" """
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' This method also processes Const and Default, to pass these values als default values to the dataclass.
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), Only one of the parameters `default` or `const` can be vaild. They are mutually exclusive.
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.
""" """
__dataclass_fields__: "t.ClassVar[dict[str, dataclasses.Field[t.Any]]]" 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
if doc is not None:
doc = textwrap.dedent(doc).strip("\n")
subcon = cs.Renamed(subcon, newdocs=doc)
if default is not Flag.MISSING:
init = True
default = default
subcon = cs.Default(subcon, default)
elif const is not Flag.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
return t.cast(
t.Optional[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 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
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: 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:
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
@__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.
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.
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: 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::
>>> 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: 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 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")
# 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
# create construct format
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 docs
dc_constr.docs = docs
# save construct format and make the class compatible to `Constructable` protocol
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.
def __getitem__(self, key: str) -> t.Any: def __getitem__(self, key: str) -> t.Any:
return getattr(self, key) return getattr(self, key)
@ -76,209 +343,43 @@ class DataclassMixin:
text.append(indentation.join(str(v).split("\n"))) text.append(indentation.join(str(v).split("\n")))
return "".join(text) return "".join(text)
if t.TYPE_CHECKING:
def csfield( @classmethod
subcon: Construct[ParsedType, t.Any], def __constr__(cls: t.Type[T]) -> "DataclassConstruct[T]":
doc: str | None = None, ...
parsed: t.Callable[[t.Any, Context], None] | 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]" = orig_subcon
default = const_subcon.value
elif isinstance(orig_subcon, cs.Default):
default_subcon: "cs.Default[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 DataclassBitStruct(DataclassStruct):
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" # type: ignore
def __init__(
self,
dc_type: type[DataclassType],
reverse: bool = False,
) -> None:
self.dc_type: type[DataclassType] = dc_type
self.reverse: bool = 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: dict[str, t.Any] = {}
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)
@override
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 # type: ignore
@override
def _encode(
self, obj: DataclassType, context: Context, path: PathType
) -> 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: 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: type[DataclassType], reverse: bool = False
) -> "cs.Transformed[DataclassType, DataclassType] | cs.Restreamed[DataclassType, DataclassType]":
r""" r"""
Makes a DataclassStruct inside a Bitwise. Makes a DataclassStruct inside a Bitwise.
See :class:`~construct.core.Bitwise` and :class:`~construct_typed.dataclass_struct.DatclassStruct` for semantics and raisable exceptions. 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 constr: TODO
:param reverse: Flag if the fields of the dataclass should be reversed :param reverse_fields: Flag if the fields of the dataclass should be reversed
Example:: Example::
DataclassBitStruct <--> Bitwise(DataclassStruct(...)) TODO:
>>> import dataclasses
>>> from construct import BitsInteger, Flag, Nibble, Padding >>> from construct import BitsInteger, Flag, Nibble, Padding
>>> from construct_typed import DataclassBitStruct, DataclassMixin, csfield >>> from construct_typed import DataclassBitStruct, csfield, construct
>>> @dataclasses.dataclass ... class TestDataclass(DataclassBitStruct):
... class TestDataclass(DataclassMixin):
... a: int = csfield(Flag) ... a: int = csfield(Flag)
... b: int = csfield(Nibble) ... b: int = csfield(Nibble)
... c: int = csfield(BitsInteger(10)) ... c: int = csfield(BitsInteger(10))
... d: None = csfield(Padding(1)) ... d: None = csfield(Padding(1))
>>> d = DataclassBitStruct(TestDataclass) >>> d = construct(TestDataclass)
>>> d.parse(b"\x01\x02") >>> d.parse(b"\x01\x02")
TestDataclass(a=False, b=0, c=129, d=None) TestDataclass(a=False, b=0, c=129, d=None)
""" """
return cs.Bitwise(DataclassStruct(dc_type, reverse))
class EnhancedDataclassMixin(DataclassMixin):
@classmethod
def format(cls):
return DataclassStruct(cls)
@classmethod @classmethod
def build(cls, obj: t.Self, **kw: dict[str, t.Any]): def __init_subclass__(
return cls.format().build(obj, **kw) cls: t.Type[T],
constr: t.Callable[
@classmethod [DataclassConstruct[T]], Construct[t.Any, t.Any]
def parse(cls, data: bytes | bytearray, **kw: dict[str, t.Any]): ] = lambda cls: cls,
return cls.format().parse(data, **kw) reverse_fields: bool = False,
) -> None:
@classmethod DataclassStruct.__init_subclass__.__func__(cls, lambda cls: cs.Bitwise(constr(cls)), reverse_fields) # type: ignore
def parse_file(cls, file: str, **kw: dict[str, t.Any]):
return cls.format().parse_file(file, **kw)
@classmethod
def parse_stream(cls, stream: t.IO[bytes], **kw: dict[str, t.Any]):
return cls.format().parse_stream(stream, **kw)
def build_self(self) -> bytes:
return self.build(self)
# support legacy names
TStruct = DataclassStruct
TBitStruct = DataclassBitStruct
TContainerMixin = DataclassMixin
TContainerBase = DataclassMixin
TStructField = csfield
sfield = csfield

File diff suppressed because it is too large Load diff

View file

@ -12,14 +12,11 @@ if t.TYPE_CHECKING:
# while type checking, the original classes are already generics, because they are defined like this in the stubs. # while type checking, the original classes are already generics, because they are defined like this in the stubs.
from construct import Adapter as Adapter from construct import Adapter as Adapter
from construct import ConstantOrContextLambda as ConstantOrContextLambda from construct import ConstantOrContextLambda as ConstantOrContextLambda
from construct import ConstantOrContextLambda2 as ConstantOrContextLambda2
from construct import Construct as Construct from construct import Construct as Construct
from construct import Context as Context from construct import Context as Context
from construct import ListContainer as ListContainer from construct import ListContainer as ListContainer
from construct import PathType as PathType from construct import PathType as PathType
from construct import Array as Array
from construct import Subconstruct as Subconstruct
from construct import Computed as Computed
else: else:
import construct as cs import construct as cs
@ -40,18 +37,22 @@ else:
class Context: class Context:
pass pass
class Array(
t.Generic[SubconParsedType, SubconBuildTypes],
cs.Array,
):
pass
class Subconstruct(t.Generic[SubconParsedType, SubconBuildTypes, ParsedType, BuildTypes], cs.Subconstruct):
pass
class Computed(t.Generic[ParsedType], cs.Computed):
pass
ConstantOrContextLambda = t.Union[ValueType, t.Callable[[Context], t.Any]] ConstantOrContextLambda = t.Union[ValueType, t.Callable[[Context], t.Any]]
ConstantOrContextLambda2 = t.Union[ValueType, t.Callable[[Context], ValueType]]
PathType = str PathType = str
@t.runtime_checkable
class Constructable(t.Protocol[ParsedType, BuildTypes]):
def __constr__(self) -> "Construct[ParsedType, BuildTypes]":
raise NotImplementedError
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.__constr__()
return constr

View file

@ -1,111 +1,147 @@
# pyright: reportAny=false
import enum import enum
import textwrap
import typing as t import typing as t
from typing_extensions import Self, override import construct as cs
from .generic_wrapper import Construct, Adapter, Context, PathType from construct_typed.generic import *
T = t.TypeVar("T")
# ## TEnum ############################################################################################################ class _EnumMeta(enum.EnumMeta):
class EnumValue: @classmethod
""" def __prepare__(
This is a helper class for adding documentation to an enum value. 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 __init__(self, value: int, doc: str | None = None) -> None: def __new__(
self.value: int = value metacls: t.Type[T], # type: ignore
self.__doc__ = doc if doc else "" __name: str,
__bases: t.Tuple[type, ...],
__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
class EnumBase(enum.IntEnum): # if the `TEnum` class is created, there are no parameters
""" if len(kwargs) == 0:
Base class for an Enum used in `construct_typed.TEnum`. return cls
This class extends the standard `enum.IntEnum` by. # extract parameters from kwargs
- missing values are automatically generated subcon: "cs.Construct[t.Any, t.Any]" = kwargs.pop("subcon", None)
- possibility to add documentation for each enum value (see `EnumValue`) 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}")
Example:: # create construct format
if TEnum in __bases:
>>> class State(EnumBase): enum_constr = TEnumConstruct(subcon, cls) # type: ignore
... Idle = 1 elif TFlags in __bases:
... Running = EnumValue(2, "This is the running state.") enum_constr = TFlagsConstruct(subcon, cls) # type: ignore
>>> State(1)
<State.Idle: 1>
>>> State["Idle"]
<State.Idle: 1>
>>> State.Idle
<State.Idle: 1>
>>> State(3) # missing value
<State.3: 3>
>>> State.Running.__doc__ # documentation
'This is the running state.'
"""
def __new__(cls, val: EnumValue | int) -> "Self":
if isinstance(val, EnumValue):
obj = int.__new__(cls, val.value)
obj._value_ = val.value
obj.__doc__ = val.__doc__
else: else:
obj = int.__new__(cls, val) raise TypeError("neither `TEnum` nor `TFlags` in bases")
obj._value_ = val
obj.__doc__ = ""
return obj
# Extend the enum type with _missing_ method. So if a enum value # 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
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.
"""
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]",
) -> None:
...
@classmethod
def __constr__(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. # not found in the enum, a new pseudo member is created.
# The idea is taken from: https://stackoverflow.com/a/57179436 # The idea is taken from: https://stackoverflow.com/a/57179436
@classmethod @classmethod
@override def _missing_(cls, value: t.Any) -> t.Optional["TEnum"]:
def _missing_(cls, value: t.Any) -> enum.Enum | None:
if isinstance(value, int): if isinstance(value, int):
pseudo_member = cls._value2member_map_.get(value, None) return cls._create_pseudo_member_(value)
if pseudo_member is None:
new_member = int.__new__(cls, value)
# I expect a name attribute to hold a string, hence str(value)
# However, new_member._name_ = value works, too
new_member._name_ = str(value)
new_member._value_ = value
new_member.__doc__ = "missing value"
pseudo_member = cls._value2member_map_.setdefault(value, new_member)
return pseudo_member
return None # will raise the ValueError in Enum.__new__ return None # will raise the ValueError in Enum.__new__
@override @classmethod
def __reduce_ex__(self, proto: t.Any) -> tuple[t.Any, ...]: def _create_pseudo_member_(cls, value: int) -> "TEnum":
""" pseudo_member = cls._value2member_map_.get(value, None) # type: ignore
Pickle enums by value instead of name (restores pre-3.11 behavior). if pseudo_member is None:
See https://github.com/python/cpython/pull/26658 for why this exists. new_member = int.__new__(cls, value)
""" # I expect a name attribute to hold a string, hence str(value)
return self.__class__, (self._value_,) # However, new_member._name_ = value works, too
new_member._name_ = str(value)
new_member._value_ = value
pseudo_member = cls._value2member_map_.setdefault(value, new_member) # type: ignore
return pseudo_member # type: ignore
EnumType = t.TypeVar("EnumType", bound=EnumBase) EnumType = t.TypeVar("EnumType", bound=TEnum)
class TEnum(Adapter[int, int, EnumType, EnumType]): class TEnumConstruct(Adapter[int, int, EnumType, EnumType]):
""" """
Typed enum. Typed enum.
""" """
def __init__(self, subcon: Construct[int, int], enum_type: type[EnumType]):
if t.TYPE_CHECKING:
def __new__(
cls, subcon: Construct[int, int], enum_type: t.Type[EnumType]
) -> "TEnumConstruct[EnumType]":
...
def __init__(self, subcon: Construct[int, int], enum_type: t.Type[EnumType]):
if not issubclass(enum_type, TEnum):
raise TypeError(
"'{}' has to be a '{}'".format(repr(enum_type), repr(TEnum))
)
# save enum type # save enum type
self.enum_type: type[EnumType] = enum_type self.enum_type = t.cast(t.Type[EnumType], enum_type) # type: ignore
# init adatper # init adatper
super(TEnum, self).__init__(subcon) # type: ignore super(TEnumConstruct, self).__init__(subcon) # type: ignore
@override
def _decode(self, obj: int, context: Context, path: PathType) -> EnumType: def _decode(self, obj: int, context: Context, path: PathType) -> EnumType:
return self.enum_type(obj) return self.enum_type(obj)
@override
def _encode( def _encode(
self, self,
obj: EnumType, obj: EnumType,
@ -119,91 +155,60 @@ class TEnum(Adapter[int, int, EnumType, EnumType]):
) )
# ## TFlagsEnum ####################################################################################################### # ## TFlags #######################################################################################################
class FlagsEnumBase(enum.IntFlag): 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]",
) -> None:
...
@classmethod
def __constr__(
cls: "t.Type[FlagsType]",
) -> "TFlagsConstruct[FlagsType]":
...
FlagsType = t.TypeVar("FlagsType", bound=TFlags)
class TFlagsConstruct(Adapter[int, int, FlagsType, FlagsType]):
""" """
Base class for an Enum used in `construct_typed.TFlagsEnum`. Typed flags.
This class extends the standard `enum.IntFlag` by.
- possibility to add documentation for each enum value (see `EnumValue`)
Example::
>>> class Option(FlagsEnumBase):
... OptOne = 1
... OptTwo = EnumValue(2, "This is option two.")
>>> Option(1)
<Option.OptOne: 1>
>>> Option["OptOne"]
<Option.OptOne: 1>
>>> Option.OptOne
<Option.OptOne: 1>
>>> Option(3)
<Option.OptTwo|OptOne: 3>
>>> Option(4)
<Option.4: 4>
>>> Option.OptTwo.__doc__ # documentation
'This is option two.'
""" """
def __new__(cls, val: EnumValue | int) -> "Self": if t.TYPE_CHECKING:
if isinstance(val, EnumValue):
obj = int.__new__(cls, val.value)
obj._value_ = val.value
obj.__doc__ = val.__doc__
else:
obj = int.__new__(cls, val)
obj._value_ = val
obj.__doc__ = ""
return obj
@classmethod def __new__(
@override cls, subcon: Construct[int, int], enum_type: t.Type[FlagsType]
def _missing_(cls, value: t.Any) -> t.Any: ) -> "TFlagsConstruct[FlagsType]":
""" ...
Returns member (possibly creating it) if one can be found for value.
"""
new_member = super()._missing_(value)
new_member.__doc__ = "missing value"
return new_member
@override def __init__(self, subcon: Construct[int, int], enum_type: t.Type[FlagsType]):
def __reduce_ex__(self, proto: t.Any) -> tuple[t.Any, ...]: if not issubclass(enum_type, TFlags):
""" raise TypeError(
Pickle enums by value instead of name (restores pre-3.11 behavior). "'{}' has to be a '{}'".format(repr(enum_type), repr(TFlags))
See https://github.com/python/cpython/pull/26658 for why this exists. )
"""
return self.__class__, (self._value_,)
FlagsEnumType = t.TypeVar("FlagsEnumType", bound=FlagsEnumBase)
class TFlagsEnum(Adapter[int, int, FlagsEnumType, FlagsEnumType]):
"""
Typed enum.
"""
def __init__(self, subcon: Construct[int, int], enum_type: type[FlagsEnumType]):
# save enum type # save enum type
self.enum_type: type[FlagsEnumType] = enum_type self.enum_type = t.cast(t.Type[FlagsType], enum_type) # type: ignore
# init adatper # init adatper
super(TFlagsEnum, self).__init__(subcon) # type: ignore super(TFlagsConstruct, self).__init__(subcon) # type: ignore
@override def _decode(self, obj: int, context: Context, path: PathType) -> FlagsType:
def _decode(self, obj: int, context: Context, path: PathType) -> FlagsEnumType:
return self.enum_type(obj) return self.enum_type(obj)
@override
def _encode( def _encode(
self, self,
obj: FlagsEnumType, obj: FlagsType,
context: Context, context: Context,
path: PathType, path: PathType,
) -> int: ) -> int:

View file

@ -1,2 +1,2 @@
version = (0, 7, 0) version = (0, 5, 2)
version_string = "0.7.0+wrapper" version_string = "0.5.2"

3
mypy.ini Normal file
View file

@ -0,0 +1,3 @@
[mypy]
strict = True
warn_unused_ignores = False

View file

@ -1,76 +0,0 @@
[build-system]
requires = ["setuptools >= 75.8.0"]
build-backend = "setuptools.build_meta"
[project]
name="construct-typing"
dynamic = ["version"]
license = { file = "LICENSE" }
description="Extension for the python package 'construct' that adds typing features"
readme = "README.md"
authors=[{ name = "Tim Riddermann" }]
requires-python = ">=3.9"
dependencies = [
"construct==2.10.70",
"typing_extensions>=4.6.0"
]
keywords = [
"construct",
"kaitai",
"declarative",
"data structure",
"struct",
"binary",
"symmetric",
"parser",
"builder",
"parsing",
"building",
"pack",
"unpack",
"packer",
"unpacker",
"bitstring",
"bytestring",
"annotation",
"type hint",
"typing",
"typed",
"bitstruct",
"PEP 561",
]
classifiers = [
"Development Status :: 3 - Alpha",
"License :: OSI Approved :: MIT License",
"Intended Audience :: Developers",
"Topic :: Software Development :: Libraries :: Python Modules",
"Topic :: Software Development :: Build Tools",
"Topic :: Software Development :: Code Generators",
"Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.9",
"Programming Language :: Python :: 3.10",
"Programming Language :: Python :: 3.11",
"Programming Language :: Python :: 3.12",
"Programming Language :: Python :: 3.13",
"Programming Language :: Python :: Implementation :: CPython",
"Typing :: Typed",
]
[project.urls]
"Homepage" = "https://github.com/timrid/construct-typing"
"Bug Reports" = "https://github.com/timrid/construct-typing/issues"
[tool.setuptools]
packages=[
"construct-stubs",
"construct-stubs.lib",
"construct_typed"
]
[tool.setuptools.dynamic]
version = {attr = "construct_typed.version.version_string"}
[tool.mypy]
strict = true
warn_unused_ignores = false

View file

@ -1,14 +1,10 @@
construct==2.10.70 construct==2.10.67
pytest>=6.2.0 pytest>=6.2.0
numpy numpy==1.21.*
arrow arrow
ruamel.yaml ruamel.yaml
cloudpickle cloudpickle
lz4 lz4
black black
isort isort
mypy mypy
cryptography
build
setuptools
wheel

64
setup.py Normal file
View file

@ -0,0 +1,64 @@
#!/usr/bin/env python
from setuptools import setup
version_string = "?.?.?"
exec(open("./construct_typed/version.py").read())
setup(
name="construct-typing",
version=version_string,
packages=["construct-stubs", "construct_typed"],
package_data={
"construct-stubs": ["*.pyi", "lib/*.pyi"],
"construct_typed": ["py.typed"],
},
license="MIT",
license_files=("LICENSE",),
description="Extension for the python package 'construct' that adds typing features",
long_description=open("README.md").read(),
long_description_content_type="text/markdown",
platforms=["POSIX", "Windows"],
url="https://github.com/timrid/construct-typing",
author="Tim Riddermann",
python_requires=">=3.7",
install_requires=["construct==2.10.67"],
keywords=[
"construct",
"kaitai",
"declarative",
"data structure",
"struct",
"binary",
"symmetric",
"parser",
"builder",
"parsing",
"building",
"pack",
"unpack",
"packer",
"unpacker",
"bitstring",
"bytestring",
"annotation",
"type hint",
"typing",
"typed",
"bitstruct",
"PEP 561",
],
classifiers=[
"Development Status :: 3 - Alpha",
"License :: OSI Approved :: MIT License",
"Intended Audience :: Developers",
"Topic :: Software Development :: Libraries :: Python Modules",
"Topic :: Software Development :: Build Tools",
"Topic :: Software Development :: Code Generators",
"Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.7",
"Programming Language :: Python :: 3.8",
"Programming Language :: Python :: 3.9",
"Programming Language :: Python :: Implementation :: CPython",
"Typing :: Typed",
],
)

View file

@ -1,170 +1,38 @@
import binascii
import io
import typing as t
import pytest import pytest
from construct import *
from construct.lib import *
import construct_typed as cst
xfail = pytest.mark.xfail xfail = pytest.mark.xfail
skip = pytest.mark.skip skip = pytest.mark.skip
skipif = pytest.mark.skipif skipif = pytest.mark.skipif
Buffer = t.Union[bytes, memoryview, bytearray] import os, math, random, collections, itertools, io, hashlib, binascii
ParsedType = t.TypeVar("ParsedType")
BuildTypes = t.TypeVar("BuildTypes")
ContainerType = t.TypeVar("ContainerType", bound=cst.TContainerMixin)
T = t.TypeVar("T")
IdentType = t.TypeVar("IdentType") from construct import *
from construct.lib import *
class ZeroIO(io.BufferedIOBase): class ZeroIO(io.BufferedIOBase):
def read(self, __size: t.Optional[int] = None) -> bytes: def read(self, __size=None):
if __size is not None: if __size is not None:
return bytes(__size) return bytes(__size)
else: else:
return bytes(0) return bytes(0)
def read1(self, __size: int = 0) -> bytes: def read1(self, __size=0):
return bytes(__size) return bytes(__size)
def ident(x: IdentType) -> IdentType: ident = lambda x: x
return x devzero = ZeroIO()
devzero: t.BinaryIO = ZeroIO() # type: ignore def raises(func, *args, **kw):
def raises(
func: t.Callable[..., t.Any], *args: t.Any, **kw: t.Any
) -> t.Union[t.Any, Exception]:
try: try:
return func(*args, **kw) return func(*args, **kw)
except Exception as e: except Exception as e:
return e.__class__ return e.__class__
@t.overload def common(format, datasample, objsample, sizesample=SizeofError, **kw):
def common(
format: cst.TStruct[ContainerType],
datasample: Buffer,
objsample: t.Union[ContainerType, t.Dict[str, t.Any]],
sizesample: t.Union[int, t.Type[Exception]] = ...,
**kw: t.Any
) -> None:
...
@t.overload
def common(
format: "Construct[ListContainer[ParsedType], t.Any]",
datasample: Buffer,
objsample: t.List[ParsedType],
sizesample: t.Union[int, t.Type[Exception]] = ...,
**kw: t.Any
) -> None:
...
@t.overload
def common(
format: "Construct[Container[t.Any], t.Any]",
datasample: Buffer,
objsample: t.Dict[str, t.Any],
sizesample: t.Union[int, t.Type[Exception]] = ...,
**kw: t.Any
) -> None:
...
@t.overload
def common(
format: "Construct[t.Union[EnumInteger, EnumIntegerString], t.Any]",
datasample: Buffer,
objsample: t.Union[int, str],
sizesample: t.Union[int, t.Type[Exception]] = ...,
**kw: t.Any
) -> None:
...
@t.overload
def common(
format: "Construct[HexDisplayedInteger, t.Any]",
datasample: Buffer,
objsample: t.Union[HexDisplayedInteger, int],
sizesample: t.Union[int, t.Type[Exception]] = ...,
**kw: t.Any
) -> None:
...
@t.overload
def common(
format: "Construct[HexDisplayedBytes, t.Any]",
datasample: Buffer,
objsample: t.Union[HexDisplayedBytes, bytes],
sizesample: t.Union[int, t.Type[Exception]] = ...,
**kw: t.Any
) -> None:
...
@t.overload
def common(
format: "Construct[HexDisplayedDict[str, t.Any], t.Any]",
datasample: Buffer,
objsample: t.Dict[str, t.Any],
sizesample: t.Union[int, t.Type[Exception]] = ...,
**kw: t.Any
) -> None:
...
@t.overload
def common(
format: "Construct[HexDumpDisplayedBytes, t.Any]",
datasample: Buffer,
objsample: t.Union[HexDumpDisplayedBytes, bytes],
sizesample: t.Union[int, t.Type[Exception]] = ...,
**kw: t.Any
) -> None:
...
@t.overload
def common(
format: "Construct[HexDumpDisplayedDict[str, t.Any], t.Any]",
datasample: Buffer,
objsample: t.Dict[str, t.Any],
sizesample: t.Union[int, t.Type[Exception]] = ...,
**kw: t.Any
) -> None:
...
@t.overload
def common(
format: "Construct[ParsedType, t.Any]",
datasample: Buffer,
objsample: ParsedType,
sizesample: t.Union[int, t.Type[Exception]] = ...,
**kw: t.Any
) -> None:
...
def common(
format: "Construct[t.Any, t.Any]",
datasample: Buffer,
objsample: t.Any,
sizesample: t.Union[int, t.Type[Exception]] = SizeofError,
**kw: t.Any
) -> None:
obj = format.parse(datasample, **kw) obj = format.parse(datasample, **kw)
assert obj == objsample assert obj == objsample
data = format.build(objsample, **kw) data = format.build(objsample, **kw)
@ -176,35 +44,35 @@ def common(
size = format.sizeof(**kw) size = format.sizeof(**kw)
assert size == sizesample assert size == sizesample
else: else:
size_ex = raises(format.sizeof, **kw) size = raises(format.sizeof, **kw)
assert size_ex == sizesample assert size == sizesample
def setattrs(obj: T, **kwargs: t.Any) -> T: def setattrs(obj, **kwargs):
"""Set multiple named values of an object""" """ Set multiple named values of an object """
for name, value in kwargs.items(): for name, value in kwargs.items():
setattr(obj, name, value) setattr(obj, name, value)
return obj return obj
def commonhex(format: "Construct[t.Any, t.Any]", hexdata: str) -> None: def commonhex(format, hexdata):
commonbytes(format, binascii.unhexlify(hexdata)) commonbytes(format, binascii.unhexlify(hexdata))
def commondumpdeprecated(format: "Construct[t.Any, t.Any]", filename: str) -> None: def commondumpdeprecated(format, filename):
filename = "tests/deprecated_gallery/blobs/" + filename filename = "tests/deprecated_gallery/blobs/" + filename
with open(filename, "rb") as f: with open(filename, "rb") as f:
data = f.read() data = f.read()
commonbytes(format, data) commonbytes(format, data)
def commondump(format: "Construct[t.Any, t.Any]", filename: str) -> None: def commondump(format, filename):
filename = "tests/gallery/blobs/" + filename filename = "tests/gallery/blobs/" + filename
with open(filename, "rb") as f: with open(filename, "rb") as f:
data = f.read() data = f.read()
commonbytes(format, data) commonbytes(format, data)
def commonbytes(format: "Construct[t.Any, t.Any]", data: bytes) -> None: def commonbytes(format, data):
obj = format.parse(data) obj = format.parse(data)
format.build(obj) data2 = format.build(obj)

View file

@ -0,0 +1,109 @@
import typing as t
from construct import *
from construct.lib import *
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.DataclassStruct)
T = t.TypeVar("T")
IdentType = t.TypeVar("IdentType")
def ident(p1: IdentType) -> IdentType: ...
devzero: t.BinaryIO
def raises(
func: t.Callable[..., t.Any], *args: t.Any, **kw: t.Any
) -> t.Union[t.Any, Exception]: ...
@t.overload
def common(
format: ContainerType,
datasample: Buffer,
objsample: t.Union[ContainerType, t.Dict[str, t.Any]],
sizesample: t.Union[int, t.Type[Exception]] = ...,
**kw: t.Any
) -> None: ...
@t.overload
def common(
format: Construct[ListContainer[ParsedType], t.Any],
datasample: Buffer,
objsample: t.List[ParsedType],
sizesample: t.Union[int, t.Type[Exception]] = ...,
**kw: t.Any
) -> None: ...
@t.overload
def common(
format: Construct[Container[t.Any], t.Any],
datasample: Buffer,
objsample: t.Dict[str, t.Any],
sizesample: t.Union[int, t.Type[Exception]] = ...,
**kw: t.Any
) -> None: ...
@t.overload
def common(
format: Construct[t.Union[EnumInteger, EnumIntegerString], t.Any],
datasample: Buffer,
objsample: t.Union[int, str],
sizesample: t.Union[int, t.Type[Exception]] = ...,
**kw: t.Any
) -> None: ...
@t.overload
def common(
format: Construct[HexDisplayedInteger, t.Any],
datasample: Buffer,
objsample: t.Union[HexDisplayedInteger, int],
sizesample: t.Union[int, t.Type[Exception]] = ...,
**kw: t.Any
) -> None: ...
@t.overload
def common(
format: Construct[HexDisplayedBytes, t.Any],
datasample: Buffer,
objsample: t.Union[HexDisplayedBytes, bytes],
sizesample: t.Union[int, t.Type[Exception]] = ...,
**kw: t.Any
) -> None: ...
@t.overload
def common(
format: Construct[HexDisplayedDict[str, t.Any], t.Any],
datasample: Buffer,
objsample: t.Dict[str, t.Any],
sizesample: t.Union[int, t.Type[Exception]] = ...,
**kw: t.Any
) -> None: ...
@t.overload
def common(
format: Construct[HexDumpDisplayedBytes, t.Any],
datasample: Buffer,
objsample: t.Union[HexDumpDisplayedBytes, bytes],
sizesample: t.Union[int, t.Type[Exception]] = ...,
**kw: t.Any
) -> None: ...
@t.overload
def common(
format: Construct[HexDumpDisplayedDict[str, t.Any], t.Any],
datasample: Buffer,
objsample: t.Dict[str, t.Any],
sizesample: t.Union[int, t.Type[Exception]] = ...,
**kw: t.Any
) -> None: ...
@t.overload
def common(
format: Construct[ParsedType, t.Any],
datasample: Buffer,
objsample: ParsedType,
sizesample: t.Union[int, t.Type[Exception]] = ...,
**kw: t.Any
) -> None: ...
def setattrs(obj: T, **kwargs: t.Any) -> T: ...
def commonhex(format: Construct[t.Any, t.Any], hexdata: str) -> None: ...
def commondumpdeprecated(
format: Construct[t.Any, t.Any], filename: str
) -> None: ...
def commondump(format: Construct[t.Any, t.Any], filename: str) -> None: ...
def commonbytes(
format: Construct[ParsedType, t.Any], data: ParsedType
) -> None: ...

View file

@ -1,6 +1,6 @@
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
# mypy: no-warn-unused-ignores
from .declarativeunittest import raises, common, ident, devzero from .declarativeunittest import raises, common, commonhex, commondumpdeprecated, commondump, commonbytes, ident, devzero
from construct.core import * from construct.core import *
from construct import * from construct import *
from construct.lib import * from construct.lib import *
@ -151,29 +151,17 @@ def test_formatfield_bool_issue_901() -> None:
assert d.sizeof() == 1 assert d.sizeof() == 1
def test_bytesinteger() -> None: def test_bytesinteger() -> None:
d = BytesInteger(0)
assert raises(d.parse, b"") == IntegerError
assert raises(d.build, 0) == IntegerError
d = BytesInteger(4, signed=True, swapped=False) d = BytesInteger(4, signed=True, swapped=False)
common(d, b"\x01\x02\x03\x04", 0x01020304, 4) common(d, b"\x01\x02\x03\x04", 0x01020304, 4)
common(d, b"\xff\xff\xff\xff", -1, 4) common(d, b"\xff\xff\xff\xff", -1, 4)
d = BytesInteger(4, signed=False, swapped=this.swapped) d = BytesInteger(4, signed=False, swapped=this.swapped)
common(d, b"\x01\x02\x03\x04", 0x01020304, 4, swapped=False) common(d, b"\x01\x02\x03\x04", 0x01020304, 4, swapped=False)
common(d, b"\x04\x03\x02\x01", 0x01020304, 4, swapped=True) common(d, b"\x04\x03\x02\x01", 0x01020304, 4, swapped=True)
assert raises(BytesInteger(-1).parse, b"") == IntegerError
assert raises(BytesInteger(-1).build, 0) == IntegerError
assert raises(BytesInteger(8).build, None) == IntegerError
assert raises(BytesInteger(8, signed=False).build, -1) == IntegerError
assert raises(BytesInteger(8, True).build, -2**64) == IntegerError
assert raises(BytesInteger(8, True).build, 2**64) == IntegerError
assert raises(BytesInteger(8, False).build, -2**64) == IntegerError
assert raises(BytesInteger(8, False).build, 2**64) == IntegerError
assert raises(BytesInteger(this.missing).sizeof) == SizeofError assert raises(BytesInteger(this.missing).sizeof) == SizeofError
assert raises(BytesInteger(4, signed=False).build, -1) == IntegerError
common(BytesInteger(0), b"", 0, 0)
def test_bitsinteger() -> None: def test_bitsinteger() -> None:
d = BitsInteger(0)
assert raises(d.parse, b"") == IntegerError
assert raises(d.build, 0) == IntegerError
d = BitsInteger(8) d = BitsInteger(8)
common(d, b"\x01\x01\x01\x01\x01\x01\x01\x01", 255, 8) common(d, b"\x01\x01\x01\x01\x01\x01\x01\x01", 255, 8)
d = BitsInteger(8, signed=True) d = BitsInteger(8, signed=True)
@ -183,17 +171,9 @@ def test_bitsinteger() -> None:
d = BitsInteger(16, swapped=this.swapped) d = BitsInteger(16, swapped=this.swapped)
common(d, b"\x01\x01\x01\x01\x01\x01\x01\x01\x00\x00\x00\x00\x00\x00\x00\x00", 0xff00, 16, swapped=False) common(d, b"\x01\x01\x01\x01\x01\x01\x01\x01\x00\x00\x00\x00\x00\x00\x00\x00", 0xff00, 16, swapped=False)
common(d, b"\x00\x00\x00\x00\x00\x00\x00\x00\x01\x01\x01\x01\x01\x01\x01\x01", 0xff00, 16, swapped=True) common(d, b"\x00\x00\x00\x00\x00\x00\x00\x00\x01\x01\x01\x01\x01\x01\x01\x01", 0xff00, 16, swapped=True)
assert raises(BitsInteger(-1).parse, b"") == IntegerError
assert raises(BitsInteger(-1).build, 0) == IntegerError
assert raises(BitsInteger(5, swapped=True).parse, bytes(5)) == IntegerError
assert raises(BitsInteger(5, swapped=True).build, 0) == IntegerError
assert raises(BitsInteger(8).build, None) == IntegerError
assert raises(BitsInteger(8, signed=False).build, -1) == IntegerError
assert raises(BitsInteger(8, True).build, -2**64) == IntegerError
assert raises(BitsInteger(8, True).build, 2**64) == IntegerError
assert raises(BitsInteger(8, False).build, -2**64) == IntegerError
assert raises(BitsInteger(8, False).build, 2**64) == IntegerError
assert raises(BitsInteger(this.missing).sizeof) == SizeofError assert raises(BitsInteger(this.missing).sizeof) == SizeofError
assert raises(BitsInteger(8, signed=False).build, -1) == IntegerError
common(BitsInteger(0), b"", 0, 0)
def test_varint() -> None: def test_varint() -> None:
d = VarInt d = VarInt
@ -244,8 +224,8 @@ def test_paddedstring() -> None:
common(PaddedString(100, e), data, s, 100) common(PaddedString(100, e), data, s, 100)
for e in ["ascii","utf8","utf16","utf-16-le","utf32","utf-32-le"]: for e in ["ascii","utf8","utf16","utf-16-le","utf32","utf-32-le"]:
assert PaddedString(10, e).sizeof() == 10 PaddedString(10, e).sizeof() == 10
assert PaddedString(this.n, e).sizeof(n=10) == 10 PaddedString(this.n, e).sizeof(n=10) == 10
def test_pascalstring() -> None: def test_pascalstring() -> None:
for e,_ in [("utf8",1),("utf16",2),("utf_16_le",2),("utf32",4),("utf_32_le",4)]: for e,_ in [("utf8",1),("utf16",2),("utf_16_le",2),("utf32",4),("utf_32_le",4)]:
@ -256,8 +236,8 @@ def test_pascalstring() -> None:
common(PascalString(sc, e), sc.build(0), u"") common(PascalString(sc, e), sc.build(0), u"")
for e in ["utf8","utf16","utf-16-le","utf32","utf-32-le","ascii"]: for e in ["utf8","utf16","utf-16-le","utf32","utf-32-le","ascii"]:
assert raises(PascalString(Byte, e).sizeof) == SizeofError raises(PascalString(Byte, e).sizeof) == SizeofError
assert raises(PascalString(VarInt, e).sizeof) == SizeofError raises(PascalString(VarInt, e).sizeof) == SizeofError
def test_cstring() -> None: def test_cstring() -> None:
s = u"" s = u""
@ -266,12 +246,12 @@ def test_cstring() -> None:
common(CString(e), s.encode(e)+bytes(us), s) common(CString(e), s.encode(e)+bytes(us), s)
common(CString(e), bytes(us), u"") common(CString(e), bytes(us), u"")
assert CString("utf8").build(s) == b'\xd0\x90\xd1\x84\xd0\xbe\xd0\xbd'+b"\x00" CString("utf8").build(s) == b'\xd0\x90\xd1\x84\xd0\xbe\xd0\xbd'+b"\x00"
assert CString("utf16").build(s) == b'\xff\xfe\x10\x04D\x04>\x04=\x04'+b"\x00\x00" CString("utf16").build(s) == b'\xff\xfe\x10\x04D\x04>\x04=\x04'+b"\x00\x00"
assert CString("utf32").build(s) == b'\xff\xfe\x00\x00\x10\x04\x00\x00D\x04\x00\x00>\x04\x00\x00=\x04\x00\x00'+b"\x00\x00\x00\x00" CString("utf32").build(s) == b'\xff\xfe\x00\x00\x10\x04\x00\x00D\x04\x00\x00>\x04\x00\x00=\x04\x00\x00'+b"\x00\x00\x00\x00"
for e in ["utf8","utf16","utf-16-le","utf32","utf-32-le","ascii"]: for e in ["utf8","utf16","utf-16-le","utf32","utf-32-le","ascii"]:
assert raises(CString(e).sizeof) == SizeofError raises(CString(e).sizeof) == SizeofError
def test_greedystring() -> None: def test_greedystring() -> None:
for e,_ in [("utf8",1),("utf16",2),("utf_16_le",2),("utf32",4),("utf_32_le",4)]: for e,_ in [("utf8",1),("utf16",2),("utf_16_le",2),("utf32",4),("utf_32_le",4)]:
@ -280,7 +260,7 @@ def test_greedystring() -> None:
common(GreedyString(e), b"", u"") common(GreedyString(e), b"", u"")
for e in ["utf8","utf16","utf-16-le","utf32","utf-32-le","ascii"]: for e in ["utf8","utf16","utf-16-le","utf32","utf-32-le","ascii"]:
assert raises(GreedyString(e).sizeof) == SizeofError raises(GreedyString(e).sizeof) == SizeofError
def test_string_encodings() -> None: def test_string_encodings() -> None:
# checks that "-" is replaced with "_" # checks that "-" is replaced with "_"
@ -291,7 +271,7 @@ def test_flag() -> None:
d = Flag d = Flag
common(d, b"\x00", False, 1) common(d, b"\x00", False, 1)
common(d, b"\x01", True, 1) common(d, b"\x01", True, 1)
assert d.parse(b"\xff") == True d.parse(b"\xff") == True
def test_enum() -> None: def test_enum() -> None:
d = Enum(Byte, one=1, two=2, four=4, eight=8) d = Enum(Byte, one=1, two=2, four=4, eight=8)
@ -440,11 +420,11 @@ def test_struct_proper_context() -> None:
"x"/Byte, "x"/Byte,
"inner"/Struct( "inner"/Struct(
"y"/Byte, "y"/Byte,
"a"/Computed(this._.x+1), # type: ignore "a"/Computed(this._.x+1),
"b"/Computed(this.y+2), # type: ignore "b"/Computed(this.y+2),
), ),
"c"/Computed(this.x+3), # type: ignore "c"/Computed(this.x+3),
"d"/Computed(this.inner.y+4), # type: ignore "d"/Computed(this.inner.y+4),
) )
assert d.parse(b"\x01\x0f") == Container(x=1, inner=Container(y=15, a=2, b=17), c=4, d=19) assert d.parse(b"\x01\x0f") == Container(x=1, inner=Container(y=15, a=2, b=17), c=4, d=19)
@ -531,7 +511,7 @@ def test_const() -> None:
def test_computed() -> None: def test_computed() -> None:
common(Computed(255), b"", 255, 0) common(Computed(255), b"", 255, 0)
common(Computed(lambda ctx: 255), b"", 255, 0) # type: ignore common(Computed(lambda ctx: 255), b"", 255, 0)
assert Computed(255).build(None) == b"" assert Computed(255).build(None) == b""
assert Struct(Computed(255)).build({}) == b"" assert Struct(Computed(255)).build({}) == b""
assert raises(Computed(this.missing).parse, b"") == KeyError assert raises(Computed(this.missing).parse, b"") == KeyError
@ -611,7 +591,7 @@ def test_rebuild_issue_664() -> None:
def test_default() -> None: def test_default() -> None:
d = Default(Byte, 0) d = Default(Byte, 0)
common(d, b"\xff", 255, 1) common(d, b"\xff", 255, 1)
assert d.build(None) == b"\x00" d.build(None) == b"\x00"
def test_check() -> None: def test_check() -> None:
common(Check(True), b"", None, 0) common(Check(True), b"", None, 0)
@ -657,7 +637,8 @@ def test_numpy_error() -> None:
numpy.load(io.BytesIO(b"")) # type: ignore numpy.load(io.BytesIO(b"")) # type: ignore
def test_namedtuple() -> None: def test_namedtuple() -> None:
coord = t.NamedTuple("coord", [("x", int), ("y", int), ("z", int)]) import collections
coord = collections.namedtuple("coord", "x y z")
d1 = NamedTuple("coord", "x y z", Array(3, Byte)) d1 = NamedTuple("coord", "x y z", Array(3, Byte))
common(d1, b"123", coord(49,50,51), 3) common(d1, b"123", coord(49,50,51), 3)
d2 = NamedTuple("coord", "x y z", GreedyRange(Byte)) d2 = NamedTuple("coord", "x y z", GreedyRange(Byte))
@ -727,13 +708,10 @@ def test_hexdump() -> None:
def test_hexdump_regression_issue_188() -> None: def test_hexdump_regression_issue_188() -> None:
# Hex HexDump were not inheriting subcon flags # Hex HexDump were not inheriting subcon flags
a = Hex(Const(b"MZ")) d = Struct(Hex(Const(b"MZ")))
d = Struct(a)
assert d.parse(b"MZ") == Container() assert d.parse(b"MZ") == Container()
assert d.build(dict()) == b"MZ" assert d.build(dict()) == b"MZ"
d = Struct(HexDump(Const(b"MZ")))
b = HexDump(Const(b"MZ"))
d = Struct(b)
assert d.parse(b"MZ") == Container() assert d.parse(b"MZ") == Container()
assert d.build(dict()) == b"MZ" assert d.build(dict()) == b"MZ"
@ -830,10 +808,8 @@ def test_select_buildfromnone_issue_747() -> None:
assert d.build(dict()) == b"" assert d.build(dict()) == b""
def test_if() -> None: def test_if() -> None:
d = If(True, Byte) common(If(True, Byte), b"\x01", 1, 1)
common(d, b"\x01", 1, 1) common(If(False, Byte), b"", None, 0)
d = If(False, Byte)
common(d, b"", None, 0)
def test_ifthenelse() -> None: def test_ifthenelse() -> None:
common(IfThenElse(True, Int8ub, Int16ub), b"\x01", 1, 1) common(IfThenElse(True, Int8ub, Int16ub), b"\x01", 1, 1)
@ -946,17 +922,6 @@ def test_peek() -> None:
assert d4.build(Container(a=0x01, b=0x0102)) == b"" assert d4.build(Container(a=0x01, b=0x0102)) == b""
assert d4.sizeof() == 0 assert d4.sizeof() == 0
def test_offsettedend() -> None:
d1 = Struct(
"header" / Bytes(2),
"data" / OffsettedEnd(-2, GreedyBytes),
"footer" / Bytes(2),
)
common(d1, b"\x01\x02\x03\x04\x05\x06\x07", Container(header=b'\x01\x02', data=b'\x03\x04\x05', footer=b'\x06\x07'))
d2 = OffsettedEnd(0, Byte)
assert raises(d2.sizeof) == SizeofError
def test_seek() -> None: def test_seek() -> None:
d = Seek(5) d = Seek(5)
assert d.parse(b"") == 5 assert d.parse(b"") == 5
@ -1068,14 +1033,13 @@ def test_prefixed() -> None:
common(d5, b"\x0a"+bytes(10), u"\x00"*10, SizeofError) common(d5, b"\x0a"+bytes(10), u"\x00"*10, SizeofError)
def test_prefixedarray() -> None: def test_prefixedarray() -> None:
d = PrefixedArray(Byte, Byte) common(PrefixedArray(Byte,Byte), b"\x02\x0a\x0b", [10,11], SizeofError)
common(d, b"\x02\x0a\x0b", [10,11], SizeofError) assert PrefixedArray(Byte, Byte).parse(b"\x03\x01\x02\x03") == [1,2,3]
assert d.parse(b"\x03\x01\x02\x03") == [1,2,3] assert PrefixedArray(Byte, Byte).parse(b"\x00") == []
assert d.parse(b"\x00") == [] assert PrefixedArray(Byte, Byte).build([1,2,3]) == b"\x03\x01\x02\x03"
assert d.build([1,2,3]) == b"\x03\x01\x02\x03" assert raises(PrefixedArray(Byte, Byte).parse, b"") == StreamError
assert raises(d.parse, b"") == StreamError assert raises(PrefixedArray(Byte, Byte).parse, b"\x03\x01") == StreamError
assert raises(d.parse, b"\x03\x01") == StreamError assert raises(PrefixedArray(Byte, Byte).sizeof) == SizeofError
assert raises(d.sizeof) == SizeofError
def test_fixedsized() -> None: def test_fixedsized() -> None:
d1 = FixedSized(10, Byte) d1 = FixedSized(10, Byte)
@ -1249,7 +1213,7 @@ def test_checksum() -> None:
def test_checksum_nonbytes_issue_323() -> None: def test_checksum_nonbytes_issue_323() -> None:
d = Struct( d = Struct(
"vals" / Byte[2], "vals" / Byte[2],
"checksum" / Checksum(Byte, lambda vals: int(sum(vals)) & 0xFF, this.vals), "checksum" / Checksum(Byte, lambda vals: sum(vals) & 0xFF, this.vals),
) )
assert d.parse(b"\x00\x00\x00") == Container(vals=[0, 0], checksum=0) assert d.parse(b"\x00\x00\x00") == Container(vals=[0, 0], checksum=0)
assert raises(d.parse, b"\x00\x00\x01") == ChecksumError assert raises(d.parse, b"\x00\x00\x01") == ChecksumError
@ -1365,105 +1329,6 @@ def test_compressed_prefixed() -> None:
assert st.parse(st.build(Container(one=zeros,two=zeros))) == Container(one=zeros,two=zeros) assert st.parse(st.build(Container(one=zeros,two=zeros))) == Container(one=zeros,two=zeros)
assert raises(d.sizeof) == SizeofError assert raises(d.sizeof) == SizeofError
@pytest.mark.xfail(ONWINDOWS and PYPY, reason="no wheel for 'cryptography' is currently available for pypy on windows")
def test_encryptedsym() -> None:
from cryptography.hazmat.primitives.ciphers import Cipher, algorithms, modes
key128 = b"\x10\x11\x12\x13\x14\x15\x16\x17\x18\x19\x1a\x1b\x1c\x1d\x1e\x1f"
key256 = b"\x10\x11\x12\x13\x14\x15\x16\x17\x18\x19\x1a\x1b\x1c\x1d\x1e\x1f\x10\x11\x12\x13\x14\x15\x16\x17\x18\x19\x1a\x1b\x1c\x1d\x1e\x1f"
iv = b"\x20\x21\x22\x23\x24\x25\x26\x27\x28\x29\x2a\x2b\x2c\x2d\x2e\x2f"
nonce = iv
# AES 128/256 bit - ECB
d = EncryptedSym(GreedyBytes, lambda ctx: Cipher(algorithms.AES(ctx.key), modes.ECB()))
common(d, b"\xf4\x0f\x54\xb7\x6a\x7a\xf1\xdb\x92\x73\x14\xde\x2f\xa0\x3e\x2d", b'Secret Message..', key=key128, iv=iv)
common(d, b"\x82\x6b\x01\x82\x90\x02\xa1\x9e\x35\x0a\xe2\xc3\xee\x1a\x42\xf5", b'Secret Message..', key=key256, iv=iv)
# AES 128/256 bit - CBC
d = EncryptedSym(GreedyBytes, lambda ctx: Cipher(algorithms.AES(ctx.key), modes.CBC(ctx.iv)))
common(d, b"\xba\x79\xc2\x62\x22\x08\x29\xb9\xfb\xd3\x90\xc4\x04\xb7\x55\x87", b'Secret Message..', key=key128, iv=iv)
common(d, b"\x60\xc2\x45\x0d\x7e\x41\xd4\xf8\x85\xd4\x8a\x64\xd1\x45\x49\xe3", b'Secret Message..', key=key256, iv=iv)
# AES 128/256 bit - CTR
d = EncryptedSym(GreedyBytes, lambda ctx: Cipher(algorithms.AES(ctx.key), modes.CTR(ctx.nonce)))
common(d, b"\x80\x78\xb6\x0c\x07\xf5\x0c\x90\xce\xa2\xbf\xcb\x5b\x22\xb9\xb5", b'Secret Message..', key=key128, nonce=nonce)
common(d, b"\x6a\xae\x7b\x86\x1a\xa6\xe0\x6a\x49\x02\x02\x1b\xf2\x3c\xd8\x0d", b'Secret Message..', key=key256, nonce=nonce)
assert raises(EncryptedSym(GreedyBytes, "AES").build, b"") == CipherError # type: ignore
assert raises(EncryptedSym(GreedyBytes, "AES").parse, b"") == CipherError # type: ignore
@pytest.mark.xfail(ONWINDOWS and PYPY, reason="no wheel for 'cryptography' is currently available for pypy on windows")
def test_encryptedsym_cbc_example() -> None:
from cryptography.hazmat.primitives.ciphers import Cipher, algorithms, modes
d = Struct(
"iv" / Default(Bytes(16), os.urandom(16)),
"enc_data" / EncryptedSym(
Aligned(16,
Struct(
"width" / Int16ul,
"height" / Int16ul
)
),
lambda ctx: Cipher(algorithms.AES(ctx._.key), modes.CBC(ctx.iv))
)
)
key128 = b"\x10\x11\x12\x13\x14\x15\x16\x17\x18\x19\x1a\x1b\x1c\x1d\x1e\x1f"
byts = d.build({"enc_data": {"width": 5, "height": 4}}, key=key128)
obj = d.parse(byts, key=key128)
assert obj.enc_data == Container(width=5, height=4)
@pytest.mark.xfail(ONWINDOWS and PYPY, reason="no wheel for 'cryptography' is currently available for pypy on windows")
def test_encryptedsymaead() -> None:
from cryptography.hazmat.primitives.ciphers import aead
key128 = b"\x10\x11\x12\x13\x14\x15\x16\x17\x18\x19\x1a\x1b\x1c\x1d\x1e\x1f"
key256 = b"\x10\x11\x12\x13\x14\x15\x16\x17\x18\x19\x1a\x1b\x1c\x1d\x1e\x1f\x10\x11\x12\x13\x14\x15\x16\x17\x18\x19\x1a\x1b\x1c\x1d\x1e\x1f"
nonce = b"\x20\x21\x22\x23\x24\x25\x26\x27\x28\x29\x2a\x2b\x2c\x2d\x2e\x2f"
# AES 128/256 bit - GCM
d = Struct(
"associated_data" / Bytes(21),
"data" / EncryptedSymAead(
GreedyBytes,
lambda ctx: aead.AESGCM(ctx._.key),
this._.nonce,
this.associated_data
)
)
common(
d,
b"This is authenticated\xb6\xd3\x64\x0c\x7a\x31\xaa\x16\xa3\x58\xec\x17\x39\x99\x2e\xf8\x4e\x41\x17\x76\x3f\xd1\x06\x47\x04\x9f\x42\x1c\xf4\xa9\xfd\x99\x9c\xe9",
Container(associated_data=b"This is authenticated", data=b"The secret message"),
key=key128,
nonce=nonce
)
common(
d,
b"This is authenticated\xde\xb4\x41\x79\xc8\x7f\xea\x8d\x0e\x41\xf6\x44\x2f\x93\x21\xe6\x37\xd1\xd3\x29\xa4\x97\xc3\xb5\xf4\x81\x72\xa1\x7f\x3b\x9b\x53\x24\xe4",
Container(associated_data=b"This is authenticated", data=b"The secret message"),
key=key256,
nonce=nonce
)
assert raises(EncryptedSymAead(GreedyBytes, "AESGCM", bytes(16)).build, b"") == CipherError # type: ignore
assert raises(EncryptedSymAead(GreedyBytes, "AESGCM", bytes(16)).parse, b"") == CipherError # type: ignore
@pytest.mark.xfail(ONWINDOWS and PYPY, reason="no wheel for 'cryptography' is currently available for pypy on windows")
def test_encryptedsymaead_gcm_example() -> None:
from cryptography.hazmat.primitives.ciphers import aead
d = Struct(
"nonce" / Default(Bytes(16), os.urandom(16)),
"associated_data" / Bytes(21),
"enc_data" / EncryptedSymAead(
GreedyBytes,
lambda ctx: aead.AESGCM(ctx._.key),
this.nonce,
this.associated_data
)
)
key128 = b"\x10\x11\x12\x13\x14\x15\x16\x17\x18\x19\x1a\x1b\x1c\x1d\x1e\x1f"
byts = d.build({"associated_data": b"This is authenticated", "enc_data": b"The secret message"}, key=key128)
obj = d.parse(byts, key=key128)
assert obj.enc_data == b"The secret message"
assert obj.associated_data == b"This is authenticated"
def test_rebuffered() -> None: def test_rebuffered() -> None:
data = b"0" * 1000 data = b"0" * 1000
assert Rebuffered(Array(1000,Byte)).parse_stream(io.BytesIO(data)) == [48]*1000 assert Rebuffered(Array(1000,Byte)).parse_stream(io.BytesIO(data)) == [48]*1000
@ -1679,7 +1544,7 @@ def test_operators() -> None:
assert d.docs == "description" assert d.docs == "description"
d = "description" * Byte d = "description" * Byte
assert d.docs == "description" assert d.docs == "description"
_ = """ """
description description
""" * \ """ * \
Byte Byte
@ -1821,11 +1686,9 @@ def test_from_issue_244() -> None:
assert d.parse(b"abcd") == [Container(num=97, index=0),Container(num=98, index=1),Container(num=99, index=2),Container(num=100, index=3),] assert d.parse(b"abcd") == [Container(num=97, index=0),Container(num=98, index=1),Container(num=99, index=2),Container(num=100, index=3),]
def test_from_issue_269() -> None: def test_from_issue_269() -> None:
a = If(this.enabled, Padding(2)) d = Struct("enabled" / Byte, If(this.enabled, Padding(2)))
d = Struct("enabled" / Byte, a)
assert d.build(dict(enabled=1)) == b"\x01\x00\x00" assert d.build(dict(enabled=1)) == b"\x01\x00\x00"
assert d.build(dict(enabled=0)) == b"\x00" assert d.build(dict(enabled=0)) == b"\x00"
d = Struct("enabled" / Byte, "pad" / If(this.enabled, Padding(2))) d = Struct("enabled" / Byte, "pad" / If(this.enabled, Padding(2)))
assert d.build(dict(enabled=1)) == b"\x01\x00\x00" assert d.build(dict(enabled=1)) == b"\x01\x00\x00"
assert d.build(dict(enabled=0)) == b"\x00" assert d.build(dict(enabled=0)) == b"\x00"
@ -1841,7 +1704,7 @@ def test_from_issue_324() -> None:
)), )),
"checksum" / Checksum( "checksum" / Checksum(
Byte, Byte,
lambda data: int(sum(data)) & 0xFF, lambda data: sum(data) & 0xFF,
this.vals.data this.vals.data
), ),
) )
@ -1932,11 +1795,11 @@ def test_pickling_constructs() -> None:
) )
data = bytes(100) data = bytes(100)
du = cloudpickle.loads(cloudpickle.dumps(d, protocol=-1)) # type: ignore du = cloudpickle.loads(cloudpickle.dumps(d, protocol=-1))
assert du.parse(data) == d.parse(data) assert du.parse(data) == d.parse(data)
def test_pickling_constructs_issue_894() -> None: def test_pickling_constructs_issue_894() -> None:
import cloudpickle # type: ignore import cloudpickle
fundus_header = Struct( fundus_header = Struct(
'width' / Int32un, 'width' / Int32un,
@ -1948,7 +1811,7 @@ def test_pickling_constructs_issue_894() -> None:
'img' / Int8un, 'img' / Int8un,
) )
cloudpickle.dumps(fundus_header) # type: ignore cloudpickle.dumps(fundus_header)
def test_exposing_members_attributes() -> None: def test_exposing_members_attributes() -> None:
d1 = Struct( d1 = Struct(
@ -2159,7 +2022,7 @@ def test_struct_root_topmost() -> None:
assert d.parse(b"", z=2) == Container(x=1, inner=Container(inner2=Container(x=1,z=2,zz=2))) assert d.parse(b"", z=2) == Container(x=1, inner=Container(inner2=Container(x=1,z=2,zz=2)))
def test_parsedhook_repeatersdiscard() -> None: def test_parsedhook_repeatersdiscard() -> None:
outputs: t.List[int] = [] outputs = []
def printobj1(obj: int, ctx: "Context") -> None: def printobj1(obj: int, ctx: "Context") -> None:
outputs.append(obj) outputs.append(obj)
d1 = GreedyRange(Byte * printobj1, discard=True) d1 = GreedyRange(Byte * printobj1, discard=True)

View file

@ -1,152 +1,137 @@
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
# pyright: strict # pyright: strict
import dataclasses
import enum
import textwrap
import typing as t import typing as t
import construct as cs import construct as cs
from construct_typed import (
DataclassBitStruct,
DataclassStruct,
csfield,
constr,
TEnum,
TFlags,
)
import construct_typed as cst from tests.declarativeunittest import common, raises, setattrs
from construct_typed import DataclassBitStruct, DataclassMixin, DataclassStruct, csfield
from .declarativeunittest import common, raises, setattrs
def test_dataclass_const_default() -> None: def test_dataclass_const_default() -> None:
@dataclasses.dataclass class TestDataclass(DataclassStruct):
class ConstDefaultTest(DataclassMixin): const_bytes: bytes = csfield(cs.Bytes(3), const=b"BMP")
const_bytes: bytes = csfield(cs.Const(b"BMP")) const_int: int = csfield(cs.Int8ub, const=5)
const_int: int = csfield(cs.Const(5, cs.Int8ub)) default_int: int = csfield(cs.Int8ub, default=26)
default_int: int = csfield(cs.Default(cs.Int8ub, 28)) default_lambda: t.Optional[bytes] = csfield(
default_lambda: bytes = csfield(
cs.Default(cs.Bytes(cs.this.const_int), lambda ctx: bytes(ctx.const_int)) cs.Default(cs.Bytes(cs.this.const_int), lambda ctx: bytes(ctx.const_int))
) )
a = ConstDefaultTest() obj = TestDataclass()
assert a.const_bytes == b"BMP" assert obj.const_bytes == b"BMP"
assert a.const_int == 5 assert obj.const_int == 5
assert a.default_int == 28 assert obj.default_int == 26
assert a.default_lambda == None assert obj.default_lambda == None
obj = TestDataclass(default_int=1)
assert obj.default_int == 1
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: def test_dataclass_access() -> None:
@dataclasses.dataclass class TestDataclass(DataclassStruct):
class TestTContainer(DataclassMixin): a: int = csfield(cs.Byte, const=1)
a: t.Optional[int] = csfield(cs.Const(1, cs.Byte))
b: int = csfield(cs.Int8ub) b: int = csfield(cs.Int8ub)
tcontainer = TestTContainer(b=2) obj = TestDataclass(b=2)
# tcontainer assert obj.a == 1
assert tcontainer.a == 1 assert obj["a"] == 1
assert tcontainer["a"] == 1 assert obj.b == 2
assert tcontainer.b == 2 assert obj["b"] == 2
assert tcontainer["b"] == 2
tcontainer.a = 5 obj.a = 5
assert tcontainer.a == 5 assert obj.a == 5
assert tcontainer["a"] == 5 assert obj["a"] == 5
tcontainer["a"] = 6 obj["a"] = 6
assert tcontainer.a == 6 assert obj.a == 6
assert tcontainer["a"] == 6 assert obj["a"] == 6
# wrong creation # wrong creation
assert raises(lambda: TestTContainer(a=0, b=1)) == TypeError assert raises(lambda: TestDataclass(a=0, b=1)) == TypeError # type: ignore
def test_dataclass_str_repr() -> None: def test_dataclass_str_repr() -> None:
@dataclasses.dataclass class Image(DataclassStruct):
class Image(DataclassMixin): signature: bytes = csfield(cs.Bytes(3), const=b"BMP")
signature: t.Optional[bytes] = csfield(cs.Const(b"BMP"))
width: int = csfield(cs.Int8ub) width: int = csfield(cs.Int8ub)
height: int = csfield(cs.Int8ub) height: int = csfield(cs.Int8ub)
format = DataclassStruct(Image) fmt = constr(Image)
obj = Image(width=3, height=2) obj = Image(width=3, height=2)
assert ( assert (
str(obj) str(obj)
== "Image: \n signature = b'BMP' (total 3)\n width = 3\n height = 2" == "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 ( assert (
str(obj) str(obj)
== "Image: \n signature = b'BMP' (total 3)\n width = 3\n height = 2" == "Image: \n signature = b'BMP' (total 3)\n width = 3\n height = 2"
) )
def test_dataclass_ifthenelse() -> None:
@dataclasses.dataclass
class IfThenElseTest(DataclassMixin):
test_if: t.Optional[int] = csfield(cs.If(False, cs.Int8ub))
test_ifthenelse: t.Optional[int] = csfield(
cs.IfThenElse(True, cs.Int8ub, cs.Pass)
)
a = IfThenElseTest(test_if=None, test_ifthenelse=None)
assert a.test_if == None
assert a.test_ifthenelse == None
def test_dataclass_struct() -> None: def test_dataclass_struct() -> None:
@dataclasses.dataclass class Image(DataclassStruct):
class Image(DataclassMixin):
width: int = csfield(cs.Int8ub) width: int = csfield(cs.Int8ub)
height: int = csfield(cs.Int8ub) height: int = csfield(cs.Int8ub)
pixels: bytes = csfield(cs.Bytes(cs.this.height * cs.this.width)) pixels: bytes = csfield(cs.Bytes(cs.this.height * cs.this.width))
common( common(
cst.DataclassStruct(Image), constr(Image),
b"\x01\x0212", b"\x01\x0212",
Image(width=1, height=2, pixels=b"12"), Image(width=1, height=2, pixels=b"12"),
) )
# check __getattr__ # check __getattr__
c = cst.DataclassStruct(Image) fmt = Image.__constr__()
assert c.width.name == "width" assert fmt.width.name == "width"
assert c.height.name == "height" assert fmt.height.name == "height"
assert c.width.subcon is cs.Int8ub assert fmt.width.subcon is cs.Int8ub
assert c.height.subcon is cs.Int8ub assert fmt.height.subcon is cs.Int8ub
def test_dataclass_struct_reverse() -> None: def test_dataclass_struct_reverse() -> None:
@dataclasses.dataclass class TestDataclass(DataclassStruct, reverse_fields=True):
class TestContainer(DataclassMixin):
a: int = csfield(cs.Int16ub) a: int = csfield(cs.Int16ub)
b: int = csfield(cs.Int8ub) b: int = csfield(cs.Int8ub)
common( common(
DataclassStruct(TestContainer, reverse=True), constr(TestDataclass),
b"\x02\x00\x01", b"\x02\x00\x01",
TestContainer(a=1, b=2), TestDataclass(a=1, b=2),
3, 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: def test_dataclass_struct_nested() -> None:
@dataclasses.dataclass class TestDataclass(DataclassStruct):
class TestContainer(DataclassMixin): class InnerDataclass(DataclassStruct):
@dataclasses.dataclass
class InnerDataclass(DataclassMixin):
b: int = csfield(cs.Byte) b: int = csfield(cs.Byte)
c: bytes = csfield(cs.Bytes(cs.this._.length)) c: bytes = csfield(cs.Bytes(cs.this._.length))
length: int = csfield(cs.Byte) length: int = csfield(cs.Byte)
a: InnerDataclass = csfield(DataclassStruct(InnerDataclass)) a: InnerDataclass = csfield(constr(InnerDataclass))
common( common(
DataclassStruct(TestContainer), constr(TestDataclass),
b"\x02\x01\xF1\xF2", 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")),
) )
def test_dataclass_struct_default_field() -> None: def test_dataclass_struct_default_field() -> None:
@dataclasses.dataclass class Image(DataclassStruct):
class Image(DataclassMixin):
width: int = csfield(cs.Int8ub) width: int = csfield(cs.Int8ub)
height: int = csfield(cs.Int8ub) height: int = csfield(cs.Int8ub)
pixels: t.Optional[bytes] = csfield( pixels: t.Optional[bytes] = csfield(
@ -157,80 +142,75 @@ def test_dataclass_struct_default_field() -> None:
) )
common( common(
DataclassStruct(Image), constr(Image),
b"\x02\x03\x00\x00\x00\x00\x00\x00", b"\x02\x03\x00\x00\x00\x00\x00\x00",
setattrs(Image(2, 3), pixels=bytes(6)), setattrs(Image(width=2, height=3), pixels=bytes(6)),
sample_building=Image(2, 3), sample_building=Image(width=2, height=3),
) )
def test_dataclass_struct_const_field() -> None: def test_dataclass_struct_const_field() -> None:
@dataclasses.dataclass class TestDataclass(DataclassStruct):
class TestContainer(DataclassMixin):
const_field: t.Optional[bytes] = csfield(cs.Const(b"\x00")) const_field: t.Optional[bytes] = csfield(cs.Const(b"\x00"))
common( common(
DataclassStruct(TestContainer), constr(TestDataclass),
bytes(1), bytes(1),
setattrs(TestContainer(), const_field=b"\x00"), setattrs(TestDataclass(), const_field=b"\x00"),
1, 1,
) )
assert ( assert (
raises( raises(
DataclassStruct(TestContainer).build, constr(TestDataclass).build,
setattrs(TestContainer(), const_field=b"\x01"), setattrs(TestDataclass(), const_field=b"\x01"),
) )
== cs.ConstError == cs.ConstError
) )
def test_dataclass_struct_array_field() -> None: def test_dataclass_struct_array_field() -> None:
@dataclasses.dataclass class TestDataclass(DataclassStruct):
class TestContainer(DataclassMixin):
array_field: t.List[int] = csfield(cs.Array(5, cs.Int8ub)) array_field: t.List[int] = csfield(cs.Array(5, cs.Int8ub))
common( common(
DataclassStruct(TestContainer), constr(TestDataclass),
bytes(5), bytes(5),
TestContainer(array_field=[0, 0, 0, 0, 0]), TestDataclass(array_field=[0, 0, 0, 0, 0]),
5, 5,
) )
def test_dataclass_struct_anonymus_fields_1() -> None: def test_dataclass_struct_anonymus_fields_1() -> None:
@dataclasses.dataclass class TestDataclass(DataclassStruct):
class TestContainer(DataclassMixin):
_1: t.Optional[bytes] = csfield(cs.Const(b"\x00")) _1: t.Optional[bytes] = csfield(cs.Const(b"\x00"))
_2: None = csfield(cs.Padding(1)) _2: None = csfield(cs.Padding(1))
_3: None = csfield(cs.Pass) _3: None = csfield(cs.Pass)
_4: None = csfield(cs.Terminated) _4: None = csfield(cs.Terminated)
common( common(
DataclassStruct(TestContainer), constr(TestDataclass),
bytes(2), bytes(2),
setattrs(TestContainer(), _1=b"\x00"), setattrs(TestDataclass(), _1=b"\x00"),
cs.SizeofError, cs.SizeofError,
) )
def test_dataclass_struct_anonymus_fields_2() -> None: def test_dataclass_struct_anonymus_fields_2() -> None:
@dataclasses.dataclass class TestDataclass(DataclassStruct):
class TestContainer(DataclassMixin): _1: t.Optional[int] = csfield(cs.Computed(7))
_1: int = csfield(cs.Computed(7))
_2: t.Optional[bytes] = csfield(cs.Const(b"JPEG")) _2: t.Optional[bytes] = csfield(cs.Const(b"JPEG"))
_3: None = csfield(cs.Pass) _3: None = csfield(cs.Pass)
_4: None = csfield(cs.Terminated) _4: None = csfield(cs.Terminated)
d = DataclassStruct(TestContainer) fmt = constr(TestDataclass)
assert d.build(TestContainer()) == d.build(TestContainer()) assert fmt.build(TestDataclass()) == fmt.build(TestDataclass())
def test_dataclass_struct_overloaded_method() -> None: def test_dataclass_struct_overloaded_method() -> None:
# Test dot access to some names that are not accessable via dot # Test dot access to some names that are not accessable via dot
# in the original 'cs.Container'. # in the original 'cs.Container'.
@dataclasses.dataclass class TestDataclass(DataclassStruct):
class TestContainer(DataclassMixin):
clear: int = csfield(cs.Int8ul) clear: int = csfield(cs.Int8ul)
copy: int = csfield(cs.Int8ul) copy: int = csfield(cs.Int8ul)
fromkeys: int = csfield(cs.Int8ul) fromkeys: int = csfield(cs.Int8ul)
@ -246,10 +226,10 @@ def test_dataclass_struct_overloaded_method() -> None:
update: int = csfield(cs.Int8ul) update: int = csfield(cs.Int8ul)
values: int = csfield(cs.Int8ul) values: int = csfield(cs.Int8ul)
d = DataclassStruct(TestContainer) fmt = constr(TestDataclass)
obj = d.parse( obj = fmt.parse(
d.build( fmt.build(
TestContainer( TestDataclass(
clear=1, clear=1,
copy=2, copy=2,
fromkeys=3, fromkeys=3,
@ -283,294 +263,172 @@ def test_dataclass_struct_overloaded_method() -> None:
assert obj.values == 14 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: def test_dataclass_struct_wrong_container() -> None:
@dataclasses.dataclass class TestContainer1(DataclassStruct):
class TestContainer1(DataclassMixin):
a: int = csfield(cs.Int16ub) a: int = csfield(cs.Int16ub)
b: int = csfield(cs.Int8ub) b: int = csfield(cs.Int8ub)
@dataclasses.dataclass class TestContainer2(DataclassStruct):
class TestContainer2(DataclassMixin):
a: int = csfield(cs.Int16ub) a: int = csfield(cs.Int16ub)
b: int = csfield(cs.Int8ub) b: int = csfield(cs.Int8ub)
assert ( assert raises(constr(TestContainer1).build, TestContainer2(a=1, b=2)) == TypeError
raises(DataclassStruct(TestContainer1).build, TestContainer2(a=1, b=2))
== TypeError
)
def test_dataclass_struct_doc() -> None: def test_dataclass_struct_doc() -> None:
@dataclasses.dataclass class TestDataclass1(DataclassStruct):
class TestContainer(DataclassMixin): """
a: int = csfield(cs.Int16ub, "This is the documentation of a") Documentation of TestDataclass1
b: int = csfield( """
cs.Int8ub, doc="This is the documentation of b\nwhich is multiline"
) 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( c: int = csfield(
cs.Int8ub, cs.Int8ub,
""" doc="""
This is the documentation of c This is the doc of c
which is also multiline which is also multiline
""", """,
) )
format = DataclassStruct(TestContainer) fmt1 = TestDataclass1.__constr__()
common(format, b"\x00\x01\x02\x03", TestContainer(a=1, b=2, c=3), 4) common(fmt1, b"\x00\x01\x02\x03", TestDataclass1(a=1, b=2, c=3), 4)
assert format.subcon.a.docs == "This is the documentation of a" assert fmt1.docs == "Documentation of TestDataclass1"
assert format.subcon.b.docs == "This is the documentation of b\nwhich is multiline" assert fmt1.subcon.a.docs == "This is the doc of a"
assert ( assert fmt1.subcon.b.docs == "This is the doc of b\nwhich is multiline"
format.subcon.c.docs assert fmt1.subcon.c.docs == "This is the doc of c\nwhich is also multiline"
== "This is the documentation of c\nwhich is also multiline"
) class TestDataclass2(DataclassStruct):
a: int = csfield(cs.Int16ub)
b: int = csfield(cs.Int8ub)
c: int = csfield(cs.Int8ub)
fmt2 = TestDataclass2.__constr__()
assert fmt2.docs == ""
assert fmt2.subcon.a.docs == ""
assert fmt2.subcon.b.docs == ""
assert fmt2.subcon.c.docs == ""
def test_dataclass_bitstruct() -> None: def test_dataclass_bitwise() -> None:
@dataclasses.dataclass class TestDataclass(DataclassStruct, constr=lambda cls: cs.Bitwise(cls)):
class TestContainer(DataclassMixin):
a: int = csfield(cs.BitsInteger(7)) a: int = csfield(cs.BitsInteger(7))
b: int = csfield(cs.Bit) b: int = csfield(cs.Bit)
c: int = csfield(cs.BitsInteger(8)) c: int = csfield(cs.BitsInteger(8))
print("")
common( common(
DataclassBitStruct(TestContainer), constr(TestDataclass),
b"\xFD\x12", b"\xFD\x12",
TestContainer(a=0x7E, b=1, c=0x12), TestDataclass(a=0x7E, b=1, c=0x12),
2, 2,
) )
# check __getattr__ # check __getattr__
c = DataclassStruct(TestContainer) fmt = TestDataclass.__constr__()
assert c.a.name == "a" assert fmt.subcon.a.name == "a"
assert c.b.name == "b" assert fmt.subcon.b.name == "b"
assert c.c.name == "c" assert fmt.subcon.c.name == "c"
assert isinstance(c.a.subcon, cs.BitsInteger) assert isinstance(fmt.subcon.a.subcon, cs.BitsInteger)
assert c.b.subcon is cs.Bit assert fmt.subcon.b.subcon is cs.Bit
assert isinstance(c.c.subcon, cs.BitsInteger) assert isinstance(fmt.subcon.c.subcon, cs.BitsInteger)
def test_dataclass_bitstruct() -> None:
class TestDataclass(DataclassBitStruct):
a: int = csfield(cs.BitsInteger(7))
b: int = csfield(cs.Bit)
c: int = csfield(cs.BitsInteger(8))
common(
constr(TestDataclass),
b"\xFD\x12",
TestDataclass(a=0x7E, b=1, c=0x12),
2,
)
# check __getattr__
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: def test_tenum() -> None:
class TestEnum(cst.EnumBase): class TestEnum(TEnum, subcon=cs.Byte):
one = 1 one = 1
two = 2 two = 2
four = 4 four = 4
eight = 8 eight = 8
d = cst.TEnum(cs.Byte, TestEnum) fmt = constr(TestEnum)
common(d, b"\x01", TestEnum.one, 1) common(fmt, b"\x01", TestEnum.one, 1)
common(d, b"\xff", TestEnum(255), 1) common(fmt, b"\xff", TestEnum(255), 1)
assert d.parse(b"\x01") == TestEnum.one assert fmt.parse(b"\x01") == TestEnum.one
assert d.parse(b"\x01") == 1 assert fmt.parse(b"\x01") == 1
assert int(d.parse(b"\x01")) == 1 assert int(fmt.parse(b"\x01")) == 1
assert d.parse(b"\xff") == TestEnum(255) assert fmt.parse(b"\xff") == TestEnum(255)
assert d.parse(b"\xff") == 255 assert fmt.parse(b"\xff") == 255
assert int(d.parse(b"\xff")) == 255 assert int(fmt.parse(b"\xff")) == 255
assert raises(d.build, 8) == TypeError assert raises(fmt.build, 8) == TypeError
def test_tenum_no_enumbase() -> None: def test_tenum_doc() -> None:
class E(enum.Enum): class TestEnum1(TEnum, subcon=cs.Byte):
a = 1 """
b = 2 TestEnum documentation
"""
cls = t.cast(t.Type[cst.EnumBase], E)
assert raises(lambda: cst.TEnum(cs.Byte, cls)) == TypeError
def test_tenum_asdict() -> None:
# see: https://github.com/timrid/construct-typing/issues/21
import dataclasses
import construct_typed as cst
class TestEnum(cst.EnumBase):
one = 1 one = 1
d1 = constr(TestEnum1)
assert d1.docs == "TestEnum documentation"
class TestEnum2(TEnum, subcon=cs.Byte):
two = 2 two = 2
four = 4
eight = 8
@dataclasses.dataclass d2 = constr(TestEnum2)
class SomeDataclass: assert d2.docs == ""
a: TestEnum
dc = SomeDataclass(TestEnum.one)
dc_dict = dataclasses.asdict(dc)
assert dc_dict["a"] == dc.a
assert dc_dict["a"] is dc.a
dc = SomeDataclass(TestEnum(5))
dc_dict = dataclasses.asdict(dc)
assert dc_dict["a"] == dc.a
assert dc_dict["a"] is dc.a
def test_tenum_docstring() -> None: def test_tenum_in_dataclass_struct() -> None:
class TestEnum(cst.EnumBase): class TestEnum(TEnum, subcon=cs.Int8ub):
"""
This is an test enum.
"""
Value_WithDoc = cst.EnumValue(0, doc="an enum with a documentation")
Value_WithMultilineDoc = cst.EnumValue(
1,
"""
An enum with a multiline documentation...
...next line...
""",
)
Value_NoDoc = cst.EnumValue(2)
Value_NoDoc2 = 3
assert TestEnum.__doc__ is not None
assert textwrap.dedent(TestEnum.__doc__) == textwrap.dedent(
"""
This is an test enum.
"""
)
assert TestEnum.Value_WithDoc.__doc__ == "an enum with a documentation"
assert (
TestEnum.Value_WithMultilineDoc.__doc__
== """
An enum with a multiline documentation...
...next line...
"""
)
assert TestEnum.Value_NoDoc.__doc__ == ""
assert TestEnum.Value_NoDoc2.__doc__ == ""
assert TestEnum(5).__doc__ == "missing value"
def test_dataclass_struct_wrong_enumbase() -> None:
class E1(cst.EnumBase):
a = 1 a = 1
b = 2 b = 2
class E2(cst.EnumBase): class TestDataclass(DataclassStruct):
a = 1 a: TestEnum = csfield(constr(TestEnum))
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))
b: int = csfield(cs.Int8ub) b: int = csfield(cs.Int8ub)
common( common(
DataclassStruct(TestContainer), constr(TestDataclass),
b"\x01\x02", b"\x01\x02",
TestContainer(a=TestEnum.a, b=2), TestDataclass(a=TestEnum.a, b=2),
2, 2,
) )
assert ( assert (
raises(cst.TEnum(cs.Byte, TestEnum).build, TestContainer(a=1, b=2)) == TypeError # type: ignore raises(constr(TestEnum).build, TestDataclass(a=1, b=2)) == TypeError # type: ignore
) )
def test_tenum_flags() -> None: def test_tflags() -> None:
class TestEnum(cst.FlagsEnumBase): class TestFlags(TFlags, subcon=cs.Byte):
one = 1 one = 1
two = 2 two = 2
four = 4 four = 4
eight = 8 eight = 8
d = cst.TFlagsEnum(cs.Byte, TestEnum) fmt = constr(TestFlags)
common(d, b"\x03", TestEnum.one | TestEnum.two, 1) common(fmt, b"\x03", TestFlags.one | TestFlags.two, 1)
assert d.build(TestEnum(0)) == b"\x00" assert fmt.build(TestFlags(0)) == b"\x00"
assert d.build(TestEnum.one | TestEnum.two) == b"\x03" assert fmt.build(TestFlags.one | TestFlags.two) == b"\x03"
assert d.build(TestEnum(8)) == b"\x08" assert fmt.build(TestFlags(8)) == b"\x08"
assert d.build(TestEnum(1 | 2)) == b"\x03" assert fmt.build(TestFlags(1 | 2)) == b"\x03"
assert d.build(TestEnum(255)) == b"\xff" assert fmt.build(TestFlags(255)) == b"\xff"
assert d.build(TestEnum.eight) == b"\x08" assert fmt.build(TestFlags.eight) == b"\x08"
assert raises(d.build, 2) == TypeError assert raises(fmt.build, 2) == TypeError
def test_tenum_flags_asdict() -> None:
import dataclasses
import construct_typed as cst
class TestEnum(cst.FlagsEnumBase):
one = 1
two = 2
four = 4
eight = 8
@dataclasses.dataclass
class SomeDataclass:
a: TestEnum
dc = SomeDataclass(TestEnum.one)
dc_dict = dataclasses.asdict(dc)
assert dc_dict["a"] == dc.a
assert dc_dict["a"] is dc.a
dc = SomeDataclass(TestEnum(5))
dc_dict = dataclasses.asdict(dc)
assert dc_dict["a"] == dc.a
assert dc_dict["a"] is dc.a
def test_tenum_flags_docstring() -> None:
class TestEnum(cst.FlagsEnumBase):
"""
This is an test flags enum.
"""
Value_WithDoc = cst.EnumValue(0, doc="an enum with a documentation")
Value_WithMultilineDoc = cst.EnumValue(
1,
"""
An enum with a multiline documentation...
...next line...
""",
)
Value_NoDoc = cst.EnumValue(2)
Value_NoDoc2 = 4
assert TestEnum.__doc__ is not None
assert textwrap.dedent(TestEnum.__doc__) == textwrap.dedent(
"""
This is an test flags enum.
"""
)
assert TestEnum.Value_WithDoc.__doc__ == "an enum with a documentation"
assert (
TestEnum.Value_WithMultilineDoc.__doc__
== """
An enum with a multiline documentation...
...next line...
"""
)
assert TestEnum.Value_NoDoc.__doc__ == ""
assert TestEnum.Value_NoDoc2.__doc__ == ""
assert TestEnum(8).__doc__ == "missing value"