Files
2026-07-28 19:44:58 +02:00

312 lines
6.9 KiB
Python

import struct
import typing
import types
from typing import (
Protocol,
TypeVar,
Any,
get_type_hints,
get_origin,
get_args,
)
from .adv_types import *
class PulseWire(Protocol):
def to_com(self) -> bytes:
return bytes()
@classmethod
def from_com(cls, data: bytes, offset: int = 0) -> tuple[typing.Any, int]:
return (None, offset)
class PWPError(Exception):
pass
def ensure_size(data: bytes, offset: int, size: int):
if offset + size > len(data):
raise PWPError(
f"Buffer underflow: need {size} bytes at offset {offset}, "
f"only {len(data) - offset} available"
)
def encode_value(value, typ) -> bytes:
origin = get_origin(typ)
# Optional[T]
if origin in (typing.Union, types.UnionType):
args = get_args(typ)
if type(None) in args:
if value is None:
return b"\x00"
real = next(x for x in args if x is not type(None))
return b"\x01" + encode_value(value, real)
# list[T] or list
if origin is list or typ is list:
args = get_args(typ)
item_type = args[0] if args else Any
out = len(value).to_bytes(4, "little")
for item in value:
# If item_type is Any, infer type dynamically from the item itself
actual_type = item_type if item_type is not Any else type(item)
out += encode_value(item, actual_type)
return out
# tuple[T, ...]
if origin is tuple or typ is tuple:
args = get_args(typ)
out = len(value).to_bytes(4, "little")
for item, item_type in zip(value, args if args else [type(x) for x in value]):
out += encode_value(item, item_type)
return out
# dict[K, V]
if origin is dict or typ is dict:
args = get_args(typ)
key_type, value_type = args if args else (Any, Any)
out = len(value).to_bytes(4, "little")
for k, v in value.items():
out += encode_value(k, key_type if key_type is not Any else type(k))
out += encode_value(v, value_type if value_type is not Any else type(v))
return out
# Primitives
if typ is bool:
return b"\x01" if value else b"\x00"
if typ is int:
return struct.pack("<q", value)
if typ is float:
return struct.pack("<d", value)
if typ is str:
raw = value.encode()
return len(raw).to_bytes(4, "little") + raw
if typ is bytes:
return len(value).to_bytes(4, "little") + value
# Custom PWP
if hasattr(value, "to_com"):
return value.to_com()
raise PWPError(f"Cannot encode type {typ}")
def decode_value(data: bytes, offset: int, typ):
origin = get_origin(typ)
# Optional
if origin in (typing.Union, types.UnionType):
args = get_args(typ)
if type(None) in args:
ensure_size(data, offset, 1)
present = data[offset]
offset += 1
if present == 0:
return None, offset
real = next(x for x in args if x is not type(None))
return decode_value(data, offset, real)
# list
if origin is list or typ is list:
args = get_args(typ)
item_type = args[0] if args else Any
ensure_size(data, offset, 4)
count = int.from_bytes(data[offset:offset+4], "little")
offset += 4
result = []
for _ in range(count):
item, offset = decode_value(data, offset, item_type)
result.append(item)
return result, offset
# Primitives & Custom fallbacks...
if typ is bool:
ensure_size(data, offset, 1)
return data[offset] != 0, offset + 1
if typ is int:
ensure_size(data, offset, 8)
return struct.unpack_from("<q", data, offset)[0], offset + 8
if typ is float:
ensure_size(data, offset, 8)
return struct.unpack_from("<d", data, offset)[0], offset + 8
if typ is str:
ensure_size(data, offset, 4)
length = int.from_bytes(data[offset:offset+4], "little")
offset += 4
ensure_size(data, offset, length)
return data[offset:offset+length].decode(), offset + length
if typ is bytes:
ensure_size(data, offset, 4)
length = int.from_bytes(data[offset:offset+4], "little")
offset += 4
ensure_size(data, offset, length)
return data[offset:offset+length], offset + length
if hasattr(typ, "from_com"):
return typ.from_com(data, offset)
raise PWPError(f"Cannot decode type {typ}")
T = TypeVar(
"T",
bound=type[PulseWire]
)
def pwp(cls: T) -> T:
fields = get_type_hints(cls)
def to_com(self):
out = b""
for name, typ in fields.items():
if not hasattr(self, name):
raise PWPError(
f"Missing field {name}"
)
out += encode_value(
getattr(self, name),
typ
)
return out
@classmethod
def from_com(
cls,
data,
offset=0
):
obj = cls.__new__(cls)
for name, typ in fields.items():
value, offset = decode_value(
data,
offset,
typ
)
setattr(
obj,
name,
value
)
return obj, offset
cls.to_com = to_com
cls.from_com = from_com
return cls
def pwp_enum(cls: T) -> T:
variants = {}
index = 0
for name, variant in list(cls.__dict__.items()):
if (
isinstance(variant, type)
and variant.__module__ == cls.__module__
):
variant.__annotations__ = {
"_id": u8,
**getattr(
variant,
"__annotations__",
{}
)
}
variant._id = u8(index)
variant = pwp(variant)
variants[index] = variant
setattr(
cls,
name,
variant
)
index += 1
def to_com(self):
for idx, variant in variants.items():
if isinstance(self, variant):
return (
idx.to_bytes(1, "little")
+ self.to_com()
)
raise PWPError(
f"Unknown enum variant {type(self)}"
)
@classmethod
def from_com(
enum_cls,
data,
offset=0
):
ensure_size(
data,
offset,
1
)
enum_id = data[offset]
# offset += 1
if enum_id not in variants:
raise PWPError(
f"Unknown enum id {enum_id}"
)
return variants[enum_id].from_com(
data,
offset
)
cls.to_com = to_com
cls.from_com = from_com
return cls