feat: 实现 Python 与 Java SDK 连接收发及其余接口
This commit is contained in:
+43
-1
@@ -465,7 +465,49 @@
|
||||
|
||||
## SDK 二 S2
|
||||
|
||||
暂无。
|
||||
### S2-PY/JAVA 1–3 2026-09-30
|
||||
|
||||
1. **HiveMQ「websocket 模块」用 netty-codec-http 显式依赖**
|
||||
- 原条款:DEVELOPMENT 2.3「HiveMQ MQTT Client,加上 websocket 模块」。
|
||||
- 实际做法:依赖 `com.hivemq:hivemq-mqtt-client:1.3.5`,并额外声明 `io.netty:netty-codec-http`;连接时用 `webSocketConfig().subprotocol("mqtt")`。未使用独立 artifact `hivemq-mqtt-client-websocket`(Maven Central 上该坐标未作为独立稳定模块发布)。
|
||||
- 原因:与 HiveMQ 官方 WebSocket 用法一致,满足子协议 `mqtt`。
|
||||
- 备选方案:若日后官方拆出独立 websocket 模块再改坐标。
|
||||
- 影响:无行为差异。
|
||||
|
||||
2. **JSON 库选型**
|
||||
- 原条款:未指定 Java JSON 库。
|
||||
- 实际做法:Java 用 Gson(`disableHtmlEscaping`);Python 用标准库 `json`(`ensure_ascii=False`)。
|
||||
- 原因:满足「不转义 HTML / 非 ASCII」;不引入过重依赖。
|
||||
- 备选方案:Jackson。
|
||||
- 影响:无。
|
||||
|
||||
3. **假传输单测,未接真实服务器**
|
||||
- 原条款:任务 1–3 单元测试用假传输;任务 4 才做接入清单。
|
||||
- 实际做法:Python `FakeTransport`、Java `FakeTransport` 覆盖 Clean Start、去重再 ack、本地超限、令牌回调、重交不改 `send_at_ms` / 消息号;未做对真实服务器的接入清单(任务 4)。
|
||||
- 原因:本波范围。
|
||||
- 备选方案:无。
|
||||
- 影响:真实联调留待 S2 任务 4。
|
||||
|
||||
4. **Python 发布元数据**
|
||||
- 原条款:`license = { file = "LICENSE" }` 与专有分类。
|
||||
- 实际做法:`pyproject.toml` 已按此写;包内复制仓库根 `LICENSE`。未配置/执行 PyPI 发布。
|
||||
- 原因:发布在阶段 3。
|
||||
- 备选方案:无。
|
||||
- 影响:无。
|
||||
|
||||
5. **Java 编译器用 JDK 21,目标字节码 8**
|
||||
- 原条款:字节码目标 Java 8。
|
||||
- 实际做法:`maven.compiler.release=8`,本机用 Temurin 21 编译。
|
||||
- 原因:环境已有 JDK 21。
|
||||
- 备选方案:用 JDK 8 工具链。
|
||||
- 影响:无。
|
||||
|
||||
6. **Paho / HiveMQ 库内自动重连关闭,退避自管**
|
||||
- 原条款:四种 SDK 同一套重连:1s 起加倍上限 30s ±30% 抖动,稳定 60s 恢复;每次 Clean Start、会话过期 0。
|
||||
- 实际做法:两端均由 SDK 连接循环实现退避与停止条件;HiveMQ 不启库内 automaticReconnect;Paho 每次 `connect(..., clean_start=True)` 并设 `SessionExpiryInterval=0`。
|
||||
- 原因:与 Go/JS 要求一致,避免两套重连。
|
||||
- 备选方案:依赖库自带重连再改 Clean Start(易漏)。
|
||||
- 影响:无。
|
||||
|
||||
## 测试交付 Q
|
||||
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
target/
|
||||
.idea/
|
||||
*.iml
|
||||
.classpath
|
||||
.project
|
||||
.settings/
|
||||
@@ -0,0 +1,10 @@
|
||||
Copyright (c) 2026 Nixevol. All rights reserved.
|
||||
|
||||
本仓库的源代码、文档、各语言 SDK 和构建产物(包括发布的软件包和 Docker 镜像)均为专有软件。
|
||||
源代码和发布物公开可读,不代表授予任何使用许可。未经版权所有者书面许可,不得使用、复制、
|
||||
修改、合并、发布、分发、再许可或出售其任何部分。
|
||||
|
||||
This repository, including its source code, documentation, SDKs and build artifacts (including
|
||||
published packages and Docker images), is proprietary software. Public visibility does not grant
|
||||
any license. No part of it may be used, copied, modified, merged, published, distributed,
|
||||
sublicensed or sold without prior written permission from the copyright holder.
|
||||
@@ -0,0 +1,34 @@
|
||||
# NixMsg Java / Android SDK
|
||||
|
||||
坐标:`asia.asio.nixmsg:nixmsg-sdk`
|
||||
包名:`asia.asio.nixmsg`
|
||||
字节码目标:Java 8
|
||||
接口:`CompletableFuture`
|
||||
|
||||
## Android
|
||||
|
||||
最低 API **24**。MQTT 长连接由应用自行放入**前台服务**,SDK 不创建也不托管服务生命周期。
|
||||
|
||||
## 依赖
|
||||
|
||||
Maven / Gradle 仓库:
|
||||
|
||||
```
|
||||
https://git.asio.asia/api/packages/nixevol/maven
|
||||
```
|
||||
|
||||
HiveMQ MQTT Client(含 WebSocket:`webSocketConfig` + `netty-codec-http`)。
|
||||
|
||||
## 最小示例
|
||||
|
||||
```java
|
||||
Client c = new Client();
|
||||
c.onSession(token -> { /* 应用保存 */ });
|
||||
c.onMessage(msg -> System.out.println(msg.id + " " + msg.body.data));
|
||||
c.connect("ws://127.0.0.1:7443/mqtt", "device-1", "secret", null)
|
||||
.thenCompose(v -> c.send(new Types.Target("endpoint", "device-2"), new Types.Body("hello"), new Types.SendOptions()))
|
||||
.join();
|
||||
c.close();
|
||||
```
|
||||
|
||||
许可证见 `LICENSE`(专有)。
|
||||
@@ -0,0 +1,90 @@
|
||||
<?xml version="1.0" encoding="UTF-8"?>
|
||||
<project xmlns="http://maven.apache.org/POM/4.0.0"
|
||||
xmlns:xsi="http://www.w3.org/2001/XMLSchema-instance"
|
||||
xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 http://maven.apache.org/xsd/maven-4.0.0.xsd">
|
||||
<modelVersion>4.0.0</modelVersion>
|
||||
|
||||
<groupId>asia.asio.nixmsg</groupId>
|
||||
<artifactId>nixmsg-sdk</artifactId>
|
||||
<version>0.1.0</version>
|
||||
<packaging>jar</packaging>
|
||||
<name>nixmsg-sdk</name>
|
||||
<description>NixMsg Java/Android SDK</description>
|
||||
|
||||
<licenses>
|
||||
<license>
|
||||
<name>Proprietary</name>
|
||||
<url>https://git.asio.asia/nixevol/NixMsg/src/branch/main/LICENSE</url>
|
||||
<distribution>repo</distribution>
|
||||
</license>
|
||||
</licenses>
|
||||
|
||||
<properties>
|
||||
<project.build.sourceEncoding>UTF-8</project.build.sourceEncoding>
|
||||
<maven.compiler.release>8</maven.compiler.release>
|
||||
<hivemq.mqtt.version>1.3.5</hivemq.mqtt.version>
|
||||
<junit.version>4.13.2</junit.version>
|
||||
</properties>
|
||||
|
||||
<dependencies>
|
||||
<dependency>
|
||||
<groupId>com.hivemq</groupId>
|
||||
<artifactId>hivemq-mqtt-client</artifactId>
|
||||
<version>${hivemq.mqtt.version}</version>
|
||||
</dependency>
|
||||
<!-- WebSocket 传输所需 Netty HTTP 编解码(HiveMQ 文档中的 websocket 能力) -->
|
||||
<dependency>
|
||||
<groupId>io.netty</groupId>
|
||||
<artifactId>netty-codec-http</artifactId>
|
||||
<version>4.1.118.Final</version>
|
||||
</dependency>
|
||||
<dependency>
|
||||
<groupId>com.google.code.gson</groupId>
|
||||
<artifactId>gson</artifactId>
|
||||
<version>2.11.0</version>
|
||||
</dependency>
|
||||
<dependency>
|
||||
<groupId>junit</groupId>
|
||||
<artifactId>junit</artifactId>
|
||||
<version>${junit.version}</version>
|
||||
<scope>test</scope>
|
||||
</dependency>
|
||||
</dependencies>
|
||||
|
||||
<build>
|
||||
<plugins>
|
||||
<plugin>
|
||||
<groupId>org.apache.maven.plugins</groupId>
|
||||
<artifactId>maven-compiler-plugin</artifactId>
|
||||
<version>3.13.0</version>
|
||||
<configuration>
|
||||
<release>8</release>
|
||||
</configuration>
|
||||
</plugin>
|
||||
<plugin>
|
||||
<groupId>org.apache.maven.plugins</groupId>
|
||||
<artifactId>maven-surefire-plugin</artifactId>
|
||||
<version>3.5.2</version>
|
||||
</plugin>
|
||||
<plugin>
|
||||
<groupId>org.apache.maven.plugins</groupId>
|
||||
<artifactId>maven-jar-plugin</artifactId>
|
||||
<version>3.4.2</version>
|
||||
</plugin>
|
||||
</plugins>
|
||||
</build>
|
||||
|
||||
<repositories>
|
||||
<repository>
|
||||
<id>central</id>
|
||||
<url>https://repo.maven.apache.org/maven2</url>
|
||||
</repository>
|
||||
</repositories>
|
||||
|
||||
<distributionManagement>
|
||||
<repository>
|
||||
<id>asio-gitea</id>
|
||||
<url>https://git.asio.asia/api/packages/nixevol/maven</url>
|
||||
</repository>
|
||||
</distributionManagement>
|
||||
</project>
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,15 @@
|
||||
package asia.asio.nixmsg;
|
||||
|
||||
/** SDK 错误,code 与协议错误码或连接原因对齐。 */
|
||||
public class NixMsgException extends RuntimeException {
|
||||
private final String code;
|
||||
|
||||
public NixMsgException(String code, String message) {
|
||||
super(message == null || message.isEmpty() ? code : message);
|
||||
this.code = code;
|
||||
}
|
||||
|
||||
public String getCode() {
|
||||
return code;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,87 @@
|
||||
package asia.asio.nixmsg;
|
||||
|
||||
import com.google.gson.Gson;
|
||||
import com.google.gson.GsonBuilder;
|
||||
import com.google.gson.ToNumberPolicy;
|
||||
|
||||
import java.net.URI;
|
||||
import java.nio.charset.StandardCharsets;
|
||||
import java.util.Map;
|
||||
|
||||
final class Protocol {
|
||||
static final Gson GSON = new GsonBuilder()
|
||||
.disableHtmlEscaping()
|
||||
.setObjectToNumberStrategy(ToNumberPolicy.LONG_OR_DOUBLE)
|
||||
.create();
|
||||
|
||||
private Protocol() {}
|
||||
|
||||
static byte[] dumps(Object obj) {
|
||||
return GSON.toJson(obj).getBytes(StandardCharsets.UTF_8);
|
||||
}
|
||||
|
||||
@SuppressWarnings("unchecked")
|
||||
static Map<String, Object> loads(byte[] data) {
|
||||
return GSON.fromJson(new String(data, StandardCharsets.UTF_8), Map.class);
|
||||
}
|
||||
|
||||
static String upTopic(String endpointId) {
|
||||
return "nix/c/" + endpointId + "/up";
|
||||
}
|
||||
|
||||
static String downTopic(String endpointId) {
|
||||
return "nix/c/" + endpointId + "/down";
|
||||
}
|
||||
|
||||
static String normalizeMqttWsUrl(String url) {
|
||||
String raw = url.trim();
|
||||
if (!raw.contains("://")) {
|
||||
raw = "ws://" + raw;
|
||||
}
|
||||
URI u = URI.create(raw);
|
||||
String path = u.getPath() == null ? "" : u.getPath();
|
||||
if (!path.endsWith("/mqtt")) {
|
||||
if (path.endsWith("/")) {
|
||||
path = path + "mqtt";
|
||||
} else if (path.isEmpty()) {
|
||||
path = "/mqtt";
|
||||
} else {
|
||||
path = path + "/mqtt";
|
||||
}
|
||||
}
|
||||
try {
|
||||
return new URI(u.getScheme(), u.getUserInfo(), u.getHost(), u.getPort(), path, u.getQuery(), u.getFragment()).toString();
|
||||
} catch (Exception e) {
|
||||
return raw;
|
||||
}
|
||||
}
|
||||
|
||||
/** 从连接地址推出注册 HTTP 地址。 */
|
||||
static String registerUrlFromConnect(String connectUrl) {
|
||||
String raw = connectUrl.trim();
|
||||
if (!raw.contains("://")) {
|
||||
raw = "ws://" + raw;
|
||||
}
|
||||
URI u = URI.create(raw);
|
||||
String scheme = u.getScheme() == null ? "ws" : u.getScheme().toLowerCase();
|
||||
String httpScheme;
|
||||
if ("wss".equals(scheme) || "mqtts".equals(scheme) || "https".equals(scheme)) {
|
||||
httpScheme = "https";
|
||||
} else {
|
||||
httpScheme = "http";
|
||||
}
|
||||
String path = u.getPath() == null ? "" : u.getPath();
|
||||
if (path.endsWith("/mqtt")) {
|
||||
path = path.substring(0, path.length() - "/mqtt".length());
|
||||
}
|
||||
if (path.endsWith("/")) {
|
||||
path = path.substring(0, path.length() - 1);
|
||||
}
|
||||
path = path + "/api/client/register";
|
||||
try {
|
||||
return new URI(httpScheme, u.getUserInfo(), u.getHost(), u.getPort(), path, null, null).toString();
|
||||
} catch (Exception e) {
|
||||
throw new NixMsgException("bad_request", "无效连接地址");
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,406 @@
|
||||
package asia.asio.nixmsg;
|
||||
|
||||
import com.hivemq.client.mqtt.MqttClient;
|
||||
import com.hivemq.client.mqtt.MqttGlobalPublishFilter;
|
||||
import com.hivemq.client.mqtt.datatypes.MqttQos;
|
||||
import com.hivemq.client.mqtt.lifecycle.MqttDisconnectSource;
|
||||
import com.hivemq.client.mqtt.mqtt5.Mqtt5AsyncClient;
|
||||
import com.hivemq.client.mqtt.mqtt5.Mqtt5ClientBuilder;
|
||||
import com.hivemq.client.mqtt.mqtt5.message.connect.connack.Mqtt5ConnAck;
|
||||
import com.hivemq.client.mqtt.mqtt5.message.connect.connack.Mqtt5ConnAckReasonCode;
|
||||
import com.hivemq.client.mqtt.mqtt5.message.disconnect.Mqtt5Disconnect;
|
||||
import com.hivemq.client.mqtt.mqtt5.message.disconnect.Mqtt5DisconnectReasonCode;
|
||||
|
||||
import java.net.URI;
|
||||
import java.nio.charset.StandardCharsets;
|
||||
import java.util.ArrayList;
|
||||
import java.util.LinkedHashMap;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.concurrent.CopyOnWriteArrayList;
|
||||
import java.util.concurrent.TimeUnit;
|
||||
import java.util.function.BiConsumer;
|
||||
import java.util.function.Consumer;
|
||||
import java.util.function.Function;
|
||||
|
||||
/** MQTT 传输抽象。 */
|
||||
public interface Transport {
|
||||
void setHandlers(Runnable onConnected, BiConsumer<String, Boolean> onDisconnected, Consumer<byte[]> onDown);
|
||||
|
||||
void connect(ConnectParams params);
|
||||
|
||||
void subscribe(String topic);
|
||||
|
||||
void publish(String topic, byte[] payload);
|
||||
|
||||
void disconnect();
|
||||
|
||||
final class ConnectParams {
|
||||
public final String url;
|
||||
public final String clientId;
|
||||
public final String username;
|
||||
public final String password;
|
||||
public final boolean cleanStart;
|
||||
public final int sessionExpiry;
|
||||
public final boolean useTcp;
|
||||
public final long timeoutMs;
|
||||
|
||||
public ConnectParams(String url, String clientId, String username, String password,
|
||||
boolean cleanStart, int sessionExpiry, boolean useTcp, long timeoutMs) {
|
||||
this.url = url;
|
||||
this.clientId = clientId;
|
||||
this.username = username;
|
||||
this.password = password;
|
||||
this.cleanStart = cleanStart;
|
||||
this.sessionExpiry = sessionExpiry;
|
||||
this.useTcp = useTcp;
|
||||
this.timeoutMs = timeoutMs;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/** 单元测试假传输。 */
|
||||
final class FakeTransport implements Transport {
|
||||
final List<Transport.ConnectParams> connects = new ArrayList<Transport.ConnectParams>();
|
||||
final List<byte[]> publishes = new CopyOnWriteArrayList<byte[]>();
|
||||
final List<String> subscriptions = new ArrayList<String>();
|
||||
volatile String nextConnackFail;
|
||||
Map<String, Object> autoHello;
|
||||
boolean autoSendOk = true;
|
||||
final List<Function<Map<String, Object>, Map<String, Object>>> upHandlers =
|
||||
new CopyOnWriteArrayList<Function<Map<String, Object>, Map<String, Object>>>();
|
||||
|
||||
private Runnable onConnected;
|
||||
private BiConsumer<String, Boolean> onDisconnected;
|
||||
private Consumer<byte[]> onDown;
|
||||
private volatile boolean connected;
|
||||
|
||||
FakeTransport() {
|
||||
autoHello = new LinkedHashMap<String, Object>();
|
||||
autoHello.put("server_time_ms", 1750000000000L);
|
||||
autoHello.put("server_version", "0.1.0");
|
||||
autoHello.put("max_body_bytes", 262144);
|
||||
autoHello.put("max_meta_bytes", 4096);
|
||||
autoHello.put("max_frame_bytes", 786432);
|
||||
autoHello.put("max_ttl_seconds", 2592000);
|
||||
autoHello.put("max_schedule_seconds", 31536000);
|
||||
autoHello.put("ack_timeout_seconds", 300);
|
||||
autoHello.put("session_token", "nst_test_token");
|
||||
}
|
||||
|
||||
void onUp(Function<Map<String, Object>, Map<String, Object>> h) {
|
||||
upHandlers.add(h);
|
||||
}
|
||||
|
||||
@Override
|
||||
public void setHandlers(Runnable onConnected, BiConsumer<String, Boolean> onDisconnected, Consumer<byte[]> onDown) {
|
||||
this.onConnected = onConnected;
|
||||
this.onDisconnected = onDisconnected;
|
||||
this.onDown = onDown;
|
||||
}
|
||||
|
||||
@Override
|
||||
public void connect(ConnectParams params) {
|
||||
connects.add(params);
|
||||
if (nextConnackFail != null) {
|
||||
String fail = nextConnackFail;
|
||||
nextConnackFail = null;
|
||||
boolean stop = "session_invalid".equals(fail) || "bad_credentials".equals(fail) || "banned".equals(fail);
|
||||
if (onDisconnected != null) {
|
||||
onDisconnected.accept(fail, stop);
|
||||
}
|
||||
return;
|
||||
}
|
||||
connected = true;
|
||||
if (onConnected != null) {
|
||||
onConnected.run();
|
||||
}
|
||||
}
|
||||
|
||||
@Override
|
||||
public void subscribe(String topic) {
|
||||
subscriptions.add(topic);
|
||||
}
|
||||
|
||||
@Override
|
||||
public void publish(String topic, byte[] payload) {
|
||||
publishes.add(payload);
|
||||
Map<String, Object> frame = Protocol.loads(payload);
|
||||
for (Function<Map<String, Object>, Map<String, Object>> h : upHandlers) {
|
||||
Map<String, Object> resp = h.apply(frame);
|
||||
if (resp != null) {
|
||||
injectDown(Protocol.dumps(resp));
|
||||
return;
|
||||
}
|
||||
}
|
||||
String type = str(frame.get("type"));
|
||||
Object rid = frame.get("rid");
|
||||
if ("hello".equals(type) && autoHello != null) {
|
||||
Map<String, Object> resp = new LinkedHashMap<String, Object>();
|
||||
resp.put("v", 1);
|
||||
resp.put("type", "resp");
|
||||
resp.put("rid", rid);
|
||||
resp.put("ok", true);
|
||||
resp.put("data", autoHello);
|
||||
injectDown(Protocol.dumps(resp));
|
||||
return;
|
||||
}
|
||||
if ("send".equals(type) && autoSendOk) {
|
||||
Map<String, Object> data = new LinkedHashMap<String, Object>();
|
||||
data.put("id", frame.get("id"));
|
||||
Object sat = frame.get("send_at_ms");
|
||||
data.put("send_at_ms", sat == null ? 0 : sat);
|
||||
data.put("state", "dispatched");
|
||||
Map<String, Object> resp = new LinkedHashMap<String, Object>();
|
||||
resp.put("v", 1);
|
||||
resp.put("type", "resp");
|
||||
resp.put("rid", rid);
|
||||
resp.put("ok", true);
|
||||
resp.put("data", data);
|
||||
injectDown(Protocol.dumps(resp));
|
||||
return;
|
||||
}
|
||||
if ("ack".equals(type)) {
|
||||
Map<String, Object> data = new LinkedHashMap<String, Object>();
|
||||
data.put("result", "accepted");
|
||||
Map<String, Object> resp = new LinkedHashMap<String, Object>();
|
||||
resp.put("v", 1);
|
||||
resp.put("type", "resp");
|
||||
resp.put("rid", rid);
|
||||
resp.put("ok", true);
|
||||
resp.put("data", data);
|
||||
injectDown(Protocol.dumps(resp));
|
||||
return;
|
||||
}
|
||||
if (type != null && !"hello".equals(type) && !"send".equals(type) && rid != null) {
|
||||
Map<String, Object> resp = new LinkedHashMap<String, Object>();
|
||||
resp.put("v", 1);
|
||||
resp.put("type", "resp");
|
||||
resp.put("rid", rid);
|
||||
resp.put("ok", true);
|
||||
resp.put("data", new LinkedHashMap<String, Object>());
|
||||
injectDown(Protocol.dumps(resp));
|
||||
}
|
||||
}
|
||||
|
||||
@Override
|
||||
public void disconnect() {
|
||||
boolean was = connected;
|
||||
connected = false;
|
||||
if (was && onDisconnected != null) {
|
||||
onDisconnected.accept(null, false);
|
||||
}
|
||||
}
|
||||
|
||||
void injectDown(byte[] payload) {
|
||||
if (onDown != null) {
|
||||
onDown.accept(payload);
|
||||
}
|
||||
}
|
||||
|
||||
void simulateTakenOver() {
|
||||
connected = false;
|
||||
if (onDisconnected != null) {
|
||||
onDisconnected.accept("taken_over", true);
|
||||
}
|
||||
}
|
||||
|
||||
void simulateNetworkDrop() {
|
||||
connected = false;
|
||||
if (onDisconnected != null) {
|
||||
onDisconnected.accept("network", false);
|
||||
}
|
||||
}
|
||||
|
||||
private static String str(Object o) {
|
||||
return o == null ? null : String.valueOf(o);
|
||||
}
|
||||
}
|
||||
|
||||
/** HiveMQ MQTT 5 + WebSocket(子协议 mqtt)。 */
|
||||
final class HiveMqTransport implements Transport {
|
||||
private Runnable onConnected;
|
||||
private BiConsumer<String, Boolean> onDisconnected;
|
||||
private Consumer<byte[]> onDown;
|
||||
private volatile Mqtt5AsyncClient client;
|
||||
|
||||
@Override
|
||||
public void setHandlers(Runnable onConnected, BiConsumer<String, Boolean> onDisconnected, Consumer<byte[]> onDown) {
|
||||
this.onConnected = onConnected;
|
||||
this.onDisconnected = onDisconnected;
|
||||
this.onDown = onDown;
|
||||
}
|
||||
|
||||
@Override
|
||||
public void connect(final ConnectParams params) {
|
||||
disconnectQuiet();
|
||||
URI u = URI.create(params.url.contains("://") ? params.url : "ws://" + params.url);
|
||||
boolean useTcp = params.useTcp
|
||||
|| "mqtt".equalsIgnoreCase(u.getScheme())
|
||||
|| "mqtts".equalsIgnoreCase(u.getScheme());
|
||||
String host = u.getHost() == null ? "localhost" : u.getHost();
|
||||
int port = u.getPort();
|
||||
if (port < 0) {
|
||||
if (useTcp) {
|
||||
port = "mqtts".equalsIgnoreCase(u.getScheme()) ? 8883 : 1883;
|
||||
} else {
|
||||
port = "wss".equalsIgnoreCase(u.getScheme()) ? 443 : 80;
|
||||
}
|
||||
}
|
||||
|
||||
Mqtt5ClientBuilder b5 = MqttClient.builder()
|
||||
.useMqttVersion5()
|
||||
.identifier(params.clientId)
|
||||
.serverHost(host)
|
||||
.serverPort(port)
|
||||
.addDisconnectedListener(context -> {
|
||||
String reason = "network";
|
||||
boolean stop = false;
|
||||
if (context.getSource() == MqttDisconnectSource.SERVER) {
|
||||
try {
|
||||
java.lang.reflect.Method m = context.getClass().getMethod("getMqttDisconnect");
|
||||
Object disc = m.invoke(context);
|
||||
if (disc instanceof Mqtt5Disconnect) {
|
||||
Mqtt5DisconnectReasonCode rc = ((Mqtt5Disconnect) disc).getReasonCode();
|
||||
if (rc == Mqtt5DisconnectReasonCode.SESSION_TAKEN_OVER) {
|
||||
reason = "taken_over";
|
||||
stop = true;
|
||||
}
|
||||
}
|
||||
} catch (Exception ignored) {
|
||||
}
|
||||
}
|
||||
Throwable cause = context.getCause();
|
||||
if (cause != null && cause.getMessage() != null
|
||||
&& cause.getMessage().toLowerCase().contains("taken over")) {
|
||||
reason = "taken_over";
|
||||
stop = true;
|
||||
}
|
||||
BiConsumer<String, Boolean> h = onDisconnected;
|
||||
if (h != null) {
|
||||
h.accept(reason, stop);
|
||||
}
|
||||
});
|
||||
|
||||
if (!useTcp) {
|
||||
String path = u.getPath() == null || u.getPath().isEmpty() ? "/mqtt" : u.getPath();
|
||||
if (!path.endsWith("/mqtt")) {
|
||||
path = path.endsWith("/") ? path + "mqtt" : path + "/mqtt";
|
||||
}
|
||||
String serverPath = path.startsWith("/") ? path.substring(1) : path;
|
||||
b5 = b5.webSocketConfig()
|
||||
.serverPath(serverPath)
|
||||
.subprotocol("mqtt")
|
||||
.applyWebSocketConfig();
|
||||
if ("wss".equalsIgnoreCase(u.getScheme())) {
|
||||
b5 = b5.sslWithDefaultConfig();
|
||||
}
|
||||
} else if ("mqtts".equalsIgnoreCase(u.getScheme())) {
|
||||
b5 = b5.sslWithDefaultConfig();
|
||||
}
|
||||
|
||||
Mqtt5AsyncClient c = b5.buildAsync();
|
||||
this.client = c;
|
||||
c.publishes(MqttGlobalPublishFilter.ALL, publish -> {
|
||||
Consumer<byte[]> h = onDown;
|
||||
if (h != null) {
|
||||
h.accept(publish.getPayloadAsBytes());
|
||||
}
|
||||
});
|
||||
|
||||
try {
|
||||
Mqtt5ConnAck ack = c.connectWith()
|
||||
.cleanStart(params.cleanStart)
|
||||
.sessionExpiryInterval(params.sessionExpiry)
|
||||
.simpleAuth()
|
||||
.username(params.username)
|
||||
.password(params.password.getBytes(StandardCharsets.UTF_8))
|
||||
.applySimpleAuth()
|
||||
.send()
|
||||
.get(params.timeoutMs, TimeUnit.MILLISECONDS);
|
||||
if (ack.getReasonCode() == Mqtt5ConnAckReasonCode.SUCCESS) {
|
||||
Runnable h = onConnected;
|
||||
if (h != null) {
|
||||
h.run();
|
||||
}
|
||||
} else {
|
||||
String reason = classify(ack.getReasonCode());
|
||||
boolean stop = isStop(ack.getReasonCode());
|
||||
BiConsumer<String, Boolean> h = onDisconnected;
|
||||
if (h != null) {
|
||||
h.accept(reason, stop);
|
||||
}
|
||||
}
|
||||
} catch (Exception e) {
|
||||
BiConsumer<String, Boolean> h = onDisconnected;
|
||||
if (h != null) {
|
||||
h.accept("network", false);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@Override
|
||||
public void subscribe(String topic) {
|
||||
Mqtt5AsyncClient c = client;
|
||||
if (c == null) {
|
||||
throw new NixMsgException("not_connected", "未连接");
|
||||
}
|
||||
try {
|
||||
c.subscribeWith().topicFilter(topic).qos(MqttQos.AT_LEAST_ONCE).send().get(10, TimeUnit.SECONDS);
|
||||
} catch (Exception e) {
|
||||
throw new NixMsgException("busy", "订阅失败");
|
||||
}
|
||||
}
|
||||
|
||||
@Override
|
||||
public void publish(String topic, byte[] payload) {
|
||||
Mqtt5AsyncClient c = client;
|
||||
if (c == null) {
|
||||
throw new NixMsgException("not_connected", "未连接");
|
||||
}
|
||||
try {
|
||||
c.publishWith()
|
||||
.topic(topic)
|
||||
.qos(MqttQos.AT_LEAST_ONCE)
|
||||
.payload(payload)
|
||||
.send()
|
||||
.get(10, TimeUnit.SECONDS);
|
||||
} catch (Exception e) {
|
||||
throw new NixMsgException("busy", "发布失败");
|
||||
}
|
||||
}
|
||||
|
||||
@Override
|
||||
public void disconnect() {
|
||||
disconnectQuiet();
|
||||
}
|
||||
|
||||
private void disconnectQuiet() {
|
||||
Mqtt5AsyncClient c = client;
|
||||
client = null;
|
||||
if (c != null) {
|
||||
try {
|
||||
c.disconnect().get(5, TimeUnit.SECONDS);
|
||||
} catch (Exception ignored) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private static boolean isStop(Mqtt5ConnAckReasonCode code) {
|
||||
return code == Mqtt5ConnAckReasonCode.BAD_USER_NAME_OR_PASSWORD
|
||||
|| code == Mqtt5ConnAckReasonCode.NOT_AUTHORIZED
|
||||
|| code == Mqtt5ConnAckReasonCode.BANNED;
|
||||
}
|
||||
|
||||
private static String classify(Mqtt5ConnAckReasonCode code) {
|
||||
if (code == Mqtt5ConnAckReasonCode.BAD_USER_NAME_OR_PASSWORD
|
||||
|| code == Mqtt5ConnAckReasonCode.NOT_AUTHORIZED
|
||||
|| code == Mqtt5ConnAckReasonCode.BANNED) {
|
||||
return "bad_credentials";
|
||||
}
|
||||
if (code == Mqtt5ConnAckReasonCode.SERVER_UNAVAILABLE || code == Mqtt5ConnAckReasonCode.SERVER_BUSY) {
|
||||
return "busy";
|
||||
}
|
||||
return "network";
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,240 @@
|
||||
package asia.asio.nixmsg;
|
||||
|
||||
import java.nio.charset.StandardCharsets;
|
||||
import java.util.Base64;
|
||||
import java.util.Collections;
|
||||
import java.util.LinkedHashMap;
|
||||
import java.util.Map;
|
||||
import java.util.Objects;
|
||||
|
||||
public final class Types {
|
||||
private Types() {}
|
||||
|
||||
public static final String CLIENT_NAME = "java-sdk/0.1";
|
||||
public static final int DEFAULT_MAX_BODY = 262144;
|
||||
public static final int DEFAULT_MAX_META = 4096;
|
||||
public static final int DEFAULT_MAX_FRAME = 786432;
|
||||
public static final int MIN_MAX_RECEIVE = 1024;
|
||||
public static final int SEND_QUEUE_LIMIT = 1000;
|
||||
public static final int INFLIGHT_LIMIT = 100;
|
||||
public static final int DEDUP_CAPACITY = 10000;
|
||||
public static final long CONNECT_TIMEOUT_MS = 30_000L;
|
||||
public static final long BACKOFF_INITIAL_MS = 1_000L;
|
||||
public static final long BACKOFF_MAX_MS = 30_000L;
|
||||
public static final double BACKOFF_JITTER = 0.3;
|
||||
public static final long STABLE_RESET_MS = 60_000L;
|
||||
|
||||
public enum ConnectionState {
|
||||
CONNECTING, ONLINE, RECONNECTING, OFFLINE, KICKED, AUTH_FAILED
|
||||
}
|
||||
|
||||
public static final class Target {
|
||||
public final String kind;
|
||||
public final String id;
|
||||
|
||||
public Target(String kind, String id) {
|
||||
this.kind = Objects.requireNonNull(kind);
|
||||
this.id = Objects.requireNonNull(id);
|
||||
}
|
||||
|
||||
public Map<String, Object> toMap() {
|
||||
Map<String, Object> m = new LinkedHashMap<String, Object>();
|
||||
m.put("kind", kind);
|
||||
m.put("id", id);
|
||||
return m;
|
||||
}
|
||||
}
|
||||
|
||||
public static final class Body {
|
||||
public final String enc;
|
||||
public final String data;
|
||||
public String contentType;
|
||||
|
||||
public Body(String data) {
|
||||
this("utf8", data, null);
|
||||
}
|
||||
|
||||
public Body(String enc, String data, String contentType) {
|
||||
this.enc = enc;
|
||||
this.data = data;
|
||||
this.contentType = contentType;
|
||||
}
|
||||
|
||||
public static Body ofBytes(byte[] raw) {
|
||||
return new Body("base64", Base64.getEncoder().encodeToString(raw), "application/octet-stream");
|
||||
}
|
||||
|
||||
public int decodedSize() {
|
||||
if ("base64".equals(enc)) {
|
||||
return Base64.getDecoder().decode(data).length;
|
||||
}
|
||||
return data.getBytes(StandardCharsets.UTF_8).length;
|
||||
}
|
||||
|
||||
public String effectiveContentType() {
|
||||
if (contentType != null && !contentType.isEmpty()) {
|
||||
return contentType;
|
||||
}
|
||||
return "base64".equals(enc) ? "application/octet-stream" : "text/plain; charset=utf-8";
|
||||
}
|
||||
|
||||
public Map<String, Object> toMap() {
|
||||
Map<String, Object> m = new LinkedHashMap<String, Object>();
|
||||
m.put("enc", enc);
|
||||
m.put("data", data);
|
||||
m.put("content_type", effectiveContentType());
|
||||
return m;
|
||||
}
|
||||
}
|
||||
|
||||
public static final class SendOptions {
|
||||
public Long sendAtMs;
|
||||
public Long delayMs;
|
||||
public boolean keep;
|
||||
public Long ttlSeconds;
|
||||
public boolean receipt = true;
|
||||
public String talkPassword = "";
|
||||
public String contentType;
|
||||
public Map<String, Object> meta;
|
||||
public String messageId;
|
||||
}
|
||||
|
||||
public static final class SendResult {
|
||||
public final String id;
|
||||
public final long sendAtMs;
|
||||
public final String state;
|
||||
|
||||
public SendResult(String id, long sendAtMs, String state) {
|
||||
this.id = id;
|
||||
this.sendAtMs = sendAtMs;
|
||||
this.state = state;
|
||||
}
|
||||
}
|
||||
|
||||
public static final class RecallResult {
|
||||
public final String result;
|
||||
public final int recalled;
|
||||
public final int accepted;
|
||||
public final int other;
|
||||
|
||||
public RecallResult(String result, int recalled, int accepted, int other) {
|
||||
this.result = result;
|
||||
this.recalled = recalled;
|
||||
this.accepted = accepted;
|
||||
this.other = other;
|
||||
}
|
||||
}
|
||||
|
||||
public static final class IncomingMessage {
|
||||
public final String id;
|
||||
public final String from;
|
||||
public final Target to;
|
||||
public final Body body;
|
||||
public final long sendAtMs;
|
||||
public final Map<String, Object> meta;
|
||||
|
||||
public IncomingMessage(String id, String from, Target to, Body body, long sendAtMs, Map<String, Object> meta) {
|
||||
this.id = id;
|
||||
this.from = from;
|
||||
this.to = to;
|
||||
this.body = body;
|
||||
this.sendAtMs = sendAtMs;
|
||||
this.meta = meta == null ? Collections.<String, Object>emptyMap() : meta;
|
||||
}
|
||||
}
|
||||
|
||||
public static final class Receipt {
|
||||
public final String receiptId;
|
||||
public final String id;
|
||||
public final String endpointId;
|
||||
public final String state;
|
||||
public final String reason;
|
||||
public final long atMs;
|
||||
|
||||
public Receipt(String receiptId, String id, String endpointId, String state, String reason, long atMs) {
|
||||
this.receiptId = receiptId;
|
||||
this.id = id;
|
||||
this.endpointId = endpointId;
|
||||
this.state = state;
|
||||
this.reason = reason;
|
||||
this.atMs = atMs;
|
||||
}
|
||||
}
|
||||
|
||||
public static final class RevokedEvent {
|
||||
public final String id;
|
||||
public final String from;
|
||||
public final String reason;
|
||||
|
||||
public RevokedEvent(String id, String from, String reason) {
|
||||
this.id = id;
|
||||
this.from = from;
|
||||
this.reason = reason;
|
||||
}
|
||||
}
|
||||
|
||||
public static final class PresenceEvent {
|
||||
public final String id;
|
||||
public final boolean online;
|
||||
public final long atMs;
|
||||
|
||||
public PresenceEvent(String id, boolean online, long atMs) {
|
||||
this.id = id;
|
||||
this.online = online;
|
||||
this.atMs = atMs;
|
||||
}
|
||||
}
|
||||
|
||||
public static final class GroupEvent {
|
||||
public final String groupId;
|
||||
public final String event;
|
||||
public final String endpointId;
|
||||
public final long atMs;
|
||||
|
||||
public GroupEvent(String groupId, String event, String endpointId, long atMs) {
|
||||
this.groupId = groupId;
|
||||
this.event = event;
|
||||
this.endpointId = endpointId;
|
||||
this.atMs = atMs;
|
||||
}
|
||||
}
|
||||
|
||||
public static final class HelloLimits {
|
||||
public long serverTimeMs;
|
||||
public String serverVersion = "";
|
||||
public int maxBodyBytes = DEFAULT_MAX_BODY;
|
||||
public int maxMetaBytes = DEFAULT_MAX_META;
|
||||
public int maxFrameBytes = DEFAULT_MAX_FRAME;
|
||||
public long maxTtlSeconds = 2592000;
|
||||
public long maxScheduleSeconds = 31536000;
|
||||
public long ackTimeoutSeconds = 300;
|
||||
public String sessionToken = "";
|
||||
}
|
||||
|
||||
public static final class RegisterOptions {
|
||||
public String id = "";
|
||||
public String loginPassword = "";
|
||||
public String name = "";
|
||||
public String talkPassword = "";
|
||||
}
|
||||
|
||||
public static final class RegisterResult {
|
||||
public final String id;
|
||||
public final String loginPassword;
|
||||
|
||||
public RegisterResult(String id, String loginPassword) {
|
||||
this.id = id;
|
||||
this.loginPassword = loginPassword;
|
||||
}
|
||||
}
|
||||
|
||||
public static final class ConnectionEvent {
|
||||
public final ConnectionState state;
|
||||
public final String reason;
|
||||
|
||||
public ConnectionEvent(ConnectionState state, String reason) {
|
||||
this.state = state;
|
||||
this.reason = reason == null ? "" : reason;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,19 @@
|
||||
package asia.asio.nixmsg;
|
||||
|
||||
import java.security.SecureRandom;
|
||||
import java.util.UUID;
|
||||
|
||||
final class Uuid7 {
|
||||
private static final SecureRandom RAND = new SecureRandom();
|
||||
|
||||
private Uuid7() {}
|
||||
|
||||
static String next() {
|
||||
long ts = System.currentTimeMillis() & ((1L << 48) - 1);
|
||||
long randA = RAND.nextInt() & 0x0FFF;
|
||||
long randB = RAND.nextLong() & ((1L << 62) - 1);
|
||||
long msb = (ts << 16) | (0x7L << 12) | randA;
|
||||
long lsb = (0b10L << 62) | randB;
|
||||
return new UUID(msb, lsb).toString();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,276 @@
|
||||
package asia.asio.nixmsg;
|
||||
|
||||
import asia.asio.nixmsg.Types.Body;
|
||||
import asia.asio.nixmsg.Types.ConnectionState;
|
||||
import asia.asio.nixmsg.Types.IncomingMessage;
|
||||
import asia.asio.nixmsg.Types.SendOptions;
|
||||
import asia.asio.nixmsg.Types.SendResult;
|
||||
import asia.asio.nixmsg.Types.Target;
|
||||
import org.junit.After;
|
||||
import org.junit.Test;
|
||||
|
||||
import java.nio.charset.StandardCharsets;
|
||||
import java.util.ArrayList;
|
||||
import java.util.LinkedHashMap;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.concurrent.CountDownLatch;
|
||||
import java.util.concurrent.TimeUnit;
|
||||
import java.util.concurrent.atomic.AtomicInteger;
|
||||
|
||||
import static org.junit.Assert.assertEquals;
|
||||
import static org.junit.Assert.assertTrue;
|
||||
import static org.junit.Assert.fail;
|
||||
|
||||
public class ClientTest {
|
||||
private Client client;
|
||||
private FakeTransport transport;
|
||||
|
||||
@After
|
||||
public void tearDown() {
|
||||
if (client != null) {
|
||||
client.close();
|
||||
}
|
||||
}
|
||||
|
||||
private void connectOnline() {
|
||||
transport = new FakeTransport();
|
||||
client = new Client(transport);
|
||||
client.connectSync("ws://example.test/mqtt", "ep1", "secret", null, false);
|
||||
assertEquals(ConnectionState.ONLINE, client.getState());
|
||||
}
|
||||
|
||||
@Test
|
||||
public void cleanStartAndSessionExpiryEveryConnect() throws Exception {
|
||||
connectOnline();
|
||||
assertTrue(transport.connects.size() >= 1);
|
||||
Transport.ConnectParams p = transport.connects.get(0);
|
||||
assertTrue(p.cleanStart);
|
||||
assertEquals(0, p.sessionExpiry);
|
||||
assertTrue(transport.subscriptions.contains("nix/c/ep1/down"));
|
||||
|
||||
transport.simulateNetworkDrop();
|
||||
long deadline = System.currentTimeMillis() + 3000;
|
||||
while (transport.connects.size() < 2 && System.currentTimeMillis() < deadline) {
|
||||
Thread.sleep(50);
|
||||
}
|
||||
assertTrue(transport.connects.size() >= 2);
|
||||
for (Transport.ConnectParams c : transport.connects) {
|
||||
assertTrue(c.cleanStart);
|
||||
assertEquals(0, c.sessionExpiry);
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void sessionTokenCallback() {
|
||||
transport = new FakeTransport();
|
||||
client = new Client(transport);
|
||||
final List<String> tokens = new ArrayList<String>();
|
||||
client.onSession(new java.util.function.Consumer<String>() {
|
||||
@Override
|
||||
public void accept(String t) {
|
||||
tokens.add(t);
|
||||
}
|
||||
});
|
||||
client.connectSync("ws://example.test/mqtt", "ep1", "pw", null, false);
|
||||
assertEquals(1, tokens.size());
|
||||
assertEquals("nst_test_token", tokens.get(0));
|
||||
assertEquals("nst_test_token", client.getSessionToken());
|
||||
}
|
||||
|
||||
@Test
|
||||
public void dedupAckedRearrivalAcksAgain() throws Exception {
|
||||
connectOnline();
|
||||
final List<String> delivered = new ArrayList<String>();
|
||||
client.onMessage(new java.util.function.Consumer<IncomingMessage>() {
|
||||
@Override
|
||||
public void accept(IncomingMessage m) {
|
||||
delivered.add(m.id);
|
||||
}
|
||||
});
|
||||
Map<String, Object> msg = new LinkedHashMap<String, Object>();
|
||||
msg.put("v", 1);
|
||||
msg.put("type", "msg");
|
||||
msg.put("id", "m1");
|
||||
msg.put("from", "peer");
|
||||
Map<String, Object> to = new LinkedHashMap<String, Object>();
|
||||
to.put("kind", "endpoint");
|
||||
to.put("id", "ep1");
|
||||
msg.put("to", to);
|
||||
Map<String, Object> body = new LinkedHashMap<String, Object>();
|
||||
body.put("enc", "utf8");
|
||||
body.put("data", "hi");
|
||||
msg.put("body", body);
|
||||
msg.put("send_at_ms", 1);
|
||||
transport.injectDown(Protocol.dumps(msg));
|
||||
Thread.sleep(100);
|
||||
assertEquals(1, delivered.size());
|
||||
int before = transport.publishes.size();
|
||||
transport.injectDown(Protocol.dumps(msg));
|
||||
Thread.sleep(100);
|
||||
assertEquals(1, delivered.size());
|
||||
int ackCount = 0;
|
||||
for (int i = before; i < transport.publishes.size(); i++) {
|
||||
Map<String, Object> f = Protocol.loads(transport.publishes.get(i));
|
||||
if ("ack".equals(String.valueOf(f.get("type")))) {
|
||||
ackCount++;
|
||||
}
|
||||
}
|
||||
assertTrue(ackCount >= 1);
|
||||
}
|
||||
|
||||
@Test
|
||||
public void localBodyTooLarge() {
|
||||
connectOnline();
|
||||
StringBuilder sb = new StringBuilder();
|
||||
for (int i = 0; i < client.getLimits().maxBodyBytes + 1; i++) {
|
||||
sb.append('x');
|
||||
}
|
||||
try {
|
||||
client.sendSync(new Target("endpoint", "ep2"), new Body(sb.toString()), new SendOptions());
|
||||
fail("expected body_too_large");
|
||||
} catch (NixMsgException e) {
|
||||
assertEquals("body_too_large", e.getCode());
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void localFrameTooLarge() {
|
||||
transport = new FakeTransport();
|
||||
transport.autoHello.put("max_frame_bytes", 200);
|
||||
client = new Client(transport);
|
||||
client.connectSync("ws://example.test/mqtt", "ep1", "pw", null, false);
|
||||
StringBuilder sb = new StringBuilder();
|
||||
for (int i = 0; i < 40; i++) {
|
||||
sb.append("hello world ");
|
||||
}
|
||||
try {
|
||||
client.sendSync(new Target("endpoint", "ep2"), new Body(sb.toString()), new SendOptions());
|
||||
fail("expected frame_too_large");
|
||||
} catch (NixMsgException e) {
|
||||
assertEquals("frame_too_large", e.getCode());
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void resendKeepsSendAtMs() {
|
||||
transport = new FakeTransport();
|
||||
transport.autoSendOk = false;
|
||||
final AtomicInteger rateHits = new AtomicInteger();
|
||||
transport.onUp(new java.util.function.Function<Map<String, Object>, Map<String, Object>>() {
|
||||
@Override
|
||||
public Map<String, Object> apply(Map<String, Object> frame) {
|
||||
if (!"send".equals(String.valueOf(frame.get("type")))) {
|
||||
return null;
|
||||
}
|
||||
if (rateHits.getAndIncrement() == 0) {
|
||||
Map<String, Object> err = new LinkedHashMap<String, Object>();
|
||||
err.put("code", "rate_limited");
|
||||
err.put("message", "slow");
|
||||
Map<String, Object> resp = new LinkedHashMap<String, Object>();
|
||||
resp.put("v", 1);
|
||||
resp.put("type", "resp");
|
||||
resp.put("rid", frame.get("rid"));
|
||||
resp.put("ok", false);
|
||||
resp.put("error", err);
|
||||
return resp;
|
||||
}
|
||||
Map<String, Object> data = new LinkedHashMap<String, Object>();
|
||||
data.put("id", frame.get("id"));
|
||||
data.put("send_at_ms", frame.get("send_at_ms"));
|
||||
data.put("state", "scheduled");
|
||||
Map<String, Object> resp = new LinkedHashMap<String, Object>();
|
||||
resp.put("v", 1);
|
||||
resp.put("type", "resp");
|
||||
resp.put("rid", frame.get("rid"));
|
||||
resp.put("ok", true);
|
||||
resp.put("data", data);
|
||||
return resp;
|
||||
}
|
||||
});
|
||||
client = new Client(transport);
|
||||
client.connectSync("ws://example.test/mqtt", "ep1", "pw", null, false);
|
||||
SendOptions opt = new SendOptions();
|
||||
opt.sendAtMs = 1700000000000L;
|
||||
opt.messageId = "fixed-id-1";
|
||||
SendResult result = client.sendSync(new Target("endpoint", "ep2"), new Body("hi"), opt);
|
||||
assertEquals("fixed-id-1", result.id);
|
||||
int sendCount = 0;
|
||||
for (byte[] raw : transport.publishes) {
|
||||
Map<String, Object> f = Protocol.loads(raw);
|
||||
if ("send".equals(String.valueOf(f.get("type")))) {
|
||||
sendCount++;
|
||||
assertEquals("fixed-id-1", String.valueOf(f.get("id")));
|
||||
assertEquals(1700000000000L, ((Number) f.get("send_at_ms")).longValue());
|
||||
}
|
||||
}
|
||||
assertTrue(sendCount >= 2);
|
||||
}
|
||||
|
||||
@Test
|
||||
public void sessionInvalidStopsReconnect() throws Exception {
|
||||
transport = new FakeTransport();
|
||||
client = new Client(transport);
|
||||
client.connectSync("ws://example.test/mqtt", "ep1", "pw", null, false);
|
||||
transport.nextConnackFail = "bad_credentials";
|
||||
transport.simulateNetworkDrop();
|
||||
long deadline = System.currentTimeMillis() + 3000;
|
||||
while (client.getState() != ConnectionState.AUTH_FAILED && System.currentTimeMillis() < deadline) {
|
||||
Thread.sleep(50);
|
||||
}
|
||||
assertEquals(ConnectionState.AUTH_FAILED, client.getState());
|
||||
int n = transport.connects.size();
|
||||
Thread.sleep(500);
|
||||
assertEquals(n, transport.connects.size());
|
||||
}
|
||||
|
||||
@Test
|
||||
public void deliveredDuplicateIgnored() throws Exception {
|
||||
transport = new FakeTransport();
|
||||
client = new Client(transport, false, Types.DEFAULT_MAX_FRAME, Types.CLIENT_NAME, Types.CONNECT_TIMEOUT_MS);
|
||||
client.connectSync("ws://example.test/mqtt", "ep1", "pw", null, false);
|
||||
final CountDownLatch started = new CountDownLatch(1);
|
||||
final CountDownLatch release = new CountDownLatch(1);
|
||||
final List<String> seen = new ArrayList<String>();
|
||||
client.onMessage(new java.util.function.Consumer<IncomingMessage>() {
|
||||
@Override
|
||||
public void accept(IncomingMessage m) {
|
||||
seen.add(m.id);
|
||||
started.countDown();
|
||||
try {
|
||||
release.await(2, TimeUnit.SECONDS);
|
||||
} catch (InterruptedException e) {
|
||||
Thread.currentThread().interrupt();
|
||||
}
|
||||
}
|
||||
});
|
||||
Map<String, Object> msg = new LinkedHashMap<String, Object>();
|
||||
msg.put("v", 1);
|
||||
msg.put("type", "msg");
|
||||
msg.put("id", "m2");
|
||||
msg.put("from", "peer");
|
||||
Map<String, Object> to = new LinkedHashMap<String, Object>();
|
||||
to.put("kind", "endpoint");
|
||||
to.put("id", "ep1");
|
||||
msg.put("to", to);
|
||||
Map<String, Object> body = new LinkedHashMap<String, Object>();
|
||||
body.put("enc", "utf8");
|
||||
body.put("data", "x");
|
||||
msg.put("body", body);
|
||||
msg.put("send_at_ms", 1);
|
||||
final byte[] raw = Protocol.dumps(msg);
|
||||
Thread t = new Thread(new Runnable() {
|
||||
@Override
|
||||
public void run() {
|
||||
transport.injectDown(raw);
|
||||
}
|
||||
});
|
||||
t.start();
|
||||
assertTrue(started.await(2, TimeUnit.SECONDS));
|
||||
transport.injectDown(raw);
|
||||
release.countDown();
|
||||
t.join(2000);
|
||||
Thread.sleep(50);
|
||||
assertEquals(1, seen.size());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,7 @@
|
||||
.venv/
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
*.egg-info/
|
||||
.pytest_cache/
|
||||
dist/
|
||||
build/
|
||||
@@ -0,0 +1,10 @@
|
||||
Copyright (c) 2026 Nixevol. All rights reserved.
|
||||
|
||||
本仓库的源代码、文档、各语言 SDK 和构建产物(包括发布的软件包和 Docker 镜像)均为专有软件。
|
||||
源代码和发布物公开可读,不代表授予任何使用许可。未经版权所有者书面许可,不得使用、复制、
|
||||
修改、合并、发布、分发、再许可或出售其任何部分。
|
||||
|
||||
This repository, including its source code, documentation, SDKs and build artifacts (including
|
||||
published packages and Docker images), is proprietary software. Public visibility does not grant
|
||||
any license. No part of it may be used, copied, modified, merged, published, distributed,
|
||||
sublicensed or sold without prior written permission from the copyright holder.
|
||||
@@ -0,0 +1,24 @@
|
||||
# NixMsg Python SDK
|
||||
|
||||
包名 `nixmsg`,最低 Python 3.10。同步接口为主,`AsyncClient` 提供 asyncio 包装。
|
||||
|
||||
## 安装
|
||||
|
||||
```bash
|
||||
pip install nixmsg --index-url https://git.asio.asia/api/packages/nixevol/pypi/simple/
|
||||
```
|
||||
|
||||
## 最小示例
|
||||
|
||||
```python
|
||||
from nixmsg import Client, Target, Body
|
||||
|
||||
c = Client()
|
||||
c.on_session(lambda token: print("session", token))
|
||||
c.on_message(lambda msg: print("msg", msg.id, msg.body.data))
|
||||
c.connect("ws://127.0.0.1:7443/mqtt", "device-1", password="secret")
|
||||
c.send(Target(kind="endpoint", id="device-2"), Body(data="hello"))
|
||||
c.close()
|
||||
```
|
||||
|
||||
许可证见 `LICENSE`(专有)。
|
||||
@@ -0,0 +1,33 @@
|
||||
[build-system]
|
||||
requires = ["setuptools>=68", "wheel"]
|
||||
build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "nixmsg"
|
||||
version = "0.1.0"
|
||||
description = "NixMsg Python SDK"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10"
|
||||
license = { file = "LICENSE" }
|
||||
authors = [{ name = "Nixevol" }]
|
||||
classifiers = [
|
||||
"License :: Other/Proprietary License",
|
||||
"Programming Language :: Python :: 3",
|
||||
"Programming Language :: Python :: 3.10",
|
||||
"Programming Language :: Python :: 3.11",
|
||||
"Programming Language :: Python :: 3.12",
|
||||
"Programming Language :: Python :: 3.13",
|
||||
]
|
||||
dependencies = [
|
||||
"paho-mqtt>=2.0,<3",
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
dev = ["pytest>=7"]
|
||||
|
||||
[tool.setuptools.packages.find]
|
||||
where = ["src"]
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
testpaths = ["tests"]
|
||||
pythonpath = ["src"]
|
||||
@@ -0,0 +1,50 @@
|
||||
"""NixMsg Python SDK。"""
|
||||
|
||||
from .async_client import AsyncClient
|
||||
from .client import Client
|
||||
from .errors import ClosedError, NixMsgError, NotConnectedError
|
||||
from .protocol import register_url_from_connect
|
||||
from .transport import FakeTransport, PahoTransport
|
||||
from .types import (
|
||||
Body,
|
||||
ConnectionEvent,
|
||||
ConnectionState,
|
||||
GroupEvent,
|
||||
IncomingMessage,
|
||||
PresenceEvent,
|
||||
Receipt,
|
||||
RegisterOptions,
|
||||
RegisterResult,
|
||||
RecallResult,
|
||||
RevokedEvent,
|
||||
SendOptions,
|
||||
SendResult,
|
||||
Target,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"AsyncClient",
|
||||
"Body",
|
||||
"Client",
|
||||
"ClosedError",
|
||||
"ConnectionEvent",
|
||||
"ConnectionState",
|
||||
"FakeTransport",
|
||||
"GroupEvent",
|
||||
"IncomingMessage",
|
||||
"NixMsgError",
|
||||
"NotConnectedError",
|
||||
"PahoTransport",
|
||||
"PresenceEvent",
|
||||
"Receipt",
|
||||
"RegisterOptions",
|
||||
"RegisterResult",
|
||||
"RecallResult",
|
||||
"RevokedEvent",
|
||||
"SendOptions",
|
||||
"SendResult",
|
||||
"Target",
|
||||
"register_url_from_connect",
|
||||
]
|
||||
|
||||
__version__ = "0.1.0"
|
||||
@@ -0,0 +1,124 @@
|
||||
"""asyncio 包装(同一包内)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from typing import Any, Optional
|
||||
|
||||
from .client import Client
|
||||
from .types import (
|
||||
Body,
|
||||
RegisterOptions,
|
||||
RegisterResult,
|
||||
SendOptions,
|
||||
SendResult,
|
||||
Target,
|
||||
)
|
||||
|
||||
|
||||
class AsyncClient:
|
||||
"""把同步 Client 的阻塞调用丢到线程池。"""
|
||||
|
||||
def __init__(self, client: Optional[Client] = None, **kwargs: Any) -> None:
|
||||
self._client = client or Client(**kwargs)
|
||||
|
||||
@property
|
||||
def sync(self) -> Client:
|
||||
return self._client
|
||||
|
||||
def on_session(self, handler): # noqa: ANN001
|
||||
self._client.on_session(handler)
|
||||
|
||||
def on_message(self, handler): # noqa: ANN001
|
||||
self._client.on_message(handler)
|
||||
|
||||
def on_receipt(self, handler): # noqa: ANN001
|
||||
self._client.on_receipt(handler)
|
||||
|
||||
def on_revoked(self, handler): # noqa: ANN001
|
||||
self._client.on_revoked(handler)
|
||||
|
||||
def on_presence(self, handler): # noqa: ANN001
|
||||
self._client.on_presence(handler)
|
||||
|
||||
def on_group_event(self, handler): # noqa: ANN001
|
||||
self._client.on_group_event(handler)
|
||||
|
||||
def on_connection(self, handler): # noqa: ANN001
|
||||
self._client.on_connection(handler)
|
||||
|
||||
async def connect(self, url: str, endpoint_id: str, **kwargs: Any) -> None:
|
||||
await asyncio.to_thread(self._client.connect, url, endpoint_id, **kwargs)
|
||||
|
||||
async def close(self) -> None:
|
||||
await asyncio.to_thread(self._client.close)
|
||||
|
||||
async def logout(self) -> None:
|
||||
await asyncio.to_thread(self._client.logout)
|
||||
|
||||
async def send(self, to: Target, body: Body | str | bytes, options: Optional[SendOptions] = None) -> SendResult:
|
||||
return await asyncio.to_thread(self._client.send, to, body, options)
|
||||
|
||||
async def ack(self, message) -> None: # noqa: ANN001
|
||||
await asyncio.to_thread(self._client.ack, message)
|
||||
|
||||
async def recall(self, message_id: str):
|
||||
return await asyncio.to_thread(self._client.recall, message_id)
|
||||
|
||||
async def status(self, message_id: str, cursor: str = "", limit: int = 100):
|
||||
return await asyncio.to_thread(self._client.status, message_id, cursor, limit)
|
||||
|
||||
async def unlock(self, endpoint_id: str, talk_password: str):
|
||||
return await asyncio.to_thread(self._client.unlock, endpoint_id, talk_password)
|
||||
|
||||
async def presence(self, ids: list[str]):
|
||||
return await asyncio.to_thread(self._client.presence, ids)
|
||||
|
||||
async def directory(self, cursor: str = "", query: str = "", limit: int = 100):
|
||||
return await asyncio.to_thread(self._client.directory, cursor, query, limit)
|
||||
|
||||
async def watch_presence(self, ids: Optional[list[str]] = None, *, all: bool = False):
|
||||
return await asyncio.to_thread(self._client.watch_presence, ids, all=all)
|
||||
|
||||
async def get_self(self):
|
||||
return await asyncio.to_thread(self._client.get_self)
|
||||
|
||||
async def update_self(self, **kwargs: Any):
|
||||
return await asyncio.to_thread(self._client.update_self, **kwargs)
|
||||
|
||||
async def set_talk_password(self, talk_password: str):
|
||||
return await asyncio.to_thread(self._client.set_talk_password, talk_password)
|
||||
|
||||
async def change_login_password(self, old_password: str, new_password: str):
|
||||
return await asyncio.to_thread(self._client.change_login_password, old_password, new_password)
|
||||
|
||||
async def group_create(self, name: str, members: list[dict[str, str]], group_id: str = ""):
|
||||
return await asyncio.to_thread(self._client.group_create, name, members, group_id)
|
||||
|
||||
async def group_add(self, group_id: str, members: list[dict[str, str]]):
|
||||
return await asyncio.to_thread(self._client.group_add, group_id, members)
|
||||
|
||||
async def group_remove(self, group_id: str, endpoint_id: str):
|
||||
return await asyncio.to_thread(self._client.group_remove, group_id, endpoint_id)
|
||||
|
||||
async def group_leave(self, group_id: str):
|
||||
return await asyncio.to_thread(self._client.group_leave, group_id)
|
||||
|
||||
async def group_transfer(self, group_id: str, endpoint_id: str):
|
||||
return await asyncio.to_thread(self._client.group_transfer, group_id, endpoint_id)
|
||||
|
||||
async def group_rename(self, group_id: str, name: str):
|
||||
return await asyncio.to_thread(self._client.group_rename, group_id, name)
|
||||
|
||||
async def group_dissolve(self, group_id: str):
|
||||
return await asyncio.to_thread(self._client.group_dissolve, group_id)
|
||||
|
||||
async def group_list(self, cursor: str = "", limit: int = 100):
|
||||
return await asyncio.to_thread(self._client.group_list, cursor, limit)
|
||||
|
||||
async def group_get(self, group_id: str, cursor: str = "", limit: int = 100):
|
||||
return await asyncio.to_thread(self._client.group_get, group_id, cursor, limit)
|
||||
|
||||
@staticmethod
|
||||
async def register(url: str, registration_code: str, options: Optional[RegisterOptions] = None) -> RegisterResult:
|
||||
return await asyncio.to_thread(Client.register, url, registration_code, options)
|
||||
@@ -0,0 +1,988 @@
|
||||
"""NixMsg 同步客户端。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import random
|
||||
import threading
|
||||
import time
|
||||
import urllib.error
|
||||
import urllib.request
|
||||
from collections import OrderedDict
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Optional
|
||||
|
||||
from .errors import ClosedError, NixMsgError, NotConnectedError
|
||||
from .protocol import dumps, down_topic, loads, normalize_mqtt_ws_url, register_url_from_connect, up_topic
|
||||
from .transport import ConnectParams, FakeTransport, PahoTransport, Transport
|
||||
from .types import (
|
||||
BACKOFF_INITIAL_S,
|
||||
BACKOFF_JITTER,
|
||||
BACKOFF_MAX_S,
|
||||
CLIENT_NAME,
|
||||
CONNECT_TIMEOUT_S,
|
||||
DEDUP_CAPACITY,
|
||||
DEFAULT_MAX_BODY,
|
||||
DEFAULT_MAX_FRAME,
|
||||
DEFAULT_MAX_META,
|
||||
INFLIGHT_LIMIT,
|
||||
MIN_MAX_RECEIVE,
|
||||
SEND_QUEUE_LIMIT,
|
||||
STABLE_RESET_S,
|
||||
Body,
|
||||
ConnectionEvent,
|
||||
ConnectionHandler,
|
||||
ConnectionState,
|
||||
GroupEvent,
|
||||
GroupEventHandler,
|
||||
HelloLimits,
|
||||
IncomingMessage,
|
||||
MessageHandler,
|
||||
PresenceEvent,
|
||||
PresenceHandler,
|
||||
Receipt,
|
||||
ReceiptHandler,
|
||||
RegisterOptions,
|
||||
RegisterResult,
|
||||
RecallResult,
|
||||
RevokedEvent,
|
||||
RevokedHandler,
|
||||
SendOptions,
|
||||
SendResult,
|
||||
SessionHandler,
|
||||
Target,
|
||||
)
|
||||
from .uuid7 import new_uuid7
|
||||
|
||||
log = logging.getLogger("nixmsg")
|
||||
|
||||
|
||||
@dataclass
|
||||
class _PendingReq:
|
||||
rid: str
|
||||
frame: dict[str, Any]
|
||||
event: threading.Event = field(default_factory=threading.Event)
|
||||
response: Optional[dict[str, Any]] = None
|
||||
error: Optional[BaseException] = None
|
||||
is_send: bool = False
|
||||
message_id: str = ""
|
||||
|
||||
|
||||
@dataclass
|
||||
class _SendItem:
|
||||
message_id: str
|
||||
frame: dict[str, Any] # 已含 send_at_ms,重交不改
|
||||
pending: _PendingReq
|
||||
|
||||
|
||||
class _DedupState:
|
||||
DELIVERED = "delivered" # 已交应用未确认
|
||||
ACKED = "acked"
|
||||
|
||||
|
||||
class Client:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
transport: Optional[Transport] = None,
|
||||
auto_ack: bool = True,
|
||||
max_receive_bytes: int = DEFAULT_MAX_FRAME,
|
||||
client_name: str = CLIENT_NAME,
|
||||
connect_timeout_s: float = CONNECT_TIMEOUT_S,
|
||||
) -> None:
|
||||
self._transport: Transport = transport or PahoTransport()
|
||||
self._auto_ack = auto_ack
|
||||
self._max_receive_bytes = max(MIN_MAX_RECEIVE, max_receive_bytes)
|
||||
self._client_name = client_name
|
||||
self._connect_timeout_s = connect_timeout_s
|
||||
|
||||
self._session_handler: Optional[SessionHandler] = None
|
||||
self._message_handler: Optional[MessageHandler] = None
|
||||
self._receipt_handler: Optional[ReceiptHandler] = None
|
||||
self._revoked_handler: Optional[RevokedHandler] = None
|
||||
self._presence_handler: Optional[PresenceHandler] = None
|
||||
self._group_handler: Optional[GroupEventHandler] = None
|
||||
self._connection_handler: Optional[ConnectionHandler] = None
|
||||
|
||||
self._lock = threading.RLock()
|
||||
self._cb_lock = threading.Lock() # 回调串行
|
||||
self._state = ConnectionState.OFFLINE
|
||||
self._stop_reconnect = False
|
||||
self._closed = False
|
||||
self._user_close = False
|
||||
self._url = ""
|
||||
self._endpoint_id = ""
|
||||
self._password: Optional[str] = None
|
||||
self._session_token: Optional[str] = None
|
||||
self._use_token = False
|
||||
self._use_tcp = False
|
||||
self._limits = HelloLimits()
|
||||
self._clock_skew_ms = 0
|
||||
self._online_since = 0.0
|
||||
self._backoff_s = BACKOFF_INITIAL_S
|
||||
self._rid_seq = 0
|
||||
self._pending: dict[str, _PendingReq] = {}
|
||||
self._send_queue: list[_SendItem] = []
|
||||
self._inflight_sends = 0
|
||||
self._dedup: OrderedDict[str, str] = OrderedDict()
|
||||
self._receipt_seen: OrderedDict[str, bool] = OrderedDict()
|
||||
self._watch_ids: Optional[list[str]] = None
|
||||
self._watch_all = False
|
||||
self._conn_event = threading.Event()
|
||||
self._handshake_error: Optional[BaseException] = None
|
||||
self._worker: Optional[threading.Thread] = None
|
||||
self._wake = threading.Event()
|
||||
self._want_connected = False
|
||||
|
||||
self._transport.set_handlers(self._on_transport_connected, self._on_transport_disconnected, self._on_down)
|
||||
|
||||
# ---- 回调注册 ----
|
||||
def on_session(self, handler: SessionHandler) -> None:
|
||||
self._session_handler = handler
|
||||
|
||||
def on_message(self, handler: MessageHandler) -> None:
|
||||
self._message_handler = handler
|
||||
|
||||
def on_receipt(self, handler: ReceiptHandler) -> None:
|
||||
self._receipt_handler = handler
|
||||
|
||||
def on_revoked(self, handler: RevokedHandler) -> None:
|
||||
self._revoked_handler = handler
|
||||
|
||||
def on_presence(self, handler: PresenceHandler) -> None:
|
||||
self._presence_handler = handler
|
||||
|
||||
def on_group_event(self, handler: GroupEventHandler) -> None:
|
||||
self._group_handler = handler
|
||||
|
||||
def on_connection(self, handler: ConnectionHandler) -> None:
|
||||
self._connection_handler = handler
|
||||
|
||||
# ---- 连接 ----
|
||||
def connect(
|
||||
self,
|
||||
url: str,
|
||||
endpoint_id: str,
|
||||
*,
|
||||
password: Optional[str] = None,
|
||||
session_token: Optional[str] = None,
|
||||
use_tcp: bool = False,
|
||||
wait: bool = True,
|
||||
) -> None:
|
||||
if password is None and session_token is None:
|
||||
raise ValueError("需要 password 或 session_token")
|
||||
with self._lock:
|
||||
if self._closed:
|
||||
raise ClosedError()
|
||||
self._url = normalize_mqtt_ws_url(url) if not use_tcp else url
|
||||
self._endpoint_id = endpoint_id
|
||||
self._password = password
|
||||
self._session_token = session_token
|
||||
self._use_token = session_token is not None and password is None
|
||||
self._use_tcp = use_tcp
|
||||
self._stop_reconnect = False
|
||||
self._user_close = False
|
||||
self._want_connected = True
|
||||
self._handshake_error = None
|
||||
self._conn_event.clear()
|
||||
self._set_state(ConnectionState.CONNECTING)
|
||||
if self._worker is None or not self._worker.is_alive():
|
||||
self._worker = threading.Thread(target=self._run_loop, name="nixmsg-client", daemon=True)
|
||||
self._worker.start()
|
||||
self._wake.set()
|
||||
if wait:
|
||||
if not self._conn_event.wait(self._connect_timeout_s + 5):
|
||||
raise NixMsgError("busy", "连接超时")
|
||||
err = self._handshake_error
|
||||
if err:
|
||||
raise err
|
||||
if self._state not in (ConnectionState.ONLINE,):
|
||||
if self._state == ConnectionState.AUTH_FAILED:
|
||||
raise NixMsgError(
|
||||
getattr(self, "_auth_reason", "bad_credentials"),
|
||||
"认证失败",
|
||||
)
|
||||
if self._state == ConnectionState.KICKED:
|
||||
raise NixMsgError("taken_over", "会话被接管")
|
||||
raise NixMsgError("busy", f"连接未成功: {self._state.value}")
|
||||
|
||||
def close(self) -> None:
|
||||
with self._lock:
|
||||
self._user_close = True
|
||||
self._want_connected = False
|
||||
self._stop_reconnect = True
|
||||
self._closed = True
|
||||
self._fail_all_pending(ClosedError())
|
||||
self._set_state(ConnectionState.OFFLINE)
|
||||
try:
|
||||
self._transport.disconnect()
|
||||
except Exception:
|
||||
pass
|
||||
self._wake.set()
|
||||
|
||||
def logout(self) -> None:
|
||||
try:
|
||||
self._request({"type": "self.logout"}, wait=True)
|
||||
except Exception:
|
||||
pass
|
||||
with self._lock:
|
||||
self._stop_reconnect = True
|
||||
self._want_connected = False
|
||||
self._session_token = None
|
||||
self._fail_all_pending(NixMsgError("auth_failed", "已退出登录"))
|
||||
try:
|
||||
self._transport.disconnect()
|
||||
except Exception:
|
||||
pass
|
||||
self._set_state(ConnectionState.OFFLINE)
|
||||
self._wake.set()
|
||||
|
||||
# ---- 发送 / 确认 ----
|
||||
def send(self, to: Target, body: Body | str | bytes, options: Optional[SendOptions] = None) -> SendResult:
|
||||
options = options or SendOptions()
|
||||
if isinstance(body, str):
|
||||
body = Body(data=body)
|
||||
elif isinstance(body, bytes):
|
||||
import base64
|
||||
|
||||
body = Body(data=base64.b64encode(body).decode("ascii"), enc="base64")
|
||||
if options.content_type:
|
||||
body.content_type = options.content_type
|
||||
|
||||
with self._lock:
|
||||
if self._closed:
|
||||
raise ClosedError()
|
||||
if len(self._send_queue) >= SEND_QUEUE_LIMIT:
|
||||
raise NixMsgError("quota_exceeded", "发送队列已满")
|
||||
|
||||
max_body = self._limits.max_body_bytes or DEFAULT_MAX_BODY
|
||||
max_meta = self._limits.max_meta_bytes or DEFAULT_MAX_META
|
||||
max_frame = self._limits.max_frame_bytes or DEFAULT_MAX_FRAME
|
||||
|
||||
if body.decoded_size() > max_body:
|
||||
raise NixMsgError("body_too_large", "正文超限")
|
||||
|
||||
meta = options.meta or {}
|
||||
meta_bytes = dumps({"meta": meta}) # 近似;真正检查序列化后 meta 对象
|
||||
# 精确:meta 单独序列化
|
||||
meta_raw = dumps(meta) if meta else b"{}"
|
||||
if len(meta_raw) > max_meta:
|
||||
raise NixMsgError("meta_too_large", "自定义字段超限")
|
||||
|
||||
msg_id = options.message_id or new_uuid7()
|
||||
frame: dict[str, Any] = {
|
||||
"v": 1,
|
||||
"type": "send",
|
||||
"id": msg_id,
|
||||
"to": to.to_dict(),
|
||||
"body": body.to_dict(),
|
||||
}
|
||||
if meta:
|
||||
frame["meta"] = meta
|
||||
|
||||
# send_at_ms 在入队时固定,重交不重算
|
||||
if options.send_at_ms is not None:
|
||||
frame["send_at_ms"] = int(options.send_at_ms)
|
||||
elif options.send_at is not None:
|
||||
# 本机语义时间 + 偏差
|
||||
local_ms = int(options.send_at * 1000) if options.send_at < 1e12 else int(options.send_at)
|
||||
frame["send_at_ms"] = local_ms + self._clock_skew_ms
|
||||
elif options.delay_ms is not None:
|
||||
frame["delay_ms"] = int(options.delay_ms)
|
||||
|
||||
if options.keep:
|
||||
offline: dict[str, Any] = {"keep": True}
|
||||
if options.ttl_seconds is not None:
|
||||
offline["ttl_seconds"] = int(options.ttl_seconds)
|
||||
frame["offline"] = offline
|
||||
if not options.receipt:
|
||||
frame["receipt"] = False
|
||||
if options.talk_password:
|
||||
frame["talk_password"] = options.talk_password
|
||||
|
||||
# 帧大小本地检查(含 rid 占位)
|
||||
probe = dict(frame)
|
||||
probe["rid"] = "0" * 8
|
||||
raw = dumps(probe)
|
||||
if len(raw) > max_frame:
|
||||
raise NixMsgError("frame_too_large", "整帧超限")
|
||||
|
||||
pending = _PendingReq(rid="", frame=frame, is_send=True, message_id=msg_id)
|
||||
item = _SendItem(message_id=msg_id, frame=frame, pending=pending)
|
||||
self._send_queue.append(item)
|
||||
self._wake.set()
|
||||
|
||||
if not pending.event.wait(timeout=None if self._state == ConnectionState.ONLINE else 3600):
|
||||
raise NixMsgError("busy", "发送等待中断")
|
||||
if pending.error:
|
||||
raise pending.error
|
||||
assert pending.response is not None
|
||||
data = pending.response.get("data") or {}
|
||||
return SendResult(id=data.get("id", msg_id), send_at_ms=int(data.get("send_at_ms", 0)), state=str(data.get("state", "")))
|
||||
|
||||
def ack(self, message: IncomingMessage) -> None:
|
||||
self._send_ack(message.from_id, message.id, mark_acked=True)
|
||||
|
||||
# ---- 其余接口 ----
|
||||
def recall(self, message_id: str) -> RecallResult:
|
||||
resp = self._request({"type": "recall", "id": message_id})
|
||||
data = resp.get("data") or {}
|
||||
return RecallResult(
|
||||
result=str(data.get("result", "")),
|
||||
recalled=int(data.get("recalled", 0)),
|
||||
accepted=int(data.get("accepted", 0)),
|
||||
other=int(data.get("other", 0)),
|
||||
)
|
||||
|
||||
def status(self, message_id: str, cursor: str = "", limit: int = 100) -> dict[str, Any]:
|
||||
return self._request({"type": "status", "id": message_id, "cursor": cursor, "limit": limit})
|
||||
|
||||
def unlock(self, endpoint_id: str, talk_password: str) -> dict[str, Any]:
|
||||
return self._request({"type": "unlock", "endpoint_id": endpoint_id, "talk_password": talk_password})
|
||||
|
||||
def presence(self, ids: list[str]) -> dict[str, Any]:
|
||||
return self._request({"type": "presence.get", "ids": ids})
|
||||
|
||||
def directory(self, cursor: str = "", query: str = "", limit: int = 100) -> dict[str, Any]:
|
||||
return self._request({"type": "directory.list", "cursor": cursor, "limit": limit, "query": query})
|
||||
|
||||
def watch_presence(self, ids: Optional[list[str]] = None, *, all: bool = False) -> dict[str, Any]:
|
||||
with self._lock:
|
||||
self._watch_ids = list(ids) if ids is not None else None
|
||||
self._watch_all = all
|
||||
frame: dict[str, Any] = {"type": "presence.watch", "all": all}
|
||||
if ids is not None:
|
||||
frame["ids"] = ids
|
||||
return self._request(frame)
|
||||
|
||||
def get_self(self) -> dict[str, Any]:
|
||||
return self._request({"type": "self.get"})
|
||||
|
||||
def update_self(self, *, name: Optional[str] = None, default_delay_ms: Optional[int] = None) -> dict[str, Any]:
|
||||
frame: dict[str, Any] = {"type": "self.update"}
|
||||
if name is not None:
|
||||
frame["name"] = name
|
||||
if default_delay_ms is not None:
|
||||
frame["default_delay_ms"] = default_delay_ms
|
||||
return self._request(frame)
|
||||
|
||||
def set_talk_password(self, talk_password: str) -> dict[str, Any]:
|
||||
return self._request({"type": "self.talk_password", "talk_password": talk_password})
|
||||
|
||||
def change_login_password(self, old_password: str, new_password: str) -> dict[str, Any]:
|
||||
resp = self._request({"type": "self.login_password", "old_password": old_password, "new_password": new_password})
|
||||
data = resp.get("data") or {}
|
||||
token = data.get("session_token")
|
||||
if token:
|
||||
with self._lock:
|
||||
self._session_token = token
|
||||
self._use_token = True
|
||||
self._fire_session(token)
|
||||
return resp
|
||||
|
||||
def group_create(self, name: str, members: list[dict[str, str]], group_id: str = "") -> dict[str, Any]:
|
||||
frame: dict[str, Any] = {"type": "group.create", "name": name, "members": members}
|
||||
if group_id:
|
||||
frame["id"] = group_id
|
||||
else:
|
||||
frame["id"] = ""
|
||||
return self._request(frame)
|
||||
|
||||
def group_add(self, group_id: str, members: list[dict[str, str]]) -> dict[str, Any]:
|
||||
return self._request({"type": "group.add", "group_id": group_id, "members": members})
|
||||
|
||||
def group_remove(self, group_id: str, endpoint_id: str) -> dict[str, Any]:
|
||||
return self._request({"type": "group.remove", "group_id": group_id, "endpoint_id": endpoint_id})
|
||||
|
||||
def group_leave(self, group_id: str) -> dict[str, Any]:
|
||||
return self._request({"type": "group.leave", "group_id": group_id})
|
||||
|
||||
def group_transfer(self, group_id: str, endpoint_id: str) -> dict[str, Any]:
|
||||
return self._request({"type": "group.transfer", "group_id": group_id, "endpoint_id": endpoint_id})
|
||||
|
||||
def group_rename(self, group_id: str, name: str) -> dict[str, Any]:
|
||||
return self._request({"type": "group.rename", "group_id": group_id, "name": name})
|
||||
|
||||
def group_dissolve(self, group_id: str) -> dict[str, Any]:
|
||||
return self._request({"type": "group.dissolve", "group_id": group_id})
|
||||
|
||||
def group_list(self, cursor: str = "", limit: int = 100) -> dict[str, Any]:
|
||||
return self._request({"type": "group.list", "cursor": cursor, "limit": limit})
|
||||
|
||||
def group_get(self, group_id: str, cursor: str = "", limit: int = 100) -> dict[str, Any]:
|
||||
return self._request({"type": "group.get", "group_id": group_id, "cursor": cursor, "limit": limit})
|
||||
|
||||
@staticmethod
|
||||
def register(url: str, registration_code: str, options: Optional[RegisterOptions] = None) -> RegisterResult:
|
||||
options = options or RegisterOptions()
|
||||
reg_url = register_url_from_connect(url)
|
||||
body = {
|
||||
"registration_code": registration_code,
|
||||
"id": options.id or "",
|
||||
"login_password": options.login_password or "",
|
||||
"name": options.name or "",
|
||||
"talk_password": options.talk_password or "",
|
||||
}
|
||||
raw = dumps(body)
|
||||
req = urllib.request.Request(
|
||||
reg_url,
|
||||
data=raw,
|
||||
headers={"Content-Type": "application/json"},
|
||||
method="POST",
|
||||
)
|
||||
try:
|
||||
with urllib.request.urlopen(req, timeout=30) as resp:
|
||||
data = loads(resp.read())
|
||||
except urllib.error.HTTPError as e:
|
||||
try:
|
||||
data = loads(e.read())
|
||||
except Exception:
|
||||
raise NixMsgError("bad_request", f"HTTP {e.code}") from e
|
||||
err = data.get("error") or {}
|
||||
raise NixMsgError(str(err.get("code", "bad_request")), str(err.get("message", ""))) from e
|
||||
if not data.get("ok"):
|
||||
err = data.get("error") or {}
|
||||
raise NixMsgError(str(err.get("code", "bad_request")), str(err.get("message", "")))
|
||||
d = data.get("data") or {}
|
||||
return RegisterResult(id=str(d.get("id", "")), login_password=d.get("login_password"))
|
||||
|
||||
# ---- 内部:连接循环 ----
|
||||
def _run_loop(self) -> None:
|
||||
while True:
|
||||
with self._lock:
|
||||
if self._closed and not self._want_connected:
|
||||
return
|
||||
want = self._want_connected and not self._stop_reconnect
|
||||
state = self._state
|
||||
if not want:
|
||||
self._wake.wait(0.5)
|
||||
self._wake.clear()
|
||||
continue
|
||||
if state == ConnectionState.ONLINE:
|
||||
self._pump_sends()
|
||||
# 稳定 60 秒后恢复退避
|
||||
if self._online_since and time.monotonic() - self._online_since >= STABLE_RESET_S:
|
||||
self._backoff_s = BACKOFF_INITIAL_S
|
||||
self._wake.wait(0.2)
|
||||
self._wake.clear()
|
||||
continue
|
||||
# 尝试连接
|
||||
try:
|
||||
self._attempt_connect()
|
||||
except Exception as e:
|
||||
log.debug("connect attempt failed: %s", e)
|
||||
with self._lock:
|
||||
if self._state == ConnectionState.ONLINE:
|
||||
continue
|
||||
if self._stop_reconnect or not self._want_connected:
|
||||
continue
|
||||
delay = self._jitter(self._backoff_s)
|
||||
self._backoff_s = min(BACKOFF_MAX_S, self._backoff_s * 2)
|
||||
self._set_state(ConnectionState.RECONNECTING)
|
||||
self._wake.wait(delay)
|
||||
self._wake.clear()
|
||||
|
||||
def _attempt_connect(self) -> None:
|
||||
with self._lock:
|
||||
if self._stop_reconnect or not self._want_connected:
|
||||
return
|
||||
self._set_state(ConnectionState.CONNECTING if self._state == ConnectionState.OFFLINE else ConnectionState.RECONNECTING)
|
||||
self._handshake_error = None
|
||||
url = self._url
|
||||
eid = self._endpoint_id
|
||||
if self._session_token:
|
||||
cred = self._session_token
|
||||
using_token = True
|
||||
else:
|
||||
cred = self._password or ""
|
||||
using_token = False
|
||||
self._connecting_with_token = using_token
|
||||
use_tcp = self._use_tcp
|
||||
timeout = self._connect_timeout_s
|
||||
|
||||
self._conn_event.clear()
|
||||
params = ConnectParams(
|
||||
url=url,
|
||||
client_id=eid,
|
||||
username=eid,
|
||||
password=cred,
|
||||
clean_start=True,
|
||||
session_expiry=0,
|
||||
timeout_s=timeout,
|
||||
use_tcp=use_tcp,
|
||||
)
|
||||
self._transport.connect(params)
|
||||
# 等待握手完成或失败
|
||||
ok = self._conn_event.wait(timeout)
|
||||
if not ok:
|
||||
try:
|
||||
self._transport.disconnect()
|
||||
except Exception:
|
||||
pass
|
||||
with self._lock:
|
||||
self._handshake_error = NixMsgError("busy", "连接超时")
|
||||
return
|
||||
|
||||
def _on_transport_connected(self) -> None:
|
||||
try:
|
||||
topic = down_topic(self._endpoint_id)
|
||||
self._transport.subscribe(topic)
|
||||
rid = self._next_rid()
|
||||
hello = {
|
||||
"v": 1,
|
||||
"type": "hello",
|
||||
"rid": rid,
|
||||
"max_receive_bytes": self._max_receive_bytes,
|
||||
"client": self._client_name,
|
||||
}
|
||||
t0 = time.time()
|
||||
pending = _PendingReq(rid=rid, frame=hello)
|
||||
with self._lock:
|
||||
self._pending[rid] = pending
|
||||
self._transport.publish(up_topic(self._endpoint_id), dumps(hello))
|
||||
if not pending.event.wait(self._connect_timeout_s):
|
||||
raise NixMsgError("busy", "握手超时")
|
||||
if pending.error:
|
||||
raise pending.error
|
||||
assert pending.response is not None
|
||||
if not pending.response.get("ok"):
|
||||
err = (pending.response.get("error") or {})
|
||||
raise NixMsgError(str(err.get("code", "bad_request")), str(err.get("message", "")))
|
||||
t1 = time.time()
|
||||
data = pending.response.get("data") or {}
|
||||
limits = HelloLimits(
|
||||
server_time_ms=int(data.get("server_time_ms", 0)),
|
||||
server_version=str(data.get("server_version", "")),
|
||||
max_body_bytes=int(data.get("max_body_bytes", DEFAULT_MAX_BODY)),
|
||||
max_meta_bytes=int(data.get("max_meta_bytes", DEFAULT_MAX_META)),
|
||||
max_frame_bytes=int(data.get("max_frame_bytes", DEFAULT_MAX_FRAME)),
|
||||
max_ttl_seconds=int(data.get("max_ttl_seconds", 2592000)),
|
||||
max_schedule_seconds=int(data.get("max_schedule_seconds", 31536000)),
|
||||
ack_timeout_seconds=int(data.get("ack_timeout_seconds", 300)),
|
||||
session_token=str(data.get("session_token", "")),
|
||||
)
|
||||
skew = limits.server_time_ms - int(((t0 + t1) / 2) * 1000)
|
||||
with self._lock:
|
||||
self._limits = limits
|
||||
self._clock_skew_ms = skew
|
||||
self._online_since = time.monotonic()
|
||||
self._handshake_error = None
|
||||
self._set_state(ConnectionState.ONLINE)
|
||||
token = limits.session_token
|
||||
if token:
|
||||
with self._lock:
|
||||
self._session_token = token
|
||||
self._use_token = True
|
||||
self._fire_session(token)
|
||||
# 重连后恢复 presence.watch
|
||||
if self._watch_all or self._watch_ids is not None:
|
||||
try:
|
||||
frame: dict[str, Any] = {"type": "presence.watch", "all": self._watch_all}
|
||||
if self._watch_ids is not None:
|
||||
frame["ids"] = self._watch_ids
|
||||
self._request(frame, wait=False)
|
||||
except Exception:
|
||||
pass
|
||||
self._conn_event.set()
|
||||
self._wake.set()
|
||||
except Exception as e:
|
||||
with self._lock:
|
||||
self._handshake_error = e
|
||||
try:
|
||||
self._transport.disconnect()
|
||||
except Exception:
|
||||
pass
|
||||
self._conn_event.set()
|
||||
|
||||
def _on_transport_disconnected(self, reason: Optional[str], stop: bool) -> None:
|
||||
with self._lock:
|
||||
was_online = self._state == ConnectionState.ONLINE
|
||||
using_token = getattr(self, "_connecting_with_token", self._use_token)
|
||||
if reason == "taken_over":
|
||||
self._stop_reconnect = True
|
||||
self._want_connected = False
|
||||
self._fail_all_pending(NixMsgError("taken_over", "会话被接管"))
|
||||
self._set_state(ConnectionState.KICKED, "taken_over")
|
||||
self._conn_event.set()
|
||||
return
|
||||
if stop or reason in ("session_invalid", "bad_credentials", "banned"):
|
||||
# 令牌被拒 -> session_invalid;密码被拒 -> bad_credentials
|
||||
if reason in ("bad_credentials", "session_invalid", "banned") or stop:
|
||||
auth_reason = reason or "bad_credentials"
|
||||
if auth_reason == "bad_credentials" and using_token:
|
||||
auth_reason = "session_invalid"
|
||||
if auth_reason == "banned":
|
||||
auth_reason = "session_invalid" if using_token else "bad_credentials"
|
||||
self._auth_reason = auth_reason
|
||||
self._stop_reconnect = True
|
||||
self._want_connected = False
|
||||
self._fail_all_pending(NixMsgError(auth_reason, "认证失败"))
|
||||
self._handshake_error = NixMsgError(auth_reason, "认证失败")
|
||||
self._set_state(ConnectionState.AUTH_FAILED, auth_reason)
|
||||
self._conn_event.set()
|
||||
return
|
||||
# 网络 / busy:继续重连
|
||||
if self._user_close or self._closed:
|
||||
self._set_state(ConnectionState.OFFLINE)
|
||||
self._conn_event.set()
|
||||
return
|
||||
if was_online or self._state in (ConnectionState.CONNECTING, ConnectionState.RECONNECTING):
|
||||
self._set_state(ConnectionState.RECONNECTING, reason or "network")
|
||||
self._conn_event.set()
|
||||
self._wake.set()
|
||||
|
||||
def _on_down(self, payload: bytes) -> None:
|
||||
try:
|
||||
frame = loads(payload)
|
||||
except Exception:
|
||||
return
|
||||
ftype = frame.get("type")
|
||||
if ftype == "resp":
|
||||
rid = str(frame.get("rid", ""))
|
||||
with self._lock:
|
||||
pending = self._pending.pop(rid, None)
|
||||
if pending:
|
||||
pending.response = frame
|
||||
if pending.is_send:
|
||||
with self._lock:
|
||||
self._inflight_sends = max(0, self._inflight_sends - 1)
|
||||
# rate_limited 重交:清状态后不 set event
|
||||
err = (frame.get("error") or {}) if not frame.get("ok") else {}
|
||||
if not frame.get("ok") and str(err.get("code")) == "rate_limited":
|
||||
with self._lock:
|
||||
pending.rid = ""
|
||||
pending.response = None
|
||||
pending.error = None
|
||||
pending.event.clear()
|
||||
self._wake.set()
|
||||
return
|
||||
# 从发送队列移除
|
||||
with self._lock:
|
||||
self._send_queue = [it for it in self._send_queue if it.pending is not pending]
|
||||
if not frame.get("ok"):
|
||||
pending.error = NixMsgError(str(err.get("code", "bad_request")), str(err.get("message", "")))
|
||||
pending.event.set()
|
||||
self._wake.set()
|
||||
else:
|
||||
pending.event.set()
|
||||
return
|
||||
if ftype == "msg":
|
||||
self._handle_msg(frame)
|
||||
return
|
||||
if ftype == "receipt":
|
||||
self._handle_receipt(frame)
|
||||
return
|
||||
if ftype == "revoked":
|
||||
self._handle_revoked(frame)
|
||||
return
|
||||
if ftype == "presence":
|
||||
self._fire_presence(
|
||||
PresenceEvent(id=str(frame.get("id", "")), online=bool(frame.get("online")), at_ms=int(frame.get("at_ms", 0)))
|
||||
)
|
||||
return
|
||||
if ftype == "group_event":
|
||||
self._fire_group(
|
||||
GroupEvent(
|
||||
group_id=str(frame.get("group_id", "")),
|
||||
event=str(frame.get("event", "")),
|
||||
endpoint_id=str(frame.get("endpoint_id", "")),
|
||||
at_ms=int(frame.get("at_ms", 0)),
|
||||
)
|
||||
)
|
||||
return
|
||||
if ftype == "fatal":
|
||||
reason = str(frame.get("reason", "protocol"))
|
||||
with self._lock:
|
||||
self._stop_reconnect = True
|
||||
self._want_connected = False
|
||||
self._fail_all_pending(NixMsgError("fatal", reason))
|
||||
self._set_state(ConnectionState.AUTH_FAILED, reason)
|
||||
try:
|
||||
self._transport.disconnect()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _handle_msg(self, frame: dict[str, Any]) -> None:
|
||||
mid = str(frame.get("id", ""))
|
||||
from_id = str(frame.get("from", ""))
|
||||
key = f"{from_id}\0{mid}"
|
||||
with self._lock:
|
||||
st = self._dedup.get(key)
|
||||
if st == _DedupState.ACKED:
|
||||
# 已确认再到达:再 ack,不交应用
|
||||
pass
|
||||
elif st == _DedupState.DELIVERED:
|
||||
# 已交未确认:忽略
|
||||
return
|
||||
else:
|
||||
st = None
|
||||
if st == _DedupState.ACKED:
|
||||
self._send_ack(from_id, mid, mark_acked=True)
|
||||
return
|
||||
|
||||
to_raw = frame.get("to") or {}
|
||||
msg = IncomingMessage(
|
||||
id=mid,
|
||||
from_id=from_id,
|
||||
to=Target(kind=str(to_raw.get("kind", "endpoint")), id=str(to_raw.get("id", ""))),
|
||||
body=Body(
|
||||
data=str((frame.get("body") or {}).get("data", "")),
|
||||
enc=str((frame.get("body") or {}).get("enc", "utf8")),
|
||||
content_type=(frame.get("body") or {}).get("content_type"),
|
||||
),
|
||||
send_at_ms=int(frame.get("send_at_ms", 0)),
|
||||
meta=dict(frame.get("meta") or {}),
|
||||
)
|
||||
with self._lock:
|
||||
self._dedup_put(key, _DedupState.DELIVERED)
|
||||
|
||||
if not self._message_handler:
|
||||
if self._auto_ack:
|
||||
self._send_ack(from_id, mid, mark_acked=True)
|
||||
return
|
||||
|
||||
try:
|
||||
with self._cb_lock:
|
||||
self._message_handler(msg)
|
||||
except Exception:
|
||||
log.exception("on_message 回调错误,等待重推")
|
||||
with self._lock:
|
||||
self._dedup.pop(key, None)
|
||||
return
|
||||
|
||||
if self._auto_ack:
|
||||
self._send_ack(from_id, mid, mark_acked=True)
|
||||
|
||||
def _handle_receipt(self, frame: dict[str, Any]) -> None:
|
||||
rid = str(frame.get("receipt_id", ""))
|
||||
with self._lock:
|
||||
if rid in self._receipt_seen:
|
||||
return
|
||||
self._receipt_seen[rid] = True
|
||||
while len(self._receipt_seen) > DEDUP_CAPACITY:
|
||||
self._receipt_seen.popitem(last=False)
|
||||
receipt = Receipt(
|
||||
receipt_id=rid,
|
||||
id=str(frame.get("id", "")),
|
||||
endpoint_id=str(frame.get("endpoint_id", "")),
|
||||
state=str(frame.get("state", "")),
|
||||
reason=str(frame.get("reason", "")),
|
||||
at_ms=int(frame.get("at_ms", 0)),
|
||||
)
|
||||
self._fire_receipt(receipt)
|
||||
# 自动 receipt_ack
|
||||
try:
|
||||
self._request({"type": "receipt_ack", "receipt_id": rid}, wait=False)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _handle_revoked(self, frame: dict[str, Any]) -> None:
|
||||
mid = str(frame.get("id", ""))
|
||||
from_id = str(frame.get("from", ""))
|
||||
key = f"{from_id}\0{mid}"
|
||||
with self._lock:
|
||||
st = self._dedup.get(key)
|
||||
if st == _DedupState.ACKED:
|
||||
return
|
||||
if st is None:
|
||||
# 还没交给应用:直接丢弃
|
||||
return
|
||||
# 已交未确认:发撤回事件并不再确认
|
||||
self._dedup.pop(key, None)
|
||||
self._fire_revoked(RevokedEvent(id=mid, from_id=from_id, reason=str(frame.get("reason", ""))))
|
||||
|
||||
def _send_ack(self, from_id: str, message_id: str, *, mark_acked: bool) -> None:
|
||||
try:
|
||||
resp = self._request({"type": "ack", "from": from_id, "id": message_id})
|
||||
except Exception:
|
||||
return
|
||||
key = f"{from_id}\0{message_id}"
|
||||
data = resp.get("data") or {}
|
||||
result = data.get("result", "accepted")
|
||||
if result != "accepted":
|
||||
with self._lock:
|
||||
self._dedup.pop(key, None)
|
||||
self._fire_revoked(RevokedEvent(id=message_id, from_id=from_id, reason=str(result)))
|
||||
return
|
||||
if mark_acked:
|
||||
with self._lock:
|
||||
self._dedup_put(key, _DedupState.ACKED)
|
||||
|
||||
def _pump_sends(self) -> None:
|
||||
while True:
|
||||
with self._lock:
|
||||
if self._state != ConnectionState.ONLINE:
|
||||
return
|
||||
if self._inflight_sends >= INFLIGHT_LIMIT:
|
||||
return
|
||||
item = None
|
||||
for it in self._send_queue:
|
||||
if it.pending.rid == "" and it.pending.response is None and it.pending.error is None:
|
||||
item = it
|
||||
break
|
||||
if item is None:
|
||||
return
|
||||
rid = self._next_rid_locked()
|
||||
item.pending.rid = rid
|
||||
frame = dict(item.frame)
|
||||
frame["rid"] = rid
|
||||
self._pending[rid] = item.pending
|
||||
self._inflight_sends += 1
|
||||
topic = up_topic(self._endpoint_id)
|
||||
try:
|
||||
self._transport.publish(topic, dumps(frame))
|
||||
except Exception as e:
|
||||
with self._lock:
|
||||
self._pending.pop(rid, None)
|
||||
item.pending.rid = ""
|
||||
self._inflight_sends = max(0, self._inflight_sends - 1)
|
||||
item.pending.error = e
|
||||
item.pending.event.set()
|
||||
self._send_queue = [it for it in self._send_queue if it is not item]
|
||||
continue
|
||||
|
||||
def _request(self, frame: dict[str, Any], *, wait: bool = True) -> dict[str, Any]:
|
||||
with self._lock:
|
||||
if self._closed:
|
||||
raise ClosedError()
|
||||
if self._state != ConnectionState.ONLINE:
|
||||
raise NotConnectedError()
|
||||
rid = self._next_rid_locked()
|
||||
body = dict(frame)
|
||||
body["v"] = 1
|
||||
body["rid"] = rid
|
||||
pending = _PendingReq(rid=rid, frame=body)
|
||||
self._pending[rid] = pending
|
||||
topic = up_topic(self._endpoint_id)
|
||||
self._transport.publish(topic, dumps(body))
|
||||
if not wait:
|
||||
# 仍等一小会拿结果;后台请求也尽量同步完成
|
||||
pending.event.wait(timeout=30)
|
||||
with self._lock:
|
||||
self._pending.pop(rid, None)
|
||||
if pending.error:
|
||||
raise pending.error
|
||||
if pending.response is None:
|
||||
return {}
|
||||
if not pending.response.get("ok"):
|
||||
err = pending.response.get("error") or {}
|
||||
raise NixMsgError(str(err.get("code", "bad_request")), str(err.get("message", "")))
|
||||
return pending.response
|
||||
if not pending.event.wait(timeout=60):
|
||||
with self._lock:
|
||||
self._pending.pop(rid, None)
|
||||
raise NixMsgError("busy", "请求超时")
|
||||
if pending.error:
|
||||
raise pending.error
|
||||
assert pending.response is not None
|
||||
if not pending.response.get("ok"):
|
||||
err = pending.response.get("error") or {}
|
||||
raise NixMsgError(str(err.get("code", "bad_request")), str(err.get("message", "")))
|
||||
return pending.response
|
||||
|
||||
def _next_rid(self) -> str:
|
||||
with self._lock:
|
||||
return self._next_rid_locked()
|
||||
|
||||
def _next_rid_locked(self) -> str:
|
||||
self._rid_seq += 1
|
||||
return str(self._rid_seq)
|
||||
|
||||
def _dedup_put(self, key: str, state: str) -> None:
|
||||
if key in self._dedup:
|
||||
self._dedup.move_to_end(key)
|
||||
self._dedup[key] = state
|
||||
while len(self._dedup) > DEDUP_CAPACITY:
|
||||
self._dedup.popitem(last=False)
|
||||
|
||||
def _fail_all_pending(self, err: BaseException) -> None:
|
||||
for p in list(self._pending.values()):
|
||||
p.error = err
|
||||
p.event.set()
|
||||
self._pending.clear()
|
||||
for it in list(self._send_queue):
|
||||
it.pending.error = err
|
||||
it.pending.event.set()
|
||||
self._send_queue.clear()
|
||||
self._inflight_sends = 0
|
||||
|
||||
def _set_state(self, state: ConnectionState, reason: str = "") -> None:
|
||||
self._state = state
|
||||
h = self._connection_handler
|
||||
if h:
|
||||
try:
|
||||
with self._cb_lock:
|
||||
h(ConnectionEvent(state=state, reason=reason))
|
||||
except Exception:
|
||||
log.exception("on_connection 回调错误")
|
||||
|
||||
def _fire_session(self, token: str) -> None:
|
||||
h = self._session_handler
|
||||
if h:
|
||||
try:
|
||||
with self._cb_lock:
|
||||
h(token)
|
||||
except Exception:
|
||||
log.exception("on_session 回调错误")
|
||||
|
||||
def _fire_receipt(self, receipt: Receipt) -> None:
|
||||
h = self._receipt_handler
|
||||
if h:
|
||||
try:
|
||||
with self._cb_lock:
|
||||
h(receipt)
|
||||
except Exception:
|
||||
log.exception("on_receipt 回调错误")
|
||||
|
||||
def _fire_revoked(self, ev: RevokedEvent) -> None:
|
||||
h = self._revoked_handler
|
||||
if h:
|
||||
try:
|
||||
with self._cb_lock:
|
||||
h(ev)
|
||||
except Exception:
|
||||
log.exception("on_revoked 回调错误")
|
||||
|
||||
def _fire_presence(self, ev: PresenceEvent) -> None:
|
||||
h = self._presence_handler
|
||||
if h:
|
||||
try:
|
||||
with self._cb_lock:
|
||||
h(ev)
|
||||
except Exception:
|
||||
log.exception("on_presence 回调错误")
|
||||
|
||||
def _fire_group(self, ev: GroupEvent) -> None:
|
||||
h = self._group_handler
|
||||
if h:
|
||||
try:
|
||||
with self._cb_lock:
|
||||
h(ev)
|
||||
except Exception:
|
||||
log.exception("on_group_event 回调错误")
|
||||
|
||||
@staticmethod
|
||||
def _jitter(base: float) -> float:
|
||||
return max(0.0, base * (1.0 + random.uniform(-BACKOFF_JITTER, BACKOFF_JITTER)))
|
||||
|
||||
# 测试辅助
|
||||
@property
|
||||
def state(self) -> ConnectionState:
|
||||
return self._state
|
||||
|
||||
@property
|
||||
def limits(self) -> HelloLimits:
|
||||
return self._limits
|
||||
|
||||
@property
|
||||
def clock_skew_ms(self) -> int:
|
||||
return self._clock_skew_ms
|
||||
|
||||
@property
|
||||
def session_token(self) -> Optional[str]:
|
||||
return self._session_token
|
||||
|
||||
|
||||
def _auto_hello_responder(client: Client, transport: FakeTransport, *, session_token: str = "nst_test") -> None:
|
||||
"""测试辅助:对 hello 自动回成功(由测试自行调用更清晰)。"""
|
||||
_ = (client, transport, session_token)
|
||||
@@ -0,0 +1,20 @@
|
||||
"""NixMsg SDK 错误。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
class NixMsgError(Exception):
|
||||
def __init__(self, code: str, message: str = "") -> None:
|
||||
self.code = code
|
||||
self.message = message or code
|
||||
super().__init__(self.message)
|
||||
|
||||
|
||||
class NotConnectedError(NixMsgError):
|
||||
def __init__(self, message: str = "未连接") -> None:
|
||||
super().__init__("not_connected", message)
|
||||
|
||||
|
||||
class ClosedError(NixMsgError):
|
||||
def __init__(self, message: str = "已关闭") -> None:
|
||||
super().__init__("closed", message)
|
||||
@@ -0,0 +1,58 @@
|
||||
"""帧编解码与注册地址推导。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
from urllib.parse import urlparse, urlunparse
|
||||
|
||||
|
||||
def dumps(obj: dict[str, Any]) -> bytes:
|
||||
return json.dumps(obj, ensure_ascii=False, separators=(",", ":")).encode("utf-8")
|
||||
|
||||
|
||||
def loads(data: bytes | str) -> dict[str, Any]:
|
||||
if isinstance(data, bytes):
|
||||
data = data.decode("utf-8")
|
||||
return json.loads(data)
|
||||
|
||||
|
||||
def register_url_from_connect(connect_url: str) -> str:
|
||||
"""从连接地址推出注册 HTTP 地址(DEVELOPMENT 6.9)。"""
|
||||
raw = connect_url.strip()
|
||||
if "://" not in raw:
|
||||
raw = "ws://" + raw
|
||||
u = urlparse(raw)
|
||||
scheme = u.scheme.lower()
|
||||
if scheme in ("wss", "mqtts", "https"):
|
||||
http_scheme = "https"
|
||||
elif scheme in ("ws", "mqtt", "http"):
|
||||
http_scheme = "http"
|
||||
else:
|
||||
http_scheme = "https" if scheme.endswith("s") else "http"
|
||||
# 去掉 /mqtt 路径
|
||||
path = u.path or ""
|
||||
if path.endswith("/mqtt"):
|
||||
path = path[: -len("/mqtt")]
|
||||
path = path.rstrip("/") + "/api/client/register"
|
||||
return urlunparse((http_scheme, u.netloc, path, "", "", ""))
|
||||
|
||||
|
||||
def up_topic(endpoint_id: str) -> str:
|
||||
return f"nix/c/{endpoint_id}/up"
|
||||
|
||||
|
||||
def down_topic(endpoint_id: str) -> str:
|
||||
return f"nix/c/{endpoint_id}/down"
|
||||
|
||||
|
||||
def normalize_mqtt_ws_url(url: str) -> str:
|
||||
"""保证 WebSocket 路径以 /mqtt 结尾。"""
|
||||
raw = url.strip()
|
||||
if "://" not in raw:
|
||||
raw = "ws://" + raw
|
||||
u = urlparse(raw)
|
||||
path = u.path or ""
|
||||
if not path.endswith("/mqtt"):
|
||||
path = path.rstrip("/") + "/mqtt"
|
||||
return urlunparse((u.scheme, u.netloc, path, u.params, u.query, u.fragment))
|
||||
@@ -0,0 +1,350 @@
|
||||
"""MQTT 传输抽象、假传输与 paho 实现。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Callable, Optional, Protocol
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from paho.mqtt.client import CallbackAPIVersion, Client as PahoClient, MQTT_ERR_SUCCESS
|
||||
from paho.mqtt.enums import MQTTErrorCode
|
||||
from paho.mqtt.reasoncodes import ReasonCode
|
||||
|
||||
DownHandler = Callable[[bytes], None]
|
||||
ConnHandler = Callable[[], None]
|
||||
DiscHandler = Callable[[Optional[str], bool], None]
|
||||
# reason_code_str, stop_reconnect
|
||||
|
||||
|
||||
@dataclass
|
||||
class ConnectParams:
|
||||
url: str
|
||||
client_id: str
|
||||
username: str
|
||||
password: str
|
||||
clean_start: bool = True
|
||||
session_expiry: int = 0
|
||||
keep_alive: int = 30
|
||||
timeout_s: float = 30.0
|
||||
use_tcp: bool = False # 显式打开裸 TCP
|
||||
|
||||
|
||||
class Transport(Protocol):
|
||||
def set_handlers(
|
||||
self,
|
||||
on_connected: ConnHandler,
|
||||
on_disconnected: DiscHandler,
|
||||
on_down: DownHandler,
|
||||
) -> None: ...
|
||||
|
||||
def connect(self, params: ConnectParams) -> None: ...
|
||||
|
||||
def subscribe(self, topic: str) -> None: ...
|
||||
|
||||
def publish(self, topic: str, payload: bytes) -> None: ...
|
||||
|
||||
def disconnect(self) -> None: ...
|
||||
|
||||
|
||||
@dataclass
|
||||
class FakeConnectRecord:
|
||||
params: ConnectParams
|
||||
at: float = field(default_factory=time.time)
|
||||
|
||||
|
||||
class FakeTransport:
|
||||
"""单元测试用假传输:不连真实服务器。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._on_connected: Optional[ConnHandler] = None
|
||||
self._on_disconnected: Optional[DiscHandler] = None
|
||||
self._on_down: Optional[DownHandler] = None
|
||||
self.connects: list[FakeConnectRecord] = []
|
||||
self.publishes: list[tuple[str, bytes]] = []
|
||||
self.subscriptions: list[str] = []
|
||||
self.connected = False
|
||||
self.auto_accept = True
|
||||
self.next_connack_fail: Optional[str] = None # session_invalid / bad_credentials / busy
|
||||
self.auto_hello: Optional[dict] = {
|
||||
"server_time_ms": 1_750_000_000_000,
|
||||
"server_version": "0.1.0",
|
||||
"max_body_bytes": 262144,
|
||||
"max_meta_bytes": 4096,
|
||||
"max_frame_bytes": 786432,
|
||||
"max_ttl_seconds": 2592000,
|
||||
"max_schedule_seconds": 31536000,
|
||||
"ack_timeout_seconds": 300,
|
||||
"session_token": "nst_test_token",
|
||||
}
|
||||
self.auto_send_ok = True
|
||||
self._up_handlers: list[Callable[[dict], Optional[dict]]] = []
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def on_up(self, handler: Callable[[dict], Optional[dict]]) -> None:
|
||||
"""上行帧钩子:返回 dict 则作为 resp 注入;返回 None 表示不处理。"""
|
||||
self._up_handlers.append(handler)
|
||||
|
||||
def set_handlers(
|
||||
self,
|
||||
on_connected: ConnHandler,
|
||||
on_disconnected: DiscHandler,
|
||||
on_down: DownHandler,
|
||||
) -> None:
|
||||
self._on_connected = on_connected
|
||||
self._on_disconnected = on_disconnected
|
||||
self._on_down = on_down
|
||||
|
||||
def connect(self, params: ConnectParams) -> None:
|
||||
with self._lock:
|
||||
self.connects.append(FakeConnectRecord(params=params))
|
||||
fail = self.next_connack_fail
|
||||
self.next_connack_fail = None
|
||||
if fail:
|
||||
stop = fail in ("session_invalid", "bad_credentials", "banned")
|
||||
if self._on_disconnected:
|
||||
self._on_disconnected(fail, stop)
|
||||
return
|
||||
self.connected = True
|
||||
if self.auto_accept and self._on_connected:
|
||||
self._on_connected()
|
||||
|
||||
def subscribe(self, topic: str) -> None:
|
||||
self.subscriptions.append(topic)
|
||||
|
||||
def publish(self, topic: str, payload: bytes) -> None:
|
||||
self.publishes.append((topic, payload))
|
||||
import json
|
||||
|
||||
try:
|
||||
frame = json.loads(payload.decode("utf-8"))
|
||||
except Exception:
|
||||
return
|
||||
# 自定义钩子优先
|
||||
for h in list(self._up_handlers):
|
||||
try:
|
||||
resp = h(frame)
|
||||
except Exception:
|
||||
continue
|
||||
if resp is not None:
|
||||
self.inject_down(json.dumps(resp, ensure_ascii=False).encode("utf-8"))
|
||||
return
|
||||
ftype = frame.get("type")
|
||||
rid = frame.get("rid")
|
||||
if ftype == "hello" and self.auto_hello is not None:
|
||||
resp = {"v": 1, "type": "resp", "rid": rid, "ok": True, "data": dict(self.auto_hello)}
|
||||
self.inject_down(json.dumps(resp, ensure_ascii=False).encode("utf-8"))
|
||||
return
|
||||
if ftype == "send" and self.auto_send_ok:
|
||||
resp = {
|
||||
"v": 1,
|
||||
"type": "resp",
|
||||
"rid": rid,
|
||||
"ok": True,
|
||||
"data": {"id": frame.get("id"), "send_at_ms": frame.get("send_at_ms") or 0, "state": "dispatched"},
|
||||
}
|
||||
self.inject_down(json.dumps(resp, ensure_ascii=False).encode("utf-8"))
|
||||
return
|
||||
if ftype == "ack":
|
||||
resp = {"v": 1, "type": "resp", "rid": rid, "ok": True, "data": {"result": "accepted"}}
|
||||
self.inject_down(json.dumps(resp, ensure_ascii=False).encode("utf-8"))
|
||||
return
|
||||
if ftype and ftype not in ("hello", "send") and rid is not None:
|
||||
# 其它请求默认成功
|
||||
resp = {"v": 1, "type": "resp", "rid": rid, "ok": True, "data": {}}
|
||||
self.inject_down(json.dumps(resp, ensure_ascii=False).encode("utf-8"))
|
||||
|
||||
def disconnect(self) -> None:
|
||||
was = self.connected
|
||||
self.connected = False
|
||||
if was and self._on_disconnected:
|
||||
self._on_disconnected(None, False)
|
||||
|
||||
def inject_down(self, payload: bytes) -> None:
|
||||
if self._on_down:
|
||||
self._on_down(payload)
|
||||
|
||||
def simulate_taken_over(self) -> None:
|
||||
self.connected = False
|
||||
if self._on_disconnected:
|
||||
self._on_disconnected("taken_over", True)
|
||||
|
||||
def simulate_network_drop(self) -> None:
|
||||
self.connected = False
|
||||
if self._on_disconnected:
|
||||
self._on_disconnected("network", False)
|
||||
|
||||
|
||||
class PahoTransport:
|
||||
"""paho-mqtt 2.x CallbackAPIVersion.VERSION2。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._client: Optional[PahoClient] = None
|
||||
self._on_connected: Optional[ConnHandler] = None
|
||||
self._on_disconnected: Optional[DiscHandler] = None
|
||||
self._on_down: Optional[DownHandler] = None
|
||||
self._down_topic = ""
|
||||
self._loop_started = False
|
||||
|
||||
def set_handlers(
|
||||
self,
|
||||
on_connected: ConnHandler,
|
||||
on_disconnected: DiscHandler,
|
||||
on_down: DownHandler,
|
||||
) -> None:
|
||||
self._on_connected = on_connected
|
||||
self._on_disconnected = on_disconnected
|
||||
self._on_down = on_down
|
||||
|
||||
def connect(self, params: ConnectParams) -> None:
|
||||
self.disconnect()
|
||||
client = PahoClient(
|
||||
callback_api_version=CallbackAPIVersion.VERSION2,
|
||||
client_id=params.client_id,
|
||||
protocol=PahoClient.MQTTv5,
|
||||
)
|
||||
client.username_pw_set(params.username, params.password)
|
||||
client.on_connect = self._on_connect
|
||||
client.on_disconnect = self._on_disconnect
|
||||
client.on_message = self._on_message
|
||||
self._client = client
|
||||
|
||||
url = params.url
|
||||
u = urlparse(url if "://" in url else "ws://" + url)
|
||||
use_tcp = params.use_tcp or u.scheme in ("mqtt", "mqtts")
|
||||
host = u.hostname or "localhost"
|
||||
port = u.port or (8883 if u.scheme in ("wss", "mqtts") else 443 if u.scheme == "wss" else 80)
|
||||
if use_tcp:
|
||||
if not u.port:
|
||||
port = 8883 if u.scheme == "mqtts" else 1883
|
||||
tls = u.scheme == "mqtts"
|
||||
if tls:
|
||||
client.tls_set()
|
||||
props = None
|
||||
try:
|
||||
from paho.mqtt.properties import Properties
|
||||
from paho.mqtt.packettypes import PacketTypes
|
||||
|
||||
props = Properties(PacketTypes.CONNECT)
|
||||
props.SessionExpiryInterval = params.session_expiry
|
||||
except Exception:
|
||||
props = None
|
||||
client.connect(
|
||||
host,
|
||||
port,
|
||||
keepalive=params.keep_alive,
|
||||
clean_start=params.clean_start,
|
||||
properties=props,
|
||||
)
|
||||
else:
|
||||
path = u.path or "/mqtt"
|
||||
if not path.endswith("/mqtt"):
|
||||
path = path.rstrip("/") + "/mqtt"
|
||||
if not u.port:
|
||||
port = 443 if u.scheme == "wss" else 80
|
||||
if u.scheme == "wss":
|
||||
client.tls_set()
|
||||
client.ws_set_options(path=path, headers={"Sec-WebSocket-Protocol": "mqtt"})
|
||||
props = None
|
||||
try:
|
||||
from paho.mqtt.properties import Properties
|
||||
from paho.mqtt.packettypes import PacketTypes
|
||||
|
||||
props = Properties(PacketTypes.CONNECT)
|
||||
props.SessionExpiryInterval = params.session_expiry
|
||||
except Exception:
|
||||
props = None
|
||||
client.connect(
|
||||
host,
|
||||
port,
|
||||
keepalive=params.keep_alive,
|
||||
clean_start=params.clean_start,
|
||||
properties=props,
|
||||
)
|
||||
client.loop_start()
|
||||
self._loop_started = True
|
||||
# 等待连接结果由回调驱动;超时由 Client 层处理
|
||||
|
||||
def subscribe(self, topic: str) -> None:
|
||||
self._down_topic = topic
|
||||
if self._client:
|
||||
self._client.subscribe(topic, qos=1)
|
||||
|
||||
def publish(self, topic: str, payload: bytes) -> None:
|
||||
if not self._client:
|
||||
raise RuntimeError("not connected")
|
||||
info = self._client.publish(topic, payload, qos=1)
|
||||
if info.rc != MQTT_ERR_SUCCESS:
|
||||
raise RuntimeError(f"publish failed: {info.rc}")
|
||||
|
||||
def disconnect(self) -> None:
|
||||
c = self._client
|
||||
self._client = None
|
||||
if c is not None:
|
||||
try:
|
||||
c.disconnect()
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
if self._loop_started:
|
||||
c.loop_stop()
|
||||
except Exception:
|
||||
pass
|
||||
self._loop_started = False
|
||||
|
||||
def _on_connect(self, client, userdata, flags, reason_code, properties) -> None:
|
||||
code = _reason_to_int(reason_code)
|
||||
if code == 0:
|
||||
if self._on_connected:
|
||||
self._on_connected()
|
||||
return
|
||||
stop, reason = _classify_connack(code)
|
||||
if self._on_disconnected:
|
||||
self._on_disconnected(reason, stop)
|
||||
|
||||
def _on_disconnect(self, client, userdata, flags, reason_code, properties) -> None:
|
||||
code = _reason_to_int(reason_code)
|
||||
if code in (0, None):
|
||||
if self._on_disconnected:
|
||||
self._on_disconnected(None, False)
|
||||
return
|
||||
# 0x8E = 142 Session taken over
|
||||
if code == 142:
|
||||
if self._on_disconnected:
|
||||
self._on_disconnected("taken_over", True)
|
||||
return
|
||||
stop, reason = _classify_connack(code)
|
||||
if self._on_disconnected:
|
||||
self._on_disconnected(reason, stop)
|
||||
|
||||
def _on_message(self, client, userdata, msg) -> None:
|
||||
if self._on_down:
|
||||
self._on_down(bytes(msg.payload))
|
||||
|
||||
|
||||
def _reason_to_int(reason_code) -> Optional[int]:
|
||||
if reason_code is None:
|
||||
return None
|
||||
if isinstance(reason_code, int):
|
||||
return reason_code
|
||||
if isinstance(reason_code, ReasonCode):
|
||||
return int(reason_code)
|
||||
if isinstance(reason_code, MQTTErrorCode):
|
||||
return int(reason_code)
|
||||
try:
|
||||
return int(reason_code)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _classify_connack(code: int) -> tuple[bool, str]:
|
||||
# MQTT5: 0x86=134 Bad User Name or Password, 0x87=135 Not authorized, 0x8A=138 Banned
|
||||
# 0x88=136 Server unavailable, 0x89=137 Server busy -> continue
|
||||
if code in (4, 5, 134, 135, 138): # 3.1.1 4/5 and MQTT5 auth failures
|
||||
if code in (134, 4):
|
||||
return True, "bad_credentials" # 可能是令牌,Client 层再区分
|
||||
return True, "bad_credentials"
|
||||
if code in (136, 137, 0x88, 0x89):
|
||||
return False, "busy"
|
||||
return False, "network"
|
||||
@@ -0,0 +1,182 @@
|
||||
"""公共类型与常量。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from typing import Any, Callable, Optional
|
||||
|
||||
CLIENT_NAME = "python-sdk/0.1"
|
||||
|
||||
DEFAULT_MAX_BODY = 262144
|
||||
DEFAULT_MAX_META = 4096
|
||||
DEFAULT_MAX_FRAME = 786432
|
||||
MIN_MAX_RECEIVE = 1024
|
||||
SEND_QUEUE_LIMIT = 1000
|
||||
INFLIGHT_LIMIT = 100
|
||||
DEDUP_CAPACITY = 10000
|
||||
CONNECT_TIMEOUT_S = 30.0
|
||||
BACKOFF_INITIAL_S = 1.0
|
||||
BACKOFF_MAX_S = 30.0
|
||||
BACKOFF_JITTER = 0.3
|
||||
STABLE_RESET_S = 60.0
|
||||
|
||||
SESSION_TOKEN_PREFIX = "nst_"
|
||||
|
||||
|
||||
class ConnectionState(str, Enum):
|
||||
CONNECTING = "connecting"
|
||||
ONLINE = "online"
|
||||
RECONNECTING = "reconnecting"
|
||||
OFFLINE = "offline"
|
||||
KICKED = "kicked"
|
||||
AUTH_FAILED = "auth_failed"
|
||||
|
||||
|
||||
@dataclass
|
||||
class Target:
|
||||
kind: str
|
||||
id: str
|
||||
|
||||
def to_dict(self) -> dict[str, str]:
|
||||
return {"kind": self.kind, "id": self.id}
|
||||
|
||||
|
||||
@dataclass
|
||||
class Body:
|
||||
data: str
|
||||
enc: str = "utf8"
|
||||
content_type: Optional[str] = None
|
||||
|
||||
def decoded_size(self) -> int:
|
||||
if self.enc == "base64":
|
||||
import base64
|
||||
|
||||
return len(base64.b64decode(self.data, validate=False))
|
||||
return len(self.data.encode("utf-8"))
|
||||
|
||||
def effective_content_type(self) -> str:
|
||||
if self.content_type:
|
||||
return self.content_type
|
||||
if self.enc == "base64":
|
||||
return "application/octet-stream"
|
||||
return "text/plain; charset=utf-8"
|
||||
|
||||
def to_dict(self) -> dict[str, str]:
|
||||
d = {"enc": self.enc, "data": self.data}
|
||||
ct = self.effective_content_type()
|
||||
d["content_type"] = ct
|
||||
return d
|
||||
|
||||
|
||||
@dataclass
|
||||
class SendOptions:
|
||||
send_at: Optional[float] = None # 本机语义时间(秒)或与 send_at_ms 二选一
|
||||
send_at_ms: Optional[int] = None
|
||||
delay_ms: Optional[int] = None
|
||||
keep: bool = False
|
||||
ttl_seconds: Optional[int] = None
|
||||
receipt: bool = True
|
||||
talk_password: str = ""
|
||||
content_type: Optional[str] = None
|
||||
meta: Optional[dict[str, Any]] = None
|
||||
message_id: Optional[str] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class SendResult:
|
||||
id: str
|
||||
send_at_ms: int
|
||||
state: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class RecallResult:
|
||||
result: str
|
||||
recalled: int = 0
|
||||
accepted: int = 0
|
||||
other: int = 0
|
||||
|
||||
|
||||
@dataclass
|
||||
class IncomingMessage:
|
||||
id: str
|
||||
from_id: str
|
||||
to: Target
|
||||
body: Body
|
||||
send_at_ms: int
|
||||
meta: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Receipt:
|
||||
receipt_id: str
|
||||
id: str
|
||||
endpoint_id: str
|
||||
state: str
|
||||
reason: str
|
||||
at_ms: int
|
||||
|
||||
|
||||
@dataclass
|
||||
class RevokedEvent:
|
||||
id: str
|
||||
from_id: str
|
||||
reason: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class PresenceEvent:
|
||||
id: str
|
||||
online: bool
|
||||
at_ms: int
|
||||
|
||||
|
||||
@dataclass
|
||||
class GroupEvent:
|
||||
group_id: str
|
||||
event: str
|
||||
endpoint_id: str
|
||||
at_ms: int
|
||||
|
||||
|
||||
@dataclass
|
||||
class HelloLimits:
|
||||
server_time_ms: int = 0
|
||||
server_version: str = ""
|
||||
max_body_bytes: int = DEFAULT_MAX_BODY
|
||||
max_meta_bytes: int = DEFAULT_MAX_META
|
||||
max_frame_bytes: int = DEFAULT_MAX_FRAME
|
||||
max_ttl_seconds: int = 2592000
|
||||
max_schedule_seconds: int = 31536000
|
||||
ack_timeout_seconds: int = 300
|
||||
session_token: str = ""
|
||||
|
||||
|
||||
@dataclass
|
||||
class RegisterOptions:
|
||||
id: str = ""
|
||||
login_password: str = ""
|
||||
name: str = ""
|
||||
talk_password: str = ""
|
||||
|
||||
|
||||
@dataclass
|
||||
class RegisterResult:
|
||||
id: str
|
||||
login_password: Optional[str] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class ConnectionEvent:
|
||||
state: ConnectionState
|
||||
reason: str = ""
|
||||
|
||||
|
||||
SessionHandler = Callable[[str], None]
|
||||
MessageHandler = Callable[[IncomingMessage], None]
|
||||
ReceiptHandler = Callable[[Receipt], None]
|
||||
RevokedHandler = Callable[[RevokedEvent], None]
|
||||
PresenceHandler = Callable[[PresenceEvent], None]
|
||||
GroupEventHandler = Callable[[GroupEvent], None]
|
||||
ConnectionHandler = Callable[[ConnectionEvent], None]
|
||||
@@ -0,0 +1,18 @@
|
||||
"""UUIDv7(36 字符形式),兼容 Python 3.10。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import time
|
||||
import uuid
|
||||
|
||||
|
||||
def new_uuid7() -> str:
|
||||
if hasattr(uuid, "uuid7"):
|
||||
return str(uuid.uuid7()) # type: ignore[attr-defined]
|
||||
# RFC 9562 简化实现
|
||||
ts_ms = int(time.time() * 1000) & ((1 << 48) - 1)
|
||||
rand_a = int.from_bytes(os.urandom(2), "big") & 0x0FFF
|
||||
rand_b = int.from_bytes(os.urandom(8), "big") & ((1 << 62) - 1)
|
||||
value = (ts_ms << 80) | (0x7 << 76) | (rand_a << 64) | (0b10 << 62) | rand_b
|
||||
return str(uuid.UUID(int=value))
|
||||
@@ -0,0 +1,202 @@
|
||||
"""假传输单元测试:Clean Start、去重再 ack、本地超限、令牌回调、重交不改消息号。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import threading
|
||||
import time
|
||||
import unittest
|
||||
|
||||
from nixmsg import Body, Client, ConnectionState, FakeTransport, NixMsgError, SendOptions, Target
|
||||
from nixmsg.protocol import dumps
|
||||
|
||||
|
||||
class FakeTransportTests(unittest.TestCase):
|
||||
def _connect(self, transport: FakeTransport | None = None, **kwargs) -> tuple[Client, FakeTransport]:
|
||||
tr = transport or FakeTransport()
|
||||
c = Client(transport=tr, **kwargs)
|
||||
tokens: list[str] = []
|
||||
c.on_session(lambda t: tokens.append(t))
|
||||
c.connect("ws://example.test/mqtt", "ep1", password="secret", wait=True)
|
||||
self.assertEqual(c.state, ConnectionState.ONLINE)
|
||||
return c, tr
|
||||
|
||||
def test_clean_start_and_session_expiry_every_connect(self) -> None:
|
||||
tr = FakeTransport()
|
||||
c, _ = self._connect(tr)
|
||||
self.assertGreaterEqual(len(tr.connects), 1)
|
||||
p = tr.connects[0].params
|
||||
self.assertTrue(p.clean_start)
|
||||
self.assertEqual(p.session_expiry, 0)
|
||||
self.assertIn("nix/c/ep1/down", tr.subscriptions)
|
||||
# 断开再连
|
||||
tr.simulate_network_drop()
|
||||
time.sleep(0.3)
|
||||
# 等待重连至少再记一次
|
||||
deadline = time.time() + 3
|
||||
while len(tr.connects) < 2 and time.time() < deadline:
|
||||
time.sleep(0.05)
|
||||
self.assertGreaterEqual(len(tr.connects), 2)
|
||||
for rec in tr.connects:
|
||||
self.assertTrue(rec.params.clean_start)
|
||||
self.assertEqual(rec.params.session_expiry, 0)
|
||||
c.close()
|
||||
|
||||
def test_session_token_callback(self) -> None:
|
||||
tr = FakeTransport()
|
||||
tokens: list[str] = []
|
||||
c = Client(transport=tr)
|
||||
c.on_session(lambda t: tokens.append(t))
|
||||
c.connect("ws://example.test/mqtt", "ep1", password="pw")
|
||||
self.assertEqual(tokens, ["nst_test_token"])
|
||||
self.assertEqual(c.session_token, "nst_test_token")
|
||||
c.close()
|
||||
|
||||
def test_dedup_acked_rearrival_acks_again(self) -> None:
|
||||
c, tr = self._connect()
|
||||
delivered: list[str] = []
|
||||
c.on_message(lambda m: delivered.append(m.id))
|
||||
|
||||
msg = {
|
||||
"v": 1,
|
||||
"type": "msg",
|
||||
"id": "m1",
|
||||
"from": "peer",
|
||||
"to": {"kind": "endpoint", "id": "ep1"},
|
||||
"body": {"enc": "utf8", "data": "hi"},
|
||||
"send_at_ms": 1,
|
||||
}
|
||||
tr.inject_down(dumps(msg))
|
||||
time.sleep(0.1)
|
||||
self.assertEqual(delivered, ["m1"])
|
||||
# 找 ack 帧
|
||||
acks = [json.loads(p.decode()) for _, p in tr.publishes if json.loads(p.decode()).get("type") == "ack"]
|
||||
self.assertGreaterEqual(len(acks), 1)
|
||||
|
||||
before = len(tr.publishes)
|
||||
tr.inject_down(dumps(msg)) # 已确认再到达
|
||||
time.sleep(0.1)
|
||||
self.assertEqual(delivered, ["m1"]) # 不重复交应用
|
||||
acks2 = [json.loads(p.decode()) for _, p in tr.publishes[before:] if json.loads(p.decode()).get("type") == "ack"]
|
||||
self.assertGreaterEqual(len(acks2), 1) # 再 ack
|
||||
c.close()
|
||||
|
||||
def test_local_body_too_large(self) -> None:
|
||||
c, tr = self._connect()
|
||||
big = "x" * (c.limits.max_body_bytes + 1)
|
||||
with self.assertRaises(NixMsgError) as cm:
|
||||
c.send(Target(kind="endpoint", id="ep2"), Body(data=big))
|
||||
self.assertEqual(cm.exception.code, "body_too_large")
|
||||
c.close()
|
||||
|
||||
def test_local_frame_too_large(self) -> None:
|
||||
tr = FakeTransport()
|
||||
tr.auto_hello = {
|
||||
"server_time_ms": 1_750_000_000_000,
|
||||
"server_version": "0.1.0",
|
||||
"max_body_bytes": 262144,
|
||||
"max_meta_bytes": 4096,
|
||||
"max_frame_bytes": 200, # 故意很小
|
||||
"max_ttl_seconds": 2592000,
|
||||
"max_schedule_seconds": 31536000,
|
||||
"ack_timeout_seconds": 300,
|
||||
"session_token": "nst_x",
|
||||
}
|
||||
c = Client(transport=tr)
|
||||
c.connect("ws://example.test/mqtt", "ep1", password="pw")
|
||||
with self.assertRaises(NixMsgError) as cm:
|
||||
c.send(Target(kind="endpoint", id="ep2"), Body(data="hello world " * 20))
|
||||
self.assertEqual(cm.exception.code, "frame_too_large")
|
||||
c.close()
|
||||
|
||||
def test_resend_keeps_send_at_ms(self) -> None:
|
||||
tr = FakeTransport()
|
||||
rate_hits = {"n": 0}
|
||||
|
||||
def hook(frame: dict):
|
||||
if frame.get("type") != "send":
|
||||
return None
|
||||
if rate_hits["n"] == 0:
|
||||
rate_hits["n"] += 1
|
||||
return {
|
||||
"v": 1,
|
||||
"type": "resp",
|
||||
"rid": frame["rid"],
|
||||
"ok": False,
|
||||
"error": {"code": "rate_limited", "message": "slow"},
|
||||
}
|
||||
return {
|
||||
"v": 1,
|
||||
"type": "resp",
|
||||
"rid": frame["rid"],
|
||||
"ok": True,
|
||||
"data": {"id": frame["id"], "send_at_ms": frame.get("send_at_ms", 0), "state": "scheduled"},
|
||||
}
|
||||
|
||||
tr.on_up(hook)
|
||||
tr.auto_send_ok = False
|
||||
c = Client(transport=tr)
|
||||
c.connect("ws://example.test/mqtt", "ep1", password="pw")
|
||||
fixed = 1_700_000_000_000
|
||||
result = c.send(
|
||||
Target(kind="endpoint", id="ep2"),
|
||||
Body(data="hi"),
|
||||
SendOptions(send_at_ms=fixed, message_id="fixed-id-1"),
|
||||
)
|
||||
self.assertEqual(result.id, "fixed-id-1")
|
||||
sends = [json.loads(p.decode()) for _, p in tr.publishes if json.loads(p.decode()).get("type") == "send"]
|
||||
self.assertGreaterEqual(len(sends), 2)
|
||||
for s in sends:
|
||||
self.assertEqual(s["id"], "fixed-id-1")
|
||||
self.assertEqual(s["send_at_ms"], fixed)
|
||||
c.close()
|
||||
|
||||
def test_session_invalid_stops_reconnect(self) -> None:
|
||||
tr = FakeTransport()
|
||||
# 先用密码连上
|
||||
c = Client(transport=tr)
|
||||
c.connect("ws://example.test/mqtt", "ep1", password="pw")
|
||||
# 下次连接令牌失败
|
||||
tr.next_connack_fail = "bad_credentials"
|
||||
tr.simulate_network_drop()
|
||||
deadline = time.time() + 3
|
||||
while c.state != ConnectionState.AUTH_FAILED and time.time() < deadline:
|
||||
time.sleep(0.05)
|
||||
self.assertEqual(c.state, ConnectionState.AUTH_FAILED)
|
||||
n = len(tr.connects)
|
||||
time.sleep(0.5)
|
||||
self.assertEqual(len(tr.connects), n) # 不再重连
|
||||
c.close()
|
||||
|
||||
def test_delivered_duplicate_ignored(self) -> None:
|
||||
c, tr = self._connect(auto_ack=False)
|
||||
barrier = threading.Event()
|
||||
seen = []
|
||||
|
||||
def handler(m):
|
||||
seen.append(m.id)
|
||||
barrier.wait(timeout=2)
|
||||
|
||||
c.on_message(handler)
|
||||
msg = {
|
||||
"v": 1,
|
||||
"type": "msg",
|
||||
"id": "m2",
|
||||
"from": "peer",
|
||||
"to": {"kind": "endpoint", "id": "ep1"},
|
||||
"body": {"enc": "utf8", "data": "x"},
|
||||
"send_at_ms": 1,
|
||||
}
|
||||
t = threading.Thread(target=lambda: tr.inject_down(dumps(msg)))
|
||||
t.start()
|
||||
time.sleep(0.05)
|
||||
tr.inject_down(dumps(msg)) # 已交未确认
|
||||
barrier.set()
|
||||
t.join(timeout=2)
|
||||
time.sleep(0.05)
|
||||
self.assertEqual(seen, ["m2"])
|
||||
c.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user