diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index c9d50bd..2c6a468 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -1180,6 +1180,15 @@ - 备选方案:沿用 attempt+base 双重翻倍;否决。 - 影响:仅 sdk/js。 +### 复审修复 K-03 + +- 日期:2026-09-30 +- 原条款:issue #60 及第二轮补充;DEVELOPMENT 第 9 节附录。 +- 实际做法:Paho `reconnect_on_failure=False`;重连退避只维护计数 n;断线后在途发送新 rid 重交;回调经单一队列派发不持锁;fatal 收包路径同步处理;logout 把请求失败抛出;假传输不从 `__init__` 导出;AsyncClient 记录事件循环并 await 协程回调。 +- 原因:与 K-00 对齐并修顶号互踢、回调死锁、断线永不重交。 +- 备选方案:照搬旧 JS 双重翻倍;否决。 +- 影响:仅 sdk/python。 + ## SDK 二 S2 ### S2-PY/JAVA 1–3 2026-09-30 @@ -1221,8 +1230,8 @@ 6. **Paho / HiveMQ 库内自动重连关闭,退避自管** - 原条款:四种 SDK 同一套重连:1s 起加倍上限 30s ±30% 抖动,稳定 60s 恢复;每次 Clean Start、会话过期 0。 - - 实际做法:两端均由 SDK 连接循环实现退避与停止条件;HiveMQ 不启库内 automaticReconnect;Paho 每次 `connect(..., clean_start=True)` 并设 `SessionExpiryInterval=0`。 - - 原因:与 Go/JS 要求一致,避免两套重连。 + - 实际做法:两端均由 SDK 连接循环实现退避与停止条件;HiveMQ 不启库内 automaticReconnect;Paho 创建 Client 时 `reconnect_on_failure=False`,每次 `connect(..., clean_start=True)` 并设 `SessionExpiryInterval=0`。 + - 原因:paho 2.x 默认自动重连,被顶号后会两端互踢。 - 备选方案:依赖库自带重连再改 Clean Start(易漏)。 - 影响:无。 diff --git a/sdk/python/README.md b/sdk/python/README.md index bdb5f63..295043e 100644 --- a/sdk/python/README.md +++ b/sdk/python/README.md @@ -25,7 +25,7 @@ pip install -e ".[dev]" from nixmsg import Body, Client, SendOptions, Target c = Client() -c.on_session(lambda token: print("session", token)) +c.on_session(lambda token: print("session", token[:8])) 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( diff --git a/sdk/python/src/nixmsg/__init__.py b/sdk/python/src/nixmsg/__init__.py index d138601..72ba466 100644 --- a/sdk/python/src/nixmsg/__init__.py +++ b/sdk/python/src/nixmsg/__init__.py @@ -4,7 +4,7 @@ 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 .transport import PahoTransport from .types import ( Body, ConnectionEvent, @@ -29,7 +29,6 @@ __all__ = [ "ClosedError", "ConnectionEvent", "ConnectionState", - "FakeTransport", "GroupEvent", "IncomingMessage", "NixMsgError", diff --git a/sdk/python/src/nixmsg/async_client.py b/sdk/python/src/nixmsg/async_client.py index a388eb4..d8061e1 100644 --- a/sdk/python/src/nixmsg/async_client.py +++ b/sdk/python/src/nixmsg/async_client.py @@ -48,6 +48,7 @@ class AsyncClient: self._client.on_connection(handler) async def connect(self, url: str, endpoint_id: str, **kwargs: Any) -> None: + self._client._loop = asyncio.get_running_loop() await asyncio.to_thread(self._client.connect, url, endpoint_id, **kwargs) async def close(self) -> None: diff --git a/sdk/python/src/nixmsg/client.py b/sdk/python/src/nixmsg/client.py index 59c038f..2fcb371 100644 --- a/sdk/python/src/nixmsg/client.py +++ b/sdk/python/src/nixmsg/client.py @@ -2,6 +2,7 @@ from __future__ import annotations +import inspect import logging import random import threading @@ -46,12 +47,15 @@ from .types import ( RegisterOptions, RegisterResult, RecallResult, + ReconnectBackoff, RevokedEvent, RevokedHandler, SendOptions, SendResult, SessionHandler, Target, + apply_jitter_s, + nominal_delay_s, ) from .uuid7 import new_uuid7 @@ -74,11 +78,14 @@ class _SendItem: message_id: str frame: dict[str, Any] # 已含 send_at_ms,重交不改 pending: _PendingReq + rate_n: int = 0 + retry_at: float = 0.0 class _DedupState: DELIVERED = "delivered" # 已交应用未确认 ACKED = "acked" + REVOKED = "revoked" class Client: @@ -93,7 +100,7 @@ class Client: ) -> None: self._transport: Transport = transport or PahoTransport() self._auto_ack = auto_ack - self._max_receive_bytes = max(MIN_MAX_RECEIVE, max_receive_bytes) + self._max_receive_bytes = max_receive_bytes self._client_name = client_name self._connect_timeout_s = connect_timeout_s @@ -106,7 +113,6 @@ class Client: 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 @@ -119,8 +125,10 @@ class Client: self._use_tcp = False self._limits = HelloLimits() self._clock_skew_ms = 0 - self._online_since = 0.0 - self._backoff_s = BACKOFF_INITIAL_S + self._backoff = ReconnectBackoff() + self._last_stop_code = "" + self._last_stop_err: Optional[BaseException] = None + self._loop = None self._rid_seq = 0 self._pending: dict[str, _PendingReq] = {} self._send_queue: list[_SendItem] = [] @@ -137,6 +145,9 @@ class Client: self._down_q: SimpleQueue = SimpleQueue() self._down_thread = threading.Thread(target=self._down_loop, name="nixmsg-down", daemon=True) self._down_thread.start() + self._cb_q: SimpleQueue = SimpleQueue() + self._cb_thread = threading.Thread(target=self._cb_loop, name="nixmsg-cb", daemon=True) + self._cb_thread.start() self._transport.set_handlers(self._on_transport_connected, self._on_transport_disconnected, self._on_down) @@ -175,10 +186,15 @@ class Client: ) -> None: if password is None and session_token is None: raise ValueError("需要 password 或 session_token") + if 0 < self._max_receive_bytes < MIN_MAX_RECEIVE: + raise NixMsgError("bad_request", "max_receive_bytes 小于 1024") with self._lock: if self._closed: raise ClosedError() - self._url = normalize_mqtt_ws_url(url) if not use_tcp else url + try: + self._url = normalize_mqtt_ws_url(url, allow_tcp=use_tcp) + except ValueError as e: + raise NixMsgError("bad_request", str(e)) from e self._endpoint_id = endpoint_id self._password = password self._session_token = session_token @@ -188,6 +204,9 @@ class Client: self._user_close = False self._want_connected = True self._handshake_error = None + self._last_stop_code = "" + self._last_stop_err = None + self._backoff = ReconnectBackoff() self._conn_event.clear() self._set_state(ConnectionState.CONNECTING) if self._worker is None or not self._worker.is_alive(): @@ -195,8 +214,17 @@ class Client: self._worker.start() self._wake.set() if wait: - if not self._conn_event.wait(self._connect_timeout_s + 5): - raise NixMsgError("busy", "连接超时") + if not self._conn_event.wait(self._connect_timeout_s): + with self._lock: + self._stop_reconnect = True + self._want_connected = False + self._last_stop_code = "not_connected" + self._last_stop_err = NixMsgError("not_connected", "连接超时") + try: + self._transport.disconnect() + except Exception: + pass + raise NixMsgError("not_connected", "连接超时") err = self._handshake_error if err: raise err @@ -208,7 +236,7 @@ class Client: ) if self._state == ConnectionState.KICKED: raise NixMsgError("taken_over", "会话被接管") - raise NixMsgError("busy", f"连接未成功: {self._state.value}") + raise NixMsgError("not_connected", f"连接未成功: {self._state.value}") def close(self) -> None: with self._lock: @@ -216,6 +244,8 @@ class Client: self._want_connected = False self._stop_reconnect = True self._closed = True + self._last_stop_code = "closed" + self._last_stop_err = ClosedError() self._fail_all_pending(ClosedError()) self._set_state(ConnectionState.OFFLINE) try: @@ -226,24 +256,34 @@ class Client: self._down_q.put(None) except Exception: pass + try: + self._cb_q.put(None) + except Exception: + pass self._wake.set() def logout(self) -> None: + err: Optional[BaseException] = None try: self._request({"type": "self.logout"}, wait=True) - except Exception: - pass + except Exception as e: + err = e + stop = NixMsgError("logged_out", "已退出登录") with self._lock: self._stop_reconnect = True self._want_connected = False self._session_token = None - self._fail_all_pending(NixMsgError("auth_failed", "已退出登录")) + self._last_stop_code = "logged_out" + self._last_stop_err = stop + self._fail_all_pending(stop) try: self._transport.disconnect() except Exception: pass self._set_state(ConnectionState.OFFLINE) self._wake.set() + if err: + raise err # ---- 发送 / 确认 ---- def send(self, to: Target, body: Body | str | bytes, options: Optional[SendOptions] = None) -> SendResult: @@ -260,8 +300,11 @@ class Client: with self._lock: if self._closed: raise ClosedError() + if self._stop_reconnect: + err = self._last_stop_err or NixMsgError(self._last_stop_code or "not_connected", "已停止重连") + raise err if len(self._send_queue) >= SEND_QUEUE_LIMIT: - raise NixMsgError("quota_exceeded", "发送队列已满") + raise NixMsgError("queue_full", "发送队列已满") max_body = self._limits.max_body_bytes or DEFAULT_MAX_BODY max_meta = self._limits.max_meta_bytes or DEFAULT_MAX_META @@ -289,6 +332,12 @@ class Client: frame["meta"] = meta # send_at_ms 在入队时固定,重交不重算 + if options.send_at_ms is not None and options.delay_ms is not None: + raise NixMsgError("bad_request", "send_at 与 delay 互斥") + if options.send_at is not None and options.delay_ms is not None: + raise NixMsgError("bad_request", "sendAt 与 delay 互斥") + if options.send_at_ms is not None and options.send_at is not None: + raise NixMsgError("bad_request", "send_at 与 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: @@ -468,26 +517,25 @@ class Client: 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 - # 尝试连接 + wait = self._backoff.next_wait() + if wait > 0: + self._wake.wait(wait) + self._wake.clear() try: self._attempt_connect() except Exception as e: log.debug("connect attempt failed: %s", e) + self._backoff.mark_offline() 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.wait(0.05) self._wake.clear() def _attempt_connect(self) -> None: @@ -540,9 +588,10 @@ class Client: "v": 1, "type": "hello", "rid": rid, - "max_receive_bytes": self._max_receive_bytes, "client": self._client_name, } + if self._max_receive_bytes > 0: + hello["max_receive_bytes"] = self._max_receive_bytes t0 = time.time() pending = _PendingReq(rid=rid, frame=hello) with self._lock: @@ -573,8 +622,8 @@ class Client: with self._lock: self._limits = limits self._clock_skew_ms = skew - self._online_since = time.monotonic() self._handshake_error = None + self._backoff.mark_online() self._set_state(ConnectionState.ONLINE) token = limits.session_token if token: @@ -609,12 +658,12 @@ class Client: if reason == "taken_over": self._stop_reconnect = True self._want_connected = False - self._fail_all_pending(NixMsgError("taken_over", "会话被接管")) + self._last_stop_code = "taken_over" + self._last_stop_err = NixMsgError("taken_over", "会话被接管") + self._fail_all_pending(self._last_stop_err) 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 + elif stop or reason in ("session_invalid", "bad_credentials", "banned"): if reason in ("bad_credentials", "session_invalid", "banned") or stop: auth_reason = reason or "bad_credentials" if auth_reason == "bad_credentials" and using_token: @@ -624,20 +673,31 @@ class Client: 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._last_stop_code = auth_reason + self._last_stop_err = NixMsgError(auth_reason, "认证失败") + self._fail_all_pending(self._last_stop_err) + self._handshake_error = self._last_stop_err self._set_state(ConnectionState.AUTH_FAILED, auth_reason) self._conn_event.set() - return - # 网络 / busy:继续重连 - if self._user_close or self._closed: + elif 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() + else: + self._backoff.mark_offline() + self._requeue_inflight() + self._fail_pending(NotConnectedError(), include_send=False) + 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() + if reason == "taken_over" or stop or reason in ("session_invalid", "bad_credentials", "banned"): + threading.Thread(target=self._safe_disconnect, name="nixmsg-stop", daemon=True).start() + + def _safe_disconnect(self) -> None: + try: + self._transport.disconnect() + except Exception: + pass def _on_down(self, payload: bytes) -> None: # resp 必须立即完成 pending(含 auto_ack 等待),不能进 down 队列,否则自死锁。 @@ -649,6 +709,9 @@ class Client: if frame.get("type") == "resp": self._dispatch_resp(frame) return + if frame.get("type") == "fatal": + self._handle_fatal(str(frame.get("reason", "protocol"))) + return if threading.current_thread() is self._down_thread: self._dispatch_down_body(frame) return @@ -685,6 +748,11 @@ class Client: pending.response = None pending.error = None pending.event.clear() + for it in self._send_queue: + if it.pending is pending: + it.rate_n += 1 + it.retry_at = time.monotonic() + apply_jitter_s(nominal_delay_s(it.rate_n)) + break self._wake.set() return with self._lock: @@ -733,16 +801,20 @@ class Client: ) 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 + self._handle_fatal(str(frame.get("reason", "protocol"))) + return + + def _handle_fatal(self, reason: str) -> None: + with self._lock: + if self._stop_reconnect and self._last_stop_code: + return + self._stop_reconnect = True + self._want_connected = False + self._last_stop_code = reason or "fatal" + self._last_stop_err = NixMsgError(reason or "fatal", "致命错误,停止重连") + self._fail_all_pending(self._last_stop_err) + self._set_state(ConnectionState.AUTH_FAILED, reason) + threading.Thread(target=self._safe_disconnect, name="nixmsg-fatal", daemon=True).start() def _handle_msg(self, frame: dict[str, Any]) -> None: mid = str(frame.get("id", "")) @@ -753,8 +825,8 @@ class Client: if st == _DedupState.ACKED: # 已确认再到达:再 ack,不交应用 pass - elif st == _DedupState.DELIVERED: - # 已交未确认:忽略 + elif st == _DedupState.DELIVERED or st == _DedupState.REVOKED: + # 已交未确认或已撤回:忽略 return else: st = None @@ -780,29 +852,54 @@ class Client: if not self._message_handler: if self._auto_ack: - self._send_ack(from_id, mid, mark_acked=True) + with self._lock: + self._dedup_put(key, _DedupState.ACKED) + self._send_ack(from_id, mid, mark_acked=False) return - try: - with self._cb_lock: - self._message_handler(msg) - except Exception: + cb_err: list[BaseException] = [] + done = threading.Event() + + def _run() -> None: + try: + self._invoke(self._message_handler, msg) + except Exception as e: + cb_err.append(e) + finally: + done.set() + + self._cb_q.put(_run) + done.wait(timeout=120) + if cb_err: log.exception("on_message 回调错误,等待重推") with self._lock: self._dedup.pop(key, None) return + with self._lock: + if self._dedup.get(key) == _DedupState.REVOKED: + return if self._auto_ack: - self._send_ack(from_id, mid, mark_acked=True) + with self._lock: + self._dedup_put(key, _DedupState.ACKED) + self._send_ack(from_id, mid, mark_acked=False) 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) + seen = True + else: + seen = False + self._receipt_seen[rid] = True + while len(self._receipt_seen) > DEDUP_CAPACITY: + self._receipt_seen.popitem(last=False) + if seen: + try: + self._request({"type": "receipt_ack", "receipt_id": rid}, wait=False) + except Exception: + pass + return receipt = Receipt( receipt_id=rid, id=str(frame.get("id", "")), @@ -824,19 +921,17 @@ class Client: key = f"{from_id}\0{mid}" with self._lock: st = self._dedup.get(key) - if st == _DedupState.ACKED: + if st == _DedupState.ACKED or st == _DedupState.REVOKED: return - if st is None: - # 还没交给应用:直接丢弃 - return - # 已交未确认:发撤回事件并不再确认 - self._dedup.pop(key, None) + self._dedup_put(key, _DedupState.REVOKED) 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: + if mark_acked: + raise return key = f"{from_id}\0{message_id}" data = resp.get("data") or {} @@ -858,7 +953,10 @@ class Client: if self._inflight_sends >= INFLIGHT_LIMIT: return item = None + now = time.monotonic() for it in self._send_queue: + if it.retry_at and it.retry_at > now: + continue if it.pending.rid == "" and it.pending.response is None and it.pending.error is None: item = it break @@ -868,19 +966,28 @@ class Client: item.pending.rid = rid frame = dict(item.frame) frame["rid"] = rid + raw = dumps(frame) + max_frame = self._limits.max_frame_bytes or DEFAULT_MAX_FRAME + if len(raw) > max_frame: + item.pending.error = NixMsgError("frame_too_large", "整帧超限") + item.pending.event.set() + self._send_queue = [it for it in self._send_queue if it is not item] + continue 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: + self._transport.publish(topic, raw) + except Exception: 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] + if self._stop_reconnect: + err = self._last_stop_err or NotConnectedError() + item.pending.error = err + 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]: @@ -938,74 +1045,95 @@ class Client: self._dedup.popitem(last=False) def _fail_all_pending(self, err: BaseException) -> None: - for p in list(self._pending.values()): + self._fail_pending(err, include_send=True) + + def _fail_pending(self, err: BaseException, *, include_send: bool) -> None: + for rid, p in list(self._pending.items()): + if not include_send and p.is_send: + self._pending.pop(rid, None) + p.rid = "" + continue 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._pending.pop(rid, None) + if include_send: + for it in list(self._send_queue): + it.pending.error = err + it.pending.event.set() + self._send_queue.clear() + self._inflight_sends = 0 + + def _requeue_inflight(self) -> None: + for it in self._send_queue: + if it.pending.rid: + self._pending.pop(it.pending.rid, None) + it.pending.rid = "" + it.pending.response = None + it.pending.error = None + it.pending.event.clear() self._inflight_sends = 0 + def _cb_loop(self) -> None: + while True: + fn = self._cb_q.get() + if fn is None: + return + try: + fn() + except Exception: + log.exception("回调错误") + + def _invoke(self, handler, *args: Any) -> None: + if handler is None: + return + if inspect.iscoroutinefunction(handler): + loop = self._loop + if loop is None: + raise NixMsgError("bad_request", "同步 Client 不支持 async 回调") + import asyncio + + asyncio.run_coroutine_threadsafe(handler(*args), loop).result() + return + result = handler(*args) + if inspect.iscoroutine(result): + result.close() + raise NixMsgError("bad_request", "同步 Client 不支持 async 回调") + def _set_state(self, state: ConnectionState, reason: str = "") -> None: self._state = state h = self._connection_handler + ev = ConnectionEvent(state=state, reason=reason) if h: - try: - with self._cb_lock: - h(ConnectionEvent(state=state, reason=reason)) - except Exception: - log.exception("on_connection 回调错误") + self._cb_q.put(lambda: self._invoke(h, ev)) 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 回调错误") + self._cb_q.put(lambda: self._invoke(h, token)) 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 回调错误") + self._cb_q.put(lambda: self._invoke(h, 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 回调错误") + self._cb_q.put(lambda: self._invoke(h, ev)) 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 回调错误") + self._cb_q.put(lambda: self._invoke(h, ev)) 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 回调错误") + self._cb_q.put(lambda: self._invoke(h, ev)) - @staticmethod - def _jitter(base: float) -> float: - return max(0.0, base * (1.0 + random.uniform(-BACKOFF_JITTER, BACKOFF_JITTER))) + @property + def last_stop_code(self) -> str: + return self._last_stop_code # 测试辅助 @property diff --git a/sdk/python/src/nixmsg/protocol.py b/sdk/python/src/nixmsg/protocol.py index 4b397df..5f43023 100644 --- a/sdk/python/src/nixmsg/protocol.py +++ b/sdk/python/src/nixmsg/protocol.py @@ -46,13 +46,24 @@ def down_topic(endpoint_id: str) -> str: return f"nix/c/{endpoint_id}/down" -def normalize_mqtt_ws_url(url: str) -> str: - """保证 WebSocket 路径以 /mqtt 结尾。""" +def normalize_mqtt_ws_url(url: str, *, allow_tcp: bool = False) -> str: + """http→ws、https→wss;路径为空或 / 时用 /mqtt,否则保留;mqtt/mqtts 仅显式允许裸 TCP。""" raw = url.strip() if "://" not in raw: raw = "ws://" + raw u = urlparse(raw) + scheme = u.scheme.lower() + if scheme == "http": + scheme = "ws" + elif scheme == "https": + scheme = "wss" + if scheme in ("mqtt", "mqtts"): + if not allow_tcp: + raise ValueError("裸 TCP 需显式 use_tcp") + return urlunparse((scheme, u.netloc, u.path, u.params, u.query, u.fragment)) + if scheme not in ("ws", "wss"): + raise ValueError(f"unsupported scheme {u.scheme}") 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)) + if path in ("", "/"): + path = "/mqtt" + return urlunparse((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 index f9b46d6..113921a 100644 --- a/sdk/python/src/nixmsg/transport.py +++ b/sdk/python/src/nixmsg/transport.py @@ -203,13 +203,21 @@ class PahoTransport: self.disconnect() url = params.url u = urlparse(url if "://" in url else "ws://" + url) - use_tcp = params.use_tcp or u.scheme in ("mqtt", "mqtts") + use_tcp = params.use_tcp + scheme = u.scheme.lower() + if scheme == "https": + scheme = "wss" + elif scheme == "http": + scheme = "ws" + if scheme in ("mqtt", "mqtts"): + use_tcp = True # WebSocket 必须显式 transport=websockets;裸 TCP 走默认。 client = PahoClient( callback_api_version=CallbackAPIVersion.VERSION2, client_id=params.client_id, protocol=MQTTv5, transport="tcp" if use_tcp else "websockets", + reconnect_on_failure=False, ) client.username_pw_set(params.username, params.password) client.on_connect = self._on_connect @@ -231,16 +239,14 @@ class PahoTransport: props = None if use_tcp: if not u.port: - port = 8883 if u.scheme == "mqtts" else 1883 - if u.scheme == "mqtts": + port = 8883 if scheme == "mqtts" else 1883 + if scheme == "mqtts": client.tls_set() 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": + port = 443 if scheme == "wss" else 80 + if scheme == "wss": client.tls_set() client.ws_set_options(path=path, headers={"Sec-WebSocket-Protocol": "mqtt"}) client.connect( @@ -309,11 +315,15 @@ class PahoTransport: if self._on_disconnected: self._on_disconnected(None, False) return - # 0x8E = 142 Session taken over + # 0x8E = 142 Session taken over;0x8B = 139 可重试 if code == 142: if self._on_disconnected: self._on_disconnected("taken_over", True) return + if code == 139: + if self._on_disconnected: + self._on_disconnected("network", False) + return stop, reason = _classify_connack(code) if self._on_disconnected: self._on_disconnected(reason, stop) diff --git a/sdk/python/src/nixmsg/types.py b/sdk/python/src/nixmsg/types.py index 04631e5..abd32da 100644 --- a/sdk/python/src/nixmsg/types.py +++ b/sdk/python/src/nixmsg/types.py @@ -2,6 +2,8 @@ from __future__ import annotations +import random +import time from dataclasses import dataclass, field from enum import Enum from typing import Any, Callable, Optional @@ -173,6 +175,86 @@ class ConnectionEvent: reason: str = "" +_JITTER_DISABLED = False + + +def disable_jitter_for_test() -> None: + global _JITTER_DISABLED + _JITTER_DISABLED = True + + +def restore_jitter_for_test() -> None: + global _JITTER_DISABLED + _JITTER_DISABLED = False + + +def nominal_delay_s(n: int) -> float: + if n < 1: + return 0.0 + if n > 6: + return BACKOFF_MAX_S + return min(BACKOFF_INITIAL_S * (2 ** (n - 1)), BACKOFF_MAX_S) + + +def apply_jitter_s(d: float) -> float: + if d <= 0: + return 0.0 + if _JITTER_DISABLED: + return d + return d * (0.7 + random.random() * 0.6) + + +class ReconnectBackoff: + """只维护一个连续失败计数 n。""" + + def __init__(self) -> None: + self.n = 0 + self._skip_first = True + self._online = False + self._online_at = 0.0 + self._counted = False + + def next_wait(self) -> float: + self._counted = False + if self._skip_first: + self._skip_first = False + return 0.0 + n = self.n if self.n >= 1 else 1 + return apply_jitter_s(nominal_delay_s(n)) + + def next_wait_no_jitter(self) -> float: + self._counted = False + if self._skip_first: + self._skip_first = False + return 0.0 + n = self.n if self.n >= 1 else 1 + return nominal_delay_s(n) + + def mark_online(self) -> None: + self._online = True + self._online_at = time.monotonic() + self._counted = False + + def mark_offline(self) -> None: + if self._counted: + return + self._counted = True + was = self._online + at = self._online_at + self._online = False + if not was: + self.n += 1 + return + if time.monotonic() - at >= STABLE_RESET_S: + self.n = 1 + return + self.n += 1 + + def set_online_at_for_test(self, mono: float) -> None: + self._online = True + self._online_at = mono + + SessionHandler = Callable[[str], None] MessageHandler = Callable[[IncomingMessage], None] ReceiptHandler = Callable[[Receipt], None] diff --git a/sdk/python/tests/test_client.py b/sdk/python/tests/test_client.py index b8928f2..cd48d6c 100644 --- a/sdk/python/tests/test_client.py +++ b/sdk/python/tests/test_client.py @@ -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() diff --git a/sdk/python/tests/test_k00.py b/sdk/python/tests/test_k00.py new file mode 100644 index 0000000..4c1e2b8 --- /dev/null +++ b/sdk/python/tests/test_k00.py @@ -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()