fix: 按 K-00 约定修复 Python SDK 断线重交与退避
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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]
|
||||
|
||||
Reference in New Issue
Block a user