diff --git a/src/__init__.py b/src/__init__.py index bb8ad77..c2d43b9 100644 --- a/src/__init__.py +++ b/src/__init__.py @@ -49,5 +49,48 @@ class Strategy(PulseWire): def getWatchList(self): 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): - 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) diff --git a/src/plugin.py b/src/plugin.py index e8e7eb4..fa3716b 100644 --- a/src/plugin.py +++ b/src/plugin.py @@ -20,7 +20,7 @@ class StrategyMessage(PulseWire): class RequestCandlestick(PulseWire): symbol: str - interval: u8 + interval: general.CandleInterval count: u32 class Subscribe(PulseWire): @@ -39,7 +39,7 @@ class RiskMessage(PulseWire): class Log(PulseWire): log: general.EventLog - class GetWatchList: pass + class GetWatchList(PulseWire): pass class Approve(PulseWire): signal: general.Signal diff --git a/src/pwp.py b/src/pwp.py index d5b5f00..dcbecbf 100644 --- a/src/pwp.py +++ b/src/pwp.py @@ -1,12 +1,22 @@ import struct 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 * -# A Protocol representing the instance methods added by @pwp + class PulseWire(Protocol): def to_com(self) -> bytes: return bytes() + @classmethod def from_com(cls, data: bytes, offset: int = 0) -> tuple[typing.Any, int]: return (None, offset) @@ -14,22 +24,35 @@ class PulseWire(Protocol): class PWPError(Exception): 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: origin = get_origin(typ) # Optional[T] - if origin is typing.Union: + if origin in (typing.Union, types.UnionType): args = get_args(typ) if type(None) in args: if value is None: return b"\x00" - real_type = next(t for t in args if t is not type(None)) - return b"\x01" + encode_value(value, real_type) + real = next( + x for x in args + if x is not type(None) + ) + + return b"\x01" + encode_value(value, real) # list[T] - elif origin is list: + if origin is list: item_type = get_args(typ)[0] out = len(value).to_bytes(4, "little") @@ -39,7 +62,18 @@ def encode_value(value, typ) -> bytes: 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: key_type, value_type = get_args(typ) @@ -51,13 +85,11 @@ def encode_value(value, typ) -> bytes: return out - # primitives - if typ is str: - data = value.encode("utf-8") - return len(data).to_bytes(4, "little") + data - if typ is bytes: - return len(value).to_bytes(4, "little") + value + # primitives + + if typ is bool: + return b"\x01" if value else b"\x00" if typ is int: return struct.pack(" bytes: if typ is float: return struct.pack(" T: + fields = get_type_hints(cls) - def to_com(self) -> bytes: - output = b"" + + def to_com(self): + + out = b"" + 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 - def from_com(cls, data: bytes, offset: int = 0) -> tuple[typing.Any, int]: - values = {} + def from_com( + cls, + data, + offset=0 + ): + obj = cls.__new__(cls) 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 - cls.to_com = to_com + + + cls.to_com = to_com cls.from_com = from_com return cls + + def pwp_enum(cls: T) -> T: + variants = {} + index = 0 - - for name, value in cls.__dict__.items(): - if isinstance(value, type) and value.__module__ == cls.__module__: - value.__annotations__["_id"] = u8 - value._id = u8(index) - value = pwp(value) - variants[index] = value # Safely align variant registration index - setattr(cls, name, value) + + + for name, variant in list(cls.__dict__.items()): + + if ( + isinstance(variant, type) + and variant.__module__ == cls.__module__ + ): + + variant.__annotations__ = { + "_id": u8, + **getattr( + variant, + "__annotations__", + {} + ) + } + + variant._id = u8(index) + + variant = pwp(variant) + + variants[index] = variant + + setattr( + cls, + name, + variant + ) + 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 - 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] offset += 1 + 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(): - if name == "_id": - continue - value, offset = decode_value(data, offset, typ) - setattr(obj, name, value) + return variants[enum_id].from_com( + data, + offset + ) - return obj, offset + cls.to_com = to_com cls.from_com = from_com return cls \ No newline at end of file