feat: 实现 Python 与 Java SDK 连接收发及其余接口

This commit is contained in:
Nixevol
2026-09-30 07:06:30 +08:00
parent 9650cbff76
commit d357082f3d
27 changed files with 4515 additions and 1 deletions
+350
View File
@@ -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"