fix: 首次握手失败停止重连并修正 Java 限速在途计数

This commit is contained in:
Nixevol
2026-09-30 19:18:47 +08:00
parent 08c08de5e8
commit 9c6de00ea0
6 changed files with 210 additions and 12 deletions
+10
View File
@@ -1757,3 +1757,13 @@ issue #3 未关闭,`feat/fix-3-downlink-deadlock` 未合入 `main`。下面是
- 原因:原先只换 `rid` 并 `return`,连接未断时 `hello` 不再走,队列停泵,`send()` Promise 永不结束。 - 原因:原先只换 `rid` 并 `return`,连接未断时 `hello` 不再走,队列停泵,`send()` Promise 永不结束。
- 备选方案:按 `rate_limited` 同一条的退避再泵(否决本波,连接仍在线时立即重交更贴切,且避免与已有等待叠乘抖动)。 - 备选方案:按 `rate_limited` 同一条的退避再泵(否决本波,连接仍在线时立即重交更贴切,且避免与已有等待叠乘抖动)。
- 影响:仅 `sdk/js`;假传输失败一次后会再次上行且 `rid` 已变。 - 影响:仅 `sdk/js`;假传输失败一次后会再次上行且 `rid` 已变。
### 复审修复 R3-06
1. **Python/Java 首次握手失败应停止重连;Java rate_limited 在途计数只减一次**
- 日期:2026-09-30
- 原条款:DEVELOPMENT 第 9 节附录「首次连接超时或握手失败应停止重连,并向 connect() 返回未连接」;issue #70。
- 实际做法:`sdk/python` 的 `connect()` 在 `_handshake_error` 或非 ONLINE 终态时置 `_stop_reconnect`/`_want_connected=False`;内部连接/握手超时改用 `not_connected`(不再用 `busy`)。`sdk/java` 的 `connectSync` 在 `handshakeError` 时同样置 `stopReconnect`;`rate_limited` 路径去掉第二次 `inflightSends--`,只换新 rid 再排队。Python `transport.py` 将 paho 改为惰性导入,便于本机无 paho 时仍跑 FakeTransport 单测。不改 Go/JS 重连公式。
- 原因:原先抛错后工作线程仍按退避重连;Java 限速把在途计数减了两次。
- 备选方案:在后台 `_attempt_connect`/`attemptConnect` 内首次失败即停(否决:应用再次 `connect()` 才应恢复,标志应在对外 `connect` 失败路径统一置位)。
- 影响:仅 `sdk/python`、`sdk/java`;应用需再次调用 `connect()` 才会重连。
@@ -143,6 +143,12 @@ public final class Client {
return lastStopCode == null ? "" : lastStopCode; return lastStopCode == null ? "" : lastStopCode;
} }
int inflightSendsForTest() {
synchronized (lock) {
return inflightSends;
}
}
void setSendQueueLimitForTest(int n) { void setSendQueueLimitForTest(int n) {
sendQueueLimit = n; sendQueueLimit = n;
} }
@@ -211,6 +217,19 @@ public final class Client {
} }
} }
if (handshakeError != null) { if (handshakeError != null) {
synchronized (lock) {
stopReconnect = true;
wantConnected = false;
if (lastStopCode == null || lastStopCode.isEmpty()) {
if (handshakeError instanceof NixMsgException) {
lastStopCode = ((NixMsgException) handshakeError).getCode();
lastStopErr = (NixMsgException) handshakeError;
} else {
lastStopCode = "not_connected";
lastStopErr = new NixMsgException("not_connected", handshakeError.getMessage());
}
}
}
if (handshakeError instanceof NixMsgException) { if (handshakeError instanceof NixMsgException) {
throw (NixMsgException) handshakeError; throw (NixMsgException) handshakeError;
} }
@@ -229,6 +248,8 @@ public final class Client {
synchronized (lock) { synchronized (lock) {
stopReconnect = true; stopReconnect = true;
wantConnected = false; wantConnected = false;
lastStopCode = "not_connected";
lastStopErr = new NixMsgException("not_connected", "连接未成功: " + state);
} }
throw new NixMsgException("not_connected", "连接未成功: " + state); throw new NixMsgException("not_connected", "连接未成功: " + state);
} }
@@ -704,7 +725,9 @@ public final class Client {
transport.publish(Protocol.upTopic(endpointId), Protocol.dumps(hello)); transport.publish(Protocol.upTopic(endpointId), Protocol.dumps(hello));
await(p.future, connectTimeoutMs); await(p.future, connectTimeoutMs);
if (p.error != null) { if (p.error != null) {
throw p.error instanceof RuntimeException ? (RuntimeException) p.error : new NixMsgException("busy", p.error.getMessage()); throw p.error instanceof RuntimeException
? (RuntimeException) p.error
: new NixMsgException("not_connected", p.error.getMessage());
} }
if (p.response == null) { if (p.response == null) {
throw new NixMsgException("not_connected", "握手超时"); throw new NixMsgException("not_connected", "握手超时");
@@ -883,9 +906,9 @@ public final class Client {
if (!Boolean.TRUE.equals(frame.get("ok"))) { if (!Boolean.TRUE.equals(frame.get("ok"))) {
Map<String, Object> err = asMap(frame.get("error")); Map<String, Object> err = asMap(frame.get("error"));
if ("rate_limited".equals(str(err.get("code"), ""))) { if ("rate_limited".equals(str(err.get("code"), ""))) {
// 在途计数已在上方减过一次;只换新 rid 再排队,勿再减。
synchronized (lock) { synchronized (lock) {
pending.remove(rid); pending.remove(rid);
inflightSends = Math.max(0, inflightSends - 1);
p.rid = ""; p.rid = "";
p.response = null; p.response = null;
p.error = null; p.error = null;
@@ -41,7 +41,7 @@ public class K00Test {
} }
@Test @Test
public void testK00FirstConnectTimeout() { public void testK00FirstConnectTimeout() throws Exception {
FakeTransport tr = new FakeTransport(); FakeTransport tr = new FakeTransport();
tr.autoHello = null; tr.autoHello = null;
client = new Client(tr, true, Types.DEFAULT_MAX_FRAME, Types.CLIENT_NAME, 150); client = new Client(tr, true, Types.DEFAULT_MAX_FRAME, Types.CLIENT_NAME, 150);
@@ -51,6 +51,9 @@ public class K00Test {
} catch (NixMsgException e) { } catch (NixMsgException e) {
assertEquals("not_connected", e.getCode()); assertEquals("not_connected", e.getCode());
} }
int n = tr.connects.size();
Thread.sleep(350);
assertEquals("首次握手失败后不得自动重连", n, tr.connects.size());
tr.autoHello = new LinkedHashMap<String, Object>(); tr.autoHello = new LinkedHashMap<String, Object>();
tr.autoHello.put("server_time_ms", 1750000000000L); tr.autoHello.put("server_time_ms", 1750000000000L);
tr.autoHello.put("server_version", "0.1.0"); tr.autoHello.put("server_version", "0.1.0");
@@ -63,6 +66,108 @@ public class K00Test {
tr.autoHello.put("session_token", "nst_test_token"); tr.autoHello.put("session_token", "nst_test_token");
client.connectSync("ws://example.test/mqtt", "ep1", "p", null, false); client.connectSync("ws://example.test/mqtt", "ep1", "p", null, false);
assertEquals(ConnectionState.ONLINE, client.getState()); assertEquals(ConnectionState.ONLINE, client.getState());
assertTrue(tr.connects.size() > n);
}
@Test
public void testK00RateLimitedInflightOnce() throws Exception {
Types.ReconnectBackoff.disableJitterForTest();
FakeTransport tr = new FakeTransport();
tr.autoSendOk = false;
final List<String> holdRids = new CopyOnWriteArrayList<String>();
final List<String> limitedRids = new CopyOnWriteArrayList<String>();
tr.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;
}
String rid = String.valueOf(frame.get("rid"));
String id = String.valueOf(frame.get("id"));
if ("hold".equals(id)) {
holdRids.add(rid);
return null; // 保持在途
}
limitedRids.add(rid);
if (limitedRids.size() == 1) {
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", 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", rid);
resp.put("ok", true);
resp.put("data", data);
return resp;
}
});
connectOnline(tr);
Thread holder = new Thread(new Runnable() {
@Override
public void run() {
try {
SendOptions opt = new SendOptions();
opt.messageId = "hold";
opt.sendAtMs = 1700000000001L;
client.sendSync(new Target("endpoint", "ep2"), new Body("hold"), opt);
} catch (Exception ignored) {
}
}
});
holder.setDaemon(true);
holder.start();
long deadline = System.currentTimeMillis() + 2000;
while (holdRids.isEmpty() && System.currentTimeMillis() < deadline) {
Thread.sleep(10);
}
assertFalse(holdRids.isEmpty());
assertEquals(1, client.inflightSendsForTest());
final SendOptions opt = new SendOptions();
opt.messageId = "lim";
opt.sendAtMs = 1700000000002L;
Thread sender = new Thread(new Runnable() {
@Override
public void run() {
try {
client.sendSync(new Target("endpoint", "ep2"), new Body("lim"), opt);
} catch (Exception ignored) {
}
}
});
sender.setDaemon(true);
sender.start();
deadline = System.currentTimeMillis() + 3000;
while (limitedRids.size() < 1 && System.currentTimeMillis() < deadline) {
Thread.sleep(10);
}
assertTrue(limitedRids.size() >= 1);
// rate_limited 后只应减 1:仍剩 hold 那一条在途
deadline = System.currentTimeMillis() + 500;
int seen = -1;
while (System.currentTimeMillis() < deadline) {
seen = client.inflightSendsForTest();
if (seen == 1) {
break;
}
Thread.sleep(10);
}
assertEquals("rate_limited 后在途应只减 1", 1, seen);
client.close();
holder.join(1000);
sender.join(1000);
} }
@Test @Test
+16 -2
View File
@@ -227,6 +227,15 @@ class Client:
raise NixMsgError("not_connected", "连接超时") raise NixMsgError("not_connected", "连接超时")
err = self._handshake_error err = self._handshake_error
if err: if err:
with self._lock:
self._stop_reconnect = True
self._want_connected = False
if not self._last_stop_code:
code = getattr(err, "code", None) or "not_connected"
self._last_stop_code = str(code)
self._last_stop_err = err if isinstance(err, NixMsgError) else NixMsgError(
"not_connected", str(err)
)
raise err raise err
if self._state not in (ConnectionState.ONLINE,): if self._state not in (ConnectionState.ONLINE,):
if self._state == ConnectionState.AUTH_FAILED: if self._state == ConnectionState.AUTH_FAILED:
@@ -236,6 +245,11 @@ class Client:
) )
if self._state == ConnectionState.KICKED: if self._state == ConnectionState.KICKED:
raise NixMsgError("taken_over", "会话被接管") raise NixMsgError("taken_over", "会话被接管")
with self._lock:
self._stop_reconnect = True
self._want_connected = False
self._last_stop_code = "not_connected"
self._last_stop_err = NixMsgError("not_connected", f"连接未成功: {self._state.value}")
raise NixMsgError("not_connected", f"连接未成功: {self._state.value}") raise NixMsgError("not_connected", f"连接未成功: {self._state.value}")
def close(self) -> None: def close(self) -> None:
@@ -576,7 +590,7 @@ class Client:
except Exception: except Exception:
pass pass
with self._lock: with self._lock:
self._handshake_error = NixMsgError("busy", "连接超时") self._handshake_error = NixMsgError("not_connected", "连接超时")
return return
def _on_transport_connected(self) -> None: def _on_transport_connected(self) -> None:
@@ -598,7 +612,7 @@ class Client:
self._pending[rid] = pending self._pending[rid] = pending
self._transport.publish(up_topic(self._endpoint_id), dumps(hello)) self._transport.publish(up_topic(self._endpoint_id), dumps(hello))
if not pending.event.wait(self._connect_timeout_s): if not pending.event.wait(self._connect_timeout_s):
raise NixMsgError("busy", "握手超时") raise NixMsgError("not_connected", "握手超时")
if pending.error: if pending.error:
raise pending.error raise pending.error
assert pending.response is not None assert pending.response is not None
+22 -7
View File
@@ -8,16 +8,23 @@ from dataclasses import dataclass, field
from typing import Any, Callable, Optional, Protocol from typing import Any, Callable, Optional, Protocol
from urllib.parse import urlparse from urllib.parse import urlparse
from paho.mqtt.client import CallbackAPIVersion, Client as PahoClient, MQTT_ERR_SUCCESS, MQTTv5
from paho.mqtt.enums import MQTTErrorCode
from paho.mqtt.reasoncodes import ReasonCode
DownHandler = Callable[[bytes], None] DownHandler = Callable[[bytes], None]
ConnHandler = Callable[[], None] ConnHandler = Callable[[], None]
DiscHandler = Callable[[Optional[str], bool], None] DiscHandler = Callable[[Optional[str], bool], None]
# reason_code_str, stop_reconnect # reason_code_str, stop_reconnect
def _require_paho():
"""真实 MQTT 路径才加载 paho;假传输单测不依赖。"""
try:
from paho.mqtt.client import CallbackAPIVersion, Client as PahoClient, MQTT_ERR_SUCCESS, MQTTv5
from paho.mqtt.enums import MQTTErrorCode
from paho.mqtt.reasoncodes import ReasonCode
except ImportError as e:
raise ImportError("需要 paho-mqtt>=2.0(真实 MQTT 连接)") from e
return CallbackAPIVersion, PahoClient, MQTT_ERR_SUCCESS, MQTTv5, MQTTErrorCode, ReasonCode
@dataclass @dataclass
class ConnectParams: class ConnectParams:
url: str url: str
@@ -180,7 +187,7 @@ class PahoTransport:
"""paho-mqtt 2.x CallbackAPIVersion.VERSION2。""" """paho-mqtt 2.x CallbackAPIVersion.VERSION2。"""
def __init__(self) -> None: def __init__(self) -> None:
self._client: Optional[PahoClient] = None self._client: Any = None
self._on_connected: Optional[ConnHandler] = None self._on_connected: Optional[ConnHandler] = None
self._on_disconnected: Optional[DiscHandler] = None self._on_disconnected: Optional[DiscHandler] = None
self._on_down: Optional[DownHandler] = None self._on_down: Optional[DownHandler] = None
@@ -200,6 +207,7 @@ class PahoTransport:
self._on_down = on_down self._on_down = on_down
def connect(self, params: ConnectParams) -> None: def connect(self, params: ConnectParams) -> None:
CallbackAPIVersion, PahoClient, _, MQTTv5, _, _ = _require_paho()
self.disconnect() self.disconnect()
url = params.url url = params.url
u = urlparse(url if "://" in url else "ws://" + url) u = urlparse(url if "://" in url else "ws://" + url)
@@ -261,6 +269,7 @@ class PahoTransport:
# 等待连接结果由回调驱动;超时由 Client 层处理 # 等待连接结果由回调驱动;超时由 Client 层处理
def subscribe(self, topic: str) -> None: def subscribe(self, topic: str) -> None:
_, _, MQTT_ERR_SUCCESS, _, _, _ = _require_paho()
self._down_topic = topic self._down_topic = topic
if not self._client: if not self._client:
return return
@@ -273,6 +282,7 @@ class PahoTransport:
raise RuntimeError("subscribe timeout") raise RuntimeError("subscribe timeout")
def publish(self, topic: str, payload: bytes) -> None: def publish(self, topic: str, payload: bytes) -> None:
_, _, MQTT_ERR_SUCCESS, _, _, _ = _require_paho()
if not self._client: if not self._client:
raise RuntimeError("not connected") raise RuntimeError("not connected")
info = self._client.publish(topic, payload, qos=1) info = self._client.publish(topic, payload, qos=1)
@@ -338,9 +348,14 @@ def _reason_to_int(reason_code) -> Optional[int]:
return None return None
if isinstance(reason_code, int): if isinstance(reason_code, int):
return reason_code return reason_code
if isinstance(reason_code, ReasonCode): try:
_, _, _, _, MQTTErrorCode, ReasonCode = _require_paho()
except ImportError:
MQTTErrorCode = () # type: ignore[assignment,misc]
ReasonCode = () # type: ignore[assignment,misc]
if ReasonCode and isinstance(reason_code, ReasonCode):
return int(reason_code.value) return int(reason_code.value)
if isinstance(reason_code, MQTTErrorCode): if MQTTErrorCode and isinstance(reason_code, MQTTErrorCode):
return int(reason_code) return int(reason_code)
# paho 偶发其它包装 # paho 偶发其它包装
val = getattr(reason_code, "value", None) val = getattr(reason_code, "value", None)
+31
View File
@@ -30,11 +30,42 @@ class K00Tests(unittest.TestCase):
with self.assertRaises(NixMsgError) as cm: with self.assertRaises(NixMsgError) as cm:
c.connect("ws://example.test/mqtt", "ep1", password="p", wait=True) c.connect("ws://example.test/mqtt", "ep1", password="p", wait=True)
self.assertEqual(cm.exception.code, "not_connected") self.assertEqual(cm.exception.code, "not_connected")
n = len(tr.connects)
time.sleep(0.35)
self.assertEqual(len(tr.connects), n, "首次失败后不得自动重连")
tr.auto_accept = True tr.auto_accept = True
c.connect("ws://example.test/mqtt", "ep1", password="p", wait=True) c.connect("ws://example.test/mqtt", "ep1", password="p", wait=True)
self.assertEqual(c.state, ConnectionState.ONLINE) self.assertEqual(c.state, ConnectionState.ONLINE)
c.close() c.close()
def test_k00_first_hello_fail_stops_reconnect(self) -> None:
"""MQTT 已通但 hello 失败:connect 返回未连接,且后台不再连。"""
tr = FakeTransport()
tr.auto_hello = None
c = Client(transport=tr, connect_timeout_s=0.2)
with self.assertRaises(NixMsgError) as cm:
c.connect("ws://example.test/mqtt", "ep1", password="p", wait=True)
self.assertEqual(cm.exception.code, "not_connected")
n = len(tr.connects)
self.assertGreaterEqual(n, 1)
time.sleep(0.45)
self.assertEqual(len(tr.connects), n, "hello 失败后不得自动重连")
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": 786432,
"max_ttl_seconds": 2592000,
"max_schedule_seconds": 31536000,
"ack_timeout_seconds": 300,
"session_token": "nst_retry",
}
c.connect("ws://example.test/mqtt", "ep1", password="p", wait=True)
self.assertEqual(c.state, ConnectionState.ONLINE)
self.assertGreater(len(tr.connects), n)
c.close()
def test_k00_taken_over_reason(self) -> None: def test_k00_taken_over_reason(self) -> None:
c, tr = self._connect() c, tr = self._connect()
got = [] got = []