Coverage for src / monte_neo / data / websocket.py: 71%

131 statements  

« prev     ^ index     » next       coverage.py v7.13.1, created at 2026-01-28 16:27 +0200

1"""Binance WebSocket streaming utilities.""" 

2 

3from __future__ import annotations 

4 

5import time 

6from collections.abc import Callable 

7from typing import Any 

8 

9from binance.websocket.spot.websocket_stream import ( 

10 SpotWebsocketStreamClient as WebsocketClient, 

11) 

12 

13from monte_neo.utils.logger import get_logger 

14 

15logger = get_logger(__name__) 

16 

17 

18class BinanceWebsocketStreamer: 

19 """Stream Binance market data via WebSocket.""" 

20 

21 _LOWER_INTERVALS = { 

22 "1m", 

23 "3m", 

24 "5m", 

25 "15m", 

26 "30m", 

27 "1h", 

28 "2h", 

29 "4h", 

30 "6h", 

31 "8h", 

32 "12h", 

33 "1d", 

34 "3d", 

35 "1w", 

36 } 

37 _UPPER_INTERVALS = {"1M"} 

38 

39 def __init__( 

40 self, 

41 stream_url: str | None = None, 

42 callback: Callable[[dict], None] | None = None, 

43 ) -> None: 

44 """Initialize WebSocket client. 

45 

46 Args: 

47 stream_url: Optional WebSocket endpoint override. 

48 callback: Optional default handler for incoming messages. 

49 """ 

50 self._callbacks: list[Callable[[dict], None]] = [] 

51 if callback: 

52 self._callbacks.append(callback) 

53 self._client = WebsocketClient( 

54 stream_url=stream_url or "wss://stream.binance.com:9443", 

55 on_message=self._dispatch_message, 

56 on_error=self._handle_error, 

57 on_close=self._handle_close, 

58 ) 

59 self._started = False 

60 self._active_streams: set[str] = set() 

61 self._reconnect_attempts = 0 

62 self._max_reconnect_attempts = 5 

63 

64 def __enter__(self) -> BinanceWebsocketStreamer: 

65 """Enter context manager.""" 

66 self.start() 

67 return self 

68 

69 def __exit__(self, exc_type, exc, exc_tb) -> None: 

70 """Exit context manager.""" 

71 self.stop() 

72 

73 def start(self) -> None: 

74 """Start WebSocket connection.""" 

75 if not self._started: 

76 logger.info("Starting WebSocket connection...") 

77 self._client.start() 

78 self._started = True 

79 self._reconnect_attempts = 0 

80 

81 def stop(self) -> None: 

82 """Stop WebSocket connection.""" 

83 if self._started: 

84 logger.info("Stopping WebSocket connection...") 

85 try: 

86 self._started = False # Set flag first to prevent auto-reconnect 

87 self._client.stop() 

88 except Exception as exc: 

89 logger.error("Error stopping WebSocket: %s", exc) 

90 

91 def subscribe_kline( 

92 self, 

93 symbol: str, 

94 interval: str, 

95 callback: Callable[[dict], None], 

96 stream_id: int = 1, 

97 ) -> None: 

98 """Subscribe to kline stream. 

99 

100 Args: 

101 symbol: Trading pair symbol (e.g., 'BTCUSDT'). 

102 interval: Candle interval (e.g., '1m', '1h'). 

103 callback: Handler for incoming messages. 

104 stream_id: Client message id. 

105 """ 

106 normalized_symbol = self._normalize_symbol(symbol) 

107 normalized_interval = self._normalize_interval(interval) 

108 

109 stream_name = f"{normalized_symbol}@kline_{normalized_interval}" 

110 self._active_streams.add(stream_name) 

111 

112 self.start() 

113 self._register_callback(callback) 

114 # Use subscribe directly to maintain consistency with reconnection logic 

115 self._client.subscribe(stream=stream_name, id=stream_id) 

116 

117 def subscribe_mini_ticker( 

118 self, 

119 symbol: str, 

120 callback: Callable[[dict], None], 

121 stream_id: int = 2, 

122 ) -> None: 

123 """Subscribe to mini ticker stream. 

124 

125 Args: 

126 symbol: Trading pair symbol (e.g., 'BTCUSDT'). 

127 callback: Handler for incoming messages. 

128 stream_id: Client message id. 

129 """ 

130 normalized_symbol = self._normalize_symbol(symbol) 

