Files
NixMsg/sdk/python/src/nixmsg/transport.py
T

380 lines
13 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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"