From 0f82066ae7e483d74d8a4940b35678c036a9c955 Mon Sep 17 00:00:00 2001 From: bai <1145000687@qq.com> Date: Mon, 18 Dec 2023 12:20:32 +0800 Subject: [PATCH 1/6] =?UTF-8?q?=E4=BF=AE=E6=94=B9=E5=BC=82=E5=B8=B8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- python/socketd/socketd/core/async_utils.py | 84 ------------------- .../impl/AIOWebSocketClientImpl.py | 5 +- 2 files changed, 3 insertions(+), 86 deletions(-) delete mode 100644 python/socketd/socketd/core/async_utils.py diff --git a/python/socketd/socketd/core/async_utils.py b/python/socketd/socketd/core/async_utils.py deleted file mode 100644 index 024b56b3..00000000 --- a/python/socketd/socketd/core/async_utils.py +++ /dev/null @@ -1,84 +0,0 @@ -import asyncio -from concurrent.futures import ThreadPoolExecutor, as_completed, wait, ALL_COMPLETED -from common.globalconfig import global_config - -as_completed = as_completed -wait = wait -ALL_COMPLETED = ALL_COMPLETED - -ThreadPool = ThreadPoolExecutor() -# 生产最大线程10个 -if global_config.active == "prod": - ThreadPool._max_workers = 10 - - -class AsyncUtils: - - @staticmethod - def run(func): - return asyncio.run(func) - - @staticmethod - def run_coroutine(coro): - """运行一个协程并返回结果""" - loop = asyncio.get_event_loop() - task = loop.create_task(coro) - return task - - @staticmethod - async def async_core(core): - return core() - - @staticmethod - def run_futures(coro): - """运行一个协程并返回结果""" - loop = asyncio.get_event_loop() - task = loop.create_task(coro) - loop.run_until_complete(asyncio.wait(task)) - loop.close() - return task.result() - - @staticmethod - async def to_thread(func, *args, **kwargs): - """将一个阻塞型函数转换成协程任务,使用线程池异步执行""" - loop = asyncio.get_running_loop() - return await loop.run_in_executor(ThreadPool, lambda: func(*args, **kwargs)) - - @staticmethod - async def gather_concurrent(coros, limit=10): - """并发执行多个协程任务,并限制同时执行的数量""" - - async def worker(semaphore, coro): - async with semaphore: - return await coro - if limit < 0: - return await asyncio.gather(*coros) - semaphore = asyncio.Semaphore(limit) - tasks = [worker(semaphore, coro) for coro in coros] - return await asyncio.gather(*tasks) - - @staticmethod - async def run_with_timeout(coro, timeout=None): - """运行一个协程任务,并设置超时时间""" - task = asyncio.ensure_future(coro) - done, pending = await asyncio.wait([task], timeout=timeout) - if task in done: - return task.result() - else: - raise asyncio.TimeoutError() - - @staticmethod - async def repeat_task(coro, interval): - """重复运行一个协程任务,每隔一定时间间隔执行一次""" - while True: - await coro - await asyncio.sleep(interval) - - @staticmethod - async def sleep_until(timestamp): - """将当前协程挂起一段时间,直到指定的时间点""" - now = asyncio.get_event_loop().time() - if timestamp > now: - await asyncio.sleep(timestamp - now) - else: - raise ValueError("timestamp must be in the future") diff --git a/python/socketd/socketd_websocket/impl/AIOWebSocketClientImpl.py b/python/socketd/socketd_websocket/impl/AIOWebSocketClientImpl.py index 2763380c..40158d70 100644 --- a/python/socketd/socketd_websocket/impl/AIOWebSocketClientImpl.py +++ b/python/socketd/socketd_websocket/impl/AIOWebSocketClientImpl.py @@ -52,6 +52,7 @@ class AIOWebSocketClientImpl(WebSocketClientProtocol): try: await self.on_message() except Exception as e: + log.error(e) break def thread__handler(_loop): @@ -72,7 +73,7 @@ class AIOWebSocketClientImpl(WebSocketClientProtocol): await self.channel.send_connect(self.client.get_config().get_url()) await self.on_message() except Exception as e: - log.warning(str(e), exc_info=True) + log.error(str(e), exc_info=True) raise e async def on_message(self): @@ -93,7 +94,7 @@ class AIOWebSocketClientImpl(WebSocketClientProtocol): # 超时自动推出 log.debug(c) except Exception as e: - log.warning(str(e), exc_info=True) + log.error(str(e), exc_info=True) raise e def on_close(self): -- Gitee From 94a499f5e2951551bbc6e3af9b3e95f62cc42533 Mon Sep 17 00:00:00 2001 From: bai <1145000687@qq.com> Date: Thu, 21 Dec 2023 16:39:33 +0800 Subject: [PATCH 2/6] =?UTF-8?q?=E6=96=B0=E5=A2=9E=E6=B5=8B=E8=AF=95?= =?UTF-8?q?=E6=A1=88=E4=BE=8B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../test/java/benchmark/cases/ClientTest.java | 36 ++++++++++ .../test/java/benchmark/cases/ServerTest.java | 46 +++++++++++++ .../socketd/transport/core/AsyncUtil.py | 18 +++++ .../transport/core/HeartbeatHandlerDefault.py | 14 ++++ .../socketd/socketd_websocket/WsAioServer.py | 10 +-- .../impl/AIOWebSocketServerImpl.py | 2 +- python/socketd/test/FurtureTest.py | 14 ++++ python/socketd/test/TestCase00.py | 32 +++++++++ python/socketd/test/TestCase01.py | 38 +++++++++++ python/socketd/test/base_test/__init__.py | 0 .../test/cases/TestCase01_client_send.py | 65 ++++++++++++++++++ .../test/cases/TestCase02_auto_reconnect.py | 66 +++++++++++++++++++ python/socketd/test/cases/__init__.py | 0 python/socketd/test/modelu/BaseTestCase.py | 16 +++++ .../socketd/test/modelu/SimpleListenerTest.py | 23 ++++++- 15 files changed, 373 insertions(+), 7 deletions(-) create mode 100644 java/socketd-transport-test/src/test/java/benchmark/cases/ClientTest.java create mode 100644 java/socketd-transport-test/src/test/java/benchmark/cases/ServerTest.java create mode 100644 python/socketd/socketd/transport/core/AsyncUtil.py create mode 100644 python/socketd/socketd/transport/core/HeartbeatHandlerDefault.py create mode 100644 python/socketd/test/FurtureTest.py create mode 100644 python/socketd/test/TestCase00.py create mode 100644 python/socketd/test/TestCase01.py create mode 100644 python/socketd/test/base_test/__init__.py create mode 100644 python/socketd/test/cases/TestCase01_client_send.py create mode 100644 python/socketd/test/cases/TestCase02_auto_reconnect.py create mode 100644 python/socketd/test/cases/__init__.py create mode 100644 python/socketd/test/modelu/BaseTestCase.py diff --git a/java/socketd-transport-test/src/test/java/benchmark/cases/ClientTest.java b/java/socketd-transport-test/src/test/java/benchmark/cases/ClientTest.java new file mode 100644 index 00000000..64d83dc0 --- /dev/null +++ b/java/socketd-transport-test/src/test/java/benchmark/cases/ClientTest.java @@ -0,0 +1,36 @@ +package benchmark.cases; + +import org.noear.socketd.SocketD; +import org.noear.socketd.transport.client.ClientSession; +import org.noear.socketd.transport.core.Entity; +import org.noear.socketd.transport.core.Frame; +import org.noear.socketd.transport.core.Message; +import org.noear.socketd.transport.core.Session; +import org.noear.socketd.transport.core.entity.StringEntity; +import org.noear.socketd.transport.core.identifier.TimeidGenerator; + +import java.io.IOException; + +public class ClientTest { + + ClientSession clientSession; + + + public void client() throws IOException { + //client + String serverUrl = "ws" + "://127.0.0.1:" + 7779 + "/path?u=a&p=2"; + clientSession = SocketD.createClient(serverUrl) + .config(config -> config.idGenerator(new TimeidGenerator()) + .requestTimeout(5000)) + .open(); +// clientSession.send("test", new StringEntity("test")); + Entity entity = clientSession.sendAndRequest("test", new StringEntity("test")); + System.out.println(entity); + clientSession.close(); + } + + public static void main(String[] args) throws IOException { + ClientTest clientTest = new ClientTest(); + clientTest.client(); + } +} diff --git a/java/socketd-transport-test/src/test/java/benchmark/cases/ServerTest.java b/java/socketd-transport-test/src/test/java/benchmark/cases/ServerTest.java new file mode 100644 index 00000000..5d66d5b5 --- /dev/null +++ b/java/socketd-transport-test/src/test/java/benchmark/cases/ServerTest.java @@ -0,0 +1,46 @@ +package benchmark.cases; + +import org.noear.socketd.SocketD; +import org.noear.socketd.transport.core.Message; +import org.noear.socketd.transport.core.Session; +import org.noear.socketd.transport.core.entity.StringEntity; +import org.noear.socketd.transport.core.identifier.TimeidGenerator; +import org.noear.socketd.transport.core.listener.SimpleListener; +import org.noear.socketd.transport.server.Server; + +import java.io.IOException; +import java.util.concurrent.CountDownLatch; + +public class ServerTest { + + private CountDownLatch sendLatch; + + private Server server; + private Session clientSession; + + + public void server() throws IOException { + //server + server = SocketD.createServer("ws") + .config(c -> c.port(7779).idGenerator(new TimeidGenerator())) + .listen(new SimpleListener() { + @Override + public void onMessage(Session session, Message message) throws IOException { + if (message.isRequest()) { + session.replyEnd(message, new StringEntity("test")); + } else if (message.isSubscribe()) { + session.replyEnd(message, new StringEntity("test")); + } else { + sendLatch.countDown(); + } + System.out.println(message.toString()); + } + }) + .start(); + } + + public static void main(String[] args) throws IOException { + new ServerTest().server(); + } + +} diff --git a/python/socketd/socketd/transport/core/AsyncUtil.py b/python/socketd/socketd/transport/core/AsyncUtil.py new file mode 100644 index 00000000..bb95e5fa --- /dev/null +++ b/python/socketd/socketd/transport/core/AsyncUtil.py @@ -0,0 +1,18 @@ +import asyncio +from typing import Coroutine, Any + + +class AsyncUtil: + + @staticmethod + def thread_handler(_loop, fn: asyncio.Task): + """ + 静态方法 thread_handler 用于在指定的事件循环中运行一个协程函数。 + + 参数: + _loop (EventLoop): 事件循环对象,用于控制协程的执行。 + fn (Coroutine): 需要运行的协程函数。 + """ + asyncio.set_event_loop(_loop) + _loop.run_until_complete(fn) + diff --git a/python/socketd/socketd/transport/core/HeartbeatHandlerDefault.py b/python/socketd/socketd/transport/core/HeartbeatHandlerDefault.py new file mode 100644 index 00000000..eee7a824 --- /dev/null +++ b/python/socketd/socketd/transport/core/HeartbeatHandlerDefault.py @@ -0,0 +1,14 @@ +from abc import ABC + +from socketd.core.Session import Session + + +class HeartbeatHandler(ABC): + + def heartbeat(self, session: Session): ... + + +class HeartbeatHandlerDefault(HeartbeatHandler): + + async def heartbeat(self, session: Session): + await session.send_ping() diff --git a/python/socketd/socketd_websocket/WsAioServer.py b/python/socketd/socketd_websocket/WsAioServer.py index a04ec8ca..b03ecc90 100644 --- a/python/socketd/socketd_websocket/WsAioServer.py +++ b/python/socketd/socketd_websocket/WsAioServer.py @@ -17,13 +17,13 @@ class WsAioServer(ServerBase): super().__init__(config, WsAioChannelAssistant(config)) self.__loop = asyncio.get_event_loop() self.server: Serve = None - self._stop = asyncio.Future() # set this future to exit the server + self.__is_started = False async def start(self) -> 'WebSocketServer': - if self.isStarted: + if self.__is_started: raise Exception("Server started") else: - self.isStarted = True + self.__is_started = True if self._config.get_host() is not None: _server = AIOServe(ws_handler=None, host="0.0.0.0", port=self._config.get_port(), @@ -50,5 +50,5 @@ class WsAioServer(ServerBase): async def stop(self): logger.info("WsAioServer stop...") - await self.server.ws_server.close() - await self._stop + self.server.ws_server.close() + self.__is_started = False diff --git a/python/socketd/socketd_websocket/impl/AIOWebSocketServerImpl.py b/python/socketd/socketd_websocket/impl/AIOWebSocketServerImpl.py index 1e3dd1df..8d1923ce 100644 --- a/python/socketd/socketd_websocket/impl/AIOWebSocketServerImpl.py +++ b/python/socketd/socketd_websocket/impl/AIOWebSocketServerImpl.py @@ -31,7 +31,7 @@ class AIOWebSocketServerImpl(WebSocketServerProtocol, IWebSocketServer): def connection_open(self) -> None: """握手完成回调""" super().connection_open() - log.debug("AIOWebSocketServerImpl 打开握手完成回调") + log.debug("AIOWebSocketServerImpl connection_open") self.on_open(self) def handshake_handler(self): diff --git a/python/socketd/test/FurtureTest.py b/python/socketd/test/FurtureTest.py new file mode 100644 index 00000000..855f17dc --- /dev/null +++ b/python/socketd/test/FurtureTest.py @@ -0,0 +1,14 @@ +import asyncio + +import unittest + + +class FutureTest(unittest.TestCase): + + def test_wait(self): + async def _wait(): + top = asyncio.Future() + top.set_result(0) + await top + + asyncio.run(_wait()) diff --git a/python/socketd/test/TestCase00.py b/python/socketd/test/TestCase00.py new file mode 100644 index 00000000..403d8bbd --- /dev/null +++ b/python/socketd/test/TestCase00.py @@ -0,0 +1,32 @@ +import asyncio +import unittest +import sys +from loguru import logger + +from test.modelu.BaseTest import BaseTest + + +class TestCase00(unittest.TestCase): + + count = 100000 + timeout = 30 + + def __init__(self, *args, **kwargs): + logger.remove() + logger.add(sys.stderr, level="INFO") + super().__init__(*args, **kwargs) + + def test_send(self): + test = BaseTest() + loop = asyncio.new_event_loop() + try: + loop.run_until_complete(test.start()) + loop.run_until_complete(test.send(TestCase00.count)) + # loop.run_until_complete(test.send_and_request(TestCase00.count)) + # loop.run_until_complete(test.send_and_subscribe(TestCase00.count)) + except Exception as e: + pass + finally: + loop.run_until_complete(test.close()) + + diff --git a/python/socketd/test/TestCase01.py b/python/socketd/test/TestCase01.py new file mode 100644 index 00000000..4b586f57 --- /dev/null +++ b/python/socketd/test/TestCase01.py @@ -0,0 +1,38 @@ +import asyncio +import unittest +import sys +from loguru import logger + +from test.cases.TestCase01_client_send import TestCase01_client_send +from test.cases.TestCase02_auto_reconnect import TestCase02_auto_reconnect +from test.modelu.BaseTest import BaseTest + + +class TestCase01(unittest.TestCase): + + schemas = ["ws"] + + def __init__(self, *args, **kwargs): + logger.remove() + logger.add(sys.stderr, level="DEBUG") + super().__init__(*args, **kwargs) + + def test_Case01_client_send(self): + for i in range(len(TestCase01.schemas)): + t = TestCase01_client_send(TestCase01.schemas[i], 9000 + i) + try: + t.start() + t.stop() + except Exception as e: + t.on_error() + raise e + + def test_Case02_auto_reconnect(self): + for i in range(len(TestCase01.schemas)): + t = TestCase02_auto_reconnect(TestCase01.schemas[i], 9000 + i) + try: + t.start() + t.stop() + except Exception as e: + t.on_error() + raise e diff --git a/python/socketd/test/base_test/__init__.py b/python/socketd/test/base_test/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/python/socketd/test/cases/TestCase01_client_send.py b/python/socketd/test/cases/TestCase01_client_send.py new file mode 100644 index 00000000..58eb1cc0 --- /dev/null +++ b/python/socketd/test/cases/TestCase01_client_send.py @@ -0,0 +1,65 @@ +import asyncio + +from test.modelu.BaseTestCase import BaseTestCase + +import time +from websockets.legacy.server import WebSocketServer + +from socketd.core.Session import Session +from socketd.core.SocketD import SocketD +from socketd.core.config.ServerConfig import ServerConfig +from socketd.core.module.StringEntity import StringEntity +from socketd.transport.server.Server import Server +from test.modelu.SimpleListenerTest import SimpleListenerTest, config_handler, send_and_subscribe_test +from loguru import logger + + +class TestCase01_client_send(BaseTestCase): + + def __init__(self, schema, port): + super().__init__(schema, port) + self.server: Server = None + self.server_session: WebSocketServer = None + self.client_session: Session = None + self.loop = asyncio.get_event_loop() + + async def _start(self): + self.server: Server = SocketD.create_server(ServerConfig(self.schema).set_port(self.port)) + self.server_session: WebSocketServer = await self.server.config(config_handler).listen( + SimpleListenerTest()).start() + + serverUrl = self.schema + "://127.0.0.1:" + str(self.port) + "/path?u=a&p=2" + self.client_session: Session = await SocketD.create_client(serverUrl) \ + .config(config_handler).open() + await self.client_session.send_and_request("demo", StringEntity("test"), 100) + + start_time = time.monotonic() + for _ in range(3): + await self.client_session.send("demo", StringEntity("test")) + + await self.client_session.send_and_subscribe("demo", StringEntity("test"), send_and_subscribe_test, 100) + end_time = time.monotonic() + logger.info(f"Coroutine send took {(end_time - start_time) * 1000.0} monotonic to complete.") + await asyncio.sleep(3) + + def start(self): + super().start() + + self.loop.run_until_complete(self._start()) + + async def _stop(self): + if self.client_session: + await self.client_session.close() + + if self.server_session: + self.server_session.close() + if self.server: + await self.server.stop() + + def stop(self): + super().stop() + + self.loop.run_until_complete(self._stop()) + + def on_error(self): + super().on_error() diff --git a/python/socketd/test/cases/TestCase02_auto_reconnect.py b/python/socketd/test/cases/TestCase02_auto_reconnect.py new file mode 100644 index 00000000..43518734 --- /dev/null +++ b/python/socketd/test/cases/TestCase02_auto_reconnect.py @@ -0,0 +1,66 @@ +import asyncio + +from test.modelu.BaseTestCase import BaseTestCase + +from websockets.legacy.server import WebSocketServer +from loguru import logger + +from socketd.core.Session import Session +from socketd.core.SocketD import SocketD +from socketd.core.config.ServerConfig import ServerConfig +from socketd.core.module.StringEntity import StringEntity +from socketd.transport.server.Server import Server +from test.modelu.SimpleListenerTest import SimpleListenerTest, config_handler, send_and_subscribe_test + + +class TestCase02_auto_reconnect(BaseTestCase): + + def __init__(self, schema, port): + super().__init__(schema, port) + self.server: Server + self.server_session: WebSocketServer + self.client_session: Session + self.loop = asyncio.get_event_loop() + + async def _start(self): + self.server: Server = SocketD.create_server(ServerConfig(self.schema).set_port(self.port)) + _simple = SimpleListenerTest() + _server = self.server.config(config_handler).listen(_simple) + self.server_session: WebSocketServer = await _server.start() + + serverUrl = self.schema + "://127.0.0.1:" + str(self.port) + "/path?u=a&p=2" + self.client_session: Session = await SocketD.create_client(serverUrl) \ + .config(config_handler).open() + await self.client_session.send_and_request("demo", StringEntity("test"), 100) + + await self.server.stop() + del self.server_session + await asyncio.sleep(10) + self.server_session = await _server.start() + + for _ in range(3): + await self.client_session.send("demo", StringEntity("test")) + + await self.client_session.send_and_subscribe("demo", StringEntity("test"), send_and_subscribe_test, 100) + await asyncio.sleep(1) + logger.info(f"counter {_simple.server_counter.get()} ") + + def start(self): + super().start() + self.loop.run_until_complete(self._start()) + + async def _stop(self): + if self.client_session: + await self.client_session.close() + + if self.server_session: + self.server_session.close() + if self.server: + await self.server.stop() + + def stop(self): + super().stop() + self.loop.run_until_complete(self._stop()) + + def on_error(self): + super().on_error() diff --git a/python/socketd/test/cases/__init__.py b/python/socketd/test/cases/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/python/socketd/test/modelu/BaseTestCase.py b/python/socketd/test/modelu/BaseTestCase.py new file mode 100644 index 00000000..8a4d3ee1 --- /dev/null +++ b/python/socketd/test/modelu/BaseTestCase.py @@ -0,0 +1,16 @@ +from loguru import logger + +class BaseTestCase: + + def __init__(self, schema, port): + self.schema = schema + self.port = port + + def start(self): + logger.info (f"--------START {self.schema}------") + + def stop(self): + logger.info(f"--------STOP {self.schema}------") + + def on_error(self): + logger.info(f"--------ERROR {self.schema}------") diff --git a/python/socketd/test/modelu/SimpleListenerTest.py b/python/socketd/test/modelu/SimpleListenerTest.py index 33611004..3ffbb4be 100644 --- a/python/socketd/test/modelu/SimpleListenerTest.py +++ b/python/socketd/test/modelu/SimpleListenerTest.py @@ -2,20 +2,36 @@ import uuid from abc import ABC from socketd.core.Listener import Listener +from socketd.core.config.ClientConfig import ClientConfig +from socketd.core.config.ServerConfig import ServerConfig +from socketd.core.module.Entity import Entity from socketd.core.module.Message import Message from socketd.core.module.StringEntity import StringEntity +from loguru import logger + +from socketd.core.sync_api.AtomicRefer import AtomicRefer + class SimpleListenerTest(Listener, ABC): + def __init__(self): + self.server_counter = AtomicRefer(0) + def on_open(self, session): pass async def on_message(self, session, message: Message): + with self.server_counter: + self.server_counter.set(self.server_counter.get() + 1) if message.is_request(): + await session.reply(message, StringEntity("reply")) await session.reply_end(message, StringEntity("ok test")) + await session.reply(message, StringEntity("reply")) elif message.is_subscribe(): + await session.reply(message, StringEntity("reply")) await session.reply_end(message, StringEntity("ok test")) + await session.reply(message, StringEntity("reply")) def on_close(self, session): pass @@ -24,5 +40,10 @@ class SimpleListenerTest(Listener, ABC): pass -def idGenerator(config): +def config_handler(config: ServerConfig | ClientConfig) -> ServerConfig | ClientConfig: + config.set_is_thread(False) return config.id_generator(uuid.uuid4) + + +def send_and_subscribe_test(e: Entity): + logger.info(e) -- Gitee From 8a8fdaf85203decbf6f03ace9813d022b7e87f3d Mon Sep 17 00:00:00 2001 From: bai <1145000687@qq.com> Date: Thu, 21 Dec 2023 16:39:41 +0800 Subject: [PATCH 3/6] =?UTF-8?q?=E6=96=B0=E5=A2=9E=E6=B5=8B=E8=AF=95?= =?UTF-8?q?=E6=A1=88=E4=BE=8B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../src/test/java/features/CaseTest.java | 6 ++- .../transport/client/ClientChannel.java | 2 +- python/socketd/reqeusts.txt | 5 +-- python/socketd/socketd/core/ChannelDefault.py | 7 ++-- python/socketd/socketd/core/Session.py | 24 +++++++++-- python/socketd/socketd/core/SessionDefault.py | 14 ++++--- .../socketd/core/config/ClientConfig.py | 9 +--- .../socketd/socketd/core/config/ConfigBase.py | 7 ++++ .../socketd/transport/client/ClientChannel.py | 41 +++++++++++++++---- .../transport/core/CompletableFuture.py | 1 - .../socketd/transport/core/StreamBase.py | 1 - .../socketd/transport/server/Server.py | 6 +-- .../socketd/socketd_websocket/WsAioFactoy.py | 2 +- .../impl/AIOWebSocketClientImpl.py | 23 +++++++---- python/socketd/test/00_TestCase.py | 32 --------------- .../test/{ => base_test}/01_applictionTest.py | 2 +- .../test/{ => base_test}/02_ClientTest.py | 0 .../test/{ => base_test}/02_ServerTest.py | 0 .../test/{ => base_test}/04_websockets.py | 1 - 19 files changed, 98 insertions(+), 85 deletions(-) delete mode 100644 python/socketd/test/00_TestCase.py rename python/socketd/test/{ => base_test}/01_applictionTest.py (98%) rename python/socketd/test/{ => base_test}/02_ClientTest.py (100%) rename python/socketd/test/{ => base_test}/02_ServerTest.py (100%) rename python/socketd/test/{ => base_test}/04_websockets.py (99%) diff --git a/java/socketd-transport-test/src/test/java/features/CaseTest.java b/java/socketd-transport-test/src/test/java/features/CaseTest.java index f52a044f..8bbcf7b2 100644 --- a/java/socketd-transport-test/src/test/java/features/CaseTest.java +++ b/java/socketd-transport-test/src/test/java/features/CaseTest.java @@ -9,7 +9,9 @@ import org.junit.jupiter.api.Test; */ public class CaseTest { static final String[] schemas = new String[]{ - "sd:tcp-java", "sd:tcp-netty", "sd:tcp-smartsocket", + "sd:tcp-java", +// "sd:tcp-netty", +// "sd:tcp-smartsocket", "sd:ws-java", "sd:udp-java"}; @@ -17,7 +19,7 @@ public class CaseTest { public void TestCase01_client_send() throws Exception { for (int i = 0; i < schemas.length; i++) { String s1 = schemas[i]; - BaseTestCase testCase = new TestCase01_client_send(s1, 1000 + i); + BaseTestCase testCase = new TestCase01_client_send(s1, 9000 + i); try { testCase.start(); testCase.stop(); diff --git a/java/socketd/src/main/java/org/noear/socketd/transport/client/ClientChannel.java b/java/socketd/src/main/java/org/noear/socketd/transport/client/ClientChannel.java index c87c753e..546eadf9 100644 --- a/java/socketd/src/main/java/org/noear/socketd/transport/client/ClientChannel.java +++ b/java/socketd/src/main/java/org/noear/socketd/transport/client/ClientChannel.java @@ -227,7 +227,7 @@ public class ClientChannel extends ChannelBase implements Channel { * @return 是否为新链接 */ private boolean prepareCheck() throws IOException { - if (real == null || real.isValid() == false) { + if (real == null || !real.isValid()) { real = connector.connect(); return true; diff --git a/python/socketd/reqeusts.txt b/python/socketd/reqeusts.txt index f12ae6ca..ee89d343 100644 --- a/python/socketd/reqeusts.txt +++ b/python/socketd/reqeusts.txt @@ -1,5 +1,2 @@ -colorama==0.4.6 -connect==0.2 loguru==0.7.2 -nest-asyncio==1.5.8 -websockets==12.0 +websockets==12.0 \ No newline at end of file diff --git a/python/socketd/socketd/core/ChannelDefault.py b/python/socketd/socketd/core/ChannelDefault.py index ee8bb565..8d127085 100644 --- a/python/socketd/socketd/core/ChannelDefault.py +++ b/python/socketd/socketd/core/ChannelDefault.py @@ -79,10 +79,11 @@ class ChannelDefault(ChannelBase): if acceptor.is_single() or frame.get_flag() == Flag.ReplyEnd: self.acceptorMap.pop(frame.get_message().get_sid()) await asyncio.get_event_loop().run_in_executor(self.get_config().get_executor(), - lambda _m: acceptor.on_accept(_m, onError), frame.get_message()) + lambda _m: acceptor.on_accept(_m, onError), + frame.get_message()) else: - logger.debug(f"{self.get_config().getRoleName()} stream not found, sid={frame.get_message().get_sid()}, sessionId={self.get_session().get_session_id()}") - + logger.debug( + f"{self.get_config().get_role_name()} stream not found, sid={frame.get_message().get_sid()}, sessionId={self.get_session().get_session_id()}") def get_session(self) -> Session: if self.session is None: diff --git a/python/socketd/socketd/core/Session.py b/python/socketd/socketd/core/Session.py index 6615f0f6..45f46134 100644 --- a/python/socketd/socketd/core/Session.py +++ b/python/socketd/socketd/core/Session.py @@ -1,5 +1,5 @@ import abc -from typing import Union, Dict, Any, Callable, Awaitable +from typing import Union, Dict, Any, Callable, Awaitable, AsyncGenerator, Coroutine from socket import gethostbyaddr from socketd.core.Handshake import Handshake from socketd.core.module.Message import Message @@ -44,7 +44,7 @@ class Session(abc.ABC): ... @abc.abstractmethod - def send_ping(self) -> None: + def send_ping(self) -> Callable | Coroutine: ... @abc.abstractmethod @@ -55,8 +55,23 @@ class Session(abc.ABC): async def send_and_request(self, topic: str, content: Entity, timeout: int) -> Entity: ... - async def send_stream_and_request(self, event: str, content: Entity, consumer: Callable[[Entity], Awaitable[Any]], - timeout: int): ... + @abc.abstractmethod + async def send_stream_and_request(self, event: str, content: Entity, + consumer: Callable[[Entity], Awaitable[Any]] | Coroutine[Entity, Any, None], + timeout: int): + """ + 发送流和请求的抽象方法。 + + Args: + event (str): 事件名称。 + content (Entity): 内容实体。 + consumer (Callable[[Entity], Awaitable[Any]] | Coroutine[Entity, Awaitable[Any]]): 消费函数或协程,用来处理内容实体并返回异步结果。 + timeout (int): 超时时间。 + + Returns: + Awaitable[Any]: 消费函数或协程的异步结果。 + """ + ... @abc.abstractmethod async def send_and_subscribe(self, topic: str, content: Entity, consumer: Callable[[Entity], Any], @@ -71,5 +86,6 @@ class Session(abc.ABC): def reply_end(self, from_msg: Message, content: Entity) -> None: ... + @abc.abstractmethod def close(self): ... diff --git a/python/socketd/socketd/core/SessionDefault.py b/python/socketd/socketd/core/SessionDefault.py index 4080220c..c9a14225 100644 --- a/python/socketd/socketd/core/SessionDefault.py +++ b/python/socketd/socketd/core/SessionDefault.py @@ -1,7 +1,7 @@ import asyncio from abc import ABC -from typing import Callable, Awaitable, Any +from typing import Callable, Awaitable, Any, Coroutine from .SessionBase import SessionBase from .Channel import Channel @@ -34,8 +34,8 @@ class SessionDefault(SessionBase, ABC): def get_handshake(self) -> Handshake: return self.channel.get_handshake() - def send_ping(self): - self.channel.send_ping() + async def send_ping(self): + await self.channel.send_ping() async def send(self, topic: str, content: Entity): message = MessageDefault().set_sid(self.generate_id()).set_event(topic).set_entity(content) @@ -65,12 +65,16 @@ class SessionDefault(SessionBase, ABC): finally: self.channel.remove_acceptor(message.get_sid()) - async def send_stream_and_request(self, event: str, content: Entity, consumer: Callable[[Entity], Awaitable[Any]], + async def send_stream_and_request(self, event: str, content: Entity, + consumer: Callable[[Entity], Awaitable[Any]] | Coroutine[Entity, Any, None], timeout: int): message = MessageDefault().set_sid(self.generate_id()).set_event(event).set_entity(content) future: CompletableFuture[Entity] = CompletableFuture() try: - consumer(content) + if asyncio.iscoroutinefunction(consumer): + await consumer(content) + else: + self.channel.get_config().get_executor().submit(fn=consumer, args=(content,)) except Exception as e: self.channel.on_error(e) streamAcceptor = StreamRequest(message.get_sid(), timeout, future) diff --git a/python/socketd/socketd/core/config/ClientConfig.py b/python/socketd/socketd/core/config/ClientConfig.py index 5967e135..6c25b4b4 100644 --- a/python/socketd/socketd/core/config/ClientConfig.py +++ b/python/socketd/socketd/core/config/ClientConfig.py @@ -18,7 +18,6 @@ class ClientConfig(ConfigBase): self.__auto_reconnect = True self.__read_buffer_size = None self.__write_buffer_size = None - self.__is_thread = False def get_schema(self): return self.__schema @@ -70,12 +69,7 @@ class ClientConfig(ConfigBase): self.__auto_reconnect = __auto_reconnect return self - def is_thread(self, __is_thread: bool): - self.__is_thread = __is_thread - return self - def get_is_thread(self): - return self.__is_thread def __str__(self): return f"ClientConfig{{__schema='{self.__schema}', __url='{self.__url}', " \ @@ -85,5 +79,4 @@ class ClientConfig(ConfigBase): f"writeBufferSize={self.__write_buffer_size}, " \ f"autoReconnect={self.__auto_reconnect}, " \ f"maxRequests={self._max_requests}, " \ - f"maxUdpSize={self._max_udp_size}}}" \ - f"isThread={self.is_thread}}}" \ No newline at end of file + f"maxUdpSize={self._max_udp_size}}}" diff --git a/python/socketd/socketd/core/config/ConfigBase.py b/python/socketd/socketd/core/config/ConfigBase.py index 63cd82a9..86577279 100644 --- a/python/socketd/socketd/core/config/ConfigBase.py +++ b/python/socketd/socketd/core/config/ConfigBase.py @@ -27,6 +27,7 @@ class ConfigBase(Config): self._reply_timeout = 3000 self._max_requests = 10 self._max_udp_size = 2048 + self.__is_thread = False def client_mode(self): return self._client_mode @@ -128,3 +129,9 @@ class ConfigBase(Config): def set_stream_timeout(self, _stream_timeout): self._stream_timeout = _stream_timeout return self + + def set_is_thread(self, _is_thread): + self.__is_thread = _is_thread + + def get_is_thread(self): + return self.__is_thread diff --git a/python/socketd/socketd/transport/client/ClientChannel.py b/python/socketd/socketd/transport/client/ClientChannel.py index 042cb197..d1d4600a 100644 --- a/python/socketd/socketd/transport/client/ClientChannel.py +++ b/python/socketd/socketd/transport/client/ClientChannel.py @@ -1,4 +1,6 @@ +import asyncio from abc import ABC +from asyncio import Future from socketd.core.AssertsUtil import AssertsUtil from socketd.core.Channel import Channel @@ -6,6 +8,9 @@ from socketd.core.ChannelBase import ChannelBase from socketd.transport.client.ClientConnector import ClientConnector from loguru import logger +from socketd.transport.core.AsyncUtil import AsyncUtil +from socketd.transport.core.HeartbeatHandlerDefault import HeartbeatHandlerDefault + class ClientChannel(ChannelBase, ABC): def __init__(self, real: Channel, connector: ClientConnector): @@ -13,13 +18,17 @@ class ClientChannel(ChannelBase, ABC): self.real: Channel = real self.connector: ClientConnector = connector self.heartbeatHandler = connector.heartbeatHandler() + self._heartbeatScheduledFuture: Future | None = None + + if self.heartbeatHandler is None: + self.heartbeatHandler = HeartbeatHandlerDefault() - # if self.heartbeatHandler is None: - # self.heartbeatHandler = HeartbeatHandlerDefault() - connector.autoReconnect() - # if connector.autoReconnect() and self.heartbeatScheduledFuture is None: - # self.heartbeatScheduledFuture = threading.Timer(connector.heartbeatInterval(), self.heartbeatHandle) - # self.heartbeatScheduledFuture.start() + self._loop = asyncio.new_event_loop() + self.initHeartbeat() + + def __del__(self): + if self._loop: + self._loop.close() def remove_acceptor(self, sid): if self.real is not None: @@ -49,13 +58,27 @@ class ClientChannel(ChannelBase, ABC): else: return self.real.get_local_address() + def initHeartbeat(self): + if self._heartbeatScheduledFuture is not None: + self._heartbeatScheduledFuture.cancel() + + if self.connector.autoReconnect(): + async def _heartbeatScheduled(): + while True: + await asyncio.sleep(self.connector.heartbeatInterval()) + await self.heartbeat_handle() + self._heartbeatScheduledFuture = asyncio.create_task(_heartbeatScheduled()) + + self.get_config().get_executor().submit(lambda: + AsyncUtil.thread_handler(self._loop, self._heartbeatScheduledFuture)) + def heartbeat_handle(self): AssertsUtil.assert_closed(self.real) with self: try: self.prepare_send() - self.heartbeatHandler.heartbeat_handle() + self.heartbeatHandler.heartbeat(self.get_session()) except Exception as e: if self.connector.autoReconnect(): self.real.close() @@ -75,7 +98,7 @@ class ClientChannel(ChannelBase, ABC): raise e async def retrieve(self, frame, on_error): - self.real.retrieve(frame, on_error) + await self.real.retrieve(frame, on_error) def get_session(self): return self.real.get_session() @@ -84,6 +107,7 @@ class ClientChannel(ChannelBase, ABC): reason: str = "", ): try: await super().close(code, reason) + self._heartbeatScheduledFuture.cancel() if self.real is not None: await self.real.close() except Exception as e: @@ -101,4 +125,3 @@ class ClientChannel(ChannelBase, ABC): def on_error(self, error: Exception): pass - diff --git a/python/socketd/socketd/transport/core/CompletableFuture.py b/python/socketd/socketd/transport/core/CompletableFuture.py index dad3e595..322ac4b5 100644 --- a/python/socketd/socketd/transport/core/CompletableFuture.py +++ b/python/socketd/socketd/transport/core/CompletableFuture.py @@ -1,5 +1,4 @@ import asyncio -from asyncio.coroutines import iscoroutine from typing import Generic, TypeVar from loguru import logger diff --git a/python/socketd/socketd/transport/core/StreamBase.py b/python/socketd/socketd/transport/core/StreamBase.py index 6ffcc896..ac2501e7 100644 --- a/python/socketd/socketd/transport/core/StreamBase.py +++ b/python/socketd/socketd/transport/core/StreamBase.py @@ -12,7 +12,6 @@ class StreamBase(StreamInternal): self.__sid = sid self.__timeout = timeout self.__onError: Callable[[Exception], None] = None - self.insurance_future = Future() def on_error(self, error: Exception): if error: diff --git a/python/socketd/socketd/transport/server/Server.py b/python/socketd/socketd/transport/server/Server.py index 27c1b049..6d29bf7a 100644 --- a/python/socketd/socketd/transport/server/Server.py +++ b/python/socketd/socketd/transport/server/Server.py @@ -1,4 +1,4 @@ -from typing import Callable +from typing import Callable, Coroutine from asyncio import Future from websockets.sync.server import WebSocketServer @@ -15,9 +15,9 @@ class Server: def listen(self, listener: Listener) -> 'Server': ... - def start(self) -> WebSocketServer | Future: ... + def start(self) -> WebSocketServer | Coroutine: ... - def stop(self) -> Future: ... + def stop(self) -> Coroutine: ... def get_assistant(self): ... diff --git a/python/socketd/socketd_websocket/WsAioFactoy.py b/python/socketd/socketd_websocket/WsAioFactoy.py index 842ae5ea..60345565 100644 --- a/python/socketd/socketd_websocket/WsAioFactoy.py +++ b/python/socketd/socketd_websocket/WsAioFactoy.py @@ -10,7 +10,7 @@ from socketd_websocket.WsAioServer import WsAioServer class WsAioFactory(ClientFactory, ServerFactory): - def schema(self): + def schema(self) -> list[str]: return ["ws", "wss", "ws-python"] def create_server(self, serverConfig: ServerConfig) -> Server: diff --git a/python/socketd/socketd_websocket/impl/AIOWebSocketClientImpl.py b/python/socketd/socketd_websocket/impl/AIOWebSocketClientImpl.py index 40158d70..3fc25b36 100644 --- a/python/socketd/socketd_websocket/impl/AIOWebSocketClientImpl.py +++ b/python/socketd/socketd_websocket/impl/AIOWebSocketClientImpl.py @@ -12,6 +12,7 @@ from websockets import WebSocketClientProtocol, Origin, Subprotocol, HeadersLike from socketd.core.Costants import Flag from socketd.core.module.Frame import Frame from socketd.transport.client.Client import Client +from socketd.transport.core.AsyncUtil import AsyncUtil log = logger.opt() @@ -22,7 +23,7 @@ class AIOWebSocketClientImpl(WebSocketClientProtocol): self.status_state = Flag.Unknown self.client = client self.channel = ChannelDefault(self, client.get_config(), client.get_assistant()) - self.connect_read_thread: Thread = None + self.connect_read_thread: Thread | None = None def get_channel(self): return self.channel @@ -39,11 +40,17 @@ class AIOWebSocketClientImpl(WebSocketClientProtocol): return return_data def connection_open(self) -> None: - """打开握手完成回调""" + """ + 打开握手完成回调函数。 + :return: 无返回值 + """ super().connection_open() log.debug("AIOWebSocketClientImpl connection_open") async def _handler(): + """ + 异步处理函数,用于处理握手完成后的消息处理逻辑。 + """ while True: await asyncio.sleep(0) if self.closed or self.status_state == Flag.Close: @@ -54,17 +61,14 @@ class AIOWebSocketClientImpl(WebSocketClientProtocol): except Exception as e: log.error(e) break - - def thread__handler(_loop): - asyncio.set_event_loop(_loop) - _loop.run_until_complete(_handler()) - + # 如果配置中设置了使用线程,则创建一个新的事件循环,并启动一个线程来处理读取操作 if self.client.get_config().get_is_thread(): loop = asyncio.new_event_loop() if self.connect_read_thread is None: - self.connect_read_thread = Thread(target=thread__handler, args=(loop,)) + self.connect_read_thread = Thread(target=AsyncUtil.thread_handler, args=(loop,)) self.connect_read_thread.start() else: + # 否则,使用异步运行此协程,并在当前线程中运行。 asyncio.run_coroutine_threadsafe(_handler(), asyncio.get_event_loop()) async def on_open(self): @@ -79,6 +83,8 @@ class AIOWebSocketClientImpl(WebSocketClientProtocol): async def on_message(self): """处理消息""" try: + if self.status_state == Flag.Close: + return message = await self.recv() log.debug(message) frame: Frame = self.client.get_assistant().read(message) @@ -94,7 +100,6 @@ class AIOWebSocketClientImpl(WebSocketClientProtocol): # 超时自动推出 log.debug(c) except Exception as e: - log.error(str(e), exc_info=True) raise e def on_close(self): diff --git a/python/socketd/test/00_TestCase.py b/python/socketd/test/00_TestCase.py deleted file mode 100644 index 812b1789..00000000 --- a/python/socketd/test/00_TestCase.py +++ /dev/null @@ -1,32 +0,0 @@ -import asyncio -import unittest -import sys -from loguru import logger - -from test.modelu.BaseTest import BaseTest - - -class TestCase(unittest.TestCase): - - count = 10000 - timeout = 30 - - def __init__(self, *args, **kwargs): - logger.remove() - logger.add(sys.stderr, level="INFO") - super().__init__(*args, **kwargs) - - def test_send(self): - test = BaseTest() - loop = asyncio.get_event_loop() - try: - loop.run_until_complete(test.start()) - loop.run_until_complete(test.send(TestCase.count)) - loop.run_until_complete(test.send_and_request(TestCase.count)) - loop.run_until_complete(test.send_and_subscribe(TestCase.count)) - except Exception as e: - pass - finally: - test.close() - - diff --git a/python/socketd/test/01_applictionTest.py b/python/socketd/test/base_test/01_applictionTest.py similarity index 98% rename from python/socketd/test/01_applictionTest.py rename to python/socketd/test/base_test/01_applictionTest.py index c7ce4354..dbdd1b44 100644 --- a/python/socketd/test/01_applictionTest.py +++ b/python/socketd/test/base_test/01_applictionTest.py @@ -54,7 +54,7 @@ async def application_test(): server: Server = SocketD.create_server(ServerConfig("ws").set_port(9999)) server_session: WebSocketServer = await server.config(idGenerator).listen( SimpleListenerTest()).start() - + await asyncio.sleep(1) client_session: Session = await SocketD.create_client("ws://127.0.0.1:9999") \ .config(idGenerator).open() diff --git a/python/socketd/test/02_ClientTest.py b/python/socketd/test/base_test/02_ClientTest.py similarity index 100% rename from python/socketd/test/02_ClientTest.py rename to python/socketd/test/base_test/02_ClientTest.py diff --git a/python/socketd/test/02_ServerTest.py b/python/socketd/test/base_test/02_ServerTest.py similarity index 100% rename from python/socketd/test/02_ServerTest.py rename to python/socketd/test/base_test/02_ServerTest.py diff --git a/python/socketd/test/04_websockets.py b/python/socketd/test/base_test/04_websockets.py similarity index 99% rename from python/socketd/test/04_websockets.py rename to python/socketd/test/base_test/04_websockets.py index b496a2fb..fe79d115 100644 --- a/python/socketd/test/04_websockets.py +++ b/python/socketd/test/base_test/04_websockets.py @@ -50,4 +50,3 @@ class websockets_Test(unittest.TestCase): if __name__ == "__main__": websockets_Test().test_application() - pass \ No newline at end of file -- Gitee From 73e4f7f2e07b9bc369b2ae2f41a5c6434a21394d Mon Sep 17 00:00:00 2001 From: bai <1145000687@qq.com> Date: Thu, 21 Dec 2023 17:53:32 +0800 Subject: [PATCH 4/6] =?UTF-8?q?=E6=96=B0=E5=A2=9E=E6=B5=8B=E8=AF=95?= =?UTF-8?q?=E6=A1=88=E4=BE=8B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- python/socketd/socketd/core/SessionDefault.py | 5 +- python/socketd/test/TestCase01.py | 23 ++++++- .../cases/TestCase03_client_session_close.py | 67 +++++++++++++++++++ .../TestCase04_sendAndRequest_timeout.py | 60 +++++++++++++++++ .../socketd/test/modelu/SimpleListenerTest.py | 8 ++- 5 files changed, 159 insertions(+), 4 deletions(-) create mode 100644 python/socketd/test/cases/TestCase03_client_session_close.py create mode 100644 python/socketd/test/cases/TestCase04_sendAndRequest_timeout.py diff --git a/python/socketd/socketd/core/SessionDefault.py b/python/socketd/socketd/core/SessionDefault.py index c9a14225..6922f920 100644 --- a/python/socketd/socketd/core/SessionDefault.py +++ b/python/socketd/socketd/core/SessionDefault.py @@ -102,5 +102,6 @@ class SessionDefault(SessionBase, ABC): .set_entity(content)), None) async def close(self): - await self.channel.send_close() - await self.channel.close() + if self.channel.is_valid(): + await self.channel.send_close() + await self.channel.close() diff --git a/python/socketd/test/TestCase01.py b/python/socketd/test/TestCase01.py index 4b586f57..15fadfc5 100644 --- a/python/socketd/test/TestCase01.py +++ b/python/socketd/test/TestCase01.py @@ -5,11 +5,12 @@ from loguru import logger from test.cases.TestCase01_client_send import TestCase01_client_send from test.cases.TestCase02_auto_reconnect import TestCase02_auto_reconnect +from test.cases.TestCase03_client_session_close import TestCase03_client_session_close +from test.cases.TestCase04_sendAndRequest_timeout import TestCase04_sendAndRequest_timeout from test.modelu.BaseTest import BaseTest class TestCase01(unittest.TestCase): - schemas = ["ws"] def __init__(self, *args, **kwargs): @@ -36,3 +37,23 @@ class TestCase01(unittest.TestCase): except Exception as e: t.on_error() raise e + + def test_Case02_client_session_close(self): + for i in range(len(TestCase01.schemas)): + t = TestCase03_client_session_close(TestCase01.schemas[i], 9000 + i) + try: + t.start() + t.stop() + except Exception as e: + t.on_error() + raise e + + def test_Case02sendAndRequest_timeout(self): + for i in range(len(TestCase01.schemas)): + t = TestCase04_sendAndRequest_timeout(TestCase01.schemas[i], 9000 + i) + try: + t.start() + t.stop() + except Exception as e: + t.on_error() + raise e diff --git a/python/socketd/test/cases/TestCase03_client_session_close.py b/python/socketd/test/cases/TestCase03_client_session_close.py new file mode 100644 index 00000000..3e5adec4 --- /dev/null +++ b/python/socketd/test/cases/TestCase03_client_session_close.py @@ -0,0 +1,67 @@ +import asyncio + +from test.modelu.BaseTestCase import BaseTestCase + +from websockets.legacy.server import WebSocketServer +from loguru import logger + +from socketd.core.Session import Session +from socketd.core.SocketD import SocketD +from socketd.core.config.ServerConfig import ServerConfig +from socketd.core.module.StringEntity import StringEntity +from socketd.transport.server.Server import Server +from test.modelu.SimpleListenerTest import SimpleListenerTest, config_handler, send_and_subscribe_test + + +class TestCase03_client_session_close(BaseTestCase): + + def __init__(self, schema, port): + super().__init__(schema, port) + self.server: Server + self.server_session: WebSocketServer + self.client_session: Session + self.loop = asyncio.get_event_loop() + + async def _start(self): + self.server: Server = SocketD.create_server(ServerConfig(self.schema).set_port(self.port)) + _simple = SimpleListenerTest() + _server = self.server.config(config_handler).listen(_simple) + self.server_session: WebSocketServer = await _server.start() + + serverUrl = self.schema + "://127.0.0.1:" + str(self.port) + "/path?u=a&p=2" + self.client_session: Session = await SocketD.create_client(serverUrl) \ + .config(config_handler).open() + await self.client_session.send_and_request("demo", StringEntity("test"), 100) + + + try: + await self.client_session.close() + await self.client_session.send("demo", StringEntity("test")) + await self.client_session.send_and_subscribe("demo", StringEntity("test"), send_and_subscribe_test, 100) + except Exception as e: + pass + await asyncio.sleep(5) + logger.info(f"counter {_simple.server_counter.get()} close_counter {_simple.close_counter.get()}") + + def start(self): + super().start() + self.loop.run_until_complete(self._start()) + + async def _stop(self): + if self.client_session: + await self.client_session.close() + + if self.server_session: + self.server_session.close() + if self.server: + await self.server.stop() + + def stop(self): + super().stop() + self.loop.run_until_complete(self._stop()) + + def on_error(self): + try: + super().on_error() + except Exception as e: + logger.error(e) diff --git a/python/socketd/test/cases/TestCase04_sendAndRequest_timeout.py b/python/socketd/test/cases/TestCase04_sendAndRequest_timeout.py new file mode 100644 index 00000000..5fb9ad30 --- /dev/null +++ b/python/socketd/test/cases/TestCase04_sendAndRequest_timeout.py @@ -0,0 +1,60 @@ +import asyncio + +from test.modelu.BaseTestCase import BaseTestCase + +from websockets.legacy.server import WebSocketServer +from loguru import logger + +from socketd.core.Session import Session +from socketd.core.SocketD import SocketD +from socketd.core.config.ServerConfig import ServerConfig +from socketd.core.module.StringEntity import StringEntity +from socketd.transport.server.Server import Server +from test.modelu.SimpleListenerTest import SimpleListenerTest, config_handler, send_and_subscribe_test + + +class TestCase04_sendAndRequest_timeout(BaseTestCase): + + def __init__(self, schema, port): + super().__init__(schema, port) + self.server: Server + self.server_session: WebSocketServer + self.client_session: Session + self.loop = asyncio.get_event_loop() + + async def _start(self): + self.server: Server = SocketD.create_server(ServerConfig(self.schema).set_port(self.port)) + _simple = SimpleListenerTest() + _server = self.server.config(config_handler).listen(_simple) + self.server_session: WebSocketServer = await _server.start() + + serverUrl = self.schema + "://127.0.0.1:" + str(self.port) + "/path?u=a&p=2" + self.client_session: Session = await SocketD.create_client(serverUrl) \ + .config(config_handler).open() + try: + await self.client_session.send("demo", StringEntity("test")) + await self.client_session.send_and_request("demo", StringEntity("test"), 100) + except Exception as e: + pass + await asyncio.sleep(5) + logger.info(f"counter {_simple.server_counter.get()} close_counter {_simple.close_counter.get()} message {_simple.message_counter.get()}") + + def start(self): + super().start() + self.loop.run_until_complete(self._start()) + + async def _stop(self): + if self.client_session: + await self.client_session.close() + + if self.server_session: + self.server_session.close() + if self.server: + await self.server.stop() + + def stop(self): + super().stop() + self.loop.run_until_complete(self._stop()) + + def on_error(self): + super().on_error() diff --git a/python/socketd/test/modelu/SimpleListenerTest.py b/python/socketd/test/modelu/SimpleListenerTest.py index 3ffbb4be..6fd74ba0 100644 --- a/python/socketd/test/modelu/SimpleListenerTest.py +++ b/python/socketd/test/modelu/SimpleListenerTest.py @@ -17,6 +17,8 @@ class SimpleListenerTest(Listener, ABC): def __init__(self): self.server_counter = AtomicRefer(0) + self.close_counter = AtomicRefer(0) + self.message_counter = AtomicRefer(0) def on_open(self, session): pass @@ -25,6 +27,8 @@ class SimpleListenerTest(Listener, ABC): with self.server_counter: self.server_counter.set(self.server_counter.get() + 1) if message.is_request(): + with self.message_counter: + self.message_counter.set(self.message_counter.get() + 1) await session.reply(message, StringEntity("reply")) await session.reply_end(message, StringEntity("ok test")) await session.reply(message, StringEntity("reply")) @@ -34,7 +38,9 @@ class SimpleListenerTest(Listener, ABC): await session.reply(message, StringEntity("reply")) def on_close(self, session): - pass + logger.debug("客户端主动关闭了") + with self.close_counter: + self.close_counter.set(self.close_counter.get() + 1) def on_error(self, session, error): pass -- Gitee From 890c4f05fb6ffb4a4f64bab54af15d4d37790d2d Mon Sep 17 00:00:00 2001 From: bai <1145000687@qq.com> Date: Thu, 21 Dec 2023 18:32:48 +0800 Subject: [PATCH 5/6] update readme --- python/socketd/README.md | 46 ++++++++++++++++++++-------------------- 1 file changed, 23 insertions(+), 23 deletions(-) diff --git a/python/socketd/README.md b/python/socketd/README.md index 754a0a58..994ce36a 100644 --- a/python/socketd/README.md +++ b/python/socketd/README.md @@ -6,13 +6,16 @@

