Library structure
This commit is contained in:
@@ -0,0 +1,504 @@
|
||||
import struct
|
||||
import typing
|
||||
import types
|
||||
from typing import (
|
||||
Protocol,
|
||||
TypeVar,
|
||||
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]
|
||||
if origin is list:
|
||||
item_type = get_args(typ)[0]
|
||||
|
||||
out = len(value).to_bytes(4, "little")
|
||||
|
||||
for item in value:
|
||||
out += encode_value(item, item_type)
|
||||
|
||||
return out
|
||||
|
||||
# tuple[T]
|
||||
if origin is tuple:
|
||||
args = get_args(typ)
|
||||
|
||||
out = len(value).to_bytes(4, "little")
|
||||
|
||||
for item, item_type in zip(value, args):
|
||||
out += encode_value(item, item_type)
|
||||
|
||||
return out
|
||||
|
||||
# dict[K,V]
|
||||
if origin is dict:
|
||||
key_type, value_type = get_args(typ)
|
||||
|
||||
out = len(value).to_bytes(4, "little")
|
||||
|
||||
for k, v in value.items():
|
||||
out += encode_value(k, key_type)
|
||||
out += encode_value(v, value_type)
|
||||
|
||||
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:
|
||||
|
||||
item_type = get_args(typ)[0]
|
||||
|
||||
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
|
||||
|
||||
|
||||
|
||||
# tuple
|
||||
|
||||
if origin is tuple:
|
||||
|
||||
types_ = get_args(typ)
|
||||
|
||||
ensure_size(data, offset, 4)
|
||||
|
||||
count = int.from_bytes(
|
||||
data[offset:offset+4],
|
||||
"little"
|
||||
)
|
||||
|
||||
offset += 4
|
||||
|
||||
result = []
|
||||
|
||||
for i in range(count):
|
||||
item, offset = decode_value(
|
||||
data,
|
||||
offset,
|
||||
types_[i]
|
||||
)
|
||||
|
||||
result.append(item)
|
||||
|
||||
return tuple(result), offset
|
||||
|
||||
|
||||
|
||||
# dict
|
||||
|
||||
if origin is dict:
|
||||
|
||||
key_type, value_type = get_args(typ)
|
||||
|
||||
ensure_size(data, offset, 4)
|
||||
|
||||
count = int.from_bytes(
|
||||
data[offset:offset+4],
|
||||
"little"
|
||||
)
|
||||
|
||||
offset += 4
|
||||
|
||||
result = {}
|
||||
|
||||
for _ in range(count):
|
||||
|
||||
key, offset = decode_value(
|
||||
data,
|
||||
offset,
|
||||
key_type
|
||||
)
|
||||
|
||||
value, offset = decode_value(
|
||||
data,
|
||||
offset,
|
||||
value_type
|
||||
)
|
||||
|
||||
result[key] = value
|
||||
|
||||
return result, offset
|
||||
|
||||
|
||||
|
||||
# primitives
|
||||
|
||||
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
|
||||
Reference in New Issue
Block a user