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):
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)
+36 -17
View File
@@ -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("<B", int)
class u8: pass
class u8(PulseUnit): pass
@pwp_number("<H", int)
class u16: pass
class u16(PulseUnit): pass
@pwp_number("<I", int)
class u32: pass
class u32(PulseUnit): pass
@pwp_number("<Q", int)
class u64: pass
class u64(PulseUnit): pass
# Signed integers
@pwp_number("<b", int)
class i8: pass
class i8(PulseUnit): pass
@pwp_number("<h", int)
class i16: pass
class i16(PulseUnit): pass
@pwp_number("<i", int)
class i32: pass
class i32(PulseUnit): pass
@pwp_number("<q", int)
class i64: pass
class i64(PulseUnit): pass
# Floating point
@pwp_number("<f", float)
class f32: pass
class f32(PulseUnit): pass
@pwp_number("<d", float)
class f64: pass
class f64(PulseUnit): pass
class decimal:
class decimal(PulseUnit):
size = struct.calcsize("<qb")
def __init__(self, value):
@@ -75,10 +93,11 @@ class decimal:
if sign:
integer = -integer
return struct.pack("<qb", integer, scale)
@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)
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 *
import units
@pwp_enum
class MarketTrend:
class MarketTrend(PulseWire):
class Bullish: pass
class Bearish: pass
class Neutral: pass
@pwp_enum
class LogKind:
class LogKind(PulseWire):
class Info: pass
class Warn: pass
class Err: pass
class Debug: pass
@pwp_enum
class CandleInterval:
class CandleInterval(PulseWire):
class OneMinute: pass
class ThreeMinutes: pass
class FiveMinutes: pass
@@ -33,7 +33,7 @@ class CandleInterval:
class OneMonth: pass
@pwp
class Signal:
class Signal(PulseWire):
symbol: str
kind: units.Direction
confidence: f32
@@ -43,27 +43,27 @@ class Signal:
stop_loss: units.USD
@pwp
class EventLog:
class EventLog(PulseWire):
kind: LogKind
name: str
message: str
@pwp
class Position:
class Position(PulseWire):
symbol: units.Symbol
size: f64
entry_price: units.USD
profit: units.USD
@pwp
class MarketItem:
class MarketItem(PulseWire):
symbol: units.Symbol
price: units.USD
trend: f64
volume_24h: units.USD
@pwp
class Candle:
class Candle(PulseWire):
open_time: u64
close_time: u64
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 *
import units
import general
from typing import Optional
from typing import Optional, TypeAlias
@pwp
class StrategySignal:
class StrategySignal(PulseWire):
symbol: str
side: units.Direction
confidence: f32
price: Optional[f64]
@pwp_enum
class StrategyMessage:
class Log:
class StrategyMessage(PulseWire):
class Log(PulseWire):
log: general.EventLog
class GetWatchList: pass
class GetWatchList(PulseWire): pass
class RequestCandlestick:
class RequestCandlestick(PulseWire):
symbol: str
interval: u8
count: u32
class Subscribe:
class Subscribe(PulseWire):
subscription: u8
class Unsubscribe:
class Unsubscribe(PulseWire):
subscription: u8
class UnsubscribeAll: pass
class UnsubscribeAll(PulseWire): pass
class Signal:
class Signal(PulseWire):
signal: StrategySignal
@pwp_enum
class RiskMessage:
class Log:
class RiskMessage(PulseWire):
class Log(PulseWire):
log: general.EventLog
class GetWatchList: pass
class Approve:
class Approve(PulseWire):
signal: general.Signal
class Reject:
class Reject(PulseWire):
reason: str
@pwp_enum
class StrategyEngineMessage:
class Initialize: pass
class StrategyEngineMessage(PulseWire):
class Initialize(PulseWire): pass
class WatchList:
class WatchList(PulseWire):
watchlist: list[general.MarketItem]
class Command:
class Command(PulseWire):
command: str
args: list[str]
class CandleUpdate:
class CandleUpdate(PulseWire):
symbol: str
interval: general.CandleInterval
candle: general.Candle
class CandleStick:
class CandleStick(PulseWire):
symbol: str
interval: general.CandleInterval
candles: list[general.Candle]
@pwp_enum
class RiskEngineMessage:
class Initialize: pass
class RiskEngineMessage(PulseWire):
class Initialize(PulseWire): pass
class WatchList:
class WatchList(PulseWire):
watchlist: list[general.MarketItem]
class Command:
class Command(PulseWire):
command: str
args: list[str]
class Signal:
signal: general.Signal
class Signal(PulseWire):
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 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 *
# 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]: ...
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)
return cls