diff --git a/src/__init__.py b/src/__init__.py index 359fe89..bd214e0 100644 --- a/src/__init__.py +++ b/src/__init__.py @@ -1,12 +1,16 @@ -import sys import plugin +import sys +from typing import Any, Generic, TypeVar, Type, Optional import general -class PulseWire: - def __init__(self): - ... +T = TypeVar("T", bound=general.PulseWire) +TT = TypeVar("TT", bound=general.PulseWire) - def onRaw(self, data: bytes): +class PulseWire(Generic[T, TT]): + def __init__(self, wire_cls: Type[T]): + self.wire_cls = wire_cls + + def on(self, req: TT): ... def send_raw(self, data: bytes): @@ -18,24 +22,39 @@ class PulseWire: while True: len_buf = sys.stdin.buffer.read(8) - sys.stderr.write(str(len_buf)) + # Standard EOF check: zero bytes returned from read means stdin closed + if not len_buf: + break - if len(len_buf) == 0: - continue + if len(len_buf) < 8: + raise EOFError("Unexpected EOF while reading message length header") length = int.from_bytes(len_buf, "little") - if length == 0: continue buffer = sys.stdin.buffer.read(length) - if len(buffer) != length: raise EOFError("Unexpected EOF while reading payload") self.onRaw(buffer) -class Strategy(PulseWire): + def onRaw(self, data: bytes): + # Call from_com on the actual class passed in __init__ + result: tuple[TT, int] = self.wire_cls.from_com(data=data) + req, _ = result + + self.on(req) + + # Dynamic dispatch based on class name (e.g., onMyRequest) + handler = getattr(self, f"on{req.__class__.__name__}", None) + if handler and callable(handler): + handler(req) + +class Strategy(PulseWire[plugin.StrategyEngineMessage, plugin.StrategyEngineMessageType]): + def __init__(self): + super().__init__(plugin.StrategyEngineMessage) + def onInitialize(self, msg: plugin.StrategyEngineMessage.Initialize): ... @@ -51,9 +70,6 @@ class Strategy(PulseWire): def onCandlestick(self, msg: plugin.StrategyEngineMessage.Candlestick): ... - def on(self, req: plugin.StrategyEngineMessageType): - ... - def send(self, req: plugin.StrategyMessageType): self.send_raw(req.to_com()) @@ -80,22 +96,10 @@ class Strategy(PulseWire): msg.signal = signal self.send(msg) - def onRaw(self, data: bytes): - result: tuple[plugin.StrategyEngineMessageType, int] = plugin.StrategyEngineMessage.from_com(data=data) - req, _ = result +class Risk(PulseWire[plugin.RiskEngineMessage, plugin.RiskEngineMessageType]): + def __init__(self): + super().__init__(plugin.RiskEngineMessage) - self.on(req) - - handler = getattr( - self, - f"on{req.__class__.__name__}", - None - ) - - if handler: - handler(req) - -class Risk(PulseWire): def onInitialize(self, msg: plugin.RiskEngineMessage.Initialize): ... @@ -108,9 +112,6 @@ class Risk(PulseWire): def onSignal(self, msg: plugin.RiskEngineMessage.Signal): ... - def on(self, req: plugin.RiskEngineMessageType): - ... - def send(self, req: plugin.RiskMessageType): self.send_raw(req.to_com()) @@ -131,28 +132,3 @@ class Risk(PulseWire): msg = plugin.RiskMessage.Reject() msg.reason = reason self.send(msg) - - def onRaw(self, data: bytes): - result: tuple[plugin.RiskEngineMessageType, int] = plugin.RiskEngineMessage.from_com(data=data) - req, _ = result - - self.on(req) - - handler = getattr( - self, - f"on{req.__class__.__name__}", - None - ) - - if handler: - handler(req) - -class Bro(Strategy): - def __init__(self): - super().__init__() - - def onInitialize(self, msg: plugin.StrategyEngineMessage.Initialize): - self.log(general.EventLog.info("bro", "Hello, World")) - return super().onInitialize(msg) - -Bro().start() \ No newline at end of file