diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index 831164a..3a99bc7 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -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 diff --git a/sdk/java/.gitignore b/sdk/java/.gitignore new file mode 100644 index 0000000..105499d --- /dev/null +++ b/sdk/java/.gitignore @@ -0,0 +1,6 @@ +target/ +.idea/ +*.iml +.classpath +.project +.settings/ diff --git a/sdk/java/.gitkeep b/sdk/java/.gitkeep deleted file mode 100644 index e69de29..0000000 diff --git a/sdk/java/LICENSE b/sdk/java/LICENSE new file mode 100644 index 0000000..e1c9937 --- /dev/null +++ b/sdk/java/LICENSE @@ -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. diff --git a/sdk/java/README.md b/sdk/java/README.md new file mode 100644 index 0000000..c75d282 --- /dev/null +++ b/sdk/java/README.md @@ -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`(专有)。 diff --git a/sdk/java/pom.xml b/sdk/java/pom.xml new file mode 100644 index 0000000..bc5f659 --- /dev/null +++ b/sdk/java/pom.xml @@ -0,0 +1,90 @@ + + + 4.0.0 + + asia.asio.nixmsg + nixmsg-sdk + 0.1.0 + jar + nixmsg-sdk + NixMsg Java/Android SDK + + + + Proprietary + https://git.asio.asia/nixevol/NixMsg/src/branch/main/LICENSE + repo + + + + + UTF-8 + 8 + 1.3.5 + 4.13.2 + + + + + com.hivemq + hivemq-mqtt-client + ${hivemq.mqtt.version} + + + + io.netty + netty-codec-http + 4.1.118.Final + + + com.google.code.gson + gson + 2.11.0 + + + junit + junit + ${junit.version} + test + + + + + + + org.apache.maven.plugins + maven-compiler-plugin + 3.13.0 + + 8 + + + + org.apache.maven.plugins + maven-surefire-plugin + 3.5.2 + + + org.apache.maven.plugins + maven-jar-plugin + 3.4.2 + + + + + + + central + https://repo.maven.apache.org/maven2 + + + + + + asio-gitea + https://git.asio.asia/api/packages/nixevol/maven + + + diff --git a/sdk/java/src/main/java/asia/asio/nixmsg/Client.java b/sdk/java/src/main/java/asia/asio/nixmsg/Client.java new file mode 100644 index 0000000..19e9347 --- /dev/null +++ b/sdk/java/src/main/java/asia/asio/nixmsg/Client.java @@ -0,0 +1,1223 @@ +package asia.asio.nixmsg; + +import asia.asio.nixmsg.Types.Body; +import asia.asio.nixmsg.Types.ConnectionEvent; +import asia.asio.nixmsg.Types.ConnectionState; +import asia.asio.nixmsg.Types.GroupEvent; +import asia.asio.nixmsg.Types.HelloLimits; +import asia.asio.nixmsg.Types.IncomingMessage; +import asia.asio.nixmsg.Types.PresenceEvent; +import asia.asio.nixmsg.Types.Receipt; +import asia.asio.nixmsg.Types.RegisterOptions; +import asia.asio.nixmsg.Types.RegisterResult; +import asia.asio.nixmsg.Types.RecallResult; +import asia.asio.nixmsg.Types.RevokedEvent; +import asia.asio.nixmsg.Types.SendOptions; +import asia.asio.nixmsg.Types.SendResult; +import asia.asio.nixmsg.Types.Target; + +import java.io.ByteArrayOutputStream; +import java.io.InputStream; +import java.io.OutputStream; +import java.net.HttpURLConnection; +import java.net.URL; +import java.nio.charset.StandardCharsets; +import java.util.ArrayList; +import java.util.Iterator; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicLong; +import java.util.function.Consumer; +import java.util.logging.Level; +import java.util.logging.Logger; + +/** + * NixMsg Java/Android SDK。接口返回 {@link CompletableFuture}。 + * Android 长连接由应用自行放入前台服务(最低 API 24)。 + */ +public final class Client { + private static final Logger LOG = Logger.getLogger("nixmsg"); + private static final String DELIVERED = "delivered"; + private static final String ACKED = "acked"; + + private final Transport transport; + private final boolean autoAck; + private final int maxReceiveBytes; + private final String clientName; + private final long connectTimeoutMs; + + private final Object lock = new Object(); + private final Object cbLock = new Object(); + private final AtomicLong ridSeq = new AtomicLong(); + private final Map pending = new LinkedHashMap(); + private final List sendQueue = new ArrayList(); + private final LinkedHashMap dedup = new LinkedHashMap(); + private final LinkedHashMap receiptSeen = new LinkedHashMap(); + + private volatile ConnectionState state = ConnectionState.OFFLINE; + private volatile boolean stopReconnect; + private volatile boolean closed; + private volatile boolean userClose; + private volatile boolean wantConnected; + private volatile boolean connectingWithToken; + private String url = ""; + private String endpointId = ""; + private String password; + private String sessionToken; + private boolean useTcp; + private HelloLimits limits = new HelloLimits(); + private long clockSkewMs; + private long onlineSinceMs; + private long backoffMs = Types.BACKOFF_INITIAL_MS; + private int inflightSends; + private List watchIds; + private boolean watchAll; + private volatile Throwable handshakeError; + private volatile String authReason = ""; + private final AtomicBoolean connReady = new AtomicBoolean(false); + private final Object connWait = new Object(); + private Thread worker; + private final Object wake = new Object(); + + private Consumer sessionHandler; + private Consumer messageHandler; + private Consumer receiptHandler; + private Consumer revokedHandler; + private Consumer presenceHandler; + private Consumer groupHandler; + private Consumer connectionHandler; + + public Client() { + this(new HiveMqTransport(), true, Types.DEFAULT_MAX_FRAME, Types.CLIENT_NAME, Types.CONNECT_TIMEOUT_MS); + } + + public Client(Transport transport) { + this(transport, true, Types.DEFAULT_MAX_FRAME, Types.CLIENT_NAME, Types.CONNECT_TIMEOUT_MS); + } + + public Client(Transport transport, boolean autoAck, int maxReceiveBytes, String clientName, long connectTimeoutMs) { + this.transport = transport; + this.autoAck = autoAck; + this.maxReceiveBytes = Math.max(Types.MIN_MAX_RECEIVE, maxReceiveBytes); + this.clientName = clientName; + this.connectTimeoutMs = connectTimeoutMs; + this.transport.setHandlers(this::onTransportConnected, this::onTransportDisconnected, this::onDown); + } + + public void onSession(Consumer handler) { this.sessionHandler = handler; } + public void onMessage(Consumer handler) { this.messageHandler = handler; } + public void onReceipt(Consumer handler) { this.receiptHandler = handler; } + public void onRevoked(Consumer handler) { this.revokedHandler = handler; } + public void onPresence(Consumer handler) { this.presenceHandler = handler; } + public void onGroupEvent(Consumer handler) { this.groupHandler = handler; } + public void onConnection(Consumer handler) { this.connectionHandler = handler; } + + public ConnectionState getState() { return state; } + public HelloLimits getLimits() { return limits; } + public String getSessionToken() { return sessionToken; } + public long getClockSkewMs() { return clockSkewMs; } + + public CompletableFuture connect(String url, String endpointId, String password, String sessionToken) { + return connect(url, endpointId, password, sessionToken, false); + } + + public CompletableFuture connect(String url, String endpointId, String password, String sessionToken, boolean useTcp) { + return CompletableFuture.runAsync(() -> connectSync(url, endpointId, password, sessionToken, useTcp)); + } + + public void connectSync(String url, String endpointId, String password, String sessionToken, boolean useTcp) { + if (password == null && sessionToken == null) { + throw new IllegalArgumentException("需要 password 或 sessionToken"); + } + synchronized (lock) { + if (closed) { + throw new NixMsgException("closed", "已关闭"); + } + this.url = useTcp ? url : Protocol.normalizeMqttWsUrl(url); + this.endpointId = endpointId; + this.password = password; + this.sessionToken = sessionToken; + this.useTcp = useTcp; + this.stopReconnect = false; + this.userClose = false; + this.wantConnected = true; + this.handshakeError = null; + connReady.set(false); + setState(ConnectionState.CONNECTING, ""); + if (worker == null || !worker.isAlive()) { + worker = new Thread(this::runLoop, "nixmsg-client"); + worker.setDaemon(true); + worker.start(); + } + wakeUp(); + } + long deadline = System.currentTimeMillis() + connectTimeoutMs + 5000; + synchronized (connWait) { + while (!connReady.get() && System.currentTimeMillis() < deadline) { + try { + connWait.wait(200); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + break; + } + } + } + if (handshakeError != null) { + if (handshakeError instanceof RuntimeException) { + throw (RuntimeException) handshakeError; + } + throw new NixMsgException("busy", handshakeError.getMessage()); + } + if (state != ConnectionState.ONLINE) { + if (state == ConnectionState.AUTH_FAILED) { + throw new NixMsgException(authReason.isEmpty() ? "bad_credentials" : authReason, "认证失败"); + } + if (state == ConnectionState.KICKED) { + throw new NixMsgException("taken_over", "会话被接管"); + } + throw new NixMsgException("busy", "连接未成功: " + state); + } + } + + public CompletableFuture closeAsync() { + return CompletableFuture.runAsync(this::close); + } + + public void close() { + synchronized (lock) { + userClose = true; + wantConnected = false; + stopReconnect = true; + closed = true; + failAll(new NixMsgException("closed", "已关闭")); + setState(ConnectionState.OFFLINE, ""); + } + try { + transport.disconnect(); + } catch (Exception ignored) { + } + wakeUp(); + } + + public CompletableFuture logout() { + return CompletableFuture.runAsync(() -> { + try { + request(mapOf("type", "self.logout"), true); + } catch (Exception ignored) { + } + synchronized (lock) { + stopReconnect = true; + wantConnected = false; + sessionToken = null; + failAll(new NixMsgException("auth_failed", "已退出登录")); + } + try { + transport.disconnect(); + } catch (Exception ignored) { + } + setState(ConnectionState.OFFLINE, ""); + wakeUp(); + }); + } + + public CompletableFuture send(Target to, Body body, SendOptions options) { + return CompletableFuture.supplyAsync(() -> sendSync(to, body, options)); + } + + public SendResult sendSync(Target to, Body body, SendOptions options) { + if (options == null) { + options = new SendOptions(); + } + if (options.contentType != null) { + body.contentType = options.contentType; + } + final Pending pendingReq; + synchronized (lock) { + if (closed) { + throw new NixMsgException("closed", "已关闭"); + } + if (sendQueue.size() >= Types.SEND_QUEUE_LIMIT) { + throw new NixMsgException("quota_exceeded", "发送队列已满"); + } + int maxBody = limits.maxBodyBytes > 0 ? limits.maxBodyBytes : Types.DEFAULT_MAX_BODY; + int maxMeta = limits.maxMetaBytes > 0 ? limits.maxMetaBytes : Types.DEFAULT_MAX_META; + int maxFrame = limits.maxFrameBytes > 0 ? limits.maxFrameBytes : Types.DEFAULT_MAX_FRAME; + if (body.decodedSize() > maxBody) { + throw new NixMsgException("body_too_large", "正文超限"); + } + Map meta = options.meta == null ? new LinkedHashMap() : options.meta; + byte[] metaRaw = Protocol.dumps(meta); + if (metaRaw.length > maxMeta) { + throw new NixMsgException("meta_too_large", "自定义字段超限"); + } + String msgId = options.messageId != null ? options.messageId : Uuid7.next(); + Map frame = new LinkedHashMap(); + frame.put("v", 1); + frame.put("type", "send"); + frame.put("id", msgId); + frame.put("to", to.toMap()); + frame.put("body", body.toMap()); + if (!meta.isEmpty()) { + frame.put("meta", meta); + } + if (options.sendAtMs != null) { + frame.put("send_at_ms", options.sendAtMs); + } else if (options.delayMs != null) { + frame.put("delay_ms", options.delayMs); + } + if (options.keep) { + Map offline = new LinkedHashMap(); + offline.put("keep", true); + if (options.ttlSeconds != null) { + offline.put("ttl_seconds", options.ttlSeconds); + } + frame.put("offline", offline); + } + if (!options.receipt) { + frame.put("receipt", false); + } + if (options.talkPassword != null && !options.talkPassword.isEmpty()) { + frame.put("talk_password", options.talkPassword); + } + Map probe = new LinkedHashMap(frame); + probe.put("rid", "00000000"); + if (Protocol.dumps(probe).length > maxFrame) { + throw new NixMsgException("frame_too_large", "整帧超限"); + } + pendingReq = new Pending(true, msgId); + sendQueue.add(new SendItem(msgId, frame, pendingReq)); + wakeUp(); + } + await(pendingReq.future, 3600_000L); + if (pendingReq.error != null) { + if (pendingReq.error instanceof RuntimeException) { + throw (RuntimeException) pendingReq.error; + } + throw new NixMsgException("busy", pendingReq.error.getMessage()); + } + Map data = asMap(pendingReq.response.get("data")); + return new SendResult(str(data.get("id"), pendingReq.messageId), + longVal(data.get("send_at_ms"), 0), + str(data.get("state"), "")); + } + + public CompletableFuture ack(IncomingMessage message) { + return CompletableFuture.runAsync(() -> sendAck(message.from, message.id, true)); + } + + public CompletableFuture recall(String messageId) { + return CompletableFuture.supplyAsync(() -> { + Map resp = request(mapOf("type", "recall", "id", messageId), true); + Map data = asMap(resp.get("data")); + return new RecallResult(str(data.get("result"), ""), + (int) longVal(data.get("recalled"), 0), + (int) longVal(data.get("accepted"), 0), + (int) longVal(data.get("other"), 0)); + }); + } + + public CompletableFuture> status(String messageId, String cursor, int limit) { + return CompletableFuture.supplyAsync(() -> request(mapOf( + "type", "status", "id", messageId, "cursor", cursor == null ? "" : cursor, "limit", limit), true)); + } + + public CompletableFuture> unlock(String endpointId, String talkPassword) { + return CompletableFuture.supplyAsync(() -> request(mapOf( + "type", "unlock", "endpoint_id", endpointId, "talk_password", talkPassword), true)); + } + + public CompletableFuture> presence(List ids) { + return CompletableFuture.supplyAsync(() -> { + Map f = new LinkedHashMap(); + f.put("type", "presence.get"); + f.put("ids", ids); + return request(f, true); + }); + } + + public CompletableFuture> directory(String cursor, String query, int limit) { + return CompletableFuture.supplyAsync(() -> request(mapOf( + "type", "directory.list", + "cursor", cursor == null ? "" : cursor, + "query", query == null ? "" : query, + "limit", limit), true)); + } + + public CompletableFuture> watchPresence(List ids, boolean all) { + return CompletableFuture.supplyAsync(() -> { + synchronized (lock) { + watchIds = ids == null ? null : new ArrayList(ids); + watchAll = all; + } + Map f = new LinkedHashMap(); + f.put("type", "presence.watch"); + f.put("all", all); + if (ids != null) { + f.put("ids", ids); + } + return request(f, true); + }); + } + + public CompletableFuture> getSelf() { + return CompletableFuture.supplyAsync(() -> request(mapOf("type", "self.get"), true)); + } + + public CompletableFuture> updateSelf(String name, Integer defaultDelayMs) { + return CompletableFuture.supplyAsync(() -> { + Map f = new LinkedHashMap(); + f.put("type", "self.update"); + if (name != null) { + f.put("name", name); + } + if (defaultDelayMs != null) { + f.put("default_delay_ms", defaultDelayMs); + } + return request(f, true); + }); + } + + public CompletableFuture> setTalkPassword(String talkPassword) { + return CompletableFuture.supplyAsync(() -> request(mapOf("type", "self.talk_password", "talk_password", talkPassword), true)); + } + + public CompletableFuture> changeLoginPassword(String oldPassword, String newPassword) { + return CompletableFuture.supplyAsync(() -> { + Map resp = request(mapOf( + "type", "self.login_password", + "old_password", oldPassword, + "new_password", newPassword), true); + Map data = asMap(resp.get("data")); + Object token = data.get("session_token"); + if (token != null) { + synchronized (lock) { + sessionToken = String.valueOf(token); + } + fireSession(String.valueOf(token)); + } + return resp; + }); + } + + public CompletableFuture> groupCreate(String name, List> members, String groupId) { + return CompletableFuture.supplyAsync(() -> { + Map f = new LinkedHashMap(); + f.put("type", "group.create"); + f.put("name", name); + f.put("members", members); + f.put("id", groupId == null ? "" : groupId); + return request(f, true); + }); + } + + public CompletableFuture> groupAdd(String groupId, List> members) { + return CompletableFuture.supplyAsync(() -> { + Map f = new LinkedHashMap(); + f.put("type", "group.add"); + f.put("group_id", groupId); + f.put("members", members); + return request(f, true); + }); + } + + public CompletableFuture> groupRemove(String groupId, String endpointId) { + return CompletableFuture.supplyAsync(() -> request(mapOf("type", "group.remove", "group_id", groupId, "endpoint_id", endpointId), true)); + } + + public CompletableFuture> groupLeave(String groupId) { + return CompletableFuture.supplyAsync(() -> request(mapOf("type", "group.leave", "group_id", groupId), true)); + } + + public CompletableFuture> groupTransfer(String groupId, String endpointId) { + return CompletableFuture.supplyAsync(() -> request(mapOf("type", "group.transfer", "group_id", groupId, "endpoint_id", endpointId), true)); + } + + public CompletableFuture> groupRename(String groupId, String name) { + return CompletableFuture.supplyAsync(() -> request(mapOf("type", "group.rename", "group_id", groupId, "name", name), true)); + } + + public CompletableFuture> groupDissolve(String groupId) { + return CompletableFuture.supplyAsync(() -> request(mapOf("type", "group.dissolve", "group_id", groupId), true)); + } + + public CompletableFuture> groupList(String cursor, int limit) { + return CompletableFuture.supplyAsync(() -> request(mapOf("type", "group.list", "cursor", cursor == null ? "" : cursor, "limit", limit), true)); + } + + public CompletableFuture> groupGet(String groupId, String cursor, int limit) { + return CompletableFuture.supplyAsync(() -> request(mapOf( + "type", "group.get", "group_id", groupId, "cursor", cursor == null ? "" : cursor, "limit", limit), true)); + } + + public static CompletableFuture register(String connectUrl, String registrationCode, RegisterOptions options) { + return CompletableFuture.supplyAsync(() -> registerSync(connectUrl, registrationCode, options)); + } + + public static RegisterResult registerSync(String connectUrl, String registrationCode, RegisterOptions options) { + if (options == null) { + options = new RegisterOptions(); + } + String regUrl = Protocol.registerUrlFromConnect(connectUrl); + Map body = new LinkedHashMap(); + body.put("registration_code", registrationCode); + body.put("id", options.id == null ? "" : options.id); + body.put("login_password", options.loginPassword == null ? "" : options.loginPassword); + body.put("name", options.name == null ? "" : options.name); + body.put("talk_password", options.talkPassword == null ? "" : options.talkPassword); + byte[] raw = Protocol.dumps(body); + try { + HttpURLConnection conn = (HttpURLConnection) new URL(regUrl).openConnection(); + conn.setConnectTimeout(30000); + conn.setReadTimeout(30000); + conn.setRequestMethod("POST"); + conn.setDoOutput(true); + conn.setRequestProperty("Content-Type", "application/json"); + OutputStream os = conn.getOutputStream(); + os.write(raw); + os.close(); + int code = conn.getResponseCode(); + InputStream in = code >= 400 ? conn.getErrorStream() : conn.getInputStream(); + byte[] respBytes = readAll(in); + Map data = Protocol.loads(respBytes); + if (!Boolean.TRUE.equals(data.get("ok"))) { + Map err = asMap(data.get("error")); + throw new NixMsgException(str(err.get("code"), "bad_request"), str(err.get("message"), "")); + } + Map d = asMap(data.get("data")); + Object lp = d.get("login_password"); + return new RegisterResult(str(d.get("id"), ""), lp == null ? null : String.valueOf(lp)); + } catch (NixMsgException e) { + throw e; + } catch (Exception e) { + throw new NixMsgException("busy", e.getMessage()); + } + } + + // ---- loop ---- + private void runLoop() { + while (true) { + boolean want; + ConnectionState st; + synchronized (lock) { + if (closed && !wantConnected) { + return; + } + want = wantConnected && !stopReconnect; + st = state; + } + if (!want) { + waitWake(500); + continue; + } + if (st == ConnectionState.ONLINE) { + pumpSends(); + if (onlineSinceMs > 0 && System.currentTimeMillis() - onlineSinceMs >= Types.STABLE_RESET_MS) { + backoffMs = Types.BACKOFF_INITIAL_MS; + } + waitWake(200); + continue; + } + try { + attemptConnect(); + } catch (Exception e) { + LOG.log(Level.FINE, "connect attempt failed", e); + } + long delay; + synchronized (lock) { + if (state == ConnectionState.ONLINE) { + continue; + } + if (stopReconnect || !wantConnected) { + continue; + } + delay = jitter(backoffMs); + backoffMs = Math.min(Types.BACKOFF_MAX_MS, backoffMs * 2); + setState(ConnectionState.RECONNECTING, ""); + } + try { + Thread.sleep(delay); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + } + } + + private void attemptConnect() { + String u; + String eid; + String cred; + boolean usingToken; + boolean tcp; + long timeout; + synchronized (lock) { + if (stopReconnect || !wantConnected) { + return; + } + setState(state == ConnectionState.OFFLINE ? ConnectionState.CONNECTING : ConnectionState.RECONNECTING, ""); + handshakeError = null; + u = url; + eid = endpointId; + if (sessionToken != null && !sessionToken.isEmpty()) { + cred = sessionToken; + usingToken = true; + } else { + cred = password == null ? "" : password; + usingToken = false; + } + connectingWithToken = usingToken; + tcp = useTcp; + timeout = connectTimeoutMs; + connReady.set(false); + } + transport.connect(new Transport.ConnectParams(u, eid, eid, cred, true, 0, tcp, timeout)); + long deadline = System.currentTimeMillis() + timeout; + synchronized (connWait) { + while (!connReady.get() && System.currentTimeMillis() < deadline) { + try { + connWait.wait(100); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + break; + } + } + } + if (!connReady.get() || state != ConnectionState.ONLINE) { + try { + transport.disconnect(); + } catch (Exception ignored) { + } + if (handshakeError == null && state != ConnectionState.AUTH_FAILED && state != ConnectionState.KICKED) { + handshakeError = new NixMsgException("busy", "连接超时"); + } + signalConn(); + } + } + + private void onTransportConnected() { + try { + String topic = Protocol.downTopic(endpointId); + transport.subscribe(topic); + String rid = nextRid(); + Map hello = new LinkedHashMap(); + hello.put("v", 1); + hello.put("type", "hello"); + hello.put("rid", rid); + hello.put("max_receive_bytes", maxReceiveBytes); + hello.put("client", clientName); + long t0 = System.currentTimeMillis(); + Pending p = new Pending(false, ""); + synchronized (lock) { + pending.put(rid, p); + } + transport.publish(Protocol.upTopic(endpointId), Protocol.dumps(hello)); + await(p.future, connectTimeoutMs); + if (p.error != null) { + throw p.error instanceof RuntimeException ? (RuntimeException) p.error : new NixMsgException("busy", p.error.getMessage()); + } + if (p.response == null || !Boolean.TRUE.equals(p.response.get("ok"))) { + Map err = p.response == null ? new LinkedHashMap() : asMap(p.response.get("error")); + throw new NixMsgException(str(err.get("code"), "bad_request"), str(err.get("message"), "")); + } + long t1 = System.currentTimeMillis(); + Map data = asMap(p.response.get("data")); + HelloLimits lim = new HelloLimits(); + lim.serverTimeMs = longVal(data.get("server_time_ms"), 0); + lim.serverVersion = str(data.get("server_version"), ""); + lim.maxBodyBytes = (int) longVal(data.get("max_body_bytes"), Types.DEFAULT_MAX_BODY); + lim.maxMetaBytes = (int) longVal(data.get("max_meta_bytes"), Types.DEFAULT_MAX_META); + lim.maxFrameBytes = (int) longVal(data.get("max_frame_bytes"), Types.DEFAULT_MAX_FRAME); + lim.maxTtlSeconds = longVal(data.get("max_ttl_seconds"), 2592000); + lim.maxScheduleSeconds = longVal(data.get("max_schedule_seconds"), 31536000); + lim.ackTimeoutSeconds = longVal(data.get("ack_timeout_seconds"), 300); + lim.sessionToken = str(data.get("session_token"), ""); + long skew = lim.serverTimeMs - ((t0 + t1) / 2); + synchronized (lock) { + limits = lim; + clockSkewMs = skew; + onlineSinceMs = System.currentTimeMillis(); + handshakeError = null; + setState(ConnectionState.ONLINE, ""); + } + if (lim.sessionToken != null && !lim.sessionToken.isEmpty()) { + synchronized (lock) { + sessionToken = lim.sessionToken; + } + fireSession(lim.sessionToken); + } + if (watchAll || watchIds != null) { + try { + Map f = new LinkedHashMap(); + f.put("type", "presence.watch"); + f.put("all", watchAll); + if (watchIds != null) { + f.put("ids", watchIds); + } + request(f, false); + } catch (Exception ignored) { + } + } + signalConn(); + wakeUp(); + } catch (Exception e) { + handshakeError = e; + try { + transport.disconnect(); + } catch (Exception ignored) { + } + signalConn(); + } + } + + private void onTransportDisconnected(String reason, boolean stop) { + synchronized (lock) { + boolean wasOnline = state == ConnectionState.ONLINE; + boolean usingToken = connectingWithToken || (sessionToken != null && password == null); + if ("taken_over".equals(reason)) { + stopReconnect = true; + wantConnected = false; + failAll(new NixMsgException("taken_over", "会话被接管")); + setState(ConnectionState.KICKED, "taken_over"); + signalConn(); + return; + } + if (stop || "session_invalid".equals(reason) || "bad_credentials".equals(reason) || "banned".equals(reason)) { + String ar = reason == null ? "bad_credentials" : reason; + if (("bad_credentials".equals(ar) || "banned".equals(ar)) && usingToken) { + ar = "session_invalid"; + } + if ("banned".equals(ar)) { + ar = usingToken ? "session_invalid" : "bad_credentials"; + } + authReason = ar; + stopReconnect = true; + wantConnected = false; + NixMsgException err = new NixMsgException(ar, "认证失败"); + failAll(err); + handshakeError = err; + setState(ConnectionState.AUTH_FAILED, ar); + signalConn(); + return; + } + if (userClose || closed) { + setState(ConnectionState.OFFLINE, ""); + signalConn(); + return; + } + if (wasOnline || state == ConnectionState.CONNECTING || state == ConnectionState.RECONNECTING) { + setState(ConnectionState.RECONNECTING, reason == null ? "network" : reason); + } + signalConn(); + wakeUp(); + } + } + + private void onDown(byte[] payload) { + Map frame; + try { + frame = Protocol.loads(payload); + } catch (Exception e) { + return; + } + String type = str(frame.get("type"), ""); + if ("resp".equals(type)) { + String rid = str(frame.get("rid"), ""); + Pending p; + synchronized (lock) { + p = pending.remove(rid); + } + if (p != null) { + if (p.isSend) { + synchronized (lock) { + inflightSends = Math.max(0, inflightSends - 1); + } + if (!Boolean.TRUE.equals(frame.get("ok"))) { + Map err = asMap(frame.get("error")); + if ("rate_limited".equals(str(err.get("code"), ""))) { + synchronized (lock) { + p.rid = ""; + p.response = null; + p.error = null; + // 保留原 future,重交成功后再 complete + } + wakeUp(); + return; + } + p.error = new NixMsgException(str(err.get("code"), "bad_request"), str(err.get("message"), "")); + } + p.response = frame; + synchronized (lock) { + Iterator it = sendQueue.iterator(); + while (it.hasNext()) { + if (it.next().pending == p) { + it.remove(); + break; + } + } + } + p.future.complete(null); + wakeUp(); + } else { + p.response = frame; + p.future.complete(null); + } + } + return; + } + if ("msg".equals(type)) { + handleMsg(frame); + return; + } + if ("receipt".equals(type)) { + handleReceipt(frame); + return; + } + if ("revoked".equals(type)) { + handleRevoked(frame); + return; + } + if ("presence".equals(type)) { + firePresence(new PresenceEvent(str(frame.get("id"), ""), Boolean.TRUE.equals(frame.get("online")), longVal(frame.get("at_ms"), 0))); + return; + } + if ("group_event".equals(type)) { + fireGroup(new GroupEvent(str(frame.get("group_id"), ""), str(frame.get("event"), ""), + str(frame.get("endpoint_id"), ""), longVal(frame.get("at_ms"), 0))); + return; + } + if ("fatal".equals(type)) { + String reason = str(frame.get("reason"), "protocol"); + synchronized (lock) { + stopReconnect = true; + wantConnected = false; + failAll(new NixMsgException("fatal", reason)); + } + setState(ConnectionState.AUTH_FAILED, reason); + try { + transport.disconnect(); + } catch (Exception ignored) { + } + } + } + + private void handleMsg(Map frame) { + String mid = str(frame.get("id"), ""); + String from = str(frame.get("from"), ""); + String key = from + "\0" + mid; + String st; + synchronized (lock) { + st = dedup.get(key); + if (DELIVERED.equals(st)) { + return; + } + } + if (ACKED.equals(st)) { + sendAck(from, mid, true); + return; + } + Map toRaw = asMap(frame.get("to")); + Map bodyRaw = asMap(frame.get("body")); + IncomingMessage msg = new IncomingMessage( + mid, + from, + new Target(str(toRaw.get("kind"), "endpoint"), str(toRaw.get("id"), "")), + new Body(str(bodyRaw.get("enc"), "utf8"), str(bodyRaw.get("data"), ""), + bodyRaw.get("content_type") == null ? null : String.valueOf(bodyRaw.get("content_type"))), + longVal(frame.get("send_at_ms"), 0), + asMap(frame.get("meta"))); + synchronized (lock) { + dedupPut(key, DELIVERED); + } + if (messageHandler == null) { + if (autoAck) { + sendAck(from, mid, true); + } + return; + } + try { + synchronized (cbLock) { + messageHandler.accept(msg); + } + } catch (Exception e) { + LOG.log(Level.SEVERE, "onMessage 回调错误,等待重推", e); + synchronized (lock) { + dedup.remove(key); + } + return; + } + if (autoAck) { + sendAck(from, mid, true); + } + } + + private void handleReceipt(Map frame) { + String rid = str(frame.get("receipt_id"), ""); + synchronized (lock) { + if (receiptSeen.containsKey(rid)) { + return; + } + receiptSeen.put(rid, Boolean.TRUE); + while (receiptSeen.size() > Types.DEDUP_CAPACITY) { + String first = receiptSeen.keySet().iterator().next(); + receiptSeen.remove(first); + } + } + fireReceipt(new Receipt(rid, str(frame.get("id"), ""), str(frame.get("endpoint_id"), ""), + str(frame.get("state"), ""), str(frame.get("reason"), ""), longVal(frame.get("at_ms"), 0))); + try { + request(mapOf("type", "receipt_ack", "receipt_id", rid), false); + } catch (Exception ignored) { + } + } + + private void handleRevoked(Map frame) { + String mid = str(frame.get("id"), ""); + String from = str(frame.get("from"), ""); + String key = from + "\0" + mid; + synchronized (lock) { + String st = dedup.get(key); + if (ACKED.equals(st)) { + return; + } + if (st == null) { + return; + } + dedup.remove(key); + } + fireRevoked(new RevokedEvent(mid, from, str(frame.get("reason"), ""))); + } + + private void sendAck(String from, String messageId, boolean markAcked) { + try { + Map resp = request(mapOf("type", "ack", "from", from, "id", messageId), true); + Map data = asMap(resp.get("data")); + String result = str(data.get("result"), "accepted"); + String key = from + "\0" + messageId; + if (!"accepted".equals(result)) { + synchronized (lock) { + dedup.remove(key); + } + fireRevoked(new RevokedEvent(messageId, from, result)); + return; + } + if (markAcked) { + synchronized (lock) { + dedupPut(key, ACKED); + } + } + } catch (Exception ignored) { + } + } + + private void pumpSends() { + while (true) { + SendItem item; + String rid; + String topic; + byte[] raw; + synchronized (lock) { + if (state != ConnectionState.ONLINE) { + return; + } + if (inflightSends >= Types.INFLIGHT_LIMIT) { + return; + } + item = null; + for (SendItem it : sendQueue) { + if (it.pending.rid.isEmpty() && it.pending.response == null && it.pending.error == null) { + item = it; + break; + } + } + if (item == null) { + return; + } + rid = nextRidLocked(); + item.pending.rid = rid; + Map frame = new LinkedHashMap(item.frame); + frame.put("rid", rid); + pending.put(rid, item.pending); + inflightSends++; + topic = Protocol.upTopic(endpointId); + raw = Protocol.dumps(frame); + } + try { + transport.publish(topic, raw); + } catch (Exception e) { + synchronized (lock) { + pending.remove(rid); + item.pending.rid = ""; + inflightSends = Math.max(0, inflightSends - 1); + item.pending.error = e; + item.pending.future.complete(null); + sendQueue.remove(item); + } + } + } + } + + private Map request(Map frame, boolean wait) { + String rid; + Pending p; + String topic; + synchronized (lock) { + if (closed) { + throw new NixMsgException("closed", "已关闭"); + } + if (state != ConnectionState.ONLINE) { + throw new NixMsgException("not_connected", "未连接"); + } + rid = nextRidLocked(); + Map body = new LinkedHashMap(frame); + body.put("v", 1); + body.put("rid", rid); + p = new Pending(false, ""); + p.rid = rid; + pending.put(rid, p); + topic = Protocol.upTopic(endpointId); + transport.publish(topic, Protocol.dumps(body)); + } + await(p.future, wait ? 60_000L : 30_000L); + if (p.error != null) { + if (p.error instanceof RuntimeException) { + throw (RuntimeException) p.error; + } + throw new NixMsgException("busy", p.error.getMessage()); + } + if (p.response == null) { + if (!wait) { + return new LinkedHashMap(); + } + throw new NixMsgException("busy", "请求超时"); + } + if (!Boolean.TRUE.equals(p.response.get("ok"))) { + Map err = asMap(p.response.get("error")); + throw new NixMsgException(str(err.get("code"), "bad_request"), str(err.get("message"), "")); + } + return p.response; + } + + private String nextRid() { + synchronized (lock) { + return nextRidLocked(); + } + } + + private String nextRidLocked() { + return String.valueOf(ridSeq.incrementAndGet()); + } + + private void dedupPut(String key, String st) { + dedup.remove(key); + dedup.put(key, st); + while (dedup.size() > Types.DEDUP_CAPACITY) { + String first = dedup.keySet().iterator().next(); + dedup.remove(first); + } + } + + private void failAll(Throwable err) { + for (Pending p : new ArrayList(pending.values())) { + p.error = err; + p.future.complete(null); + } + pending.clear(); + for (SendItem it : new ArrayList(sendQueue)) { + it.pending.error = err; + it.pending.future.complete(null); + } + sendQueue.clear(); + inflightSends = 0; + } + + private void setState(ConnectionState st, String reason) { + state = st; + Consumer h = connectionHandler; + if (h != null) { + try { + synchronized (cbLock) { + h.accept(new ConnectionEvent(st, reason)); + } + } catch (Exception e) { + LOG.log(Level.SEVERE, "onConnection", e); + } + } + } + + private void fireSession(String token) { + Consumer h = sessionHandler; + if (h != null) { + try { + synchronized (cbLock) { + h.accept(token); + } + } catch (Exception e) { + LOG.log(Level.SEVERE, "onSession", e); + } + } + } + + private void fireReceipt(Receipt r) { + Consumer h = receiptHandler; + if (h != null) { + try { + synchronized (cbLock) { + h.accept(r); + } + } catch (Exception e) { + LOG.log(Level.SEVERE, "onReceipt", e); + } + } + } + + private void fireRevoked(RevokedEvent ev) { + Consumer h = revokedHandler; + if (h != null) { + try { + synchronized (cbLock) { + h.accept(ev); + } + } catch (Exception e) { + LOG.log(Level.SEVERE, "onRevoked", e); + } + } + } + + private void firePresence(PresenceEvent ev) { + Consumer h = presenceHandler; + if (h != null) { + try { + synchronized (cbLock) { + h.accept(ev); + } + } catch (Exception e) { + LOG.log(Level.SEVERE, "onPresence", e); + } + } + } + + private void fireGroup(GroupEvent ev) { + Consumer h = groupHandler; + if (h != null) { + try { + synchronized (cbLock) { + h.accept(ev); + } + } catch (Exception e) { + LOG.log(Level.SEVERE, "onGroupEvent", e); + } + } + } + + private void signalConn() { + connReady.set(true); + synchronized (connWait) { + connWait.notifyAll(); + } + } + + private void wakeUp() { + synchronized (wake) { + wake.notifyAll(); + } + } + + private void waitWake(long ms) { + synchronized (wake) { + try { + wake.wait(ms); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + } + } + + private static long jitter(long base) { + double j = 1.0 + (Math.random() * 2 - 1) * Types.BACKOFF_JITTER; + return Math.max(0L, (long) (base * j)); + } + + private static void await(CompletableFuture f, long timeoutMs) { + try { + f.get(timeoutMs, TimeUnit.MILLISECONDS); + } catch (Exception e) { + throw new NixMsgException("busy", "等待超时"); + } + } + + private static Map mapOf(Object... kv) { + Map m = new LinkedHashMap(); + for (int i = 0; i + 1 < kv.length; i += 2) { + m.put(String.valueOf(kv[i]), kv[i + 1]); + } + return m; + } + + @SuppressWarnings("unchecked") + private static Map asMap(Object o) { + if (o instanceof Map) { + return (Map) o; + } + return new LinkedHashMap(); + } + + private static String str(Object o, String def) { + return o == null ? def : String.valueOf(o); + } + + private static long longVal(Object o, long def) { + if (o instanceof Number) { + return ((Number) o).longValue(); + } + if (o == null) { + return def; + } + try { + return Long.parseLong(String.valueOf(o)); + } catch (Exception e) { + return def; + } + } + + private static byte[] readAll(InputStream in) throws Exception { + if (in == null) { + return new byte[0]; + } + ByteArrayOutputStream bos = new ByteArrayOutputStream(); + byte[] buf = new byte[4096]; + int n; + while ((n = in.read(buf)) >= 0) { + bos.write(buf, 0, n); + } + return bos.toByteArray(); + } + + private static final class Pending { + final boolean isSend; + final String messageId; + String rid = ""; + Map response; + Throwable error; + CompletableFuture future = new CompletableFuture(); + + Pending(boolean isSend, String messageId) { + this.isSend = isSend; + this.messageId = messageId; + } + } + + private static final class SendItem { + final String messageId; + final Map frame; + final Pending pending; + + SendItem(String messageId, Map frame, Pending pending) { + this.messageId = messageId; + this.frame = frame; + this.pending = pending; + } + } +} diff --git a/sdk/java/src/main/java/asia/asio/nixmsg/NixMsgException.java b/sdk/java/src/main/java/asia/asio/nixmsg/NixMsgException.java new file mode 100644 index 0000000..8c07fa7 --- /dev/null +++ b/sdk/java/src/main/java/asia/asio/nixmsg/NixMsgException.java @@ -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; + } +} diff --git a/sdk/java/src/main/java/asia/asio/nixmsg/Protocol.java b/sdk/java/src/main/java/asia/asio/nixmsg/Protocol.java new file mode 100644 index 0000000..79411bc --- /dev/null +++ b/sdk/java/src/main/java/asia/asio/nixmsg/Protocol.java @@ -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 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", "无效连接地址"); + } + } +} diff --git a/sdk/java/src/main/java/asia/asio/nixmsg/Transport.java b/sdk/java/src/main/java/asia/asio/nixmsg/Transport.java new file mode 100644 index 0000000..ed30519 --- /dev/null +++ b/sdk/java/src/main/java/asia/asio/nixmsg/Transport.java @@ -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 onDisconnected, Consumer 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 connects = new ArrayList(); + final List publishes = new CopyOnWriteArrayList(); + final List subscriptions = new ArrayList(); + volatile String nextConnackFail; + Map autoHello; + boolean autoSendOk = true; + final List, Map>> upHandlers = + new CopyOnWriteArrayList, Map>>(); + + private Runnable onConnected; + private BiConsumer onDisconnected; + private Consumer onDown; + private volatile boolean connected; + + FakeTransport() { + autoHello = new LinkedHashMap(); + 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> h) { + upHandlers.add(h); + } + + @Override + public void setHandlers(Runnable onConnected, BiConsumer onDisconnected, Consumer 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 frame = Protocol.loads(payload); + for (Function, Map> h : upHandlers) { + Map 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 resp = new LinkedHashMap(); + 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 data = new LinkedHashMap(); + 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 resp = new LinkedHashMap(); + 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 data = new LinkedHashMap(); + data.put("result", "accepted"); + Map resp = new LinkedHashMap(); + 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 resp = new LinkedHashMap(); + resp.put("v", 1); + resp.put("type", "resp"); + resp.put("rid", rid); + resp.put("ok", true); + resp.put("data", new LinkedHashMap()); + 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 onDisconnected; + private Consumer onDown; + private volatile Mqtt5AsyncClient client; + + @Override + public void setHandlers(Runnable onConnected, BiConsumer onDisconnected, Consumer 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 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 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 h = onDisconnected; + if (h != null) { + h.accept(reason, stop); + } + } + } catch (Exception e) { + BiConsumer 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"; + } +} diff --git a/sdk/java/src/main/java/asia/asio/nixmsg/Types.java b/sdk/java/src/main/java/asia/asio/nixmsg/Types.java new file mode 100644 index 0000000..7047603 --- /dev/null +++ b/sdk/java/src/main/java/asia/asio/nixmsg/Types.java @@ -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 toMap() { + Map m = new LinkedHashMap(); + 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 toMap() { + Map m = new LinkedHashMap(); + 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 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 meta; + + public IncomingMessage(String id, String from, Target to, Body body, long sendAtMs, Map meta) { + this.id = id; + this.from = from; + this.to = to; + this.body = body; + this.sendAtMs = sendAtMs; + this.meta = meta == null ? Collections.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; + } + } +} diff --git a/sdk/java/src/main/java/asia/asio/nixmsg/Uuid7.java b/sdk/java/src/main/java/asia/asio/nixmsg/Uuid7.java new file mode 100644 index 0000000..98a05c8 --- /dev/null +++ b/sdk/java/src/main/java/asia/asio/nixmsg/Uuid7.java @@ -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(); + } +} diff --git a/sdk/java/src/test/java/asia/asio/nixmsg/ClientTest.java b/sdk/java/src/test/java/asia/asio/nixmsg/ClientTest.java new file mode 100644 index 0000000..75b47f4 --- /dev/null +++ b/sdk/java/src/test/java/asia/asio/nixmsg/ClientTest.java @@ -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 tokens = new ArrayList(); + client.onSession(new java.util.function.Consumer() { + @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 delivered = new ArrayList(); + client.onMessage(new java.util.function.Consumer() { + @Override + public void accept(IncomingMessage m) { + delivered.add(m.id); + } + }); + Map msg = new LinkedHashMap(); + msg.put("v", 1); + msg.put("type", "msg"); + msg.put("id", "m1"); + msg.put("from", "peer"); + Map to = new LinkedHashMap(); + to.put("kind", "endpoint"); + to.put("id", "ep1"); + msg.put("to", to); + Map body = new LinkedHashMap(); + 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 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>() { + @Override + public Map apply(Map frame) { + if (!"send".equals(String.valueOf(frame.get("type")))) { + return null; + } + if (rateHits.getAndIncrement() == 0) { + Map err = new LinkedHashMap(); + err.put("code", "rate_limited"); + err.put("message", "slow"); + Map resp = new LinkedHashMap(); + 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 data = new LinkedHashMap(); + data.put("id", frame.get("id")); + data.put("send_at_ms", frame.get("send_at_ms")); + data.put("state", "scheduled"); + Map resp = new LinkedHashMap(); + 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 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 seen = new ArrayList(); + client.onMessage(new java.util.function.Consumer() { + @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 msg = new LinkedHashMap(); + msg.put("v", 1); + msg.put("type", "msg"); + msg.put("id", "m2"); + msg.put("from", "peer"); + Map to = new LinkedHashMap(); + to.put("kind", "endpoint"); + to.put("id", "ep1"); + msg.put("to", to); + Map body = new LinkedHashMap(); + 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()); + } +} diff --git a/sdk/python/.gitignore b/sdk/python/.gitignore new file mode 100644 index 0000000..4a787be --- /dev/null +++ b/sdk/python/.gitignore @@ -0,0 +1,7 @@ +.venv/ +__pycache__/ +*.py[cod] +*.egg-info/ +.pytest_cache/ +dist/ +build/ diff --git a/sdk/python/.gitkeep b/sdk/python/.gitkeep deleted file mode 100644 index e69de29..0000000 diff --git a/sdk/python/LICENSE b/sdk/python/LICENSE new file mode 100644 index 0000000..e1c9937 --- /dev/null +++ b/sdk/python/LICENSE @@ -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. diff --git a/sdk/python/README.md b/sdk/python/README.md new file mode 100644 index 0000000..cb65249 --- /dev/null +++ b/sdk/python/README.md @@ -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`(专有)。 diff --git a/sdk/python/pyproject.toml b/sdk/python/pyproject.toml new file mode 100644 index 0000000..5df902e --- /dev/null +++ b/sdk/python/pyproject.toml @@ -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"] diff --git a/sdk/python/src/nixmsg/__init__.py b/sdk/python/src/nixmsg/__init__.py new file mode 100644 index 0000000..d138601 --- /dev/null +++ b/sdk/python/src/nixmsg/__init__.py @@ -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" diff --git a/sdk/python/src/nixmsg/async_client.py b/sdk/python/src/nixmsg/async_client.py new file mode 100644 index 0000000..a388eb4 --- /dev/null +++ b/sdk/python/src/nixmsg/async_client.py @@ -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) diff --git a/sdk/python/src/nixmsg/client.py b/sdk/python/src/nixmsg/client.py new file mode 100644 index 0000000..eb7f522 --- /dev/null +++ b/sdk/python/src/nixmsg/client.py @@ -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) diff --git a/sdk/python/src/nixmsg/errors.py b/sdk/python/src/nixmsg/errors.py new file mode 100644 index 0000000..75f9f4c --- /dev/null +++ b/sdk/python/src/nixmsg/errors.py @@ -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) diff --git a/sdk/python/src/nixmsg/protocol.py b/sdk/python/src/nixmsg/protocol.py new file mode 100644 index 0000000..4b397df --- /dev/null +++ b/sdk/python/src/nixmsg/protocol.py @@ -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)) diff --git a/sdk/python/src/nixmsg/transport.py b/sdk/python/src/nixmsg/transport.py new file mode 100644 index 0000000..b92d7cf --- /dev/null +++ b/sdk/python/src/nixmsg/transport.py @@ -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" diff --git a/sdk/python/src/nixmsg/types.py b/sdk/python/src/nixmsg/types.py new file mode 100644 index 0000000..04631e5 --- /dev/null +++ b/sdk/python/src/nixmsg/types.py @@ -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] diff --git a/sdk/python/src/nixmsg/uuid7.py b/sdk/python/src/nixmsg/uuid7.py new file mode 100644 index 0000000..0de3735 --- /dev/null +++ b/sdk/python/src/nixmsg/uuid7.py @@ -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)) diff --git a/sdk/python/tests/test_client.py b/sdk/python/tests/test_client.py new file mode 100644 index 0000000..706cf4d --- /dev/null +++ b/sdk/python/tests/test_client.py @@ -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()