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

This commit is contained in:
Nixevol
2026-09-30 07:06:30 +08:00
parent 9650cbff76
commit d357082f3d
27 changed files with 4515 additions and 1 deletions
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,15 @@
package asia.asio.nixmsg;
/** SDK 错误,code 与协议错误码或连接原因对齐。 */
public class NixMsgException extends RuntimeException {
private final String code;
public NixMsgException(String code, String message) {
super(message == null || message.isEmpty() ? code : message);
this.code = code;
}
public String getCode() {
return code;
}
}
@@ -0,0 +1,87 @@
package asia.asio.nixmsg;
import com.google.gson.Gson;
import com.google.gson.GsonBuilder;
import com.google.gson.ToNumberPolicy;
import java.net.URI;
import java.nio.charset.StandardCharsets;
import java.util.Map;
final class Protocol {
static final Gson GSON = new GsonBuilder()
.disableHtmlEscaping()
.setObjectToNumberStrategy(ToNumberPolicy.LONG_OR_DOUBLE)
.create();
private Protocol() {}
static byte[] dumps(Object obj) {
return GSON.toJson(obj).getBytes(StandardCharsets.UTF_8);
}
@SuppressWarnings("unchecked")
static Map<String, Object> loads(byte[] data) {
return GSON.fromJson(new String(data, StandardCharsets.UTF_8), Map.class);
}
static String upTopic(String endpointId) {
return "nix/c/" + endpointId + "/up";
}
static String downTopic(String endpointId) {
return "nix/c/" + endpointId + "/down";
}
static String normalizeMqttWsUrl(String url) {
String raw = url.trim();
if (!raw.contains("://")) {
raw = "ws://" + raw;
}
URI u = URI.create(raw);
String path = u.getPath() == null ? "" : u.getPath();
if (!path.endsWith("/mqtt")) {
if (path.endsWith("/")) {
path = path + "mqtt";
} else if (path.isEmpty()) {
path = "/mqtt";
} else {
path = path + "/mqtt";
}
}
try {
return new URI(u.getScheme(), u.getUserInfo(), u.getHost(), u.getPort(), path, u.getQuery(), u.getFragment()).toString();
} catch (Exception e) {
return raw;
}
}
/** 从连接地址推出注册 HTTP 地址。 */
static String registerUrlFromConnect(String connectUrl) {
String raw = connectUrl.trim();
if (!raw.contains("://")) {
raw = "ws://" + raw;
}
URI u = URI.create(raw);
String scheme = u.getScheme() == null ? "ws" : u.getScheme().toLowerCase();
String httpScheme;
if ("wss".equals(scheme) || "mqtts".equals(scheme) || "https".equals(scheme)) {
httpScheme = "https";
} else {
httpScheme = "http";
}
String path = u.getPath() == null ? "" : u.getPath();
if (path.endsWith("/mqtt")) {
path = path.substring(0, path.length() - "/mqtt".length());
}
if (path.endsWith("/")) {
path = path.substring(0, path.length() - 1);
}
path = path + "/api/client/register";
try {
return new URI(httpScheme, u.getUserInfo(), u.getHost(), u.getPort(), path, null, null).toString();
} catch (Exception e) {
throw new NixMsgException("bad_request", "无效连接地址");
}
}
}
@@ -0,0 +1,406 @@
package asia.asio.nixmsg;
import com.hivemq.client.mqtt.MqttClient;
import com.hivemq.client.mqtt.MqttGlobalPublishFilter;
import com.hivemq.client.mqtt.datatypes.MqttQos;
import com.hivemq.client.mqtt.lifecycle.MqttDisconnectSource;
import com.hivemq.client.mqtt.mqtt5.Mqtt5AsyncClient;
import com.hivemq.client.mqtt.mqtt5.Mqtt5ClientBuilder;
import com.hivemq.client.mqtt.mqtt5.message.connect.connack.Mqtt5ConnAck;
import com.hivemq.client.mqtt.mqtt5.message.connect.connack.Mqtt5ConnAckReasonCode;
import com.hivemq.client.mqtt.mqtt5.message.disconnect.Mqtt5Disconnect;
import com.hivemq.client.mqtt.mqtt5.message.disconnect.Mqtt5DisconnectReasonCode;
import java.net.URI;
import java.nio.charset.StandardCharsets;
import java.util.ArrayList;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import java.util.concurrent.CopyOnWriteArrayList;
import java.util.concurrent.TimeUnit;
import java.util.function.BiConsumer;
import java.util.function.Consumer;
import java.util.function.Function;
/** MQTT 传输抽象。 */
public interface Transport {
void setHandlers(Runnable onConnected, BiConsumer<String, Boolean> onDisconnected, Consumer<byte[]> onDown);
void connect(ConnectParams params);
void subscribe(String topic);
void publish(String topic, byte[] payload);
void disconnect();
final class ConnectParams {
public final String url;
public final String clientId;
public final String username;
public final String password;
public final boolean cleanStart;
public final int sessionExpiry;
public final boolean useTcp;
public final long timeoutMs;
public ConnectParams(String url, String clientId, String username, String password,
boolean cleanStart, int sessionExpiry, boolean useTcp, long timeoutMs) {
this.url = url;
this.clientId = clientId;
this.username = username;
this.password = password;
this.cleanStart = cleanStart;
this.sessionExpiry = sessionExpiry;
this.useTcp = useTcp;
this.timeoutMs = timeoutMs;
}
}
}
/** 单元测试假传输。 */
final class FakeTransport implements Transport {
final List<Transport.ConnectParams> connects = new ArrayList<Transport.ConnectParams>();
final List<byte[]> publishes = new CopyOnWriteArrayList<byte[]>();
final List<String> subscriptions = new ArrayList<String>();
volatile String nextConnackFail;
Map<String, Object> autoHello;
boolean autoSendOk = true;
final List<Function<Map<String, Object>, Map<String, Object>>> upHandlers =
new CopyOnWriteArrayList<Function<Map<String, Object>, Map<String, Object>>>();
private Runnable onConnected;
private BiConsumer<String, Boolean> onDisconnected;
private Consumer<byte[]> onDown;
private volatile boolean connected;
FakeTransport() {
autoHello = new LinkedHashMap<String, Object>();
autoHello.put("server_time_ms", 1750000000000L);
autoHello.put("server_version", "0.1.0");
autoHello.put("max_body_bytes", 262144);
autoHello.put("max_meta_bytes", 4096);
autoHello.put("max_frame_bytes", 786432);
autoHello.put("max_ttl_seconds", 2592000);
autoHello.put("max_schedule_seconds", 31536000);
autoHello.put("ack_timeout_seconds", 300);
autoHello.put("session_token", "nst_test_token");
}
void onUp(Function<Map<String, Object>, Map<String, Object>> h) {
upHandlers.add(h);
}
@Override
public void setHandlers(Runnable onConnected, BiConsumer<String, Boolean> onDisconnected, Consumer<byte[]> onDown) {
this.onConnected = onConnected;
this.onDisconnected = onDisconnected;
this.onDown = onDown;
}
@Override
public void connect(ConnectParams params) {
connects.add(params);
if (nextConnackFail != null) {
String fail = nextConnackFail;
nextConnackFail = null;
boolean stop = "session_invalid".equals(fail) || "bad_credentials".equals(fail) || "banned".equals(fail);
if (onDisconnected != null) {
onDisconnected.accept(fail, stop);
}
return;
}
connected = true;
if (onConnected != null) {
onConnected.run();
}
}
@Override
public void subscribe(String topic) {
subscriptions.add(topic);
}
@Override
public void publish(String topic, byte[] payload) {
publishes.add(payload);
Map<String, Object> frame = Protocol.loads(payload);
for (Function<Map<String, Object>, Map<String, Object>> h : upHandlers) {
Map<String, Object> resp = h.apply(frame);
if (resp != null) {
injectDown(Protocol.dumps(resp));
return;
}
}
String type = str(frame.get("type"));
Object rid = frame.get("rid");
if ("hello".equals(type) && autoHello != null) {
Map<String, Object> resp = new LinkedHashMap<String, Object>();
resp.put("v", 1);
resp.put("type", "resp");
resp.put("rid", rid);
resp.put("ok", true);
resp.put("data", autoHello);
injectDown(Protocol.dumps(resp));
return;
}
if ("send".equals(type) && autoSendOk) {
Map<String, Object> data = new LinkedHashMap<String, Object>();
data.put("id", frame.get("id"));
Object sat = frame.get("send_at_ms");
data.put("send_at_ms", sat == null ? 0 : sat);
data.put("state", "dispatched");
Map<String, Object> resp = new LinkedHashMap<String, Object>();
resp.put("v", 1);
resp.put("type", "resp");
resp.put("rid", rid);
resp.put("ok", true);
resp.put("data", data);
injectDown(Protocol.dumps(resp));
return;
}
if ("ack".equals(type)) {
Map<String, Object> data = new LinkedHashMap<String, Object>();
data.put("result", "accepted");
Map<String, Object> resp = new LinkedHashMap<String, Object>();
resp.put("v", 1);
resp.put("type", "resp");
resp.put("rid", rid);
resp.put("ok", true);
resp.put("data", data);
injectDown(Protocol.dumps(resp));
return;
}
if (type != null && !"hello".equals(type) && !"send".equals(type) && rid != null) {
Map<String, Object> resp = new LinkedHashMap<String, Object>();
resp.put("v", 1);
resp.put("type", "resp");
resp.put("rid", rid);
resp.put("ok", true);
resp.put("data", new LinkedHashMap<String, Object>());
injectDown(Protocol.dumps(resp));
}
}
@Override
public void disconnect() {
boolean was = connected;
connected = false;
if (was && onDisconnected != null) {
onDisconnected.accept(null, false);
}
}
void injectDown(byte[] payload) {
if (onDown != null) {
onDown.accept(payload);
}
}
void simulateTakenOver() {
connected = false;
if (onDisconnected != null) {
onDisconnected.accept("taken_over", true);
}
}
void simulateNetworkDrop() {
connected = false;
if (onDisconnected != null) {
onDisconnected.accept("network", false);
}
}
private static String str(Object o) {
return o == null ? null : String.valueOf(o);
}
}
/** HiveMQ MQTT 5 + WebSocket(子协议 mqtt)。 */
final class HiveMqTransport implements Transport {
private Runnable onConnected;
private BiConsumer<String, Boolean> onDisconnected;
private Consumer<byte[]> onDown;
private volatile Mqtt5AsyncClient client;
@Override
public void setHandlers(Runnable onConnected, BiConsumer<String, Boolean> onDisconnected, Consumer<byte[]> onDown) {
this.onConnected = onConnected;
this.onDisconnected = onDisconnected;
this.onDown = onDown;
}
@Override
public void connect(final ConnectParams params) {
disconnectQuiet();
URI u = URI.create(params.url.contains("://") ? params.url : "ws://" + params.url);
boolean useTcp = params.useTcp
|| "mqtt".equalsIgnoreCase(u.getScheme())
|| "mqtts".equalsIgnoreCase(u.getScheme());
String host = u.getHost() == null ? "localhost" : u.getHost();
int port = u.getPort();
if (port < 0) {
if (useTcp) {
port = "mqtts".equalsIgnoreCase(u.getScheme()) ? 8883 : 1883;
} else {
port = "wss".equalsIgnoreCase(u.getScheme()) ? 443 : 80;
}
}
Mqtt5ClientBuilder b5 = MqttClient.builder()
.useMqttVersion5()
.identifier(params.clientId)
.serverHost(host)
.serverPort(port)
.addDisconnectedListener(context -> {
String reason = "network";
boolean stop = false;
if (context.getSource() == MqttDisconnectSource.SERVER) {
try {
java.lang.reflect.Method m = context.getClass().getMethod("getMqttDisconnect");
Object disc = m.invoke(context);
if (disc instanceof Mqtt5Disconnect) {
Mqtt5DisconnectReasonCode rc = ((Mqtt5Disconnect) disc).getReasonCode();
if (rc == Mqtt5DisconnectReasonCode.SESSION_TAKEN_OVER) {
reason = "taken_over";
stop = true;
}
}
} catch (Exception ignored) {
}
}
Throwable cause = context.getCause();
if (cause != null && cause.getMessage() != null
&& cause.getMessage().toLowerCase().contains("taken over")) {
reason = "taken_over";
stop = true;
}
BiConsumer<String, Boolean> h = onDisconnected;
if (h != null) {
h.accept(reason, stop);
}
});
if (!useTcp) {
String path = u.getPath() == null || u.getPath().isEmpty() ? "/mqtt" : u.getPath();
if (!path.endsWith("/mqtt")) {
path = path.endsWith("/") ? path + "mqtt" : path + "/mqtt";
}
String serverPath = path.startsWith("/") ? path.substring(1) : path;
b5 = b5.webSocketConfig()
.serverPath(serverPath)
.subprotocol("mqtt")
.applyWebSocketConfig();
if ("wss".equalsIgnoreCase(u.getScheme())) {
b5 = b5.sslWithDefaultConfig();
}
} else if ("mqtts".equalsIgnoreCase(u.getScheme())) {
b5 = b5.sslWithDefaultConfig();
}
Mqtt5AsyncClient c = b5.buildAsync();
this.client = c;
c.publishes(MqttGlobalPublishFilter.ALL, publish -> {
Consumer<byte[]> h = onDown;
if (h != null) {
h.accept(publish.getPayloadAsBytes());
}
});
try {
Mqtt5ConnAck ack = c.connectWith()
.cleanStart(params.cleanStart)
.sessionExpiryInterval(params.sessionExpiry)
.simpleAuth()
.username(params.username)
.password(params.password.getBytes(StandardCharsets.UTF_8))
.applySimpleAuth()
.send()
.get(params.timeoutMs, TimeUnit.MILLISECONDS);
if (ack.getReasonCode() == Mqtt5ConnAckReasonCode.SUCCESS) {
Runnable h = onConnected;
if (h != null) {
h.run();
}
} else {
String reason = classify(ack.getReasonCode());
boolean stop = isStop(ack.getReasonCode());
BiConsumer<String, Boolean> h = onDisconnected;
if (h != null) {
h.accept(reason, stop);
}
}
} catch (Exception e) {
BiConsumer<String, Boolean> h = onDisconnected;
if (h != null) {
h.accept("network", false);
}
}
}
@Override
public void subscribe(String topic) {
Mqtt5AsyncClient c = client;
if (c == null) {
throw new NixMsgException("not_connected", "未连接");
}
try {
c.subscribeWith().topicFilter(topic).qos(MqttQos.AT_LEAST_ONCE).send().get(10, TimeUnit.SECONDS);
} catch (Exception e) {
throw new NixMsgException("busy", "订阅失败");
}
}
@Override
public void publish(String topic, byte[] payload) {
Mqtt5AsyncClient c = client;
if (c == null) {
throw new NixMsgException("not_connected", "未连接");
}
try {
c.publishWith()
.topic(topic)
.qos(MqttQos.AT_LEAST_ONCE)
.payload(payload)
.send()
.get(10, TimeUnit.SECONDS);
} catch (Exception e) {
throw new NixMsgException("busy", "发布失败");
}
}
@Override
public void disconnect() {
disconnectQuiet();
}
private void disconnectQuiet() {
Mqtt5AsyncClient c = client;
client = null;
if (c != null) {
try {
c.disconnect().get(5, TimeUnit.SECONDS);
} catch (Exception ignored) {
}
}
}
private static boolean isStop(Mqtt5ConnAckReasonCode code) {
return code == Mqtt5ConnAckReasonCode.BAD_USER_NAME_OR_PASSWORD
|| code == Mqtt5ConnAckReasonCode.NOT_AUTHORIZED
|| code == Mqtt5ConnAckReasonCode.BANNED;
}
private static String classify(Mqtt5ConnAckReasonCode code) {
if (code == Mqtt5ConnAckReasonCode.BAD_USER_NAME_OR_PASSWORD
|| code == Mqtt5ConnAckReasonCode.NOT_AUTHORIZED
|| code == Mqtt5ConnAckReasonCode.BANNED) {
return "bad_credentials";
}
if (code == Mqtt5ConnAckReasonCode.SERVER_UNAVAILABLE || code == Mqtt5ConnAckReasonCode.SERVER_BUSY) {
return "busy";
}
return "network";
}
}
@@ -0,0 +1,240 @@
package asia.asio.nixmsg;
import java.nio.charset.StandardCharsets;
import java.util.Base64;
import java.util.Collections;
import java.util.LinkedHashMap;
import java.util.Map;
import java.util.Objects;
public final class Types {
private Types() {}
public static final String CLIENT_NAME = "java-sdk/0.1";
public static final int DEFAULT_MAX_BODY = 262144;
public static final int DEFAULT_MAX_META = 4096;
public static final int DEFAULT_MAX_FRAME = 786432;
public static final int MIN_MAX_RECEIVE = 1024;
public static final int SEND_QUEUE_LIMIT = 1000;
public static final int INFLIGHT_LIMIT = 100;
public static final int DEDUP_CAPACITY = 10000;
public static final long CONNECT_TIMEOUT_MS = 30_000L;
public static final long BACKOFF_INITIAL_MS = 1_000L;
public static final long BACKOFF_MAX_MS = 30_000L;
public static final double BACKOFF_JITTER = 0.3;
public static final long STABLE_RESET_MS = 60_000L;
public enum ConnectionState {
CONNECTING, ONLINE, RECONNECTING, OFFLINE, KICKED, AUTH_FAILED
}
public static final class Target {
public final String kind;
public final String id;
public Target(String kind, String id) {
this.kind = Objects.requireNonNull(kind);
this.id = Objects.requireNonNull(id);
}
public Map<String, Object> toMap() {
Map<String, Object> m = new LinkedHashMap<String, Object>();
m.put("kind", kind);
m.put("id", id);
return m;
}
}
public static final class Body {
public final String enc;
public final String data;
public String contentType;
public Body(String data) {
this("utf8", data, null);
}
public Body(String enc, String data, String contentType) {
this.enc = enc;
this.data = data;
this.contentType = contentType;
}
public static Body ofBytes(byte[] raw) {
return new Body("base64", Base64.getEncoder().encodeToString(raw), "application/octet-stream");
}
public int decodedSize() {
if ("base64".equals(enc)) {
return Base64.getDecoder().decode(data).length;
}
return data.getBytes(StandardCharsets.UTF_8).length;
}
public String effectiveContentType() {
if (contentType != null && !contentType.isEmpty()) {
return contentType;
}
return "base64".equals(enc) ? "application/octet-stream" : "text/plain; charset=utf-8";
}
public Map<String, Object> toMap() {
Map<String, Object> m = new LinkedHashMap<String, Object>();
m.put("enc", enc);
m.put("data", data);
m.put("content_type", effectiveContentType());
return m;
}
}
public static final class SendOptions {
public Long sendAtMs;
public Long delayMs;
public boolean keep;
public Long ttlSeconds;
public boolean receipt = true;
public String talkPassword = "";
public String contentType;
public Map<String, Object> meta;
public String messageId;
}
public static final class SendResult {
public final String id;
public final long sendAtMs;
public final String state;
public SendResult(String id, long sendAtMs, String state) {
this.id = id;
this.sendAtMs = sendAtMs;
this.state = state;
}
}
public static final class RecallResult {
public final String result;
public final int recalled;
public final int accepted;
public final int other;
public RecallResult(String result, int recalled, int accepted, int other) {
this.result = result;
this.recalled = recalled;
this.accepted = accepted;
this.other = other;
}
}
public static final class IncomingMessage {
public final String id;
public final String from;
public final Target to;
public final Body body;
public final long sendAtMs;
public final Map<String, Object> meta;
public IncomingMessage(String id, String from, Target to, Body body, long sendAtMs, Map<String, Object> meta) {
this.id = id;
this.from = from;
this.to = to;
this.body = body;
this.sendAtMs = sendAtMs;
this.meta = meta == null ? Collections.<String, Object>emptyMap() : meta;
}
}
public static final class Receipt {
public final String receiptId;
public final String id;
public final String endpointId;
public final String state;
public final String reason;
public final long atMs;
public Receipt(String receiptId, String id, String endpointId, String state, String reason, long atMs) {
this.receiptId = receiptId;
this.id = id;
this.endpointId = endpointId;
this.state = state;
this.reason = reason;
this.atMs = atMs;
}
}
public static final class RevokedEvent {
public final String id;
public final String from;
public final String reason;
public RevokedEvent(String id, String from, String reason) {
this.id = id;
this.from = from;
this.reason = reason;
}
}
public static final class PresenceEvent {
public final String id;
public final boolean online;
public final long atMs;
public PresenceEvent(String id, boolean online, long atMs) {
this.id = id;
this.online = online;
this.atMs = atMs;
}
}
public static final class GroupEvent {
public final String groupId;
public final String event;
public final String endpointId;
public final long atMs;
public GroupEvent(String groupId, String event, String endpointId, long atMs) {
this.groupId = groupId;
this.event = event;
this.endpointId = endpointId;
this.atMs = atMs;
}
}
public static final class HelloLimits {
public long serverTimeMs;
public String serverVersion = "";
public int maxBodyBytes = DEFAULT_MAX_BODY;
public int maxMetaBytes = DEFAULT_MAX_META;
public int maxFrameBytes = DEFAULT_MAX_FRAME;
public long maxTtlSeconds = 2592000;
public long maxScheduleSeconds = 31536000;
public long ackTimeoutSeconds = 300;
public String sessionToken = "";
}
public static final class RegisterOptions {
public String id = "";
public String loginPassword = "";
public String name = "";
public String talkPassword = "";
}
public static final class RegisterResult {
public final String id;
public final String loginPassword;
public RegisterResult(String id, String loginPassword) {
this.id = id;
this.loginPassword = loginPassword;
}
}
public static final class ConnectionEvent {
public final ConnectionState state;
public final String reason;
public ConnectionEvent(ConnectionState state, String reason) {
this.state = state;
this.reason = reason == null ? "" : reason;
}
}
}
@@ -0,0 +1,19 @@
package asia.asio.nixmsg;
import java.security.SecureRandom;
import java.util.UUID;
final class Uuid7 {
private static final SecureRandom RAND = new SecureRandom();
private Uuid7() {}
static String next() {
long ts = System.currentTimeMillis() & ((1L << 48) - 1);
long randA = RAND.nextInt() & 0x0FFF;
long randB = RAND.nextLong() & ((1L << 62) - 1);
long msb = (ts << 16) | (0x7L << 12) | randA;
long lsb = (0b10L << 62) | randB;
return new UUID(msb, lsb).toString();
}
}
@@ -0,0 +1,276 @@
package asia.asio.nixmsg;
import asia.asio.nixmsg.Types.Body;
import asia.asio.nixmsg.Types.ConnectionState;
import asia.asio.nixmsg.Types.IncomingMessage;
import asia.asio.nixmsg.Types.SendOptions;
import asia.asio.nixmsg.Types.SendResult;
import asia.asio.nixmsg.Types.Target;
import org.junit.After;
import org.junit.Test;
import java.nio.charset.StandardCharsets;
import java.util.ArrayList;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import java.util.concurrent.CountDownLatch;
import java.util.concurrent.TimeUnit;
import java.util.concurrent.atomic.AtomicInteger;
import static org.junit.Assert.assertEquals;
import static org.junit.Assert.assertTrue;
import static org.junit.Assert.fail;
public class ClientTest {
private Client client;
private FakeTransport transport;
@After
public void tearDown() {
if (client != null) {
client.close();
}
}
private void connectOnline() {
transport = new FakeTransport();
client = new Client(transport);
client.connectSync("ws://example.test/mqtt", "ep1", "secret", null, false);
assertEquals(ConnectionState.ONLINE, client.getState());
}
@Test
public void cleanStartAndSessionExpiryEveryConnect() throws Exception {
connectOnline();
assertTrue(transport.connects.size() >= 1);
Transport.ConnectParams p = transport.connects.get(0);
assertTrue(p.cleanStart);
assertEquals(0, p.sessionExpiry);
assertTrue(transport.subscriptions.contains("nix/c/ep1/down"));
transport.simulateNetworkDrop();
long deadline = System.currentTimeMillis() + 3000;
while (transport.connects.size() < 2 && System.currentTimeMillis() < deadline) {
Thread.sleep(50);
}
assertTrue(transport.connects.size() >= 2);
for (Transport.ConnectParams c : transport.connects) {
assertTrue(c.cleanStart);
assertEquals(0, c.sessionExpiry);
}
}
@Test
public void sessionTokenCallback() {
transport = new FakeTransport();
client = new Client(transport);
final List<String> tokens = new ArrayList<String>();
client.onSession(new java.util.function.Consumer<String>() {
@Override
public void accept(String t) {
tokens.add(t);
}
});
client.connectSync("ws://example.test/mqtt", "ep1", "pw", null, false);
assertEquals(1, tokens.size());
assertEquals("nst_test_token", tokens.get(0));
assertEquals("nst_test_token", client.getSessionToken());
}
@Test
public void dedupAckedRearrivalAcksAgain() throws Exception {
connectOnline();
final List<String> delivered = new ArrayList<String>();
client.onMessage(new java.util.function.Consumer<IncomingMessage>() {
@Override
public void accept(IncomingMessage m) {
delivered.add(m.id);
}
});
Map<String, Object> msg = new LinkedHashMap<String, Object>();
msg.put("v", 1);
msg.put("type", "msg");
msg.put("id", "m1");
msg.put("from", "peer");
Map<String, Object> to = new LinkedHashMap<String, Object>();
to.put("kind", "endpoint");
to.put("id", "ep1");
msg.put("to", to);
Map<String, Object> body = new LinkedHashMap<String, Object>();
body.put("enc", "utf8");
body.put("data", "hi");
msg.put("body", body);
msg.put("send_at_ms", 1);
transport.injectDown(Protocol.dumps(msg));
Thread.sleep(100);
assertEquals(1, delivered.size());
int before = transport.publishes.size();
transport.injectDown(Protocol.dumps(msg));
Thread.sleep(100);
assertEquals(1, delivered.size());
int ackCount = 0;
for (int i = before; i < transport.publishes.size(); i++) {
Map<String, Object> f = Protocol.loads(transport.publishes.get(i));
if ("ack".equals(String.valueOf(f.get("type")))) {
ackCount++;
}
}
assertTrue(ackCount >= 1);
}
@Test
public void localBodyTooLarge() {
connectOnline();
StringBuilder sb = new StringBuilder();
for (int i = 0; i < client.getLimits().maxBodyBytes + 1; i++) {
sb.append('x');
}
try {
client.sendSync(new Target("endpoint", "ep2"), new Body(sb.toString()), new SendOptions());
fail("expected body_too_large");
} catch (NixMsgException e) {
assertEquals("body_too_large", e.getCode());
}
}
@Test
public void localFrameTooLarge() {
transport = new FakeTransport();
transport.autoHello.put("max_frame_bytes", 200);
client = new Client(transport);
client.connectSync("ws://example.test/mqtt", "ep1", "pw", null, false);
StringBuilder sb = new StringBuilder();
for (int i = 0; i < 40; i++) {
sb.append("hello world ");
}
try {
client.sendSync(new Target("endpoint", "ep2"), new Body(sb.toString()), new SendOptions());
fail("expected frame_too_large");
} catch (NixMsgException e) {
assertEquals("frame_too_large", e.getCode());
}
}
@Test
public void resendKeepsSendAtMs() {
transport = new FakeTransport();
transport.autoSendOk = false;
final AtomicInteger rateHits = new AtomicInteger();
transport.onUp(new java.util.function.Function<Map<String, Object>, Map<String, Object>>() {
@Override
public Map<String, Object> apply(Map<String, Object> frame) {
if (!"send".equals(String.valueOf(frame.get("type")))) {
return null;
}
if (rateHits.getAndIncrement() == 0) {
Map<String, Object> err = new LinkedHashMap<String, Object>();
err.put("code", "rate_limited");
err.put("message", "slow");
Map<String, Object> resp = new LinkedHashMap<String, Object>();
resp.put("v", 1);
resp.put("type", "resp");
resp.put("rid", frame.get("rid"));
resp.put("ok", false);
resp.put("error", err);
return resp;
}
Map<String, Object> data = new LinkedHashMap<String, Object>();
data.put("id", frame.get("id"));
data.put("send_at_ms", frame.get("send_at_ms"));
data.put("state", "scheduled");
Map<String, Object> resp = new LinkedHashMap<String, Object>();
resp.put("v", 1);
resp.put("type", "resp");
resp.put("rid", frame.get("rid"));
resp.put("ok", true);
resp.put("data", data);
return resp;
}
});
client = new Client(transport);
client.connectSync("ws://example.test/mqtt", "ep1", "pw", null, false);
SendOptions opt = new SendOptions();
opt.sendAtMs = 1700000000000L;
opt.messageId = "fixed-id-1";
SendResult result = client.sendSync(new Target("endpoint", "ep2"), new Body("hi"), opt);
assertEquals("fixed-id-1", result.id);
int sendCount = 0;
for (byte[] raw : transport.publishes) {
Map<String, Object> f = Protocol.loads(raw);
if ("send".equals(String.valueOf(f.get("type")))) {
sendCount++;
assertEquals("fixed-id-1", String.valueOf(f.get("id")));
assertEquals(1700000000000L, ((Number) f.get("send_at_ms")).longValue());
}
}
assertTrue(sendCount >= 2);
}
@Test
public void sessionInvalidStopsReconnect() throws Exception {
transport = new FakeTransport();
client = new Client(transport);
client.connectSync("ws://example.test/mqtt", "ep1", "pw", null, false);
transport.nextConnackFail = "bad_credentials";
transport.simulateNetworkDrop();
long deadline = System.currentTimeMillis() + 3000;
while (client.getState() != ConnectionState.AUTH_FAILED && System.currentTimeMillis() < deadline) {
Thread.sleep(50);
}
assertEquals(ConnectionState.AUTH_FAILED, client.getState());
int n = transport.connects.size();
Thread.sleep(500);
assertEquals(n, transport.connects.size());
}
@Test
public void deliveredDuplicateIgnored() throws Exception {
transport = new FakeTransport();
client = new Client(transport, false, Types.DEFAULT_MAX_FRAME, Types.CLIENT_NAME, Types.CONNECT_TIMEOUT_MS);
client.connectSync("ws://example.test/mqtt", "ep1", "pw", null, false);
final CountDownLatch started = new CountDownLatch(1);
final CountDownLatch release = new CountDownLatch(1);
final List<String> seen = new ArrayList<String>();
client.onMessage(new java.util.function.Consumer<IncomingMessage>() {
@Override
public void accept(IncomingMessage m) {
seen.add(m.id);
started.countDown();
try {
release.await(2, TimeUnit.SECONDS);
} catch (InterruptedException e) {
Thread.currentThread().interrupt();
}
}
});
Map<String, Object> msg = new LinkedHashMap<String, Object>();
msg.put("v", 1);
msg.put("type", "msg");
msg.put("id", "m2");
msg.put("from", "peer");
Map<String, Object> to = new LinkedHashMap<String, Object>();
to.put("kind", "endpoint");
to.put("id", "ep1");
msg.put("to", to);
Map<String, Object> body = new LinkedHashMap<String, Object>();
body.put("enc", "utf8");
body.put("data", "x");
msg.put("body", body);
msg.put("send_at_ms", 1);
final byte[] raw = Protocol.dumps(msg);
Thread t = new Thread(new Runnable() {
@Override
public void run() {
transport.injectDown(raw);
}
});
t.start();
assertTrue(started.await(2, TimeUnit.SECONDS));
transport.injectDown(raw);
release.countDown();
t.join(2000);
Thread.sleep(50);
assertEquals(1, seen.size());
}
}