380 lines
13 KiB
Python
380 lines
13 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
|
||
|
||
DownHandler = Callable[[bytes], None]
|
||
ConnHandler = Callable[[], None]
|
||
DiscHandler = Callable[[Optional[str], bool], None]
|
||
# reason_code_str, stop_reconnect
|
||
|
||
|
||
def _require_paho():
|
||
"""真实 MQTT 路径才加载 paho;假传输单测不依赖。"""
|
||
try:
|
||
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
|
||
except ImportError as e:
|
||
raise ImportError("需要 paho-mqtt>=2.0(真实 MQTT 连接)") from e
|
||
return CallbackAPIVersion, PahoClient, MQTT_ERR_SUCCESS, MQTTv5, MQTTErrorCode, ReasonCode
|
||
|
||
|
||
@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: Any = 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:
|
||
CallbackAPIVersion, PahoClient, _, MQTTv5, _, _ = _require_paho()
|
||
self.disconnect()
|
||
url = params.url
|
||
u = urlparse(url if "://" in url else "ws://" + url)
|
||
use_tcp = params.use_tcp
|
||
scheme = u.scheme.lower()
|
||
if scheme == "https":
|
||
scheme = "wss"
|
||
elif scheme == "http":
|
||
scheme = "ws"
|
||
if scheme in ("mqtt", "mqtts"):
|
||
use_tcp = True
|
||
# WebSocket 必须显式 transport=websockets;裸 TCP 走默认。
|
||
client = PahoClient(
|
||
callback_api_version=CallbackAPIVersion.VERSION2,
|
||
client_id=params.client_id,
|
||
protocol=MQTTv5,
|
||
transport="tcp" if use_tcp else "websockets",
|
||
reconnect_on_failure=False,
|
||
)
|
||
client.username_pw_set(params.username, params.password)
|
||
client.on_connect = self._on_connect
|
||
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 scheme == "mqtts" else 1883
|
||
if scheme == "mqtts":
|
||
client.tls_set()
|
||
else:
|
||
path = u.path or "/mqtt"
|
||
if not u.port:
|
||
port = 443 if scheme == "wss" else 80
|
||
if scheme == "wss":
|
||
client.tls_set()
|
||
client.ws_set_options(path=path, headers={"Sec-WebSocket-Protocol": "mqtt"})
|
||
client.connect(
|
||
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:
|
||
_, _, MQTT_ERR_SUCCESS, _, _, _ = _require_paho()
|
||
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:
|
||
_, _, MQTT_ERR_SUCCESS, _, _, _ = _require_paho()
|
||
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;0x8B = 139 可重试
|
||
if code == 142:
|
||
if self._on_disconnected:
|
||
self._on_disconnected("taken_over", True)
|
||
return
|
||
if code == 139:
|
||
if self._on_disconnected:
|
||
self._on_disconnected("network", False)
|
||
return
|
||
stop, reason = _classify_connack(code)
|
||
if self._on_disconnected:
|
||
self._on_disconnected(reason, stop)
|
||
|
||
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
|
||
try:
|
||
_, _, _, _, MQTTErrorCode, ReasonCode = _require_paho()
|
||
except ImportError:
|
||
MQTTErrorCode = () # type: ignore[assignment,misc]
|
||
ReasonCode = () # type: ignore[assignment,misc]
|
||
if ReasonCode and isinstance(reason_code, ReasonCode):
|
||
return int(reason_code.value)
|
||
if MQTTErrorCode and 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"
|