1159 lines
45 KiB
Python
1159 lines
45 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:
|
|
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", "会话被接管")
|
|
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("busy", "连接超时")
|
|
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("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", "")))
|
|
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)
|