feat: 实现 Python 与 Java SDK 连接收发及其余接口
This commit is contained in:
@@ -0,0 +1,7 @@
|
||||
.venv/
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
*.egg-info/
|
||||
.pytest_cache/
|
||||
dist/
|
||||
build/
|
||||
@@ -0,0 +1,10 @@
|
||||
Copyright (c) 2026 Nixevol. All rights reserved.
|
||||
|
||||
本仓库的源代码、文档、各语言 SDK 和构建产物(包括发布的软件包和 Docker 镜像)均为专有软件。
|
||||
源代码和发布物公开可读,不代表授予任何使用许可。未经版权所有者书面许可,不得使用、复制、
|
||||
修改、合并、发布、分发、再许可或出售其任何部分。
|
||||
|
||||
This repository, including its source code, documentation, SDKs and build artifacts (including
|
||||
published packages and Docker images), is proprietary software. Public visibility does not grant
|
||||
any license. No part of it may be used, copied, modified, merged, published, distributed,
|
||||
sublicensed or sold without prior written permission from the copyright holder.
|
||||
@@ -0,0 +1,24 @@
|
||||
# NixMsg Python SDK
|
||||
|
||||
包名 `nixmsg`,最低 Python 3.10。同步接口为主,`AsyncClient` 提供 asyncio 包装。
|
||||
|
||||
## 安装
|
||||
|
||||
```bash
|
||||
pip install nixmsg --index-url https://git.asio.asia/api/packages/nixevol/pypi/simple/
|
||||
```
|
||||
|
||||
## 最小示例
|
||||
|
||||
```python
|
||||
from nixmsg import Client, Target, Body
|
||||
|
||||
c = Client()
|
||||
c.on_session(lambda token: print("session", token))
|
||||
c.on_message(lambda msg: print("msg", msg.id, msg.body.data))
|
||||
c.connect("ws://127.0.0.1:7443/mqtt", "device-1", password="secret")
|
||||
c.send(Target(kind="endpoint", id="device-2"), Body(data="hello"))
|
||||
c.close()
|
||||
```
|
||||
|
||||
许可证见 `LICENSE`(专有)。
|
||||
@@ -0,0 +1,33 @@
|
||||
[build-system]
|
||||
requires = ["setuptools>=68", "wheel"]
|
||||
build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "nixmsg"
|
||||
version = "0.1.0"
|
||||
description = "NixMsg Python SDK"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10"
|
||||
license = { file = "LICENSE" }
|
||||
authors = [{ name = "Nixevol" }]
|
||||
classifiers = [
|
||||
"License :: Other/Proprietary License",
|
||||
"Programming Language :: Python :: 3",
|
||||
"Programming Language :: Python :: 3.10",
|
||||
"Programming Language :: Python :: 3.11",
|
||||
"Programming Language :: Python :: 3.12",
|
||||
"Programming Language :: Python :: 3.13",
|
||||
]
|
||||
dependencies = [
|
||||
"paho-mqtt>=2.0,<3",
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
dev = ["pytest>=7"]
|
||||
|
||||
[tool.setuptools.packages.find]
|
||||
where = ["src"]
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
testpaths = ["tests"]
|
||||
pythonpath = ["src"]
|
||||
@@ -0,0 +1,50 @@
|
||||
"""NixMsg Python SDK。"""
|
||||
|
||||
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 .types import (
|
||||
Body,
|
||||
ConnectionEvent,
|
||||
ConnectionState,
|
||||
GroupEvent,
|
||||
IncomingMessage,
|
||||
PresenceEvent,
|
||||
Receipt,
|
||||
RegisterOptions,
|
||||
RegisterResult,
|
||||
RecallResult,
|
||||
RevokedEvent,
|
||||
SendOptions,
|
||||
SendResult,
|
||||
Target,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"AsyncClient",
|
||||
"Body",
|
||||
"Client",
|
||||
"ClosedError",
|
||||
"ConnectionEvent",
|
||||
"ConnectionState",
|
||||
"FakeTransport",
|
||||
"GroupEvent",
|
||||
"IncomingMessage",
|
||||
"NixMsgError",
|
||||
"NotConnectedError",
|
||||
"PahoTransport",
|
||||
"PresenceEvent",
|
||||
"Receipt",
|
||||
"RegisterOptions",
|
||||
"RegisterResult",
|
||||
"RecallResult",
|
||||
"RevokedEvent",
|
||||
"SendOptions",
|
||||
"SendResult",
|
||||
"Target",
|
||||
"register_url_from_connect",
|
||||
]
|
||||
|
||||
__version__ = "0.1.0"
|
||||
@@ -0,0 +1,124 @@
|
||||
"""asyncio 包装(同一包内)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from typing import Any, Optional
|
||||
|
||||
from .client import Client
|
||||
from .types import (
|
||||
Body,
|
||||
RegisterOptions,
|
||||
RegisterResult,
|
||||
SendOptions,
|
||||
SendResult,
|
||||
Target,
|
||||
)
|
||||
|
||||
|
||||
class AsyncClient:
|
||||
"""把同步 Client 的阻塞调用丢到线程池。"""
|
||||
|
||||
def __init__(self, client: Optional[Client] = None, **kwargs: Any) -> None:
|
||||
self._client = client or Client(**kwargs)
|
||||
|
||||
@property
|
||||
def sync(self) -> Client:
|
||||
return self._client
|
||||
|
||||
def on_session(self, handler): # noqa: ANN001
|
||||
self._client.on_session(handler)
|
||||
|
||||
def on_message(self, handler): # noqa: ANN001
|
||||
self._client.on_message(handler)
|
||||
|
||||
def on_receipt(self, handler): # noqa: ANN001
|
||||
self._client.on_receipt(handler)
|
||||
|
||||
def on_revoked(self, handler): # noqa: ANN001
|
||||
self._client.on_revoked(handler)
|
||||
|
||||
def on_presence(self, handler): # noqa: ANN001
|
||||
self._client.on_presence(handler)
|
||||
|
||||
def on_group_event(self, handler): # noqa: ANN001
|
||||
self._client.on_group_event(handler)
|
||||
|
||||
def on_connection(self, handler): # noqa: ANN001
|
||||
self._client.on_connection(handler)
|
||||
|
||||
async def connect(self, url: str, endpoint_id: str, **kwargs: Any) -> None:
|
||||
await asyncio.to_thread(self._client.connect, url, endpoint_id, **kwargs)
|
||||
|
||||
async def close(self) -> None:
|
||||
await asyncio.to_thread(self._client.close)
|
||||
|
||||
async def logout(self) -> None:
|
||||
await asyncio.to_thread(self._client.logout)
|
||||
|
||||
async def send(self, to: Target, body: Body | str | bytes, options: Optional[SendOptions] = None) -> SendResult:
|
||||
return await asyncio.to_thread(self._client.send, to, body, options)
|
||||
|
||||
async def ack(self, message) -> None: # noqa: ANN001
|
||||
await asyncio.to_thread(self._client.ack, message)
|
||||
|
||||
async def recall(self, message_id: str):
|
||||
return await asyncio.to_thread(self._client.recall, message_id)
|
||||
|
||||
async def status(self, message_id: str, cursor: str = "", limit: int = 100):
|
||||
return await asyncio.to_thread(self._client.status, message_id, cursor, limit)
|
||||
|
||||
async def unlock(self, endpoint_id: str, talk_password: str):
|
||||
return await asyncio.to_thread(self._client.unlock, endpoint_id, talk_password)
|
||||
|
||||
async def presence(self, ids: list[str]):
|
||||
return await asyncio.to_thread(self._client.presence, ids)
|
||||
|
||||
async def directory(self, cursor: str = "", query: str = "", limit: int = 100):
|
||||
return await asyncio.to_thread(self._client.directory, cursor, query, limit)
|
||||
|
||||
async def watch_presence(self, ids: Optional[list[str]] = None, *, all: bool = False):
|
||||
return await asyncio.to_thread(self._client.watch_presence, ids, all=all)
|
||||
|
||||
async def get_self(self):
|
||||
return await asyncio.to_thread(self._client.get_self)
|
||||
|
||||
async def update_self(self, **kwargs: Any):
|
||||
return await asyncio.to_thread(self._client.update_self, **kwargs)
|
||||
|
||||
async def set_talk_password(self, talk_password: str):
|
||||
return await asyncio.to_thread(self._client.set_talk_password, talk_password)
|
||||
|
||||
async def change_login_password(self, old_password: str, new_password: str):
|
||||
return await asyncio.to_thread(self._client.change_login_password, old_password, new_password)
|
||||
|
||||
async def group_create(self, name: str, members: list[dict[str, str]], group_id: str = ""):
|
||||
return await asyncio.to_thread(self._client.group_create, name, members, group_id)
|
||||
|
||||
async def group_add(self, group_id: str, members: list[dict[str, str]]):
|
||||
return await asyncio.to_thread(self._client.group_add, group_id, members)
|
||||
|
||||
async def group_remove(self, group_id: str, endpoint_id: str):
|
||||
return await asyncio.to_thread(self._client.group_remove, group_id, endpoint_id)
|
||||
|
||||
async def group_leave(self, group_id: str):
|
||||
return await asyncio.to_thread(self._client.group_leave, group_id)
|
||||
|
||||
async def group_transfer(self, group_id: str, endpoint_id: str):
|
||||
return await asyncio.to_thread(self._client.group_transfer, group_id, endpoint_id)
|
||||
|
||||
async def group_rename(self, group_id: str, name: str):
|
||||
return await asyncio.to_thread(self._client.group_rename, group_id, name)
|
||||
|
||||
async def group_dissolve(self, group_id: str):
|
||||
return await asyncio.to_thread(self._client.group_dissolve, group_id)
|
||||
|
||||
async def group_list(self, cursor: str = "", limit: int = 100):
|
||||
return await asyncio.to_thread(self._client.group_list, cursor, limit)
|
||||
|
||||
async def group_get(self, group_id: str, cursor: str = "", limit: int = 100):
|
||||
return await asyncio.to_thread(self._client.group_get, group_id, cursor, limit)
|
||||
|
||||
@staticmethod
|
||||
async def register(url: str, registration_code: str, options: Optional[RegisterOptions] = None) -> RegisterResult:
|
||||
return await asyncio.to_thread(Client.register, url, registration_code, options)
|
||||
@@ -0,0 +1,988 @@
|
||||
"""NixMsg 同步客户端。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
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 .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,
|
||||
RevokedEvent,
|
||||
RevokedHandler,
|
||||
SendOptions,
|
||||
SendResult,
|
||||
SessionHandler,
|
||||
Target,
|
||||
)
|
||||
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
|
||||
|
||||
|
||||
class _DedupState:
|
||||
DELIVERED = "delivered" # 已交应用未确认
|
||||
ACKED = "acked"
|
||||
|
||||
|
||||
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(MIN_MAX_RECEIVE, 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._cb_lock = threading.Lock() # 回调串行
|
||||
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._online_since = 0.0
|
||||
self._backoff_s = BACKOFF_INITIAL_S
|
||||
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._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")
|
||||
with self._lock:
|
||||
if self._closed:
|
||||
raise ClosedError()
|
||||
self._url = normalize_mqtt_ws_url(url) if not use_tcp else url
|
||||
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._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 + 5):
|
||||
raise NixMsgError("busy", "连接超时")
|
||||
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("busy", 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._fail_all_pending(ClosedError())
|
||||
self._set_state(ConnectionState.OFFLINE)
|
||||
try:
|
||||
self._transport.disconnect()
|
||||
except Exception:
|
||||
pass
|
||||
self._wake.set()
|
||||
|
||||
def logout(self) -> None:
|
||||
try:
|
||||
self._request({"type": "self.logout"}, wait=True)
|
||||
except Exception:
|
||||
pass
|
||||
with self._lock:
|
||||
self._stop_reconnect = True
|
||||
self._want_connected = False
|
||||
self._session_token = None
|
||||
self._fail_all_pending(NixMsgError("auth_failed", "已退出登录"))
|
||||
try:
|
||||
self._transport.disconnect()
|
||||
except Exception:
|
||||
pass
|
||||
self._set_state(ConnectionState.OFFLINE)
|
||||
self._wake.set()
|
||||
|
||||
# ---- 发送 / 确认 ----
|
||||
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 len(self._send_queue) >= SEND_QUEUE_LIMIT:
|
||||
raise NixMsgError("quota_exceeded", "发送队列已满")
|
||||
|
||||
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:
|
||||
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()
|
||||
# 稳定 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
|
||||
# 尝试连接
|
||||
try:
|
||||
self._attempt_connect()
|
||||
except Exception as e:
|
||||
log.debug("connect attempt failed: %s", e)
|
||||
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.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,
|
||||
"max_receive_bytes": self._max_receive_bytes,
|
||||
"client": self._client_name,
|
||||
}
|
||||
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._online_since = time.monotonic()
|
||||
self._handshake_error = None
|
||||
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._fail_all_pending(NixMsgError("taken_over", "会话被接管"))
|
||||
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
|
||||
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._fail_all_pending(NixMsgError(auth_reason, "认证失败"))
|
||||
self._handshake_error = NixMsgError(auth_reason, "认证失败")
|
||||
self._set_state(ConnectionState.AUTH_FAILED, auth_reason)
|
||||
self._conn_event.set()
|
||||
return
|
||||
# 网络 / busy:继续重连
|
||||
if 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()
|
||||
|
||||
def _on_down(self, payload: bytes) -> None:
|
||||
try:
|
||||
frame = loads(payload)
|
||||
except Exception:
|
||||
return
|
||||
ftype = frame.get("type")
|
||||
if ftype == "resp":
|
||||
rid = str(frame.get("rid", ""))
|
||||
with self._lock:
|
||||
pending = self._pending.pop(rid, None)
|
||||
if pending:
|
||||
pending.response = frame
|
||||
if pending.is_send:
|
||||
with self._lock:
|
||||
self._inflight_sends = max(0, self._inflight_sends - 1)
|
||||
# rate_limited 重交:清状态后不 set event
|
||||
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()
|
||||
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()
|
||||
return
|
||||
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":
|
||||
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
|
||||
|
||||
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:
|
||||
# 已交未确认:忽略
|
||||
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:
|
||||
self._send_ack(from_id, mid, mark_acked=True)
|
||||
return
|
||||
|
||||
try:
|
||||
with self._cb_lock:
|
||||
self._message_handler(msg)
|
||||
except Exception:
|
||||
log.exception("on_message 回调错误,等待重推")
|
||||
with self._lock:
|
||||
self._dedup.pop(key, None)
|
||||
return
|
||||
|
||||
if self._auto_ack:
|
||||
self._send_ack(from_id, mid, mark_acked=True)
|
||||
|
||||
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)
|
||||
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:
|
||||
return
|
||||
if st is None:
|
||||
# 还没交给应用:直接丢弃
|
||||
return
|
||||
# 已交未确认:发撤回事件并不再确认
|
||||
self._dedup.pop(key, None)
|
||||
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:
|
||||
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
|
||||
for it in self._send_queue:
|
||||
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
|
||||
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:
|
||||
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]
|
||||
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:
|
||||
for p in list(self._pending.values()):
|
||||
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._inflight_sends = 0
|
||||
|
||||
def _set_state(self, state: ConnectionState, reason: str = "") -> None:
|
||||
self._state = state
|
||||
h = self._connection_handler
|
||||
if h:
|
||||
try:
|
||||
with self._cb_lock:
|
||||
h(ConnectionEvent(state=state, reason=reason))
|
||||
except Exception:
|
||||
log.exception("on_connection 回调错误")
|
||||
|
||||
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 回调错误")
|
||||
|
||||
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 回调错误")
|
||||
|
||||
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 回调错误")
|
||||
|
||||
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 回调错误")
|
||||
|
||||
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 回调错误")
|
||||
|
||||
@staticmethod
|
||||
def _jitter(base: float) -> float:
|
||||
return max(0.0, base * (1.0 + random.uniform(-BACKOFF_JITTER, BACKOFF_JITTER)))
|
||||
|
||||
# 测试辅助
|
||||
@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)
|
||||
@@ -0,0 +1,20 @@
|
||||
"""NixMsg SDK 错误。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
class NixMsgError(Exception):
|
||||
def __init__(self, code: str, message: str = "") -> None:
|
||||
self.code = code
|
||||
self.message = message or code
|
||||
super().__init__(self.message)
|
||||
|
||||
|
||||
class NotConnectedError(NixMsgError):
|
||||
def __init__(self, message: str = "未连接") -> None:
|
||||
super().__init__("not_connected", message)
|
||||
|
||||
|
||||
class ClosedError(NixMsgError):
|
||||
def __init__(self, message: str = "已关闭") -> None:
|
||||
super().__init__("closed", message)
|
||||
@@ -0,0 +1,58 @@
|
||||
"""帧编解码与注册地址推导。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
from urllib.parse import urlparse, urlunparse
|
||||
|
||||
|
||||
def dumps(obj: dict[str, Any]) -> bytes:
|
||||
return json.dumps(obj, ensure_ascii=False, separators=(",", ":")).encode("utf-8")
|
||||
|
||||
|
||||
def loads(data: bytes | str) -> dict[str, Any]:
|
||||
if isinstance(data, bytes):
|
||||
data = data.decode("utf-8")
|
||||
return json.loads(data)
|
||||
|
||||
|
||||
def register_url_from_connect(connect_url: str) -> str:
|
||||
"""从连接地址推出注册 HTTP 地址(DEVELOPMENT 6.9)。"""
|
||||
raw = connect_url.strip()
|
||||
if "://" not in raw:
|
||||
raw = "ws://" + raw
|
||||
u = urlparse(raw)
|
||||
scheme = u.scheme.lower()
|
||||
if scheme in ("wss", "mqtts", "https"):
|
||||
http_scheme = "https"
|
||||
elif scheme in ("ws", "mqtt", "http"):
|
||||
http_scheme = "http"
|
||||
else:
|
||||
http_scheme = "https" if scheme.endswith("s") else "http"
|
||||
# 去掉 /mqtt 路径
|
||||
path = u.path or ""
|
||||
if path.endswith("/mqtt"):
|
||||
path = path[: -len("/mqtt")]
|
||||
path = path.rstrip("/") + "/api/client/register"
|
||||
return urlunparse((http_scheme, u.netloc, path, "", "", ""))
|
||||
|
||||
|
||||
def up_topic(endpoint_id: str) -> str:
|
||||
return f"nix/c/{endpoint_id}/up"
|
||||
|
||||
|
||||
def down_topic(endpoint_id: str) -> str:
|
||||
return f"nix/c/{endpoint_id}/down"
|
||||
|
||||
|
||||
def normalize_mqtt_ws_url(url: str) -> str:
|
||||
"""保证 WebSocket 路径以 /mqtt 结尾。"""
|
||||
raw = url.strip()
|
||||
if "://" not in raw:
|
||||
raw = "ws://" + raw
|
||||
u = urlparse(raw)
|
||||
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))
|
||||
@@ -0,0 +1,350 @@
|
||||
"""MQTT 传输抽象、假传输与 paho 实现。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Callable, Optional, Protocol
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from paho.mqtt.client import CallbackAPIVersion, Client as PahoClient, MQTT_ERR_SUCCESS
|
||||
from paho.mqtt.enums import MQTTErrorCode
|
||||
from paho.mqtt.reasoncodes import ReasonCode
|
||||
|
||||
DownHandler = Callable[[bytes], None]
|
||||
ConnHandler = Callable[[], None]
|
||||
DiscHandler = Callable[[Optional[str], bool], None]
|
||||
# reason_code_str, stop_reconnect
|
||||
|
||||
|
||||
@dataclass
|
||||
class ConnectParams:
|
||||
url: str
|
||||
client_id: str
|
||||
username: str
|
||||
password: str
|
||||
clean_start: bool = True
|
||||
session_expiry: int = 0
|
||||
keep_alive: int = 30
|
||||
timeout_s: float = 30.0
|
||||
use_tcp: bool = False # 显式打开裸 TCP
|
||||
|
||||
|
||||
class Transport(Protocol):
|
||||
def set_handlers(
|
||||
self,
|
||||
on_connected: ConnHandler,
|
||||
on_disconnected: DiscHandler,
|
||||
on_down: DownHandler,
|
||||
) -> None: ...
|
||||
|
||||
def connect(self, params: ConnectParams) -> None: ...
|
||||
|
||||
def subscribe(self, topic: str) -> None: ...
|
||||
|
||||
def publish(self, topic: str, payload: bytes) -> None: ...
|
||||
|
||||
def disconnect(self) -> None: ...
|
||||
|
||||
|
||||
@dataclass
|
||||
class FakeConnectRecord:
|
||||
params: ConnectParams
|
||||
at: float = field(default_factory=time.time)
|
||||
|
||||
|
||||
class FakeTransport:
|
||||
"""单元测试用假传输:不连真实服务器。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._on_connected: Optional[ConnHandler] = None
|
||||
self._on_disconnected: Optional[DiscHandler] = None
|
||||
self._on_down: Optional[DownHandler] = None
|
||||
self.connects: list[FakeConnectRecord] = []
|
||||
self.publishes: list[tuple[str, bytes]] = []
|
||||
self.subscriptions: list[str] = []
|
||||
self.connected = False
|
||||
self.auto_accept = True
|
||||
self.next_connack_fail: Optional[str] = None # session_invalid / bad_credentials / busy
|
||||
self.auto_hello: Optional[dict] = {
|
||||
"server_time_ms": 1_750_000_000_000,
|
||||
"server_version": "0.1.0",
|
||||
"max_body_bytes": 262144,
|
||||
"max_meta_bytes": 4096,
|
||||
"max_frame_bytes": 786432,
|
||||
"max_ttl_seconds": 2592000,
|
||||
"max_schedule_seconds": 31536000,
|
||||
"ack_timeout_seconds": 300,
|
||||
"session_token": "nst_test_token",
|
||||
}
|
||||
self.auto_send_ok = True
|
||||
self._up_handlers: list[Callable[[dict], Optional[dict]]] = []
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def on_up(self, handler: Callable[[dict], Optional[dict]]) -> None:
|
||||
"""上行帧钩子:返回 dict 则作为 resp 注入;返回 None 表示不处理。"""
|
||||
self._up_handlers.append(handler)
|
||||
|
||||
def set_handlers(
|
||||
self,
|
||||
on_connected: ConnHandler,
|
||||
on_disconnected: DiscHandler,
|
||||
on_down: DownHandler,
|
||||
) -> None:
|
||||
self._on_connected = on_connected
|
||||
self._on_disconnected = on_disconnected
|
||||
self._on_down = on_down
|
||||
|
||||
def connect(self, params: ConnectParams) -> None:
|
||||
with self._lock:
|
||||
self.connects.append(FakeConnectRecord(params=params))
|
||||
fail = self.next_connack_fail
|
||||
self.next_connack_fail = None
|
||||
if fail:
|
||||
stop = fail in ("session_invalid", "bad_credentials", "banned")
|
||||
if self._on_disconnected:
|
||||
self._on_disconnected(fail, stop)
|
||||
return
|
||||
self.connected = True
|
||||
if self.auto_accept and self._on_connected:
|
||||
self._on_connected()
|
||||
|
||||
def subscribe(self, topic: str) -> None:
|
||||
self.subscriptions.append(topic)
|
||||
|
||||
def publish(self, topic: str, payload: bytes) -> None:
|
||||
self.publishes.append((topic, payload))
|
||||
import json
|
||||
|
||||
try:
|
||||
frame = json.loads(payload.decode("utf-8"))
|
||||
except Exception:
|
||||
return
|
||||
# 自定义钩子优先
|
||||
for h in list(self._up_handlers):
|
||||
try:
|
||||
resp = h(frame)
|
||||
except Exception:
|
||||
continue
|
||||
if resp is not None:
|
||||
self.inject_down(json.dumps(resp, ensure_ascii=False).encode("utf-8"))
|
||||
return
|
||||
ftype = frame.get("type")
|
||||
rid = frame.get("rid")
|
||||
if ftype == "hello" and self.auto_hello is not None:
|
||||
resp = {"v": 1, "type": "resp", "rid": rid, "ok": True, "data": dict(self.auto_hello)}
|
||||
self.inject_down(json.dumps(resp, ensure_ascii=False).encode("utf-8"))
|
||||
return
|
||||
if ftype == "send" and self.auto_send_ok:
|
||||
resp = {
|
||||
"v": 1,
|
||||
"type": "resp",
|
||||
"rid": rid,
|
||||
"ok": True,
|
||||
"data": {"id": frame.get("id"), "send_at_ms": frame.get("send_at_ms") or 0, "state": "dispatched"},
|
||||
}
|
||||
self.inject_down(json.dumps(resp, ensure_ascii=False).encode("utf-8"))
|
||||
return
|
||||
if ftype == "ack":
|
||||
resp = {"v": 1, "type": "resp", "rid": rid, "ok": True, "data": {"result": "accepted"}}
|
||||
self.inject_down(json.dumps(resp, ensure_ascii=False).encode("utf-8"))
|
||||
return
|
||||
if ftype and ftype not in ("hello", "send") and rid is not None:
|
||||
# 其它请求默认成功
|
||||
resp = {"v": 1, "type": "resp", "rid": rid, "ok": True, "data": {}}
|
||||
self.inject_down(json.dumps(resp, ensure_ascii=False).encode("utf-8"))
|
||||
|
||||
def disconnect(self) -> None:
|
||||
was = self.connected
|
||||
self.connected = False
|
||||
if was and self._on_disconnected:
|
||||
self._on_disconnected(None, False)
|
||||
|
||||
def inject_down(self, payload: bytes) -> None:
|
||||
if self._on_down:
|
||||
self._on_down(payload)
|
||||
|
||||
def simulate_taken_over(self) -> None:
|
||||
self.connected = False
|
||||
if self._on_disconnected:
|
||||
self._on_disconnected("taken_over", True)
|
||||
|
||||
def simulate_network_drop(self) -> None:
|
||||
self.connected = False
|
||||
if self._on_disconnected:
|
||||
self._on_disconnected("network", False)
|
||||
|
||||
|
||||
class PahoTransport:
|
||||
"""paho-mqtt 2.x CallbackAPIVersion.VERSION2。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._client: Optional[PahoClient] = None
|
||||
self._on_connected: Optional[ConnHandler] = None
|
||||
self._on_disconnected: Optional[DiscHandler] = None
|
||||
self._on_down: Optional[DownHandler] = None
|
||||
self._down_topic = ""
|
||||
self._loop_started = False
|
||||
|
||||
def set_handlers(
|
||||
self,
|
||||
on_connected: ConnHandler,
|
||||
on_disconnected: DiscHandler,
|
||||
on_down: DownHandler,
|
||||
) -> None:
|
||||
self._on_connected = on_connected
|
||||
self._on_disconnected = on_disconnected
|
||||
self._on_down = on_down
|
||||
|
||||
def connect(self, params: ConnectParams) -> None:
|
||||
self.disconnect()
|
||||
client = PahoClient(
|
||||
callback_api_version=CallbackAPIVersion.VERSION2,
|
||||
client_id=params.client_id,
|
||||
protocol=PahoClient.MQTTv5,
|
||||
)
|
||||
client.username_pw_set(params.username, params.password)
|
||||
client.on_connect = self._on_connect
|
||||
client.on_disconnect = self._on_disconnect
|
||||
client.on_message = self._on_message
|
||||
self._client = client
|
||||
|
||||
url = params.url
|
||||
u = urlparse(url if "://" in url else "ws://" + url)
|
||||
use_tcp = params.use_tcp or u.scheme in ("mqtt", "mqtts")
|
||||
host = u.hostname or "localhost"
|
||||
port = u.port or (8883 if u.scheme in ("wss", "mqtts") else 443 if u.scheme == "wss" else 80)
|
||||
if use_tcp:
|
||||
if not u.port:
|
||||
port = 8883 if u.scheme == "mqtts" else 1883
|
||||
tls = u.scheme == "mqtts"
|
||||
if tls:
|
||||
client.tls_set()
|
||||
props = None
|
||||
try:
|
||||
from paho.mqtt.properties import Properties
|
||||
from paho.mqtt.packettypes import PacketTypes
|
||||
|
||||
props = Properties(PacketTypes.CONNECT)
|
||||
props.SessionExpiryInterval = params.session_expiry
|
||||
except Exception:
|
||||
props = None
|
||||
client.connect(
|
||||
host,
|
||||
port,
|
||||
keepalive=params.keep_alive,
|
||||
clean_start=params.clean_start,
|
||||
properties=props,
|
||||
)
|
||||
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":
|
||||
client.tls_set()
|
||||
client.ws_set_options(path=path, headers={"Sec-WebSocket-Protocol": "mqtt"})
|
||||
props = None
|
||||
try:
|
||||
from paho.mqtt.properties import Properties
|
||||
from paho.mqtt.packettypes import PacketTypes
|
||||
|
||||
props = Properties(PacketTypes.CONNECT)
|
||||
props.SessionExpiryInterval = params.session_expiry
|
||||
except Exception:
|
||||
props = None
|
||||
client.connect(
|
||||
host,
|
||||
port,
|
||||
keepalive=params.keep_alive,
|
||||
clean_start=params.clean_start,
|
||||
properties=props,
|
||||
)
|
||||
client.loop_start()
|
||||
self._loop_started = True
|
||||
# 等待连接结果由回调驱动;超时由 Client 层处理
|
||||
|
||||
def subscribe(self, topic: str) -> None:
|
||||
self._down_topic = topic
|
||||
if self._client:
|
||||
self._client.subscribe(topic, qos=1)
|
||||
|
||||
def publish(self, topic: str, payload: bytes) -> None:
|
||||
if not self._client:
|
||||
raise RuntimeError("not connected")
|
||||
info = self._client.publish(topic, payload, qos=1)
|
||||
if info.rc != MQTT_ERR_SUCCESS:
|
||||
raise RuntimeError(f"publish failed: {info.rc}")
|
||||
|
||||
def disconnect(self) -> None:
|
||||
c = self._client
|
||||
self._client = None
|
||||
if c is not None:
|
||||
try:
|
||||
c.disconnect()
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
if self._loop_started:
|
||||
c.loop_stop()
|
||||
except Exception:
|
||||
pass
|
||||
self._loop_started = False
|
||||
|
||||
def _on_connect(self, client, userdata, flags, reason_code, properties) -> None:
|
||||
code = _reason_to_int(reason_code)
|
||||
if code == 0:
|
||||
if self._on_connected:
|
||||
self._on_connected()
|
||||
return
|
||||
stop, reason = _classify_connack(code)
|
||||
if self._on_disconnected:
|
||||
self._on_disconnected(reason, stop)
|
||||
|
||||
def _on_disconnect(self, client, userdata, flags, reason_code, properties) -> None:
|
||||
code = _reason_to_int(reason_code)
|
||||
if code in (0, None):
|
||||
if self._on_disconnected:
|
||||
self._on_disconnected(None, False)
|
||||
return
|
||||
# 0x8E = 142 Session taken over
|
||||
if code == 142:
|
||||
if self._on_disconnected:
|
||||
self._on_disconnected("taken_over", True)
|
||||
return
|
||||
stop, reason = _classify_connack(code)
|
||||
if self._on_disconnected:
|
||||
self._on_disconnected(reason, stop)
|
||||
|
||||
def _on_message(self, client, userdata, msg) -> None:
|
||||
if self._on_down:
|
||||
self._on_down(bytes(msg.payload))
|
||||
|
||||
|
||||
def _reason_to_int(reason_code) -> Optional[int]:
|
||||
if reason_code is None:
|
||||
return None
|
||||
if isinstance(reason_code, int):
|
||||
return reason_code
|
||||
if isinstance(reason_code, ReasonCode):
|
||||
return int(reason_code)
|
||||
if isinstance(reason_code, MQTTErrorCode):
|
||||
return int(reason_code)
|
||||
try:
|
||||
return int(reason_code)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _classify_connack(code: int) -> tuple[bool, str]:
|
||||
# MQTT5: 0x86=134 Bad User Name or Password, 0x87=135 Not authorized, 0x8A=138 Banned
|
||||
# 0x88=136 Server unavailable, 0x89=137 Server busy -> continue
|
||||
if code in (4, 5, 134, 135, 138): # 3.1.1 4/5 and MQTT5 auth failures
|
||||
if code in (134, 4):
|
||||
return True, "bad_credentials" # 可能是令牌,Client 层再区分
|
||||
return True, "bad_credentials"
|
||||
if code in (136, 137, 0x88, 0x89):
|
||||
return False, "busy"
|
||||
return False, "network"
|
||||
@@ -0,0 +1,182 @@
|
||||
"""公共类型与常量。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from typing import Any, Callable, Optional
|
||||
|
||||
CLIENT_NAME = "python-sdk/0.1"
|
||||
|
||||
DEFAULT_MAX_BODY = 262144
|
||||
DEFAULT_MAX_META = 4096
|
||||
DEFAULT_MAX_FRAME = 786432
|
||||
MIN_MAX_RECEIVE = 1024
|
||||
SEND_QUEUE_LIMIT = 1000
|
||||
INFLIGHT_LIMIT = 100
|
||||
DEDUP_CAPACITY = 10000
|
||||
CONNECT_TIMEOUT_S = 30.0
|
||||
BACKOFF_INITIAL_S = 1.0
|
||||
BACKOFF_MAX_S = 30.0
|
||||
BACKOFF_JITTER = 0.3
|
||||
STABLE_RESET_S = 60.0
|
||||
|
||||
SESSION_TOKEN_PREFIX = "nst_"
|
||||
|
||||
|
||||
class ConnectionState(str, Enum):
|
||||
CONNECTING = "connecting"
|
||||
ONLINE = "online"
|
||||
RECONNECTING = "reconnecting"
|
||||
OFFLINE = "offline"
|
||||
KICKED = "kicked"
|
||||
AUTH_FAILED = "auth_failed"
|
||||
|
||||
|
||||
@dataclass
|
||||
class Target:
|
||||
kind: str
|
||||
id: str
|
||||
|
||||
def to_dict(self) -> dict[str, str]:
|
||||
return {"kind": self.kind, "id": self.id}
|
||||
|
||||
|
||||
@dataclass
|
||||
class Body:
|
||||
data: str
|
||||
enc: str = "utf8"
|
||||
content_type: Optional[str] = None
|
||||
|
||||
def decoded_size(self) -> int:
|
||||
if self.enc == "base64":
|
||||
import base64
|
||||
|
||||
return len(base64.b64decode(self.data, validate=False))
|
||||
return len(self.data.encode("utf-8"))
|
||||
|
||||
def effective_content_type(self) -> str:
|
||||
if self.content_type:
|
||||
return self.content_type
|
||||
if self.enc == "base64":
|
||||
return "application/octet-stream"
|
||||
return "text/plain; charset=utf-8"
|
||||
|
||||
def to_dict(self) -> dict[str, str]:
|
||||
d = {"enc": self.enc, "data": self.data}
|
||||
ct = self.effective_content_type()
|
||||
d["content_type"] = ct
|
||||
return d
|
||||
|
||||
|
||||
@dataclass
|
||||
class SendOptions:
|
||||
send_at: Optional[float] = None # 本机语义时间(秒)或与 send_at_ms 二选一
|
||||
send_at_ms: Optional[int] = None
|
||||
delay_ms: Optional[int] = None
|
||||
keep: bool = False
|
||||
ttl_seconds: Optional[int] = None
|
||||
receipt: bool = True
|
||||
talk_password: str = ""
|
||||
content_type: Optional[str] = None
|
||||
meta: Optional[dict[str, Any]] = None
|
||||
message_id: Optional[str] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class SendResult:
|
||||
id: str
|
||||
send_at_ms: int
|
||||
state: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class RecallResult:
|
||||
result: str
|
||||
recalled: int = 0
|
||||
accepted: int = 0
|
||||
other: int = 0
|
||||
|
||||
|
||||
@dataclass
|
||||
class IncomingMessage:
|
||||
id: str
|
||||
from_id: str
|
||||
to: Target
|
||||
body: Body
|
||||
send_at_ms: int
|
||||
meta: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Receipt:
|
||||
receipt_id: str
|
||||
id: str
|
||||
endpoint_id: str
|
||||
state: str
|
||||
reason: str
|
||||
at_ms: int
|
||||
|
||||
|
||||
@dataclass
|
||||
class RevokedEvent:
|
||||
id: str
|
||||
from_id: str
|
||||
reason: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class PresenceEvent:
|
||||
id: str
|
||||
online: bool
|
||||
at_ms: int
|
||||
|
||||
|
||||
@dataclass
|
||||
class GroupEvent:
|
||||
group_id: str
|
||||
event: str
|
||||
endpoint_id: str
|
||||
at_ms: int
|
||||
|
||||
|
||||
@dataclass
|
||||
class HelloLimits:
|
||||
server_time_ms: int = 0
|
||||
server_version: str = ""
|
||||
max_body_bytes: int = DEFAULT_MAX_BODY
|
||||
max_meta_bytes: int = DEFAULT_MAX_META
|
||||
max_frame_bytes: int = DEFAULT_MAX_FRAME
|
||||
max_ttl_seconds: int = 2592000
|
||||
max_schedule_seconds: int = 31536000
|
||||
ack_timeout_seconds: int = 300
|
||||
session_token: str = ""
|
||||
|
||||
|
||||
@dataclass
|
||||
class RegisterOptions:
|
||||
id: str = ""
|
||||
login_password: str = ""
|
||||
name: str = ""
|
||||
talk_password: str = ""
|
||||
|
||||
|
||||
@dataclass
|
||||
class RegisterResult:
|
||||
id: str
|
||||
login_password: Optional[str] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class ConnectionEvent:
|
||||
state: ConnectionState
|
||||
reason: str = ""
|
||||
|
||||
|
||||
SessionHandler = Callable[[str], None]
|
||||
MessageHandler = Callable[[IncomingMessage], None]
|
||||
ReceiptHandler = Callable[[Receipt], None]
|
||||
RevokedHandler = Callable[[RevokedEvent], None]
|
||||
PresenceHandler = Callable[[PresenceEvent], None]
|
||||
GroupEventHandler = Callable[[GroupEvent], None]
|
||||
ConnectionHandler = Callable[[ConnectionEvent], None]
|
||||
@@ -0,0 +1,18 @@
|
||||
"""UUIDv7(36 字符形式),兼容 Python 3.10。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import time
|
||||
import uuid
|
||||
|
||||
|
||||
def new_uuid7() -> str:
|
||||
if hasattr(uuid, "uuid7"):
|
||||
return str(uuid.uuid7()) # type: ignore[attr-defined]
|
||||
# RFC 9562 简化实现
|
||||
ts_ms = int(time.time() * 1000) & ((1 << 48) - 1)
|
||||
rand_a = int.from_bytes(os.urandom(2), "big") & 0x0FFF
|
||||
rand_b = int.from_bytes(os.urandom(8), "big") & ((1 << 62) - 1)
|
||||
value = (ts_ms << 80) | (0x7 << 76) | (rand_a << 64) | (0b10 << 62) | rand_b
|
||||
return str(uuid.UUID(int=value))
|
||||
@@ -0,0 +1,202 @@
|
||||
"""假传输单元测试:Clean Start、去重再 ack、本地超限、令牌回调、重交不改消息号。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import threading
|
||||
import time
|
||||
import unittest
|
||||
|
||||
from nixmsg import Body, Client, ConnectionState, FakeTransport, NixMsgError, SendOptions, Target
|
||||
from nixmsg.protocol import dumps
|
||||
|
||||
|
||||
class FakeTransportTests(unittest.TestCase):
|
||||
def _connect(self, transport: FakeTransport | None = None, **kwargs) -> tuple[Client, FakeTransport]:
|
||||
tr = transport or FakeTransport()
|
||||
c = Client(transport=tr, **kwargs)
|
||||
tokens: list[str] = []
|
||||
c.on_session(lambda t: tokens.append(t))
|
||||
c.connect("ws://example.test/mqtt", "ep1", password="secret", wait=True)
|
||||
self.assertEqual(c.state, ConnectionState.ONLINE)
|
||||
return c, tr
|
||||
|
||||
def test_clean_start_and_session_expiry_every_connect(self) -> None:
|
||||
tr = FakeTransport()
|
||||
c, _ = self._connect(tr)
|
||||
self.assertGreaterEqual(len(tr.connects), 1)
|
||||
p = tr.connects[0].params
|
||||
self.assertTrue(p.clean_start)
|
||||
self.assertEqual(p.session_expiry, 0)
|
||||
self.assertIn("nix/c/ep1/down", tr.subscriptions)
|
||||
# 断开再连
|
||||
tr.simulate_network_drop()
|
||||
time.sleep(0.3)
|
||||
# 等待重连至少再记一次
|
||||
deadline = time.time() + 3
|
||||
while len(tr.connects) < 2 and time.time() < deadline:
|
||||
time.sleep(0.05)
|
||||
self.assertGreaterEqual(len(tr.connects), 2)
|
||||
for rec in tr.connects:
|
||||
self.assertTrue(rec.params.clean_start)
|
||||
self.assertEqual(rec.params.session_expiry, 0)
|
||||
c.close()
|
||||
|
||||
def test_session_token_callback(self) -> None:
|
||||
tr = FakeTransport()
|
||||
tokens: list[str] = []
|
||||
c = Client(transport=tr)
|
||||
c.on_session(lambda t: tokens.append(t))
|
||||
c.connect("ws://example.test/mqtt", "ep1", password="pw")
|
||||
self.assertEqual(tokens, ["nst_test_token"])
|
||||
self.assertEqual(c.session_token, "nst_test_token")
|
||||
c.close()
|
||||
|
||||
def test_dedup_acked_rearrival_acks_again(self) -> None:
|
||||
c, tr = self._connect()
|
||||
delivered: list[str] = []
|
||||
c.on_message(lambda m: delivered.append(m.id))
|
||||
|
||||
msg = {
|
||||
"v": 1,
|
||||
"type": "msg",
|
||||
"id": "m1",
|
||||
"from": "peer",
|
||||
"to": {"kind": "endpoint", "id": "ep1"},
|
||||
"body": {"enc": "utf8", "data": "hi"},
|
||||
"send_at_ms": 1,
|
||||
}
|
||||
tr.inject_down(dumps(msg))
|
||||
time.sleep(0.1)
|
||||
self.assertEqual(delivered, ["m1"])
|
||||
# 找 ack 帧
|
||||
acks = [json.loads(p.decode()) for _, p in tr.publishes if json.loads(p.decode()).get("type") == "ack"]
|
||||
self.assertGreaterEqual(len(acks), 1)
|
||||
|
||||
before = len(tr.publishes)
|
||||
tr.inject_down(dumps(msg)) # 已确认再到达
|
||||
time.sleep(0.1)
|
||||
self.assertEqual(delivered, ["m1"]) # 不重复交应用
|
||||
acks2 = [json.loads(p.decode()) for _, p in tr.publishes[before:] if json.loads(p.decode()).get("type") == "ack"]
|
||||
self.assertGreaterEqual(len(acks2), 1) # 再 ack
|
||||
c.close()
|
||||
|
||||
def test_local_body_too_large(self) -> None:
|
||||
c, tr = self._connect()
|
||||
big = "x" * (c.limits.max_body_bytes + 1)
|
||||
with self.assertRaises(NixMsgError) as cm:
|
||||
c.send(Target(kind="endpoint", id="ep2"), Body(data=big))
|
||||
self.assertEqual(cm.exception.code, "body_too_large")
|
||||
c.close()
|
||||
|
||||
def test_local_frame_too_large(self) -> None:
|
||||
tr = FakeTransport()
|
||||
tr.auto_hello = {
|
||||
"server_time_ms": 1_750_000_000_000,
|
||||
"server_version": "0.1.0",
|
||||
"max_body_bytes": 262144,
|
||||
"max_meta_bytes": 4096,
|
||||
"max_frame_bytes": 200, # 故意很小
|
||||
"max_ttl_seconds": 2592000,
|
||||
"max_schedule_seconds": 31536000,
|
||||
"ack_timeout_seconds": 300,
|
||||
"session_token": "nst_x",
|
||||
}
|
||||
c = Client(transport=tr)
|
||||
c.connect("ws://example.test/mqtt", "ep1", password="pw")
|
||||
with self.assertRaises(NixMsgError) as cm:
|
||||
c.send(Target(kind="endpoint", id="ep2"), Body(data="hello world " * 20))
|
||||
self.assertEqual(cm.exception.code, "frame_too_large")
|
||||
c.close()
|
||||
|
||||
def test_resend_keeps_send_at_ms(self) -> None:
|
||||
tr = FakeTransport()
|
||||
rate_hits = {"n": 0}
|
||||
|
||||
def hook(frame: dict):
|
||||
if frame.get("type") != "send":
|
||||
return None
|
||||
if rate_hits["n"] == 0:
|
||||
rate_hits["n"] += 1
|
||||
return {
|
||||
"v": 1,
|
||||
"type": "resp",
|
||||
"rid": frame["rid"],
|
||||
"ok": False,
|
||||
"error": {"code": "rate_limited", "message": "slow"},
|
||||
}
|
||||
return {
|
||||
"v": 1,
|
||||
"type": "resp",
|
||||
"rid": frame["rid"],
|
||||
"ok": True,
|
||||
"data": {"id": frame["id"], "send_at_ms": frame.get("send_at_ms", 0), "state": "scheduled"},
|
||||
}
|
||||
|
||||
tr.on_up(hook)
|
||||
tr.auto_send_ok = False
|
||||
c = Client(transport=tr)
|
||||
c.connect("ws://example.test/mqtt", "ep1", password="pw")
|
||||
fixed = 1_700_000_000_000
|
||||
result = c.send(
|
||||
Target(kind="endpoint", id="ep2"),
|
||||
Body(data="hi"),
|
||||
SendOptions(send_at_ms=fixed, message_id="fixed-id-1"),
|
||||
)
|
||||
self.assertEqual(result.id, "fixed-id-1")
|
||||
sends = [json.loads(p.decode()) for _, p in tr.publishes if json.loads(p.decode()).get("type") == "send"]
|
||||
self.assertGreaterEqual(len(sends), 2)
|
||||
for s in sends:
|
||||
self.assertEqual(s["id"], "fixed-id-1")
|
||||
self.assertEqual(s["send_at_ms"], fixed)
|
||||
c.close()
|
||||
|
||||
def test_session_invalid_stops_reconnect(self) -> None:
|
||||
tr = FakeTransport()
|
||||
# 先用密码连上
|
||||
c = Client(transport=tr)
|
||||
c.connect("ws://example.test/mqtt", "ep1", password="pw")
|
||||
# 下次连接令牌失败
|
||||
tr.next_connack_fail = "bad_credentials"
|
||||
tr.simulate_network_drop()
|
||||
deadline = time.time() + 3
|
||||
while c.state != ConnectionState.AUTH_FAILED and time.time() < deadline:
|
||||
time.sleep(0.05)
|
||||
self.assertEqual(c.state, ConnectionState.AUTH_FAILED)
|
||||
n = len(tr.connects)
|
||||
time.sleep(0.5)
|
||||
self.assertEqual(len(tr.connects), n) # 不再重连
|
||||
c.close()
|
||||
|
||||
def test_delivered_duplicate_ignored(self) -> None:
|
||||
c, tr = self._connect(auto_ack=False)
|
||||
barrier = threading.Event()
|
||||
seen = []
|
||||
|
||||
def handler(m):
|
||||
seen.append(m.id)
|
||||
barrier.wait(timeout=2)
|
||||
|
||||
c.on_message(handler)
|
||||
msg = {
|
||||
"v": 1,
|
||||
"type": "msg",
|
||||
"id": "m2",
|
||||
"from": "peer",
|
||||
"to": {"kind": "endpoint", "id": "ep1"},
|
||||
"body": {"enc": "utf8", "data": "x"},
|
||||
"send_at_ms": 1,
|
||||
}
|
||||
t = threading.Thread(target=lambda: tr.inject_down(dumps(msg)))
|
||||
t.start()
|
||||
time.sleep(0.05)
|
||||
tr.inject_down(dumps(msg)) # 已交未确认
|
||||
barrier.set()
|
||||
t.join(timeout=2)
|
||||
time.sleep(0.05)
|
||||
self.assertEqual(seen, ["m2"])
|
||||
c.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user