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 0000000000000000000000000000000000000000..64d83dc06a6c779bce9b20dcf61f91f1a50fb3bb
--- /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 0000000000000000000000000000000000000000..5d66d5b5da90dd5c050bf0d7fffbdb2fb5e3a3b4
--- /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/java/socketd-transport-test/src/test/java/features/CaseTest.java b/java/socketd-transport-test/src/test/java/features/CaseTest.java
index f52a044f08a9d355078ba46b3332f1d83b3e4d48..efaa89da66629badfa854769f87f7dac9ed2c54f 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-transport-test/src/test/java/features/cases/TestCase14_file.java b/java/socketd-transport-test/src/test/java/features/cases/TestCase14_file.java
index 08ac7b544eb175f31d60f8029ea5e7c3eb2e8752..21c46d17559a33e87e6b319b5068adf089bf86fc 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/client/ClientChannel.java b/java/socketd/src/main/java/org/noear/socketd/transport/client/ClientChannel.java
index c87c753e0ef5639c721b484c5a906ba6daa0ecf1..546eadf92fb0f8cc1c830e97cc9a8123f37cfc7a 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/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 c17ebae144048e94144a4fd6fc5f8d7391277f44..1edfe7d7154decedc801873cfb2f6166ba5082ef 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 4fad65c13d8913a9677f9af2464e60cb1995e502..0b91bc8a8f0aaeddad511d3401abf5d1f2427e49 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 e95dc7c0d2b56460f28f1ba4d46400ca2c37d6a1..ac9aa3e51f87daee7f6c6127f5a174e248e7d4fe 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 754a0a586d41c65b205275a5343fdf02eff33a0d..9097f7d0ff121a72597546e4b50d1b8cc5468e59 100644
--- a/python/socketd/README.md
+++ b/python/socketd/README.md
@@ -6,13 +6,16 @@
- python3.10+
-
-
-
-
-
-
+
+
+
+
+
+
+
+
+
+
@@ -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,47 @@ 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()
+* 学习
+
+请点击:[《快速入门与学习》](_docs/)。Java 之外的语言与平台会尽快跟进(欢迎有兴趣的同学加入社区)
- await asyncio.sleep(3)
+* 规划情况了解
- client_session: Session = SocketD.create_client("ws://127.0.0.1:9999") \
+| 语言或平台 | 客户端 | 服务端 | 备注 |
+|--------|-----|----|----------------------|
+| java | 已完成 | 已完成 | 支持 tcp, udp, ws 通讯架构 |
+| js | 开发中 | / | 支持 ws 通讯架构 |
+| 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()
- for _ in range(100000):
+ 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()
- asyncio.get_event_loop().run_forever()
+ 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/reqeusts.txt b/python/socketd/reqeusts.txt
index f12ae6cabb6af195854006731015aa644ff1e379..ee89d343e3b4d3eba283039db8b1ed0db4cd11ce 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 ee8bb56537229cf651a13a397a236e2c9fa8e383..51ddf634c3dcd938b4ccdef9bc2d9ba0cbcf7de6 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
@@ -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/Costants.py b/python/socketd/socketd/core/Costants.py
index f756f66ae45d8a4f39abe6da4a2e087843a3a9b4..c481a9fa4cc87efda6288a9d97b23423fa070b66 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 6615f0f66c8cbaa4522a291f52ba5731ddb9fed1..283627463233e475dad1aa146820addc37afc4d8 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,19 +44,34 @@ class Session(abc.ABC):
...
@abc.abstractmethod
- def send_ping(self) -> None:
+ def send_ping(self) -> Callable | Coroutine:
...
@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:
...
- 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 4080220ce8699c9f27feec75a9f3a7405252f135..58576121e4f625ac52bdc97468efc4ec32d09bb6 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
@@ -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
@@ -34,8 +35,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)
@@ -54,23 +55,27 @@ 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
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)
@@ -98,5 +103,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/socketd/core/async_utils.py b/python/socketd/socketd/core/async_utils.py
deleted file mode 100644
index 024b56b341e416fa7d3a4339d0df02a46106e957..0000000000000000000000000000000000000000
--- 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/core/config/ClientConfig.py b/python/socketd/socketd/core/config/ClientConfig.py
index 5967e1351d026c2f6b4a21ddd8a3cc64e8a3e8f3..6c25b4b4d60698556622b0b9c1a337f3d7b49315 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 63cd82a99cba2696e4dd9e87972a6ee80fd4ab92..bd44a73cf92ff8dce21f6641b7b06b9208bd5d4c 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):
@@ -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/core/handler/FragmentAggregator.py b/python/socketd/socketd/core/handler/FragmentAggregator.py
new file mode 100644
index 0000000000000000000000000000000000000000..172b4d5937304fb238245fe5f941fb8b96663224
--- /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 0000000000000000000000000000000000000000..0068fa941b42373ec5797b213d8e90c0979fcc2e
--- /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 024b292b147f3475d7d39d84129b6305df14d361..3b581ce320f9aa1b24db0d4b97547bdc7d9c8073 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 0000000000000000000000000000000000000000..b62df04e0def9b016600b99e99c62b00517853f6
--- /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 8935130ad19d021cfadef3f665a15381a905a6a0..1c64717151cac926724bcd821a0ce4c586ddc2d2 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 88fb483fa17b58fda5900e37ce47f8a6a2851911..016d4ef5a8d745eda3beb5c64520be9a19b4d6af 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 0000000000000000000000000000000000000000..9550b0d104e93da95975258c39395293915111ad
--- /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 4016b0f53bee506eb36300e58d7ad3e49bd175c0..1fa2720cb9c1f2e19e1273351d9cdfcf0774becc 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 c6e08c962d413ff9357d5d6d16704d8dc849008c..7bc8a44eea903102e19b845983ff91d291369e0e 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 544998674bf369466b509a0af3cd1098c5c121a2..fe54caf4d279f407779771355c29a8762c80060c 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 0000000000000000000000000000000000000000..03d550a7f33cd7998bb522aa994d7dac37159e02
--- /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 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/python/socketd/socketd/transport/CodecByteBuffer.py b/python/socketd/socketd/transport/CodecByteBuffer.py
index 25a0b56d60a0fb982dbfa3af27bbb35a1dd8094c..52c14f84a35bd29ab7fc035ce91e11749360f72d 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/transport/client/ClientChannel.py b/python/socketd/socketd/transport/client/ClientChannel.py
index 042cb197a38071de27fb05c5f6297d45c066a2b9..d1d4600a71fd14663d8f2c1f768da467aa485e23 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/AsyncUtil.py b/python/socketd/socketd/transport/core/AsyncUtil.py
new file mode 100644
index 0000000000000000000000000000000000000000..bb95e5fac5168b18839b1aa7181c431e45af94d9
--- /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/CompletableFuture.py b/python/socketd/socketd/transport/core/CompletableFuture.py
index dad3e595f9ea02a6c9cac447b9bb044b2a0c19f0..322ac4b5dca04dc1a9f2d5d7f5520bfe88f43eb0 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/HeartbeatHandlerDefault.py b/python/socketd/socketd/transport/core/HeartbeatHandlerDefault.py
new file mode 100644
index 0000000000000000000000000000000000000000..eee7a824d2da7c436e1bb5d1dd9c2014ea41ee1c
--- /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/transport/core/StreamBase.py b/python/socketd/socketd/transport/core/StreamBase.py
index 6ffcc89618cee5b6b5f4688657c42be259fee692..ac2501e7469cae4e80f94deaf214e30dea4f9fe6 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 27c1b0499210b90852d10548044805b896529596..6d29bf7a6c17a15abbeea4e3b50c637b33d9ca76 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/WsAioChannelAssistant.py b/python/socketd/socketd_websocket/WsAioChannelAssistant.py
index 1ebda398381d497129c94344607643a1c51ce703..15f2e4e84ab4bce6f512c8fe5724e34e36d48452 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/socketd_websocket/WsAioFactoy.py b/python/socketd/socketd_websocket/WsAioFactoy.py
index 842ae5ea72b2260a83cd2dc843de54adb5e1f3bf..6034556582ec88a0930e0b7565dd395835a9189f 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/WsAioServer.py b/python/socketd/socketd_websocket/WsAioServer.py
index a04ec8ca9936441fc7e9a80583fe3122afa29306..b03ecc90ade09ee2bcff6d40da90e8f5371c50a6 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/AIOWebSocketClientImpl.py b/python/socketd/socketd_websocket/impl/AIOWebSocketClientImpl.py
index 2763380c540c6e1d7f500974c03d7f66d0225d4f..3fc25b3646eaf412bf6313f5ad51a931a220bad5 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:
@@ -52,18 +59,16 @@ class AIOWebSocketClientImpl(WebSocketClientProtocol):
try:
await self.on_message()
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):
@@ -72,12 +77,14 @@ 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):
"""处理消息"""
try:
+ if self.status_state == Flag.Close:
+ return
message = await self.recv()
log.debug(message)
frame: Frame = self.client.get_assistant().read(message)
@@ -93,7 +100,6 @@ class AIOWebSocketClientImpl(WebSocketClientProtocol):
# 超时自动推出
log.debug(c)
except Exception as e:
- log.warning(str(e), exc_info=True)
raise e
def on_close(self):
diff --git a/python/socketd/socketd_websocket/impl/AIOWebSocketServerImpl.py b/python/socketd/socketd_websocket/impl/AIOWebSocketServerImpl.py
index 1e3dd1dff4f83b4d77b2dbef83818f980396753c..8d1923ced0434b0704a3baa040abdd53215f2fe0 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 0000000000000000000000000000000000000000..855f17dc80b24935382413c1da4b113bf2709497
--- /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/00_TestCase.py b/python/socketd/test/TestCase00.py
similarity index 55%
rename from python/socketd/test/00_TestCase.py
rename to python/socketd/test/TestCase00.py
index 812b1789b168340700a53c6b396973fcd8ab063b..403d8bbd2f9430d9211abe1bee3d7695133c64e4 100644
--- a/python/socketd/test/00_TestCase.py
+++ b/python/socketd/test/TestCase00.py
@@ -6,9 +6,9 @@ from loguru import logger
from test.modelu.BaseTest import BaseTest
-class TestCase(unittest.TestCase):
+class TestCase00(unittest.TestCase):
- count = 10000
+ count = 100000
timeout = 30
def __init__(self, *args, **kwargs):
@@ -18,15 +18,15 @@ class TestCase(unittest.TestCase):
def test_send(self):
test = BaseTest()
- loop = asyncio.get_event_loop()
+ loop = asyncio.new_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))
+ 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:
- test.close()
+ 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 0000000000000000000000000000000000000000..93b8a79ba857eaf86794a05b5707cf35f893636b
--- /dev/null
+++ b/python/socketd/test/TestCase01.py
@@ -0,0 +1,70 @@
+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.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
+
+
+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
+
+ 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:
+ t.start()
+ t.stop()
+ except Exception as e:
+ t.on_error()
+ raise e
+
+ def test_Case04sendAndRequest_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
+
+ 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/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 c7ce4354421a0d799ea068c77dde5db85361e6e5..dbdd1b440b8f47a2adbbe1abeeb21623c637b7ed 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 b496a2fbfc402328f502999c6f925e1aea203937..fe79d115e941ea7607bc50f95f5420a4dbbd39b9 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
diff --git a/python/socketd/test/base_test/__init__.py b/python/socketd/test/base_test/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
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 0000000000000000000000000000000000000000..58eb1cc0ddbc1028a291683eb9c22f694cc9540c
--- /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 0000000000000000000000000000000000000000..43518734bb50a77439e19decc6991576f92b3a65
--- /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/TestCase03_client_session_close.py b/python/socketd/test/cases/TestCase03_client_session_close.py
new file mode 100644
index 0000000000000000000000000000000000000000..3e5adec4ac0074ce977f11bebd9866b26f45020a
--- /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 0000000000000000000000000000000000000000..ef20312eee2a97690fbb4bdfe1f9cfd8da98c840
--- /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
+
+
+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/cases/TestCase05_file.py b/python/socketd/test/cases/TestCase05_file.py
new file mode 100644
index 0000000000000000000000000000000000000000..883232ed5e84a42e38413dddc6e99ec296ee1f58
--- /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()
diff --git a/python/socketd/test/cases/__init__.py b/python/socketd/test/cases/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/python/socketd/test/modelu/BaseTestCase.py b/python/socketd/test/modelu/BaseTestCase.py
new file mode 100644
index 0000000000000000000000000000000000000000..8a4d3ee1b2e33c1b9e024aef7b24161f343f63a3
--- /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 336110041f4d301b07a5ecbd928bf69051d889f6..6fd74ba010b92b6c6ca323712bd51a8a092de8ed 100644
--- a/python/socketd/test/modelu/SimpleListenerTest.py
+++ b/python/socketd/test/modelu/SimpleListenerTest.py
@@ -2,27 +2,54 @@ 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)
+ self.close_counter = AtomicRefer(0)
+ self.message_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():
+ 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"))
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
+ logger.debug("客户端主动关闭了")
+ with self.close_counter:
+ self.close_counter.set(self.close_counter.get() + 1)
def on_error(self, session, error):
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)