Improved typing
This commit is contained in:
+2
-9
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
Reference in New Issue
Block a user