diff --git a/.github/workflows/main.yml b/.github/workflows/main.yml index 4a3e3db..e09f947 100644 --- a/.github/workflows/main.yml +++ b/.github/workflows/main.yml @@ -1,6 +1,10 @@ name: CI -on: [push, pull_request] +on: + push: + pull_request: + workflow_dispatch: + workflow_call: jobs: build: @@ -8,7 +12,7 @@ jobs: strategy: matrix: os: ['ubuntu-latest', 'windows-latest'] - python-version: [ '3.7', '3.8', '3.9' ] + python-version: [ '3.9', '3.10', '3.11', '3.12', '3.13' ] runs-on: ${{ matrix.os }} name: OS ${{ matrix.os }}, Python ${{ matrix.python-version }} @@ -16,25 +20,26 @@ jobs: steps: # Checks out a copy of your repository on the machine - name: Checkout code - uses: actions/checkout@v1 + uses: actions/checkout@v3 # Setup python - name: Setup python - uses: actions/setup-python@v1 + uses: actions/setup-python@v4 with: python-version: ${{ matrix.python-version }} architecture: x64 # Setup node.js (for pyright) - name: Setup node.js (for pyright) - uses: actions/setup-node@v2 + uses: actions/setup-node@v3 with: - node-version: '14' + node-version: 16 # Install pyright - name: Install pyright run: | npm install -g pyright + pyright --version # Install this package - name: Install this package @@ -61,3 +66,30 @@ jobs: - name: Run pyright run: | 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/ \ No newline at end of file diff --git a/.github/workflows/python-publish.yml b/.github/workflows/python-publish.yml index 4e1ef42..ea14263 100644 --- a/.github/workflows/python-publish.yml +++ b/.github/workflows/python-publish.yml @@ -1,6 +1,3 @@ -# 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 on: @@ -8,24 +5,26 @@ on: types: [created] jobs: - deploy: + create_wheel_and_sdist: + name: create_wheel_and_sdist + uses: ./.github/workflows/main.yml + deploy: + needs: [ create_wheel_and_sdist ] runs-on: ubuntu-latest + + environment: pypi + permissions: + id-token: write # IMPORTANT: this permission is mandatory for Trusted Publishing steps: - - uses: actions/checkout@v2 - - name: Set up Python - uses: actions/setup-python@v2 + - uses: actions/checkout@v3 + + - name: Download artifacts + uses: actions/download-artifact@v4 with: - python-version: '3.x' - - name: Install dependencies - run: | - python -m pip install --upgrade pip - 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/* + name: Package-Distributions-construct-typing + path: ./dist + + - name: Publish package distributions to PyPI + uses: pypa/gh-action-pypi-publish@release/v1 diff --git a/.gitignore b/.gitignore index b3d4398..1d9e0fe 100644 --- a/.gitignore +++ b/.gitignore @@ -129,3 +129,6 @@ dmypy.json example_737 example_888 example_ksy.ksy + +# Test stuff +devtest/ \ No newline at end of file diff --git a/.vscode/launch.json b/.vscode/launch.json index 8655026..d94289d 100644 --- a/.vscode/launch.json +++ b/.vscode/launch.json @@ -9,6 +9,12 @@ "type": "python", "request": "launch", "program": "${file}", + "console": "integratedTerminal" + }, + { + "name": "Debug Tests", + "type": "python", + "request": "test", "console": "integratedTerminal", "justMyCode": false } diff --git a/.vscode/settings.json b/.vscode/settings.json index b265945..75b5040 100644 --- a/.vscode/settings.json +++ b/.vscode/settings.json @@ -1,14 +1,6 @@ { + // static analysis "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.autoImportCompletions": false, "python.analysis.diagnosticSeverityOverrides": { @@ -16,7 +8,16 @@ "reportUntypedNamedTuple": "information", }, - // configure pytest + // formating + "python.formatting.provider": "black", + + // sorting + "python.sortImports.path": "isort", + "python.sortImports.args": [ + "--profile=black", + ], + + // tests "python.testing.pytestArgs": [ "tests" ], diff --git a/README.md b/README.md index b11989c..c4bac18 100644 --- a/README.md +++ b/README.md @@ -1,3 +1,12 @@ +## 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 [![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) diff --git a/construct-stubs/__init__.pyi b/construct-stubs/__init__.pyi index 858384d..2cad48a 100644 --- a/construct-stubs/__init__.pyi +++ b/construct-stubs/__init__.pyi @@ -3,6 +3,10 @@ from construct.debug import * from construct.expr import * from construct.lib import * from construct.version import * +from construct import lib + +__author__: str +__version__: str #=============================================================================== # exposed names diff --git a/construct-stubs/core.pyi b/construct-stubs/core.pyi index 010f471..7ed1af6 100644 --- a/construct-stubs/core.pyi +++ b/construct-stubs/core.pyi @@ -17,6 +17,10 @@ from construct.lib import ( ListType, RebufferedBytesIO, ) +from cryptography.hazmat.primitives.ciphers import Cipher +from cryptography.hazmat.primitives.ciphers.aead import AESCCM, AESGCM, ChaCha20Poly1305 +from cryptography.hazmat.primitives.ciphers.modes import Mode +from typing_extensions import Buffer, TypeAlias # unfortunately, there are a few duplications with "typing", e.g. Union and Optional, which is why the t. prefix must be used everywhere @@ -25,7 +29,8 @@ from construct.lib import ( # - Higher Kinded Types: https://github.com/python/typing/issues/548 # - Higher Kinded Types: https://sobolevn.me/2020/10/higher-kinded-types-in-python -StreamType = t.BinaryIO +ReadableBuffer: TypeAlias = Buffer +StreamType = t.IO[bytes] FilenameType = t.Union[str, bytes, os.PathLike[str], os.PathLike[bytes]] PathType = str ContextKWType = t.Any @@ -65,25 +70,37 @@ class RawCopyError(ConstructError): ... class RotationError(ConstructError): ... class ChecksumError(ConstructError): ... class CancelParsing(ConstructError): ... +class CipherError(ConstructError): ... # =============================================================================== # used internally # =============================================================================== def stream_read( - stream: t.BinaryIO, length: int, path: t.Optional[PathType] + stream: StreamType, length: int, path: t.Optional[PathType] ) -> bytes: ... -def stream_read_entire(stream: t.BinaryIO, path: t.Optional[PathType]) -> bytes: ... +def stream_read_entire(stream: StreamType, path: t.Optional[PathType]) -> bytes: ... def stream_write( - stream: t.BinaryIO, data: bytes, length: int, path: t.Optional[PathType] + stream: StreamType, data: bytes, length: int, path: t.Optional[PathType] ) -> None: ... def stream_seek( - stream: t.BinaryIO, offset: int, whence: int, path: t.Optional[PathType] + stream: StreamType, offset: int, whence: int, path: t.Optional[PathType] ) -> int: ... -def stream_tell(stream: t.BinaryIO, path: t.Optional[PathType]) -> int: ... -def stream_size(stream: t.BinaryIO) -> int: ... -def stream_iseof(stream: t.BinaryIO) -> bool: ... +def stream_tell(stream: StreamType, path: t.Optional[PathType]) -> int: ... +def stream_size(stream: StreamType) -> int: ... +def stream_iseof(stream: StreamType) -> bool: ... def evaluate(param: ConstantOrContextLambda2[T], context: Context) -> T: ... +class BytesIOWithOffsets(io.BytesIO): + @staticmethod + def from_reading( + stream: StreamType, length: int, path: PathType + ) -> BytesIOWithOffsets: ... + def __init__( + self, contents: bytes, parent_stream: StreamType, offset: int + ) -> None: ... + def tell(self) -> int: ... + def seek(self, offset: int, whence: int = ...) -> int: ... + # =============================================================================== # abstract constructs # =============================================================================== @@ -95,7 +112,7 @@ class Construct(t.Generic[ParsedType, BuildTypes]): docs: str flagbuildnone: bool parsed: t.Optional[t.Callable[[ParsedType, Context], None]] - def parse(self, data: bytes, **contextkw: ContextKWType) -> ParsedType: ... + def parse(self, data: ReadableBuffer, **contextkw: ContextKWType) -> ParsedType: ... def parse_stream( self, stream: StreamType, **contextkw: ContextKWType ) -> ParsedType: ... @@ -105,15 +122,17 @@ class Construct(t.Generic[ParsedType, BuildTypes]): def build(self, obj: BuildTypes, **contextkw: ContextKWType) -> bytes: ... def build_stream( self, obj: BuildTypes, stream: StreamType, **contextkw: ContextKWType - ) -> bytes: ... + ) -> None: ... def build_file( self, obj: BuildTypes, filename: FilenameType, **contextkw: ContextKWType - ) -> bytes: ... + ) -> None: ... def sizeof(self, **contextkw: ContextKWType) -> int: ... def compile( self, filename: FilenameType = ... ) -> Construct[ParsedType, BuildTypes]: ... - def benchmark(self, sampledata: bytes, filename: FilenameType = ...) -> str: ... + def benchmark( + self, sampledata: ReadableBuffer, filename: FilenameType = ... + ) -> str: ... def export_ksy( self, schemaname: str = ..., filename: FilenameType = ... ) -> str: ... @@ -129,20 +148,22 @@ class Construct(t.Generic[ParsedType, BuildTypes]): self, other: t.Union[str, bytes, t.Callable[[ParsedType, Context], None]], ) -> Renamed[ParsedType, BuildTypes]: ... - def __add__( - self, other: Construct[t.Any, t.Any] - ) -> Struct[Container[t.Any], t.Optional[t.Dict[str, t.Any]]]: ... - def __rshift__( - self, other: Construct[t.Any, t.Any] - ) -> Sequence[ListContainer[t.Any], t.Optional[t.List[t.Any]]]: ... - def __getitem__( - self, count: t.Union[int, t.Callable[[Context], int]] - ) -> Array[ + def __add__(self, other: Construct[t.Any, t.Any]) -> Struct: ... + def __rshift__(self, other: Construct[t.Any, t.Any]) -> Sequence: ... + def __getitem__(self, count: t.Union[int, t.Callable[[Context], int]]) -> Array[ ParsedType, BuildTypes, - ListContainer[ParsedType], - t.List[BuildTypes], ]: ... + def _parse( + self, stream: StreamType, context: Context, path: PathType + ) -> ParsedType: ... + def _parsereport( + self, stream: StreamType, context: Context, path: PathType + ) -> ParsedType: ... + def _build( + self, obj: BuildTypes, stream: StreamType, context: Context, path: PathType + ) -> int: ... + def _sizeof(self, context: Context, path: PathType) -> int: ... @t.type_check_only class Context(Container[t.Any]): @@ -169,22 +190,23 @@ class Subconstruct( ): subcon: Construct[SubconParsedType, SubconBuildTypes] @t.overload - def __new__( - cls, subcon: Construct[SubconParsedType, SubconBuildTypes] - ) -> Subconstruct[ - SubconParsedType, SubconBuildTypes, SubconParsedType, SubconBuildTypes - ]: ... + def __init__( + self, + subcon: Construct[SubconParsedType, SubconBuildTypes], + ) -> None: ... @t.overload - def __new__( - cls, *args: t.Any, **kwargs: t.Any - ) -> Subconstruct[SubconParsedType, SubconBuildTypes, ParsedType, BuildTypes]: ... + def __init__( # type: ignore + self, + *args: t.Any, + **kwargs: t.Any, + ) -> None: ... class Adapter( Subconstruct[SubconParsedType, SubconBuildTypes, ParsedType, BuildTypes], ): - def __new__( - cls, subcon: Construct[SubconParsedType, SubconBuildTypes] - ) -> Adapter[SubconParsedType, SubconBuildTypes, ParsedType, BuildTypes]: ... + def __init__( + self, subcon: Construct[SubconParsedType, SubconBuildTypes] + ) -> None: ... def _decode( self, obj: SubconBuildTypes, context: Context, path: PathType ) -> ParsedType: ... @@ -211,28 +233,34 @@ class Tunnel( def _decode(self, data: bytes, context: Context, path: PathType) -> bytes: ... def _encode(self, data: bytes, context: Context, path: PathType) -> bytes: ... -# TODO: Compiled +class Compiled(Construct[t.Any, t.Any]): + source: t.Optional[str] + defersubcon: t.Optional[Construct[t.Any, t.Any]] + parsefunc: t.Callable[[StreamType, Context], t.Any] + buildfunc: t.Callable[[t.Any, StreamType, Context], t.Any] + def __init__( + self, + parsefunc: t.Callable[[StreamType, Context], t.Any], + buildfunc: t.Callable[[t.Any, StreamType, Context], t.Any], + ) -> None: ... # =============================================================================== # bytes and bits # =============================================================================== -class Bytes(Construct[ParsedType, BuildTypes]): +class Bytes(Construct[bytes, t.Union[bytes, bytearray, int]]): length: ConstantOrContextLambda[int] - def __new__( - cls, length: ConstantOrContextLambda[int] - ) -> Bytes[bytes, t.Union[bytes, int]]: ... + def __init__( + self, + length: ConstantOrContextLambda[int], + ) -> None: ... -GreedyBytes: Construct[bytes, bytes] +GreedyBytes: Construct[bytes, t.Union[bytes, bytearray]] -def Bitwise( - subcon: Construct[SubconParsedType, SubconBuildTypes] -) -> t.Union[ +def Bitwise(subcon: Construct[SubconParsedType, SubconBuildTypes]) -> t.Union[ Transformed[SubconParsedType, SubconBuildTypes], Restreamed[SubconParsedType, SubconBuildTypes], ]: ... -def Bytewise( - subcon: Construct[SubconParsedType, SubconBuildTypes] -) -> t.Union[ +def Bytewise(subcon: Construct[SubconParsedType, SubconBuildTypes]) -> t.Union[ Transformed[SubconParsedType, SubconBuildTypes], Restreamed[SubconParsedType, SubconBuildTypes], ]: ... @@ -250,46 +278,61 @@ class FormatField(Construct[ParsedType, BuildTypes]): FORMAT_BOOL = t.Literal["?"] @t.overload def __new__( - cls, endianity: str, format: FORMAT_INT + cls: "type[FormatField[int, int]]", + endianity: str, + format: FORMAT_INT, ) -> FormatField[int, int]: ... @t.overload def __new__( - cls, endianity: str, format: FORMAT_FLOAT + cls: "type[FormatField[float, float]]", + endianity: str, + format: FORMAT_FLOAT, ) -> FormatField[float, float]: ... @t.overload def __new__( - cls, endianity: str, format: FORMAT_BOOL + cls: "type[FormatField[bool, bool]]", + endianity: str, + format: FORMAT_BOOL, ) -> FormatField[bool, bool]: ... @t.overload - def __new__(cls, endianity: str, format: str) -> FormatField[t.Any, t.Any]: ... + def __new__( + cls: "type[FormatField[t.Any, t.Any]]", + endianity: str, + format: str, + ) -> FormatField[t.Any, t.Any]: ... + else: - def __new__(cls, endianity: str, format: str) -> FormatField[t.Any, t.Any]: ... + def __new__( + cls: "type[FormatField[t.Any, t.Any]]", + endianity: str, + format: str, + ) -> FormatField[t.Any, t.Any]: ... -class BytesInteger(Construct[ParsedType, BuildTypes]): +class BytesInteger(Construct[int, int]): length: ConstantOrContextLambda[int] signed: bool swapped: ConstantOrContextLambda[bool] - def __new__( - cls, + def __init__( + self, length: ConstantOrContextLambda[int], signed: bool = ..., swapped: ConstantOrContextLambda[bool] = ..., - ) -> BytesInteger[int, int]: ... + ) -> None: ... -class BitsInteger(Construct[ParsedType, BuildTypes]): +class BitsInteger(Construct[int, int]): length: ConstantOrContextLambda[int] signed: bool swapped: ConstantOrContextLambda[bool] - def __new__( - cls, + def __init__( + self, length: ConstantOrContextLambda[int], signed: bool = ..., swapped: ConstantOrContextLambda[bool] = ..., - ) -> BitsInteger[int, int]: ... + ) -> None: ... -Bit: BitsInteger[int, int] -Nibble: BitsInteger[int, int] -Octet: BitsInteger[int, int] +Bit: BitsInteger +Nibble: BitsInteger +Octet: BitsInteger Int8ub: FormatField[int, int] Int16ub: FormatField[int, int] @@ -335,12 +378,12 @@ Half: FormatField[float, float] Single: FormatField[float, float] Double: FormatField[float, float] -Int24ub: BytesInteger[int, int] -Int24ul: BytesInteger[int, int] -Int24un: BytesInteger[int, int] -Int24sb: BytesInteger[int, int] -Int24sl: BytesInteger[int, int] -Int24sn: BytesInteger[int, int] +Int24ub: BytesInteger +Int24ul: BytesInteger +Int24un: BytesInteger +Int24sb: BytesInteger +Int24sl: BytesInteger +Int24sn: BytesInteger VarInt: Construct[int, int] ZigZag: Construct[int, int] @@ -348,7 +391,9 @@ ZigZag: Construct[int, int] # =============================================================================== # strings # =============================================================================== -class StringEncoded(Construct[ParsedType, BuildTypes]): +possiblestringencodings: t.Dict[str, int] + +class StringEncoded(Construct[str, str]): if sys.version_info >= (3, 8): ENCODING_1 = t.Literal["ascii", "utf8", "utf_8", "u8"] ENCODING_2 = t.Literal["utf16", "utf_16", "u16", "utf_16_be", "utf_16_le"] @@ -357,18 +402,20 @@ class StringEncoded(Construct[ParsedType, BuildTypes]): else: ENCODING = str encoding: ENCODING - def __new__( - cls, subcon: Construct[ParsedType, BuildTypes], encoding: ENCODING - ) -> StringEncoded[str, str]: ... + def __init__( + self, + subcon: Construct[bytes, bytes], + encoding: ENCODING, + ) -> None: ... def PaddedString( length: ConstantOrContextLambda[int], encoding: StringEncoded.ENCODING -) -> StringEncoded[str, str]: ... +) -> StringEncoded: ... def PascalString( lengthfield: Construct[int, int], encoding: StringEncoded.ENCODING -) -> StringEncoded[str, str]: ... -def CString(encoding: StringEncoded.ENCODING) -> StringEncoded[str, str]: ... -def GreedyString(encoding: StringEncoded.ENCODING) -> StringEncoded[str, str]: ... +) -> StringEncoded: ... +def CString(encoding: StringEncoded.ENCODING) -> StringEncoded: ... +def GreedyString(encoding: StringEncoded.ENCODING) -> StringEncoded: ... # =============================================================================== # mappings @@ -381,58 +428,68 @@ class EnumIntegerString(str): @staticmethod def new(intvalue: int, stringvalue: str) -> EnumIntegerString: ... -class Enum(Adapter[int, int, ParsedType, BuildTypes]): +class Enum( + Adapter[int, int, t.Union[EnumInteger, EnumIntegerString], t.Union[int, str]] +): encmapping: t.Dict[str, int] decmapping: t.Dict[int, EnumIntegerString] ksymapping: t.Dict[int, str] - def __new__( - cls, + def __init__( + self, subcon: Construct[int, int], *merge: t.Union[t.Type[enum.IntEnum], t.Type[enum.IntFlag]], - **mapping: int - ) -> Enum[t.Union[EnumInteger, EnumIntegerString], t.Union[int, str]]: ... + **mapping: int, + ) -> None: ... def __getattr__(self, name: str) -> EnumIntegerString: ... class BitwisableString(str): def __or__(self, other: BitwisableString) -> BitwisableString: ... -class FlagsEnum(Adapter[int, int, ParsedType, BuildTypes]): +class FlagsEnum( + Adapter[int, int, Container[bool], t.Union[int, str, t.Dict[str, bool]]] +): flags: t.Dict[str, int] reverseflags: t.Dict[int, str] - def __new__( - cls, + def __init__( + self, subcon: Construct[int, int], *merge: t.Union[t.Type[enum.IntEnum], t.Type[enum.IntFlag]], - **flags: int - ) -> FlagsEnum[Container[bool], t.Union[int, str, t.Dict[str, bool]]]: ... + **flags: int, + ) -> None: ... def __getattr__(self, name: str) -> BitwisableString: ... class Mapping(Adapter[SubconParsedType, SubconBuildTypes, t.Any, t.Any]): decmapping: t.Dict[int, str] encmapping: t.Dict[str, int] - def __new__( - cls, + def __init__( + self, subcon: Construct[SubconParsedType, SubconBuildTypes], mapping: t.Dict[t.Any, t.Any], - ) -> Mapping[t.Any, t.Any]: ... + ) -> None: ... # =============================================================================== # structures and sequences # =============================================================================== # this can maybe made better when variadic generics are available -class Struct(Construct[ParsedType, BuildTypes]): +class Struct(Construct[Container[t.Any], t.Optional[t.Dict[str, t.Any]]]): subcons: t.List[Construct[t.Any, t.Any]] - def __new__( - cls, *subcons: Construct[t.Any, t.Any], **subconskw: Construct[t.Any, t.Any] - ) -> Struct[Container[t.Any], t.Optional[t.Dict[str, t.Any]]]: ... + _subcons: t.Dict[str, Construct[t.Any, t.Any]] + def __init__( + self, + *subcons: Construct[t.Any, t.Any], + **subconskw: Construct[t.Any, t.Any], + ) -> None: ... def __getattr__(self, name: str) -> t.Any: ... # this can maybe made better when variadic generics are available -class Sequence(Construct[ParsedType, BuildTypes]): +class Sequence(Construct[ListContainer[t.Any], t.Optional[t.List[t.Any]]]): subcons: t.List[Construct[t.Any, t.Any]] - def __new__( - cls, *subcons: Construct[t.Any, t.Any], **subconskw: Construct[t.Any, t.Any] - ) -> Sequence[ListContainer[t.Any], t.Optional[t.List[t.Any]]]: ... + _subcons: t.Dict[str, Construct[t.Any, t.Any]] + def __init__( + self, + *subcons: Construct[t.Any, t.Any], + **subconskw: Construct[t.Any, t.Any], + ) -> None: ... def __getattr__(self, name: str) -> t.Any: ... # =============================================================================== @@ -442,48 +499,40 @@ class Array( Subconstruct[ SubconParsedType, SubconBuildTypes, - ParsedType, - BuildTypes, + ListContainer[SubconParsedType], # type: ignore + t.List[SubconBuildTypes], # type: ignore ] ): count: ConstantOrContextLambda[int] discard: bool - def __new__( - cls, + def __init__( + self, count: ConstantOrContextLambda[int], subcon: Construct[SubconParsedType, SubconBuildTypes], discard: bool = ..., - ) -> Array[ - SubconParsedType, - SubconBuildTypes, - ListContainer[SubconParsedType], - t.List[SubconBuildTypes], - ]: ... + ) -> None: ... class GreedyRange( Subconstruct[ SubconParsedType, SubconBuildTypes, - ParsedType, - BuildTypes, + ListContainer[SubconParsedType], # type: ignore + t.List[SubconBuildTypes], # type: ignore ] ): discard: bool - def __new__( - cls, subcon: Construct[SubconParsedType, SubconBuildTypes], discard: bool = ... - ) -> GreedyRange[ - SubconParsedType, - SubconBuildTypes, - ListContainer[SubconParsedType], - t.List[SubconBuildTypes], - ]: ... + def __init__( + self, + subcon: Construct[SubconParsedType, SubconBuildTypes], + discard: bool = ..., + ) -> None: ... class RepeatUntil( Subconstruct[ SubconParsedType, SubconBuildTypes, - ListContainer[SubconParsedType], - t.List[SubconBuildTypes], + ListContainer[SubconParsedType], # type: ignore + t.List[SubconBuildTypes], # type: ignore ] ): predicate: t.Union[ @@ -520,67 +569,69 @@ class Renamed( # =============================================================================== # miscellaneous # =============================================================================== -class Const(Subconstruct[SubconParsedType, SubconBuildTypes, ParsedType, BuildTypes]): - value: SubconBuildTypes +class Const(Subconstruct[t.Any, t.Any, ParsedType, BuildTypes]): + value: BuildTypes @t.overload def __new__( - cls, + cls: "type[Const[bytes, t.Optional[bytes]]]", value: bytes, - ) -> Const[None, None, bytes, t.Optional[bytes]]: ... + ) -> Const[bytes, t.Optional[bytes]]: ... @t.overload def __new__( - cls, + cls: "type[Const[SubconParsedType, t.Optional[SubconBuildTypes]]]", value: SubconBuildTypes, subcon: Construct[SubconParsedType, SubconBuildTypes], - ) -> Const[None, None, SubconParsedType, t.Optional[SubconBuildTypes]]: ... + ) -> Const[SubconParsedType, t.Optional[SubconBuildTypes]]: ... -class Computed(Construct[ParsedType, BuildTypes]): +class Computed(Construct[ParsedType, None]): func: ConstantOrContextLambda2[ParsedType] - @t.overload - def __new__( - cls, func: ConstantOrContextLambda2[ParsedType] - ) -> Computed[ParsedType, None]: ... - @t.overload - def __new__( - cls, func: ConstantOrContextLambda2[t.Any] - ) -> Computed[t.Any, None]: ... + def __init__( + self, + func: ConstantOrContextLambda2[ParsedType], + ) -> None: ... Index: Construct[int, t.Any] -class Rebuild(Subconstruct[SubconParsedType, SubconBuildTypes, ParsedType, BuildTypes]): +class Rebuild(Subconstruct[SubconParsedType, SubconBuildTypes, SubconParsedType, None]): func: ConstantOrContextLambda[SubconBuildTypes] - def __new__( - cls, + def __init__( + self, subcon: Construct[SubconParsedType, SubconBuildTypes], func: ConstantOrContextLambda[SubconBuildTypes], - ) -> Rebuild[SubconParsedType, SubconBuildTypes, SubconParsedType, None]: ... + ) -> None: ... -class Default(Subconstruct[SubconParsedType, SubconBuildTypes, ParsedType, BuildTypes]): - value: ConstantOrContextLambda[SubconBuildTypes] - def __new__( - cls, - subcon: Construct[SubconParsedType, SubconBuildTypes], - value: ConstantOrContextLambda[SubconBuildTypes], - ) -> Default[ +class Default( + Subconstruct[ SubconParsedType, SubconBuildTypes, SubconParsedType, t.Optional[SubconBuildTypes], - ]: ... + ] +): + value: ConstantOrContextLambda[SubconBuildTypes] + def __init__( + self, + subcon: Construct[SubconParsedType, SubconBuildTypes], + value: ConstantOrContextLambda[SubconBuildTypes], + ) -> None: ... -class Check(Construct[ParsedType, BuildTypes]): +class Check(Construct[None, None]): func: ConstantOrContextLambda[bool] - def __new__(cls, func: ConstantOrContextLambda[bool]) -> Check[None, None]: ... + def __init__( + self, + func: ConstantOrContextLambda[bool], + ) -> None: ... -Error: Construct[None, None] +Error: Construct[t.NoReturn, t.NoReturn] class FocusedSeq(Construct[t.Any, t.Any]): subcons: t.List[Construct[t.Any, t.Any]] + _subcons: t.Dict[str, Construct[t.Any, t.Any]] def __init__( self, parsebuildfrom: ConstantOrContextLambda[str], *subcons: Construct[t.Any, t.Any], - **subconskw: Construct[t.Any, t.Any] + **subconskw: Construct[t.Any, t.Any], ) -> None: ... def __getattr__(self, name: str) -> t.Any: ... @@ -592,24 +643,19 @@ class NamedTuple( Adapter[ SubconParsedType, SubconBuildTypes, - ParsedType, - BuildTypes, + t.Tuple[t.Any, ...], + t.Union[t.Tuple[t.Any, ...], t.List[t.Any], t.Dict[str, t.Any]], ] ): tuplename: str tuplefields: str factory: Construct[SubconParsedType, SubconBuildTypes] - def __new__( - cls, + def __init__( + self, tuplename: str, tuplefields: str, subcon: Construct[SubconParsedType, SubconBuildTypes], - ) -> NamedTuple[ - SubconParsedType, - SubconBuildTypes, - t.Tuple[t.Any, ...], - t.Union[t.Tuple[t.Any, ...], t.List[t.Any], t.Dict[str, t.Any]], - ]: ... + ) -> None: ... if sys.version_info >= (3, 8): MSDOS = t.Literal["msdos"] @@ -640,63 +686,60 @@ def Timestamp( K = t.TypeVar("K") V = t.TypeVar("V") -class Hex(Adapter[SubconParsedType, SubconBuildTypes, ParsedType, BuildTypes]): +class Hex(Adapter[t.Any, t.Any, ParsedType, BuildTypes]): @t.overload def __new__( - cls, subcon: Construct[int, BuildTypes] - ) -> Hex[int, BuildTypes, HexDisplayedInteger, BuildTypes]: ... + cls: "type[Hex[HexDisplayedInteger, BuildTypes]]", + subcon: Construct[int, BuildTypes], + ) -> Hex[HexDisplayedInteger, BuildTypes]: ... @t.overload def __new__( - cls, subcon: Construct[bytes, BuildTypes] - ) -> Hex[bytes, BuildTypes, HexDisplayedBytes, BuildTypes]: ... + cls: "type[Hex[HexDisplayedBytes, BuildTypes]]", + subcon: Construct[bytes, BuildTypes], + ) -> Hex[HexDisplayedBytes, BuildTypes]: ... @t.overload def __new__( - cls, subcon: Construct[RawCopyObj[SubconParsedType], BuildTypes] + cls: "type[Hex[HexDisplayedDict[str, t.Union[int, bytes, SubconParsedType]], BuildTypes,]]", + subcon: Construct[RawCopyObj[SubconParsedType], BuildTypes], ) -> Hex[ - RawCopyObj[SubconParsedType], - BuildTypes, HexDisplayedDict[str, t.Union[int, bytes, SubconParsedType]], BuildTypes, ]: ... @t.overload def __new__( - cls, subcon: Construct[Container[t.Any], BuildTypes] - ) -> Hex[ - Container[t.Any], BuildTypes, HexDisplayedDict[str, t.Any], BuildTypes - ]: ... + cls: "type[Hex[HexDisplayedDict[str, t.Any], BuildTypes]]", + subcon: Construct[Container[t.Any], BuildTypes], + ) -> Hex[HexDisplayedDict[str, t.Any], BuildTypes]: ... @t.overload def __new__( - cls, subcon: Construct[SubconParsedType, SubconBuildTypes] - ) -> Hex[ - SubconParsedType, SubconBuildTypes, SubconParsedType, SubconBuildTypes - ]: ... + cls: "type[Hex[SubconParsedType, SubconBuildTypes]]", + subcon: Construct[SubconParsedType, SubconBuildTypes], + ) -> Hex[SubconParsedType, SubconBuildTypes]: ... -class HexDump(Adapter[SubconParsedType, SubconBuildTypes, ParsedType, BuildTypes]): +class HexDump(Adapter[t.Any, t.Any, ParsedType, BuildTypes]): @t.overload def __new__( - cls, subcon: Construct[bytes, BuildTypes] - ) -> HexDump[bytes, BuildTypes, HexDumpDisplayedBytes, BuildTypes]: ... + cls: "type[HexDump[HexDumpDisplayedBytes, BuildTypes]]", + subcon: Construct[bytes, BuildTypes], + ) -> HexDump[HexDumpDisplayedBytes, BuildTypes]: ... @t.overload def __new__( - cls, subcon: Construct[RawCopyObj[SubconParsedType], BuildTypes] + cls: "type[HexDump[HexDumpDisplayedDict[str, t.Union[int, bytes, SubconParsedType]],BuildTypes,]]", + subcon: Construct[RawCopyObj[SubconParsedType], BuildTypes], ) -> HexDump[ - RawCopyObj[SubconParsedType], - BuildTypes, HexDumpDisplayedDict[str, t.Union[int, bytes, SubconParsedType]], BuildTypes, ]: ... @t.overload def __new__( - cls, subcon: Construct[Container[t.Any], BuildTypes] - ) -> HexDump[ - Container[t.Any], BuildTypes, HexDumpDisplayedDict[str, t.Any], BuildTypes - ]: ... + cls: "type[HexDump[HexDumpDisplayedDict[str, t.Any], BuildTypes]]", + subcon: Construct[Container[t.Any], BuildTypes], + ) -> HexDump[HexDumpDisplayedDict[str, t.Any], BuildTypes]: ... @t.overload def __new__( - cls, subcon: Construct[SubconParsedType, SubconBuildTypes] - ) -> HexDump[ - SubconParsedType, SubconBuildTypes, SubconParsedType, SubconBuildTypes - ]: ... + cls: "type[HexDump[SubconParsedType, SubconBuildTypes]]", + subcon: Construct[SubconParsedType, SubconBuildTypes], + ) -> HexDump[SubconParsedType, SubconBuildTypes]: ... # =============================================================================== # conditional @@ -705,49 +748,62 @@ class HexDump(Adapter[SubconParsedType, SubconBuildTypes, ParsedType, BuildTypes class Union(Construct[Container[t.Any], t.Dict[str, t.Any]]): parsefrom: t.Optional[ConstantOrContextLambda[t.Union[int, str]]] subcons: t.List[Construct[t.Any, t.Any]] + _subcons: t.Dict[str, Construct[t.Any, t.Any]] def __init__( self, parsefrom: t.Optional[ConstantOrContextLambda[t.Union[int, str]]], *subcons: Construct[t.Any, t.Any], - **subconskw: Construct[t.Any, t.Any] + **subconskw: Construct[t.Any, t.Any], ) -> None: ... def __getattr__(self, name: str) -> t.Any: ... # this can maybe made better when variadic generics are available -class Select(Construct[ParsedType, BuildTypes]): +class Select(Construct[t.Any, t.Any]): subcons: t.List[Construct[t.Any, t.Any]] - def __new__( - cls, *subcons: Construct[t.Any, t.Any], **subconskw: Construct[t.Any, t.Any] - ) -> Select[t.Any, t.Any]: ... + def __init__( + self, + *subcons: Construct[t.Any, t.Any], + **subconskw: Construct[t.Any, t.Any], + ) -> None: ... def Optional( subcon: Construct[SubconParsedType, SubconBuildTypes] -) -> Select[t.Union[SubconParsedType, None], t.Union[SubconBuildTypes, None]]: ... +) -> Construct[t.Union[SubconParsedType, None], t.Union[SubconBuildTypes, None]]: ... ThenParsedType = t.TypeVar("ThenParsedType") ThenBuildTypes = t.TypeVar("ThenBuildTypes") ElseParsedType = t.TypeVar("ElseParsedType") ElseBuildTypes = t.TypeVar("ElseBuildTypes") -# This does not represent the original code, but it is the only solution that works good with pyright -class _IfThenElse(Construct[ParsedType, BuildTypes]): +class IfThenElse(Construct[ParsedType, BuildTypes]): condfunc: ConstantOrContextLambda[bool] - thensubcon: Construct[ParsedType, BuildTypes] - elsesubcon: Construct[ParsedType, BuildTypes] + thensubcon: Construct[t.Any, t.Any] + elsesubcon: Construct[t.Any, t.Any] + @t.overload + def __new__( + cls: "type[IfThenElse[t.Union[ThenParsedType, ElseParsedType], t.Union[ThenBuildTypes, ElseBuildTypes]]]", + condfunc: ConstantOrContextLambda[bool], + thensubcon: Construct[ThenParsedType, ThenBuildTypes], + elsesubcon: Construct[ElseParsedType, ElseBuildTypes], + ) -> "IfThenElse[t.Union[ThenParsedType, ElseParsedType], t.Union[ThenBuildTypes, ElseBuildTypes]]": ... + @t.overload + def __new__( + cls: "type[IfThenElse[t.Any, t.Any]]", + condfunc: ConstantOrContextLambda[bool], + thensubcon: Construct[t.Any, t.Any], + elsesubcon: Construct[t.Any, t.Any], + ) -> "IfThenElse[t.Any, t.Any]": ... -def IfThenElse( - condfunc: ConstantOrContextLambda[bool], - thensubcon: Construct[ThenParsedType, ThenBuildTypes], - elsesubcon: Construct[ElseParsedType, ElseBuildTypes], -) -> _IfThenElse[ - t.Union[ThenParsedType, ElseParsedType], t.Union[ThenBuildTypes, ElseBuildTypes] -]: ... def If( condfunc: ConstantOrContextLambda[bool], subcon: Construct[ThenParsedType, ThenBuildTypes], -) -> _IfThenElse[t.Union[ThenParsedType, None], t.Union[ThenBuildTypes, None]]: ... +) -> IfThenElse[t.Optional[ThenParsedType], t.Optional[ThenBuildTypes]]: ... SwitchType = t.TypeVar("SwitchType") +SwitchParsedType = t.TypeVar("SwitchParsedType") +SwitchBuildTypes = t.TypeVar("SwitchBuildTypes") +SwitchDefaultParsedType = t.TypeVar("SwitchDefaultParsedType") +SwitchDefaultBuildTypes = t.TypeVar("SwitchDefaultBuildTypes") class Switch(Construct[ParsedType, BuildTypes]): keyfunc: ConstantOrContextLambda[t.Any] @@ -755,22 +811,39 @@ class Switch(Construct[ParsedType, BuildTypes]): default: Construct[t.Any, t.Any] @t.overload def __new__( - cls, + cls: "type[Switch[SwitchParsedType | None, SwitchBuildTypes | None]]", keyfunc: ConstantOrContextLambda[SwitchType], - cases: t.Dict[SwitchType, Construct[int, int]], - default: t.Optional[Construct[int, int]] = ..., - ) -> Switch[int, t.Optional[int]]: ... + cases: dict[t.Any, Construct[SwitchParsedType, SwitchBuildTypes]], + default: None = ..., + ) -> Switch[SwitchParsedType | None, SwitchBuildTypes | None]: ... @t.overload def __new__( - cls, - keyfunc: ConstantOrContextLambda[t.Any], - cases: t.Dict[t.Any, Construct[t.Any, t.Any]], - default: t.Optional[Construct[t.Any, t.Any]] = ..., + cls: "type[Switch[SwitchParsedType, SwitchBuildTypes]]", + keyfunc: ConstantOrContextLambda[SwitchType], + cases: dict[t.Any, Construct[SwitchParsedType, SwitchBuildTypes]], + default: Construct[t.NoReturn, t.NoReturn], + ) -> Switch[SwitchParsedType, SwitchBuildTypes]: ... + @t.overload + def __new__( + cls: "type[Switch[SwitchParsedType | SwitchDefaultParsedType, SwitchBuildTypes | SwitchDefaultBuildTypes]]", + keyfunc: ConstantOrContextLambda[SwitchType], + cases: dict[t.Any, Construct[SwitchParsedType, SwitchBuildTypes]], + default: Construct[SwitchDefaultParsedType, SwitchDefaultBuildTypes], + ) -> Switch[SwitchParsedType | SwitchDefaultParsedType, SwitchBuildTypes | SwitchDefaultBuildTypes]: ... + @t.overload + def __new__( + cls: "type[Switch[t.Any, t.Any]]", + keyfunc: ConstantOrContextLambda[SwitchType], + cases: dict[t.Any, Construct[t.Any, t.Any]], + default: Construct[t.Any, t.Any] | None = ..., ) -> Switch[t.Any, t.Any]: ... -class StopIf(Construct[ParsedType, BuildTypes]): +class StopIf(Construct[None, None]): condfunc: ConstantOrContextLambda[bool] - def __new__(cls, condfunc: ConstantOrContextLambda[bool]) -> StopIf[None, None]: ... + def __init__( + self, + condfunc: ConstantOrContextLambda[bool], + ) -> None: ... # =============================================================================== # alignment and padding @@ -806,8 +879,8 @@ class Aligned( def AlignedStruct( modulus: ConstantOrContextLambda[int], *subcons: Construct[t.Any, t.Any], - **subconskw: Construct[t.Any, t.Any] -) -> Struct[Container[t.Any], t.Optional[t.Dict[str, t.Any]]]: ... + **subconskw: Construct[t.Any, t.Any], +) -> Struct: ... def BitStruct( *subcons: Construct[t.Any, t.Any], **subconskw: Construct[t.Any, t.Any] ) -> t.Union[ @@ -830,16 +903,28 @@ class Pointer( stream: t.Optional[t.Callable[[Context], StreamType]] = ..., ) -> None: ... -class Peek(Subconstruct[SubconParsedType, SubconBuildTypes, ParsedType, BuildTypes]): - def __new__( - cls, - subcon: Construct[SubconParsedType, SubconBuildTypes], - ) -> Peek[ +class Peek( + Subconstruct[ SubconParsedType, SubconBuildTypes, SubconParsedType, t.Union[SubconBuildTypes, None], - ]: ... + ] +): + def __init__( + self, + subcon: Construct[SubconParsedType, SubconBuildTypes], + ) -> None: ... + +class OffsettedEnd( + Subconstruct[SubconParsedType, SubconBuildTypes, SubconParsedType, SubconBuildTypes] +): + endoffset: ConstantOrContextLambda[int] + def __init__( + self, + endoffset: ConstantOrContextLambda[int], + subcon: Construct[SubconParsedType, SubconBuildTypes], + ) -> None: ... class Seek(Construct[int, None]): at: ConstantOrContextLambda[int] @@ -869,22 +954,23 @@ class RawCopyObj(t.Generic[ParsedType], Container[t.Any]): offset2: int length: int -class RawCopy(Subconstruct[SubconParsedType, SubconBuildTypes, ParsedType, BuildTypes]): - def __new__( - cls, subcon: Construct[SubconParsedType, SubconBuildTypes] - ) -> RawCopy[ +class RawCopy( + Subconstruct[ SubconParsedType, SubconBuildTypes, RawCopyObj[SubconParsedType], t.Optional[t.Dict[str, t.Union[SubconBuildTypes, bytes]]], - ]: ... + ] +): + def __init__( + self, + subcon: Construct[SubconParsedType, SubconBuildTypes], + ) -> None: ... def ByteSwapped( subcon: Construct[SubconParsedType, SubconBuildTypes] ) -> Transformed[SubconParsedType, SubconBuildTypes]: ... -def BitsSwapped( - subcon: Construct[SubconParsedType, SubconBuildTypes] -) -> t.Union[ +def BitsSwapped(subcon: Construct[SubconParsedType, SubconBuildTypes]) -> t.Union[ Transformed[SubconParsedType, SubconBuildTypes], Restreamed[SubconParsedType, SubconBuildTypes], ]: ... @@ -907,8 +993,6 @@ def PrefixedArray( ) -> Array[ SubconParsedType, SubconBuildTypes, - ListContainer[SubconParsedType], - t.List[SubconBuildTypes], ]: ... class FixedSized( @@ -942,7 +1026,9 @@ class NullStripped( ): pad: bytes def __init__( - self, subcon: Construct[SubconParsedType, SubconBuildTypes], pad: bytes = ... + self, + subcon: Construct[SubconParsedType, SubconBuildTypes], + pad: bytes = ..., ) -> None: ... class RestreamData( @@ -994,26 +1080,26 @@ class Restreamed( ) -> None: ... class ProcessXor( - Subconstruct[SubconParsedType, SubconBuildTypes, SubconParsedType, SubconParsedType] + Subconstruct[SubconParsedType, SubconBuildTypes, SubconParsedType, SubconBuildTypes] ): padfunc: ConstantOrContextLambda2[t.Union[int, bytes]] - def __new__( - cls, + def __init__( + self, padfunc: ConstantOrContextLambda2[t.Union[int, bytes]], subcon: Construct[SubconParsedType, SubconBuildTypes], - ) -> ProcessXor[SubconParsedType, SubconBuildTypes]: ... + ) -> None: ... class ProcessRotateLeft( - Subconstruct[SubconParsedType, SubconBuildTypes, SubconParsedType, SubconParsedType] + Subconstruct[SubconParsedType, SubconBuildTypes, SubconParsedType, SubconBuildTypes] ): amount: ConstantOrContextLambda2[int] group: ConstantOrContextLambda2[int] - def __new__( - cls, + def __init__( + self, amount: ConstantOrContextLambda2[int], group: ConstantOrContextLambda2[int], subcon: Construct[SubconParsedType, SubconBuildTypes], - ) -> ProcessRotateLeft[SubconParsedType, SubconBuildTypes]: ... + ) -> None: ... T = t.TypeVar("T") @@ -1056,32 +1142,58 @@ class Rebuffered( tailcutoff: t.Optional[int] = ..., ) -> None: ... +class EncryptedSym(Tunnel[SubconParsedType, SubconBuildTypes]): + cipher: ConstantOrContextLambda2[Cipher[Mode]] + def __init__( + self, + subcon: Construct[SubconParsedType, SubconBuildTypes], + cipher: ConstantOrContextLambda2[Cipher[Mode]], + ) -> None: ... + +class EncryptedSymAead(Tunnel[SubconParsedType, SubconBuildTypes]): + cipher: ConstantOrContextLambda2[t.Union[AESGCM, AESCCM, ChaCha20Poly1305]] + nonce: ConstantOrContextLambda2[bytes] + associated_data: ConstantOrContextLambda2[bytes] + def __init__( + self, + subcon: Construct[SubconParsedType, SubconBuildTypes], + cipher: ConstantOrContextLambda2[t.Union[AESGCM, AESCCM, ChaCha20Poly1305]], + nonce: ConstantOrContextLambda2[bytes], + associated_data: ConstantOrContextLambda2[bytes] = ..., + ) -> None: ... + # =============================================================================== # lazy equivalents # =============================================================================== -class Lazy(Subconstruct[SubconParsedType, SubconBuildTypes, ParsedType, BuildTypes]): - def __new__( - cls, - subcon: Construct[SubconParsedType, SubconBuildTypes], - ) -> Lazy[ +class Lazy( + Subconstruct[ SubconParsedType, SubconBuildTypes, t.Callable[[], SubconParsedType], t.Union[t.Callable[[], SubconParsedType], SubconParsedType], - ]: ... + ] +): + def __init__( + self, + subcon: Construct[SubconParsedType, SubconBuildTypes], + ) -> None: ... class LazyContainer(t.Generic[ContainerType], t.Dict[str, ContainerType]): def __getattr__(self, name: str) -> ContainerType: ... def __getitem__(self, index: t.Union[str, int]) -> ContainerType: ... - def keys(self) -> t.Iterator[str]: ... - def values(self) -> t.List[ContainerType]: ... - def items(self) -> t.List[t.Tuple[str, ContainerType]]: ... + def keys(self) -> t.Iterator[str]: ... # type: ignore + def values(self) -> t.List[ContainerType]: ... # type: ignore + def items(self) -> t.List[t.Tuple[str, ContainerType]]: ... # type: ignore -class LazyStruct(Construct[ParsedType, BuildTypes]): +class LazyStruct(Construct[LazyContainer[t.Any], t.Optional[t.Dict[str, t.Any]]]): subcons: t.List[Construct[t.Any, t.Any]] - def __new__( - cls, *subcons: Construct[t.Any, t.Any], **subconskw: Construct[t.Any, t.Any] - ) -> LazyStruct[LazyContainer[t.Any], t.Optional[t.Dict[str, t.Any]]]: ... + _subcons: t.Dict[str, Construct[t.Any, t.Any]] + _subconsindexes: t.Dict[str, int] + def __init__( + self, + *subcons: Construct[t.Any, t.Any], + **subconskw: Construct[t.Any, t.Any], + ) -> None: ... def __getattr__(self, name: str) -> t.Any: ... class LazyListContainer(t.List[ListType]): ... @@ -1090,56 +1202,50 @@ class LazyArray( Subconstruct[ SubconParsedType, SubconBuildTypes, - ParsedType, - BuildTypes, + ListContainer[SubconParsedType], # type: ignore + t.List[SubconBuildTypes], # type: ignore ] ): count: ConstantOrContextLambda[int] - def __new__( - cls, + def __init__( + self, count: ConstantOrContextLambda[int], subcon: Construct[SubconParsedType, SubconBuildTypes], - ) -> LazyArray[ - SubconParsedType, - SubconBuildTypes, - ListContainer[SubconParsedType], - t.List[SubconBuildTypes], - ]: ... + ) -> None: ... class LazyBound(Construct[ParsedType, BuildTypes]): subconfunc: t.Callable[[], Construct[ParsedType, BuildTypes]] - def __new__( - cls, subconfunc: t.Callable[[], Construct[ParsedType, BuildTypes]] - ) -> LazyBound[ParsedType, BuildTypes]: ... + def __init__( + self, + subconfunc: t.Callable[[], Construct[ParsedType, BuildTypes]], + ) -> None: ... # =============================================================================== # adapters and validators # =============================================================================== class ExprAdapter(Adapter[SubconParsedType, SubconBuildTypes, ParsedType, BuildTypes]): - def __new__( - cls, + def __init__( + self, subcon: Construct[SubconParsedType, SubconBuildTypes], decoder: t.Callable[[SubconParsedType, Context], ParsedType], encoder: t.Callable[[BuildTypes, Context], SubconBuildTypes], - ) -> ExprAdapter[SubconParsedType, SubconBuildTypes, ParsedType, BuildTypes]: ... + ) -> None: ... class ExprSymmetricAdapter( ExprAdapter[SubconParsedType, SubconBuildTypes, ParsedType, BuildTypes] ): - def __new__( - cls, + def __init__( + self, subcon: Construct[SubconParsedType, SubconBuildTypes], encoder: t.Callable[[BuildTypes, Context], SubconBuildTypes], - ) -> ExprSymmetricAdapter[ - SubconParsedType, SubconBuildTypes, ParsedType, BuildTypes - ]: ... + ) -> None: ... class ExprValidator(Validator[SubconParsedType, SubconBuildTypes]): - def __new__( - cls, + def __init__( + self, subcon: Construct[SubconParsedType, SubconBuildTypes], validator: t.Callable[[SubconParsedType, Context], bool], - ) -> ExprValidator[SubconParsedType, SubconBuildTypes]: ... + ) -> None: ... def OneOf( subcon: Construct[SubconParsedType, SubconBuildTypes], @@ -1157,22 +1263,23 @@ def Filter( ]: ... class Slicing( - Adapter[SubconParsedType, SubconBuildTypes, SubconParsedType, SubconBuildTypes] + Adapter[ + SubconParsedType, + SubconBuildTypes, + ListContainer[SubconParsedType], # type: ignore + t.List[SubconBuildTypes], # type: ignore + ] ): - def __new__( - cls, + def __init__( + self, subcon: t.Union[ Array[ SubconParsedType, SubconBuildTypes, - ListContainer[SubconParsedType], - t.List[SubconBuildTypes], ], GreedyRange[ SubconParsedType, SubconBuildTypes, - ListContainer[SubconParsedType], - t.List[SubconBuildTypes], ], ], count: int, @@ -1180,28 +1287,24 @@ class Slicing( stop: t.Optional[int], step: int = ..., empty: t.Optional[SubconParsedType] = ..., - ) -> Slicing[ListContainer[SubconParsedType], t.List[SubconBuildTypes]]: ... + ) -> None: ... class Indexing( Adapter[SubconParsedType, SubconBuildTypes, SubconParsedType, SubconBuildTypes] ): - def __new__( - cls, + def __init__( + self, subcon: t.Union[ Array[ SubconParsedType, SubconBuildTypes, - ListContainer[SubconParsedType], - t.List[SubconBuildTypes], ], GreedyRange[ SubconParsedType, SubconBuildTypes, - ListContainer[SubconParsedType], - t.List[SubconBuildTypes], ], ], count: int, index: int, empty: t.Optional[SubconParsedType] = ..., - ) -> Indexing[SubconParsedType, SubconBuildTypes]: ... + ) -> None: ... diff --git a/construct-stubs/expr.pyi b/construct-stubs/expr.pyi index 8450a44..a7c1a1a 100644 --- a/construct-stubs/expr.pyi +++ b/construct-stubs/expr.pyi @@ -1,4 +1,3 @@ -import operator import typing as t from construct.core import * @@ -470,7 +469,7 @@ class ExprMixin(t.Generic[ReturnType], object): @t.overload def __eq__(self: ExprMixin[float], other: ConstOrCallable[float]) -> BinExpr[bool]: ... @t.overload - def __eq__(self, other: t.Any) -> BinExpr[t.Any]: ... + def __eq__(self, other: ConstOrCallable[t.Any]) -> BinExpr[t.Any]: ... # type: ignore # __ne__ ########################################################################################################### @t.overload @@ -488,7 +487,7 @@ class ExprMixin(t.Generic[ReturnType], object): @t.overload def __ne__(self: ExprMixin[float], other: ConstOrCallable[float]) -> BinExpr[bool]: ... @t.overload - def __ne__(self, other: t.Any) -> BinExpr[t.Any]: ... + def __ne__(self, other: t.Any) -> BinExpr[t.Any]: ... # type: ignore # __neg__ ########################################################################################################## @t.overload @@ -498,7 +497,7 @@ class ExprMixin(t.Generic[ReturnType], object): @t.overload def __neg__(self: ExprMixin[float]) -> BinExpr[float]: ... @t.overload - def __neg__(self) -> UniExpr[t.Any]: ... + def __neg__(self) -> BinExpr[t.Any]: ... # __pos__ ########################################################################################################## @t.overload @@ -508,7 +507,7 @@ class ExprMixin(t.Generic[ReturnType], object): @t.overload def __pos__(self: ExprMixin[float]) -> BinExpr[float]: ... @t.overload - def __pos__(self) -> UniExpr[t.Any]: ... + def __pos__(self) -> BinExpr[t.Any]: ... # __invert__ ####################################################################################################### @t.overload @@ -516,7 +515,7 @@ class ExprMixin(t.Generic[ReturnType], object): @t.overload def __invert__(self: ExprMixin[bool]) -> BinExpr[int]: ... @t.overload - def __invert__(self) -> UniExpr[t.Any]: ... + def __invert__(self) -> BinExpr[t.Any]: ... # __inv__ ########################################################################################################## def __inv__(self) -> UniExpr[t.Any]: ... @@ -543,7 +542,7 @@ class Path2(ExprMixin[ReturnType]): class FuncPath(ExprMixin[ReturnType]): - def __init__(self, func: t.Callable[[t.Any], t.Any], operand: t.Optional[t.Any] = ...) -> None: ... + def __init__(self, func: t.Callable[[t.Any], ReturnType], operand: t.Optional[t.Any] = ...) -> None: ... def __call__(self, operand: t.Any, *args: t.Any) -> ReturnType: ... diff --git a/construct-stubs/lib/containers.pyi b/construct-stubs/lib/containers.pyi index a50033a..37efc75 100644 --- a/construct-stubs/lib/containers.pyi +++ b/construct-stubs/lib/containers.pyi @@ -19,7 +19,7 @@ def recursion_lock( class Container(t.Generic[ContainerType], t.Dict[str, ContainerType]): def __getattr__(self, name: str) -> ContainerType: ... - def update( + def update( # type: ignore self, seqordict: t.Union[t.Dict[str, ContainerType], t.Tuple[str, ContainerType]], ) -> None: ... diff --git a/construct-stubs/lib/hex.pyi b/construct-stubs/lib/hex.pyi index afa985f..a39d918 100644 --- a/construct-stubs/lib/hex.pyi +++ b/construct-stubs/lib/hex.pyi @@ -1,6 +1,5 @@ import typing as t - class HexDisplayedInteger(int): ... class HexDisplayedBytes(bytes): ... @@ -10,3 +9,6 @@ V = t.TypeVar("V") class HexDisplayedDict(t.Dict[K, V]): ... class HexDumpDisplayedBytes(bytes): ... class HexDumpDisplayedDict(t.Dict[K, V]): ... + +def hexdump(data: bytes, linesize: int) -> str: ... +def hexundump(data: str, linesize: int) -> bytes: ... diff --git a/construct-stubs/lib/py3compat.pyi b/construct-stubs/lib/py3compat.pyi index f105096..c86f2f5 100644 --- a/construct-stubs/lib/py3compat.pyi +++ b/construct-stubs/lib/py3compat.pyi @@ -1,5 +1,6 @@ import typing as t +PY: t.Tuple[int, int] PY2: bool PY3: bool PYPY: bool diff --git a/construct_typed/__init__.py b/construct_typed/__init__.py index 537de1a..f594f2b 100644 --- a/construct_typed/__init__.py +++ b/construct_typed/__init__.py @@ -1,32 +1,55 @@ -from construct_typed.generic import constr -from construct_typed.dataclass_struct import ( +from .dataclass_struct import ( DataclassBitStruct, + DataclassMixin, DataclassStruct, - csfield + TBitStruct, + TContainerBase, + TContainerMixin, + TStruct, + TStructField, + csfield, + sfield, + EnhancedDataclassMixin ) -from construct_typed.generic import ( +from .generic_wrapper import ( Adapter, ConstantOrContextLambda, + ConstantOrContextLambda2, Construct, Context, ListContainer, PathType, + Array, + Subconstruct, + Computed, ) -from construct_typed.tenum import TEnum, TFlags, TEnumConstruct, TFlagsConstruct +from .tenum import EnumBase, EnumValue, FlagsEnumBase, TEnum, TFlagsEnum __all__ = [ "DataclassBitStruct", + "DataclassMixin", "DataclassStruct", - "constr", + "TBitStruct", + "TContainerBase", + "TContainerMixin", + "TStruct", + "TStructField", "csfield", + "sfield", + "EnhancedDataclassMixin", + "EnumBase", + "EnumValue", + "FlagsEnumBase", "TEnum", - "TEnumConstruct", - "TFlags", - "TFlagsConstruct", + "TFlagsEnum", "Adapter", "ConstantOrContextLambda", + "ConstantOrContextLambda2", "Construct", "Context", "ListContainer", "PathType", + "Array", + "Subconstruct", + "Computed" ] diff --git a/construct_typed/dataclass_struct.py b/construct_typed/dataclass_struct.py index 2e23dbc..e626085 100644 --- a/construct_typed/dataclass_struct.py +++ b/construct_typed/dataclass_struct.py @@ -1,10 +1,9 @@ # -*- coding: utf-8 -*- # pyright: strict +# pyright: reportIncompatibleVariableOverride=false, reportAny=false import dataclasses -import sys import textwrap import typing as t -import enum import construct as cs from construct.lib.containers import ( @@ -13,291 +12,25 @@ from construct.lib.containers import ( recursion_lock, ) from construct.lib.py3compat import bytestringtype, reprstring, unicodestringtype +from typing_extensions import override -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") +from .generic_wrapper import Adapter, Construct, Context, ParsedType, PathType -# 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]: +class DataclassMixin: """ - Helper method for "DataclassStruct" and "DataclassBitStruct" to create the dataclass fields. + Mixin for the dataclasses which are passed to "DataclassStruct" and "DataclassBitStruct". - This method also processes Const and Default, to pass these values als default values to the dataclass. - - Only one of the parameters `default` or `const` can be vaild. They are mutually exclusive. + Note: This implementation is different to the 'cs.Container' of the original 'construct' + library. In the original 'cs.Container' some names like "update", "keys", "items", ... can + only accessed via key access (square brackets) and not via attribute access (dot operator), + because they are also method names. This implementation is based on "dataclasses.dataclass" + which only uses modul-level instead of instance-level helper methods.So no instance-level + methods exists and every name can be used. """ - if (default is not Flag.MISSING) and (const is not Flag.MISSING): - raise ValueError("default and const are mutally exclusive") + __dataclass_fields__: "t.ClassVar[dict[str, dataclasses.Field[t.Any]]]" - # 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: return getattr(self, key) @@ -343,43 +76,209 @@ class DataclassStruct: text.append(indentation.join(str(v).split("\n"))) return "".join(text) - if t.TYPE_CHECKING: - @classmethod - def __constr__(cls: t.Type[T]) -> "DataclassConstruct[T]": - ... +def csfield( + subcon: Construct[ParsedType, t.Any], + 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}, + ), + ) -class DataclassBitStruct(DataclassStruct): +DataclassType = t.TypeVar("DataclassType", bound=DataclassMixin) + + +class DataclassStruct(Adapter[t.Any, t.Any, DataclassType, DataclassType]): + """ + Adapter for a dataclasses for optimised type hints / static autocompletion in comparision to the original Struct. + + Before this construct can be created a dataclasses.dataclass type must be created, which must also derive from DataclassMixin. In this dataclass all fields must be assigned to a construct type using csfield. + + Internally, all fields are converted to a Struct, which does the actual parsing/building. + + Parses to a dataclasses.dataclass instance, and builds from such instance. Size is the sum of all subcon sizes, unless any subcon raises SizeofError. + + :param dc_type: Type of the dataclass, which also inherits from DataclassMixin + :param reverse: Flag if the fields of the dataclass should be reversed + + Example:: + + >>> import dataclasses + >>> from construct import Bytes, Int8ub, this + >>> from construct_typed import DataclassMixin, DataclassStruct, csfield + >>> @dataclasses.dataclass + ... class Image(DataclassMixin): + ... width: int = csfield(Int8ub) + ... height: int = csfield(Int8ub) + ... pixels: bytes = csfield(Bytes(this.height * this.width)) + >>> d = DataclassStruct(Image) + >>> d.parse(b"\x01\x0212") + Image(width=1, height=2, pixels=b'12') + """ + + subcon: "cs.Struct" # 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""" Makes a DataclassStruct inside a Bitwise. See :class:`~construct.core.Bitwise` and :class:`~construct_typed.dataclass_struct.DatclassStruct` for semantics and raisable exceptions. - :param constr: TODO - :param reverse_fields: Flag if the fields of the dataclass should be reversed + :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:: - TODO: + DataclassBitStruct <--> Bitwise(DataclassStruct(...)) + >>> import dataclasses >>> from construct import BitsInteger, Flag, Nibble, Padding - >>> from construct_typed import DataclassBitStruct, csfield, construct - ... class TestDataclass(DataclassBitStruct): + >>> from construct_typed import DataclassBitStruct, DataclassMixin, csfield + >>> @dataclasses.dataclass + ... class TestDataclass(DataclassMixin): ... a: int = csfield(Flag) ... b: int = csfield(Nibble) ... c: int = csfield(BitsInteger(10)) ... d: None = csfield(Padding(1)) - >>> d = construct(TestDataclass) + >>> d = DataclassBitStruct(TestDataclass) >>> d.parse(b"\x01\x02") 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 - 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: - DataclassStruct.__init_subclass__.__func__(cls, lambda cls: cs.Bitwise(constr(cls)), reverse_fields) # type: ignore + def build(cls, obj: t.Self, **kw: dict[str, t.Any]): + return cls.format().build(obj, **kw) + + @classmethod + def parse(cls, data: bytes | bytearray, **kw: dict[str, t.Any]): + return cls.format().parse(data, **kw) + + @classmethod + 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 diff --git a/construct_typed/dataclasses_py310.py b/construct_typed/dataclasses_py310.py deleted file mode 100644 index a2f76d0..0000000 --- a/construct_typed/dataclasses_py310.py +++ /dev/null @@ -1,1462 +0,0 @@ -# type: ignore -import re -import sys -import copy -import types -import inspect -import keyword -import builtins -import functools -import _thread -from types import FunctionType -from dataclasses import ( - _FIELD, - _FIELD_CLASSVAR, - _FIELD_INITVAR, - _FIELDS, - _PARAMS, - _POST_INIT_NAME, - _EMPTY_METADATA, - MISSING, - _HAS_DEFAULT_FACTORY, - FrozenInstanceError -) - - -__all__ = ['dataclass', - 'field', - 'Field', - 'FrozenInstanceError', - 'InitVar', - 'KW_ONLY', - 'MISSING', - - # Helper functions. - 'fields', - 'asdict', - 'astuple', - 'make_dataclass', - 'replace', - 'is_dataclass', - ] - -# Conditions for adding methods. The boxes indicate what action the -# dataclass decorator takes. For all of these tables, when I talk -# about init=, repr=, eq=, order=, unsafe_hash=, or frozen=, I'm -# referring to the arguments to the @dataclass decorator. When -# checking if a dunder method already exists, I mean check for an -# entry in the class's __dict__. I never check to see if an attribute -# is defined in a base class. - -# Key: -# +=========+=========================================+ -# + Value | Meaning | -# +=========+=========================================+ -# | | No action: no method is added. | -# +---------+-----------------------------------------+ -# | add | Generated method is added. | -# +---------+-----------------------------------------+ -# | raise | TypeError is raised. | -# +---------+-----------------------------------------+ -# | None | Attribute is set to None. | -# +=========+=========================================+ - -# __init__ -# -# +--- init= parameter -# | -# v | | | -# | no | yes | <--- class has __init__ in __dict__? -# +=======+=======+=======+ -# | False | | | -# +-------+-------+-------+ -# | True | add | | <- the default -# +=======+=======+=======+ - -# __repr__ -# -# +--- repr= parameter -# | -# v | | | -# | no | yes | <--- class has __repr__ in __dict__? -# +=======+=======+=======+ -# | False | | | -# +-------+-------+-------+ -# | True | add | | <- the default -# +=======+=======+=======+ - - -# __setattr__ -# __delattr__ -# -# +--- frozen= parameter -# | -# v | | | -# | no | yes | <--- class has __setattr__ or __delattr__ in __dict__? -# +=======+=======+=======+ -# | False | | | <- the default -# +-------+-------+-------+ -# | True | add | raise | -# +=======+=======+=======+ -# Raise because not adding these methods would break the "frozen-ness" -# of the class. - -# __eq__ -# -# +--- eq= parameter -# | -# v | | | -# | no | yes | <--- class has __eq__ in __dict__? -# +=======+=======+=======+ -# | False | | | -# +-------+-------+-------+ -# | True | add | | <- the default -# +=======+=======+=======+ - -# __lt__ -# __le__ -# __gt__ -# __ge__ -# -# +--- order= parameter -# | -# v | | | -# | no | yes | <--- class has any comparison method in __dict__? -# +=======+=======+=======+ -# | False | | | <- the default -# +-------+-------+-------+ -# | True | add | raise | -# +=======+=======+=======+ -# Raise because to allow this case would interfere with using -# functools.total_ordering. - -# __hash__ - -# +------------------- unsafe_hash= parameter -# | +----------- eq= parameter -# | | +--- frozen= parameter -# | | | -# v v v | | | -# | no | yes | <--- class has explicitly defined __hash__ -# +=======+=======+=======+========+========+ -# | False | False | False | | | No __eq__, use the base class __hash__ -# +-------+-------+-------+--------+--------+ -# | False | False | True | | | No __eq__, use the base class __hash__ -# +-------+-------+-------+--------+--------+ -# | False | True | False | None | | <-- the default, not hashable -# +-------+-------+-------+--------+--------+ -# | False | True | True | add | | Frozen, so hashable, allows override -# +-------+-------+-------+--------+--------+ -# | True | False | False | add | raise | Has no __eq__, but hashable -# +-------+-------+-------+--------+--------+ -# | True | False | True | add | raise | Has no __eq__, but hashable -# +-------+-------+-------+--------+--------+ -# | True | True | False | add | raise | Not frozen, but hashable -# +-------+-------+-------+--------+--------+ -# | True | True | True | add | raise | Frozen, so hashable -# +=======+=======+=======+========+========+ -# For boxes that are blank, __hash__ is untouched and therefore -# inherited from the base class. If the base is object, then -# id-based hashing is used. -# -# Note that a class may already have __hash__=None if it specified an -# __eq__ method in the class body (not one that was created by -# @dataclass). -# -# See _hash_action (below) for a coded version of this table. - -# __match_args__ -# -# +--- match_args= parameter -# | -# v | | | -# | no | yes | <--- class has __match_args__ in __dict__? -# +=======+=======+=======+ -# | False | | | -# +-------+-------+-------+ -# | True | add | | <- the default -# +=======+=======+=======+ -# __match_args__ is always added unless the class already defines it. It is a -# tuple of __init__ parameter names; non-init fields must be matched by keyword. - - -# Raised when an attempt is made to modify a frozen class. -# class FrozenInstanceError(AttributeError): pass - -# A sentinel object for default values to signal that a default -# factory will be used. This is given a nice repr() which will appear -# in the function signature of dataclasses' constructors. -# class _HAS_DEFAULT_FACTORY_CLASS: -# def __repr__(self): -# return '' -# _HAS_DEFAULT_FACTORY = _HAS_DEFAULT_FACTORY_CLASS() - -# A sentinel object to detect if a parameter is supplied or not. Use -# a class to give it a better repr. -# class _MISSING_TYPE: -# pass -# MISSING = _MISSING_TYPE() - -# A sentinel object to indicate that following fields are keyword-only by -# default. Use a class to give it a better repr. -class _KW_ONLY_TYPE: - pass -KW_ONLY = _KW_ONLY_TYPE() - -# Since most per-field metadata will be unused, create an empty -# read-only proxy that can be shared among all fields. -# _EMPTY_METADATA = types.MappingProxyType({}) - -# Markers for the various kinds of fields and pseudo-fields. -# class _FIELD_BASE: -# def __init__(self, name): -# self.name = name -# def __repr__(self): -# return self.name -# _FIELD = _FIELD_BASE('_FIELD') -# _FIELD_CLASSVAR = _FIELD_BASE('_FIELD_CLASSVAR') -# _FIELD_INITVAR = _FIELD_BASE('_FIELD_INITVAR') - -# The name of an attribute on the class where we store the Field -# objects. Also used to check if a class is a Data Class. -# _FIELDS = '__dataclass_fields__' - -# The name of an attribute on the class that stores the parameters to -# @dataclass. -# _PARAMS = '__dataclass_params__' - -# The name of the function, that if it exists, is called at the end of -# __init__. -# _POST_INIT_NAME = '__post_init__' - -# String regex that string annotations for ClassVar or InitVar must match. -# Allows "identifier.identifier[" or "identifier[". -# https://bugs.python.org/issue33453 for details. -_MODULE_IDENTIFIER_RE = re.compile(r'^(?:\s*(\w+)\s*\.)?\s*(\w+)') - -class InitVar: - __slots__ = ('type', ) - - def __init__(self, type): - self.type = type - - def __repr__(self): - if isinstance(self.type, type): - type_name = self.type.__name__ - else: - # typing objects, e.g. List[int] - type_name = repr(self.type) - return f'dataclasses.InitVar[{type_name}]' - - def __class_getitem__(cls, type): - return InitVar(type) - -# Instances of Field are only ever created from within this module, -# and only from the field() function, although Field instances are -# exposed externally as (conceptually) read-only objects. -# -# name and type are filled in after the fact, not in __init__. -# They're not known at the time this class is instantiated, but it's -# convenient if they're available later. -# -# When cls._FIELDS is filled in with a list of Field objects, the name -# and type fields will have been populated. -class Field: - __slots__ = ('name', - 'type', - 'default', - 'default_factory', - 'repr', - 'hash', - 'init', - 'compare', - 'metadata', - 'kw_only', - '_field_type', # Private: not to be used by user code. - ) - - def __init__(self, default, default_factory, init, repr, hash, compare, - metadata, kw_only): - self.name = None - self.type = None - self.default = default - self.default_factory = default_factory - self.init = init - self.repr = repr - self.hash = hash - self.compare = compare - self.metadata = (_EMPTY_METADATA - if metadata is None else - types.MappingProxyType(metadata)) - self.kw_only = kw_only - self._field_type = None - - def __repr__(self): - return ('Field(' - f'name={self.name!r},' - f'type={self.type!r},' - f'default={self.default!r},' - f'default_factory={self.default_factory!r},' - f'init={self.init!r},' - f'repr={self.repr!r},' - f'hash={self.hash!r},' - f'compare={self.compare!r},' - f'metadata={self.metadata!r},' - f'kw_only={self.kw_only!r},' - f'_field_type={self._field_type}' - ')') - - # This is used to support the PEP 487 __set_name__ protocol in the - # case where we're using a field that contains a descriptor as a - # default value. For details on __set_name__, see - # https://www.python.org/dev/peps/pep-0487/#implementation-details. - # - # Note that in _process_class, this Field object is overwritten - # with the default value, so the end result is a descriptor that - # had __set_name__ called on it at the right time. - def __set_name__(self, owner, name): - func = getattr(type(self.default), '__set_name__', None) - if func: - # There is a __set_name__ method on the descriptor, call - # it. - func(self.default, owner, name) - - -class _DataclassParams: - __slots__ = ('init', - 'repr', - 'eq', - 'order', - 'unsafe_hash', - 'frozen', - ) - - def __init__(self, init, repr, eq, order, unsafe_hash, frozen): - self.init = init - self.repr = repr - self.eq = eq - self.order = order - self.unsafe_hash = unsafe_hash - self.frozen = frozen - - def __repr__(self): - return ('_DataclassParams(' - f'init={self.init!r},' - f'repr={self.repr!r},' - f'eq={self.eq!r},' - f'order={self.order!r},' - f'unsafe_hash={self.unsafe_hash!r},' - f'frozen={self.frozen!r}' - ')') - - -# This function is used instead of exposing Field creation directly, -# so that a type checker can be told (via overloads) that this is a -# function whose type depends on its parameters. -def field(*, default=MISSING, default_factory=MISSING, init=True, repr=True, - hash=None, compare=True, metadata=None, kw_only=MISSING): - """Return an object to identify dataclass fields. - - default is the default value of the field. default_factory is a - 0-argument function called to initialize a field's value. If init - is true, the field will be a parameter to the class's __init__() - function. If repr is true, the field will be included in the - object's repr(). If hash is true, the field will be included in the - object's hash(). If compare is true, the field will be used in - comparison functions. metadata, if specified, must be a mapping - which is stored but not otherwise examined by dataclass. If kw_only - is true, the field will become a keyword-only parameter to - __init__(). - - It is an error to specify both default and default_factory. - """ - - if default is not MISSING and default_factory is not MISSING: - raise ValueError('cannot specify both default and default_factory') - return Field(default, default_factory, init, repr, hash, compare, - metadata, kw_only) - - -def _fields_in_init_order(fields): - # Returns the fields as __init__ will output them. It returns 2 tuples: - # the first for normal args, and the second for keyword args. - - return (tuple(f for f in fields if f.init and not f.kw_only), - tuple(f for f in fields if f.init and f.kw_only) - ) - - -def _tuple_str(obj_name, fields): - # Return a string representing each field of obj_name as a tuple - # member. So, if fields is ['x', 'y'] and obj_name is "self", - # return "(self.x,self.y)". - - # Special case for the 0-tuple. - if not fields: - return '()' - # Note the trailing comma, needed if this turns out to be a 1-tuple. - return f'({",".join([f"{obj_name}.{f.name}" for f in fields])},)' - - -# This function's logic is copied from "recursive_repr" function in -# reprlib module to avoid dependency. -def _recursive_repr(user_function): - # Decorator to make a repr function return "..." for a recursive - # call. - repr_running = set() - - @functools.wraps(user_function) - def wrapper(self): - key = id(self), _thread.get_ident() - if key in repr_running: - return '...' - repr_running.add(key) - try: - result = user_function(self) - finally: - repr_running.discard(key) - return result - return wrapper - - -def _create_fn(name, args, body, *, globals=None, locals=None, - return_type=MISSING): - # Note that we mutate locals when exec() is called. Caller - # beware! The only callers are internal to this module, so no - # worries about external callers. - if locals is None: - locals = {} - if 'BUILTINS' not in locals: - locals['BUILTINS'] = builtins - return_annotation = '' - if return_type is not MISSING: - locals['_return_type'] = return_type - return_annotation = '->_return_type' - args = ','.join(args) - body = '\n'.join(f' {b}' for b in body) - - # Compute the text of the entire function. - txt = f' def {name}({args}){return_annotation}:\n{body}' - - local_vars = ', '.join(locals.keys()) - txt = f"def __create_fn__({local_vars}):\n{txt}\n return {name}" - ns = {} - exec(txt, globals, ns) - return ns['__create_fn__'](**locals) - - -def _field_assign(frozen, name, value, self_name): - # If we're a frozen class, then assign to our fields in __init__ - # via object.__setattr__. Otherwise, just use a simple - # assignment. - # - # self_name is what "self" is called in this function: don't - # hard-code "self", since that might be a field name. - if frozen: - return f'BUILTINS.object.__setattr__({self_name},{name!r},{value})' - return f'{self_name}.{name}={value}' - - -def _field_init(f, frozen, globals, self_name, slots): - # Return the text of the line in the body of __init__ that will - # initialize this field. - - default_name = f'_dflt_{f.name}' - if f.default_factory is not MISSING: - if f.init: - # This field has a default factory. If a parameter is - # given, use it. If not, call the factory. - globals[default_name] = f.default_factory - value = (f'{default_name}() ' - f'if {f.name} is _HAS_DEFAULT_FACTORY ' - f'else {f.name}') - else: - # This is a field that's not in the __init__ params, but - # has a default factory function. It needs to be - # initialized here by calling the factory function, - # because there's no other way to initialize it. - - # For a field initialized with a default=defaultvalue, the - # class dict just has the default value - # (cls.fieldname=defaultvalue). But that won't work for a - # default factory, the factory must be called in __init__ - # and we must assign that to self.fieldname. We can't - # fall back to the class dict's value, both because it's - # not set, and because it might be different per-class - # (which, after all, is why we have a factory function!). - - globals[default_name] = f.default_factory - value = f'{default_name}()' - else: - # No default factory. - if f.init: - if f.default is MISSING: - # There's no default, just do an assignment. - value = f.name - elif f.default is not MISSING: - globals[default_name] = f.default - value = f.name - else: - # If the class has slots, then initialize this field. - if slots and f.default is not MISSING: - globals[default_name] = f.default - value = default_name - else: - # This field does not need initialization: reading from it will - # just use the class attribute that contains the default. - # Signify that to the caller by returning None. - return None - - # Only test this now, so that we can create variables for the - # default. However, return None to signify that we're not going - # to actually do the assignment statement for InitVars. - if f._field_type is _FIELD_INITVAR: - return None - - # Now, actually generate the field assignment. - return _field_assign(frozen, f.name, value, self_name) - - -def _init_param(f): - # Return the __init__ parameter string for this field. For - # example, the equivalent of 'x:int=3' (except instead of 'int', - # reference a variable set to int, and instead of '3', reference a - # variable set to 3). - if f.default is MISSING and f.default_factory is MISSING: - # There's no default, and no default_factory, just output the - # variable name and type. - default = '' - elif f.default is not MISSING: - # There's a default, this will be the name that's used to look - # it up. - default = f'=_dflt_{f.name}' - elif f.default_factory is not MISSING: - # There's a factory function. Set a marker. - default = '=_HAS_DEFAULT_FACTORY' - return f'{f.name}:_type_{f.name}{default}' - - -def _init_fn(fields, std_fields, kw_only_fields, frozen, has_post_init, - self_name, globals, slots): - # fields contains both real fields and InitVar pseudo-fields. - - # Make sure we don't have fields without defaults following fields - # with defaults. This actually would be caught when exec-ing the - # function source code, but catching it here gives a better error - # message, and future-proofs us in case we build up the function - # using ast. - - seen_default = False - for f in std_fields: - # Only consider the non-kw-only fields in the __init__ call. - if f.init: - if not (f.default is MISSING and f.default_factory is MISSING): - seen_default = True - elif seen_default: - raise TypeError(f'non-default argument {f.name!r} ' - 'follows default argument') - - locals = {f'_type_{f.name}': f.type for f in fields} - locals.update({ - 'MISSING': MISSING, - '_HAS_DEFAULT_FACTORY': _HAS_DEFAULT_FACTORY, - }) - - body_lines = [] - for f in fields: - line = _field_init(f, frozen, locals, self_name, slots) - # line is None means that this field doesn't require - # initialization (it's a pseudo-field). Just skip it. - if line: - body_lines.append(line) - - # Does this class have a post-init function? - if has_post_init: - params_str = ','.join(f.name for f in fields - if f._field_type is _FIELD_INITVAR) - body_lines.append(f'{self_name}.{_POST_INIT_NAME}({params_str})') - - # If no body lines, use 'pass'. - if not body_lines: - body_lines = ['pass'] - - _init_params = [_init_param(f) for f in std_fields] - if kw_only_fields: - # Add the keyword-only args. Because the * can only be added if - # there's at least one keyword-only arg, there needs to be a test here - # (instead of just concatenting the lists together). - _init_params += ['*'] - _init_params += [_init_param(f) for f in kw_only_fields] - return _create_fn('__init__', - [self_name] + _init_params, - body_lines, - locals=locals, - globals=globals, - return_type=None) - - -def _repr_fn(fields, globals): - fn = _create_fn('__repr__', - ('self',), - ['return self.__class__.__qualname__ + f"(' + - ', '.join([f"{f.name}={{self.{f.name}!r}}" - for f in fields]) + - ')"'], - globals=globals) - return _recursive_repr(fn) - - -def _frozen_get_del_attr(cls, fields, globals): - locals = {'cls': cls, - 'FrozenInstanceError': FrozenInstanceError} - if fields: - fields_str = '(' + ','.join(repr(f.name) for f in fields) + ',)' - else: - # Special case for the zero-length tuple. - fields_str = '()' - return (_create_fn('__setattr__', - ('self', 'name', 'value'), - (f'if type(self) is cls or name in {fields_str}:', - ' raise FrozenInstanceError(f"cannot assign to field {name!r}")', - f'super(cls, self).__setattr__(name, value)'), - locals=locals, - globals=globals), - _create_fn('__delattr__', - ('self', 'name'), - (f'if type(self) is cls or name in {fields_str}:', - ' raise FrozenInstanceError(f"cannot delete field {name!r}")', - f'super(cls, self).__delattr__(name)'), - locals=locals, - globals=globals), - ) - - -def _cmp_fn(name, op, self_tuple, other_tuple, globals): - # Create a comparison function. If the fields in the object are - # named 'x' and 'y', then self_tuple is the string - # '(self.x,self.y)' and other_tuple is the string - # '(other.x,other.y)'. - - return _create_fn(name, - ('self', 'other'), - [ 'if other.__class__ is self.__class__:', - f' return {self_tuple}{op}{other_tuple}', - 'return NotImplemented'], - globals=globals) - - -def _hash_fn(fields, globals): - self_tuple = _tuple_str('self', fields) - return _create_fn('__hash__', - ('self',), - [f'return hash({self_tuple})'], - globals=globals) - - -def _is_classvar(a_type, typing): - # This test uses a typing internal class, but it's the best way to - # test if this is a ClassVar. - return (a_type is typing.ClassVar - or (type(a_type) is typing._GenericAlias - and a_type.__origin__ is typing.ClassVar)) - - -def _is_initvar(a_type, dataclasses): - # The module we're checking against is the module we're - # currently in (dataclasses.py). - return (a_type is dataclasses.InitVar - or type(a_type) is dataclasses.InitVar) - -def _is_kw_only(a_type, dataclasses): - return a_type is dataclasses.KW_ONLY - - -def _is_type(annotation, cls, a_module, a_type, is_type_predicate): - # Given a type annotation string, does it refer to a_type in - # a_module? For example, when checking that annotation denotes a - # ClassVar, then a_module is typing, and a_type is - # typing.ClassVar. - - # It's possible to look up a_module given a_type, but it involves - # looking in sys.modules (again!), and seems like a waste since - # the caller already knows a_module. - - # - annotation is a string type annotation - # - cls is the class that this annotation was found in - # - a_module is the module we want to match - # - a_type is the type in that module we want to match - # - is_type_predicate is a function called with (obj, a_module) - # that determines if obj is of the desired type. - - # Since this test does not do a local namespace lookup (and - # instead only a module (global) lookup), there are some things it - # gets wrong. - - # With string annotations, cv0 will be detected as a ClassVar: - # CV = ClassVar - # @dataclass - # class C0: - # cv0: CV - - # But in this example cv1 will not be detected as a ClassVar: - # @dataclass - # class C1: - # CV = ClassVar - # cv1: CV - - # In C1, the code in this function (_is_type) will look up "CV" in - # the module and not find it, so it will not consider cv1 as a - # ClassVar. This is a fairly obscure corner case, and the best - # way to fix it would be to eval() the string "CV" with the - # correct global and local namespaces. However that would involve - # a eval() penalty for every single field of every dataclass - # that's defined. It was judged not worth it. - - match = _MODULE_IDENTIFIER_RE.match(annotation) - if match: - ns = None - module_name = match.group(1) - if not module_name: - # No module name, assume the class's module did - # "from dataclasses import InitVar". - ns = sys.modules.get(cls.__module__).__dict__ - else: - # Look up module_name in the class's module. - module = sys.modules.get(cls.__module__) - if module and module.__dict__.get(module_name) is a_module: - ns = sys.modules.get(a_type.__module__).__dict__ - if ns and is_type_predicate(ns.get(match.group(2)), a_module): - return True - return False - - -def _get_field(cls, a_name, a_type, default_kw_only): - # Return a Field object for this field name and type. ClassVars and - # InitVars are also returned, but marked as such (see f._field_type). - # default_kw_only is the value of kw_only to use if there isn't a field() - # that defines it. - - # If the default value isn't derived from Field, then it's only a - # normal default value. Convert it to a Field(). - default = getattr(cls, a_name, MISSING) - if isinstance(default, Field): - f = default - else: - if isinstance(default, types.MemberDescriptorType): - # This is a field in __slots__, so it has no default value. - default = MISSING - f = field(default=default) - - # Only at this point do we know the name and the type. Set them. - f.name = a_name - f.type = a_type - - # Assume it's a normal field until proven otherwise. We're next - # going to decide if it's a ClassVar or InitVar, everything else - # is just a normal field. - f._field_type = _FIELD - - # In addition to checking for actual types here, also check for - # string annotations. get_type_hints() won't always work for us - # (see https://github.com/python/typing/issues/508 for example), - # plus it's expensive and would require an eval for every string - # annotation. So, make a best effort to see if this is a ClassVar - # or InitVar using regex's and checking that the thing referenced - # is actually of the correct type. - - # For the complete discussion, see https://bugs.python.org/issue33453 - - # If typing has not been imported, then it's impossible for any - # annotation to be a ClassVar. So, only look for ClassVar if - # typing has been imported by any module (not necessarily cls's - # module). - typing = sys.modules.get('typing') - if typing: - if (_is_classvar(a_type, typing) - or (isinstance(f.type, str) - and _is_type(f.type, cls, typing, typing.ClassVar, - _is_classvar))): - f._field_type = _FIELD_CLASSVAR - - # If the type is InitVar, or if it's a matching string annotation, - # then it's an InitVar. - if f._field_type is _FIELD: - # The module we're checking against is the module we're - # currently in (dataclasses.py). - dataclasses = sys.modules[__name__] - if (_is_initvar(a_type, dataclasses) - or (isinstance(f.type, str) - and _is_type(f.type, cls, dataclasses, dataclasses.InitVar, - _is_initvar))): - f._field_type = _FIELD_INITVAR - - # Validations for individual fields. This is delayed until now, - # instead of in the Field() constructor, since only here do we - # know the field name, which allows for better error reporting. - - # Special restrictions for ClassVar and InitVar. - if f._field_type in (_FIELD_CLASSVAR, _FIELD_INITVAR): - if f.default_factory is not MISSING: - raise TypeError(f'field {f.name} cannot have a ' - 'default factory') - # Should I check for other field settings? default_factory - # seems the most serious to check for. Maybe add others. For - # example, how about init=False (or really, - # init=)? It makes no sense for - # ClassVar and InitVar to specify init=. - - # kw_only validation and assignment. - if f._field_type in (_FIELD, _FIELD_INITVAR): - # For real and InitVar fields, if kw_only wasn't specified use the - # default value. - if f.kw_only is MISSING: - f.kw_only = default_kw_only - else: - # Make sure kw_only isn't set for ClassVars - assert f._field_type is _FIELD_CLASSVAR - if f.kw_only is not MISSING: - raise TypeError(f'field {f.name} is a ClassVar but specifies ' - 'kw_only') - - # For real fields, disallow mutable defaults for known types. - if f._field_type is _FIELD and isinstance(f.default, (list, dict, set)): - raise ValueError(f'mutable default {type(f.default)} for field ' - f'{f.name} is not allowed: use default_factory') - - return f - -def _set_qualname(cls, value): - # Ensure that the functions returned from _create_fn uses the proper - # __qualname__ (the class they belong to). - if isinstance(value, FunctionType): - value.__qualname__ = f"{cls.__qualname__}.{value.__name__}" - return value - -def _set_new_attribute(cls, name, value): - # Never overwrites an existing attribute. Returns True if the - # attribute already exists. - if name in cls.__dict__: - return True - _set_qualname(cls, value) - setattr(cls, name, value) - return False - - -# Decide if/how we're going to create a hash function. Key is -# (unsafe_hash, eq, frozen, does-hash-exist). Value is the action to -# take. The common case is to do nothing, so instead of providing a -# function that is a no-op, use None to signify that. - -def _hash_set_none(cls, fields, globals): - return None - -def _hash_add(cls, fields, globals): - flds = [f for f in fields if (f.compare if f.hash is None else f.hash)] - return _set_qualname(cls, _hash_fn(flds, globals)) - -def _hash_exception(cls, fields, globals): - # Raise an exception. - raise TypeError(f'Cannot overwrite attribute __hash__ ' - f'in class {cls.__name__}') - -# -# +-------------------------------------- unsafe_hash? -# | +------------------------------- eq? -# | | +------------------------ frozen? -# | | | +---------------- has-explicit-hash? -# | | | | -# | | | | +------- action -# | | | | | -# v v v v v -_hash_action = {(False, False, False, False): None, - (False, False, False, True ): None, - (False, False, True, False): None, - (False, False, True, True ): None, - (False, True, False, False): _hash_set_none, - (False, True, False, True ): None, - (False, True, True, False): _hash_add, - (False, True, True, True ): None, - (True, False, False, False): _hash_add, - (True, False, False, True ): _hash_exception, - (True, False, True, False): _hash_add, - (True, False, True, True ): _hash_exception, - (True, True, False, False): _hash_add, - (True, True, False, True ): _hash_exception, - (True, True, True, False): _hash_add, - (True, True, True, True ): _hash_exception, - } -# See https://bugs.python.org/issue32929#msg312829 for an if-statement -# version of this table. - - -def _process_class(cls, init, repr, eq, order, unsafe_hash, frozen, - match_args, kw_only, slots): - # Now that dicts retain insertion order, there's no reason to use - # an ordered dict. I am leveraging that ordering here, because - # derived class fields overwrite base class fields, but the order - # is defined by the base class, which is found first. - fields = {} - - if cls.__module__ in sys.modules: - globals = sys.modules[cls.__module__].__dict__ - else: - # Theoretically this can happen if someone writes - # a custom string to cls.__module__. In which case - # such dataclass won't be fully introspectable - # (w.r.t. typing.get_type_hints) but will still function - # correctly. - globals = {} - - setattr(cls, _PARAMS, _DataclassParams(init, repr, eq, order, - unsafe_hash, frozen)) - - # Find our base classes in reverse MRO order, and exclude - # ourselves. In reversed order so that more derived classes - # override earlier field definitions in base classes. As long as - # we're iterating over them, see if any are frozen. - any_frozen_base = False - has_dataclass_bases = False - for b in cls.__mro__[-1:0:-1]: - # Only process classes that have been processed by our - # decorator. That is, they have a _FIELDS attribute. - base_fields = getattr(b, _FIELDS, None) - if base_fields is not None: - has_dataclass_bases = True - for f in base_fields.values(): - fields[f.name] = f - if getattr(b, _PARAMS).frozen: - any_frozen_base = True - - # Annotations that are defined in this class (not in base - # classes). If __annotations__ isn't present, then this class - # adds no new annotations. We use this to compute fields that are - # added by this class. - # - # Fields are found from cls_annotations, which is guaranteed to be - # ordered. Default values are from class attributes, if a field - # has a default. If the default value is a Field(), then it - # contains additional info beyond (and possibly including) the - # actual default value. Pseudo-fields ClassVars and InitVars are - # included, despite the fact that they're not real fields. That's - # dealt with later. - cls_annotations = cls.__dict__.get('__annotations__', {}) - - # Now find fields in our class. While doing so, validate some - # things, and set the default values (as class attributes) where - # we can. - cls_fields = [] - # Get a reference to this module for the _is_kw_only() test. - KW_ONLY_seen = False - dataclasses = sys.modules[__name__] - for name, type in cls_annotations.items(): - # See if this is a marker to change the value of kw_only. - if (_is_kw_only(type, dataclasses) - or (isinstance(type, str) - and _is_type(type, cls, dataclasses, dataclasses.KW_ONLY, - _is_kw_only))): - # Switch the default to kw_only=True, and ignore this - # annotation: it's not a real field. - if KW_ONLY_seen: - raise TypeError(f'{name!r} is KW_ONLY, but KW_ONLY ' - 'has already been specified') - KW_ONLY_seen = True - kw_only = True - else: - # Otherwise it's a field of some type. - cls_fields.append(_get_field(cls, name, type, kw_only)) - - for f in cls_fields: - fields[f.name] = f - - # If the class attribute (which is the default value for this - # field) exists and is of type 'Field', replace it with the - # real default. This is so that normal class introspection - # sees a real default value, not a Field. - if isinstance(getattr(cls, f.name, None), Field): - if f.default is MISSING: - # If there's no default, delete the class attribute. - # This happens if we specify field(repr=False), for - # example (that is, we specified a field object, but - # no default value). Also if we're using a default - # factory. The class attribute should not be set at - # all in the post-processed class. - delattr(cls, f.name) - else: - setattr(cls, f.name, f.default) - - # Do we have any Field members that don't also have annotations? - for name, value in cls.__dict__.items(): - if isinstance(value, Field) and not name in cls_annotations: - raise TypeError(f'{name!r} is a field but has no type annotation') - - # Check rules that apply if we are derived from any dataclasses. - if has_dataclass_bases: - # Raise an exception if any of our bases are frozen, but we're not. - if any_frozen_base and not frozen: - raise TypeError('cannot inherit non-frozen dataclass from a ' - 'frozen one') - - # Raise an exception if we're frozen, but none of our bases are. - if not any_frozen_base and frozen: - raise TypeError('cannot inherit frozen dataclass from a ' - 'non-frozen one') - - # Remember all of the fields on our class (including bases). This - # also marks this class as being a dataclass. - setattr(cls, _FIELDS, fields) - - # Was this class defined with an explicit __hash__? Note that if - # __eq__ is defined in this class, then python will automatically - # set __hash__ to None. This is a heuristic, as it's possible - # that such a __hash__ == None was not auto-generated, but it - # close enough. - class_hash = cls.__dict__.get('__hash__', MISSING) - has_explicit_hash = not (class_hash is MISSING or - (class_hash is None and '__eq__' in cls.__dict__)) - - # If we're generating ordering methods, we must be generating the - # eq methods. - if order and not eq: - raise ValueError('eq must be true if order is true') - - # Include InitVars and regular fields (so, not ClassVars). This is - # initialized here, outside of the "if init:" test, because std_init_fields - # is used with match_args, below. - all_init_fields = [f for f in fields.values() - if f._field_type in (_FIELD, _FIELD_INITVAR)] - (std_init_fields, - kw_only_init_fields) = _fields_in_init_order(all_init_fields) - - if init: - # Does this class have a post-init function? - has_post_init = hasattr(cls, _POST_INIT_NAME) - - _set_new_attribute(cls, '__init__', - _init_fn(all_init_fields, - std_init_fields, - kw_only_init_fields, - frozen, - has_post_init, - # The name to use for the "self" - # param in __init__. Use "self" - # if possible. - '__dataclass_self__' if 'self' in fields - else 'self', - globals, - slots, - )) - - # Get the fields as a list, and include only real fields. This is - # used in all of the following methods. - field_list = [f for f in fields.values() if f._field_type is _FIELD] - - if repr: - flds = [f for f in field_list if f.repr] - _set_new_attribute(cls, '__repr__', _repr_fn(flds, globals)) - - if eq: - # Create __eq__ method. There's no need for a __ne__ method, - # since python will call __eq__ and negate it. - flds = [f for f in field_list if f.compare] - self_tuple = _tuple_str('self', flds) - other_tuple = _tuple_str('other', flds) - _set_new_attribute(cls, '__eq__', - _cmp_fn('__eq__', '==', - self_tuple, other_tuple, - globals=globals)) - - if order: - # Create and set the ordering methods. - flds = [f for f in field_list if f.compare] - self_tuple = _tuple_str('self', flds) - other_tuple = _tuple_str('other', flds) - for name, op in [('__lt__', '<'), - ('__le__', '<='), - ('__gt__', '>'), - ('__ge__', '>='), - ]: - if _set_new_attribute(cls, name, - _cmp_fn(name, op, self_tuple, other_tuple, - globals=globals)): - raise TypeError(f'Cannot overwrite attribute {name} ' - f'in class {cls.__name__}. Consider using ' - 'functools.total_ordering') - - if frozen: - for fn in _frozen_get_del_attr(cls, field_list, globals): - if _set_new_attribute(cls, fn.__name__, fn): - raise TypeError(f'Cannot overwrite attribute {fn.__name__} ' - f'in class {cls.__name__}') - - # Decide if/how we're going to create a hash function. - hash_action = _hash_action[bool(unsafe_hash), - bool(eq), - bool(frozen), - has_explicit_hash] - if hash_action: - # No need to call _set_new_attribute here, since by the time - # we're here the overwriting is unconditional. - cls.__hash__ = hash_action(cls, field_list, globals) - - if not getattr(cls, '__doc__'): - # Create a class doc-string. - cls.__doc__ = (cls.__name__ + - str(inspect.signature(cls)).replace(' -> None', '')) - - if match_args: - # I could probably compute this once - _set_new_attribute(cls, '__match_args__', - tuple(f.name for f in std_init_fields)) - - if slots: - cls = _add_slots(cls, frozen) - - return cls - - -# _dataclass_getstate and _dataclass_setstate are needed for pickling frozen -# classes with slots. These could be slighly more performant if we generated -# the code instead of iterating over fields. But that can be a project for -# another day, if performance becomes an issue. -def _dataclass_getstate(self): - return [getattr(self, f.name) for f in fields(self)] - - -def _dataclass_setstate(self, state): - for field, value in zip(fields(self), state): - # use setattr because dataclass may be frozen - object.__setattr__(self, field.name, value) - - -def _add_slots(cls, is_frozen): - # Need to create a new class, since we can't set __slots__ - # after a class has been created. - - # Make sure __slots__ isn't already set. - if '__slots__' in cls.__dict__: - raise TypeError(f'{cls.__name__} already specifies __slots__') - - # Create a new dict for our new class. - cls_dict = dict(cls.__dict__) - field_names = tuple(f.name for f in fields(cls)) - cls_dict['__slots__'] = field_names - for field_name in field_names: - # Remove our attributes, if present. They'll still be - # available in _MARKER. - cls_dict.pop(field_name, None) - - # Remove __dict__ itself. - cls_dict.pop('__dict__', None) - - # And finally create the class. - qualname = getattr(cls, '__qualname__', None) - cls = type(cls)(cls.__name__, cls.__bases__, cls_dict) - if qualname is not None: - cls.__qualname__ = qualname - - if is_frozen: - # Need this for pickling frozen classes with slots. - cls.__getstate__ = _dataclass_getstate - cls.__setstate__ = _dataclass_setstate - - return cls - - -def dataclass(cls=None, /, *, init=True, repr=True, eq=True, order=False, - unsafe_hash=False, frozen=False, match_args=True, - kw_only=False, slots=False): - """Returns the same class as was passed in, with dunder methods - added based on the fields defined in the class. - - Examines PEP 526 __annotations__ to determine fields. - - If init is true, an __init__() method is added to the class. If - repr is true, a __repr__() method is added. If order is true, rich - comparison dunder methods are added. If unsafe_hash is true, a - __hash__() method function is added. If frozen is true, fields may - not be assigned to after instance creation. If match_args is true, - the __match_args__ tuple is added. If kw_only is true, then by - default all fields are keyword-only. If slots is true, an - __slots__ attribute is added. - """ - - def wrap(cls): - return _process_class(cls, init, repr, eq, order, unsafe_hash, - frozen, match_args, kw_only, slots) - - # See if we're being called as @dataclass or @dataclass(). - if cls is None: - # We're called with parens. - return wrap - - # We're called as @dataclass without parens. - return wrap(cls) - - -def fields(class_or_instance): - """Return a tuple describing the fields of this dataclass. - - Accepts a dataclass or an instance of one. Tuple elements are of - type Field. - """ - - # Might it be worth caching this, per class? - try: - fields = getattr(class_or_instance, _FIELDS) - except AttributeError: - raise TypeError('must be called with a dataclass type or instance') - - # Exclude pseudo-fields. Note that fields is sorted by insertion - # order, so the order of the tuple is as the fields were defined. - return tuple(f for f in fields.values() if f._field_type is _FIELD) - - -def _is_dataclass_instance(obj): - """Returns True if obj is an instance of a dataclass.""" - return hasattr(type(obj), _FIELDS) - - -def is_dataclass(obj): - """Returns True if obj is a dataclass or an instance of a - dataclass.""" - cls = obj if isinstance(obj, type) else type(obj) - return hasattr(cls, _FIELDS) - - -def asdict(obj, *, dict_factory=dict): - """Return the fields of a dataclass instance as a new dictionary mapping - field names to field values. - - Example usage: - - @dataclass - class C: - x: int - y: int - - c = C(1, 2) - assert asdict(c) == {'x': 1, 'y': 2} - - If given, 'dict_factory' will be used instead of built-in dict. - The function applies recursively to field values that are - dataclass instances. This will also look into built-in containers: - tuples, lists, and dicts. - """ - if not _is_dataclass_instance(obj): - raise TypeError("asdict() should be called on dataclass instances") - return _asdict_inner(obj, dict_factory) - - -def _asdict_inner(obj, dict_factory): - if _is_dataclass_instance(obj): - result = [] - for f in fields(obj): - value = _asdict_inner(getattr(obj, f.name), dict_factory) - result.append((f.name, value)) - return dict_factory(result) - elif isinstance(obj, tuple) and hasattr(obj, '_fields'): - # obj is a namedtuple. Recurse into it, but the returned - # object is another namedtuple of the same type. This is - # similar to how other list- or tuple-derived classes are - # treated (see below), but we just need to create them - # differently because a namedtuple's __init__ needs to be - # called differently (see bpo-34363). - - # I'm not using namedtuple's _asdict() - # method, because: - # - it does not recurse in to the namedtuple fields and - # convert them to dicts (using dict_factory). - # - I don't actually want to return a dict here. The main - # use case here is json.dumps, and it handles converting - # namedtuples to lists. Admittedly we're losing some - # information here when we produce a json list instead of a - # dict. Note that if we returned dicts here instead of - # namedtuples, we could no longer call asdict() on a data - # structure where a namedtuple was used as a dict key. - - return type(obj)(*[_asdict_inner(v, dict_factory) for v in obj]) - elif isinstance(obj, (list, tuple)): - # Assume we can create an object of this type by passing in a - # generator (which is not true for namedtuples, handled - # above). - return type(obj)(_asdict_inner(v, dict_factory) for v in obj) - elif isinstance(obj, dict): - return type(obj)((_asdict_inner(k, dict_factory), - _asdict_inner(v, dict_factory)) - for k, v in obj.items()) - else: - return copy.deepcopy(obj) - - -def astuple(obj, *, tuple_factory=tuple): - """Return the fields of a dataclass instance as a new tuple of field values. - - Example usage:: - - @dataclass - class C: - x: int - y: int - - c = C(1, 2) - assert astuple(c) == (1, 2) - - If given, 'tuple_factory' will be used instead of built-in tuple. - The function applies recursively to field values that are - dataclass instances. This will also look into built-in containers: - tuples, lists, and dicts. - """ - - if not _is_dataclass_instance(obj): - raise TypeError("astuple() should be called on dataclass instances") - return _astuple_inner(obj, tuple_factory) - - -def _astuple_inner(obj, tuple_factory): - if _is_dataclass_instance(obj): - result = [] - for f in fields(obj): - value = _astuple_inner(getattr(obj, f.name), tuple_factory) - result.append(value) - return tuple_factory(result) - elif isinstance(obj, tuple) and hasattr(obj, '_fields'): - # obj is a namedtuple. Recurse into it, but the returned - # object is another namedtuple of the same type. This is - # similar to how other list- or tuple-derived classes are - # treated (see below), but we just need to create them - # differently because a namedtuple's __init__ needs to be - # called differently (see bpo-34363). - return type(obj)(*[_astuple_inner(v, tuple_factory) for v in obj]) - elif isinstance(obj, (list, tuple)): - # Assume we can create an object of this type by passing in a - # generator (which is not true for namedtuples, handled - # above). - return type(obj)(_astuple_inner(v, tuple_factory) for v in obj) - elif isinstance(obj, dict): - return type(obj)((_astuple_inner(k, tuple_factory), _astuple_inner(v, tuple_factory)) - for k, v in obj.items()) - else: - return copy.deepcopy(obj) - - -def make_dataclass(cls_name, fields, *, bases=(), namespace=None, init=True, - repr=True, eq=True, order=False, unsafe_hash=False, - frozen=False, match_args=True, kw_only=False, slots=False): - """Return a new dynamically created dataclass. - - The dataclass name will be 'cls_name'. 'fields' is an iterable - of either (name), (name, type) or (name, type, Field) objects. If type is - omitted, use the string 'typing.Any'. Field objects are created by - the equivalent of calling 'field(name, type [, Field-info])'. - - C = make_dataclass('C', ['x', ('y', int), ('z', int, field(init=False))], bases=(Base,)) - - is equivalent to: - - @dataclass - class C(Base): - x: 'typing.Any' - y: int - z: int = field(init=False) - - For the bases and namespace parameters, see the builtin type() function. - - The parameters init, repr, eq, order, unsafe_hash, and frozen are passed to - dataclass(). - """ - - if namespace is None: - namespace = {} - - # While we're looking through the field names, validate that they - # are identifiers, are not keywords, and not duplicates. - seen = set() - annotations = {} - defaults = {} - for item in fields: - if isinstance(item, str): - name = item - tp = 'typing.Any' - elif len(item) == 2: - name, tp, = item - elif len(item) == 3: - name, tp, spec = item - defaults[name] = spec - else: - raise TypeError(f'Invalid field: {item!r}') - - if not isinstance(name, str) or not name.isidentifier(): - raise TypeError(f'Field names must be valid identifiers: {name!r}') - if keyword.iskeyword(name): - raise TypeError(f'Field names must not be keywords: {name!r}') - if name in seen: - raise TypeError(f'Field name duplicated: {name!r}') - - seen.add(name) - annotations[name] = tp - - # Update 'ns' with the user-supplied namespace plus our calculated values. - def exec_body_callback(ns): - ns.update(namespace) - ns.update(defaults) - ns['__annotations__'] = annotations - - # We use `types.new_class()` instead of simply `type()` to allow dynamic creation - # of generic dataclasses. - cls = types.new_class(cls_name, bases, {}, exec_body_callback) - - # Apply the normal decorator. - return dataclass(cls, init=init, repr=repr, eq=eq, order=order, - unsafe_hash=unsafe_hash, frozen=frozen, - match_args=match_args, kw_only=kw_only, slots=slots) - - -def replace(obj, /, **changes): - """Return a new object replacing specified fields with new values. - - This is especially useful for frozen classes. Example usage: - - @dataclass(frozen=True) - class C: - x: int - y: int - - c = C(1, 2) - c1 = replace(c, x=3) - assert c1.x == 3 and c1.y == 2 - """ - - # We're going to mutate 'changes', but that's okay because it's a - # new dict, even if called with 'replace(obj, **my_changes)'. - - if not _is_dataclass_instance(obj): - raise TypeError("replace() should be called on dataclass instances") - - # It's an error to have init=False fields in 'changes'. - # If a field is not in 'changes', read its value from the provided obj. - - for f in getattr(obj, _FIELDS).values(): - # Only consider normal fields or InitVars. - if f._field_type is _FIELD_CLASSVAR: - continue - - if not f.init: - # Error if this field is specified in changes. - if f.name in changes: - raise ValueError(f'field {f.name} is declared with ' - 'init=False, it cannot be specified with ' - 'replace()') - continue - - if f.name not in changes: - if f._field_type is _FIELD_INITVAR and f.default is MISSING: - raise ValueError(f"InitVar {f.name!r} " - 'must be specified with replace()') - changes[f.name] = getattr(obj, f.name) - - # Create the new object, which calls __init__() and - # __post_init__() (if defined), using all of the init fields we've - # added and/or left in 'changes'. If there are values supplied in - # changes that aren't fields, this will correctly raise a - # TypeError. - return obj.__class__(**changes) \ No newline at end of file diff --git a/construct_typed/generic.py b/construct_typed/generic_wrapper.py similarity index 69% rename from construct_typed/generic.py rename to construct_typed/generic_wrapper.py index b988444..cd4788b 100644 --- a/construct_typed/generic.py +++ b/construct_typed/generic_wrapper.py @@ -12,11 +12,14 @@ if t.TYPE_CHECKING: # 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 ConstantOrContextLambda as ConstantOrContextLambda + from construct import ConstantOrContextLambda2 as ConstantOrContextLambda2 from construct import Construct as Construct from construct import Context as Context from construct import ListContainer as ListContainer 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: import construct as cs @@ -37,22 +40,18 @@ else: class Context: 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]] + ConstantOrContextLambda2 = t.Union[ValueType, t.Callable[[Context], ValueType]] 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 diff --git a/construct_typed/tenum.py b/construct_typed/tenum.py index 5cb9116..4417c6b 100644 --- a/construct_typed/tenum.py +++ b/construct_typed/tenum.py @@ -1,147 +1,111 @@ +# pyright: reportAny=false import enum -import textwrap import typing as t -import construct as cs +from typing_extensions import Self, override -from construct_typed.generic import * - -T = t.TypeVar("T") +from .generic_wrapper import Construct, Adapter, Context, PathType -class _EnumMeta(enum.EnumMeta): - @classmethod - def __prepare__( - metacls, # type: ignore - __name: str, - __bases: t.Tuple[type, ...], - **kwargs: t.Any, - ) -> t.Mapping[str, object]: - # This method is needed, because the original __prepare__ method does not accept kwargs. - return super().__prepare__(__name, __bases) +# ## TEnum ############################################################################################################ +class EnumValue: + """ + This is a helper class for adding documentation to an enum value. + """ - def __new__( - metacls: t.Type[T], # type: ignore - __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") + def __init__(self, value: int, doc: str | None = None) -> None: + self.value: int = value + self.__doc__ = doc if doc else "" - # create new enum object - cls: T = super().__new__(metacls, __name, __bases, __namespace) # type: ignore - # if the `TEnum` class is created, there are no parameters - if len(kwargs) == 0: - return cls +class EnumBase(enum.IntEnum): + """ + Base class for an Enum used in `construct_typed.TEnum`. - # extract parameters from kwargs - subcon: "cs.Construct[t.Any, t.Any]" = kwargs.pop("subcon", None) - if not isinstance(subcon, cs.Construct): # type: ignore - raise ValueError( - f"`subcon` parameter has to be an `Construct` object but is {type(subcon)}" - ) - if len(kwargs) > 0: # check remaining parameters - unsupp_parm = ", ".join([f"'{k}'" for k in kwargs.keys()]) - raise ValueError(f"unsupported parameter(s) detected: {unsupp_parm}") + This class extends the standard `enum.IntEnum` by. + - missing values are automatically generated + - possibility to add documentation for each enum value (see `EnumValue`) - # create construct format - if TEnum in __bases: - enum_constr = TEnumConstruct(subcon, cls) # type: ignore - elif TFlags in __bases: - enum_constr = TFlagsConstruct(subcon, cls) # type: ignore + Example:: + + >>> class State(EnumBase): + ... Idle = 1 + ... Running = EnumValue(2, "This is the running state.") + + >>> State(1) + + + >>> State["Idle"] + + + >>> State.Idle + + + >>> State(3) # missing value + + + >>> 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: - raise TypeError("neither `TEnum` nor `TFlags` in bases") + obj = int.__new__(cls, val) + obj._value_ = val + obj.__doc__ = "" + return obj - # 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 + # Extend the enum type with _missing_ method. So if a enum value # not found in the enum, a new pseudo member is created. # The idea is taken from: https://stackoverflow.com/a/57179436 @classmethod - def _missing_(cls, value: t.Any) -> t.Optional["TEnum"]: + @override + def _missing_(cls, value: t.Any) -> enum.Enum | None: if isinstance(value, int): - return cls._create_pseudo_member_(value) + pseudo_member = cls._value2member_map_.get(value, None) + 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__ - @classmethod - def _create_pseudo_member_(cls, value: int) -> "TEnum": - pseudo_member = cls._value2member_map_.get(value, None) # type: ignore - if pseudo_member is None: - new_member = int.__new__(cls, value) - # 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 - pseudo_member = cls._value2member_map_.setdefault(value, new_member) # type: ignore - return pseudo_member # type: ignore + @override + def __reduce_ex__(self, proto: t.Any) -> tuple[t.Any, ...]: + """ + Pickle enums by value instead of name (restores pre-3.11 behavior). + See https://github.com/python/cpython/pull/26658 for why this exists. + """ + return self.__class__, (self._value_,) -EnumType = t.TypeVar("EnumType", bound=TEnum) +EnumType = t.TypeVar("EnumType", bound=EnumBase) -class TEnumConstruct(Adapter[int, int, EnumType, EnumType]): +class TEnum(Adapter[int, int, EnumType, EnumType]): """ Typed enum. """ - - 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)) - ) - + def __init__(self, subcon: Construct[int, int], enum_type: type[EnumType]): # save enum type - self.enum_type = t.cast(t.Type[EnumType], enum_type) # type: ignore + self.enum_type: type[EnumType] = enum_type # init adatper - super(TEnumConstruct, self).__init__(subcon) # type: ignore + super(TEnum, self).__init__(subcon) # type: ignore + @override def _decode(self, obj: int, context: Context, path: PathType) -> EnumType: return self.enum_type(obj) + @override def _encode( self, obj: EnumType, @@ -155,60 +119,91 @@ class TEnumConstruct(Adapter[int, int, EnumType, EnumType]): ) -# ## TFlags ####################################################################################################### -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]): +# ## TFlagsEnum ####################################################################################################### +class FlagsEnumBase(enum.IntFlag): """ - Typed flags. + Base class for an Enum used in `construct_typed.TFlagsEnum`. + + 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"] + + + >>> Option.OptOne + + + >>> Option(3) + + + >>> Option(4) + + + >>> Option.OptTwo.__doc__ # documentation + 'This is option two.' """ - if t.TYPE_CHECKING: + 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: + obj = int.__new__(cls, val) + obj._value_ = val + obj.__doc__ = "" + return obj - def __new__( - cls, subcon: Construct[int, int], enum_type: t.Type[FlagsType] - ) -> "TFlagsConstruct[FlagsType]": - ... + @classmethod + @override + def _missing_(cls, value: t.Any) -> t.Any: + """ + 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 - def __init__(self, subcon: Construct[int, int], enum_type: t.Type[FlagsType]): - if not issubclass(enum_type, TFlags): - raise TypeError( - "'{}' has to be a '{}'".format(repr(enum_type), repr(TFlags)) - ) + @override + def __reduce_ex__(self, proto: t.Any) -> tuple[t.Any, ...]: + """ + Pickle enums by value instead of name (restores pre-3.11 behavior). + 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 - self.enum_type = t.cast(t.Type[FlagsType], enum_type) # type: ignore + self.enum_type: type[FlagsEnumType] = enum_type # init adatper - super(TFlagsConstruct, self).__init__(subcon) # type: ignore + super(TFlagsEnum, self).__init__(subcon) # type: ignore - def _decode(self, obj: int, context: Context, path: PathType) -> FlagsType: + @override + def _decode(self, obj: int, context: Context, path: PathType) -> FlagsEnumType: return self.enum_type(obj) + @override def _encode( self, - obj: FlagsType, + obj: FlagsEnumType, context: Context, path: PathType, ) -> int: diff --git a/construct_typed/version.py b/construct_typed/version.py index 9498baa..38a2845 100644 --- a/construct_typed/version.py +++ b/construct_typed/version.py @@ -1,2 +1,2 @@ -version = (0, 5, 2) -version_string = "0.5.2" +version = (0, 7, 0) +version_string = "0.7.0+wrapper" diff --git a/mypy.ini b/mypy.ini deleted file mode 100644 index 3412486..0000000 --- a/mypy.ini +++ /dev/null @@ -1,3 +0,0 @@ -[mypy] -strict = True -warn_unused_ignores = False \ No newline at end of file diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..8a2689b --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,76 @@ + +[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 \ No newline at end of file diff --git a/requirements.txt b/requirements.txt index 2d70b89..2514c06 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,10 +1,14 @@ -construct==2.10.67 +construct==2.10.70 pytest>=6.2.0 -numpy==1.21.* +numpy arrow ruamel.yaml cloudpickle lz4 black isort -mypy \ No newline at end of file +mypy +cryptography +build +setuptools +wheel diff --git a/setup.py b/setup.py deleted file mode 100644 index 6cdeb7a..0000000 --- a/setup.py +++ /dev/null @@ -1,64 +0,0 @@ -#!/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", - ], -) diff --git a/tests/declarativeunittest.py b/tests/declarativeunittest.py index 1d1be0c..ed7fb59 100644 --- a/tests/declarativeunittest.py +++ b/tests/declarativeunittest.py @@ -1,38 +1,170 @@ +import binascii +import io +import typing as t + import pytest +from construct import * +from construct.lib import * + +import construct_typed as cst xfail = pytest.mark.xfail skip = pytest.mark.skip skipif = pytest.mark.skipif -import os, math, random, collections, itertools, io, hashlib, binascii +Buffer = t.Union[bytes, memoryview, bytearray] +ParsedType = t.TypeVar("ParsedType") +BuildTypes = t.TypeVar("BuildTypes") +ContainerType = t.TypeVar("ContainerType", bound=cst.TContainerMixin) +T = t.TypeVar("T") -from construct import * -from construct.lib import * +IdentType = t.TypeVar("IdentType") class ZeroIO(io.BufferedIOBase): - def read(self, __size=None): + def read(self, __size: t.Optional[int] = None) -> bytes: if __size is not None: return bytes(__size) else: return bytes(0) - def read1(self, __size=0): + def read1(self, __size: int = 0) -> bytes: return bytes(__size) -ident = lambda x: x -devzero = ZeroIO() +def ident(x: IdentType) -> IdentType: + return x -def raises(func, *args, **kw): +devzero: t.BinaryIO = ZeroIO() # type: ignore + + +def raises( + func: t.Callable[..., t.Any], *args: t.Any, **kw: t.Any +) -> t.Union[t.Any, Exception]: try: return func(*args, **kw) except Exception as e: return e.__class__ -def common(format, datasample, objsample, sizesample=SizeofError, **kw): +@t.overload +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) assert obj == objsample data = format.build(objsample, **kw) @@ -44,35 +176,35 @@ def common(format, datasample, objsample, sizesample=SizeofError, **kw): size = format.sizeof(**kw) assert size == sizesample else: - size = raises(format.sizeof, **kw) - assert size == sizesample + size_ex = raises(format.sizeof, **kw) + assert size_ex == sizesample -def setattrs(obj, **kwargs): - """ Set multiple named values of an object """ +def setattrs(obj: T, **kwargs: t.Any) -> T: + """Set multiple named values of an object""" for name, value in kwargs.items(): setattr(obj, name, value) return obj -def commonhex(format, hexdata): +def commonhex(format: "Construct[t.Any, t.Any]", hexdata: str) -> None: commonbytes(format, binascii.unhexlify(hexdata)) -def commondumpdeprecated(format, filename): +def commondumpdeprecated(format: "Construct[t.Any, t.Any]", filename: str) -> None: filename = "tests/deprecated_gallery/blobs/" + filename with open(filename, "rb") as f: data = f.read() commonbytes(format, data) -def commondump(format, filename): +def commondump(format: "Construct[t.Any, t.Any]", filename: str) -> None: filename = "tests/gallery/blobs/" + filename with open(filename, "rb") as f: data = f.read() commonbytes(format, data) -def commonbytes(format, data): +def commonbytes(format: "Construct[t.Any, t.Any]", data: bytes) -> None: obj = format.parse(data) - data2 = format.build(obj) + format.build(obj) diff --git a/tests/declarativeunittest.pyi b/tests/declarativeunittest.pyi deleted file mode 100644 index 5afd044..0000000 --- a/tests/declarativeunittest.pyi +++ /dev/null @@ -1,109 +0,0 @@ -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: ... diff --git a/tests/test_core.py b/tests/test_core.py index 6f65772..602a899 100644 --- a/tests/test_core.py +++ b/tests/test_core.py @@ -1,6 +1,6 @@ # -*- coding: utf-8 -*- - -from .declarativeunittest import raises, common, commonhex, commondumpdeprecated, commondump, commonbytes, ident, devzero +# mypy: no-warn-unused-ignores +from .declarativeunittest import raises, common, ident, devzero from construct.core import * from construct import * from construct.lib import * @@ -151,17 +151,29 @@ def test_formatfield_bool_issue_901() -> None: assert d.sizeof() == 1 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) common(d, b"\x01\x02\x03\x04", 0x01020304, 4) common(d, b"\xff\xff\xff\xff", -1, 4) d = BytesInteger(4, signed=False, swapped=this.swapped) common(d, b"\x01\x02\x03\x04", 0x01020304, 4, swapped=False) 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(4, signed=False).build, -1) == IntegerError - common(BytesInteger(0), b"", 0, 0) def test_bitsinteger() -> None: + d = BitsInteger(0) + assert raises(d.parse, b"") == IntegerError + assert raises(d.build, 0) == IntegerError d = BitsInteger(8) common(d, b"\x01\x01\x01\x01\x01\x01\x01\x01", 255, 8) d = BitsInteger(8, signed=True) @@ -171,9 +183,17 @@ def test_bitsinteger() -> None: 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"\x00\x00\x00\x00\x00\x00\x00\x00\x01\x01\x01\x01\x01\x01\x01\x01", 0xff00, 16, swapped=True) - assert raises(BitsInteger(this.missing).sizeof) == SizeofError + 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 - common(BitsInteger(0), b"", 0, 0) + 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 def test_varint() -> None: d = VarInt @@ -224,8 +244,8 @@ def test_paddedstring() -> None: common(PaddedString(100, e), data, s, 100) for e in ["ascii","utf8","utf16","utf-16-le","utf32","utf-32-le"]: - PaddedString(10, e).sizeof() == 10 - PaddedString(this.n, e).sizeof(n=10) == 10 + assert PaddedString(10, e).sizeof() == 10 + assert PaddedString(this.n, e).sizeof(n=10) == 10 def test_pascalstring() -> None: for e,_ in [("utf8",1),("utf16",2),("utf_16_le",2),("utf32",4),("utf_32_le",4)]: @@ -236,8 +256,8 @@ def test_pascalstring() -> None: common(PascalString(sc, e), sc.build(0), u"") for e in ["utf8","utf16","utf-16-le","utf32","utf-32-le","ascii"]: - raises(PascalString(Byte, e).sizeof) == SizeofError - raises(PascalString(VarInt, e).sizeof) == SizeofError + assert raises(PascalString(Byte, e).sizeof) == SizeofError + assert raises(PascalString(VarInt, e).sizeof) == SizeofError def test_cstring() -> None: s = u"" @@ -246,12 +266,12 @@ def test_cstring() -> None: common(CString(e), s.encode(e)+bytes(us), s) common(CString(e), bytes(us), u"") - CString("utf8").build(s) == b'\xd0\x90\xd1\x84\xd0\xbe\xd0\xbd'+b"\x00" - CString("utf16").build(s) == b'\xff\xfe\x10\x04D\x04>\x04=\x04'+b"\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" + assert 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" + 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" for e in ["utf8","utf16","utf-16-le","utf32","utf-32-le","ascii"]: - raises(CString(e).sizeof) == SizeofError + assert raises(CString(e).sizeof) == SizeofError def test_greedystring() -> None: for e,_ in [("utf8",1),("utf16",2),("utf_16_le",2),("utf32",4),("utf_32_le",4)]: @@ -260,7 +280,7 @@ def test_greedystring() -> None: common(GreedyString(e), b"", u"") for e in ["utf8","utf16","utf-16-le","utf32","utf-32-le","ascii"]: - raises(GreedyString(e).sizeof) == SizeofError + assert raises(GreedyString(e).sizeof) == SizeofError def test_string_encodings() -> None: # checks that "-" is replaced with "_" @@ -271,7 +291,7 @@ def test_flag() -> None: d = Flag common(d, b"\x00", False, 1) common(d, b"\x01", True, 1) - d.parse(b"\xff") == True + assert d.parse(b"\xff") == True def test_enum() -> None: d = Enum(Byte, one=1, two=2, four=4, eight=8) @@ -420,11 +440,11 @@ def test_struct_proper_context() -> None: "x"/Byte, "inner"/Struct( "y"/Byte, - "a"/Computed(this._.x+1), - "b"/Computed(this.y+2), + "a"/Computed(this._.x+1), # type: ignore + "b"/Computed(this.y+2), # type: ignore ), - "c"/Computed(this.x+3), - "d"/Computed(this.inner.y+4), + "c"/Computed(this.x+3), # type: ignore + "d"/Computed(this.inner.y+4), # type: ignore ) assert d.parse(b"\x01\x0f") == Container(x=1, inner=Container(y=15, a=2, b=17), c=4, d=19) @@ -511,7 +531,7 @@ def test_const() -> None: def test_computed() -> None: common(Computed(255), b"", 255, 0) - common(Computed(lambda ctx: 255), b"", 255, 0) + common(Computed(lambda ctx: 255), b"", 255, 0) # type: ignore assert Computed(255).build(None) == b"" assert Struct(Computed(255)).build({}) == b"" assert raises(Computed(this.missing).parse, b"") == KeyError @@ -591,7 +611,7 @@ def test_rebuild_issue_664() -> None: def test_default() -> None: d = Default(Byte, 0) common(d, b"\xff", 255, 1) - d.build(None) == b"\x00" + assert d.build(None) == b"\x00" def test_check() -> None: common(Check(True), b"", None, 0) @@ -637,8 +657,7 @@ def test_numpy_error() -> None: numpy.load(io.BytesIO(b"")) # type: ignore def test_namedtuple() -> None: - import collections - coord = collections.namedtuple("coord", "x y z") + coord = t.NamedTuple("coord", [("x", int), ("y", int), ("z", int)]) d1 = NamedTuple("coord", "x y z", Array(3, Byte)) common(d1, b"123", coord(49,50,51), 3) d2 = NamedTuple("coord", "x y z", GreedyRange(Byte)) @@ -708,10 +727,13 @@ def test_hexdump() -> None: def test_hexdump_regression_issue_188() -> None: # Hex HexDump were not inheriting subcon flags - d = Struct(Hex(Const(b"MZ"))) + a = Hex(Const(b"MZ")) + d = Struct(a) assert d.parse(b"MZ") == Container() 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.build(dict()) == b"MZ" @@ -808,8 +830,10 @@ def test_select_buildfromnone_issue_747() -> None: assert d.build(dict()) == b"" def test_if() -> None: - common(If(True, Byte), b"\x01", 1, 1) - common(If(False, Byte), b"", None, 0) + d = If(True, Byte) + common(d, b"\x01", 1, 1) + d = If(False, Byte) + common(d, b"", None, 0) def test_ifthenelse() -> None: common(IfThenElse(True, Int8ub, Int16ub), b"\x01", 1, 1) @@ -922,6 +946,17 @@ def test_peek() -> None: assert d4.build(Container(a=0x01, b=0x0102)) == b"" 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: d = Seek(5) assert d.parse(b"") == 5 @@ -1033,13 +1068,14 @@ def test_prefixed() -> None: common(d5, b"\x0a"+bytes(10), u"\x00"*10, SizeofError) def test_prefixedarray() -> None: - common(PrefixedArray(Byte,Byte), b"\x02\x0a\x0b", [10,11], SizeofError) - assert PrefixedArray(Byte, Byte).parse(b"\x03\x01\x02\x03") == [1,2,3] - assert PrefixedArray(Byte, Byte).parse(b"\x00") == [] - assert PrefixedArray(Byte, Byte).build([1,2,3]) == b"\x03\x01\x02\x03" - assert raises(PrefixedArray(Byte, Byte).parse, b"") == StreamError - assert raises(PrefixedArray(Byte, Byte).parse, b"\x03\x01") == StreamError - assert raises(PrefixedArray(Byte, Byte).sizeof) == SizeofError + d = PrefixedArray(Byte, Byte) + common(d, b"\x02\x0a\x0b", [10,11], SizeofError) + assert d.parse(b"\x03\x01\x02\x03") == [1,2,3] + assert d.parse(b"\x00") == [] + assert d.build([1,2,3]) == b"\x03\x01\x02\x03" + assert raises(d.parse, b"") == StreamError + assert raises(d.parse, b"\x03\x01") == StreamError + assert raises(d.sizeof) == SizeofError def test_fixedsized() -> None: d1 = FixedSized(10, Byte) @@ -1213,7 +1249,7 @@ def test_checksum() -> None: def test_checksum_nonbytes_issue_323() -> None: d = Struct( "vals" / Byte[2], - "checksum" / Checksum(Byte, lambda vals: sum(vals) & 0xFF, this.vals), + "checksum" / Checksum(Byte, lambda vals: int(sum(vals)) & 0xFF, this.vals), ) assert d.parse(b"\x00\x00\x00") == Container(vals=[0, 0], checksum=0) assert raises(d.parse, b"\x00\x00\x01") == ChecksumError @@ -1329,6 +1365,105 @@ def test_compressed_prefixed() -> None: assert st.parse(st.build(Container(one=zeros,two=zeros))) == Container(one=zeros,two=zeros) 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: data = b"0" * 1000 assert Rebuffered(Array(1000,Byte)).parse_stream(io.BytesIO(data)) == [48]*1000 @@ -1544,7 +1679,7 @@ def test_operators() -> None: assert d.docs == "description" d = "description" * Byte assert d.docs == "description" - """ + _ = """ description """ * \ Byte @@ -1686,9 +1821,11 @@ 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),] def test_from_issue_269() -> None: - d = Struct("enabled" / Byte, If(this.enabled, Padding(2))) + a = 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=0)) == b"\x00" + 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=0)) == b"\x00" @@ -1704,7 +1841,7 @@ def test_from_issue_324() -> None: )), "checksum" / Checksum( Byte, - lambda data: sum(data) & 0xFF, + lambda data: int(sum(data)) & 0xFF, this.vals.data ), ) @@ -1795,11 +1932,11 @@ def test_pickling_constructs() -> None: ) data = bytes(100) - du = cloudpickle.loads(cloudpickle.dumps(d, protocol=-1)) + du = cloudpickle.loads(cloudpickle.dumps(d, protocol=-1)) # type: ignore assert du.parse(data) == d.parse(data) def test_pickling_constructs_issue_894() -> None: - import cloudpickle + import cloudpickle # type: ignore fundus_header = Struct( 'width' / Int32un, @@ -1811,7 +1948,7 @@ def test_pickling_constructs_issue_894() -> None: 'img' / Int8un, ) - cloudpickle.dumps(fundus_header) + cloudpickle.dumps(fundus_header) # type: ignore def test_exposing_members_attributes() -> None: d1 = Struct( @@ -2022,7 +2159,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))) def test_parsedhook_repeatersdiscard() -> None: - outputs = [] + outputs: t.List[int] = [] def printobj1(obj: int, ctx: "Context") -> None: outputs.append(obj) d1 = GreedyRange(Byte * printobj1, discard=True) diff --git a/tests/test_typed.py b/tests/test_typed.py index ed578c8..df8cc84 100644 --- a/tests/test_typed.py +++ b/tests/test_typed.py @@ -1,137 +1,152 @@ # -*- coding: utf-8 -*- # pyright: strict +import dataclasses +import enum +import textwrap import typing as t import construct as cs -from construct_typed import ( - DataclassBitStruct, - DataclassStruct, - csfield, - constr, - TEnum, - TFlags, -) -from tests.declarativeunittest import common, raises, setattrs +import construct_typed as cst +from construct_typed import DataclassBitStruct, DataclassMixin, DataclassStruct, csfield + +from .declarativeunittest import common, raises, setattrs def test_dataclass_const_default() -> None: - class TestDataclass(DataclassStruct): - const_bytes: bytes = csfield(cs.Bytes(3), const=b"BMP") - const_int: int = csfield(cs.Int8ub, const=5) - default_int: int = csfield(cs.Int8ub, default=26) - default_lambda: t.Optional[bytes] = csfield( + @dataclasses.dataclass + class ConstDefaultTest(DataclassMixin): + const_bytes: bytes = csfield(cs.Const(b"BMP")) + const_int: int = csfield(cs.Const(5, cs.Int8ub)) + default_int: int = csfield(cs.Default(cs.Int8ub, 28)) + default_lambda: bytes = csfield( cs.Default(cs.Bytes(cs.this.const_int), lambda ctx: bytes(ctx.const_int)) ) - obj = TestDataclass() - assert obj.const_bytes == b"BMP" - assert obj.const_int == 5 - assert obj.default_int == 26 - assert obj.default_lambda == None - obj = TestDataclass(default_int=1) - assert obj.default_int == 1 - - 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) + a = ConstDefaultTest() + assert a.const_bytes == b"BMP" + assert a.const_int == 5 + assert a.default_int == 28 + assert a.default_lambda == None def test_dataclass_access() -> None: - class TestDataclass(DataclassStruct): - a: int = csfield(cs.Byte, const=1) + @dataclasses.dataclass + class TestTContainer(DataclassMixin): + a: t.Optional[int] = csfield(cs.Const(1, cs.Byte)) b: int = csfield(cs.Int8ub) - obj = TestDataclass(b=2) + tcontainer = TestTContainer(b=2) - assert obj.a == 1 - assert obj["a"] == 1 - assert obj.b == 2 - assert obj["b"] == 2 + # tcontainer + assert tcontainer.a == 1 + assert tcontainer["a"] == 1 + assert tcontainer.b == 2 + assert tcontainer["b"] == 2 - obj.a = 5 - assert obj.a == 5 - assert obj["a"] == 5 - obj["a"] = 6 - assert obj.a == 6 - assert obj["a"] == 6 + tcontainer.a = 5 + assert tcontainer.a == 5 + assert tcontainer["a"] == 5 + tcontainer["a"] = 6 + assert tcontainer.a == 6 + assert tcontainer["a"] == 6 # wrong creation - assert raises(lambda: TestDataclass(a=0, b=1)) == TypeError # type: ignore + assert raises(lambda: TestTContainer(a=0, b=1)) == TypeError def test_dataclass_str_repr() -> None: - class Image(DataclassStruct): - signature: bytes = csfield(cs.Bytes(3), const=b"BMP") + @dataclasses.dataclass + class Image(DataclassMixin): + signature: t.Optional[bytes] = csfield(cs.Const(b"BMP")) width: int = csfield(cs.Int8ub) height: int = csfield(cs.Int8ub) - fmt = constr(Image) + format = DataclassStruct(Image) obj = Image(width=3, height=2) assert ( str(obj) == "Image: \n signature = b'BMP' (total 3)\n width = 3\n height = 2" ) - obj = fmt.parse(fmt.build(obj)) + obj = format.parse(format.build(obj)) assert ( str(obj) == "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: - class Image(DataclassStruct): + @dataclasses.dataclass + class Image(DataclassMixin): width: int = csfield(cs.Int8ub) height: int = csfield(cs.Int8ub) pixels: bytes = csfield(cs.Bytes(cs.this.height * cs.this.width)) common( - constr(Image), + cst.DataclassStruct(Image), b"\x01\x0212", Image(width=1, height=2, pixels=b"12"), ) # check __getattr__ - fmt = Image.__constr__() - assert fmt.width.name == "width" - assert fmt.height.name == "height" - assert fmt.width.subcon is cs.Int8ub - assert fmt.height.subcon is cs.Int8ub + c = cst.DataclassStruct(Image) + assert c.width.name == "width" + assert c.height.name == "height" + assert c.width.subcon is cs.Int8ub + assert c.height.subcon is cs.Int8ub def test_dataclass_struct_reverse() -> None: - class TestDataclass(DataclassStruct, reverse_fields=True): + @dataclasses.dataclass + class TestContainer(DataclassMixin): a: int = csfield(cs.Int16ub) b: int = csfield(cs.Int8ub) common( - constr(TestDataclass), + DataclassStruct(TestContainer, reverse=True), b"\x02\x00\x01", - TestDataclass(a=1, b=2), + TestContainer(a=1, b=2), 3, ) + normal = DataclassStruct(TestContainer) + reverse = DataclassStruct(TestContainer, reverse=True) + assert str(normal.parse(b"\x00\x01\x02")) == str(reverse.parse(b"\x02\x00\x01")) def test_dataclass_struct_nested() -> None: - class TestDataclass(DataclassStruct): - class InnerDataclass(DataclassStruct): + @dataclasses.dataclass + class TestContainer(DataclassMixin): + @dataclasses.dataclass + class InnerDataclass(DataclassMixin): b: int = csfield(cs.Byte) c: bytes = csfield(cs.Bytes(cs.this._.length)) length: int = csfield(cs.Byte) - a: InnerDataclass = csfield(constr(InnerDataclass)) + a: InnerDataclass = csfield(DataclassStruct(InnerDataclass)) common( - constr(TestDataclass), + DataclassStruct(TestContainer), b"\x02\x01\xF1\xF2", - TestDataclass(length=2, a=TestDataclass.InnerDataclass(b=1, c=b"\xF1\xF2")), + TestContainer(length=2, a=TestContainer.InnerDataclass(b=1, c=b"\xF1\xF2")), ) def test_dataclass_struct_default_field() -> None: - class Image(DataclassStruct): + @dataclasses.dataclass + class Image(DataclassMixin): width: int = csfield(cs.Int8ub) height: int = csfield(cs.Int8ub) pixels: t.Optional[bytes] = csfield( @@ -142,75 +157,80 @@ def test_dataclass_struct_default_field() -> None: ) common( - constr(Image), + DataclassStruct(Image), b"\x02\x03\x00\x00\x00\x00\x00\x00", - setattrs(Image(width=2, height=3), pixels=bytes(6)), - sample_building=Image(width=2, height=3), + setattrs(Image(2, 3), pixels=bytes(6)), + sample_building=Image(2, 3), ) def test_dataclass_struct_const_field() -> None: - class TestDataclass(DataclassStruct): + @dataclasses.dataclass + class TestContainer(DataclassMixin): const_field: t.Optional[bytes] = csfield(cs.Const(b"\x00")) common( - constr(TestDataclass), + DataclassStruct(TestContainer), bytes(1), - setattrs(TestDataclass(), const_field=b"\x00"), + setattrs(TestContainer(), const_field=b"\x00"), 1, ) assert ( raises( - constr(TestDataclass).build, - setattrs(TestDataclass(), const_field=b"\x01"), + DataclassStruct(TestContainer).build, + setattrs(TestContainer(), const_field=b"\x01"), ) == cs.ConstError ) def test_dataclass_struct_array_field() -> None: - class TestDataclass(DataclassStruct): + @dataclasses.dataclass + class TestContainer(DataclassMixin): array_field: t.List[int] = csfield(cs.Array(5, cs.Int8ub)) common( - constr(TestDataclass), + DataclassStruct(TestContainer), bytes(5), - TestDataclass(array_field=[0, 0, 0, 0, 0]), + TestContainer(array_field=[0, 0, 0, 0, 0]), 5, ) def test_dataclass_struct_anonymus_fields_1() -> None: - class TestDataclass(DataclassStruct): + @dataclasses.dataclass + class TestContainer(DataclassMixin): _1: t.Optional[bytes] = csfield(cs.Const(b"\x00")) _2: None = csfield(cs.Padding(1)) _3: None = csfield(cs.Pass) _4: None = csfield(cs.Terminated) common( - constr(TestDataclass), + DataclassStruct(TestContainer), bytes(2), - setattrs(TestDataclass(), _1=b"\x00"), + setattrs(TestContainer(), _1=b"\x00"), cs.SizeofError, ) def test_dataclass_struct_anonymus_fields_2() -> None: - class TestDataclass(DataclassStruct): - _1: t.Optional[int] = csfield(cs.Computed(7)) + @dataclasses.dataclass + class TestContainer(DataclassMixin): + _1: int = csfield(cs.Computed(7)) _2: t.Optional[bytes] = csfield(cs.Const(b"JPEG")) _3: None = csfield(cs.Pass) _4: None = csfield(cs.Terminated) - fmt = constr(TestDataclass) - assert fmt.build(TestDataclass()) == fmt.build(TestDataclass()) + d = DataclassStruct(TestContainer) + assert d.build(TestContainer()) == d.build(TestContainer()) def test_dataclass_struct_overloaded_method() -> None: # Test dot access to some names that are not accessable via dot # in the original 'cs.Container'. - class TestDataclass(DataclassStruct): + @dataclasses.dataclass + class TestContainer(DataclassMixin): clear: int = csfield(cs.Int8ul) copy: int = csfield(cs.Int8ul) fromkeys: int = csfield(cs.Int8ul) @@ -226,10 +246,10 @@ def test_dataclass_struct_overloaded_method() -> None: update: int = csfield(cs.Int8ul) values: int = csfield(cs.Int8ul) - fmt = constr(TestDataclass) - obj = fmt.parse( - fmt.build( - TestDataclass( + d = DataclassStruct(TestContainer) + obj = d.parse( + d.build( + TestContainer( clear=1, copy=2, fromkeys=3, @@ -263,172 +283,294 @@ def test_dataclass_struct_overloaded_method() -> None: assert obj.values == 14 +def test_dataclass_struct_no_dataclass() -> None: + class TestContainer(DataclassMixin): + a: int = csfield(cs.Int16ub) + b: int = csfield(cs.Int8ub) + + assert raises(lambda: DataclassStruct(TestContainer)) == TypeError + + +def test_dataclass_struct_no_DataclassMixin() -> None: + @dataclasses.dataclass + class TestContainer: + a: int = csfield(cs.Int16ub) + b: int = csfield(cs.Int8ub) + + cls = t.cast(t.Type[DataclassMixin], TestContainer) + assert raises(lambda: DataclassStruct(cls)) == TypeError + + def test_dataclass_struct_wrong_container() -> None: - class TestContainer1(DataclassStruct): + @dataclasses.dataclass + class TestContainer1(DataclassMixin): a: int = csfield(cs.Int16ub) b: int = csfield(cs.Int8ub) - class TestContainer2(DataclassStruct): + @dataclasses.dataclass + class TestContainer2(DataclassMixin): a: int = csfield(cs.Int16ub) b: int = csfield(cs.Int8ub) - assert raises(constr(TestContainer1).build, TestContainer2(a=1, b=2)) == TypeError + assert ( + raises(DataclassStruct(TestContainer1).build, TestContainer2(a=1, b=2)) + == TypeError + ) def test_dataclass_struct_doc() -> None: - class TestDataclass1(DataclassStruct): - """ - Documentation of TestDataclass1 - """ - - 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") + @dataclasses.dataclass + class TestContainer(DataclassMixin): + a: int = csfield(cs.Int16ub, "This is the documentation of a") + b: int = csfield( + cs.Int8ub, doc="This is the documentation of b\nwhich is multiline" + ) c: int = csfield( cs.Int8ub, - doc=""" - This is the doc of c + """ + This is the documentation of c which is also multiline """, ) - fmt1 = TestDataclass1.__constr__() - common(fmt1, b"\x00\x01\x02\x03", TestDataclass1(a=1, b=2, c=3), 4) + format = DataclassStruct(TestContainer) + common(format, b"\x00\x01\x02\x03", TestContainer(a=1, b=2, c=3), 4) - assert fmt1.docs == "Documentation of TestDataclass1" - assert fmt1.subcon.a.docs == "This is the doc of a" - assert fmt1.subcon.b.docs == "This is the doc of b\nwhich is multiline" - assert fmt1.subcon.c.docs == "This is the doc of c\nwhich is also multiline" - - class 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_bitwise() -> None: - class TestDataclass(DataclassStruct, constr=lambda cls: cs.Bitwise(cls)): - a: int = csfield(cs.BitsInteger(7)) - b: int = csfield(cs.Bit) - c: int = csfield(cs.BitsInteger(8)) - - common( - constr(TestDataclass), - b"\xFD\x12", - TestDataclass(a=0x7E, b=1, c=0x12), - 2, + assert format.subcon.a.docs == "This is the documentation of a" + assert format.subcon.b.docs == "This is the documentation of b\nwhich is multiline" + assert ( + format.subcon.c.docs + == "This is the documentation of c\nwhich is also multiline" ) - # 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_dataclass_bitstruct() -> None: - class TestDataclass(DataclassBitStruct): + @dataclasses.dataclass + class TestContainer(DataclassMixin): a: int = csfield(cs.BitsInteger(7)) b: int = csfield(cs.Bit) c: int = csfield(cs.BitsInteger(8)) + print("") + common( - constr(TestDataclass), + DataclassBitStruct(TestContainer), b"\xFD\x12", - TestDataclass(a=0x7E, b=1, c=0x12), + TestContainer(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) + c = DataclassStruct(TestContainer) + assert c.a.name == "a" + assert c.b.name == "b" + assert c.c.name == "c" + assert isinstance(c.a.subcon, cs.BitsInteger) + assert c.b.subcon is cs.Bit + assert isinstance(c.c.subcon, cs.BitsInteger) def test_tenum() -> None: - class TestEnum(TEnum, subcon=cs.Byte): + class TestEnum(cst.EnumBase): one = 1 two = 2 four = 4 eight = 8 - fmt = constr(TestEnum) + d = cst.TEnum(cs.Byte, TestEnum) - common(fmt, b"\x01", TestEnum.one, 1) - common(fmt, b"\xff", TestEnum(255), 1) - assert fmt.parse(b"\x01") == TestEnum.one - assert fmt.parse(b"\x01") == 1 - assert int(fmt.parse(b"\x01")) == 1 - assert fmt.parse(b"\xff") == TestEnum(255) - assert fmt.parse(b"\xff") == 255 - assert int(fmt.parse(b"\xff")) == 255 - assert raises(fmt.build, 8) == TypeError + common(d, b"\x01", TestEnum.one, 1) + common(d, b"\xff", TestEnum(255), 1) + assert d.parse(b"\x01") == TestEnum.one + assert d.parse(b"\x01") == 1 + assert int(d.parse(b"\x01")) == 1 + assert d.parse(b"\xff") == TestEnum(255) + assert d.parse(b"\xff") == 255 + assert int(d.parse(b"\xff")) == 255 + assert raises(d.build, 8) == TypeError -def test_tenum_doc() -> None: - class TestEnum1(TEnum, subcon=cs.Byte): - """ - TestEnum documentation - """ - - one = 1 - - d1 = constr(TestEnum1) - assert d1.docs == "TestEnum documentation" - - class TestEnum2(TEnum, subcon=cs.Byte): - two = 2 - - d2 = constr(TestEnum2) - assert d2.docs == "" - - -def test_tenum_in_dataclass_struct() -> None: - class TestEnum(TEnum, subcon=cs.Int8ub): +def test_tenum_no_enumbase() -> None: + class E(enum.Enum): a = 1 b = 2 - class TestDataclass(DataclassStruct): - a: TestEnum = csfield(constr(TestEnum)) + 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 + 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_docstring() -> None: + class TestEnum(cst.EnumBase): + """ + 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 + b = 2 + + class E2(cst.EnumBase): + a = 1 + b = 2 + + assert raises(cst.TEnum(cs.Byte, E1).build, E2.a) == TypeError + + +def test_tenum_in_tstruct() -> None: + class TestEnum(cst.EnumBase): + a = 1 + b = 2 + + @dataclasses.dataclass + class TestContainer(DataclassMixin): + a: TestEnum = csfield(cst.TEnum(cs.Int8ub, TestEnum)) b: int = csfield(cs.Int8ub) common( - constr(TestDataclass), + DataclassStruct(TestContainer), b"\x01\x02", - TestDataclass(a=TestEnum.a, b=2), + TestContainer(a=TestEnum.a, b=2), 2, ) assert ( - raises(constr(TestEnum).build, TestDataclass(a=1, b=2)) == TypeError # type: ignore + raises(cst.TEnum(cs.Byte, TestEnum).build, TestContainer(a=1, b=2)) == TypeError # type: ignore ) -def test_tflags() -> None: - class TestFlags(TFlags, subcon=cs.Byte): +def test_tenum_flags() -> None: + class TestEnum(cst.FlagsEnumBase): one = 1 two = 2 four = 4 eight = 8 - fmt = constr(TestFlags) - common(fmt, b"\x03", TestFlags.one | TestFlags.two, 1) - assert fmt.build(TestFlags(0)) == b"\x00" - assert fmt.build(TestFlags.one | TestFlags.two) == b"\x03" - assert fmt.build(TestFlags(8)) == b"\x08" - assert fmt.build(TestFlags(1 | 2)) == b"\x03" - assert fmt.build(TestFlags(255)) == b"\xff" - assert fmt.build(TestFlags.eight) == b"\x08" - assert raises(fmt.build, 2) == TypeError + d = cst.TFlagsEnum(cs.Byte, TestEnum) + common(d, b"\x03", TestEnum.one | TestEnum.two, 1) + assert d.build(TestEnum(0)) == b"\x00" + assert d.build(TestEnum.one | TestEnum.two) == b"\x03" + assert d.build(TestEnum(8)) == b"\x08" + assert d.build(TestEnum(1 | 2)) == b"\x03" + assert d.build(TestEnum(255)) == b"\xff" + assert d.build(TestEnum.eight) == b"\x08" + assert raises(d.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"