Improved pwp

This commit is contained in:
2026-07-28 15:54:17 +02:00
parent aea78778f6
commit 2b26ce46fc
3 changed files with 408 additions and 80 deletions
+44 -1
View File
@@ -49,5 +49,48 @@ class Strategy(PulseWire):
def getWatchList(self): def getWatchList(self):
self.send(plugin.StrategyMessage.GetWatchList()) self.send(plugin.StrategyMessage.GetWatchList())
def getCandleStick(self, symbol: str, interval: general.CandleInterval, count: general.u32):
msg = plugin.StrategyMessage.RequestCandlestick()
msg.symbol = symbol
msg.interval = interval
msg.count = count
self.send(msg)
def unsubscribeAll(self):
self.send(plugin.StrategyMessage.UnsubscribeAll())
def signal(self, signal: plugin.StrategySignal):
msg = plugin.StrategyMessage.Signal()
msg.signal = signal
self.send(msg)
def on_raw_engine_request(self, data: bytes): def on_raw_engine_request(self, data: bytes):
com = plugin.StrategyEngineMessage.from_com(data=data) self.on_request(plugin.StrategyEngineMessage.from_com(data=data)) #type: ignore
class Risk(PulseWire):
def on_request(self, req: plugin.RiskEngineMessage):
pass
def send(self, req: plugin.RiskMessageType):
self.send_raw(req.to_com())
def log(self, log: general.EventLog):
msg = plugin.RiskMessage.Log()
msg.log = log
self.send(msg)
def getWatchList(self):
self.send(plugin.RiskMessage.GetWatchList())
def approveSignal(self, signal: general.Signal):
msg = plugin.RiskMessage.Approve()
msg.signal = signal
self.send(msg)
def rejectSignal(self, reason: str):
msg = plugin.RiskMessage.Reject()
msg.reason = reason
self.send(msg)
def on_raw_engine_request(self, data: bytes):
com = plugin.RiskEngineMessage.from_com(data=data)
+2 -2
View File
@@ -20,7 +20,7 @@ class StrategyMessage(PulseWire):
class RequestCandlestick(PulseWire): class RequestCandlestick(PulseWire):
symbol: str symbol: str
interval: u8 interval: general.CandleInterval
count: u32 count: u32
class Subscribe(PulseWire): class Subscribe(PulseWire):
@@ -39,7 +39,7 @@ class RiskMessage(PulseWire):
class Log(PulseWire): class Log(PulseWire):
log: general.EventLog log: general.EventLog
class GetWatchList: pass class GetWatchList(PulseWire): pass
class Approve(PulseWire): class Approve(PulseWire):
signal: general.Signal signal: general.Signal
+362 -77
View File
@@ -1,12 +1,22 @@
import struct import struct
import typing import typing
from typing import Protocol, TypeVar, Union, get_type_hints, get_origin, get_args import types
from typing import (
Protocol,
TypeVar,
get_type_hints,
get_origin,
get_args,
)
from adv_types import * from adv_types import *
# 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:
return bytes() return bytes()
@classmethod @classmethod
def from_com(cls, data: bytes, offset: int = 0) -> tuple[typing.Any, int]: def from_com(cls, data: bytes, offset: int = 0) -> tuple[typing.Any, int]:
return (None, offset) return (None, offset)
@@ -14,22 +24,35 @@ class PulseWire(Protocol):
class PWPError(Exception): class PWPError(Exception):
pass pass
def ensure_size(data: bytes, offset: int, size: int):
if offset + size > len(data):
raise PWPError(
f"Buffer underflow: need {size} bytes at offset {offset}, "
f"only {len(data) - offset} available"
)
def encode_value(value, typ) -> bytes: def encode_value(value, typ) -> bytes:
origin = get_origin(typ) origin = get_origin(typ)
# Optional[T] # Optional[T]
if origin is typing.Union: if origin in (typing.Union, types.UnionType):
args = get_args(typ) args = get_args(typ)
if type(None) in args: if type(None) in args:
if value is None: if value is None:
return b"\x00" return b"\x00"
real_type = next(t for t in args if t is not type(None)) real = next(
return b"\x01" + encode_value(value, real_type) x for x in args
if x is not type(None)
)
return b"\x01" + encode_value(value, real)
# list[T] # list[T]
elif origin is list: if origin is list:
item_type = get_args(typ)[0] item_type = get_args(typ)[0]
out = len(value).to_bytes(4, "little") out = len(value).to_bytes(4, "little")
@@ -39,7 +62,18 @@ def encode_value(value, typ) -> bytes:
return out return out
# dict[K, V] # tuple[T]
if origin is tuple:
args = get_args(typ)
out = len(value).to_bytes(4, "little")
for item, item_type in zip(value, args):
out += encode_value(item, item_type)
return out
# dict[K,V]
if origin is dict: if origin is dict:
key_type, value_type = get_args(typ) key_type, value_type = get_args(typ)
@@ -51,13 +85,11 @@ def encode_value(value, typ) -> bytes:
return out return out
# primitives
if typ is str:
data = value.encode("utf-8")
return len(data).to_bytes(4, "little") + data
if typ is bytes: # primitives
return len(value).to_bytes(4, "little") + value
if typ is bool:
return b"\x01" if value else b"\x00"
if typ is int: if typ is int:
return struct.pack("<q", value) return struct.pack("<q", value)
@@ -65,156 +97,409 @@ def encode_value(value, typ) -> bytes:
if typ is float: if typ is float:
return struct.pack("<d", value) return struct.pack("<d", value)
if typ is bool: if typ is str:
return b"\x01" if value else b"\x00" raw = value.encode()
if hasattr(typ, "to_com"): return (
len(raw).to_bytes(4, "little")
+ raw
)
if typ is bytes:
return (
len(value).to_bytes(4, "little")
+ value
)
# custom PWP
if hasattr(value, "to_com"):
return value.to_com() return value.to_com()
raise PWPError(f"Unsupported type: {typ}")
raise PWPError(
f"Cannot encode type {typ}"
)
def decode_value(data: bytes, offset: int, typ): def decode_value(data: bytes, offset: int, typ):
origin = get_origin(typ) origin = get_origin(typ)
# Optional[T]
if origin is typing.Union: # Optional
if origin in (typing.Union, types.UnionType):
args = get_args(typ) args = get_args(typ)
if type(None) in args: if type(None) in args:
ensure_size(data, offset, 1)
present = data[offset] present = data[offset]
offset += 1 offset += 1
if present == 0: if present == 0:
return None, offset return None, offset
real_type = next(t for t in args if t is not type(None)) real = next(
return decode_value(data, offset, real_type) x for x in args
if x is not type(None)
)
return decode_value(
data,
offset,
real
)
# list
# list[T]
if origin is list: if origin is list:
item_type = get_args(typ)[0] item_type = get_args(typ)[0]
count = int.from_bytes(data[offset:offset+4], "little") ensure_size(data, offset, 4)
count = int.from_bytes(
data[offset:offset+4],
"little"
)
offset += 4 offset += 4
result = [] result = []
for _ in range(count): for _ in range(count):
value, offset = decode_value(data, offset, item_type) item, offset = decode_value(
result.append(value) data,
offset,
item_type
)
result.append(item)
return result, offset return result, offset
# dict[K,V]
# tuple
if origin is tuple:
types_ = get_args(typ)
ensure_size(data, offset, 4)
count = int.from_bytes(
data[offset:offset+4],
"little"
)
offset += 4
result = []
for i in range(count):
item, offset = decode_value(
data,
offset,
types_[i]
)
result.append(item)
return tuple(result), offset
# dict
if origin is dict: if origin is dict:
key_type, value_type = get_args(typ) key_type, value_type = get_args(typ)
count = int.from_bytes(data[offset:offset+4], "little") ensure_size(data, offset, 4)
count = int.from_bytes(
data[offset:offset+4],
"little"
)
offset += 4 offset += 4
result = {} result = {}
for _ in range(count): for _ in range(count):
key, offset = decode_value(data, offset, key_type)
value, offset = decode_value(data, offset, value_type) key, offset = decode_value(
data,
offset,
key_type
)
value, offset = decode_value(
data,
offset,
value_type
)
result[key] = value result[key] = value
return result, offset return result, offset
# primitives # primitives
if typ is str:
length = int.from_bytes(data[offset:offset+4], "little") if typ is bool:
offset += 4
ensure_size(data, offset, 1)
return data[offset] != 0, offset + 1
if typ is int:
ensure_size(data, offset, 8)
return ( return (
data[offset:offset+length].decode("utf-8"), struct.unpack_from(
offset + length, "<q",
data,
offset
)[0],
offset + 8
)
if typ is float:
ensure_size(data, offset, 8)
return (
struct.unpack_from(
"<d",
data,
offset
)[0],
offset + 8
)
if typ is str:
ensure_size(data, offset, 4)
length = int.from_bytes(
data[offset:offset+4],
"little"
) )
if typ is bytes:
length = int.from_bytes(data[offset:offset+4], "little")
offset += 4 offset += 4
ensure_size(
data,
offset,
length
)
return (
data[offset:offset+length].decode(),
offset + length
)
if typ is bytes:
ensure_size(data, offset, 4)
length = int.from_bytes(
data[offset:offset+4],
"little"
)
offset += 4
ensure_size(
data,
offset,
length
)
return ( return (
data[offset:offset+length], data[offset:offset+length],
offset + length, offset + length
) )
if typ is int:
return struct.unpack_from("<q", data, offset)[0], offset + 8
if typ is float:
return struct.unpack_from("<d", data, offset)[0], offset + 8
if typ is bool:
return data[offset] != 0, offset + 1
# nested PWP class
if hasattr(typ, "from_com"): if hasattr(typ, "from_com"):
return typ.from_com(data, offset) return typ.from_com(
data,
offset
)
raise PWPError(
f"Cannot decode type {typ}"
)
T = TypeVar(
"T",
bound=type[PulseWire]
)
raise PWPError(f"Unsupported type: {typ}")
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)
def to_com(self) -> bytes:
output = b"" def to_com(self):
out = b""
for name, typ in fields.items(): for name, typ in fields.items():
output += encode_value(getattr(self, name), typ)
return output if not hasattr(self, name):
raise PWPError(
f"Missing field {name}"
)
out += encode_value(
getattr(self, name),
typ
)
return out
@classmethod @classmethod
def from_com(cls, data: bytes, offset: int = 0) -> tuple[typing.Any, int]: def from_com(
values = {} cls,
data,
offset=0
):
obj = cls.__new__(cls) obj = cls.__new__(cls)
for name, typ in fields.items(): for name, typ in fields.items():
value, offset = decode_value(data, offset, typ)
setattr(obj, name, value) value, offset = decode_value(
data,
offset,
typ
)
setattr(
obj,
name,
value
)
return obj, offset return obj, offset
cls.to_com = to_com
cls.to_com = to_com
cls.from_com = from_com cls.from_com = from_com
return cls return cls
def pwp_enum(cls: T) -> T: def pwp_enum(cls: T) -> T:
variants = {} variants = {}
index = 0 index = 0
for name, value in cls.__dict__.items():
if isinstance(value, type) and value.__module__ == cls.__module__: for name, variant in list(cls.__dict__.items()):
value.__annotations__["_id"] = u8
value._id = u8(index) if (
value = pwp(value) isinstance(variant, type)
variants[index] = value # Safely align variant registration index and variant.__module__ == cls.__module__
setattr(cls, name, value) ):
variant.__annotations__ = {
"_id": u8,
**getattr(
variant,
"__annotations__",
{}
)
}
variant._id = u8(index)
variant = pwp(variant)
variants[index] = variant
setattr(
cls,
name,
variant
)
index += 1 index += 1
def to_com(self):
for idx, variant in variants.items():
if isinstance(self, variant):
return (
idx.to_bytes(1, "little")
+ PulseWire.to_com(self)
)
raise PWPError(
f"Unknown enum variant {type(self)}"
)
@classmethod @classmethod
def from_com(enum_cls, data: bytes, offset: int = 0) -> tuple[typing.Any, int]: def from_com(
enum_cls,
data,
offset=0
):
ensure_size(
data,
offset,
1
)
enum_id = data[offset] enum_id = data[offset]
offset += 1 offset += 1
if enum_id not in variants: if enum_id not in variants:
raise PWPError(f"Unknown enum id: {enum_id}") raise PWPError(
f"Unknown enum id {enum_id}"
)
variant = variants[enum_id]
obj = variant.__new__(variant)
fields = get_type_hints(variant)
for name, typ in fields.items(): return variants[enum_id].from_com(
if name == "_id": data,
continue offset
value, offset = decode_value(data, offset, typ) )
setattr(obj, name, value)
return obj, offset
cls.to_com = to_com
cls.from_com = from_com cls.from_com = from_com
return cls return cls