feat: 实现 Python 与 Java SDK 连接收发及其余接口
This commit is contained in:
@@ -0,0 +1,350 @@
|
||||
"""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
|
||||
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
|
||||
|
||||
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()
|
||||
client = PahoClient(
|
||||
callback_api_version=CallbackAPIVersion.VERSION2,
|
||||
client_id=params.client_id,
|
||||
protocol=PahoClient.MQTTv5,
|
||||
)
|
||||
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
|
||||
self._client = client
|
||||
|
||||
url = params.url
|
||||
u = urlparse(url if "://" in url else "ws://" + url)
|
||||
use_tcp = params.use_tcp or u.scheme in ("mqtt", "mqtts")
|
||||
host = u.hostname or "localhost"
|
||||
port = u.port or (8883 if u.scheme in ("wss", "mqtts") else 443 if u.scheme == "wss" else 80)
|
||||
if use_tcp:
|
||||
if not u.port:
|
||||
port = 8883 if u.scheme == "mqtts" else 1883
|
||||
tls = u.scheme == "mqtts"
|
||||
if tls:
|
||||
client.tls_set()
|
||||
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
|
||||
client.connect(
|
||||
host,
|
||||
port,
|
||||
keepalive=params.keep_alive,
|
||||
clean_start=params.clean_start,
|
||||
properties=props,
|
||||
)
|
||||
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"})
|
||||
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
|
||||
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 self._client:
|
||||
self._client.subscribe(topic, qos=1)
|
||||
|
||||
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:
|
||||
if self._on_connected:
|
||||
self._on_connected()
|
||||
return
|
||||
stop, reason = _classify_connack(code)
|
||||
if self._on_disconnected:
|
||||
self._on_disconnected(reason, stop)
|
||||
|
||||
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)
|
||||
if isinstance(reason_code, MQTTErrorCode):
|
||||
return int(reason_code)
|
||||
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"
|
||||
Reference in New Issue
Block a user