对真实 nixmsg 跑 DEVELOPMENT 第 9 节接入清单(跳过仅 JS 跨域),补 README/示例,并修 Paho/HiveMQ 真机联调死锁与鉴权分类。
355 lines
12 KiB
Python
355 lines
12 KiB
Python
"""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, MQTTv5
|
|
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
|
|
self._sub_event = threading.Event()
|
|
self._sub_mid: Optional[int] = None
|
|
|
|
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()
|
|
url = params.url
|
|
u = urlparse(url if "://" in url else "ws://" + url)
|
|
use_tcp = params.use_tcp or u.scheme in ("mqtt", "mqtts")
|
|
# WebSocket 必须显式 transport=websockets;裸 TCP 走默认。
|
|
client = PahoClient(
|
|
callback_api_version=CallbackAPIVersion.VERSION2,
|
|
client_id=params.client_id,
|
|
protocol=MQTTv5,
|
|
transport="tcp" if use_tcp else "websockets",
|
|
)
|
|
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
|
|
client.on_subscribe = self._on_subscribe
|
|
self._client = client
|
|
|
|
host = u.hostname or "localhost"
|
|
port = u.port or (8883 if u.scheme in ("wss", "mqtts") else 443 if u.scheme == "wss" else 80)
|
|
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
|
|
if use_tcp:
|
|
if not u.port:
|
|
port = 8883 if u.scheme == "mqtts" else 1883
|
|
if u.scheme == "mqtts":
|
|
client.tls_set()
|
|
else:
|
|
path = u.path or "/mqtt"
|
|
if not path.endswith("/mqtt"):
|
|
path = path.rstrip("/") + "/mqtt"
|
|
if not u.port:
|
|
port = 443 if u.scheme == "wss" else 80
|
|
if u.scheme == "wss":
|
|
client.tls_set()
|
|
client.ws_set_options(path=path, headers={"Sec-WebSocket-Protocol": "mqtt"})
|
|
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 not self._client:
|
|
return
|
|
self._sub_event.clear()
|
|
result, mid = self._client.subscribe(topic, qos=1)
|
|
if result != MQTT_ERR_SUCCESS:
|
|
raise RuntimeError(f"subscribe failed: {result}")
|
|
self._sub_mid = mid
|
|
if not self._sub_event.wait(10):
|
|
raise RuntimeError("subscribe timeout")
|
|
|
|
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:
|
|
# 不在 loop 线程里同步做 subscribe+等待,否则会卡死 SUBACK
|
|
if self._on_connected:
|
|
threading.Thread(target=self._on_connected, name="nixmsg-on-connected", daemon=True).start()
|
|
return
|
|
stop, reason = _classify_connack(code if code is not None else -1)
|
|
if self._on_disconnected:
|
|
self._on_disconnected(reason, stop)
|
|
|
|
def _on_subscribe(self, client, userdata, mid, reason_codes, properties) -> None:
|
|
if self._sub_mid is None or mid == self._sub_mid:
|
|
self._sub_event.set()
|
|
|
|
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.value)
|
|
if isinstance(reason_code, MQTTErrorCode):
|
|
return int(reason_code)
|
|
# paho 偶发其它包装
|
|
val = getattr(reason_code, "value", None)
|
|
if isinstance(val, int):
|
|
return val
|
|
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"
|