stx-python 0.6.0rc1__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/_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}