Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -110,6 +110,8 @@ passphrase = ""
- WebSocketAPI
- Run test/WsPrivateTest.py for private websocket channels
- Run test/WsPublicTest.py for public websocket channels
- Pass `proxy="http://127.0.0.1:7890"` to `WsPublicAsync` or
`WsPrivateAsync` when the WebSocket connection must use an explicit proxy
- Use different URLs for different environment
- Live trading URLs: https://www.okx.com/docs-v5/en/#overview-production-trading-services
- Demo trading URLs: https://www.okx.com/docs-v5/en/#overview-demo-trading-services
Expand Down
12 changes: 7 additions & 5 deletions okx/websocket/WebSocketFactory.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,3 @@
import asyncio
import logging
import ssl

Expand All @@ -10,21 +9,24 @@

class WebSocketFactory:

def __init__(self, url):
def __init__(self, url, proxy=None):
self.url = url
self.proxy = proxy
self.websocket = None
self.loop = asyncio.get_event_loop()

async def connect(self):
ssl_context = ssl.create_default_context()
ssl_context.load_verify_locations(certifi.where())
try:
self.websocket = await websockets.connect(self.url, ssl=ssl_context)
connect_kwargs = {"ssl": ssl_context}
if self.proxy is not None:
connect_kwargs["proxy"] = self.proxy
self.websocket = await websockets.connect(self.url, **connect_kwargs)
logger.info("WebSocket connection established.")
return self.websocket
except Exception as e:
logger.error(f"Error connecting to WebSocket: {e}")
return None
raise

async def close(self):
if self.websocket:
Expand Down
11 changes: 6 additions & 5 deletions okx/websocket/WsPrivateAsync.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,12 +12,12 @@


class WsPrivateAsync:
def __init__(self, apiKey, passphrase, secretKey, url, useServerTime=None, debug=False):
def __init__(self, apiKey, passphrase, secretKey, url, useServerTime=None, debug=False, proxy=None):
self.url = url
self.subscriptions = set()
self.callback = None
self.loop = asyncio.get_event_loop()
self.factory = WebSocketFactory(url)
self.loop = None
self.factory = WebSocketFactory(url, proxy=proxy)
self.apiKey = apiKey
self.passphrase = passphrase
self.secretKey = secretKey
Expand Down Expand Up @@ -201,12 +201,13 @@ async def start(self):
logger.debug("Connecting to WebSocket...")
else:
logger.info("Connecting to WebSocket...")
self.loop = asyncio.get_running_loop()
await self.connect()
return self.loop.create_task(self.consume())

def stop_sync(self):
if self.loop.is_running():
if self.loop is not None and self.loop.is_running():
future = asyncio.run_coroutine_threadsafe(self.stop(), self.loop)
future.result(timeout=10)
else:
self.loop.run_until_complete(self.stop())
asyncio.run(self.stop())
11 changes: 6 additions & 5 deletions okx/websocket/WsPublicAsync.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,12 +11,12 @@


class WsPublicAsync:
def __init__(self, url, apiKey='', passphrase='', secretKey='', debug=False):
def __init__(self, url, apiKey='', passphrase='', secretKey='', debug=False, proxy=None):
self.url = url
self.subscriptions = set()
self.callback = None
self.loop = asyncio.get_event_loop()
self.factory = WebSocketFactory(url)
self.loop = None
self.factory = WebSocketFactory(url, proxy=proxy)
self.websocket = None
self.debug = debug
# Credentials for business channel login
Expand Down Expand Up @@ -122,12 +122,13 @@ async def start(self):
logger.debug("Connecting to WebSocket...")
else:
logger.info("Connecting to WebSocket...")
self.loop = asyncio.get_running_loop()
await self.connect()
return self.loop.create_task(self.consume())

def stop_sync(self):
if self.loop.is_running():
if self.loop is not None and self.loop.is_running():
future = asyncio.run_coroutine_threadsafe(self.stop(), self.loop)
future.result(timeout=10)
else:
self.loop.run_until_complete(self.stop())
asyncio.run(self.stop())
72 changes: 72 additions & 0 deletions test/unit/okx/websocket/test_websocket_factory.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,72 @@
"""Unit tests for the WebSocket transport factory."""

import asyncio
import unittest
from unittest.mock import AsyncMock, patch

from okx.websocket.WebSocketFactory import WebSocketFactory


TEST_WS_URL = "wss://test.example.com/ws/v5/public"


class TestWebSocketFactory(unittest.TestCase):

def test_connect_forwards_explicit_proxy(self):
websocket = AsyncMock()

async def run_test():
with patch(
"okx.websocket.WebSocketFactory.websockets.connect",
new_callable=AsyncMock,
return_value=websocket,
) as connect:
factory = WebSocketFactory(TEST_WS_URL, proxy="http://127.0.0.1:7890")

result = await factory.connect()

self.assertIs(result, websocket)
self.assertEqual(connect.await_args.args, (TEST_WS_URL,))
self.assertEqual(connect.await_args.kwargs["proxy"], "http://127.0.0.1:7890")
self.assertIn("ssl", connect.await_args.kwargs)

asyncio.run(run_test())

def test_connect_omits_proxy_when_not_configured(self):
websocket = AsyncMock()

async def run_test():
with patch(
"okx.websocket.WebSocketFactory.websockets.connect",
new_callable=AsyncMock,
return_value=websocket,
) as connect:
factory = WebSocketFactory(TEST_WS_URL)

await factory.connect()

self.assertNotIn("proxy", connect.await_args.kwargs)

asyncio.run(run_test())

def test_connect_propagates_original_exception(self):
error = OSError("proxy connection failed")

async def run_test():
with patch(
"okx.websocket.WebSocketFactory.websockets.connect",
new_callable=AsyncMock,
side_effect=error,
):
factory = WebSocketFactory(TEST_WS_URL)

with self.assertRaises(OSError) as raised:
await factory.connect()

self.assertIs(raised.exception, error)

asyncio.run(run_test())


if __name__ == "__main__":
unittest.main()
Loading