From 20ef21339986bc9b0ae196c8a411552afac88f1f Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Tue, 28 Jul 2026 14:38:53 +0200 Subject: [PATCH] Improved typing --- src/__init__.py | 29 ++++++++++++++++++++++++++- src/pwp.py | 53 +++++++++++++++++++++---------------------------- 2 files changed, 51 insertions(+), 31 deletions(-) diff --git a/src/__init__.py b/src/__init__.py index 1ec6b30..55cd644 100644 --- a/src/__init__.py +++ b/src/__init__.py @@ -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) \ No newline at end of file + 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()) diff --git a/src/pwp.py b/src/pwp.py index 3981724..5fa2f11 100644 --- a/src/pwp.py +++ b/src/pwp.py @@ -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 \ No newline at end of file + return typing.cast(T, cls) \ No newline at end of file