Improved typing

This commit is contained in:
2026-07-28 15:39:35 +02:00
parent 20ef213399
commit aea78778f6
5 changed files with 119 additions and 75 deletions
+2 -9
View File
@@ -38,7 +38,7 @@ class Strategy(PulseWire):
def on_request(self, req: plugin.StrategyEngineMessage): def on_request(self, req: plugin.StrategyEngineMessage):
pass pass
def send(self, req: plugin.StrategyMessage): def send(self, req: plugin.StrategyMessageType):
self.send_raw(req.to_com()) self.send_raw(req.to_com())
def log(self, log: general.EventLog): def log(self, log: general.EventLog):
@@ -50,11 +50,4 @@ class Strategy(PulseWire):
self.send(plugin.StrategyMessage.GetWatchList()) self.send(plugin.StrategyMessage.GetWatchList())
def on_raw_engine_request(self, data: bytes): def on_raw_engine_request(self, data: bytes):
com = plugin.StrategyEngineMessage.from_com(data) com = plugin.StrategyEngineMessage.from_com(data=data)
def test(req: plugin.StrategyMessage.type):
print(req, plugin.StrategyMessage.type)
pass
test(plugin.StrategyMessage.GetWatchList())
+36 -17
View File
@@ -1,13 +1,32 @@
import struct import struct
from decimal import Decimal 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): def pwp_number(fmt, value_type):
size = struct.calcsize(fmt) size = struct.calcsize(fmt)
def decorator(cls): def decorator(cls: B) -> B:
def __init__(self, value):
self.value = value_type(value)
def to_com(self): def to_com(self):
return struct.pack(fmt, self.value) return struct.pack(fmt, self.value)
@@ -16,7 +35,6 @@ def pwp_number(fmt, value_type):
value = struct.unpack_from(fmt, data, offset)[0] value = struct.unpack_from(fmt, data, offset)[0]
return cls(value), offset + size return cls(value), offset + size
cls.__init__ = __init__
cls.to_com = to_com cls.to_com = to_com
cls.from_com = from_com cls.from_com = from_com
@@ -26,38 +44,38 @@ def pwp_number(fmt, value_type):
# Unsigned integers # Unsigned integers
@pwp_number("<B", int) @pwp_number("<B", int)
class u8: pass class u8(PulseUnit): pass
@pwp_number("<H", int) @pwp_number("<H", int)
class u16: pass class u16(PulseUnit): pass
@pwp_number("<I", int) @pwp_number("<I", int)
class u32: pass class u32(PulseUnit): pass
@pwp_number("<Q", int) @pwp_number("<Q", int)
class u64: pass class u64(PulseUnit): pass
# Signed integers # Signed integers
@pwp_number("<b", int) @pwp_number("<b", int)
class i8: pass class i8(PulseUnit): pass
@pwp_number("<h", int) @pwp_number("<h", int)
class i16: pass class i16(PulseUnit): pass
@pwp_number("<i", int) @pwp_number("<i", int)
class i32: pass class i32(PulseUnit): pass
@pwp_number("<q", int) @pwp_number("<q", int)
class i64: pass class i64(PulseUnit): pass
# Floating point # Floating point
@pwp_number("<f", float) @pwp_number("<f", float)
class f32: pass class f32(PulseUnit): pass
@pwp_number("<d", float) @pwp_number("<d", float)
class f64: pass class f64(PulseUnit): pass
class decimal: class decimal(PulseUnit):
size = struct.calcsize("<qb") size = struct.calcsize("<qb")
def __init__(self, value): def __init__(self, value):
@@ -75,10 +93,11 @@ class decimal:
if sign: if sign:
integer = -integer integer = -integer
return struct.pack("<qb", integer, scale) return struct.pack("<qb", integer, scale)
@classmethod @classmethod
def from_com(cls, data, offset=0): def from_com(cls, data: bytes, offset: int = 0) -> tuple[typing.Any, int]:
integer, scale = struct.unpack_from("<qb", data, offset) integer, scale = struct.unpack_from("<qb", data, offset)
value = Decimal(integer) / (Decimal(10) ** scale) value = Decimal(integer) / (Decimal(10) ** scale)
+9 -9
View File
@@ -1,22 +1,22 @@
from pwp import pwp, pwp_enum from pwp import pwp, pwp_enum, PulseWire
from adv_types import * from adv_types import *
import units import units
@pwp_enum @pwp_enum
class MarketTrend: class MarketTrend(PulseWire):
class Bullish: pass class Bullish: pass
class Bearish: pass class Bearish: pass
class Neutral: pass class Neutral: pass
@pwp_enum @pwp_enum
class LogKind: class LogKind(PulseWire):
class Info: pass class Info: pass
class Warn: pass class Warn: pass
class Err: pass class Err: pass
class Debug: pass class Debug: pass
@pwp_enum @pwp_enum
class CandleInterval: class CandleInterval(PulseWire):
class OneMinute: pass class OneMinute: pass
class ThreeMinutes: pass class ThreeMinutes: pass
class FiveMinutes: pass class FiveMinutes: pass
@@ -33,7 +33,7 @@ class CandleInterval:
class OneMonth: pass class OneMonth: pass
@pwp @pwp
class Signal: class Signal(PulseWire):
symbol: str symbol: str
kind: units.Direction kind: units.Direction
confidence: f32 confidence: f32
@@ -43,27 +43,27 @@ class Signal:
stop_loss: units.USD stop_loss: units.USD
@pwp @pwp
class EventLog: class EventLog(PulseWire):
kind: LogKind kind: LogKind
name: str name: str
message: str message: str
@pwp @pwp
class Position: class Position(PulseWire):
symbol: units.Symbol symbol: units.Symbol
size: f64 size: f64
entry_price: units.USD entry_price: units.USD
profit: units.USD profit: units.USD
@pwp @pwp
class MarketItem: class MarketItem(PulseWire):
symbol: units.Symbol symbol: units.Symbol
price: units.USD price: units.USD
trend: f64 trend: f64
volume_24h: units.USD volume_24h: units.USD
@pwp @pwp
class Candle: class Candle(PulseWire):
open_time: u64 open_time: u64
close_time: u64 close_time: u64
coin: str coin: str
+60 -27
View File
@@ -1,83 +1,116 @@
from pwp import pwp, pwp_enum from pwp import pwp, pwp_enum, PulseWire
from adv_types import * from adv_types import *
import units import units
import general import general
from typing import Optional from typing import Optional, TypeAlias
@pwp @pwp
class StrategySignal: class StrategySignal(PulseWire):
symbol: str symbol: str
side: units.Direction side: units.Direction
confidence: f32 confidence: f32
price: Optional[f64] price: Optional[f64]
@pwp_enum @pwp_enum
class StrategyMessage: class StrategyMessage(PulseWire):
class Log: class Log(PulseWire):
log: general.EventLog log: general.EventLog
class GetWatchList: pass class GetWatchList(PulseWire): pass
class RequestCandlestick: class RequestCandlestick(PulseWire):
symbol: str symbol: str
interval: u8 interval: u8
count: u32 count: u32
class Subscribe: class Subscribe(PulseWire):
subscription: u8 subscription: u8
class Unsubscribe: class Unsubscribe(PulseWire):
subscription: u8 subscription: u8
class UnsubscribeAll: pass class UnsubscribeAll(PulseWire): pass
class Signal: class Signal(PulseWire):
signal: StrategySignal signal: StrategySignal
@pwp_enum @pwp_enum
class RiskMessage: class RiskMessage(PulseWire):
class Log: class Log(PulseWire):
log: general.EventLog log: general.EventLog
class GetWatchList: pass class GetWatchList: pass
class Approve: class Approve(PulseWire):
signal: general.Signal signal: general.Signal
class Reject: class Reject(PulseWire):
reason: str reason: str
@pwp_enum @pwp_enum
class StrategyEngineMessage: class StrategyEngineMessage(PulseWire):
class Initialize: pass class Initialize(PulseWire): pass
class WatchList: class WatchList(PulseWire):
watchlist: list[general.MarketItem] watchlist: list[general.MarketItem]
class Command: class Command(PulseWire):
command: str command: str
args: list[str] args: list[str]
class CandleUpdate: class CandleUpdate(PulseWire):
symbol: str symbol: str
interval: general.CandleInterval interval: general.CandleInterval
candle: general.Candle candle: general.Candle
class CandleStick: class CandleStick(PulseWire):
symbol: str symbol: str
interval: general.CandleInterval interval: general.CandleInterval
candles: list[general.Candle] candles: list[general.Candle]
@pwp_enum @pwp_enum
class RiskEngineMessage: class RiskEngineMessage(PulseWire):
class Initialize: pass class Initialize(PulseWire): pass
class WatchList: class WatchList(PulseWire):
watchlist: list[general.MarketItem] watchlist: list[general.MarketItem]
class Command: class Command(PulseWire):
command: str command: str
args: list[str] args: list[str]
class Signal: class Signal(PulseWire):
signal: general.Signal signal: general.Signal
StrategyMessageType: TypeAlias = (
StrategyMessage.Log
| StrategyMessage.GetWatchList
| StrategyMessage.RequestCandlestick
| StrategyMessage.Subscribe
| StrategyMessage.Unsubscribe
| StrategyMessage.UnsubscribeAll
| StrategyMessage.Signal
)
RiskMessageType: TypeAlias = (
RiskMessage.Log
| RiskMessage.GetWatchList
| RiskMessage.Approve
| RiskMessage.Reject
)
StrategyEngineMessageType: TypeAlias = (
StrategyEngineMessage.Initialize
| StrategyEngineMessage.WatchList
| StrategyEngineMessage.Command
| StrategyEngineMessage.CandleUpdate
| StrategyEngineMessage.CandleStick
)
RiskEngineMessageType: TypeAlias = (
RiskEngineMessage.Initialize
| RiskEngineMessage.WatchList
| RiskEngineMessage.Command
| RiskEngineMessage.Signal
)
+12 -13
View File
@@ -1,12 +1,15 @@
import struct import struct
import typing import typing
from typing import Protocol, Type, TypeVar, Union, get_type_hints, get_origin, get_args from typing import Protocol, 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 # A Protocol representing the instance methods added by @pwp
class PulseWire(Protocol): class PulseWire(Protocol):
def to_com(self) -> bytes: ... def to_com(self) -> bytes:
def from_com(self, data: bytes, offset: int = 0) -> tuple[typing.Any, int]: ... return bytes()
@classmethod
def from_com(cls, data: bytes, offset: int = 0) -> tuple[typing.Any, int]:
return (None, offset)
class PWPError(Exception): class PWPError(Exception):
pass pass
@@ -153,7 +156,7 @@ 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]) T = TypeVar("T", bound=type[PulseWire])
def pwp(cls: T) -> T: def pwp(cls: T) -> T:
fields = get_type_hints(cls) fields = get_type_hints(cls)
@@ -177,18 +180,15 @@ def pwp(cls: T) -> T:
cls.to_com = to_com cls.to_com = to_com
cls.from_com = from_com cls.from_com = from_com
return typing.cast(T, cls) return cls
def pwp_enum(cls: T) -> T: def pwp_enum(cls: T) -> T:
variants = {} variants = {}
index = 0 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__: if isinstance(value, type) and value.__module__ == cls.__module__:
value.__annotations__ = { value.__annotations__["_id"] = u8
"_id": u8,
**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[index] = value # Safely align variant registration index
@@ -215,7 +215,6 @@ def pwp_enum(cls: T) -> T:
return obj, offset return obj, offset
cls.from_com = from_com # type: ignore cls.from_com = from_com
cls.type = Union[tuple(variants.values())] # type: ignore
return typing.cast(T, cls) return cls