stx-python 0.6.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- stx/__init__.py +61 -0
- stx/_async_client.py +923 -0
- stx/_base.py +64 -0
- stx/_client.py +502 -0
- stx/_config.py +107 -0
- stx/_http.py +131 -0
- stx/_money.py +122 -0
- stx/_operations.py +259 -0
- stx/_paging.py +59 -0
- stx/_results.py +29 -0
- stx/_retry.py +80 -0
- stx/_settings.py +203 -0
- stx/_signing.py +229 -0
- stx/_version.py +12 -0
- stx/_ws.py +1016 -0
- stx/enums.py +13 -0
- stx/exceptions.py +156 -0
- stx/models.py +2034 -0
- stx/py.typed +0 -0
- stx_python-0.6.0.dist-info/METADATA +121 -0
- stx_python-0.6.0.dist-info/RECORD +24 -0
- stx_python-0.6.0.dist-info/WHEEL +5 -0
- stx_python-0.6.0.dist-info/licenses/LICENSE +21 -0
- stx_python-0.6.0.dist-info/top_level.txt +1 -0
stx/_ws.py
ADDED
|
@@ -0,0 +1,1016 @@
|
|
|
1
|
+
"""``STXWebSocket``: Phoenix channels over one signed socket.
|
|
2
|
+
|
|
3
|
+
One socket carries every topic. Each join method returns a
|
|
4
|
+
:class:`Channel`; messages reach you through an ``on_message`` callback,
|
|
5
|
+
by iterating the channel (``async for msg in channel``), or both.
|
|
6
|
+
|
|
7
|
+
What the client does for you:
|
|
8
|
+
|
|
9
|
+
* signs the handshake with your API key (``GET /socket/websocket``, query
|
|
10
|
+
string dropped), with a fresh signature on every connect;
|
|
11
|
+
* sends the socket heartbeat on the ``phoenix`` topic and reconnects when a
|
|
12
|
+
heartbeat goes unanswered;
|
|
13
|
+
* pings each joined channel on a timer; on ``orders`` with
|
|
14
|
+
``cancel_on_disconnect`` the ping runs at 60% of the ``ping_timeout`` the
|
|
15
|
+
server granted, which is a separate deadline from the socket heartbeat;
|
|
16
|
+
* reconnects with backoff after a drop, re-signs, and rejoins every channel
|
|
17
|
+
with its current filters (including ``market_updates`` watches), so the
|
|
18
|
+
snapshot events arrive again;
|
|
19
|
+
* converts the cents in ``markets`` and ``market_updates`` to dollar
|
|
20
|
+
strings (see ``stx._money``).
|
|
21
|
+
|
|
22
|
+
What stays yours: reconciling after a reconnect. Pass ``on_reconnect`` to
|
|
23
|
+
be told when to take a fresh REST snapshot.
|
|
24
|
+
"""
|
|
25
|
+
|
|
26
|
+
from __future__ import annotations
|
|
27
|
+
|
|
28
|
+
import asyncio
|
|
29
|
+
import inspect
|
|
30
|
+
import itertools
|
|
31
|
+
import json
|
|
32
|
+
import logging
|
|
33
|
+
import random
|
|
34
|
+
from dataclasses import dataclass, field
|
|
35
|
+
from typing import (
|
|
36
|
+
Any,
|
|
37
|
+
AsyncIterator,
|
|
38
|
+
Awaitable,
|
|
39
|
+
Callable,
|
|
40
|
+
Dict,
|
|
41
|
+
List,
|
|
42
|
+
Optional,
|
|
43
|
+
Sequence,
|
|
44
|
+
Union,
|
|
45
|
+
)
|
|
46
|
+
|
|
47
|
+
import websockets
|
|
48
|
+
|
|
49
|
+
from stx._money import convert_market_payload, convert_markets_frame
|
|
50
|
+
from stx._settings import _SENTINEL as _UNSET
|
|
51
|
+
from stx._signing import Signer
|
|
52
|
+
from stx._version import USER_AGENT
|
|
53
|
+
from stx.exceptions import STXChannelException, STXConfigException, STXTransportException
|
|
54
|
+
|
|
55
|
+
logger = logging.getLogger("stx.ws")
|
|
56
|
+
|
|
57
|
+
PHOENIX_TOPIC = "phoenix"
|
|
58
|
+
|
|
59
|
+
# Events that carry state on join, by channel name (topic prefix).
|
|
60
|
+
SNAPSHOT_EVENTS: Dict[str, tuple] = {
|
|
61
|
+
"orders": ("all_orders",),
|
|
62
|
+
"fills": ("all_trades",),
|
|
63
|
+
"positions": ("all_positions",),
|
|
64
|
+
"balances": ("balances",),
|
|
65
|
+
"account": ("all_orders", "all_trades", "all_positions", "balances"),
|
|
66
|
+
"user_info": ("user_updated",),
|
|
67
|
+
}
|
|
68
|
+
|
|
69
|
+
# Channels whose join reply itself is the snapshot.
|
|
70
|
+
REPLY_SNAPSHOT_CHANNELS = frozenset({"market_stats"})
|
|
71
|
+
|
|
72
|
+
ACCOUNT_CHANNELS = frozenset(
|
|
73
|
+
{"orders", "fills", "positions", "settlements", "balances", "account", "user_info"}
|
|
74
|
+
)
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
@dataclass
|
|
78
|
+
class ChannelMessage:
|
|
79
|
+
"""One pushed frame.
|
|
80
|
+
|
|
81
|
+
``payload`` is the event's JSON object as documented for the channel,
|
|
82
|
+
with money and quantities as strings. ``channel`` is the topic without
|
|
83
|
+
the user id (``"orders"`` for ``orders:<user_id>``).
|
|
84
|
+
"""
|
|
85
|
+
|
|
86
|
+
topic: str
|
|
87
|
+
event: str
|
|
88
|
+
payload: Any
|
|
89
|
+
ref: Optional[str] = None
|
|
90
|
+
join_ref: Optional[str] = None
|
|
91
|
+
|
|
92
|
+
@property
|
|
93
|
+
def channel(self) -> str:
|
|
94
|
+
return self.topic.split(":", 1)[0]
|
|
95
|
+
|
|
96
|
+
@property
|
|
97
|
+
def is_snapshot(self) -> bool:
|
|
98
|
+
"""``True`` for the state-on-join events (``all_orders``, ``balances``...)."""
|
|
99
|
+
return self.event in SNAPSHOT_EVENTS.get(self.channel, ())
|
|
100
|
+
|
|
101
|
+
|
|
102
|
+
MessageHandler = Callable[[ChannelMessage], Union[Awaitable[None], None]]
|
|
103
|
+
|
|
104
|
+
|
|
105
|
+
def _handshake_headers_kwarg() -> str:
|
|
106
|
+
"""``additional_headers`` on websockets >= 13, ``extra_headers`` before."""
|
|
107
|
+
try:
|
|
108
|
+
params = inspect.signature(websockets.connect).parameters
|
|
109
|
+
except (TypeError, ValueError): # pragma: no cover
|
|
110
|
+
return "extra_headers"
|
|
111
|
+
return "additional_headers" if "additional_headers" in params else "extra_headers"
|
|
112
|
+
|
|
113
|
+
|
|
114
|
+
@dataclass
|
|
115
|
+
class ReconnectPolicy:
|
|
116
|
+
"""Backoff between reconnect attempts. ``max_attempts=None`` never gives up."""
|
|
117
|
+
|
|
118
|
+
initial_backoff: float = 0.5
|
|
119
|
+
max_backoff: float = 30.0
|
|
120
|
+
max_attempts: Optional[int] = None
|
|
121
|
+
jitter: bool = True
|
|
122
|
+
|
|
123
|
+
def delay(self, attempt: int) -> float:
|
|
124
|
+
base = min(self.initial_backoff * (2 ** (attempt - 1)), self.max_backoff)
|
|
125
|
+
return base * (0.5 + random.random() * 0.5) if self.jitter else base
|
|
126
|
+
|
|
127
|
+
|
|
128
|
+
class Channel:
|
|
129
|
+
"""One joined topic.
|
|
130
|
+
|
|
131
|
+
Attributes:
|
|
132
|
+
topic: the wire topic, e.g. ``orders:<user_id>``.
|
|
133
|
+
name: the topic without the user id, e.g. ``orders``.
|
|
134
|
+
join_payload: what the next (re)join sends; updated by the
|
|
135
|
+
``select_*`` and ``watch`` helpers so a reconnect keeps them.
|
|
136
|
+
reply: the ``response`` object from the latest join reply, e.g.
|
|
137
|
+
``{"selected_market_ids": None}``.
|
|
138
|
+
snapshots: the latest snapshot payload per snapshot event name.
|
|
139
|
+
"""
|
|
140
|
+
|
|
141
|
+
def __init__(
|
|
142
|
+
self,
|
|
143
|
+
ws: STXWebSocket,
|
|
144
|
+
topic: str,
|
|
145
|
+
join_payload: Dict[str, Any],
|
|
146
|
+
on_message: Optional[MessageHandler],
|
|
147
|
+
*,
|
|
148
|
+
ping_interval: Optional[float],
|
|
149
|
+
queue_size: int,
|
|
150
|
+
) -> None:
|
|
151
|
+
self.ws = ws
|
|
152
|
+
self.topic = topic
|
|
153
|
+
self.name = topic.split(":", 1)[0]
|
|
154
|
+
self.join_payload = dict(join_payload)
|
|
155
|
+
self.on_message = on_message
|
|
156
|
+
self.ping_interval = ping_interval
|
|
157
|
+
self.reply: Dict[str, Any] = {}
|
|
158
|
+
self.snapshots: Dict[str, Any] = {}
|
|
159
|
+
self.watches: List[str] = []
|
|
160
|
+
self.join_ref: Optional[str] = None
|
|
161
|
+
self.joined = False
|
|
162
|
+
self.closed = False
|
|
163
|
+
self._queue: asyncio.Queue[Optional[ChannelMessage]] = asyncio.Queue(queue_size)
|
|
164
|
+
self._snapshot_event = asyncio.Event()
|
|
165
|
+
self._ping_task: Optional[asyncio.Task] = None
|
|
166
|
+
|
|
167
|
+
def __repr__(self) -> str:
|
|
168
|
+
state = "joined" if self.joined else ("closed" if self.closed else "joining")
|
|
169
|
+
return f"Channel({self.topic!r}, {state})"
|
|
170
|
+
|
|
171
|
+
# -- receiving ------------------------------------------------------
|
|
172
|
+
|
|
173
|
+
def __aiter__(self) -> AsyncIterator[ChannelMessage]:
|
|
174
|
+
return self._iter()
|
|
175
|
+
|
|
176
|
+
async def _iter(self) -> AsyncIterator[ChannelMessage]:
|
|
177
|
+
while True:
|
|
178
|
+
msg = await self._queue.get()
|
|
179
|
+
if msg is None:
|
|
180
|
+
return
|
|
181
|
+
yield msg
|
|
182
|
+
|
|
183
|
+
async def next(self, timeout: Optional[float] = None) -> ChannelMessage:
|
|
184
|
+
"""The next message on this channel. Raises ``asyncio.TimeoutError``."""
|
|
185
|
+
msg = await asyncio.wait_for(self._queue.get(), timeout)
|
|
186
|
+
if msg is None:
|
|
187
|
+
raise STXChannelException(f"{self.topic} is closed", topic=self.topic)
|
|
188
|
+
return msg
|
|
189
|
+
|
|
190
|
+
async def wait_snapshot(self, timeout: Optional[float] = 10.0) -> Dict[str, Any]:
|
|
191
|
+
"""Wait for the state-on-join and return it as ``{event: payload}``.
|
|
192
|
+
|
|
193
|
+
``orders`` gives ``{"all_orders": {...}}``; ``account`` waits for all
|
|
194
|
+
four of its snapshots. ``market_stats`` returns the join reply, which
|
|
195
|
+
carries the series. Channels with no snapshot (``settlements``,
|
|
196
|
+
``ticker``, ``trades``, ``orderbook``, ``markets``,
|
|
197
|
+
``market_updates``) raise ``ValueError``.
|
|
198
|
+
"""
|
|
199
|
+
if self.name in REPLY_SNAPSHOT_CHANNELS:
|
|
200
|
+
return {"reply": self.reply}
|
|
201
|
+
if self.name not in SNAPSHOT_EVENTS:
|
|
202
|
+
raise ValueError(f"{self.name} sends no snapshot on join")
|
|
203
|
+
await asyncio.wait_for(self._snapshot_event.wait(), timeout)
|
|
204
|
+
return dict(self.snapshots)
|
|
205
|
+
|
|
206
|
+
async def _deliver(self, msg: ChannelMessage) -> None:
|
|
207
|
+
expected = SNAPSHOT_EVENTS.get(self.name, ())
|
|
208
|
+
if msg.event in expected:
|
|
209
|
+
self.snapshots[msg.event] = msg.payload
|
|
210
|
+
if all(e in self.snapshots for e in expected):
|
|
211
|
+
self._snapshot_event.set()
|
|
212
|
+
if self.on_message is not None:
|
|
213
|
+
try:
|
|
214
|
+
result = self.on_message(msg)
|
|
215
|
+
if inspect.isawaitable(result):
|
|
216
|
+
await result
|
|
217
|
+
except Exception: # a bad callback must not kill the reader
|
|
218
|
+
logger.exception("on_message for %s raised", self.topic)
|
|
219
|
+
if self._queue.full():
|
|
220
|
+
# Drop the oldest so a slow consumer sees current state.
|
|
221
|
+
try:
|
|
222
|
+
self._queue.get_nowait()
|
|
223
|
+
except asyncio.QueueEmpty: # pragma: no cover
|
|
224
|
+
pass
|
|
225
|
+
logger.warning("%s queue full; dropped the oldest message", self.topic)
|
|
226
|
+
self._queue.put_nowait(msg)
|
|
227
|
+
|
|
228
|
+
# -- sending --------------------------------------------------------
|
|
229
|
+
|
|
230
|
+
async def push(self, event: str, payload: Any = None, timeout: Optional[float] = None) -> Any:
|
|
231
|
+
"""Send ``event`` on this channel and return the reply's ``response``.
|
|
232
|
+
|
|
233
|
+
Raises ``STXChannelException`` when the reply status is ``error``.
|
|
234
|
+
"""
|
|
235
|
+
if not self.joined or self.join_ref is None:
|
|
236
|
+
raise STXChannelException(f"{self.topic} is not joined", topic=self.topic)
|
|
237
|
+
return await self.ws._request(self.topic, event, {} if payload is None else payload,
|
|
238
|
+
self.join_ref, timeout) # fmt: skip
|
|
239
|
+
|
|
240
|
+
async def ping(self) -> Any:
|
|
241
|
+
"""Channel ``ping``. On ``orders`` with cancel-on-disconnect armed this
|
|
242
|
+
resets the cancel deadline; on ``account`` and ``orders`` it keeps the
|
|
243
|
+
session alive."""
|
|
244
|
+
return await self.push("ping", {})
|
|
245
|
+
|
|
246
|
+
async def select_market_ids(self, market_ids: Optional[Sequence[str]]) -> Any:
|
|
247
|
+
"""Change the ``market_ids`` filter without rejoining.
|
|
248
|
+
|
|
249
|
+
On ``orders``, ``fills``, ``positions``, ``settlements`` and
|
|
250
|
+
``account`` ``None`` clears the filter; ``orderbook`` and
|
|
251
|
+
``market_stats`` require at least one id. Returns the reply, whose
|
|
252
|
+
``selected_market_ids`` is what the server applied.
|
|
253
|
+
"""
|
|
254
|
+
ids = None if market_ids is None else list(market_ids)
|
|
255
|
+
reply = await self.push("select_market_ids", {"market_ids": ids})
|
|
256
|
+
self.join_payload["market_ids"] = ids
|
|
257
|
+
return reply
|
|
258
|
+
|
|
259
|
+
async def select_filters(self, **filters: Optional[Sequence[str]]) -> Any:
|
|
260
|
+
"""``ticker`` (``sports``, ``competitions``) and ``trades``
|
|
261
|
+
(``market_ids``, ``event_ids``): change filters without rejoining."""
|
|
262
|
+
payload = {k: (None if v is None else list(v)) for k, v in filters.items()}
|
|
263
|
+
reply = await self.push("select_filters", payload)
|
|
264
|
+
self.join_payload.update(payload)
|
|
265
|
+
return reply
|
|
266
|
+
|
|
267
|
+
async def select_rule_filters(self, rule_filters: Optional[Sequence[str]]) -> Any:
|
|
268
|
+
"""``markets``: change the ``rules`` filter; ``None`` disables it."""
|
|
269
|
+
value = None if rule_filters is None else list(rule_filters)
|
|
270
|
+
reply = await self.push("select_rule_filters", {"rule_filters": value})
|
|
271
|
+
self.join_payload["rule_filters"] = value
|
|
272
|
+
return reply
|
|
273
|
+
|
|
274
|
+
async def select_message_types(self, message_types: Optional[Sequence[str]]) -> Any:
|
|
275
|
+
"""``markets``: receive ``market_created``, ``market_updated`` or both."""
|
|
276
|
+
value = None if message_types is None else list(message_types)
|
|
277
|
+
reply = await self.push("select_message_types", {"message_types": value})
|
|
278
|
+
self.join_payload["message_types"] = value
|
|
279
|
+
return reply
|
|
280
|
+
|
|
281
|
+
async def watch(self, market_ids: Sequence[str]) -> Any:
|
|
282
|
+
"""``market_updates``: start receiving ``created``/``updated`` for
|
|
283
|
+
these markets. Re-sent automatically after a reconnect. Returns the
|
|
284
|
+
reply, whose ``subscriptions.watches`` lists what is watched."""
|
|
285
|
+
ids = list(market_ids)
|
|
286
|
+
reply = await self.push("watch", ids)
|
|
287
|
+
for mid in ids:
|
|
288
|
+
if mid not in self.watches:
|
|
289
|
+
self.watches.append(mid)
|
|
290
|
+
return reply
|
|
291
|
+
|
|
292
|
+
async def request_series(self, market_ids: Sequence[str], range: str = "all") -> Any:
|
|
293
|
+
"""``market_stats``: fetch history at ``range`` (``day``, ``week``,
|
|
294
|
+
``month``, ``all``) without changing the subscription."""
|
|
295
|
+
return await self.push("request_series", {"market_ids": list(market_ids), "range": range})
|
|
296
|
+
|
|
297
|
+
async def leave(self) -> None:
|
|
298
|
+
"""Leave the topic. The channel stops receiving and its iterator ends."""
|
|
299
|
+
await self.ws._leave(self)
|
|
300
|
+
|
|
301
|
+
# -- internals ------------------------------------------------------
|
|
302
|
+
|
|
303
|
+
def _start_pinger(self) -> None:
|
|
304
|
+
self._stop_pinger()
|
|
305
|
+
interval = self.ping_interval
|
|
306
|
+
granted = self.reply.get("ping_timeout") if isinstance(self.reply, dict) else None
|
|
307
|
+
if isinstance(granted, int) and self.join_payload.get("cancel_on_disconnect"):
|
|
308
|
+
# Cancel-on-disconnect deadline: ping at 60% of what the server
|
|
309
|
+
# granted, whatever the general interval is.
|
|
310
|
+
cod = granted / 1000.0 * 0.6
|
|
311
|
+
interval = cod if interval is None else min(interval, cod)
|
|
312
|
+
if interval:
|
|
313
|
+
self._ping_task = asyncio.ensure_future(self._ping_loop(interval))
|
|
314
|
+
|
|
315
|
+
def _stop_pinger(self) -> None:
|
|
316
|
+
if self._ping_task is not None and not self._ping_task.done():
|
|
317
|
+
self._ping_task.cancel()
|
|
318
|
+
self._ping_task = None
|
|
319
|
+
|
|
320
|
+
async def _ping_loop(self, interval: float) -> None:
|
|
321
|
+
while self.joined and not self.closed:
|
|
322
|
+
await asyncio.sleep(interval)
|
|
323
|
+
if not self.joined:
|
|
324
|
+
return
|
|
325
|
+
try:
|
|
326
|
+
await self.push("ping", {}, timeout=max(interval, 5.0))
|
|
327
|
+
except asyncio.CancelledError:
|
|
328
|
+
raise
|
|
329
|
+
except Exception as exc:
|
|
330
|
+
logger.warning("%s ping failed: %s", self.topic, exc)
|
|
331
|
+
|
|
332
|
+
def _reset_for_rejoin(self) -> None:
|
|
333
|
+
self.joined = False
|
|
334
|
+
self.snapshots.clear()
|
|
335
|
+
self._snapshot_event.clear()
|
|
336
|
+
self._stop_pinger()
|
|
337
|
+
|
|
338
|
+
def _close(self) -> None:
|
|
339
|
+
self.closed = True
|
|
340
|
+
self.joined = False
|
|
341
|
+
self._stop_pinger()
|
|
342
|
+
try:
|
|
343
|
+
self._queue.put_nowait(None)
|
|
344
|
+
except asyncio.QueueFull:
|
|
345
|
+
self._queue.get_nowait()
|
|
346
|
+
self._queue.put_nowait(None)
|
|
347
|
+
|
|
348
|
+
|
|
349
|
+
@dataclass
|
|
350
|
+
class _Pending:
|
|
351
|
+
future: asyncio.Future[Any]
|
|
352
|
+
topic: str
|
|
353
|
+
event: str
|
|
354
|
+
|
|
355
|
+
|
|
356
|
+
@dataclass
|
|
357
|
+
class _State:
|
|
358
|
+
pending: Dict[str, _Pending] = field(default_factory=dict)
|
|
359
|
+
|
|
360
|
+
|
|
361
|
+
class STXWebSocket:
|
|
362
|
+
"""Async Phoenix-channels client for the documented STX topics.
|
|
363
|
+
|
|
364
|
+
Build it from an ``AsyncSTX`` client, which supplies the host, the key
|
|
365
|
+
and your user id::
|
|
366
|
+
|
|
367
|
+
async with AsyncSTX(profile="us-demo") as client:
|
|
368
|
+
async with client.websocket() as ws:
|
|
369
|
+
book = await ws.orderbook(["<market-id>"], on_message=print)
|
|
370
|
+
orders = await ws.orders()
|
|
371
|
+
print(await orders.wait_snapshot())
|
|
372
|
+
await ws.run_forever()
|
|
373
|
+
|
|
374
|
+
or standalone with the same settings arguments as ``AsyncSTX``.
|
|
375
|
+
|
|
376
|
+
Args:
|
|
377
|
+
rest: an ``AsyncSTX`` to take host, key and user id from.
|
|
378
|
+
user_id: your user id, if you already have it (skips ``GET /me``).
|
|
379
|
+
heartbeat_interval: seconds between socket heartbeats. The server
|
|
380
|
+
closes a socket silent for 60 s.
|
|
381
|
+
channel_ping_interval: seconds between channel ``ping`` frames on
|
|
382
|
+
every joined topic; ``None`` disables them. ``orders`` with
|
|
383
|
+
cancel-on-disconnect pings faster, from the granted timeout.
|
|
384
|
+
reconnect: reconnect after a drop (default ``True``).
|
|
385
|
+
reconnect_policy: backoff for reconnects.
|
|
386
|
+
on_reconnect: called (sync or async) after every successful
|
|
387
|
+
reconnect and rejoin: the moment to call ``orders()`` and
|
|
388
|
+
anything else you show again.
|
|
389
|
+
join_timeout: seconds to wait for a join or push reply.
|
|
390
|
+
queue_size: messages buffered per channel for ``async for``; the
|
|
391
|
+
oldest is dropped when full.
|
|
392
|
+
"""
|
|
393
|
+
|
|
394
|
+
def __init__(
|
|
395
|
+
self,
|
|
396
|
+
*,
|
|
397
|
+
rest: Any = None,
|
|
398
|
+
region: Any = _UNSET,
|
|
399
|
+
env: Any = _UNSET,
|
|
400
|
+
host: Any = _UNSET,
|
|
401
|
+
key_id: Any = _UNSET,
|
|
402
|
+
private_key: Any = _UNSET,
|
|
403
|
+
signer: Optional[Signer] = None,
|
|
404
|
+
profile: Any = _UNSET,
|
|
405
|
+
verify_tls: Any = _UNSET,
|
|
406
|
+
user_id: Optional[str] = None,
|
|
407
|
+
heartbeat_interval: float = 25.0,
|
|
408
|
+
channel_ping_interval: Optional[float] = 30.0,
|
|
409
|
+
reconnect: bool = True,
|
|
410
|
+
reconnect_policy: Optional[ReconnectPolicy] = None,
|
|
411
|
+
on_reconnect: Optional[Callable[[], Union[Awaitable[None], None]]] = None,
|
|
412
|
+
join_timeout: float = 10.0,
|
|
413
|
+
queue_size: int = 10_000,
|
|
414
|
+
) -> None:
|
|
415
|
+
if rest is None:
|
|
416
|
+
from stx._async_client import AsyncSTX
|
|
417
|
+
|
|
418
|
+
rest = AsyncSTX(
|
|
419
|
+
region=region,
|
|
420
|
+
env=env,
|
|
421
|
+
host=host,
|
|
422
|
+
key_id=key_id,
|
|
423
|
+
private_key=private_key,
|
|
424
|
+
signer=signer,
|
|
425
|
+
profile=profile,
|
|
426
|
+
verify_tls=verify_tls,
|
|
427
|
+
)
|
|
428
|
+
self._owns_rest = True
|
|
429
|
+
else:
|
|
430
|
+
self._owns_rest = False
|
|
431
|
+
self.rest = rest
|
|
432
|
+
self.url: str = rest.socket_url
|
|
433
|
+
self.credentials = rest.credentials
|
|
434
|
+
self.verify_tls: bool = rest.verify_tls
|
|
435
|
+
self._user_id = user_id
|
|
436
|
+
self.heartbeat_interval = heartbeat_interval
|
|
437
|
+
self.channel_ping_interval = channel_ping_interval
|
|
438
|
+
self.reconnect = reconnect
|
|
439
|
+
self.reconnect_policy = reconnect_policy or ReconnectPolicy()
|
|
440
|
+
self.on_reconnect = on_reconnect
|
|
441
|
+
self.join_timeout = join_timeout
|
|
442
|
+
self.queue_size = queue_size
|
|
443
|
+
|
|
444
|
+
self.channels: Dict[str, Channel] = {}
|
|
445
|
+
self._conn: Any = None
|
|
446
|
+
self._refs = itertools.count(1)
|
|
447
|
+
self._state = _State()
|
|
448
|
+
self._reader: Optional[asyncio.Task] = None
|
|
449
|
+
self._heartbeat: Optional[asyncio.Task] = None
|
|
450
|
+
self._supervisor: Optional[asyncio.Task] = None
|
|
451
|
+
# Created on first use inside the running loop (Python 3.9 binds
|
|
452
|
+
# asyncio primitives to the loop current at construction).
|
|
453
|
+
self._connected_ev: Optional[asyncio.Event] = None
|
|
454
|
+
self._closed = False
|
|
455
|
+
self._closed_ev: Optional[asyncio.Event] = None
|
|
456
|
+
self._pending_heartbeat: Optional[str] = None
|
|
457
|
+
self._tasks: set[asyncio.Task] = set()
|
|
458
|
+
self.reconnects = 0
|
|
459
|
+
|
|
460
|
+
def __repr__(self) -> str:
|
|
461
|
+
return f"STXWebSocket({self.url!r}, channels={list(self.channels)})"
|
|
462
|
+
|
|
463
|
+
# ------------------------------------------------------------------
|
|
464
|
+
# Lifecycle
|
|
465
|
+
# ------------------------------------------------------------------
|
|
466
|
+
|
|
467
|
+
async def __aenter__(self) -> STXWebSocket:
|
|
468
|
+
await self.connect()
|
|
469
|
+
return self
|
|
470
|
+
|
|
471
|
+
async def __aexit__(self, *exc: Any) -> None:
|
|
472
|
+
await self.close()
|
|
473
|
+
|
|
474
|
+
@property
|
|
475
|
+
def _connected(self) -> asyncio.Event:
|
|
476
|
+
if self._connected_ev is None:
|
|
477
|
+
self._connected_ev = asyncio.Event()
|
|
478
|
+
return self._connected_ev
|
|
479
|
+
|
|
480
|
+
@property
|
|
481
|
+
def _closed_event(self) -> asyncio.Event:
|
|
482
|
+
if self._closed_ev is None:
|
|
483
|
+
self._closed_ev = asyncio.Event()
|
|
484
|
+
return self._closed_ev
|
|
485
|
+
|
|
486
|
+
@property
|
|
487
|
+
def connected(self) -> bool:
|
|
488
|
+
return self._connected_ev is not None and self._connected_ev.is_set()
|
|
489
|
+
|
|
490
|
+
async def connect(self) -> None:
|
|
491
|
+
"""Open the socket. Idempotent."""
|
|
492
|
+
if self._closed:
|
|
493
|
+
raise STXTransportException("This STXWebSocket is closed; create a new one.")
|
|
494
|
+
if self._connected.is_set():
|
|
495
|
+
return
|
|
496
|
+
await self._open()
|
|
497
|
+
self._supervisor = asyncio.ensure_future(self._supervise())
|
|
498
|
+
|
|
499
|
+
async def close(self) -> None:
|
|
500
|
+
"""Leave every channel and close the socket. Safe to call twice."""
|
|
501
|
+
if self._closed:
|
|
502
|
+
return
|
|
503
|
+
self._closed = True
|
|
504
|
+
for channel in list(self.channels.values()):
|
|
505
|
+
channel.closed = True
|
|
506
|
+
if channel.joined and self._connected.is_set():
|
|
507
|
+
try:
|
|
508
|
+
await asyncio.wait_for(
|
|
509
|
+
self._request(channel.topic, "phx_leave", {}, channel.join_ref, 2.0), 3.0
|
|
510
|
+
)
|
|
511
|
+
except Exception:
|
|
512
|
+
pass
|
|
513
|
+
channel._close()
|
|
514
|
+
self.channels.clear()
|
|
515
|
+
for task in (self._supervisor, self._heartbeat, self._reader, *self._tasks):
|
|
516
|
+
if task is not None and not task.done():
|
|
517
|
+
task.cancel()
|
|
518
|
+
try:
|
|
519
|
+
await task
|
|
520
|
+
except (asyncio.CancelledError, Exception):
|
|
521
|
+
pass
|
|
522
|
+
await self._close_conn()
|
|
523
|
+
self._fail_pending(STXTransportException("socket closed"))
|
|
524
|
+
self._closed_event.set()
|
|
525
|
+
if self._owns_rest:
|
|
526
|
+
await self.rest.close()
|
|
527
|
+
|
|
528
|
+
async def run_forever(self) -> None:
|
|
529
|
+
"""Block until :meth:`close` is called (or reconnects are exhausted)."""
|
|
530
|
+
await self._closed_event.wait()
|
|
531
|
+
|
|
532
|
+
async def user_id(self) -> str:
|
|
533
|
+
"""Your user id, from ``user_id=`` or ``GET /api/v1/me``."""
|
|
534
|
+
if self._user_id is None:
|
|
535
|
+
self._user_id = await self.rest.user_id()
|
|
536
|
+
return self._user_id
|
|
537
|
+
|
|
538
|
+
# ------------------------------------------------------------------
|
|
539
|
+
# Public market channels
|
|
540
|
+
# ------------------------------------------------------------------
|
|
541
|
+
|
|
542
|
+
async def orderbook(
|
|
543
|
+
self, market_ids: Sequence[str], *, on_message: Optional[MessageHandler] = None
|
|
544
|
+
) -> Channel:
|
|
545
|
+
"""``orderbook``: the aggregated book, one ``book`` push per market.
|
|
546
|
+
|
|
547
|
+
Each push is a full snapshot of that market's book; replace what you
|
|
548
|
+
hold rather than merging. ``market_ids`` is required.
|
|
549
|
+
"""
|
|
550
|
+
if not market_ids:
|
|
551
|
+
raise ValueError("orderbook needs at least one market id")
|
|
552
|
+
return await self.join("orderbook", {"market_ids": list(market_ids)}, on_message)
|
|
553
|
+
|
|
554
|
+
async def ticker(
|
|
555
|
+
self,
|
|
556
|
+
*,
|
|
557
|
+
sports: Optional[Sequence[str]] = None,
|
|
558
|
+
competitions: Optional[Sequence[str]] = None,
|
|
559
|
+
on_message: Optional[MessageHandler] = None,
|
|
560
|
+
) -> Channel:
|
|
561
|
+
"""``ticker``: a ``ticker`` push whenever a market's price, top of
|
|
562
|
+
book, volume or open interest moves. No snapshot on join."""
|
|
563
|
+
payload = _drop_none({"sports": _opt_list(sports), "competitions": _opt_list(competitions)})
|
|
564
|
+
return await self.join("ticker", payload, on_message)
|
|
565
|
+
|
|
566
|
+
async def trades(
|
|
567
|
+
self,
|
|
568
|
+
*,
|
|
569
|
+
market_ids: Optional[Sequence[str]] = None,
|
|
570
|
+
event_ids: Optional[Sequence[str]] = None,
|
|
571
|
+
on_message: Optional[MessageHandler] = None,
|
|
572
|
+
) -> Channel:
|
|
573
|
+
"""``trades``: every execution on the exchange, anonymised. ``action``
|
|
574
|
+
is the taker's side. Not your fills: see :meth:`fills`."""
|
|
575
|
+
payload = _drop_none(
|
|
576
|
+
{"market_ids": _opt_list(market_ids), "event_ids": _opt_list(event_ids)}
|
|
577
|
+
)
|
|
578
|
+
return await self.join("trades", payload, on_message)
|
|
579
|
+
|
|
580
|
+
async def markets(
|
|
581
|
+
self,
|
|
582
|
+
*,
|
|
583
|
+
rule_filters: Optional[Sequence[str]] = None,
|
|
584
|
+
message_types: Optional[Sequence[str]] = None,
|
|
585
|
+
on_message: Optional[MessageHandler] = None,
|
|
586
|
+
) -> Channel:
|
|
587
|
+
"""``markets``: ``market_created`` and ``market_updated`` for every
|
|
588
|
+
market. Each payload maps market id to a market object;
|
|
589
|
+
``market_updated`` carries only the changed fields. Prices arrive in
|
|
590
|
+
cents on the wire and are converted to dollar strings here."""
|
|
591
|
+
payload = _drop_none(
|
|
592
|
+
{"rule_filters": _opt_list(rule_filters), "message_types": _opt_list(message_types)}
|
|
593
|
+
)
|
|
594
|
+
return await self.join("markets", payload, on_message)
|
|
595
|
+
|
|
596
|
+
async def market_stats(
|
|
597
|
+
self,
|
|
598
|
+
market_ids: Sequence[str],
|
|
599
|
+
*,
|
|
600
|
+
range: Optional[str] = None,
|
|
601
|
+
on_message: Optional[MessageHandler] = None,
|
|
602
|
+
) -> Channel:
|
|
603
|
+
"""``market_stats``: a price series per market. The history is in the
|
|
604
|
+
join reply (``channel.reply["markets"]``); ``market_stats`` pushes
|
|
605
|
+
changed buckets (upsert by ``timestamp_us``) and
|
|
606
|
+
``market_stats_snapshot`` replaces a series."""
|
|
607
|
+
if not market_ids:
|
|
608
|
+
raise ValueError("market_stats needs at least one market id")
|
|
609
|
+
payload = _drop_none({"market_ids": list(market_ids), "range": range})
|
|
610
|
+
return await self.join("market_stats", payload, on_message)
|
|
611
|
+
|
|
612
|
+
async def market_updates(
|
|
613
|
+
self,
|
|
614
|
+
watch: Optional[Sequence[str]] = None,
|
|
615
|
+
*,
|
|
616
|
+
on_message: Optional[MessageHandler] = None,
|
|
617
|
+
) -> Channel:
|
|
618
|
+
"""``market_updates``: ``created`` and ``updated`` for the markets you
|
|
619
|
+
watch. Nothing arrives until you watch something; pass ``watch=`` or
|
|
620
|
+
call ``channel.watch([...])``. Prices are converted from cents to
|
|
621
|
+
dollar strings here."""
|
|
622
|
+
channel = await self.join("market_updates", {}, on_message)
|
|
623
|
+
if watch:
|
|
624
|
+
await channel.watch(watch)
|
|
625
|
+
return channel
|
|
626
|
+
|
|
627
|
+
# ------------------------------------------------------------------
|
|
628
|
+
# Account channels
|
|
629
|
+
# ------------------------------------------------------------------
|
|
630
|
+
|
|
631
|
+
async def orders(
|
|
632
|
+
self,
|
|
633
|
+
*,
|
|
634
|
+
market_ids: Optional[Sequence[str]] = None,
|
|
635
|
+
cancel_on_disconnect: bool = False,
|
|
636
|
+
ping_timeout: Optional[int] = None,
|
|
637
|
+
on_message: Optional[MessageHandler] = None,
|
|
638
|
+
) -> Channel:
|
|
639
|
+
"""``orders:{user_id}``: ``all_orders`` on join, then ``new_open_order``.
|
|
640
|
+
|
|
641
|
+
``cancel_on_disconnect=True`` arms cancel-on-disconnect for orders
|
|
642
|
+
placed with ``cancel_on_disconnect=True``. ``ping_timeout`` is in
|
|
643
|
+
milliseconds, clamped by the server to 5000 to 20000; the granted
|
|
644
|
+
value is in ``channel.reply["ping_timeout"]`` and the SDK pings at
|
|
645
|
+
60% of it for as long as the channel is joined.
|
|
646
|
+
"""
|
|
647
|
+
payload: Dict[str, Any] = _drop_none({"market_ids": _opt_list(market_ids)})
|
|
648
|
+
if cancel_on_disconnect:
|
|
649
|
+
payload["cancel_on_disconnect"] = True
|
|
650
|
+
payload["ping_timeout"] = int(ping_timeout if ping_timeout is not None else 5000)
|
|
651
|
+
return await self._join_account("orders", payload, on_message)
|
|
652
|
+
|
|
653
|
+
async def fills(
|
|
654
|
+
self,
|
|
655
|
+
*,
|
|
656
|
+
market_ids: Optional[Sequence[str]] = None,
|
|
657
|
+
on_message: Optional[MessageHandler] = None,
|
|
658
|
+
) -> Channel:
|
|
659
|
+
"""``fills:{user_id}``: ``all_trades`` on join, then one ``trade`` per
|
|
660
|
+
execution or status change."""
|
|
661
|
+
payload = _drop_none({"market_ids": _opt_list(market_ids)})
|
|
662
|
+
return await self._join_account("fills", payload, on_message)
|
|
663
|
+
|
|
664
|
+
async def positions(
|
|
665
|
+
self,
|
|
666
|
+
*,
|
|
667
|
+
market_ids: Optional[Sequence[str]] = None,
|
|
668
|
+
on_message: Optional[MessageHandler] = None,
|
|
669
|
+
) -> Channel:
|
|
670
|
+
"""``positions:{user_id}``: ``all_positions`` on join, then
|
|
671
|
+
``updated_positions`` deltas with only the changed positions."""
|
|
672
|
+
payload = _drop_none({"market_ids": _opt_list(market_ids)})
|
|
673
|
+
return await self._join_account("positions", payload, on_message)
|
|
674
|
+
|
|
675
|
+
async def settlements(
|
|
676
|
+
self,
|
|
677
|
+
*,
|
|
678
|
+
market_ids: Optional[Sequence[str]] = None,
|
|
679
|
+
on_message: Optional[MessageHandler] = None,
|
|
680
|
+
) -> Channel:
|
|
681
|
+
"""``settlements:{user_id}``: ``new_settlements`` as they are recorded.
|
|
682
|
+
No snapshot; history is ``AsyncSTX.settlements()``."""
|
|
683
|
+
payload = _drop_none({"market_ids": _opt_list(market_ids)})
|
|
684
|
+
return await self._join_account("settlements", payload, on_message)
|
|
685
|
+
|
|
686
|
+
async def balances(
|
|
687
|
+
self,
|
|
688
|
+
*,
|
|
689
|
+
account_id: Optional[str] = None,
|
|
690
|
+
on_message: Optional[MessageHandler] = None,
|
|
691
|
+
) -> Channel:
|
|
692
|
+
"""``balances:{user_id}``: ``balances`` on join, then ``update`` and
|
|
693
|
+
``payment_update``. ``account_id`` picks one of your accounts."""
|
|
694
|
+
payload = _drop_none({"account_id": account_id})
|
|
695
|
+
return await self._join_account("balances", payload, on_message)
|
|
696
|
+
|
|
697
|
+
async def account(
|
|
698
|
+
self,
|
|
699
|
+
*,
|
|
700
|
+
market_ids: Optional[Sequence[str]] = None,
|
|
701
|
+
on_message: Optional[MessageHandler] = None,
|
|
702
|
+
) -> Channel:
|
|
703
|
+
"""``account:{user_id}``: everything the five channels above carry, on
|
|
704
|
+
one join. Do not also join a per-type channel (you would get every
|
|
705
|
+
message twice), and use ``orders`` if you need cancel-on-disconnect."""
|
|
706
|
+
payload = _drop_none({"market_ids": _opt_list(market_ids)})
|
|
707
|
+
return await self._join_account("account", payload, on_message)
|
|
708
|
+
|
|
709
|
+
async def user_info(self, *, on_message: Optional[MessageHandler] = None) -> Channel:
|
|
710
|
+
"""``user_info:{user_id}``: ``user_updated`` right after joining, then
|
|
711
|
+
on every profile change."""
|
|
712
|
+
return await self._join_account("user_info", {}, on_message)
|
|
713
|
+
|
|
714
|
+
# ------------------------------------------------------------------
|
|
715
|
+
# Generic join
|
|
716
|
+
# ------------------------------------------------------------------
|
|
717
|
+
|
|
718
|
+
async def join(
|
|
719
|
+
self,
|
|
720
|
+
topic: str,
|
|
721
|
+
payload: Optional[Dict[str, Any]] = None,
|
|
722
|
+
on_message: Optional[MessageHandler] = None,
|
|
723
|
+
*,
|
|
724
|
+
ping_interval: Any = _UNSET,
|
|
725
|
+
) -> Channel:
|
|
726
|
+
"""Join ``topic`` with ``payload`` and wait for the reply.
|
|
727
|
+
|
|
728
|
+
Raises ``STXChannelException`` with the server's reason (for example
|
|
729
|
+
``market_ids_required`` or ``unauthorized``) if the join is refused.
|
|
730
|
+
"""
|
|
731
|
+
if topic in self.channels and not self.channels[topic].closed:
|
|
732
|
+
raise STXChannelException(f"Already joined {topic}", topic=topic)
|
|
733
|
+
await self.connect()
|
|
734
|
+
interval = self.channel_ping_interval if ping_interval is _UNSET else ping_interval
|
|
735
|
+
channel = Channel(
|
|
736
|
+
self, topic, payload or {}, on_message,
|
|
737
|
+
ping_interval=interval, queue_size=self.queue_size,
|
|
738
|
+
) # fmt: skip
|
|
739
|
+
self.channels[topic] = channel
|
|
740
|
+
try:
|
|
741
|
+
await self._join(channel)
|
|
742
|
+
except BaseException:
|
|
743
|
+
self.channels.pop(topic, None)
|
|
744
|
+
channel._close()
|
|
745
|
+
raise
|
|
746
|
+
return channel
|
|
747
|
+
|
|
748
|
+
async def _join_account(
|
|
749
|
+
self, name: str, payload: Dict[str, Any], on_message: Optional[MessageHandler]
|
|
750
|
+
) -> Channel:
|
|
751
|
+
if self.credentials is None:
|
|
752
|
+
raise STXConfigException(
|
|
753
|
+
f"The {name} channel is private and needs an API key; the socket "
|
|
754
|
+
"must be signed. Configure key_id and private_key."
|
|
755
|
+
)
|
|
756
|
+
return await self.join(f"{name}:{await self.user_id()}", payload, on_message)
|
|
757
|
+
|
|
758
|
+
async def _join(self, channel: Channel) -> None:
|
|
759
|
+
ref = str(next(self._refs))
|
|
760
|
+
channel.join_ref = ref
|
|
761
|
+
channel._reset_for_rejoin()
|
|
762
|
+
# A join's ref doubles as its join_ref, as Phoenix clients do.
|
|
763
|
+
reply = await self._request(channel.topic, "phx_join", channel.join_payload, ref, ref=ref)
|
|
764
|
+
channel.reply = reply if isinstance(reply, dict) else {"response": reply}
|
|
765
|
+
channel.joined = True
|
|
766
|
+
if channel.watches:
|
|
767
|
+
await channel.push("watch", list(channel.watches))
|
|
768
|
+
channel._start_pinger()
|
|
769
|
+
|
|
770
|
+
async def _leave(self, channel: Channel) -> None:
|
|
771
|
+
# Marked closed first so the server's phx_close is not taken as a
|
|
772
|
+
# reason to rejoin.
|
|
773
|
+
channel.closed = True
|
|
774
|
+
if channel.joined and self._connected.is_set():
|
|
775
|
+
try:
|
|
776
|
+
await self._request(channel.topic, "phx_leave", {}, channel.join_ref, 5.0)
|
|
777
|
+
except Exception:
|
|
778
|
+
pass
|
|
779
|
+
channel._close()
|
|
780
|
+
self.channels.pop(channel.topic, None)
|
|
781
|
+
|
|
782
|
+
# ------------------------------------------------------------------
|
|
783
|
+
# Wire
|
|
784
|
+
# ------------------------------------------------------------------
|
|
785
|
+
|
|
786
|
+
async def _open(self) -> None:
|
|
787
|
+
# A fresh signature per connect: the timestamp must be within 30 s.
|
|
788
|
+
headers = {"User-Agent": USER_AGENT}
|
|
789
|
+
if self.credentials is not None:
|
|
790
|
+
headers.update(self.credentials.headers("GET", "/socket/websocket"))
|
|
791
|
+
kwargs: Dict[str, Any] = {_handshake_headers_kwarg(): headers, "max_size": None}
|
|
792
|
+
if self.url.startswith("wss://") and not self.verify_tls:
|
|
793
|
+
import ssl
|
|
794
|
+
|
|
795
|
+
ctx = ssl.create_default_context()
|
|
796
|
+
ctx.check_hostname = False
|
|
797
|
+
ctx.verify_mode = ssl.CERT_NONE
|
|
798
|
+
kwargs["ssl"] = ctx
|
|
799
|
+
try:
|
|
800
|
+
self._conn = await asyncio.wait_for(websockets.connect(self.url, **kwargs), 15.0)
|
|
801
|
+
except Exception as exc:
|
|
802
|
+
raise STXTransportException(f"Could not open {self.url}: {exc}") from exc
|
|
803
|
+
self._pending_heartbeat = None
|
|
804
|
+
self._connected.set()
|
|
805
|
+
self._reader = asyncio.ensure_future(self._read_loop())
|
|
806
|
+
self._heartbeat = asyncio.ensure_future(self._heartbeat_loop())
|
|
807
|
+
logger.debug("connected to %s", self.url)
|
|
808
|
+
|
|
809
|
+
async def _close_conn(self) -> None:
|
|
810
|
+
self._connected.clear()
|
|
811
|
+
conn, self._conn = self._conn, None
|
|
812
|
+
if conn is not None:
|
|
813
|
+
try:
|
|
814
|
+
await conn.close()
|
|
815
|
+
except Exception:
|
|
816
|
+
pass
|
|
817
|
+
|
|
818
|
+
async def _send(self, frame: List[Any]) -> None:
|
|
819
|
+
if self._conn is None or not self._connected.is_set():
|
|
820
|
+
raise STXTransportException("socket is not connected")
|
|
821
|
+
await self._conn.send(json.dumps(frame))
|
|
822
|
+
|
|
823
|
+
async def _request(
|
|
824
|
+
self,
|
|
825
|
+
topic: str,
|
|
826
|
+
event: str,
|
|
827
|
+
payload: Any,
|
|
828
|
+
join_ref: Optional[str],
|
|
829
|
+
timeout: Optional[float] = None,
|
|
830
|
+
ref: Optional[str] = None,
|
|
831
|
+
) -> Any:
|
|
832
|
+
ref = ref or str(next(self._refs))
|
|
833
|
+
future: asyncio.Future[Any] = asyncio.get_running_loop().create_future()
|
|
834
|
+
self._state.pending[ref] = _Pending(future, topic, event)
|
|
835
|
+
try:
|
|
836
|
+
await self._send([join_ref, ref, topic, event, payload])
|
|
837
|
+
return await asyncio.wait_for(future, timeout or self.join_timeout)
|
|
838
|
+
except asyncio.TimeoutError:
|
|
839
|
+
raise STXChannelException(
|
|
840
|
+
f"No reply to {event} on {topic} within {timeout or self.join_timeout}s",
|
|
841
|
+
topic=topic,
|
|
842
|
+
) from None
|
|
843
|
+
finally:
|
|
844
|
+
self._state.pending.pop(ref, None)
|
|
845
|
+
|
|
846
|
+
def _fail_pending(self, exc: Exception) -> None:
|
|
847
|
+
for pending in list(self._state.pending.values()):
|
|
848
|
+
if not pending.future.done():
|
|
849
|
+
pending.future.set_exception(exc)
|
|
850
|
+
self._state.pending.clear()
|
|
851
|
+
|
|
852
|
+
async def _read_loop(self) -> None:
|
|
853
|
+
conn = self._conn
|
|
854
|
+
try:
|
|
855
|
+
async for raw in conn:
|
|
856
|
+
try:
|
|
857
|
+
frame = json.loads(raw)
|
|
858
|
+
join_ref, ref, topic, event, payload = frame
|
|
859
|
+
except (ValueError, TypeError):
|
|
860
|
+
logger.warning("unparseable frame: %r", raw[:200])
|
|
861
|
+
continue
|
|
862
|
+
await self._dispatch(join_ref, ref, topic, event, payload)
|
|
863
|
+
except asyncio.CancelledError:
|
|
864
|
+
raise
|
|
865
|
+
except Exception as exc:
|
|
866
|
+
logger.info("socket read ended: %s", exc)
|
|
867
|
+
finally:
|
|
868
|
+
if self._conn is conn:
|
|
869
|
+
self._connected.clear()
|
|
870
|
+
|
|
871
|
+
async def _dispatch(
|
|
872
|
+
self, join_ref: Any, ref: Any, topic: str, event: str, payload: Any
|
|
873
|
+
) -> None:
|
|
874
|
+
if event == "phx_reply":
|
|
875
|
+
if topic == PHOENIX_TOPIC and ref == self._pending_heartbeat:
|
|
876
|
+
self._pending_heartbeat = None
|
|
877
|
+
pending = self._state.pending.get(str(ref)) if ref is not None else None
|
|
878
|
+
if pending is not None and not pending.future.done():
|
|
879
|
+
status = payload.get("status") if isinstance(payload, dict) else None
|
|
880
|
+
response = payload.get("response") if isinstance(payload, dict) else payload
|
|
881
|
+
if status == "ok":
|
|
882
|
+
pending.future.set_result(response)
|
|
883
|
+
else:
|
|
884
|
+
reason = response.get("reason") if isinstance(response, dict) else None
|
|
885
|
+
pending.future.set_exception(
|
|
886
|
+
STXChannelException(
|
|
887
|
+
f"{pending.event} on {topic} failed: {reason or response}",
|
|
888
|
+
topic=topic,
|
|
889
|
+
reply=response,
|
|
890
|
+
)
|
|
891
|
+
)
|
|
892
|
+
return
|
|
893
|
+
channel = self.channels.get(topic)
|
|
894
|
+
if channel is None:
|
|
895
|
+
return
|
|
896
|
+
if (
|
|
897
|
+
join_ref is not None
|
|
898
|
+
and channel.join_ref is not None
|
|
899
|
+
and str(join_ref) != channel.join_ref
|
|
900
|
+
):
|
|
901
|
+
return # a frame for an earlier join of this topic
|
|
902
|
+
if event in ("phx_error", "phx_close"):
|
|
903
|
+
if not channel.closed and not self._closed:
|
|
904
|
+
logger.warning("%s: server sent %s; rejoining", topic, event)
|
|
905
|
+
channel._reset_for_rejoin()
|
|
906
|
+
task = asyncio.ensure_future(self._rejoin_later(channel))
|
|
907
|
+
self._tasks.add(task)
|
|
908
|
+
task.add_done_callback(self._tasks.discard)
|
|
909
|
+
return
|
|
910
|
+
if channel.name == "markets":
|
|
911
|
+
payload = convert_markets_frame(payload)
|
|
912
|
+
elif channel.name == "market_updates" and isinstance(payload, dict):
|
|
913
|
+
payload = convert_market_payload(payload)
|
|
914
|
+
await channel._deliver(ChannelMessage(topic, event, payload, ref, join_ref))
|
|
915
|
+
|
|
916
|
+
async def _rejoin_later(self, channel: Channel) -> None:
|
|
917
|
+
attempt = 1
|
|
918
|
+
while not channel.closed and not self._closed:
|
|
919
|
+
await asyncio.sleep(self.reconnect_policy.delay(attempt))
|
|
920
|
+
if not self._connected.is_set():
|
|
921
|
+
return # the socket reconnect will rejoin it
|
|
922
|
+
try:
|
|
923
|
+
await self._join(channel)
|
|
924
|
+
return
|
|
925
|
+
except STXChannelException as exc:
|
|
926
|
+
logger.error("rejoin %s failed: %s", channel.topic, exc)
|
|
927
|
+
if isinstance(exc.reply, dict) and exc.reply.get("reason") == "unauthorized":
|
|
928
|
+
channel._close()
|
|
929
|
+
return
|
|
930
|
+
attempt += 1
|
|
931
|
+
|
|
932
|
+
async def _heartbeat_loop(self) -> None:
|
|
933
|
+
while self._connected.is_set():
|
|
934
|
+
await asyncio.sleep(self.heartbeat_interval)
|
|
935
|
+
if self._pending_heartbeat is not None:
|
|
936
|
+
logger.warning("heartbeat unanswered; reconnecting")
|
|
937
|
+
await self._close_conn()
|
|
938
|
+
return
|
|
939
|
+
ref = str(next(self._refs))
|
|
940
|
+
self._pending_heartbeat = ref
|
|
941
|
+
try:
|
|
942
|
+
await self._send([None, ref, PHOENIX_TOPIC, "heartbeat", {}])
|
|
943
|
+
except Exception:
|
|
944
|
+
await self._close_conn()
|
|
945
|
+
return
|
|
946
|
+
|
|
947
|
+
async def _supervise(self) -> None:
|
|
948
|
+
# Waits for the socket to drop, then reconnects and rejoins.
|
|
949
|
+
while not self._closed:
|
|
950
|
+
reader = self._reader
|
|
951
|
+
if reader is not None:
|
|
952
|
+
try:
|
|
953
|
+
await reader
|
|
954
|
+
except asyncio.CancelledError:
|
|
955
|
+
if self._closed:
|
|
956
|
+
return
|
|
957
|
+
except Exception:
|
|
958
|
+
pass
|
|
959
|
+
if self._closed:
|
|
960
|
+
return
|
|
961
|
+
await self._close_conn()
|
|
962
|
+
if self._heartbeat is not None:
|
|
963
|
+
self._heartbeat.cancel()
|
|
964
|
+
self._fail_pending(STXTransportException("socket dropped"))
|
|
965
|
+
for channel in self.channels.values():
|
|
966
|
+
channel._reset_for_rejoin()
|
|
967
|
+
if not self.reconnect:
|
|
968
|
+
await self._shutdown_after_drop()
|
|
969
|
+
return
|
|
970
|
+
if not await self._reconnect():
|
|
971
|
+
await self._shutdown_after_drop()
|
|
972
|
+
return
|
|
973
|
+
|
|
974
|
+
async def _reconnect(self) -> bool:
|
|
975
|
+
attempt = 1
|
|
976
|
+
while not self._closed:
|
|
977
|
+
policy = self.reconnect_policy
|
|
978
|
+
if policy.max_attempts is not None and attempt > policy.max_attempts:
|
|
979
|
+
logger.error("giving up after %d reconnect attempts", attempt - 1)
|
|
980
|
+
return False
|
|
981
|
+
await asyncio.sleep(policy.delay(attempt))
|
|
982
|
+
try:
|
|
983
|
+
await self._open()
|
|
984
|
+
for channel in list(self.channels.values()):
|
|
985
|
+
if channel.closed:
|
|
986
|
+
continue
|
|
987
|
+
try:
|
|
988
|
+
await self._join(channel)
|
|
989
|
+
except STXChannelException as exc:
|
|
990
|
+
logger.error("rejoin %s failed: %s", channel.topic, exc)
|
|
991
|
+
self.reconnects += 1
|
|
992
|
+
logger.info("reconnected (%d)", self.reconnects)
|
|
993
|
+
if self.on_reconnect is not None:
|
|
994
|
+
result = self.on_reconnect()
|
|
995
|
+
if inspect.isawaitable(result):
|
|
996
|
+
await result
|
|
997
|
+
return True
|
|
998
|
+
except STXTransportException as exc:
|
|
999
|
+
logger.warning("reconnect attempt %d failed: %s", attempt, exc)
|
|
1000
|
+
await self._close_conn()
|
|
1001
|
+
attempt += 1
|
|
1002
|
+
return False
|
|
1003
|
+
|
|
1004
|
+
async def _shutdown_after_drop(self) -> None:
|
|
1005
|
+
for channel in list(self.channels.values()):
|
|
1006
|
+
channel._close()
|
|
1007
|
+
self._closed = True
|
|
1008
|
+
self._closed_event.set()
|
|
1009
|
+
|
|
1010
|
+
|
|
1011
|
+
def _opt_list(value: Optional[Sequence[str]]) -> Optional[List[str]]:
|
|
1012
|
+
return None if value is None else list(value)
|
|
1013
|
+
|
|
1014
|
+
|
|
1015
|
+
def _drop_none(payload: Dict[str, Any]) -> Dict[str, Any]:
|
|
1016
|
+
return {k: v for k, v in payload.items() if v is not None}
|