diff --git a/README.md b/README.md index 824ebd74..039dbc35 100644 --- a/README.md +++ b/README.md @@ -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 diff --git a/okx/websocket/WebSocketFactory.py b/okx/websocket/WebSocketFactory.py index ff983e6b..83503685 100644 --- a/okx/websocket/WebSocketFactory.py +++ b/okx/websocket/WebSocketFactory.py @@ -1,4 +1,3 @@ -import asyncio import logging import ssl @@ -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: diff --git a/okx/websocket/WsPrivateAsync.py b/okx/websocket/WsPrivateAsync.py index 2c99db37..874658f5 100644 --- a/okx/websocket/WsPrivateAsync.py +++ b/okx/websocket/WsPrivateAsync.py @@ -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 @@ -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()) diff --git a/okx/websocket/WsPublicAsync.py b/okx/websocket/WsPublicAsync.py index 72045f5a..b0ae7b0c 100644 --- a/okx/websocket/WsPublicAsync.py +++ b/okx/websocket/WsPublicAsync.py @@ -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 @@ -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()) diff --git a/test/unit/okx/websocket/test_websocket_factory.py b/test/unit/okx/websocket/test_websocket_factory.py new file mode 100644 index 00000000..a4975d75 --- /dev/null +++ b/test/unit/okx/websocket/test_websocket_factory.py @@ -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() diff --git a/test/unit/okx/websocket/test_ws_private_async.py b/test/unit/okx/websocket/test_ws_private_async.py index 88756952..c7ed046f 100644 --- a/test/unit/okx/websocket/test_ws_private_async.py +++ b/test/unit/okx/websocket/test_ws_private_async.py @@ -44,7 +44,22 @@ def test_init_with_required_params(self): self.assertEqual(ws.url, TEST_WS_URL) self.assertFalse(ws.useServerTime) self.assertFalse(ws.debug) - mock_factory.assert_called_once_with(TEST_WS_URL) + mock_factory.assert_called_once_with(TEST_WS_URL, proxy=None) + + def test_init_forwards_proxy_to_factory(self): + """An explicit proxy is forwarded to the WebSocket transport.""" + with patch.object(ws_private_module, 'WebSocketFactory') as mock_factory: + proxy = 'http://127.0.0.1:7890' + + WsPrivateAsync( + apiKey=_STUB_ID, + passphrase=_STUB_PHRASE, + secretKey=_STUB_SIGN, + url=TEST_WS_URL, + proxy=proxy + ) + + mock_factory.assert_called_once_with(TEST_WS_URL, proxy=proxy) def test_init_with_debug_enabled(self): """Test initialization with debug mode enabled""" @@ -130,7 +145,7 @@ async def run_test(): self.assertEqual(payload["args"], params) self.assertNotIn("id", payload) - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) def test_subscribe_with_id(self): """Test subscribe with id parameter""" @@ -159,7 +174,7 @@ async def run_test(): self.assertEqual(payload["op"], "subscribe") self.assertEqual(payload["id"], "sub001") - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) class TestWsPrivateAsyncUnsubscribe(unittest.TestCase): @@ -187,7 +202,7 @@ async def run_test(): self.assertEqual(payload["args"], params) self.assertNotIn("id", payload) - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) def test_unsubscribe_with_id(self): """Test unsubscribe with id parameter""" @@ -210,7 +225,7 @@ async def run_test(): self.assertEqual(payload["op"], "unsubscribe") self.assertEqual(payload["id"], "unsub001") - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) class TestWsPrivateAsyncSend(unittest.TestCase): @@ -240,7 +255,7 @@ async def run_test(): self.assertEqual(payload["args"], args) self.assertNotIn("id", payload) - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) def test_send_with_id(self): """Test generic send method with id""" @@ -263,7 +278,7 @@ async def run_test(): self.assertEqual(payload["op"], "custom_op") self.assertEqual(payload["id"], "send001") - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) class TestWsPrivateAsyncOrderMethods(unittest.TestCase): @@ -306,7 +321,7 @@ async def run_test(): self.assertEqual(payload["args"], order_args) self.assertEqual(payload["id"], "order001") - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) def test_place_order_without_id(self): """Test place_order without id parameter""" @@ -321,7 +336,7 @@ async def run_test(): self.assertEqual(payload["op"], "order") self.assertNotIn("id", payload) - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) def test_batch_orders_sends_correct_payload(self): """Test batch_orders sends correct operation""" @@ -341,7 +356,7 @@ async def run_test(): self.assertEqual(payload["args"], order_args) self.assertEqual(payload["id"], "batch001") - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) def test_batch_orders_without_id(self): """Test batch_orders without id parameter""" @@ -356,7 +371,7 @@ async def run_test(): self.assertEqual(payload["op"], "batch-orders") self.assertNotIn("id", payload) - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) def test_cancel_order_sends_correct_payload(self): """Test cancel_order sends correct operation""" @@ -373,7 +388,7 @@ async def run_test(): self.assertEqual(payload["args"], cancel_args) self.assertEqual(payload["id"], "cancel001") - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) def test_cancel_order_without_id(self): """Test cancel_order without id parameter""" @@ -388,7 +403,7 @@ async def run_test(): self.assertEqual(payload["op"], "cancel-order") self.assertNotIn("id", payload) - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) def test_batch_cancel_orders_sends_correct_payload(self): """Test batch_cancel_orders sends correct operation""" @@ -408,7 +423,7 @@ async def run_test(): self.assertEqual(payload["args"], cancel_args) self.assertEqual(payload["id"], "batchCancel001") - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) def test_batch_cancel_orders_without_id(self): """Test batch_cancel_orders without id parameter""" @@ -423,7 +438,7 @@ async def run_test(): self.assertEqual(payload["op"], "batch-cancel-orders") self.assertNotIn("id", payload) - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) def test_amend_order_sends_correct_payload(self): """Test amend_order sends correct operation""" @@ -445,7 +460,7 @@ async def run_test(): self.assertEqual(payload["args"], amend_args) self.assertEqual(payload["id"], "amend001") - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) def test_amend_order_without_id(self): """Test amend_order without id parameter""" @@ -460,7 +475,7 @@ async def run_test(): self.assertEqual(payload["op"], "amend-order") self.assertNotIn("id", payload) - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) def test_batch_amend_orders_sends_correct_payload(self): """Test batch_amend_orders sends correct operation""" @@ -480,7 +495,7 @@ async def run_test(): self.assertEqual(payload["args"], amend_args) self.assertEqual(payload["id"], "batchAmend001") - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) def test_batch_amend_orders_without_id(self): """Test batch_amend_orders without id parameter""" @@ -495,7 +510,7 @@ async def run_test(): self.assertEqual(payload["op"], "batch-amend-orders") self.assertNotIn("id", payload) - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) def test_mass_cancel_sends_correct_payload(self): """Test mass_cancel sends correct operation""" @@ -515,7 +530,7 @@ async def run_test(): self.assertEqual(payload["args"], mass_cancel_args) self.assertEqual(payload["id"], "massCancel001") - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) def test_mass_cancel_without_id(self): """Test mass_cancel without id parameter""" @@ -530,7 +545,7 @@ async def run_test(): self.assertEqual(payload["op"], "mass-cancel") self.assertNotIn("id", payload) - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) class TestWsPrivateAsyncLogin(unittest.TestCase): @@ -562,7 +577,7 @@ async def run_test(): secretKey=_STUB_SIGN ) - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) class TestWsPrivateAsyncStartStop(unittest.TestCase): @@ -586,7 +601,7 @@ async def run_test(): await ws.stop() mock_factory_instance.close.assert_called_once() - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) def test_start_returns_task(self): """start() returns the consume task so the caller can retain/await it (GH#116)""" @@ -605,7 +620,7 @@ async def run_test(): self.assertIsInstance(task, asyncio.Task) await task - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) class _RaisingAsyncIterator: @@ -651,7 +666,7 @@ async def run_test(): await ws.consume() with patch.object(ws_private_module.logger, 'error') as mock_log_error: - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) mock_log_error.assert_called_once() self.assertEqual(received[0], '{"data": "ok"}') @@ -683,7 +698,7 @@ async def __anext__(self): async def run_test(): await ws.consume() - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) self.assertEqual(received, ['a', 'b']) diff --git a/test/unit/okx/websocket/test_ws_public_async.py b/test/unit/okx/websocket/test_ws_public_async.py index 92c4d87f..80d3764e 100644 --- a/test/unit/okx/websocket/test_ws_public_async.py +++ b/test/unit/okx/websocket/test_ws_public_async.py @@ -38,6 +38,16 @@ def test_init_with_url(self): self.assertEqual(ws.secretKey, '') self.assertFalse(ws.debug) self.assertFalse(ws.isLoggedIn) + mock_factory.assert_called_once_with(TEST_WS_URL, proxy=None) + + def test_init_forwards_proxy_to_factory(self): + """An explicit proxy is forwarded to the WebSocket transport.""" + with patch.object(ws_public_module, 'WebSocketFactory') as mock_factory: + proxy = 'http://127.0.0.1:7890' + + WsPublicAsync(url=TEST_WS_URL, proxy=proxy) + + mock_factory.assert_called_once_with(TEST_WS_URL, proxy=proxy) def test_init_with_credentials(self): """Test initialization with all credentials for business channel""" @@ -85,7 +95,7 @@ async def run_test(): await ws.login() self.assertIn("apiKey, secretKey and passphrase are required for login", str(context.exception)) - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) def test_login_with_credentials_success(self): """Test successful login with valid credentials""" @@ -116,7 +126,7 @@ async def run_test(): ) mock_websocket.send.assert_called_once() - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) class TestWsPublicAsyncSubscribe(unittest.TestCase): @@ -144,7 +154,7 @@ async def run_test(): self.assertEqual(payload["args"], params) self.assertNotIn("id", payload) - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) def test_subscribe_with_id(self): """Test subscribe with id parameter""" @@ -165,7 +175,7 @@ async def run_test(): self.assertEqual(payload["args"], params) self.assertEqual(payload["id"], "sub001") - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) def test_subscribe_with_multiple_channels(self): """Test subscribe with multiple channels""" @@ -186,7 +196,7 @@ async def run_test(): self.assertEqual(len(payload["args"]), 2) self.assertEqual(payload["id"], "multi001") - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) class TestWsPublicAsyncUnsubscribe(unittest.TestCase): @@ -209,7 +219,7 @@ async def run_test(): self.assertEqual(payload["args"], params) self.assertNotIn("id", payload) - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) def test_unsubscribe_with_id(self): """Test unsubscribe with id parameter""" @@ -227,7 +237,7 @@ async def run_test(): self.assertEqual(payload["op"], "unsubscribe") self.assertEqual(payload["id"], "unsub001") - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) class TestWsPublicAsyncSend(unittest.TestCase): @@ -252,7 +262,7 @@ async def run_test(): self.assertEqual(payload["args"], args) self.assertNotIn("id", payload) - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) def test_send_with_id(self): """Test generic send method with id""" @@ -270,7 +280,7 @@ async def run_test(): self.assertEqual(payload["op"], "custom_op") self.assertEqual(payload["id"], "send001") - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) def test_send_without_callback(self): """Test send method without callback (preserves existing callback)""" @@ -288,7 +298,7 @@ async def run_test(): # Callback should remain unchanged self.assertEqual(ws.callback, existing_callback) - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) def test_send_with_new_callback_replaces_existing(self): """Test send method with new callback replaces existing callback""" @@ -306,7 +316,7 @@ async def run_test(): await ws.send("custom_op", args, callback=new_callback) self.assertEqual(ws.callback, new_callback) - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) class TestWsPublicAsyncStartStop(unittest.TestCase): @@ -325,7 +335,7 @@ async def run_test(): await ws.stop() mock_factory_instance.close.assert_called_once() - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) def test_start_returns_task(self): """start() returns the consume task so the caller can retain/await it (GH#116)""" @@ -339,7 +349,7 @@ async def run_test(): self.assertIsInstance(task, asyncio.Task) await task - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) class _RaisingAsyncIterator: @@ -377,7 +387,7 @@ async def run_test(): await ws.consume() with patch.object(ws_public_module.logger, 'error') as mock_log_error: - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) mock_log_error.assert_called_once() # Normal message delivered first, then the injected error event @@ -410,7 +420,7 @@ async def __anext__(self): async def run_test(): await ws.consume() - asyncio.get_event_loop().run_until_complete(run_test()) + asyncio.run(run_test()) self.assertEqual(received, ['a', 'b'])