"""假传输单元测试: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()