fix: 按 K-00 约定修复 Python SDK 断线重交与退避

This commit is contained in:
Nixevol
2026-09-30 16:24:13 +08:00
parent 55aa0ccf53
commit 537cef0e32
10 changed files with 564 additions and 126 deletions
+3 -1
View File
@@ -7,8 +7,9 @@ import threading
import time
import unittest
from nixmsg import Body, Client, ConnectionState, FakeTransport, NixMsgError, SendOptions, Target
from nixmsg import Body, Client, ConnectionState, NixMsgError, SendOptions, Target
from nixmsg.protocol import dumps
from nixmsg.transport import FakeTransport
class FakeTransportTests(unittest.TestCase):
@@ -48,6 +49,7 @@ class FakeTransportTests(unittest.TestCase):
c = Client(transport=tr)
c.on_session(lambda t: tokens.append(t))
c.connect("ws://example.test/mqtt", "ep1", password="pw")
time.sleep(0.05)
self.assertEqual(tokens, ["nst_test_token"])
self.assertEqual(c.session_token, "nst_test_token")
c.close()
+196
View File
@@ -0,0 +1,196 @@
"""K-00 / K-03 对外语义与退避单测。"""
from __future__ import annotations
import json
import threading
import time
import unittest
from nixmsg import Body, Client, ConnectionState, NixMsgError, SendOptions, Target
from nixmsg.protocol import dumps, normalize_mqtt_ws_url
from nixmsg.transport import FakeTransport
from nixmsg.types import ReconnectBackoff, disable_jitter_for_test, restore_jitter_for_test
class K00Tests(unittest.TestCase):
def tearDown(self) -> None:
restore_jitter_for_test()
def _connect(self, tr: FakeTransport | None = None, **kwargs) -> tuple[Client, FakeTransport]:
tr = tr or FakeTransport()
c = Client(transport=tr, **kwargs)
c.connect("ws://example.test/mqtt", "ep1", password="p", wait=True)
return c, tr
def test_k00_first_connect_timeout(self) -> None:
tr = FakeTransport()
tr.auto_accept = False
c = Client(transport=tr, connect_timeout_s=0.15)
with self.assertRaises(NixMsgError) as cm:
c.connect("ws://example.test/mqtt", "ep1", password="p", wait=True)
self.assertEqual(cm.exception.code, "not_connected")
tr.auto_accept = True
c.connect("ws://example.test/mqtt", "ep1", password="p", wait=True)
self.assertEqual(c.state, ConnectionState.ONLINE)
c.close()
def test_k00_taken_over_reason(self) -> None:
c, tr = self._connect()
got = []
c.on_connection(lambda ev: got.append(ev.reason) if ev.state == ConnectionState.KICKED else None)
tr.simulate_taken_over()
deadline = time.time() + 1
while time.time() < deadline and "taken_over" not in got:
time.sleep(0.01)
self.assertIn("taken_over", got)
self.assertEqual(c.last_stop_code, "taken_over")
def test_k00_send_after_stopped(self) -> None:
c, tr = self._connect()
tr.simulate_taken_over()
time.sleep(0.05)
with self.assertRaises(NixMsgError) as cm:
c.send(Target(kind="endpoint", id="b"), Body(data="x"))
self.assertEqual(cm.exception.code, "taken_over")
def test_k00_logout_returns_error(self) -> None:
c, tr = self._connect()
tr.simulate_network_drop()
time.sleep(0.05)
with self.assertRaises(NixMsgError) as cm:
c.logout()
self.assertEqual(cm.exception.code, "not_connected")
with self.assertRaises(NixMsgError):
c.send(Target(kind="endpoint", id="b"), Body(data="x"))
def test_k00_send_at_and_delay_conflict(self) -> None:
c, tr = self._connect()
with self.assertRaises(NixMsgError) as cm:
c.send(
Target(kind="endpoint", id="b"),
Body(data="x"),
SendOptions(send_at=1.0, delay_ms=1000),
)
self.assertEqual(cm.exception.code, "bad_request")
c.close()
def test_k00_max_receive_bytes_min(self) -> None:
tr = FakeTransport()
c = Client(transport=tr, max_receive_bytes=512)
with self.assertRaises(NixMsgError) as cm:
c.connect("ws://example.test/mqtt", "ep1", password="p")
self.assertEqual(cm.exception.code, "bad_request")
def test_k00_url_mapping(self) -> None:
u = normalize_mqtt_ws_url("https://host:7443/")
self.assertTrue(u.startswith("wss://"))
self.assertTrue(u.endswith("/mqtt"))
u = normalize_mqtt_ws_url("http://host/app")
self.assertTrue(u.startswith("ws://"))
self.assertIn("/app", u)
self.assertFalse(u.endswith("/app/mqtt"))
with self.assertRaises(ValueError):
normalize_mqtt_ws_url("mqtt://host:1883", allow_tcp=False)
self.assertIn("mqtt://", normalize_mqtt_ws_url("mqtt://host:1883", allow_tcp=True))
def test_k00_reconnect_backoff(self) -> None:
b = ReconnectBackoff()
self.assertEqual(b.next_wait_no_jitter(), 0.0)
got = []
for _ in range(6):
b.mark_offline()
got.append(b.next_wait_no_jitter())
self.assertEqual(got, [1.0, 2.0, 4.0, 8.0, 16.0, 30.0])
b2 = ReconnectBackoff()
b2.next_wait_no_jitter()
b2.mark_online()
b2.set_online_at_for_test(time.monotonic() - 61)
b2.mark_offline()
self.assertEqual(b2.next_wait_no_jitter(), 1.0)
def test_k00_rate_limited_new_rid(self) -> None:
disable_jitter_for_test()
tr = FakeTransport()
rids: list[str] = []
def hook(frame: dict):
if frame.get("type") != "send":
return None
rid = str(frame["rid"])
rids.append(rid)
if len(rids) < 3:
return {
"v": 1,
"type": "resp",
"rid": rid,
"ok": False,
"error": {"code": "rate_limited", "message": "slow"},
}
return {
"v": 1,
"type": "resp",
"rid": 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")
c.send(
Target(kind="endpoint", id="ep2"),
Body(data="hi"),
SendOptions(send_at_ms=1_700_000_000_000, message_id="id1"),
)
self.assertEqual(len(set(rids)), 3)
c.close()
def test_k00_inflight_resend(self) -> None:
tr = FakeTransport()
tr.auto_send_ok = False
first = {"rid": "", "id": ""}
def hook(frame: dict):
if frame.get("type") != "send":
return None
if not first["rid"]:
first["rid"] = str(frame["rid"])
first["id"] = str(frame["id"])
threading.Thread(target=tr.simulate_network_drop, daemon=True).start()
return None
if str(frame["rid"]) != first["rid"]:
self.assertEqual(frame["id"], first["id"])
return {
"v": 1,
"type": "resp",
"rid": frame["rid"],
"ok": True,
"data": {"id": frame["id"], "send_at_ms": frame.get("send_at_ms", 0), "state": "accepted"},
}
return None
tr.on_up(hook)
c = Client(transport=tr)
c.connect("ws://example.test/mqtt", "ep1", password="pw")
c.send(
Target(kind="endpoint", id="ep2"),
Body(data="hi"),
SendOptions(send_at_ms=1_700_000_000_111, message_id="keep"),
)
c.close()
def test_k00_duration_int64(self) -> None:
raw = json.loads('{"id":"m1","send_at_ms":123,"state":"scheduled"}')
self.assertEqual(raw["send_at_ms"], 123)
self.assertEqual(30 * 24 * 3600 * 1000, 2592000000)
def test_k00_keepalive_default(self) -> None:
c, tr = self._connect()
self.assertEqual(tr.connects[0].params.keep_alive, 30)
c.close()
if __name__ == "__main__":
unittest.main()