131 stream_name = f"{normalized_symbol}@miniTicker" 

132 self._active_streams.add(stream_name) 

133 

134 self.start() 

135 self._register_callback(callback) 

136 self._client.subscribe(stream=stream_name, id=stream_id) 

137 

138 def subscribe_streams( 

139 self, 

140 streams: list[str], 

141 callback: Callable[[dict], None], 

142 stream_id: int = 3, 

143 ) -> None: 

144 """Subscribe to multiple raw stream names. 

145 

146 Args: 

147 streams: Raw stream names (e.g., ['btcusdt@trade']). 

148 callback: Handler for incoming messages. 

149 stream_id: Client message id. 

150 """ 

151 normalized_streams = self._normalize_streams(streams) 

152 for stream in normalized_streams: 

153 self._active_streams.add(stream) 

154 

155 self.start() 

156 self._register_callback(callback) 

157 self._client.subscribe(stream=normalized_streams, id=stream_id) 

158 

159 def _normalize_symbol(self, symbol: str) -> str: 

160 normalized = symbol.strip() 

161 if not normalized: 

162 raise ValueError("Symbol must be non-empty") 

163 return normalized.lower() 

164 

165 def _normalize_interval(self, interval: str) -> str: 

166 normalized = interval.strip() 

167 if normalized in self._UPPER_INTERVALS: 

168 return normalized 

169 lower = normalized.lower() 

170 if lower in self._LOWER_INTERVALS: 

171 return lower 

172 raise ValueError(f"Unsupported interval: {interval}") 

173 

174 def _normalize_streams(self, streams: list[str]) -> list[str]: 

175 if not streams: 

176 raise ValueError("Streams list must be non-empty") 

177 normalized = [stream.strip() for stream in streams] 

178 if any(not stream for stream in normalized): 

179 raise ValueError("Streams list contains empty stream") 

180 return normalized 

181 

182 def _register_callback(self, callback: Callable[[dict], None]) -> None: 

183 if callback not in self._callbacks: 

184 self._callbacks.append(callback) 

185 

186 def _dispatch_message(self, _, message: Any) -> None: 

187 if not self._callbacks: 

188 return 

189 if isinstance(message, str): 

190 import json 

191 try: 

192 payload = json.loads(message) 

193 except Exception: 

194 payload = {"message": message} 

195 elif isinstance(message, dict): 

196 payload = message 

197 else: 

198 payload = {"message": message} 

199 

200 for callback in self._callbacks: 

201 try: 

202 callback(payload) 

203 except Exception as exc: 

204 logger.exception("WebSocket callback error: %s", exc) 

205 

206 def _handle_error(self, error: Any) -> None: 

207 logger.error("WebSocket error: %s", error) 

208 self._attempt_reconnect() 

209 

210 def _handle_close(self, *args) -> None: 

211 """Handle WebSocket close event.""" 

212 logger.info("WebSocket connection closed") 

213 if self._started: 

214 logger.warning("Unexpected close, attempting reconnect...") 

215 self._attempt_reconnect() 

216 

217 def _attempt_reconnect(self) -> None: 

218 """Attempt to reconnect to WebSocket stream with backoff.""" 

219 if not self._started: 

220 return 

221 

222 if self._reconnect_attempts >= self._max_reconnect_attempts: 

223 logger.error( 

224 "Max reconnect attempts (%d) reached. Giving up.", 

225 self._max_reconnect_attempts, 

226 ) 

227 self._started = False 

228 return 

229 

230 self._reconnect_attempts += 1 

231 delay = min(2**self._reconnect_attempts, 60) # Exponential backoff 

232 logger.info( 

233 "Attempting reconnect %d/%d in %ds...", 

234 self._reconnect_attempts, 

235 self._max_reconnect_attempts, 

236 delay, 

237 ) 

238 

239 time.sleep(delay) 

240 

241 try: 

242 # Force stop old connection to be safe 

243 try: 

244 self._client.stop() 

245 except Exception: 

246 pass 

247 

248 self._client.start() 

249 

250 # Resubscribe to active streams 

251 if self._active_streams: 

252 logger.info("Resubscribing to %d streams...", len(self._active_streams)) 

253 self._client.subscribe(stream=list(self._active_streams)) 

254 

255 logger.info("Reconnect successful") 

256 

257 except Exception as exc: 

258 logger.error("Reconnect failed: %s", exc) 

259 # If immediate reconnect failed, try again recursively 

260 self._attempt_reconnect()