Split exchange and strategy modules for future Binance support.
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
@@ -1,207 +1,5 @@
|
||||
"""OKX 公共 WebSocket:永续 + 期权 books5 / mark-price。只读。"""
|
||||
"""兼容层:OKX WS 在 exchange.okx.ws。"""
|
||||
|
||||
from __future__ import annotations
|
||||
from ..exchange.okx.ws import OkxPublicWs
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
from typing import Any
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import websockets
|
||||
from websockets.asyncio.client import ClientConnection
|
||||
|
||||
from .book_cache import BookCache
|
||||
from .instruments import safe_float
|
||||
from .types import BookLevel
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class OkxPublicWs:
|
||||
def __init__(
|
||||
self,
|
||||
url: str,
|
||||
cache: BookCache,
|
||||
*,
|
||||
proxy: str | None = None,
|
||||
ping_interval: float = 20.0,
|
||||
) -> None:
|
||||
self.url = url
|
||||
self.cache = cache
|
||||
self.proxy = (proxy or "").strip() or None
|
||||
self.ping_interval = ping_interval
|
||||
self._inst_ids: list[str] = []
|
||||
self._task: asyncio.Task[None] | None = None
|
||||
self._stop = asyncio.Event()
|
||||
self._subscribed: set[str] = set()
|
||||
|
||||
def set_instruments(self, inst_ids: list[str]) -> None:
|
||||
self._inst_ids = [i for i in inst_ids if i]
|
||||
|
||||
async def start(self) -> None:
|
||||
if self._task and not self._task.done():
|
||||
return
|
||||
self._stop.clear()
|
||||
self._task = asyncio.create_task(self._run_forever(), name="okx-public-ws")
|
||||
|
||||
async def stop(self) -> None:
|
||||
self._stop.set()
|
||||
if self._task:
|
||||
self._task.cancel()
|
||||
try:
|
||||
await self._task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
self._task = None
|
||||
self.cache.set_connected(False)
|
||||
|
||||
async def resubscribe(self, inst_ids: list[str]) -> None:
|
||||
self.set_instruments(inst_ids)
|
||||
self._stop.set()
|
||||
await asyncio.sleep(0)
|
||||
self._stop.clear()
|
||||
if self._task and not self._task.done():
|
||||
self._task.cancel()
|
||||
try:
|
||||
await self._task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
self._task = asyncio.create_task(self._run_forever(), name="okx-public-ws")
|
||||
|
||||
async def _open_connection(self) -> ClientConnection:
|
||||
if not self.proxy:
|
||||
return await websockets.connect(
|
||||
self.url,
|
||||
ping_interval=None,
|
||||
max_size=2**22,
|
||||
open_timeout=20,
|
||||
)
|
||||
|
||||
from python_socks.async_.asyncio import Proxy
|
||||
|
||||
parsed = urlparse(self.url)
|
||||
host = parsed.hostname or "ws.okx.com"
|
||||
port = parsed.port or (443 if parsed.scheme == "wss" else 80)
|
||||
sock = await Proxy.from_url(self.proxy).connect(dest_host=host, dest_port=port)
|
||||
return await websockets.connect(
|
||||
self.url,
|
||||
sock=sock,
|
||||
server_hostname=host,
|
||||
ping_interval=None,
|
||||
max_size=2**22,
|
||||
open_timeout=20,
|
||||
)
|
||||
|
||||
async def _run_forever(self) -> None:
|
||||
backoff = 1.0
|
||||
while not self._stop.is_set():
|
||||
try:
|
||||
async with await self._open_connection() as ws:
|
||||
self.cache.set_connected(True)
|
||||
backoff = 1.0
|
||||
await self._subscribe(ws)
|
||||
waiter = asyncio.create_task(self._stop.wait())
|
||||
reader = asyncio.create_task(self._read_loop(ws))
|
||||
pinger = asyncio.create_task(self._ping_loop(ws))
|
||||
done, pending = await asyncio.wait(
|
||||
{waiter, reader, pinger},
|
||||
return_when=asyncio.FIRST_COMPLETED,
|
||||
)
|
||||
for t in pending:
|
||||
t.cancel()
|
||||
for t in done:
|
||||
exc = t.exception()
|
||||
if exc and not isinstance(exc, asyncio.CancelledError):
|
||||
raise exc
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.warning("OKX WS disconnected: %s", e)
|
||||
self.cache.set_connected(False)
|
||||
try:
|
||||
await asyncio.wait_for(self._stop.wait(), timeout=backoff)
|
||||
break
|
||||
except asyncio.TimeoutError:
|
||||
backoff = min(backoff * 2, 30.0)
|
||||
|
||||
self.cache.set_connected(False)
|
||||
|
||||
async def _subscribe(self, ws: ClientConnection) -> None:
|
||||
args: list[dict[str, str]] = []
|
||||
for inst in self._inst_ids:
|
||||
args.append({"channel": "books5", "instId": inst})
|
||||
args.append({"channel": "mark-price", "instId": inst})
|
||||
if not args:
|
||||
return
|
||||
payload = {"op": "subscribe", "args": args}
|
||||
await ws.send(json.dumps(payload))
|
||||
self._subscribed = {a["instId"] for a in args}
|
||||
logger.info("OKX WS subscribed: %s", sorted(self._subscribed))
|
||||
|
||||
async def _ping_loop(self, ws: ClientConnection) -> None:
|
||||
while True:
|
||||
await asyncio.sleep(self.ping_interval)
|
||||
await ws.send("ping")
|
||||
|
||||
async def _read_loop(self, ws: ClientConnection) -> None:
|
||||
try:
|
||||
async for raw in ws:
|
||||
if raw == "pong":
|
||||
continue
|
||||
if isinstance(raw, bytes):
|
||||
raw = raw.decode("utf-8", errors="ignore")
|
||||
if raw == "ping":
|
||||
await ws.send("pong")
|
||||
continue
|
||||
try:
|
||||
msg = json.loads(raw)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
self._handle_message(msg)
|
||||
except websockets.exceptions.ConnectionClosed:
|
||||
return
|
||||
|
||||
def _handle_message(self, msg: dict[str, Any]) -> None:
|
||||
if msg.get("event") in ("subscribe", "error", "channel-conn-count"):
|
||||
if msg.get("event") == "error":
|
||||
logger.error("OKX WS error: %s", msg)
|
||||
return
|
||||
arg = msg.get("arg") or {}
|
||||
channel = str(arg.get("channel") or "")
|
||||
inst_id = str(arg.get("instId") or "")
|
||||
data = msg.get("data") or []
|
||||
if not inst_id or not data:
|
||||
return
|
||||
row = data[0] if isinstance(data[0], dict) else None
|
||||
if row is None:
|
||||
return
|
||||
|
||||
if channel == "books5":
|
||||
ts = safe_float(row.get("ts"))
|
||||
self.cache.upsert_book(
|
||||
inst_id,
|
||||
bids=_levels(row.get("bids") or []),
|
||||
asks=_levels(row.get("asks") or []),
|
||||
ts_ms=int(ts) if ts is not None else None,
|
||||
)
|
||||
elif channel == "mark-price":
|
||||
ts = safe_float(row.get("ts"))
|
||||
self.cache.set_mark_px(
|
||||
inst_id,
|
||||
safe_float(row.get("markPx")),
|
||||
ts_ms=int(ts) if ts is not None else None,
|
||||
)
|
||||
|
||||
|
||||
def _levels(raw: list[Any]) -> list[BookLevel]:
|
||||
out: list[BookLevel] = []
|
||||
for item in raw:
|
||||
if not isinstance(item, (list, tuple)) or len(item) < 2:
|
||||
continue
|
||||
px = safe_float(item[0])
|
||||
sz = safe_float(item[1])
|
||||
if px is None or sz is None or px <= 0 or sz <= 0:
|
||||
continue
|
||||
out.append(BookLevel(px=px, sz=sz))
|
||||
return out
|
||||
__all__ = ["OkxPublicWs"]
|
||||
|
||||
Reference in New Issue
Block a user