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