refactor: 重构底层代码,抽离conn,调整消息处理器并创建传输层接口。

feature: 支持mqtt非桥接版本。
This commit is contained in:
caixypromise
2025-09-14 03:00:50 +08:00
parent d04ec9d510
commit 1ba556988f
44 changed files with 6455 additions and 66 deletions
@@ -0,0 +1,279 @@
import asyncio
import json
import time
import uuid
from typing import Dict, Any, Optional, Callable
from config.logger import setup_logging
logger = setup_logging()
class MQTTConnection:
"""
MQTT连接处理类:管理单个MQTT客户端连接
处理MQTT协议消息和会话管理
"""
def __init__(self, socket, connection_id: int, mqtt_server):
self.socket = socket
self.connection_id = connection_id
self.mqtt_server = mqtt_server
# 连接信息
self.client_id = None
self.device_id = None
self.username = None
self.password = None
self.session_id = None
# 协议状态
self.is_connected_flag = False
self.keep_alive_interval = 0
self.last_activity = time.time()
# 消息处理
self.message_callback = None
self.reply_topic = None
# UDP相关
self.udp_config = None
# 任务管理
self.keep_alive_task = None
self._closed = False
# 创建MQTT协议处理器
from core.protocols.mqtt_protocol import MQTTProtocol
self.protocol = MQTTProtocol(socket)
self._setup_protocol_handlers()
def _setup_protocol_handlers(self):
"""设置协议事件处理"""
self.protocol.on('connect', self._handle_connect)
self.protocol.on('publish', self._handle_publish)
self.protocol.on('subscribe', self._handle_subscribe)
self.protocol.on('disconnect', self._handle_disconnect)
self.protocol.on('close', self._handle_close)
self.protocol.on('error', self._handle_error)
async def _handle_connect(self, connect_data: Dict[str, Any]):
"""处理CONNECT消息"""
try:
self.client_id = connect_data['clientId']
self.username = connect_data.get('username')
self.password = connect_data.get('password')
self.keep_alive_interval = connect_data.get('keepAlive', 0) * 1000 # 转换为毫秒
logger.info(f"MQTT客户端连接: {self.client_id}")
# 解析客户端ID获取设备信息
if not self._parse_client_id():
await self.protocol.send_connack(1) # 连接被拒绝
await self.close()
return
# 生成会话ID
self.session_id = str(uuid.uuid4())
# 设置回复主题
self.reply_topic = f"devices/p2p/{self.device_id.replace(':', '_')}"
# 发送连接确认
await self.protocol.send_connack(0) # 连接接受
self.is_connected_flag = True
# 启动心跳检查
if self.keep_alive_interval > 0:
self.keep_alive_task = asyncio.create_task(self._keep_alive_check())
# 通知服务器新连接
await self.mqtt_server.on_client_connected(self)
except Exception as e:
logger.error(f"处理CONNECT消息失败: {e}")
await self.close()
def _parse_client_id(self) -> bool:
"""解析客户端ID获取设备信息"""
try:
# 支持格式: GID_test@@@mac_address@@@uuid 或 GID_test@@@mac_address
parts = self.client_id.split('@@@')
if len(parts) >= 2:
self.group_id = parts[0]
# 将下划线替换为冒号格式的MAC地址
self.device_id = parts[1].replace('_', ':')
if len(parts) >= 3:
self.uuid = parts[2]
return True
else:
logger.error(f"无效的客户端ID格式: {self.client_id}")
return False
except Exception as e:
logger.error(f"解析客户端ID失败: {e}")
return False
async def _handle_publish(self, publish_data: Dict[str, Any]):
"""处理PUBLISH消息"""
try:
topic = publish_data['topic']
payload = publish_data['payload']
logger.debug(f"收到MQTT发布消息: topic={topic}, payload={payload}")
# 更新活动时间
self.last_activity = time.time()
# 解析JSON消息
try:
message_data = json.loads(payload)
# 处理不同类型的消息
if message_data.get('type') == 'hello':
await self._handle_hello_message(message_data)
else:
# 其他消息通过回调处理
if self.message_callback:
self.message_callback(topic, payload)
except json.JSONDecodeError:
logger.error(f"MQTT消息JSON解析失败: {payload}")
except Exception as e:
logger.error(f"处理PUBLISH消息失败: {e}")
async def _handle_hello_message(self, message_data: Dict[str, Any]):
"""处理hello消息,初始化UDP配置"""
try:
# 生成UDP加密配置
import os
self.udp_config = {
'key': os.urandom(16),
'encryption': 'aes-128-ctr',
'server': self.mqtt_server.public_ip,
'port': self.mqtt_server.udp_port
}
# 构造hello回复
hello_reply = {
'type': 'hello',
'version': message_data.get('version', 3),
'session_id': self.session_id,
'transport': 'udp',
'udp': {
'server': self.udp_config['server'],
'port': self.udp_config['port'],
'encryption': self.udp_config['encryption'],
'key': self.udp_config['key'].hex(),
'nonce': '00' * 16 # 临时nonce
},
'audio_params': message_data.get('audio_params', {})
}
# 发送回复
await self.send_message(self.reply_topic, json.dumps(hello_reply))
logger.info(f"MQTT Hello消息处理完成: {self.client_id}")
except Exception as e:
logger.error(f"处理hello消息失败: {e}")
async def _handle_subscribe(self, subscribe_data: Dict[str, Any]):
"""处理SUBSCRIBE消息"""
try:
topic = subscribe_data['topic']
packet_id = subscribe_data['packetId']
logger.debug(f"客户端订阅主题: {topic}")
# 发送订阅确认
await self.protocol.send_suback(packet_id, 0)
except Exception as e:
logger.error(f"处理SUBSCRIBE消息失败: {e}")
async def _handle_disconnect(self):
"""处理DISCONNECT消息"""
logger.info(f"客户端主动断开连接: {self.client_id}")
await self.close()
async def _handle_close(self):
"""处理连接关闭"""
logger.info(f"MQTT连接关闭: {self.client_id}")
await self.close()
async def _handle_error(self, error):
"""处理连接错误"""
logger.error(f"MQTT连接错误: {self.client_id}, error: {error}")
await self.close()
async def _keep_alive_check(self):
"""心跳检查任务"""
try:
while self.is_connected_flag and not self._closed:
await asyncio.sleep(self.keep_alive_interval / 1000 / 2) # 检查间隔为心跳间隔的一半
current_time = time.time()
if current_time - self.last_activity > self.keep_alive_interval / 1000 * 1.5:
logger.info(f"MQTT客户端心跳超时: {self.client_id}")
await self.close()
break
except asyncio.CancelledError:
pass
except Exception as e:
logger.error(f"心跳检查任务出错: {e}")
def set_message_callback(self, callback: Callable[[str, str], None]):
"""设置消息接收回调"""
self.message_callback = callback
async def send_message(self, topic: str, payload: str):
"""发送MQTT消息"""
if self._closed or not self.is_connected_flag:
return
try:
await self.protocol.send_publish(topic, payload, qos=0)
logger.debug(f"发送MQTT消息: topic={topic}, payload={payload}")
except Exception as e:
logger.error(f"发送MQTT消息失败: {e}")
def is_connected(self) -> bool:
"""检查连接状态"""
return self.is_connected_flag and not self._closed
async def close(self):
"""关闭连接"""
if self._closed:
return
self._closed = True
self.is_connected_flag = False
# 取消心跳检查任务
if self.keep_alive_task and not self.keep_alive_task.done():
self.keep_alive_task.cancel()
try:
await self.keep_alive_task
except asyncio.CancelledError:
pass
# 通知服务器连接关闭
try:
await self.mqtt_server.on_client_disconnected(self)
except Exception as e:
logger.error(f"通知服务器连接关闭失败: {e}")
# 关闭协议处理器
try:
await self.protocol.close()
except Exception as e:
logger.error(f"关闭MQTT协议处理器失败: {e}")
logger.info(f"MQTT连接已关闭: {self.client_id}")
@@ -0,0 +1,432 @@
import asyncio
from typing import Dict, Any, Callable
from config.logger import setup_logging
logger = setup_logging()
# MQTT 固定头部的类型
class PacketType:
CONNECT = 1
CONNACK = 2
PUBLISH = 3
SUBSCRIBE = 8
SUBACK = 9
PINGREQ = 12
PINGRESP = 13
DISCONNECT = 14
class MQTTProtocol:
"""
MQTT协议处理器:负责MQTT协议的解析和封装
"""
def __init__(self, socket):
self.socket = socket
self.buffer = b''
self.event_handlers = {}
self.is_connected = False
self.keep_alive_interval = 0
self.last_activity = 0
# 启动消息处理任务
self._processing_task = asyncio.create_task(self._process_messages())
def on(self, event: str, handler: Callable):
"""注册事件处理器"""
self.event_handlers[event] = handler
def emit(self, event: str, *args, **kwargs):
"""触发事件"""
handler = self.event_handlers.get(event)
if handler:
if asyncio.iscoroutinefunction(handler):
asyncio.create_task(handler(*args, **kwargs))
else:
handler(*args, **kwargs)
async def _process_messages(self):
"""处理消息的主循环"""
try:
while True:
# 从socket读取数据
data = await self._read_socket()
if not data:
break
# 添加到缓冲区
self.buffer += data
# 处理缓冲区中的消息
await self._process_buffer()
except asyncio.CancelledError:
pass
except Exception as e:
logger.error(f"MQTT消息处理循环出错: {e}")
self.emit('error', e)
finally:
self.emit('close')
async def _read_socket(self) -> bytes:
"""从socket读取数据"""
try:
# 使用asyncio的socket读取
loop = asyncio.get_event_loop()
data = await loop.sock_recv(self.socket, 4096)
return data
except Exception as e:
logger.error(f"读取socket数据失败: {e}")
return b''
async def _process_buffer(self):
"""处理缓冲区中的消息"""
while len(self.buffer) >= 2: # 至少需要2字节开始解析
try:
# 解析消息
message_length, message = self._parse_message()
if message_length == 0:
break # 消息不完整,等待更多数据
# 从缓冲区移除已处理的消息
self.buffer = self.buffer[message_length:]
# 处理消息
await self._handle_message(message)
except Exception as e:
logger.error(f"处理MQTT消息失败: {e}")
self.emit('protocolError', e)
break
def _parse_message(self) -> tuple[int, Dict[str, Any]]:
"""解析MQTT消息"""
if len(self.buffer) < 2:
return 0, {}
# 获取消息类型
first_byte = self.buffer[0]
packet_type = (first_byte >> 4)
# 解析剩余长度
remaining_length, bytes_read = self._decode_remaining_length()
if remaining_length == -1:
return 0, {} # 长度解析失败,等待更多数据
# 计算完整消息长度
total_length = 1 + bytes_read + remaining_length
if len(self.buffer) < total_length:
return 0, {} # 消息不完整
# 提取消息数据
message_data = self.buffer[:total_length]
# 根据消息类型解析
if packet_type == PacketType.CONNECT:
message = self._parse_connect(message_data)
elif packet_type == PacketType.PUBLISH:
message = self._parse_publish(message_data)
elif packet_type == PacketType.SUBSCRIBE:
message = self._parse_subscribe(message_data)
elif packet_type == PacketType.PINGREQ:
message = {'type': 'pingreq'}
elif packet_type == PacketType.DISCONNECT:
message = {'type': 'disconnect'}
else:
logger.warning(f"未处理的MQTT消息类型: {packet_type}")
message = {'type': 'unknown', 'packet_type': packet_type}
return total_length, message
def _decode_remaining_length(self) -> tuple[int, int]:
"""解码剩余长度字段"""
multiplier = 1
value = 0
bytes_read = 0
while bytes_read < 4 and bytes_read + 1 < len(self.buffer):
digit = self.buffer[bytes_read + 1]
bytes_read += 1
value += (digit & 127) * multiplier
multiplier *= 128
if (digit & 128) == 0:
break
else:
if bytes_read >= 4:
return -1, 0 # 长度字段过长
return -1, 0 # 数据不完整
return value, bytes_read
def _encode_remaining_length(self, length: int) -> bytes:
"""编码剩余长度字段"""
result = bytearray()
while True:
digit = length % 128
length = length // 128
if length > 0:
digit |= 0x80
result.append(digit)
if length == 0:
break
return bytes(result)
def _parse_connect(self, message_data: bytes) -> Dict[str, Any]:
"""解析CONNECT消息"""
try:
# 跳过固定头部和剩余长度
_, bytes_read = self._decode_remaining_length()
pos = 1 + bytes_read
# 协议名长度
protocol_length = int.from_bytes(message_data[pos:pos+2], 'big')
pos += 2
# 协议名
protocol = message_data[pos:pos+protocol_length].decode('utf-8')
pos += protocol_length
# 协议级别
protocol_level = message_data[pos]
pos += 1
# 连接标志
connect_flags = message_data[pos]
has_username = (connect_flags & 0x80) != 0
has_password = (connect_flags & 0x40) != 0
pos += 1
# 保持连接时间
keep_alive = int.from_bytes(message_data[pos:pos+2], 'big')
pos += 2
# 客户端ID
client_id_length = int.from_bytes(message_data[pos:pos+2], 'big')
pos += 2
client_id = message_data[pos:pos+client_id_length].decode('utf-8')
pos += client_id_length
# 用户名(如果存在)
username = ''
if has_username:
username_length = int.from_bytes(message_data[pos:pos+2], 'big')
pos += 2
username = message_data[pos:pos+username_length].decode('utf-8')
pos += username_length
# 密码(如果存在)
password = ''
if has_password:
password_length = int.from_bytes(message_data[pos:pos+2], 'big')
pos += 2
password = message_data[pos:pos+password_length].decode('utf-8')
pos += password_length
return {
'type': 'connect',
'protocol': protocol,
'protocolLevel': protocol_level,
'clientId': client_id,
'keepAlive': keep_alive,
'username': username,
'password': password
}
except Exception as e:
logger.error(f"解析CONNECT消息失败: {e}")
raise
def _parse_publish(self, message_data: bytes) -> Dict[str, Any]:
"""解析PUBLISH消息"""
try:
# 获取QoS等标志
first_byte = message_data[0]
qos = (first_byte & 0x06) >> 1
dup = (first_byte & 0x08) != 0
retain = (first_byte & 0x01) != 0
# 跳过固定头部和剩余长度
_, bytes_read = self._decode_remaining_length()
pos = 1 + bytes_read
# 主题长度
topic_length = int.from_bytes(message_data[pos:pos+2], 'big')
pos += 2
# 主题
topic = message_data[pos:pos+topic_length].decode('utf-8')
pos += topic_length
# 消息IDQoS > 0时存在)
packet_id = None
if qos > 0:
packet_id = int.from_bytes(message_data[pos:pos+2], 'big')
pos += 2
# 有效载荷
payload = message_data[pos:].decode('utf-8')
return {
'type': 'publish',
'topic': topic,
'payload': payload,
'qos': qos,
'dup': dup,
'retain': retain,
'packetId': packet_id
}
except Exception as e:
logger.error(f"解析PUBLISH消息失败: {e}")
raise
def _parse_subscribe(self, message_data: bytes) -> Dict[str, Any]:
"""解析SUBSCRIBE消息"""
try:
# 跳过固定头部和剩余长度
_, bytes_read = self._decode_remaining_length()
pos = 1 + bytes_read
# 消息ID
packet_id = int.from_bytes(message_data[pos:pos+2], 'big')
pos += 2
# 主题长度
topic_length = int.from_bytes(message_data[pos:pos+2], 'big')
pos += 2
# 主题
topic = message_data[pos:pos+topic_length].decode('utf-8')
pos += topic_length
# QoS
qos = message_data[pos]
return {
'type': 'subscribe',
'packetId': packet_id,
'topic': topic,
'qos': qos
}
except Exception as e:
logger.error(f"解析SUBSCRIBE消息失败: {e}")
raise
async def _handle_message(self, message: Dict[str, Any]):
"""处理解析后的消息"""
message_type = message.get('type')
if message_type == 'connect':
self.keep_alive_interval = message.get('keepAlive', 0)
self.is_connected = True
self.emit('connect', message)
elif message_type == 'publish':
self.emit('publish', message)
elif message_type == 'subscribe':
self.emit('subscribe', message)
elif message_type == 'pingreq':
await self.send_pingresp()
elif message_type == 'disconnect':
self.emit('disconnect')
else:
logger.warning(f"未处理的消息类型: {message_type}")
async def send_connack(self, return_code: int = 0, session_present: bool = False):
"""发送CONNACK消息"""
packet = bytearray([
PacketType.CONNACK << 4, # 固定头部
2, # 剩余长度
1 if session_present else 0, # 连接确认标志
return_code # 返回码
])
await self._send_packet(packet)
async def send_publish(self, topic: str, payload: str, qos: int = 0,
dup: bool = False, retain: bool = False, packet_id: int = None):
"""发送PUBLISH消息"""
# 构造固定头部
first_byte = PacketType.PUBLISH << 4
if dup:
first_byte |= 0x08
if qos > 0:
first_byte |= (qos << 1)
if retain:
first_byte |= 0x01
# 构造可变头部和载荷
topic_bytes = topic.encode('utf-8')
payload_bytes = payload.encode('utf-8')
variable_header = bytearray()
variable_header.extend(len(topic_bytes).to_bytes(2, 'big'))
variable_header.extend(topic_bytes)
if qos > 0 and packet_id is not None:
variable_header.extend(packet_id.to_bytes(2, 'big'))
# 计算剩余长度
remaining_length = len(variable_header) + len(payload_bytes)
remaining_length_bytes = self._encode_remaining_length(remaining_length)
# 构造完整消息
packet = bytearray([first_byte])
packet.extend(remaining_length_bytes)
packet.extend(variable_header)
packet.extend(payload_bytes)
await self._send_packet(packet)
async def send_suback(self, packet_id: int, return_code: int = 0):
"""发送SUBACK消息"""
packet = bytearray([
PacketType.SUBACK << 4, # 固定头部
3, # 剩余长度
packet_id >> 8, # 消息ID高字节
packet_id & 0xFF, # 消息ID低字节
return_code # 返回码
])
await self._send_packet(packet)
async def send_pingresp(self):
"""发送PINGRESP消息"""
packet = bytearray([
PacketType.PINGRESP << 4, # 固定头部
0 # 剩余长度
])
await self._send_packet(packet)
async def _send_packet(self, packet: bytearray):
"""发送数据包"""
try:
loop = asyncio.get_event_loop()
await loop.sock_sendall(self.socket, bytes(packet))
except Exception as e:
logger.error(f"发送MQTT数据包失败: {e}")
raise
async def close(self):
"""关闭协议处理器"""
if hasattr(self, '_processing_task') and not self._processing_task.done():
self._processing_task.cancel()
try:
await self._processing_task
except asyncio.CancelledError:
pass
try:
self.socket.close()
except Exception as e:
logger.error(f"关闭socket失败: {e}")