From aea78778f6c9e87c22fd57e342acd8157e5f100b Mon Sep 17 00:00:00 2001 From: Klesti Selimaj Date: Tue, 28 Jul 2026 15:39:35 +0200 Subject: [PATCH] Improved typing --- src/__init__.py | 11 ++---- src/adv_types.py | 53 +++++++++++++++++++---------- src/general.py | 18 +++++----- src/plugin.py | 87 +++++++++++++++++++++++++++++++++--------------- src/pwp.py | 25 +++++++------- 5 files changed, 119 insertions(+), 75 deletions(-) diff --git a/src/__init__.py b/src/__init__.py index 55cd644..bb8ad77 100644 --- a/src/__init__.py +++ b/src/__init__.py @@ -38,7 +38,7 @@ class Strategy(PulseWire): def on_request(self, req: plugin.StrategyEngineMessage): pass - def send(self, req: plugin.StrategyMessage): + def send(self, req: plugin.StrategyMessageType): self.send_raw(req.to_com()) def log(self, log: general.EventLog): @@ -50,11 +50,4 @@ class Strategy(PulseWire): 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()) + com = plugin.StrategyEngineMessage.from_com(data=data) diff --git a/src/adv_types.py b/src/adv_types.py index 60c24d3..39d6c9e 100644 --- a/src/adv_types.py +++ b/src/adv_types.py @@ -1,13 +1,32 @@ import struct from decimal import Decimal +import typing +T = typing.TypeVar("T") +class PulseUnit(typing.Generic[T]): + value: T + + def __init__(self, value: T): + self.value = value + + def __class_getitem__(cls, item): + new_cls = type( + f"{cls.__name__}[{item.__name__}]", + (cls,), + {"_type": item} + ) + return new_cls + + def to_com(self) -> bytes: ... + + @classmethod + def from_com(cls, data: bytes, offset: int = 0) -> tuple[typing.Any, int]: ... + +B = typing.TypeVar("B", bound=type[PulseUnit]) def pwp_number(fmt, value_type): size = struct.calcsize(fmt) - def decorator(cls): - def __init__(self, value): - self.value = value_type(value) - + def decorator(cls: B) -> B: def to_com(self): return struct.pack(fmt, self.value) @@ -16,7 +35,6 @@ def pwp_number(fmt, value_type): value = struct.unpack_from(fmt, data, offset)[0] return cls(value), offset + size - cls.__init__ = __init__ cls.to_com = to_com cls.from_com = from_com @@ -26,38 +44,38 @@ def pwp_number(fmt, value_type): # Unsigned integers @pwp_number(" tuple[typing.Any, int]: integer, scale = struct.unpack_from(" bytes: ... - def from_com(self, data: bytes, offset: int = 0) -> tuple[typing.Any, int]: ... + 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 @@ -153,7 +156,7 @@ def decode_value(data: bytes, offset: int, typ): raise PWPError(f"Unsupported type: {typ}") -T = TypeVar("T", bound=Type[typing.Any]) +T = TypeVar("T", bound=type[PulseWire]) def pwp(cls: T) -> T: fields = get_type_hints(cls) @@ -177,18 +180,15 @@ def pwp(cls: T) -> T: cls.to_com = to_com cls.from_com = from_com - return typing.cast(T, cls) + return cls def pwp_enum(cls: T) -> T: variants = {} index = 0 - for idx, (name, value) in enumerate(cls.__dict__.items()): + for name, value in cls.__dict__.items(): if isinstance(value, type) and value.__module__ == cls.__module__: - value.__annotations__ = { - "_id": u8, - **getattr(value, "__annotations__", {}) - } + value.__annotations__["_id"] = u8 value._id = u8(index) value = pwp(value) variants[index] = value # Safely align variant registration index @@ -215,7 +215,6 @@ def pwp_enum(cls: T) -> T: return obj, offset - cls.from_com = from_com # type: ignore - cls.type = Union[tuple(variants.values())] # type: ignore + cls.from_com = from_com - return typing.cast(T, cls) \ No newline at end of file + return cls \ No newline at end of file