Files
NixMsg/sdk/python/src/nixmsg/client.py
T

1173 lines
46 KiB
Python

"""NixMsg 同步客户端。"""
from __future__ import annotations
import inspect
import logging
import random
import threading
import time
import urllib.error
import urllib.request
from collections import OrderedDict
from dataclasses import dataclass, field
from typing import Any, Optional
from queue import SimpleQueue
from .errors import ClosedError, NixMsgError, NotConnectedError
from .protocol import dumps, down_topic, loads, normalize_mqtt_ws_url, register_url_from_connect, up_topic
from .transport import ConnectParams, FakeTransport, PahoTransport, Transport
from .types import (
BACKOFF_INITIAL_S,
BACKOFF_JITTER,
BACKOFF_MAX_S,
CLIENT_NAME,
CONNECT_TIMEOUT_S,
DEDUP_CAPACITY,
DEFAULT_MAX_BODY,
DEFAULT_MAX_FRAME,
DEFAULT_MAX_META,
INFLIGHT_LIMIT,
MIN_MAX_RECEIVE,
SEND_QUEUE_LIMIT,
STABLE_RESET_S,
Body,
ConnectionEvent,
ConnectionHandler,
ConnectionState,
GroupEvent,
GroupEventHandler,
HelloLimits,
IncomingMessage,
MessageHandler,
PresenceEvent,
PresenceHandler,
Receipt,
ReceiptHandler,
RegisterOptions,
RegisterResult,
RecallResult,
ReconnectBackoff,
RevokedEvent,
RevokedHandler,
SendOptions,
SendResult,
SessionHandler,
Target,
apply_jitter_s,
nominal_delay_s,
)
from .uuid7 import new_uuid7
log = logging.getLogger("nixmsg")
@dataclass
class _PendingReq:
rid: str
frame: dict[str, Any]
event: threading.Event = field(default_factory=threading.Event)
response: Optional[dict[str, Any]] = None
error: Optional[BaseException] = None
is_send: bool = False
message_id: str = ""
@dataclass
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:
def __init__(
self,
*,
transport: Optional[Transport] = None,
auto_ack: bool = True,
max_receive_bytes: int = DEFAULT_MAX_FRAME,
client_name: str = CLIENT_NAME,
connect_timeout_s: float = CONNECT_TIMEOUT_S,
) -> None:
self._transport: Transport = transport or PahoTransport()
self._auto_ack = auto_ack
self._max_receive_bytes = max_receive_bytes
self._client_name = client_name
self._connect_timeout_s = connect_timeout_s
self._session_handler: Optional[SessionHandler] = None
self._message_handler: Optional[MessageHandler] = None
self._receipt_handler: Optional[ReceiptHandler] = None
self._revoked_handler: Optional[RevokedHandler] = None
self._presence_handler: Optional[PresenceHandler] = None
self._group_handler: Optional[GroupEventHandler] = None
self._connection_handler: Optional[ConnectionHandler] = None
self._lock = threading.RLock()
self._state = ConnectionState.OFFLINE
self._stop_reconnect = False
self._closed = False
self._user_close = False
self._url = ""
self._endpoint_id = ""
self._password: Optional[str] = None
self._session_token: Optional[str] = None
self._use_token = False
self._use_tcp = False
self._limits = HelloLimits()
self._clock_skew_ms = 0
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] = []
self._inflight_sends = 0
self._dedup: OrderedDict[str, str] = OrderedDict()
self._receipt_seen: OrderedDict[str, bool] = OrderedDict()
self._watch_ids: Optional[list[str]] = None
self._watch_all = False
self._conn_event = threading.Event()
self._handshake_error: Optional[BaseException] = None
self._worker: Optional[threading.Thread] = None
self._wake = threading.Event()
self._want_connected = False
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)
# ---- 回调注册 ----
def on_session(self, handler: SessionHandler) -> None:
self._session_handler = handler
def on_message(self, handler: MessageHandler) -> None:
self._message_handler = handler
def on_receipt(self, handler: ReceiptHandler) -> None:
self._receipt_handler = handler
def on_revoked(self, handler: RevokedHandler) -> None:
self._revoked_handler = handler
def on_presence(self, handler: PresenceHandler) -> None:
self._presence_handler = handler
def on_group_event(self, handler: GroupEventHandler) -> None:
self._group_handler = handler
def on_connection(self, handler: ConnectionHandler) -> None:
self._connection_handler = handler
# ---- 连接 ----
def connect(
self,
url: str,
endpoint_id: str,
*,
password: Optional[str] = None,
session_token: Optional[str] = None,
use_tcp: bool = False,
wait: bool = True,
) -> 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()
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
self._use_token = session_token is not None and password is None
self._use_tcp = use_tcp
self._stop_reconnect = False
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():
self._worker = threading.Thread(target=self._run_loop, name="nixmsg-client", daemon=True)
self._worker.start()
self._wake.set()
if wait:
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:
with self._lock:
self._stop_reconnect = True
self._want_connected = False
if not self._last_stop_code:
code = getattr(err, "code", None) or "not_connected"
self._last_stop_code = str(code)
self._last_stop_err = err if isinstance(err, NixMsgError) else NixMsgError(
"not_connected", str(err)
)
raise err
if self._state not in (ConnectionState.ONLINE,):
if self._state == ConnectionState.AUTH_FAILED:
raise NixMsgError(
getattr(self, "_auth_reason", "bad_credentials"),
"认证失败",
)
if self._state == ConnectionState.KICKED:
raise NixMsgError("taken_over", "会话被接管")
with self._lock:
self._stop_reconnect = True
self._want_connected = False
self._last_stop_code = "not_connected"
self._last_stop_err = NixMsgError("not_connected", f"连接未成功: {self._state.value}")
raise NixMsgError("not_connected", f"连接未成功: {self._state.value}")
def close(self) -> None:
with self._lock:
self._user_close = True
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:
self._transport.disconnect()
except Exception:
pass
try:
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 as e:
err = e
stop = NixMsgError("logged_out", "已退出登录")
with self._lock:
self._stop_reconnect = True
self._want_connected = False
self._session_token = None
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:
options = options or SendOptions()
if isinstance(body, str):
body = Body(data=body)
elif isinstance(body, bytes):
import base64
body = Body(data=base64.b64encode(body).decode("ascii"), enc="base64")
if options.content_type:
body.content_type = options.content_type
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("queue_full", "发送队列已满")
max_body = self._limits.max_body_bytes or DEFAULT_MAX_BODY
max_meta = self._limits.max_meta_bytes or DEFAULT_MAX_META
max_frame = self._limits.max_frame_bytes or DEFAULT_MAX_FRAME
if body.decoded_size() > max_body:
raise NixMsgError("body_too_large", "正文超限")
meta = options.meta or {}
meta_bytes = dumps({"meta": meta}) # 近似;真正检查序列化后 meta 对象
# 精确:meta 单独序列化
meta_raw = dumps(meta) if meta else b"{}"
if len(meta_raw) > max_meta:
raise NixMsgError("meta_too_large", "自定义字段超限")
msg_id = options.message_id or new_uuid7()
frame: dict[str, Any] = {
"v": 1,
"type": "send",
"id": msg_id,
"to": to.to_dict(),
"body": body.to_dict(),
}
if meta:
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:
# 本机语义时间 + 偏差
local_ms = int(options.send_at * 1000) if options.send_at < 1e12 else int(options.send_at)
frame["send_at_ms"] = local_ms + self._clock_skew_ms
elif options.delay_ms is not None:
frame["delay_ms"] = int(options.delay_ms)
if options.keep:
offline: dict[str, Any] = {"keep": True}
if options.ttl_seconds is not None:
offline["ttl_seconds"] = int(options.ttl_seconds)
frame["offline"] = offline
if not options.receipt:
frame["receipt"] = False
if options.talk_password:
frame["talk_password"] = options.talk_password
# 帧大小本地检查(含 rid 占位)
probe = dict(frame)
probe["rid"] = "0" * 8
raw = dumps(probe)
if len(raw) > max_frame:
raise NixMsgError("frame_too_large", "整帧超限")
pending = _PendingReq(rid="", frame=frame, is_send=True, message_id=msg_id)
item = _SendItem(message_id=msg_id, frame=frame, pending=pending)
self._send_queue.append(item)
self._wake.set()
if not pending.event.wait(timeout=None if self._state == ConnectionState.ONLINE else 3600):
raise NixMsgError("busy", "发送等待中断")
if pending.error:
raise pending.error
assert pending.response is not None
data = pending.response.get("data") or {}
return SendResult(id=data.get("id", msg_id), send_at_ms=int(data.get("send_at_ms", 0)), state=str(data.get("state", "")))
def ack(self, message: IncomingMessage) -> None:
self._send_ack(message.from_id, message.id, mark_acked=True)
# ---- 其余接口 ----
def recall(self, message_id: str) -> RecallResult:
resp = self._request({"type": "recall", "id": message_id})
data = resp.get("data") or {}
return RecallResult(
result=str(data.get("result", "")),
recalled=int(data.get("recalled", 0)),
accepted=int(data.get("accepted", 0)),
other=int(data.get("other", 0)),
)
def status(self, message_id: str, cursor: str = "", limit: int = 100) -> dict[str, Any]:
return self._request({"type": "status", "id": message_id, "cursor": cursor, "limit": limit})
def unlock(self, endpoint_id: str, talk_password: str) -> dict[str, Any]:
return self._request({"type": "unlock", "endpoint_id": endpoint_id, "talk_password": talk_password})
def presence(self, ids: list[str]) -> dict[str, Any]:
return self._request({"type": "presence.get", "ids": ids})
def directory(self, cursor: str = "", query: str = "", limit: int = 100) -> dict[str, Any]:
return self._request({"type": "directory.list", "cursor": cursor, "limit": limit, "query": query})
def watch_presence(self, ids: Optional[list[str]] = None, *, all: bool = False) -> dict[str, Any]:
with self._lock:
self._watch_ids = list(ids) if ids is not None else None
self._watch_all = all
frame: dict[str, Any] = {"type": "presence.watch", "all": all}
if ids is not None:
frame["ids"] = ids
return self._request(frame)
def get_self(self) -> dict[str, Any]:
return self._request({"type": "self.get"})
def update_self(self, *, name: Optional[str] = None, default_delay_ms: Optional[int] = None) -> dict[str, Any]:
frame: dict[str, Any] = {"type": "self.update"}
if name is not None:
frame["name"] = name
if default_delay_ms is not None:
frame["default_delay_ms"] = default_delay_ms
return self._request(frame)
def set_talk_password(self, talk_password: str) -> dict[str, Any]:
return self._request({"type": "self.talk_password", "talk_password": talk_password})
def change_login_password(self, old_password: str, new_password: str) -> dict[str, Any]:
resp = self._request({"type": "self.login_password", "old_password": old_password, "new_password": new_password})
data = resp.get("data") or {}
token = data.get("session_token")
if token:
with self._lock:
self._session_token = token
self._use_token = True
self._fire_session(token)
return resp
def group_create(self, name: str, members: list[dict[str, str]], group_id: str = "") -> dict[str, Any]:
frame: dict[str, Any] = {"type": "group.create", "name": name, "members": members}
if group_id:
frame["id"] = group_id
else:
frame["id"] = ""
return self._request(frame)
def group_add(self, group_id: str, members: list[dict[str, str]]) -> dict[str, Any]:
return self._request({"type": "group.add", "group_id": group_id, "members": members})
def group_remove(self, group_id: str, endpoint_id: str) -> dict[str, Any]:
return self._request({"type": "group.remove", "group_id": group_id, "endpoint_id": endpoint_id})
def group_leave(self, group_id: str) -> dict[str, Any]:
return self._request({"type": "group.leave", "group_id": group_id})
def group_transfer(self, group_id: str, endpoint_id: str) -> dict[str, Any]:
return self._request({"type": "group.transfer", "group_id": group_id, "endpoint_id": endpoint_id})
def group_rename(self, group_id: str, name: str) -> dict[str, Any]:
return self._request({"type": "group.rename", "group_id": group_id, "name": name})
def group_dissolve(self, group_id: str) -> dict[str, Any]:
return self._request({"type": "group.dissolve", "group_id": group_id})
def group_list(self, cursor: str = "", limit: int = 100) -> dict[str, Any]:
return self._request({"type": "group.list", "cursor": cursor, "limit": limit})
def group_get(self, group_id: str, cursor: str = "", limit: int = 100) -> dict[str, Any]:
return self._request({"type": "group.get", "group_id": group_id, "cursor": cursor, "limit": limit})
@staticmethod
def register(url: str, registration_code: str, options: Optional[RegisterOptions] = None) -> RegisterResult:
options = options or RegisterOptions()
reg_url = register_url_from_connect(url)
body = {
"registration_code": registration_code,
"id": options.id or "",
"login_password": options.login_password or "",
"name": options.name or "",
"talk_password": options.talk_password or "",
}
raw = dumps(body)
req = urllib.request.Request(
reg_url,
data=raw,
headers={"Content-Type": "application/json"},
method="POST",
)
try:
with urllib.request.urlopen(req, timeout=30) as resp:
data = loads(resp.read())
except urllib.error.HTTPError as e:
try:
data = loads(e.read())
except Exception:
raise NixMsgError("bad_request", f"HTTP {e.code}") from e
err = data.get("error") or {}
raise NixMsgError(str(err.get("code", "bad_request")), str(err.get("message", ""))) from e
if not data.get("ok"):
err = data.get("error") or {}
raise NixMsgError(str(err.get("code", "bad_request")), str(err.get("message", "")))
d = data.get("data") or {}
return RegisterResult(id=str(d.get("id", "")), login_password=d.get("login_password"))
# ---- 内部:连接循环 ----
def _run_loop(self) -> None:
while True:
with self._lock:
if self._closed and not self._want_connected:
return
want = self._want_connected and not self._stop_reconnect
state = self._state
if not want:
self._wake.wait(0.5)
self._wake.clear()
continue
if state == ConnectionState.ONLINE:
self._pump_sends()
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
self._set_state(ConnectionState.RECONNECTING)
self._wake.wait(0.05)
self._wake.clear()
def _attempt_connect(self) -> None:
with self._lock:
if self._stop_reconnect or not self._want_connected:
return
self._set_state(ConnectionState.CONNECTING if self._state == ConnectionState.OFFLINE else ConnectionState.RECONNECTING)
self._handshake_error = None
url = self._url
eid = self._endpoint_id
if self._session_token:
cred = self._session_token
using_token = True
else:
cred = self._password or ""
using_token = False
self._connecting_with_token = using_token
use_tcp = self._use_tcp
timeout = self._connect_timeout_s
self._conn_event.clear()
params = ConnectParams(
url=url,
client_id=eid,
username=eid,
password=cred,
clean_start=True,
session_expiry=0,
timeout_s=timeout,
use_tcp=use_tcp,
)
self._transport.connect(params)
# 等待握手完成或失败
ok = self._conn_event.wait(timeout)
if not ok:
try:
self._transport.disconnect()
except Exception:
pass
with self._lock:
self._handshake_error = NixMsgError("not_connected", "连接超时")
return
def _on_transport_connected(self) -> None:
try:
topic = down_topic(self._endpoint_id)
self._transport.subscribe(topic)
rid = self._next_rid()
hello = {
"v": 1,
"type": "hello",
"rid": rid,
"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:
self._pending[rid] = pending
self._transport.publish(up_topic(self._endpoint_id), dumps(hello))
if not pending.event.wait(self._connect_timeout_s):
raise NixMsgError("not_connected", "握手超时")
if pending.error:
raise pending.error
assert pending.response is not None
if not pending.response.get("ok"):
err = (pending.response.get("error") or {})
raise NixMsgError(str(err.get("code", "bad_request")), str(err.get("message", "")))
t1 = time.time()
data = pending.response.get("data") or {}
limits = HelloLimits(
server_time_ms=int(data.get("server_time_ms", 0)),
server_version=str(data.get("server_version", "")),
max_body_bytes=int(data.get("max_body_bytes", DEFAULT_MAX_BODY)),
max_meta_bytes=int(data.get("max_meta_bytes", DEFAULT_MAX_META)),
max_frame_bytes=int(data.get("max_frame_bytes", DEFAULT_MAX_FRAME)),
max_ttl_seconds=int(data.get("max_ttl_seconds", 2592000)),
max_schedule_seconds=int(data.get("max_schedule_seconds", 31536000)),
ack_timeout_seconds=int(data.get("ack_timeout_seconds", 300)),
session_token=str(data.get("session_token", "")),
)
skew = limits.server_time_ms - int(((t0 + t1) / 2) * 1000)
with self._lock:
self._limits = limits
self._clock_skew_ms = skew
self._handshake_error = None
self._backoff.mark_online()
self._set_state(ConnectionState.ONLINE)
token = limits.session_token
if token:
with self._lock:
self._session_token = token
self._use_token = True
self._fire_session(token)
# 重连后恢复 presence.watch
if self._watch_all or self._watch_ids is not None:
try:
frame: dict[str, Any] = {"type": "presence.watch", "all": self._watch_all}
if self._watch_ids is not None:
frame["ids"] = self._watch_ids
self._request(frame, wait=False)
except Exception:
pass
self._conn_event.set()
self._wake.set()
except Exception as e:
with self._lock:
self._handshake_error = e
try:
self._transport.disconnect()
except Exception:
pass
self._conn_event.set()
def _on_transport_disconnected(self, reason: Optional[str], stop: bool) -> None:
with self._lock:
was_online = self._state == ConnectionState.ONLINE
using_token = getattr(self, "_connecting_with_token", self._use_token)
if reason == "taken_over":
self._stop_reconnect = True
self._want_connected = False
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()
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:
auth_reason = "session_invalid"
if auth_reason == "banned":
auth_reason = "session_invalid" if using_token else "bad_credentials"
self._auth_reason = auth_reason
self._stop_reconnect = True
self._want_connected = False
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()
elif self._user_close or self._closed:
self._set_state(ConnectionState.OFFLINE)
self._conn_event.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 队列,否则自死锁。
try:
frame = loads(payload)
except Exception:
self._down_q.put(payload)
return
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
self._down_q.put(payload)
def _down_loop(self) -> None:
while True:
payload = self._down_q.get()
if payload is None:
return
try:
frame = loads(payload)
except Exception:
continue
try:
self._dispatch_down_body(frame)
except Exception:
log.exception("处理下行帧失败")
def _dispatch_resp(self, frame: dict[str, Any]) -> None:
rid = str(frame.get("rid", ""))
with self._lock:
pending = self._pending.pop(rid, None)
if not pending:
return
pending.response = frame
if pending.is_send:
with self._lock:
self._inflight_sends = max(0, self._inflight_sends - 1)
err = (frame.get("error") or {}) if not frame.get("ok") else {}
if not frame.get("ok") and str(err.get("code")) == "rate_limited":
with self._lock:
pending.rid = ""
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:
self._send_queue = [it for it in self._send_queue if it.pending is not pending]
if not frame.get("ok"):
pending.error = NixMsgError(str(err.get("code", "bad_request")), str(err.get("message", "")))
pending.event.set()
self._wake.set()
else:
pending.event.set()
def _dispatch_down(self, payload: bytes) -> None:
try:
frame = loads(payload)
except Exception:
return
if frame.get("type") == "resp":
self._dispatch_resp(frame)
return
self._dispatch_down_body(frame)
def _dispatch_down_body(self, frame: dict[str, Any]) -> None:
ftype = frame.get("type")
if ftype == "msg":
self._handle_msg(frame)
return
if ftype == "receipt":
self._handle_receipt(frame)
return
if ftype == "revoked":
self._handle_revoked(frame)
return
if ftype == "presence":
self._fire_presence(
PresenceEvent(id=str(frame.get("id", "")), online=bool(frame.get("online")), at_ms=int(frame.get("at_ms", 0)))
)
return
if ftype == "group_event":
self._fire_group(
GroupEvent(
group_id=str(frame.get("group_id", "")),
event=str(frame.get("event", "")),
endpoint_id=str(frame.get("endpoint_id", "")),
at_ms=int(frame.get("at_ms", 0)),
)
)
return
if ftype == "fatal":
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", ""))
from_id = str(frame.get("from", ""))
key = f"{from_id}\0{mid}"
with self._lock:
st = self._dedup.get(key)
if st == _DedupState.ACKED:
# 已确认再到达:再 ack,不交应用
pass
elif st == _DedupState.DELIVERED or st == _DedupState.REVOKED:
# 已交未确认或已撤回:忽略
return
else:
st = None
if st == _DedupState.ACKED:
self._send_ack(from_id, mid, mark_acked=True)
return
to_raw = frame.get("to") or {}
msg = IncomingMessage(
id=mid,
from_id=from_id,
to=Target(kind=str(to_raw.get("kind", "endpoint")), id=str(to_raw.get("id", ""))),
body=Body(
data=str((frame.get("body") or {}).get("data", "")),
enc=str((frame.get("body") or {}).get("enc", "utf8")),
content_type=(frame.get("body") or {}).get("content_type"),
),
send_at_ms=int(frame.get("send_at_ms", 0)),
meta=dict(frame.get("meta") or {}),
)
with self._lock:
self._dedup_put(key, _DedupState.DELIVERED)
if not self._message_handler:
if self._auto_ack:
with self._lock:
self._dedup_put(key, _DedupState.ACKED)
self._send_ack(from_id, mid, mark_acked=False)
return
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:
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:
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", "")),
endpoint_id=str(frame.get("endpoint_id", "")),
state=str(frame.get("state", "")),
reason=str(frame.get("reason", "")),
at_ms=int(frame.get("at_ms", 0)),
)
self._fire_receipt(receipt)
# 自动 receipt_ack
try:
self._request({"type": "receipt_ack", "receipt_id": rid}, wait=False)
except Exception:
pass
def _handle_revoked(self, frame: dict[str, Any]) -> None:
mid = str(frame.get("id", ""))
from_id = str(frame.get("from", ""))
key = f"{from_id}\0{mid}"
with self._lock:
st = self._dedup.get(key)
if st == _DedupState.ACKED or st == _DedupState.REVOKED:
return
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 {}
result = data.get("result", "accepted")
if result != "accepted":
with self._lock:
self._dedup.pop(key, None)
self._fire_revoked(RevokedEvent(id=message_id, from_id=from_id, reason=str(result)))
return
if mark_acked:
with self._lock:
self._dedup_put(key, _DedupState.ACKED)
def _pump_sends(self) -> None:
while True:
with self._lock:
if self._state != ConnectionState.ONLINE:
return
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
if item is None:
return
rid = self._next_rid_locked()
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, raw)
except Exception:
with self._lock:
self._pending.pop(rid, None)
item.pending.rid = ""
self._inflight_sends = max(0, self._inflight_sends - 1)
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]:
with self._lock:
if self._closed:
raise ClosedError()
if self._state != ConnectionState.ONLINE:
raise NotConnectedError()
rid = self._next_rid_locked()
body = dict(frame)
body["v"] = 1
body["rid"] = rid
pending = _PendingReq(rid=rid, frame=body)
self._pending[rid] = pending
topic = up_topic(self._endpoint_id)
self._transport.publish(topic, dumps(body))
if not wait:
# 仍等一小会拿结果;后台请求也尽量同步完成
pending.event.wait(timeout=30)
with self._lock:
self._pending.pop(rid, None)
if pending.error:
raise pending.error
if pending.response is None:
return {}
if not pending.response.get("ok"):
err = pending.response.get("error") or {}
raise NixMsgError(str(err.get("code", "bad_request")), str(err.get("message", "")))
return pending.response
if not pending.event.wait(timeout=60):
with self._lock:
self._pending.pop(rid, None)
raise NixMsgError("busy", "请求超时")
if pending.error:
raise pending.error
assert pending.response is not None
if not pending.response.get("ok"):
err = pending.response.get("error") or {}
raise NixMsgError(str(err.get("code", "bad_request")), str(err.get("message", "")))
return pending.response
def _next_rid(self) -> str:
with self._lock:
return self._next_rid_locked()
def _next_rid_locked(self) -> str:
self._rid_seq += 1
return str(self._rid_seq)
def _dedup_put(self, key: str, state: str) -> None:
if key in self._dedup:
self._dedup.move_to_end(key)
self._dedup[key] = state
while len(self._dedup) > DEDUP_CAPACITY:
self._dedup.popitem(last=False)
def _fail_all_pending(self, err: BaseException) -> None:
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.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:
self._cb_q.put(lambda: self._invoke(h, ev))
def _fire_session(self, token: str) -> None:
h = self._session_handler
if h:
self._cb_q.put(lambda: self._invoke(h, token))
def _fire_receipt(self, receipt: Receipt) -> None:
h = self._receipt_handler
if h:
self._cb_q.put(lambda: self._invoke(h, receipt))
def _fire_revoked(self, ev: RevokedEvent) -> None:
h = self._revoked_handler
if h:
self._cb_q.put(lambda: self._invoke(h, ev))
def _fire_presence(self, ev: PresenceEvent) -> None:
h = self._presence_handler
if h:
self._cb_q.put(lambda: self._invoke(h, ev))
def _fire_group(self, ev: GroupEvent) -> None:
h = self._group_handler
if h:
self._cb_q.put(lambda: self._invoke(h, ev))
@property
def last_stop_code(self) -> str:
return self._last_stop_code
# 测试辅助
@property
def state(self) -> ConnectionState:
return self._state
@property
def limits(self) -> HelloLimits:
return self._limits
@property
def clock_skew_ms(self) -> int:
return self._clock_skew_ms
@property
def session_token(self) -> Optional[str]:
return self._session_token
def _auto_hello_responder(client: Client, transport: FakeTransport, *, session_token: str = "nst_test") -> None:
"""测试辅助:对 hello 自动回成功(由测试自行调用更清晰)。"""
_ = (client, transport, session_token)