- python3.10+ -

- -
-

- - + + socketd + + + python10 + + + qq + +

@@ -26,7 +29,7 @@ Socket.D 是一种新的通讯应用协议,也是一个开发框架。可以 ### 主要特性 -* 异步通讯,由带语义的事件消息驱动 +* 异步通讯,由带语义的路由消息驱动 * 语言无关,使用二进制通信协议(支持 tcp, ws, udp)。支持多语言、多平台 * 背压流控,请求时不让你把服务端发死了 * 断线重连,自动连接恢复 @@ -75,7 +78,7 @@ sd:ws://19.10.2.3:1023/path?u=noear&t=1234 ``` //udp only <2k -[len:int][flag:int][sid:str(<64)][\n][event:str(<512)][\n][metaString:str(<4k)][\n][data:byte(<16m)] +[len:int][flag:int][sid:str(<64)][\n][route:str(<512)][\n][metaString:str(<4k)][\n][data:byte(<16m)] ``` * 指令流 @@ -103,20 +106,17 @@ sd:ws://19.10.2.3:1023/path?u=noear&t=1234 ### 快速入门与学习 -- 快速入手 -```python -async def appliction_test(): - server = SocketD.create_server(ServerConfig("ws").setPort(9999)) - server_session: Serve = server.config(idGenerator).listen( - SimpleListenerTest()).start() +* 学习 - await asyncio.sleep(3) +请点击:[《快速入门与学习》](_docs/)。Java 之外的语言与平台会尽快跟进(欢迎有兴趣的同学加入社区) + +* 规划情况了解 + +| 语言或平台 | 客户端 | 服务端 | 备注 | +|--------|-----|----|----------------------| +| java | 已完成 | 已完成 | 支持 tcp, udp, ws 通讯架构 | +| js | 开发中 | / | 支持 ws 通讯架构 | +| python | 开发中 | / | 支持 ws 通讯架构 | +| 其它 | 计划中 | 计划中 | | - client_session: Session = SocketD.create_client("ws://127.0.0.1:9999") \ - .config(idGenerator).open() - for _ in range(100000): - await client_session.send("demo", StringEntity("test")) - await client_session.close() - asyncio.get_event_loop().run_forever() -``` -- Gitee From 8bad9bd7d29e69576e51b75cf01c7f9c8992e94e Mon Sep 17 00:00:00 2001 From: bai <1145000687@qq.com> Date: Fri, 22 Dec 2023 18:34:24 +0800 Subject: [PATCH 6/6] =?UTF-8?q?=E4=BF=AE=E6=94=B9=E5=88=86=E9=85=8D?= =?UTF-8?q?=E4=B8=8A=E4=BC=A0=E9=80=BB=E8=BE=91=EF=BC=8C=E4=BF=AE=E6=94=B9?= =?UTF-8?q?=E6=95=B0=E6=8D=AE=E6=A0=BC=E5=BC=8F?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../src/test/java/features/CaseTest.java | 4 +- .../java/features/cases/TestCase14_file.java | 8 +- .../transport/core/FragmentAggregator.java | 2 +- .../fragment/FragmentAggregatorDefault.java | 4 +- .../fragment/FragmentAggregatorTempfile.java | 4 +- python/socketd/README.md | 30 ++++++ python/socketd/socketd/core/ChannelDefault.py | 6 +- python/socketd/socketd/core/Costants.py | 27 ++++- python/socketd/socketd/core/Session.py | 4 +- python/socketd/socketd/core/SessionDefault.py | 5 +- .../socketd/socketd/core/config/ConfigBase.py | 2 +- .../core/handler/FragmentAggregator.py | 51 +++++++++ .../core/handler/FragmentAggregatorDefault.py | 53 ++++++++++ .../core/handler/FragmentHandlerDefault.py | 80 +++----------- .../socketd/core/handler/FragmentHolder.py | 8 ++ python/socketd/socketd/core/module/Entity.py | 16 ++- .../socketd/core/module/EntityDefault.py | 24 +++-- .../socketd/socketd/core/module/FileEntity.py | 12 +++ python/socketd/socketd/core/module/Frame.py | 6 +- python/socketd/socketd/core/module/Message.py | 16 +-- .../socketd/core/module/MessageDefault.py | 27 +++++ .../socketd/exception/SocketdExecption.py | 9 ++ python/socketd/socketd/exception/__init__.py | 0 .../socketd/transport/CodecByteBuffer.py | 8 +- .../WsAioChannelAssistant.py | 1 + python/socketd/test/TestCase01.py | 15 ++- .../TestCase04_sendAndRequest_timeout.py | 2 +- python/socketd/test/cases/TestCase05_file.py | 100 ++++++++++++++++++ 28 files changed, 411 insertions(+), 113 deletions(-) create mode 100644 python/socketd/socketd/core/handler/FragmentAggregator.py create mode 100644 python/socketd/socketd/core/handler/FragmentAggregatorDefault.py create mode 100644 python/socketd/socketd/core/handler/FragmentHolder.py create mode 100644 python/socketd/socketd/core/module/FileEntity.py create mode 100644 python/socketd/socketd/exception/SocketdExecption.py create mode 100644 python/socketd/socketd/exception/__init__.py create mode 100644 python/socketd/test/cases/TestCase05_file.py diff --git a/java/socketd-transport-test/src/test/java/features/CaseTest.java b/java/socketd-transport-test/src/test/java/features/CaseTest.java index 8bbcf7b2..efaa89da 100644 --- a/java/socketd-transport-test/src/test/java/features/CaseTest.java +++ b/java/socketd-transport-test/src/test/java/features/CaseTest.java @@ -10,8 +10,8 @@ import org.junit.jupiter.api.Test; public class CaseTest { static final String[] schemas = new String[]{ "sd:tcp-java", -// "sd:tcp-netty", -// "sd:tcp-smartsocket", + "sd:tcp-netty", + "sd:tcp-smartsocket", "sd:ws-java", "sd:udp-java"}; diff --git a/java/socketd-transport-test/src/test/java/features/cases/TestCase14_file.java b/java/socketd-transport-test/src/test/java/features/cases/TestCase14_file.java index 08ac7b54..21c46d17 100644 --- a/java/socketd-transport-test/src/test/java/features/cases/TestCase14_file.java +++ b/java/socketd-transport-test/src/test/java/features/cases/TestCase14_file.java @@ -52,7 +52,7 @@ public class TestCase14_file extends BaseTestCase { if (fileName != null) { System.out.println(fileName); - File fileNew = new File("/Users/noear/Downloads/socketd-upload.mov"); + File fileNew = new File("./test.mp4"); fileNew.delete(); fileNew.createNewFile(); @@ -79,7 +79,7 @@ public class TestCase14_file extends BaseTestCase { .config(c -> c.fragmentSize(1024 * 1024)) .open(); - clientSession.send("/user/upload", new FileEntity(new File("/Users/noear/Movies/snack3-rce-poc.mov"))); + clientSession.send("/user/upload", new FileEntity(new File("\\C:\\Users\\bai\\Pictures\\46c7a111437ea55469f1f5f5b35c3e55.mp4"))); Thread.sleep(10000); @@ -87,8 +87,8 @@ public class TestCase14_file extends BaseTestCase { System.out.println("counter: " + messageCounter.get()); Assertions.assertEquals(messageCounter.get(), 1, getSchema() + ":server 收的消息数量对不上"); - File file = new File("/Users/noear/Downloads/socketd-upload.mov"); - assert file.length() > 1024 * 1024 * 10; + File file = new File("./test.mp4"); + assert file.exists(); } @Override diff --git a/java/socketd/src/main/java/org/noear/socketd/transport/core/FragmentAggregator.java b/java/socketd/src/main/java/org/noear/socketd/transport/core/FragmentAggregator.java index c17ebae1..1edfe7d7 100644 --- a/java/socketd/src/main/java/org/noear/socketd/transport/core/FragmentAggregator.java +++ b/java/socketd/src/main/java/org/noear/socketd/transport/core/FragmentAggregator.java @@ -8,7 +8,7 @@ import java.io.IOException; * @author noear * @since 2.1 */ -public interface FragmentAggregator { +public interface FragmentAggregatorDefault { /** * 获取流Id */ diff --git a/java/socketd/src/main/java/org/noear/socketd/transport/core/fragment/FragmentAggregatorDefault.java b/java/socketd/src/main/java/org/noear/socketd/transport/core/fragment/FragmentAggregatorDefault.java index 4fad65c1..0b91bc8a 100644 --- a/java/socketd/src/main/java/org/noear/socketd/transport/core/fragment/FragmentAggregatorDefault.java +++ b/java/socketd/src/main/java/org/noear/socketd/transport/core/fragment/FragmentAggregatorDefault.java @@ -1,7 +1,7 @@ package org.noear.socketd.transport.core.fragment; import org.noear.socketd.transport.core.EntityMetas; -import org.noear.socketd.transport.core.FragmentAggregator; +import org.noear.socketd.transport.core.FragmentAggregatorDefault; import org.noear.socketd.transport.core.Frame; import org.noear.socketd.exception.SocketdCodecException; import org.noear.socketd.transport.core.MessageInternal; @@ -20,7 +20,7 @@ import java.util.List; * @author noear * @since 2.0 */ -public class FragmentAggregatorDefault implements FragmentAggregator { +public class FragmentAggregatorDefault implements FragmentAggregatorDefault { //主导消息 private MessageInternal main; //分片列表 diff --git a/java/socketd/src/main/java/org/noear/socketd/transport/core/fragment/FragmentAggregatorTempfile.java b/java/socketd/src/main/java/org/noear/socketd/transport/core/fragment/FragmentAggregatorTempfile.java index e95dc7c0..ac9aa3e5 100644 --- a/java/socketd/src/main/java/org/noear/socketd/transport/core/fragment/FragmentAggregatorTempfile.java +++ b/java/socketd/src/main/java/org/noear/socketd/transport/core/fragment/FragmentAggregatorTempfile.java @@ -2,7 +2,7 @@ package org.noear.socketd.transport.core.fragment; import org.noear.socketd.exception.SocketdCodecException; import org.noear.socketd.transport.core.EntityMetas; -import org.noear.socketd.transport.core.FragmentAggregator; +import org.noear.socketd.transport.core.FragmentAggregatorDefault; import org.noear.socketd.transport.core.Frame; import org.noear.socketd.transport.core.MessageInternal; import org.noear.socketd.transport.core.entity.TempfileEntity; @@ -21,7 +21,7 @@ import java.nio.channels.FileChannel; * @author noear * @since 2.0 */ -public class FragmentAggregatorTempfile implements FragmentAggregator { +public class FragmentAggregatorTempfile implements FragmentAggregatorDefault { //主导消息 private MessageInternal main; //数据流大小 diff --git a/python/socketd/README.md b/python/socketd/README.md index 994ce36a..9097f7d0 100644 --- a/python/socketd/README.md +++ b/python/socketd/README.md @@ -119,4 +119,34 @@ sd:ws://19.10.2.3:1023/path?u=noear&t=1234 | python | 开发中 | / | 支持 ws 通讯架构 | | 其它 | 计划中 | 计划中 | | +## 快速入门 +```python +async def application(): + # 服务端 + server: Server = SocketD.create_server(ServerConfig("ws").set_port(9999)) + server_session: WebSocketServer = await server.config(idGenerator).listen( + SimpleListenerTest()).start() + + # 客户端 + client_session: Session = await SocketD.create_client("ws://127.0.0.1:9999") \ + .config(idGenerator).open() + + start_time = time.monotonic() + for _ in range(100): + await client_session.send("demo", StringEntity("test")) + await client_session.send_and_request("demo", StringEntity("test"), 100) + await client_session.send_and_subscribe("demo", StringEntity("test"), send_and_subscribe_test, 100) + end_time = time.monotonic() + logger.info(f"Coroutine send took {(end_time - start_time) * 1000.0} monotonic to complete.") + await client_session.close() + server_session.close() + await server.stop() + +asyncio.run(application()) +``` +非windows用户可以引入以下依赖,性能提升2倍,asyncio(协程)性能直逼Go +```python +import uvloop +asyncio.set_event_loop_policy(uvloop.EventLoopPolicy()) +``` \ No newline at end of file diff --git a/python/socketd/socketd/core/ChannelDefault.py b/python/socketd/socketd/core/ChannelDefault.py index 8d127085..51ddf634 100644 --- a/python/socketd/socketd/core/ChannelDefault.py +++ b/python/socketd/socketd/core/ChannelDefault.py @@ -60,9 +60,9 @@ class ChannelDefault(ChannelBase): message.get_entity()) if fragmentEntity is not None: fragmentFrame = Frame(frame.get_flag(), MessageDefault() - .flag(frame.get_flag()) - .sid(message.get_sid()) - .entity(fragmentEntity)) + .set_flag(frame.get_flag()) + .set_sid(message.get_sid()) + .set_entity(fragmentEntity)) await self.assistant.write(self.source, fragmentFrame) else: return diff --git a/python/socketd/socketd/core/Costants.py b/python/socketd/socketd/core/Costants.py index f756f66a..c481a9fa 100644 --- a/python/socketd/socketd/core/Costants.py +++ b/python/socketd/socketd/core/Costants.py @@ -6,7 +6,7 @@ from typing import Callable Function = Callable -class Flag(Enum): +class Flag: Unknown = 0 Connect = 10 Connack = 11 @@ -44,6 +44,31 @@ class Flag(Enum): else: return Flag.Unknown + @staticmethod + def name(code): + if code == 10: + return "Connect" + elif code == 11: + return"Connack" + elif code == 20: + return "Ping" + elif code == 21: + return "Pong" + elif code == 30: + return "Close" + elif code == 40: + return "Message" + elif code == 41: + return "Request" + elif code == 42: + return "Subscribe" + elif code == 48: + return "Reply" + elif code == 49: + return "ReplyEnd" + else: + return "Unknown" + class Constants: DEF_SID = "" diff --git a/python/socketd/socketd/core/Session.py b/python/socketd/socketd/core/Session.py index 45f46134..28362746 100644 --- a/python/socketd/socketd/core/Session.py +++ b/python/socketd/socketd/core/Session.py @@ -48,11 +48,11 @@ class Session(abc.ABC): ... @abc.abstractmethod - async def send(self, topic: str, content: Entity) -> None: + async def send(self, event: str, content: Entity) -> None: ... @abc.abstractmethod - async def send_and_request(self, topic: str, content: Entity, timeout: int) -> Entity: + async def send_and_request(self, event: str, content: Entity, timeout: int) -> Entity: ... @abc.abstractmethod diff --git a/python/socketd/socketd/core/SessionDefault.py b/python/socketd/socketd/core/SessionDefault.py index 6922f920..58576121 100644 --- a/python/socketd/socketd/core/SessionDefault.py +++ b/python/socketd/socketd/core/SessionDefault.py @@ -11,6 +11,7 @@ from .module.Message import Message from .module.Frame import Frame from .Costants import Flag from .module.MessageDefault import MessageDefault +from ..exception.SocketdExecption import SocketDException from ..transport.core.CompletableFuture import CompletableFuture from ..transport.core.StreamRequest import StreamRequest from ..transport.core.StreamSubscribe import StreamSubscribe @@ -54,11 +55,11 @@ class SessionDefault(SessionBase, ABC): return await future.get(timeout) except asyncio.TimeoutError as e: if self.channel.is_valid(): - raise Exception(f"Request reply timeout>{timeout} " + raise SocketDException(f"Request reply timeout>{timeout} " f"sessionId={self.channel.get_session().get_session_id()} " f"event={event} sid={message.get_sid()}") else: - raise Exception(f"This channel is closed sessionId={self.channel.get_session().get_session_id()} " + raise SocketDException(f"This channel is closed sessionId={self.channel.get_session().get_session_id()} " f"event={event} sid={message.get_sid()}") except Exception as e: raise e diff --git a/python/socketd/socketd/core/config/ConfigBase.py b/python/socketd/socketd/core/config/ConfigBase.py index 86577279..bd44a73c 100644 --- a/python/socketd/socketd/core/config/ConfigBase.py +++ b/python/socketd/socketd/core/config/ConfigBase.py @@ -5,9 +5,9 @@ from typing import Callable from concurrent.futures import ThreadPoolExecutor from .Config import Config -from socketd.core.handler.FragmentHandlerDefault import FragmentHandlerDefault from socketd.transport.Codec import Codec from socketd.transport.CodecByteBuffer import CodecByteBuffer +from ..handler.FragmentHandlerDefault import FragmentHandlerDefault class ConfigBase(Config): diff --git a/python/socketd/socketd/core/handler/FragmentAggregator.py b/python/socketd/socketd/core/handler/FragmentAggregator.py new file mode 100644 index 00000000..172b4d59 --- /dev/null +++ b/python/socketd/socketd/core/handler/FragmentAggregator.py @@ -0,0 +1,51 @@ +from socketd.core.module.Frame import Frame +from socketd.core.module.Message import Message + + +class FragmentAggregator: + + def get_sid(self) -> str: + """ + 获取sid + + :return: str, 返回sid字符串 + """ + ... + + def get_data_stream_size(self) -> int: + """ + 获取数据流的大小 + + :return: int, 数据流的大小 + """ + ... + + def get_data_length(self) -> int: + """ + 获取数据的长度 + + :return: int, 数据的长度 + """ + ... + + def add(self, index: int, message: Message): + ''' + 添加消息到指定索引位置 + + Args: + index (int): 指定的索引位置 + message (Message): 要添加的消息 + + Returns: + None + ''' + ... + + def get(self) -> Frame: ... + + """ + 获取一个Frame对象。 + + Returns: + Frame: 返回获取到的Frame对象。 + """ diff --git a/python/socketd/socketd/core/handler/FragmentAggregatorDefault.py b/python/socketd/socketd/core/handler/FragmentAggregatorDefault.py new file mode 100644 index 00000000..0068fa94 --- /dev/null +++ b/python/socketd/socketd/core/handler/FragmentAggregatorDefault.py @@ -0,0 +1,53 @@ +from socketd.core.module.EntityDefault import EntityDefault +from socketd.core.module.Frame import Frame +from socketd.core.module.Message import Message +from .FragmentAggregator import FragmentAggregator +from .FragmentHolder import FragmentHolder +from ..Buffer import Buffer +from ..module.Entity import EntityMetas + +from ..module.MessageDefault import MessageDefault +from ...exception.SocketdExecption import SocketDException + + +class FragmentAggregatorDefault(FragmentAggregator): + """ + 分片聚合器 + """ + def __init__(self, frame: Message): + self.__fragments: list[FragmentHolder] = [] + self.__main: Message = frame + data_length: int = frame.get_meta(EntityMetas.META_DATA_LENGTH) + if data_length is None or type(data_length) != int: + raise SocketDException(f"Missing {EntityMetas.META_DATA_LENGTH} meta, event= {frame.get_event()}") + self.__data_length = data_length + self.__data_stream_size = 0 + + def add(self, index: int, message: Message): + self.__fragments.insert(index, FragmentHolder(index, message)) + self.__data_stream_size += message.get_data_size() + + def get_sid(self) -> str: + return self.__main.get_sid() + + def get_data_length(self) -> int: + return self.__data_length + + def get_data_stream_size(self) -> int: + return self.__data_stream_size + + def get(self) -> Frame: + self.__fragments.sort(key=lambda x: x.index) + + byte_buffer: Buffer = Buffer() + + for fragment in self.__fragments: + byte_buffer.write(fragment.message.get_data().getvalue()) + + return Frame(self.__main.get_flag(), + MessageDefault() + .set_flag(self.__main.get_flag()) + .set_sid(self.__main.get_sid()) + .set_entity(EntityDefault().set_meta_map(self.__main.get_entity().get_meta_map()) + .set_data(byte_buffer)) + ) diff --git a/python/socketd/socketd/core/handler/FragmentHandlerDefault.py b/python/socketd/socketd/core/handler/FragmentHandlerDefault.py index 024b292b..3b581ce3 100644 --- a/python/socketd/socketd/core/handler/FragmentHandlerDefault.py +++ b/python/socketd/socketd/core/handler/FragmentHandlerDefault.py @@ -1,89 +1,41 @@ import pickle -from io import BytesIO - -from socketd.core.config.Config import Config from socketd.core.module.Entity import Entity, EntityMetas from socketd.core.module.EntityDefault import EntityDefault from socketd.core.module.Frame import Frame -from socketd.core.module.Message import Message +from .FragmentAggregatorDefault import FragmentAggregatorDefault from .FragmentHandler import FragmentHandler -from ..module.MessageDefault import MessageDefault +from ..Buffer import Buffer +from ..Channel import Channel +from ..module.Message import Message class FragmentHandlerDefault(FragmentHandler): - def __init__(self): - pass - - def nextFragment(self, config: Config, fragmentIndex: int, entity: Entity) -> Entity: - # fragmentIndex.set(fragmentIndex.get() + 1) - fragmentBuf = BytesIO() - # IoUtils.transferTo(entity.getData(), fragmentBuf, 0, Config.MAX_SIZE_FRAGMENT) - pickle.dump(entity, fragmentBuf) + def nextFragment(self, channel: Channel, fragmentIndex: int, message: Message) -> Entity | None: + fragmentBuf = Buffer() + pickle.dump(message, fragmentBuf) fragmentBytes = fragmentBuf.getbuffer() if len(fragmentBytes) == 0: return None + fragmentEntity = EntityDefault().set_data(fragmentBytes) if fragmentIndex == 1: - fragmentEntity.metaMap(entity.get_meta_map()) - fragmentEntity.putMeta(EntityMetas.META_DATA_FRAGMENT_IDX, str(fragmentIndex)) + fragmentEntity.set_meta_map(message.get_meta_map()) + fragmentEntity.put_meta(EntityMetas.META_DATA_FRAGMENT_IDX, str(fragmentIndex)) return fragmentEntity - def aggrFragment(self, channel, index: int, frame: Frame) -> Frame: - aggregator = channel.getAttachment(frame.get_message().get_sid()) + def aggrFragment(self, channel: Channel, index: int, message: Message) -> Frame | None: + aggregator = channel.get_attachment(message.get_sid()) if aggregator is None: - aggregator = FragmentAggregator(frame) - channel.setAttachment(frame.get_message().get_sid(), aggregator) + aggregator = FragmentAggregatorDefault(message) + channel.set_attachment(message.get_sid(), aggregator) - aggregator.add(index, frame) + aggregator.add(index, message) if aggregator.getDataLength() > aggregator.getDataStreamSize(): return None # Length is not enough, wait for the next fragment package else: + channel.set_attachment(message.get_sid(), None) return aggregator.get() # Reset as a merged frame - - -class FragmentAggregator(): - def __init__(self, frame: Frame): - self.fragments = [] - self.sid = frame.get_message().get_sid() - self.main: Frame = frame - self.message: Message = frame.get_message() - - def add(self, index: int, frame: Frame): - self.fragments.append((index, frame)) - - def getDataLength(self) -> int: - length = 0 - for fragment in self.fragments: - length += len(fragment[1].getEntity().getData()) - return length - - def getDataStreamSize(self) -> int: - if len(self.fragments) == 0: - return 0 - else: - return len(self.fragments[0][1].getEntity().getData()) - - def get(self) -> Frame: - length = self.getDataLength() - entityBytes = bytearray(length) - - for fragment in self.fragments: - entity = fragment[1].getEntity() - data = entity.getData() - index = fragment[0] - size = len(data) - start = index * size - end = start + size - entityBytes[start:end] = data - - return Frame(self.main.get_flag(), - MessageDefault() - .set_flag(self.main.flag) - .set_sid(self.message.get_sid()) - .set_entity(EntityDefault().set_meta_map(self.main.get_message().get_entity().get_meta_map()) - .set_data(entityBytes)) - ) diff --git a/python/socketd/socketd/core/handler/FragmentHolder.py b/python/socketd/socketd/core/handler/FragmentHolder.py new file mode 100644 index 00000000..b62df04e --- /dev/null +++ b/python/socketd/socketd/core/handler/FragmentHolder.py @@ -0,0 +1,8 @@ +from socketd.core.module.Message import Message + + +class FragmentHolder: + + def __init__(self, index: int, message: Message): + self.index = index + self.message = message diff --git a/python/socketd/socketd/core/module/Entity.py b/python/socketd/socketd/core/module/Entity.py index 8935130a..1c647171 100644 --- a/python/socketd/socketd/core/module/Entity.py +++ b/python/socketd/socketd/core/module/Entity.py @@ -1,5 +1,5 @@ -from typing import Dict, Optional -from io import FileIO +from typing import Dict, Optional, Any +from io import BytesIO class Entity: @@ -9,13 +9,13 @@ class Entity: def get_meta_map(self) -> Dict[str, str]: raise NotImplementedError - def get_meta(self, name: str) -> Optional[str]: + def get_meta(self, name: str) -> Optional[Any]: raise NotImplementedError def get_meta_or_default(self, name: str, default: str) -> str: raise NotImplementedError - def get_data(self) -> FileIO: + def get_data(self) -> BytesIO: raise NotImplementedError def get_data_as_string(self) -> str: @@ -24,9 +24,17 @@ class Entity: def get_data_size(self) -> int: raise NotImplementedError + def get_data_as_bytes(self) -> bytes: + raise NotImplementedError + class EntityMetas: META_SOCKETD_VERSION = "SocketD-Version" META_DATA_LENGTH = "Data-Length" META_DATA_FRAGMENT_IDX = "Data-Fragment-Idx" META_DATA_DISPOSITION_FILENAME = "Data-Disposition-Filename" + + +class Reply: + + def is_end(self) -> bool: ... diff --git a/python/socketd/socketd/core/module/EntityDefault.py b/python/socketd/socketd/core/module/EntityDefault.py index 88fb483f..016d4ef5 100644 --- a/python/socketd/socketd/core/module/EntityDefault.py +++ b/python/socketd/socketd/core/module/EntityDefault.py @@ -1,15 +1,18 @@ from abc import ABC import pickle +from io import BytesIO +from typing import Any from .Entity import Entity +from ..Costants import Constants class EntityDefault(Entity, ABC): def __init__(self): self.meta_map = None - self.meta_string = "_dEF__mET_a__sTRING" + self.meta_string = Constants.DEF_META_STRING self.meta_stringChanged = False - self.data: bytes = None + self.data: BytesIO = Constants.DEF_DATA self.data_size = 0 def set_meta_string(self, meta_string): @@ -56,17 +59,17 @@ class EntityDefault(Entity, ABC): self.get_meta_map()[name] = val self.meta_stringChanged = True - def get_meta(self, name): + def get_meta(self, name) -> Any: return self.get_meta_map().get(name) - def get_metaOr_default(self, name, default_val): + def get_meta_or_default(self, name, default_val): return self.get_meta_map().get(name, default_val) - def set_data(self, data): - if type(data) != bytes: - self.data = pickle.loads(data) - else: + def set_data(self, data: bytes | bytearray | memoryview | BytesIO): + if type(data) == BytesIO: self.data = data + else: + self.data = BytesIO(data) self.data_size = len(data) return self @@ -74,7 +77,10 @@ class EntityDefault(Entity, ABC): return self.data def get_data_as_string(self): - return str(self.data, 'utf-8') # _assuming data is of type bytes + return str(self.data.getvalue(), 'utf-8') # _assuming data is of type bytes + + def get_data_as_bytes(self) -> bytes: + return self.data.getvalue() def get_data_size(self): return self.data_size diff --git a/python/socketd/socketd/core/module/FileEntity.py b/python/socketd/socketd/core/module/FileEntity.py new file mode 100644 index 00000000..9550b0d1 --- /dev/null +++ b/python/socketd/socketd/core/module/FileEntity.py @@ -0,0 +1,12 @@ + +from socketd.core.module.Entity import EntityMetas +from socketd.core.module.EntityDefault import EntityDefault + + +class FileEntity(EntityDefault): + + def __init__(self, byteIO: bytes, filename: str): + super().__init__() + self._file = byteIO + self.set_data(byteIO) + self.set_meta(EntityMetas.META_DATA_DISPOSITION_FILENAME, filename) diff --git a/python/socketd/socketd/core/module/Frame.py b/python/socketd/socketd/core/module/Frame.py index 4016b0f5..1fa2720c 100644 --- a/python/socketd/socketd/core/module/Frame.py +++ b/python/socketd/socketd/core/module/Frame.py @@ -4,15 +4,15 @@ from ..Costants import Flag class Frame: - def __init__(self, flag: Flag, message: Message): + def __init__(self, flag: int, message: Message): self.flag = flag self.message = message - def get_flag(self) -> Flag: + def get_flag(self) -> int: return self.flag def get_message(self) -> Message: return self.message def __str__(self) -> str: - return f"Frame{{flag={self.flag}, message={self.message}}}" + return f"Frame{{flag={Flag.name(self.flag)}, message={self.message}}}" diff --git a/python/socketd/socketd/core/module/Message.py b/python/socketd/socketd/core/module/Message.py index c6e08c96..7bc8a44e 100644 --- a/python/socketd/socketd/core/module/Message.py +++ b/python/socketd/socketd/core/module/Message.py @@ -6,24 +6,28 @@ from .Entity import Entity class Message(Entity): @abstractmethod def is_request(self) -> bool: - pass + ... @abstractmethod def is_subscribe(self) -> bool: - pass + ... @abstractmethod def is_close(self) -> bool: - pass + ... @abstractmethod def get_sid(self) -> str: - pass + ... @abstractmethod def get_event(self) -> str: - pass + ... @abstractmethod def get_entity(self) -> Entity: - pass + ... + + @abstractmethod + def get_flag(self) -> int: + ... diff --git a/python/socketd/socketd/core/module/MessageDefault.py b/python/socketd/socketd/core/module/MessageDefault.py index 54499867..fe54caf4 100644 --- a/python/socketd/socketd/core/module/MessageDefault.py +++ b/python/socketd/socketd/core/module/MessageDefault.py @@ -1,3 +1,6 @@ +from io import BytesIO +from typing import Optional, Any, Dict + from .Entity import Entity from .Message import Message from ..Costants import Constants, Flag @@ -49,3 +52,27 @@ class MessageDefault(Message): def __str__(self): return f"Message{{sid='{self.sid}', event='{self.event}', entity={self.entity}}}" + + def get_meta_string(self) -> str: + return self.entity.get_meta_string() + + def get_meta_map(self) -> Dict[str, str]: + return self.entity.get_meta_map() + + def get_meta(self, name: str) -> Optional[Any]: + return self.entity.get_meta(name) + + def get_meta_or_default(self, name: str, default: str) -> str: + return self.entity.get_meta_or_default(name, default) + + def get_data(self) -> BytesIO: + return self.entity.get_data() + + def get_data_as_string(self) -> str: + return self.entity.get_data_as_string() + + def get_data_size(self) -> int: + return self.entity.get_data_size() + + def get_data_as_bytes(self) -> bytes: + return self.entity.get_data_as_bytes() \ No newline at end of file diff --git a/python/socketd/socketd/exception/SocketdExecption.py b/python/socketd/socketd/exception/SocketdExecption.py new file mode 100644 index 00000000..03d550a7 --- /dev/null +++ b/python/socketd/socketd/exception/SocketdExecption.py @@ -0,0 +1,9 @@ + +class SocketDException(RuntimeError): + + def __init__(self, message): + super().__init__(self) + self.message = message + + def __str__(self): + return self.message diff --git a/python/socketd/socketd/exception/__init__.py b/python/socketd/socketd/exception/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/python/socketd/socketd/transport/CodecByteBuffer.py b/python/socketd/socketd/transport/CodecByteBuffer.py index 25a0b56d..52c14f84 100644 --- a/python/socketd/socketd/transport/CodecByteBuffer.py +++ b/python/socketd/socketd/transport/CodecByteBuffer.py @@ -21,7 +21,7 @@ class CodecByteBuffer(Codec): # length target.put_int(_len) # flag - target.put_int(frame.flag.value) + target.put_int(frame.flag) target.flush() return target @@ -48,7 +48,7 @@ class CodecByteBuffer(Codec): target.put_int(len1) # flag - target.put_int(frame.flag.value) + target.put_int(frame.flag) # sid target.write(sidB) @@ -64,7 +64,7 @@ class CodecByteBuffer(Codec): # data if frame.message.get_entity().get_data() is not None: - target.write(frame.message.get_entity().get_data()) + target.write(frame.message.get_entity().get_data().getvalue()) target.flush() return target @@ -98,7 +98,7 @@ class CodecByteBuffer(Codec): for i in range(dataRealSize - Config.MAX_SIZE_FRAGMENT): buffer.read() else: - data = buffer.read(dataRealSize) + data = bytearray(buffer.read(dataRealSize)) message = MessageDefault().set_sid(sid).set_event(topic).set_entity( EntityDefault().set_meta_string(metaString).set_data(data) diff --git a/python/socketd/socketd_websocket/WsAioChannelAssistant.py b/python/socketd/socketd_websocket/WsAioChannelAssistant.py index 1ebda398..15f2e4e8 100644 --- a/python/socketd/socketd_websocket/WsAioChannelAssistant.py +++ b/python/socketd/socketd_websocket/WsAioChannelAssistant.py @@ -28,6 +28,7 @@ class WsAioChannelAssistant(ChannelAssistant): return target.state == State.OPEN async def close(self, target: WebSocketServerProtocol) -> None: + # await target.wait_closed() # 等待消息 await target.close() def get_remote_address(self, target: WebSocketServerProtocol) -> str: diff --git a/python/socketd/test/TestCase01.py b/python/socketd/test/TestCase01.py index 15fadfc5..93b8a79b 100644 --- a/python/socketd/test/TestCase01.py +++ b/python/socketd/test/TestCase01.py @@ -7,6 +7,7 @@ from test.cases.TestCase01_client_send import TestCase01_client_send from test.cases.TestCase02_auto_reconnect import TestCase02_auto_reconnect from test.cases.TestCase03_client_session_close import TestCase03_client_session_close from test.cases.TestCase04_sendAndRequest_timeout import TestCase04_sendAndRequest_timeout +from test.cases.TestCase05_file import TestCase05_file from test.modelu.BaseTest import BaseTest @@ -38,7 +39,7 @@ class TestCase01(unittest.TestCase): t.on_error() raise e - def test_Case02_client_session_close(self): + def test_Case03_client_session_close(self): for i in range(len(TestCase01.schemas)): t = TestCase03_client_session_close(TestCase01.schemas[i], 9000 + i) try: @@ -48,7 +49,7 @@ class TestCase01(unittest.TestCase): t.on_error() raise e - def test_Case02sendAndRequest_timeout(self): + def test_Case04sendAndRequest_timeout(self): for i in range(len(TestCase01.schemas)): t = TestCase04_sendAndRequest_timeout(TestCase01.schemas[i], 9000 + i) try: @@ -57,3 +58,13 @@ class TestCase01(unittest.TestCase): except Exception as e: t.on_error() raise e + + def test_Case05_file(self): + for i in range(len(TestCase01.schemas)): + t = TestCase05_file(TestCase01.schemas[i], 9000 + i) + try: + t.start() + t.stop() + except Exception as e: + t.on_error() + raise e \ No newline at end of file diff --git a/python/socketd/test/cases/TestCase04_sendAndRequest_timeout.py b/python/socketd/test/cases/TestCase04_sendAndRequest_timeout.py index 5fb9ad30..ef20312e 100644 --- a/python/socketd/test/cases/TestCase04_sendAndRequest_timeout.py +++ b/python/socketd/test/cases/TestCase04_sendAndRequest_timeout.py @@ -10,7 +10,7 @@ from socketd.core.SocketD import SocketD from socketd.core.config.ServerConfig import ServerConfig from socketd.core.module.StringEntity import StringEntity from socketd.transport.server.Server import Server -from test.modelu.SimpleListenerTest import SimpleListenerTest, config_handler, send_and_subscribe_test +from test.modelu.SimpleListenerTest import SimpleListenerTest, config_handler class TestCase04_sendAndRequest_timeout(BaseTestCase): diff --git a/python/socketd/test/cases/TestCase05_file.py b/python/socketd/test/cases/TestCase05_file.py new file mode 100644 index 00000000..883232ed --- /dev/null +++ b/python/socketd/test/cases/TestCase05_file.py @@ -0,0 +1,100 @@ +import asyncio +import time +from abc import ABC + +from websockets.legacy.server import WebSocketServer +from loguru import logger +from pathlib import Path + +from socketd.core.Listener import Listener +from socketd.core.module.Entity import EntityMetas +from socketd.core.module.FileEntity import FileEntity +from socketd.core.module.Message import Message +from test.modelu.BaseTestCase import BaseTestCase + +from socketd.core.Session import Session +from socketd.core.SocketD import SocketD +from socketd.core.config.ServerConfig import ServerConfig +from socketd.core.module.StringEntity import StringEntity +from socketd.transport.server.Server import Server +from test.modelu.SimpleListenerTest import config_handler +from socketd.core.sync_api.AtomicRefer import AtomicRefer + + +class SimpleListenerTest(Listener, ABC): + + def __init__(self): + self.message_counter = AtomicRefer(0) + + def on_open(self, session): + pass + + async def on_message(self, session, message: Message): + logger.debug(message) + with self.message_counter: + self.message_counter.set(self.message_counter.get() + 1) + + file_name = message.get_meta(EntityMetas.META_DATA_DISPOSITION_FILENAME) + out_file_name = "./test" + if file_name: + logger.debug(f"file_name {file_name}") + with open(out_file_name, "wb") as f: + f.write(message.get_data_as_bytes()) + + path = Path(out_file_name) + assert path.exists() + + def on_close(self, session): + pass + + def on_error(self, session, error): + pass + + +class TestCase05_file(BaseTestCase): + + def __init__(self, schema, port): + super().__init__(schema, port) + self.server: Server + self.server_session: WebSocketServer + self.client_session: Session + self.loop = asyncio.get_event_loop() + + async def _start(self): + self.server: Server = SocketD.create_server(ServerConfig(self.schema).set_port(self.port)) + _simple = SimpleListenerTest() + _server = self.server.config(config_handler).listen(_simple) + self.server_session: WebSocketServer = await _server.start() + await asyncio.sleep(1) + serverUrl = self.schema + "://127.0.0.1:" + str(self.port) + "/path?u=a&p=2" + self.client_session: Session = await SocketD.create_client(serverUrl) \ + .config(config_handler).open() + try: + with open(r"C:\Users\bai\Pictures\46c7a111437ea55469f1f5f5b35c3e55.mp4", "rb") as f: + await self.client_session.send("/path?u=a&p=2", FileEntity(f.read(), "test")) + except Exception as e: + logger.error(e) + raise e + logger.info( + f" message {_simple.message_counter.get()}") + + def start(self): + super().start() + self.loop.run_until_complete(self._start()) + time.sleep(2) + + async def _stop(self): + if self.client_session: + await self.client_session.close() + + if self.server_session: + self.server_session.close() + if self.server: + await self.server.stop() + + def stop(self): + super().stop() + self.loop.run_until_complete(self._stop()) + + def on_error(self): + super().on_error() -- Gitee