Improved typing
This commit is contained in:
+28
-1
@@ -1,5 +1,7 @@
|
||||
import sys
|
||||
import struct
|
||||
import plugin
|
||||
import general
|
||||
|
||||
class PulseWire:
|
||||
def __init__(self):
|
||||
@@ -30,4 +32,29 @@ class PulseWire:
|
||||
if len(buffer) != length:
|
||||
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())
|
||||
|
||||
+23
-30
@@ -1,8 +1,13 @@
|
||||
import struct
|
||||
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 *
|
||||
|
||||
# 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):
|
||||
pass
|
||||
|
||||
@@ -148,81 +153,69 @@ def decode_value(data: bytes, offset: int, typ):
|
||||
|
||||
raise PWPError(f"Unsupported type: {typ}")
|
||||
|
||||
|
||||
def pwp(cls):
|
||||
T = TypeVar("T", bound=Type[typing.Any])
|
||||
def pwp(cls: T) -> T:
|
||||
fields = get_type_hints(cls)
|
||||
|
||||
def to_com(self):
|
||||
def to_com(self) -> bytes:
|
||||
output = b""
|
||||
|
||||
for name, typ in fields.items():
|
||||
output += encode_value(getattr(self, name), typ)
|
||||
|
||||
return output
|
||||
|
||||
@classmethod
|
||||
def from_com(cls, data):
|
||||
def from_com(cls, data: bytes, offset: int = 0) -> tuple[typing.Any, int]:
|
||||
values = {}
|
||||
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
|
||||
return obj, offset
|
||||
|
||||
cls.to_com = to_com
|
||||
cls.to_com = to_com
|
||||
cls.from_com = from_com
|
||||
|
||||
return cls
|
||||
return typing.cast(T, cls)
|
||||
|
||||
def pwp_enum(cls):
|
||||
def pwp_enum(cls: T) -> T:
|
||||
variants = {}
|
||||
|
||||
index = 0
|
||||
|
||||
for idx, (name, value) in enumerate(cls.__dict__.items()):
|
||||
if isinstance(value, type) and value.__module__ == cls.__module__:
|
||||
value.__annotations__ = {
|
||||
"_id": u8,
|
||||
**getattr(value, "__annotations__", {})
|
||||
}
|
||||
|
||||
value._id = u8(index)
|
||||
|
||||
value = pwp(value)
|
||||
|
||||
variants[idx] = value
|
||||
variants[index] = value # Safely align variant registration index
|
||||
setattr(cls, name, value)
|
||||
|
||||
index += 1
|
||||
|
||||
@classmethod
|
||||
def from_com(enum_cls, data):
|
||||
enum_id = data[0]
|
||||
def from_com(enum_cls, data: bytes, offset: int = 0) -> tuple[typing.Any, int]:
|
||||
enum_id = data[offset]
|
||||
offset += 1
|
||||
|
||||
if enum_id not in variants:
|
||||
raise PWPError(f"Unknown enum id: {enum_id}")
|
||||
|
||||
variant = variants[enum_id]
|
||||
|
||||
# Decode the rest of the data
|
||||
obj = variant.__new__(variant)
|
||||
|
||||
offset = 1 # skip enum id byte
|
||||
|
||||
fields = get_type_hints(variant)
|
||||
|
||||
for name, typ in fields.items():
|
||||
if name == "_id":
|
||||
continue
|
||||
|
||||
value, offset = decode_value(data, offset, typ)
|
||||
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