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(" 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