"""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"