fix: 按 K-00 约定修复 Python SDK 断线重交与退避
This commit is contained in:
+11
-2
@@ -857,6 +857,15 @@
|
|||||||
- 备选方案:沿用 attempt+base 双重翻倍;否决。
|
- 备选方案:沿用 attempt+base 双重翻倍;否决。
|
||||||
- 影响:仅 sdk/js。
|
- 影响:仅 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
|
## SDK 二 S2
|
||||||
|
|
||||||
### S2-PY/JAVA 1–3 2026-09-30
|
### S2-PY/JAVA 1–3 2026-09-30
|
||||||
@@ -898,8 +907,8 @@
|
|||||||
|
|
||||||
6. **Paho / HiveMQ 库内自动重连关闭,退避自管**
|
6. **Paho / HiveMQ 库内自动重连关闭,退避自管**
|
||||||
- 原条款:四种 SDK 同一套重连:1s 起加倍上限 30s ±30% 抖动,稳定 60s 恢复;每次 Clean Start、会话过期 0。
|
- 原条款:四种 SDK 同一套重连:1s 起加倍上限 30s ±30% 抖动,稳定 60s 恢复;每次 Clean Start、会话过期 0。
|
||||||
- 实际做法:两端均由 SDK 连接循环实现退避与停止条件;HiveMQ 不启库内 automaticReconnect;Paho 每次 `connect(..., clean_start=True)` 并设 `SessionExpiryInterval=0`。
|
- 实际做法:两端均由 SDK 连接循环实现退避与停止条件;HiveMQ 不启库内 automaticReconnect;Paho 创建 Client 时 `reconnect_on_failure=False`,每次 `connect(..., clean_start=True)` 并设 `SessionExpiryInterval=0`。
|
||||||
- 原因:与 Go/JS 要求一致,避免两套重连。
|
- 原因:paho 2.x 默认自动重连,被顶号后会两端互踢。
|
||||||
- 备选方案:依赖库自带重连再改 Clean Start(易漏)。
|
- 备选方案:依赖库自带重连再改 Clean Start(易漏)。
|
||||||
- 影响:无。
|
- 影响:无。
|
||||||
|
|
||||||
|
|||||||
@@ -25,7 +25,7 @@ pip install -e ".[dev]"
|
|||||||
from nixmsg import Body, Client, SendOptions, Target
|
from nixmsg import Body, Client, SendOptions, Target
|
||||||
|
|
||||||
c = Client()
|
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.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.connect("ws://127.0.0.1:7443/mqtt", "device-1", password="secret")
|
||||||
c.send(
|
c.send(
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ from .async_client import AsyncClient
|
|||||||
from .client import Client
|
from .client import Client
|
||||||
from .errors import ClosedError, NixMsgError, NotConnectedError
|
from .errors import ClosedError, NixMsgError, NotConnectedError
|
||||||
from .protocol import register_url_from_connect
|
from .protocol import register_url_from_connect
|
||||||
from .transport import FakeTransport, PahoTransport
|
from .transport import PahoTransport
|
||||||
from .types import (
|
from .types import (
|
||||||
Body,
|
Body,
|
||||||
ConnectionEvent,
|
ConnectionEvent,
|
||||||
@@ -29,7 +29,6 @@ __all__ = [
|
|||||||
"ClosedError",
|
"ClosedError",
|
||||||
"ConnectionEvent",
|
"ConnectionEvent",
|
||||||
"ConnectionState",
|
"ConnectionState",
|
||||||
"FakeTransport",
|
|
||||||
"GroupEvent",
|
"GroupEvent",
|
||||||
"IncomingMessage",
|
"IncomingMessage",
|
||||||
"NixMsgError",
|
"NixMsgError",
|
||||||
|
|||||||
@@ -48,6 +48,7 @@ class AsyncClient:
|
|||||||
self._client.on_connection(handler)
|
self._client.on_connection(handler)
|
||||||
|
|
||||||
async def connect(self, url: str, endpoint_id: str, **kwargs: Any) -> None:
|
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)
|
await asyncio.to_thread(self._client.connect, url, endpoint_id, **kwargs)
|
||||||
|
|
||||||
async def close(self) -> None:
|
async def close(self) -> None:
|
||||||
|
|||||||
+217
-89
@@ -2,6 +2,7 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import inspect
|
||||||
import logging
|
import logging
|
||||||
import random
|
import random
|
||||||
import threading
|
import threading
|
||||||
@@ -46,12 +47,15 @@ from .types import (
|
|||||||
RegisterOptions,
|
RegisterOptions,
|
||||||
RegisterResult,
|
RegisterResult,
|
||||||
RecallResult,
|
RecallResult,
|
||||||
|
ReconnectBackoff,
|
||||||
RevokedEvent,
|
RevokedEvent,
|
||||||
RevokedHandler,
|
RevokedHandler,
|
||||||
SendOptions,
|
SendOptions,
|
||||||
SendResult,
|
SendResult,
|
||||||
SessionHandler,
|
SessionHandler,
|
||||||
Target,
|
Target,
|
||||||
|
apply_jitter_s,
|
||||||
|
nominal_delay_s,
|
||||||
)
|
)
|
||||||
from .uuid7 import new_uuid7
|
from .uuid7 import new_uuid7
|
||||||
|
|
||||||
@@ -74,11 +78,14 @@ class _SendItem:
|
|||||||
message_id: str
|
message_id: str
|
||||||
frame: dict[str, Any] # 已含 send_at_ms,重交不改
|
frame: dict[str, Any] # 已含 send_at_ms,重交不改
|
||||||
pending: _PendingReq
|
pending: _PendingReq
|
||||||
|
rate_n: int = 0
|
||||||
|
retry_at: float = 0.0
|
||||||
|
|
||||||
|
|
||||||
class _DedupState:
|
class _DedupState:
|
||||||
DELIVERED = "delivered" # 已交应用未确认
|
DELIVERED = "delivered" # 已交应用未确认
|
||||||
ACKED = "acked"
|
ACKED = "acked"
|
||||||
|
REVOKED = "revoked"
|
||||||
|
|
||||||
|
|
||||||
class Client:
|
class Client:
|
||||||
@@ -93,7 +100,7 @@ class Client:
|
|||||||
) -> None:
|
) -> None:
|
||||||
self._transport: Transport = transport or PahoTransport()
|
self._transport: Transport = transport or PahoTransport()
|
||||||
self._auto_ack = auto_ack
|
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._client_name = client_name
|
||||||
self._connect_timeout_s = connect_timeout_s
|
self._connect_timeout_s = connect_timeout_s
|
||||||
|
|
||||||
@@ -106,7 +113,6 @@ class Client:
|
|||||||
self._connection_handler: Optional[ConnectionHandler] = None
|
self._connection_handler: Optional[ConnectionHandler] = None
|
||||||
|
|
||||||
self._lock = threading.RLock()
|
self._lock = threading.RLock()
|
||||||
self._cb_lock = threading.Lock() # 回调串行
|
|
||||||
self._state = ConnectionState.OFFLINE
|
self._state = ConnectionState.OFFLINE
|
||||||
self._stop_reconnect = False
|
self._stop_reconnect = False
|
||||||
self._closed = False
|
self._closed = False
|
||||||
@@ -119,8 +125,10 @@ class Client:
|
|||||||
self._use_tcp = False
|
self._use_tcp = False
|
||||||
self._limits = HelloLimits()
|
self._limits = HelloLimits()
|
||||||
self._clock_skew_ms = 0
|
self._clock_skew_ms = 0
|
||||||
self._online_since = 0.0
|
self._backoff = ReconnectBackoff()
|
||||||
self._backoff_s = BACKOFF_INITIAL_S
|
self._last_stop_code = ""
|
||||||
|
self._last_stop_err: Optional[BaseException] = None
|
||||||
|
self._loop = None
|
||||||
self._rid_seq = 0
|
self._rid_seq = 0
|
||||||
self._pending: dict[str, _PendingReq] = {}
|
self._pending: dict[str, _PendingReq] = {}
|
||||||
self._send_queue: list[_SendItem] = []
|
self._send_queue: list[_SendItem] = []
|
||||||
@@ -137,6 +145,9 @@ class Client:
|
|||||||
self._down_q: SimpleQueue = SimpleQueue()
|
self._down_q: SimpleQueue = SimpleQueue()
|
||||||
self._down_thread = threading.Thread(target=self._down_loop, name="nixmsg-down", daemon=True)
|
self._down_thread = threading.Thread(target=self._down_loop, name="nixmsg-down", daemon=True)
|
||||||
self._down_thread.start()
|
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)
|
self._transport.set_handlers(self._on_transport_connected, self._on_transport_disconnected, self._on_down)
|
||||||
|
|
||||||
@@ -175,10 +186,15 @@ class Client:
|
|||||||
) -> None:
|
) -> None:
|
||||||
if password is None and session_token is None:
|
if password is None and session_token is None:
|
||||||
raise ValueError("需要 password 或 session_token")
|
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:
|
with self._lock:
|
||||||
if self._closed:
|
if self._closed:
|
||||||
raise ClosedError()
|
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._endpoint_id = endpoint_id
|
||||||
self._password = password
|
self._password = password
|
||||||
self._session_token = session_token
|
self._session_token = session_token
|
||||||
@@ -188,6 +204,9 @@ class Client:
|
|||||||
self._user_close = False
|
self._user_close = False
|
||||||
self._want_connected = True
|
self._want_connected = True
|
||||||
self._handshake_error = None
|
self._handshake_error = None
|
||||||
|
self._last_stop_code = ""
|
||||||
|
self._last_stop_err = None
|
||||||
|
self._backoff = ReconnectBackoff()
|
||||||
self._conn_event.clear()
|
self._conn_event.clear()
|
||||||
self._set_state(ConnectionState.CONNECTING)
|
self._set_state(ConnectionState.CONNECTING)
|
||||||
if self._worker is None or not self._worker.is_alive():
|
if self._worker is None or not self._worker.is_alive():
|
||||||
@@ -195,8 +214,17 @@ class Client:
|
|||||||
self._worker.start()
|
self._worker.start()
|
||||||
self._wake.set()
|
self._wake.set()
|
||||||
if wait:
|
if wait:
|
||||||
if not self._conn_event.wait(self._connect_timeout_s + 5):
|
if not self._conn_event.wait(self._connect_timeout_s):
|
||||||
raise NixMsgError("busy", "连接超时")
|
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
|
err = self._handshake_error
|
||||||
if err:
|
if err:
|
||||||
raise err
|
raise err
|
||||||
@@ -208,7 +236,7 @@ class Client:
|
|||||||
)
|
)
|
||||||
if self._state == ConnectionState.KICKED:
|
if self._state == ConnectionState.KICKED:
|
||||||
raise NixMsgError("taken_over", "会话被接管")
|
raise NixMsgError("taken_over", "会话被接管")
|
||||||
raise NixMsgError("busy", f"连接未成功: {self._state.value}")
|
raise NixMsgError("not_connected", f"连接未成功: {self._state.value}")
|
||||||
|
|
||||||
def close(self) -> None:
|
def close(self) -> None:
|
||||||
with self._lock:
|
with self._lock:
|
||||||
@@ -216,6 +244,8 @@ class Client:
|
|||||||
self._want_connected = False
|
self._want_connected = False
|
||||||
self._stop_reconnect = True
|
self._stop_reconnect = True
|
||||||
self._closed = True
|
self._closed = True
|
||||||
|
self._last_stop_code = "closed"
|
||||||
|
self._last_stop_err = ClosedError()
|
||||||
self._fail_all_pending(ClosedError())
|
self._fail_all_pending(ClosedError())
|
||||||
self._set_state(ConnectionState.OFFLINE)
|
self._set_state(ConnectionState.OFFLINE)
|
||||||
try:
|
try:
|
||||||
@@ -226,24 +256,34 @@ class Client:
|
|||||||
self._down_q.put(None)
|
self._down_q.put(None)
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
try:
|
||||||
|
self._cb_q.put(None)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
self._wake.set()
|
self._wake.set()
|
||||||
|
|
||||||
def logout(self) -> None:
|
def logout(self) -> None:
|
||||||
|
err: Optional[BaseException] = None
|
||||||
try:
|
try:
|
||||||
self._request({"type": "self.logout"}, wait=True)
|
self._request({"type": "self.logout"}, wait=True)
|
||||||
except Exception:
|
except Exception as e:
|
||||||
pass
|
err = e
|
||||||
|
stop = NixMsgError("logged_out", "已退出登录")
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self._stop_reconnect = True
|
self._stop_reconnect = True
|
||||||
self._want_connected = False
|
self._want_connected = False
|
||||||
self._session_token = None
|
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:
|
try:
|
||||||
self._transport.disconnect()
|
self._transport.disconnect()
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
self._set_state(ConnectionState.OFFLINE)
|
self._set_state(ConnectionState.OFFLINE)
|
||||||
self._wake.set()
|
self._wake.set()
|
||||||
|
if err:
|
||||||
|
raise err
|
||||||
|
|
||||||
# ---- 发送 / 确认 ----
|
# ---- 发送 / 确认 ----
|
||||||
def send(self, to: Target, body: Body | str | bytes, options: Optional[SendOptions] = None) -> SendResult:
|
def send(self, to: Target, body: Body | str | bytes, options: Optional[SendOptions] = None) -> SendResult:
|
||||||
@@ -260,8 +300,11 @@ class Client:
|
|||||||
with self._lock:
|
with self._lock:
|
||||||
if self._closed:
|
if self._closed:
|
||||||
raise ClosedError()
|
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:
|
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_body = self._limits.max_body_bytes or DEFAULT_MAX_BODY
|
||||||
max_meta = self._limits.max_meta_bytes or DEFAULT_MAX_META
|
max_meta = self._limits.max_meta_bytes or DEFAULT_MAX_META
|
||||||
@@ -289,6 +332,12 @@ class Client:
|
|||||||
frame["meta"] = meta
|
frame["meta"] = meta
|
||||||
|
|
||||||
# send_at_ms 在入队时固定,重交不重算
|
# 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:
|
if options.send_at_ms is not None:
|
||||||
frame["send_at_ms"] = int(options.send_at_ms)
|
frame["send_at_ms"] = int(options.send_at_ms)
|
||||||
elif options.send_at is not None:
|
elif options.send_at is not None:
|
||||||
@@ -468,26 +517,25 @@ class Client:
|
|||||||
continue
|
continue
|
||||||
if state == ConnectionState.ONLINE:
|
if state == ConnectionState.ONLINE:
|
||||||
self._pump_sends()
|
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.wait(0.2)
|
||||||
self._wake.clear()
|
self._wake.clear()
|
||||||
continue
|
continue
|
||||||
# 尝试连接
|
wait = self._backoff.next_wait()
|
||||||
|
if wait > 0:
|
||||||
|
self._wake.wait(wait)
|
||||||
|
self._wake.clear()
|
||||||
try:
|
try:
|
||||||
self._attempt_connect()
|
self._attempt_connect()
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
log.debug("connect attempt failed: %s", e)
|
log.debug("connect attempt failed: %s", e)
|
||||||
|
self._backoff.mark_offline()
|
||||||
with self._lock:
|
with self._lock:
|
||||||
if self._state == ConnectionState.ONLINE:
|
if self._state == ConnectionState.ONLINE:
|
||||||
continue
|
continue
|
||||||
if self._stop_reconnect or not self._want_connected:
|
if self._stop_reconnect or not self._want_connected:
|
||||||
continue
|
continue
|
||||||
delay = self._jitter(self._backoff_s)
|
|
||||||
self._backoff_s = min(BACKOFF_MAX_S, self._backoff_s * 2)
|
|
||||||
self._set_state(ConnectionState.RECONNECTING)
|
self._set_state(ConnectionState.RECONNECTING)
|
||||||
self._wake.wait(delay)
|
self._wake.wait(0.05)
|
||||||
self._wake.clear()
|
self._wake.clear()
|
||||||
|
|
||||||
def _attempt_connect(self) -> None:
|
def _attempt_connect(self) -> None:
|
||||||
@@ -540,9 +588,10 @@ class Client:
|
|||||||
"v": 1,
|
"v": 1,
|
||||||
"type": "hello",
|
"type": "hello",
|
||||||
"rid": rid,
|
"rid": rid,
|
||||||
"max_receive_bytes": self._max_receive_bytes,
|
|
||||||
"client": self._client_name,
|
"client": self._client_name,
|
||||||
}
|
}
|
||||||
|
if self._max_receive_bytes > 0:
|
||||||
|
hello["max_receive_bytes"] = self._max_receive_bytes
|
||||||
t0 = time.time()
|
t0 = time.time()
|
||||||
pending = _PendingReq(rid=rid, frame=hello)
|
pending = _PendingReq(rid=rid, frame=hello)
|
||||||
with self._lock:
|
with self._lock:
|
||||||
@@ -573,8 +622,8 @@ class Client:
|
|||||||
with self._lock:
|
with self._lock:
|
||||||
self._limits = limits
|
self._limits = limits
|
||||||
self._clock_skew_ms = skew
|
self._clock_skew_ms = skew
|
||||||
self._online_since = time.monotonic()
|
|
||||||
self._handshake_error = None
|
self._handshake_error = None
|
||||||
|
self._backoff.mark_online()
|
||||||
self._set_state(ConnectionState.ONLINE)
|
self._set_state(ConnectionState.ONLINE)
|
||||||
token = limits.session_token
|
token = limits.session_token
|
||||||
if token:
|
if token:
|
||||||
@@ -609,12 +658,12 @@ class Client:
|
|||||||
if reason == "taken_over":
|
if reason == "taken_over":
|
||||||
self._stop_reconnect = True
|
self._stop_reconnect = True
|
||||||
self._want_connected = False
|
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._set_state(ConnectionState.KICKED, "taken_over")
|
||||||
self._conn_event.set()
|
self._conn_event.set()
|
||||||
return
|
elif stop or reason in ("session_invalid", "bad_credentials", "banned"):
|
||||||
if stop or reason in ("session_invalid", "bad_credentials", "banned"):
|
|
||||||
# 令牌被拒 -> session_invalid;密码被拒 -> bad_credentials
|
|
||||||
if reason in ("bad_credentials", "session_invalid", "banned") or stop:
|
if reason in ("bad_credentials", "session_invalid", "banned") or stop:
|
||||||
auth_reason = reason or "bad_credentials"
|
auth_reason = reason or "bad_credentials"
|
||||||
if auth_reason == "bad_credentials" and using_token:
|
if auth_reason == "bad_credentials" and using_token:
|
||||||
@@ -624,20 +673,31 @@ class Client:
|
|||||||
self._auth_reason = auth_reason
|
self._auth_reason = auth_reason
|
||||||
self._stop_reconnect = True
|
self._stop_reconnect = True
|
||||||
self._want_connected = False
|
self._want_connected = False
|
||||||
self._fail_all_pending(NixMsgError(auth_reason, "认证失败"))
|
self._last_stop_code = auth_reason
|
||||||
self._handshake_error = NixMsgError(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._set_state(ConnectionState.AUTH_FAILED, auth_reason)
|
||||||
self._conn_event.set()
|
self._conn_event.set()
|
||||||
return
|
elif self._user_close or self._closed:
|
||||||
# 网络 / busy:继续重连
|
|
||||||
if self._user_close or self._closed:
|
|
||||||
self._set_state(ConnectionState.OFFLINE)
|
self._set_state(ConnectionState.OFFLINE)
|
||||||
self._conn_event.set()
|
self._conn_event.set()
|
||||||
return
|
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):
|
if was_online or self._state in (ConnectionState.CONNECTING, ConnectionState.RECONNECTING):
|
||||||
self._set_state(ConnectionState.RECONNECTING, reason or "network")
|
self._set_state(ConnectionState.RECONNECTING, reason or "network")
|
||||||
self._conn_event.set()
|
self._conn_event.set()
|
||||||
self._wake.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:
|
def _on_down(self, payload: bytes) -> None:
|
||||||
# resp 必须立即完成 pending(含 auto_ack 等待),不能进 down 队列,否则自死锁。
|
# resp 必须立即完成 pending(含 auto_ack 等待),不能进 down 队列,否则自死锁。
|
||||||
@@ -649,6 +709,9 @@ class Client:
|
|||||||
if frame.get("type") == "resp":
|
if frame.get("type") == "resp":
|
||||||
self._dispatch_resp(frame)
|
self._dispatch_resp(frame)
|
||||||
return
|
return
|
||||||
|
if frame.get("type") == "fatal":
|
||||||
|
self._handle_fatal(str(frame.get("reason", "protocol")))
|
||||||
|
return
|
||||||
if threading.current_thread() is self._down_thread:
|
if threading.current_thread() is self._down_thread:
|
||||||
self._dispatch_down_body(frame)
|
self._dispatch_down_body(frame)
|
||||||
return
|
return
|
||||||
@@ -685,6 +748,11 @@ class Client:
|
|||||||
pending.response = None
|
pending.response = None
|
||||||
pending.error = None
|
pending.error = None
|
||||||
pending.event.clear()
|
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()
|
self._wake.set()
|
||||||
return
|
return
|
||||||
with self._lock:
|
with self._lock:
|
||||||
@@ -733,16 +801,20 @@ class Client:
|
|||||||
)
|
)
|
||||||
return
|
return
|
||||||
if ftype == "fatal":
|
if ftype == "fatal":
|
||||||
reason = str(frame.get("reason", "protocol"))
|
self._handle_fatal(str(frame.get("reason", "protocol")))
|
||||||
|
return
|
||||||
|
|
||||||
|
def _handle_fatal(self, reason: str) -> None:
|
||||||
with self._lock:
|
with self._lock:
|
||||||
|
if self._stop_reconnect and self._last_stop_code:
|
||||||
|
return
|
||||||
self._stop_reconnect = True
|
self._stop_reconnect = True
|
||||||
self._want_connected = False
|
self._want_connected = False
|
||||||
self._fail_all_pending(NixMsgError("fatal", reason))
|
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)
|
self._set_state(ConnectionState.AUTH_FAILED, reason)
|
||||||
try:
|
threading.Thread(target=self._safe_disconnect, name="nixmsg-fatal", daemon=True).start()
|
||||||
self._transport.disconnect()
|
|
||||||
except Exception:
|
|
||||||
pass
|
|
||||||
|
|
||||||
def _handle_msg(self, frame: dict[str, Any]) -> None:
|
def _handle_msg(self, frame: dict[str, Any]) -> None:
|
||||||
mid = str(frame.get("id", ""))
|
mid = str(frame.get("id", ""))
|
||||||
@@ -753,8 +825,8 @@ class Client:
|
|||||||
if st == _DedupState.ACKED:
|
if st == _DedupState.ACKED:
|
||||||
# 已确认再到达:再 ack,不交应用
|
# 已确认再到达:再 ack,不交应用
|
||||||
pass
|
pass
|
||||||
elif st == _DedupState.DELIVERED:
|
elif st == _DedupState.DELIVERED or st == _DedupState.REVOKED:
|
||||||
# 已交未确认:忽略
|
# 已交未确认或已撤回:忽略
|
||||||
return
|
return
|
||||||
else:
|
else:
|
||||||
st = None
|
st = None
|
||||||
@@ -780,29 +852,54 @@ class Client:
|
|||||||
|
|
||||||
if not self._message_handler:
|
if not self._message_handler:
|
||||||
if self._auto_ack:
|
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
|
return
|
||||||
|
|
||||||
|
cb_err: list[BaseException] = []
|
||||||
|
done = threading.Event()
|
||||||
|
|
||||||
|
def _run() -> None:
|
||||||
try:
|
try:
|
||||||
with self._cb_lock:
|
self._invoke(self._message_handler, msg)
|
||||||
self._message_handler(msg)
|
except Exception as e:
|
||||||
except Exception:
|
cb_err.append(e)
|
||||||
|
finally:
|
||||||
|
done.set()
|
||||||
|
|
||||||
|
self._cb_q.put(_run)
|
||||||
|
done.wait(timeout=120)
|
||||||
|
if cb_err:
|
||||||
log.exception("on_message 回调错误,等待重推")
|
log.exception("on_message 回调错误,等待重推")
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self._dedup.pop(key, None)
|
self._dedup.pop(key, None)
|
||||||
return
|
return
|
||||||
|
|
||||||
|
with self._lock:
|
||||||
|
if self._dedup.get(key) == _DedupState.REVOKED:
|
||||||
|
return
|
||||||
if self._auto_ack:
|
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:
|
def _handle_receipt(self, frame: dict[str, Any]) -> None:
|
||||||
rid = str(frame.get("receipt_id", ""))
|
rid = str(frame.get("receipt_id", ""))
|
||||||
with self._lock:
|
with self._lock:
|
||||||
if rid in self._receipt_seen:
|
if rid in self._receipt_seen:
|
||||||
return
|
seen = True
|
||||||
|
else:
|
||||||
|
seen = False
|
||||||
self._receipt_seen[rid] = True
|
self._receipt_seen[rid] = True
|
||||||
while len(self._receipt_seen) > DEDUP_CAPACITY:
|
while len(self._receipt_seen) > DEDUP_CAPACITY:
|
||||||
self._receipt_seen.popitem(last=False)
|
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 = Receipt(
|
||||||
receipt_id=rid,
|
receipt_id=rid,
|
||||||
id=str(frame.get("id", "")),
|
id=str(frame.get("id", "")),
|
||||||
@@ -824,19 +921,17 @@ class Client:
|
|||||||
key = f"{from_id}\0{mid}"
|
key = f"{from_id}\0{mid}"
|
||||||
with self._lock:
|
with self._lock:
|
||||||
st = self._dedup.get(key)
|
st = self._dedup.get(key)
|
||||||
if st == _DedupState.ACKED:
|
if st == _DedupState.ACKED or st == _DedupState.REVOKED:
|
||||||
return
|
return
|
||||||
if st is None:
|
self._dedup_put(key, _DedupState.REVOKED)
|
||||||
# 还没交给应用:直接丢弃
|
|
||||||
return
|
|
||||||
# 已交未确认:发撤回事件并不再确认
|
|
||||||
self._dedup.pop(key, None)
|
|
||||||
self._fire_revoked(RevokedEvent(id=mid, from_id=from_id, reason=str(frame.get("reason", ""))))
|
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:
|
def _send_ack(self, from_id: str, message_id: str, *, mark_acked: bool) -> None:
|
||||||
try:
|
try:
|
||||||
resp = self._request({"type": "ack", "from": from_id, "id": message_id})
|
resp = self._request({"type": "ack", "from": from_id, "id": message_id})
|
||||||
except Exception:
|
except Exception:
|
||||||
|
if mark_acked:
|
||||||
|
raise
|
||||||
return
|
return
|
||||||
key = f"{from_id}\0{message_id}"
|
key = f"{from_id}\0{message_id}"
|
||||||
data = resp.get("data") or {}
|
data = resp.get("data") or {}
|
||||||
@@ -858,7 +953,10 @@ class Client:
|
|||||||
if self._inflight_sends >= INFLIGHT_LIMIT:
|
if self._inflight_sends >= INFLIGHT_LIMIT:
|
||||||
return
|
return
|
||||||
item = None
|
item = None
|
||||||
|
now = time.monotonic()
|
||||||
for it in self._send_queue:
|
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:
|
if it.pending.rid == "" and it.pending.response is None and it.pending.error is None:
|
||||||
item = it
|
item = it
|
||||||
break
|
break
|
||||||
@@ -868,17 +966,26 @@ class Client:
|
|||||||
item.pending.rid = rid
|
item.pending.rid = rid
|
||||||
frame = dict(item.frame)
|
frame = dict(item.frame)
|
||||||
frame["rid"] = rid
|
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._pending[rid] = item.pending
|
||||||
self._inflight_sends += 1
|
self._inflight_sends += 1
|
||||||
topic = up_topic(self._endpoint_id)
|
topic = up_topic(self._endpoint_id)
|
||||||
try:
|
try:
|
||||||
self._transport.publish(topic, dumps(frame))
|
self._transport.publish(topic, raw)
|
||||||
except Exception as e:
|
except Exception:
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self._pending.pop(rid, None)
|
self._pending.pop(rid, None)
|
||||||
item.pending.rid = ""
|
item.pending.rid = ""
|
||||||
self._inflight_sends = max(0, self._inflight_sends - 1)
|
self._inflight_sends = max(0, self._inflight_sends - 1)
|
||||||
item.pending.error = e
|
if self._stop_reconnect:
|
||||||
|
err = self._last_stop_err or NotConnectedError()
|
||||||
|
item.pending.error = err
|
||||||
item.pending.event.set()
|
item.pending.event.set()
|
||||||
self._send_queue = [it for it in self._send_queue if it is not item]
|
self._send_queue = [it for it in self._send_queue if it is not item]
|
||||||
continue
|
continue
|
||||||
@@ -938,74 +1045,95 @@ class Client:
|
|||||||
self._dedup.popitem(last=False)
|
self._dedup.popitem(last=False)
|
||||||
|
|
||||||
def _fail_all_pending(self, err: BaseException) -> None:
|
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.error = err
|
||||||
p.event.set()
|
p.event.set()
|
||||||
self._pending.clear()
|
self._pending.pop(rid, None)
|
||||||
|
if include_send:
|
||||||
for it in list(self._send_queue):
|
for it in list(self._send_queue):
|
||||||
it.pending.error = err
|
it.pending.error = err
|
||||||
it.pending.event.set()
|
it.pending.event.set()
|
||||||
self._send_queue.clear()
|
self._send_queue.clear()
|
||||||
self._inflight_sends = 0
|
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:
|
def _set_state(self, state: ConnectionState, reason: str = "") -> None:
|
||||||
self._state = state
|
self._state = state
|
||||||
h = self._connection_handler
|
h = self._connection_handler
|
||||||
|
ev = ConnectionEvent(state=state, reason=reason)
|
||||||
if h:
|
if h:
|
||||||
try:
|
self._cb_q.put(lambda: self._invoke(h, ev))
|
||||||
with self._cb_lock:
|
|
||||||
h(ConnectionEvent(state=state, reason=reason))
|
|
||||||
except Exception:
|
|
||||||
log.exception("on_connection 回调错误")
|
|
||||||
|
|
||||||
def _fire_session(self, token: str) -> None:
|
def _fire_session(self, token: str) -> None:
|
||||||
h = self._session_handler
|
h = self._session_handler
|
||||||
if h:
|
if h:
|
||||||
try:
|
self._cb_q.put(lambda: self._invoke(h, token))
|
||||||
with self._cb_lock:
|
|
||||||
h(token)
|
|
||||||
except Exception:
|
|
||||||
log.exception("on_session 回调错误")
|
|
||||||
|
|
||||||
def _fire_receipt(self, receipt: Receipt) -> None:
|
def _fire_receipt(self, receipt: Receipt) -> None:
|
||||||
h = self._receipt_handler
|
h = self._receipt_handler
|
||||||
if h:
|
if h:
|
||||||
try:
|
self._cb_q.put(lambda: self._invoke(h, receipt))
|
||||||
with self._cb_lock:
|
|
||||||
h(receipt)
|
|
||||||
except Exception:
|
|
||||||
log.exception("on_receipt 回调错误")
|
|
||||||
|
|
||||||
def _fire_revoked(self, ev: RevokedEvent) -> None:
|
def _fire_revoked(self, ev: RevokedEvent) -> None:
|
||||||
h = self._revoked_handler
|
h = self._revoked_handler
|
||||||
if h:
|
if h:
|
||||||
try:
|
self._cb_q.put(lambda: self._invoke(h, ev))
|
||||||
with self._cb_lock:
|
|
||||||
h(ev)
|
|
||||||
except Exception:
|
|
||||||
log.exception("on_revoked 回调错误")
|
|
||||||
|
|
||||||
def _fire_presence(self, ev: PresenceEvent) -> None:
|
def _fire_presence(self, ev: PresenceEvent) -> None:
|
||||||
h = self._presence_handler
|
h = self._presence_handler
|
||||||
if h:
|
if h:
|
||||||
try:
|
self._cb_q.put(lambda: self._invoke(h, ev))
|
||||||
with self._cb_lock:
|
|
||||||
h(ev)
|
|
||||||
except Exception:
|
|
||||||
log.exception("on_presence 回调错误")
|
|
||||||
|
|
||||||
def _fire_group(self, ev: GroupEvent) -> None:
|
def _fire_group(self, ev: GroupEvent) -> None:
|
||||||
h = self._group_handler
|
h = self._group_handler
|
||||||
if h:
|
if h:
|
||||||
try:
|
self._cb_q.put(lambda: self._invoke(h, ev))
|
||||||
with self._cb_lock:
|
|
||||||
h(ev)
|
|
||||||
except Exception:
|
|
||||||
log.exception("on_group_event 回调错误")
|
|
||||||
|
|
||||||
@staticmethod
|
@property
|
||||||
def _jitter(base: float) -> float:
|
def last_stop_code(self) -> str:
|
||||||
return max(0.0, base * (1.0 + random.uniform(-BACKOFF_JITTER, BACKOFF_JITTER)))
|
return self._last_stop_code
|
||||||
|
|
||||||
# 测试辅助
|
# 测试辅助
|
||||||
@property
|
@property
|
||||||
|
|||||||
@@ -46,13 +46,24 @@ def down_topic(endpoint_id: str) -> str:
|
|||||||
return f"nix/c/{endpoint_id}/down"
|
return f"nix/c/{endpoint_id}/down"
|
||||||
|
|
||||||
|
|
||||||
def normalize_mqtt_ws_url(url: str) -> str:
|
def normalize_mqtt_ws_url(url: str, *, allow_tcp: bool = False) -> str:
|
||||||
"""保证 WebSocket 路径以 /mqtt 结尾。"""
|
"""http→ws、https→wss;路径为空或 / 时用 /mqtt,否则保留;mqtt/mqtts 仅显式允许裸 TCP。"""
|
||||||
raw = url.strip()
|
raw = url.strip()
|
||||||
if "://" not in raw:
|
if "://" not in raw:
|
||||||
raw = "ws://" + raw
|
raw = "ws://" + raw
|
||||||
u = urlparse(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 ""
|
path = u.path or ""
|
||||||
if not path.endswith("/mqtt"):
|
if path in ("", "/"):
|
||||||
path = path.rstrip("/") + "/mqtt"
|
path = "/mqtt"
|
||||||
return urlunparse((u.scheme, u.netloc, path, u.params, u.query, u.fragment))
|
return urlunparse((scheme, u.netloc, path, u.params, u.query, u.fragment))
|
||||||
|
|||||||
@@ -203,13 +203,21 @@ class PahoTransport:
|
|||||||
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)
|
||||||
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 走默认。
|
# WebSocket 必须显式 transport=websockets;裸 TCP 走默认。
|
||||||
client = PahoClient(
|
client = PahoClient(
|
||||||
callback_api_version=CallbackAPIVersion.VERSION2,
|
callback_api_version=CallbackAPIVersion.VERSION2,
|
||||||
client_id=params.client_id,
|
client_id=params.client_id,
|
||||||
protocol=MQTTv5,
|
protocol=MQTTv5,
|
||||||
transport="tcp" if use_tcp else "websockets",
|
transport="tcp" if use_tcp else "websockets",
|
||||||
|
reconnect_on_failure=False,
|
||||||
)
|
)
|
||||||
client.username_pw_set(params.username, params.password)
|
client.username_pw_set(params.username, params.password)
|
||||||
client.on_connect = self._on_connect
|
client.on_connect = self._on_connect
|
||||||
@@ -231,16 +239,14 @@ class PahoTransport:
|
|||||||
props = None
|
props = None
|
||||||
if use_tcp:
|
if use_tcp:
|
||||||
if not u.port:
|
if not u.port:
|
||||||
port = 8883 if u.scheme == "mqtts" else 1883
|
port = 8883 if scheme == "mqtts" else 1883
|
||||||
if u.scheme == "mqtts":
|
if scheme == "mqtts":
|
||||||
client.tls_set()
|
client.tls_set()
|
||||||
else:
|
else:
|
||||||
path = u.path or "/mqtt"
|
path = u.path or "/mqtt"
|
||||||
if not path.endswith("/mqtt"):
|
|
||||||
path = path.rstrip("/") + "/mqtt"
|
|
||||||
if not u.port:
|
if not u.port:
|
||||||
port = 443 if u.scheme == "wss" else 80
|
port = 443 if scheme == "wss" else 80
|
||||||
if u.scheme == "wss":
|
if scheme == "wss":
|
||||||
client.tls_set()
|
client.tls_set()
|
||||||
client.ws_set_options(path=path, headers={"Sec-WebSocket-Protocol": "mqtt"})
|
client.ws_set_options(path=path, headers={"Sec-WebSocket-Protocol": "mqtt"})
|
||||||
client.connect(
|
client.connect(
|
||||||
@@ -309,11 +315,15 @@ class PahoTransport:
|
|||||||
if self._on_disconnected:
|
if self._on_disconnected:
|
||||||
self._on_disconnected(None, False)
|
self._on_disconnected(None, False)
|
||||||
return
|
return
|
||||||
# 0x8E = 142 Session taken over
|
# 0x8E = 142 Session taken over;0x8B = 139 可重试
|
||||||
if code == 142:
|
if code == 142:
|
||||||
if self._on_disconnected:
|
if self._on_disconnected:
|
||||||
self._on_disconnected("taken_over", True)
|
self._on_disconnected("taken_over", True)
|
||||||
return
|
return
|
||||||
|
if code == 139:
|
||||||
|
if self._on_disconnected:
|
||||||
|
self._on_disconnected("network", False)
|
||||||
|
return
|
||||||
stop, reason = _classify_connack(code)
|
stop, reason = _classify_connack(code)
|
||||||
if self._on_disconnected:
|
if self._on_disconnected:
|
||||||
self._on_disconnected(reason, stop)
|
self._on_disconnected(reason, stop)
|
||||||
|
|||||||
@@ -2,6 +2,8 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import random
|
||||||
|
import time
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from typing import Any, Callable, Optional
|
from typing import Any, Callable, Optional
|
||||||
@@ -173,6 +175,86 @@ class ConnectionEvent:
|
|||||||
reason: str = ""
|
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]
|
SessionHandler = Callable[[str], None]
|
||||||
MessageHandler = Callable[[IncomingMessage], None]
|
MessageHandler = Callable[[IncomingMessage], None]
|
||||||
ReceiptHandler = Callable[[Receipt], None]
|
ReceiptHandler = Callable[[Receipt], None]
|
||||||
|
|||||||
@@ -7,8 +7,9 @@ import threading
|
|||||||
import time
|
import time
|
||||||
import unittest
|
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.protocol import dumps
|
||||||
|
from nixmsg.transport import FakeTransport
|
||||||
|
|
||||||
|
|
||||||
class FakeTransportTests(unittest.TestCase):
|
class FakeTransportTests(unittest.TestCase):
|
||||||
@@ -48,6 +49,7 @@ class FakeTransportTests(unittest.TestCase):
|
|||||||
c = Client(transport=tr)
|
c = Client(transport=tr)
|
||||||
c.on_session(lambda t: tokens.append(t))
|
c.on_session(lambda t: tokens.append(t))
|
||||||
c.connect("ws://example.test/mqtt", "ep1", password="pw")
|
c.connect("ws://example.test/mqtt", "ep1", password="pw")
|
||||||
|
time.sleep(0.05)
|
||||||
self.assertEqual(tokens, ["nst_test_token"])
|
self.assertEqual(tokens, ["nst_test_token"])
|
||||||
self.assertEqual(c.session_token, "nst_test_token")
|
self.assertEqual(c.session_token, "nst_test_token")
|
||||||
c.close()
|
c.close()
|
||||||
|
|||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user