Improved typing
This commit is contained in:
@@ -1,5 +1,7 @@
|
|||||||
import sys
|
import sys
|
||||||
import struct
|
import struct
|
||||||
|
import plugin
|
||||||
|
import general
|
||||||
|
|
||||||
class PulseWire:
|
class PulseWire:
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
@@ -31,3 +33,28 @@ class PulseWire:
|
|||||||
raise EOFError("Unexpected EOF while reading payload")
|
raise EOFError("Unexpected EOF while reading payload")
|
||||||
|
|
||||||
self.on_raw_engine_request(buffer)
|
self.on_raw_engine_request(buffer)
|
||||||
|
|
||||||
|
class Strategy(PulseWire):
|
||||||
|
def on_request(self, req: plugin.StrategyEngineMessage):
|
||||||
|
pass
|
||||||
|
|
||||||
|
def send(self, req: plugin.StrategyMessage):
|
||||||
|
self.send_raw(req.to_com())
|
||||||
|
|
||||||
|
def log(self, log: general.EventLog):
|
||||||
|
msg = plugin.StrategyMessage.Log()
|
||||||
|
msg.log = log
|
||||||
|
self.send(msg)
|
||||||
|
|
||||||
|
def getWatchList(self):
|
||||||
|
self.send(plugin.StrategyMessage.GetWatchList())
|
||||||
|
|
||||||
|
def on_raw_engine_request(self, data: bytes):
|
||||||
|
com = plugin.StrategyEngineMessage.from_com(data)
|
||||||
|
|
||||||
|
|
||||||
|
def test(req: plugin.StrategyMessage.type):
|
||||||
|
print(req, plugin.StrategyMessage.type)
|
||||||
|
pass
|
||||||
|
|
||||||
|
test(plugin.StrategyMessage.GetWatchList())
|
||||||
|
|||||||
+22
-29
@@ -1,8 +1,13 @@
|
|||||||
import struct
|
import struct
|
||||||
import typing
|
import typing
|
||||||
from typing import get_type_hints, get_origin, get_args
|
from typing import Protocol, Type, TypeVar, Union, get_type_hints, get_origin, get_args
|
||||||
from adv_types import *
|
from adv_types import *
|
||||||
|
|
||||||
|
# A Protocol representing the instance methods added by @pwp
|
||||||
|
class PulseWire(Protocol):
|
||||||
|
def to_com(self) -> bytes: ...
|
||||||
|
def from_com(self, data: bytes, offset: int = 0) -> tuple[typing.Any, int]: ...
|
||||||
|
|
||||||
class PWPError(Exception):
|
class PWPError(Exception):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
@@ -148,81 +153,69 @@ def decode_value(data: bytes, offset: int, typ):
|
|||||||
|
|
||||||
raise PWPError(f"Unsupported type: {typ}")
|
raise PWPError(f"Unsupported type: {typ}")
|
||||||
|
|
||||||
|
T = TypeVar("T", bound=Type[typing.Any])
|
||||||
def pwp(cls):
|
def pwp(cls: T) -> T:
|
||||||
fields = get_type_hints(cls)
|
fields = get_type_hints(cls)
|
||||||
|
|
||||||
def to_com(self):
|
def to_com(self) -> bytes:
|
||||||
output = b""
|
output = b""
|
||||||
|
|
||||||
for name, typ in fields.items():
|
for name, typ in fields.items():
|
||||||
output += encode_value(getattr(self, name), typ)
|
output += encode_value(getattr(self, name), typ)
|
||||||
|
|
||||||
return output
|
return output
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_com(cls, data):
|
def from_com(cls, data: bytes, offset: int = 0) -> tuple[typing.Any, int]:
|
||||||
values = {}
|
values = {}
|
||||||
offset = 0
|
|
||||||
|
|
||||||
obj = cls.__new__(cls)
|
obj = cls.__new__(cls)
|
||||||
|
|
||||||
for name, typ in fields.items():
|
for name, typ in fields.items():
|
||||||
value, offset = decode_value(data, offset, typ)
|
value, offset = decode_value(data, offset, typ)
|
||||||
setattr(obj, name, value)
|
setattr(obj, name, value)
|
||||||
|
|
||||||
return obj
|
return obj, offset
|
||||||
|
|
||||||
cls.to_com = to_com
|
cls.to_com = to_com
|
||||||
cls.from_com = from_com
|
cls.from_com = from_com
|
||||||
|
|
||||||
return cls
|
return typing.cast(T, cls)
|
||||||
|
|
||||||
def pwp_enum(cls):
|
def pwp_enum(cls: T) -> T:
|
||||||
variants = {}
|
variants = {}
|
||||||
|
|
||||||
index = 0
|
index = 0
|
||||||
|
|
||||||
for idx, (name, value) in enumerate(cls.__dict__.items()):
|
for idx, (name, value) in enumerate(cls.__dict__.items()):
|
||||||
if isinstance(value, type) and value.__module__ == cls.__module__:
|
if isinstance(value, type) and value.__module__ == cls.__module__:
|
||||||
value.__annotations__ = {
|
value.__annotations__ = {
|
||||||
"_id": u8,
|
"_id": u8,
|
||||||
**getattr(value, "__annotations__", {})
|
**getattr(value, "__annotations__", {})
|
||||||
}
|
}
|
||||||
|
|
||||||
value._id = u8(index)
|
value._id = u8(index)
|
||||||
|
|
||||||
value = pwp(value)
|
value = pwp(value)
|
||||||
|
variants[index] = value # Safely align variant registration index
|
||||||
variants[idx] = value
|
|
||||||
setattr(cls, name, value)
|
setattr(cls, name, value)
|
||||||
|
|
||||||
index += 1
|
index += 1
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_com(enum_cls, data):
|
def from_com(enum_cls, data: bytes, offset: int = 0) -> tuple[typing.Any, int]:
|
||||||
enum_id = data[0]
|
enum_id = data[offset]
|
||||||
|
offset += 1
|
||||||
|
|
||||||
if enum_id not in variants:
|
if enum_id not in variants:
|
||||||
raise PWPError(f"Unknown enum id: {enum_id}")
|
raise PWPError(f"Unknown enum id: {enum_id}")
|
||||||
|
|
||||||
variant = variants[enum_id]
|
variant = variants[enum_id]
|
||||||
|
|
||||||
# Decode the rest of the data
|
|
||||||
obj = variant.__new__(variant)
|
obj = variant.__new__(variant)
|
||||||
|
|
||||||
offset = 1 # skip enum id byte
|
|
||||||
|
|
||||||
fields = get_type_hints(variant)
|
fields = get_type_hints(variant)
|
||||||
|
|
||||||
for name, typ in fields.items():
|
for name, typ in fields.items():
|
||||||
if name == "_id":
|
if name == "_id":
|
||||||
continue
|
continue
|
||||||
|
|
||||||
value, offset = decode_value(data, offset, typ)
|
value, offset = decode_value(data, offset, typ)
|
||||||
setattr(obj, name, value)
|
setattr(obj, name, value)
|
||||||
|
|
||||||
return obj
|
return obj, offset
|
||||||
|
|
||||||
cls.from_com = from_com
|
cls.from_com = from_com # type: ignore
|
||||||
|
cls.type = Union[tuple(variants.values())] # type: ignore
|
||||||
|
|
||||||
return cls
|
return typing.cast(T, cls)
|
||||||
Reference in New Issue
Block a user