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
+1 -1
View File
@@ -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(
+1 -2
View File
@@ -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",
+1
View File
@@ -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:
+235 -107
View File
@@ -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
+16 -5
View File
@@ -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))
+18 -8
View File
@@ -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)
+82
View File
@@ -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]
+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()