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