Improved typing

This commit is contained in:
2026-07-28 14:38:53 +02:00
parent 85cad507a4
commit 20ef213399
2 changed files with 51 additions and 31 deletions
+27
View File
@@ -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
View File
@@ -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)