feat: 实现 Python 与 Java SDK 连接收发及其余接口

This commit is contained in:
Nixevol
2026-09-30 07:06:30 +08:00
parent 9650cbff76
commit d357082f3d
27 changed files with 4515 additions and 1 deletions
+43 -1
View File
@@ -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
+6
View File
@@ -0,0 +1,6 @@
target/
.idea/
*.iml
.classpath
.project
.settings/
View File
+10
View File
@@ -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.
+34
View File
@@ -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`(专有)。
+90
View File
@@ -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());
}
}
+7
View File
@@ -0,0 +1,7 @@
.venv/
__pycache__/
*.py[cod]
*.egg-info/
.pytest_cache/
dist/
build/
View File
+10
View File
@@ -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.
+24
View File
@@ -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`(专有)。
+33
View File
@@ -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"]
+50
View File
@@ -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"
+124
View File
@@ -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)
+988
View File
@@ -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)
+20
View File
@@ -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)
+58
View File
@@ -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))
+350
View File
@@ -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"
+182
View File
@@ -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]
+18
View File
@@ -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))
+202
View File
@@ -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()