Files
xiaozhi-esp32-server/main/xiaozhi-server/core/servers/mqtt_server.py
T

1168 lines
45 KiB
Python
Raw Normal View History

import asyncio
import ipaddress
import json
import secrets
import socket
import time
import weakref
from typing import Dict, Any, Set
from config.logger import setup_logging
from core.protocols.mqtt_connection import MQTTConnection
from core.transport.mqtt_transport import MQTTTransport, UDPAudioHandler
from core.services.connection_service import ConnectionService
from core.services.native_mqtt_connection_registry import (
NativeMqttConnectionRegistry,
)
from core.services.native_mqtt_call_manager import NativeMqttCallManager
from core.utils.mqtt_auth import normalize_signature_key, parse_mqtt_endpoint
from cryptography.hazmat.primitives.ciphers import Cipher, algorithms, modes
from cryptography.hazmat.backends import default_backend
logger = setup_logging()
class MQTTServer:
"""
原生MQTT服务器:直接处理MQTT协议连接
集成到xiaozhi-server架构中
"""
def __init__(self, config: Dict[str, Any]):
self.config = config
self.logger = setup_logging()
# 服务器配置
server_config = config.get('mqtt_server', {})
self.mqtt_port = server_config.get('port', 1883)
self.udp_port = server_config.get('udp_port', self.mqtt_port)
self.host = server_config.get('host', '0.0.0.0')
self.public_endpoint = self._resolve_public_host(server_config)
self.udp_bind_host = self.host
self._udp_bind_host_config = server_config.get('udp_bind_host')
self.signature_key = self._resolve_signature_key(config)
if not self.signature_key:
raise ValueError(
"启用原生MQTT必须配置mqtt_server.signature_key或"
"server.mqtt_signature_key"
)
self.message_queue_size = int(server_config.get('message_queue_size', 128))
self.business_ready_timeout = float(
server_config.get('business_ready_timeout', 30) or 30
)
self.close_timeout = max(
0.1, float(server_config.get('close_timeout', 2) or 2)
)
self.shutdown_timeout = max(
self.close_timeout,
float(server_config.get('shutdown_timeout', 10) or 10),
)
self.goodbye_timeout = max(
0.1, float(server_config.get('goodbye_timeout', 1) or 1)
)
self.max_connections = int(server_config.get('max_connections', 1000))
self.max_pending_connections = int(
server_config.get('max_pending_connections', 128)
)
self.max_payload_size = int(server_config.get('max_payload_size', 8192))
# 连接管理
self.connections: Dict[int, MQTTConnection] = {}
self.client_id_map: Dict[str, int] = {}
self.udp_handlers: Dict[int, UDPAudioHandler] = {}
self.connection_id_counter = 0
self._client_id_lock = asyncio.Lock()
self._client_locks = weakref.WeakValueDictionary()
self._client_reservations: Dict[str, int] = {}
self._pending_connection_ids: Set[int] = set()
self._cleanup_tasks: Dict[int, asyncio.Task] = {}
self._draining_tasks: Set[asyncio.Task] = set()
self._connection_handler_tasks: Set[asyncio.Task] = set()
# 服务器实例
self.mqtt_server = None
self.udp_server = None
# 连接服务
self.connection_service = ConnectionService(config)
self.connection_service.server = self
# 活跃连接管理
self.active_transports: Set[MQTTTransport] = set()
self.connection_registry = NativeMqttConnectionRegistry()
self.call_manager = NativeMqttCallManager(
self.connection_registry,
timeout_seconds=server_config.get("call_timeout", 60),
silence_frame=self._create_silence_frame(),
)
# 心跳检查
self.heartbeat_task = None
self.heartbeat_interval = int(server_config.get('heartbeat_interval', 30))
self._stop_event = asyncio.Event()
self._started_event = asyncio.Event()
self._is_running = False
self._stopping = False
self._shutdown_lock = asyncio.Lock()
self._shutdown_task = None
async def start(self):
"""启动MQTT服务器"""
if self._is_running:
logger.warning("MQTT服务器已经在运行中")
return
self._stop_event.clear()
self._started_event.clear()
self._stopping = False
try:
# 启动MQTT TCP服务器
await self._start_mqtt_server()
# 启动UDP服务器
await self._start_udp_server()
# 启动心跳检查
self.heartbeat_task = asyncio.create_task(self._heartbeat_check())
self._is_running = True
self._started_event.set()
logger.info(f"MQTT服务器启动成功: {self.host}:{self.mqtt_port}")
logger.info(f"UDP服务器启动成功: {self.udp_bind_host}:{self.udp_port}")
# 与 WebSocket 服务器保持一致,由上层持有该长期任务。
await self._stop_event.wait()
except Exception as e:
logger.error(f"启动MQTT服务器失败: {e}")
await self._await_shutdown_owner()
raise
finally:
self._is_running = False
self._started_event.clear()
async def _start_mqtt_server(self):
"""启动MQTT TCP服务器"""
self.mqtt_server = await asyncio.start_server(
self._accept_mqtt_connection,
self.host,
self.mqtt_port
)
if self.mqtt_port == 0 and self.mqtt_server.sockets:
self.mqtt_port = self.mqtt_server.sockets[0].getsockname()[1]
async def _start_udp_server(self):
"""启动UDP服务器"""
loop = asyncio.get_event_loop()
sock = None
bind_candidates = self._udp_bind_candidates()
for index, bind_host in enumerate(bind_candidates):
candidate = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
try:
candidate.bind((bind_host, self.udp_port))
except OSError as error:
candidate.close()
if index + 1 >= len(bind_candidates):
raise
logger.warning(
"UDP无法绑定到对外地址 {},回退监听 {}: {}",
bind_host,
bind_candidates[index + 1],
error,
)
continue
sock = candidate
self.udp_bind_host = bind_host
break
if sock is None:
raise RuntimeError("UDP socket bind failed")
sock.setblocking(False)
# 创建UDP协议处理器
transport, protocol = await loop.create_datagram_endpoint(
lambda: UDPProtocol(self),
sock=sock
)
self.udp_server = (transport, protocol)
if self.udp_port == 0:
self.udp_port = transport.get_extra_info("sockname")[1]
def _udp_bind_candidates(self):
"""Prefer a local advertised IPv4 address so UDP replies keep that source IP."""
explicit_bind_host = str(self._udp_bind_host_config or "").strip()
if explicit_bind_host:
return [explicit_bind_host]
wildcard_hosts = {"", "0.0.0.0"}
if self.host not in wildcard_hosts:
return [self.host]
try:
public_ip = ipaddress.ip_address(self.public_endpoint)
except ValueError:
return [self.host]
if public_ip.version == 4 and not public_ip.is_unspecified:
return [self.public_endpoint, self.host]
return [self.host]
def _accept_mqtt_connection(self, reader, writer) -> None:
"""Track every accepted callback before its coroutine can yield."""
task = asyncio.create_task(
self._handle_mqtt_connection(reader, writer),
name="mqtt-connection-handler",
)
self._connection_handler_tasks.add(task)
def consume(completed_task):
self._connection_handler_tasks.discard(completed_task)
MQTTConnection._consume_background_task(
completed_task, "MQTT连接处理任务"
)
task.add_done_callback(consume)
async def _handle_mqtt_connection(self, reader, writer):
"""处理新的MQTT连接"""
if self._stopping:
await self._close_stream_writer(writer, "MQTT停服期间连接")
return
if (
self.max_pending_connections > 0
and len(self._pending_connection_ids) >= self.max_pending_connections
):
logger.warning(
"MQTT待认证连接数达到上限: {}", self.max_pending_connections
)
await self._close_stream_writer(writer, "MQTT待认证连接")
return
connection_id = self._generate_connection_id()
try:
# 获取客户端地址
client_addr = writer.get_extra_info('peername')
logger.info(f"新MQTT连接: {client_addr}, connection_id: {connection_id}")
# 创建UDP音频处理器
udp_handler = UDPAudioHandler(
connection_id,
self,
{}, # 加密配置将在hello消息中设置
allowed_remote_ip=client_addr[0] if client_addr else None,
)
# 创建MQTT连接处理器
mqtt_connection = MQTTConnection(
writer.get_extra_info('socket'),
connection_id,
self,
udp_handler,
reader=reader,
writer=writer
)
mqtt_connection.business_task = asyncio.current_task()
udp_handler.set_audio_interceptor(
lambda payload, timestamp: self.call_manager.route_audio(
mqtt_connection.device_id, payload, timestamp
)
)
# 创建MQTT传输层
transport = MQTTTransport(mqtt_connection, udp_handler)
# 注册连接
self.connections[connection_id] = mqtt_connection
self.active_transports.add(transport)
self._pending_connection_ids.add(connection_id)
try:
await asyncio.wait_for(
mqtt_connection.connect_processed_event.wait(), timeout=10
)
except asyncio.TimeoutError:
logger.warning(f"MQTT CONNECT超时: connection_id={connection_id}")
return
finally:
self._pending_connection_ids.discard(connection_id)
if not mqtt_connection.connect_accepted:
return
# 提取连接头信息
headers = {
'x-real-ip': client_addr[0] if client_addr else 'unknown',
'connection-type': 'mqtt'
}
try:
# 使用ConnectionService处理连接
await self.connection_service.handle_connection(transport, headers)
except Exception as e:
logger.error(f"ConnectionService处理MQTT连接失败: {e}")
except Exception as e:
logger.error(f"处理MQTT连接失败: {e}")
finally:
# 清理连接
await self._cleanup_connection(connection_id)
async def _cleanup_connection(self, connection_id: int):
"""Await the single cleanup owner without coupling it to caller cancellation."""
if (
connection_id not in self._cleanup_tasks
and not self._connection_has_resources(connection_id)
):
return
cleanup_task = self._ensure_cleanup_task(connection_id)
if cleanup_task is asyncio.current_task():
return
await asyncio.shield(cleanup_task)
def _ensure_cleanup_task(self, connection_id: int) -> asyncio.Task:
cleanup_task = self._cleanup_tasks.get(connection_id)
if cleanup_task is None:
cleanup_task = asyncio.create_task(
self._cleanup_connection_impl(connection_id),
name=f"mqtt-cleanup-{connection_id}",
)
self._cleanup_tasks[connection_id] = cleanup_task
return cleanup_task
def _connection_has_resources(self, connection_id: int) -> bool:
if connection_id in self._pending_connection_ids:
return True
if connection_id in self.connections:
return True
if connection_id in self.client_id_map.values():
return True
if connection_id in self._client_reservations.values():
return True
return any(
getattr(
getattr(transport, "_mqtt_connection", None),
"connection_id",
None,
)
== connection_id
for transport in self.active_transports
)
async def _cleanup_connection_impl(self, connection_id: int):
"""Physically release one connection before removing its registry entries."""
connection = self.connections.get(connection_id)
transports_to_remove = [
transport
for transport in list(self.active_transports)
if (
hasattr(transport, '_mqtt_connection')
and transport._mqtt_connection.connection_id == connection_id
)
]
try:
self._pending_connection_ids.discard(connection_id)
if connection is not None:
if (
connection.udp_config
and connection.is_connected()
and connection.reply_topic
):
try:
await asyncio.wait_for(
connection.notify_device_idle(),
timeout=self.goodbye_timeout,
)
except asyncio.TimeoutError:
logger.warning("MQTT连接关闭前发送goodbye超时")
except Exception as e:
logger.warning("MQTT连接关闭前发送goodbye失败: {}", e)
try:
await asyncio.wait_for(
connection.close(), timeout=self.close_timeout
)
except asyncio.TimeoutError:
logger.warning(
"MQTT连接关闭超时,强制中止socket: {}", connection_id
)
abort = getattr(connection.protocol, "abort", None)
if callable(abort):
abort()
close_task = getattr(connection, "_close_task", None)
if close_task is not None and not close_task.done():
self.track_draining_task(
close_task,
f"MQTT连接关闭任务({connection_id})",
)
# UDP 使用每个 Hello 独立的随机 cookie,清理该物理连接的全部别名。
udp_handler = connection.udp_handler if connection is not None else None
if udp_handler is not None:
for session_id, registered in list(self.udp_handlers.items()):
if registered is udp_handler:
self.udp_handlers.pop(session_id, None)
try:
await asyncio.wait_for(
udp_handler.close(), timeout=self.close_timeout
)
except asyncio.TimeoutError:
logger.warning(
"MQTT UDP处理器关闭超时: {}", connection_id
)
for transport in transports_to_remove:
try:
await asyncio.wait_for(
transport.close(), timeout=self.close_timeout
)
except asyncio.TimeoutError:
logger.warning("MQTT transport关闭超时: {}", connection_id)
abort = getattr(
getattr(transport, "_mqtt_connection", None),
"protocol",
None,
)
abort = getattr(abort, "abort", None)
if callable(abort):
abort()
except asyncio.CancelledError:
raise
except Exception as e:
logger.error(f"清理MQTT连接失败: {e}")
finally:
# A cleanup owner may be cancelled during process shutdown. Always
# force the physical socket closed before publishing registry removal.
if connection is not None:
abort = getattr(connection.protocol, "abort", None)
if callable(abort):
abort()
close_task = getattr(connection, "_close_task", None)
if close_task is not None and not close_task.done():
self.track_draining_task(
close_task,
f"MQTT连接关闭任务({connection_id})",
)
for transport in transports_to_remove:
self.active_transports.discard(transport)
async with self._client_id_lock:
for client_id, mapped_id in list(self.client_id_map.items()):
if mapped_id == connection_id:
self.client_id_map.pop(client_id, None)
for client_id, reserved_id in list(
self._client_reservations.items()
):
if reserved_id == connection_id:
self._client_reservations.pop(client_id, None)
self.connections.pop(connection_id, None)
current_task = asyncio.current_task()
if self._cleanup_tasks.get(connection_id) is current_task:
self._cleanup_tasks.pop(connection_id, None)
logger.info(f"MQTT连接清理完成: {connection_id}")
async def _close_stream_writer(self, writer, label: str) -> None:
"""Bound close for sockets rejected before MQTTConnection exists."""
writer.close()
try:
await asyncio.wait_for(
writer.wait_closed(), timeout=self.close_timeout
)
except asyncio.TimeoutError:
logger.warning("{}关闭超时,强制中止transport", label)
transport = getattr(writer, "transport", None)
if transport is not None:
transport.abort()
except Exception:
pass
def track_draining_task(self, task: asyncio.Task, label: str) -> None:
"""Keep cancellation-resistant work visible until it actually exits."""
if task is None or task.done():
if task is not None:
MQTTConnection._consume_background_task(task, label)
return
if task in self._draining_tasks:
return
self._draining_tasks.add(task)
def consume(completed_task):
self._draining_tasks.discard(completed_task)
MQTTConnection._consume_background_task(completed_task, label)
task.add_done_callback(consume)
def _generate_connection_id(self) -> int:
"""生成连接ID"""
self.connection_id_counter += 1
return self.connection_id_counter
async def _get_client_lock(self, client_id: str) -> asyncio.Lock:
async with self._client_id_lock:
lock = self._client_locks.get(client_id)
if lock is None:
lock = asyncio.Lock()
self._client_locks[client_id] = lock
return lock
async def on_client_connected(
self, mqtt_connection: MQTTConnection, accept_callback
) -> bool:
"""Atomically replace one clientId before acknowledging the new owner."""
client_id = mqtt_connection.client_id
if not client_id or self._stopping:
return False
client_lock = await self._get_client_lock(client_id)
async with client_lock:
if mqtt_connection._closed or self._stopping:
return False
reservation_active = False
async with self._client_id_lock:
old_id = self.client_id_map.get(client_id)
identity_count = len(
set(self.client_id_map) | set(self._client_reservations)
)
if (
old_id is None
and client_id not in self._client_reservations
and self.max_connections > 0
and identity_count >= self.max_connections
):
logger.warning(
"MQTT活跃客户端数达到上限,拒绝clientId: {}", client_id
)
return False
self._client_reservations[client_id] = (
mqtt_connection.connection_id
)
reservation_active = True
try:
if old_id and old_id != mqtt_connection.connection_id:
logger.info(f"检测到重复clientId,关闭旧连接: {client_id}")
cleanup_task = self._ensure_cleanup_task(old_id)
try:
await asyncio.shield(cleanup_task)
except asyncio.CancelledError as cancelled:
# Keep the per-client lock and capacity reservation until
# the predecessor's physical close barrier has completed.
try:
await asyncio.shield(cleanup_task)
except Exception as cleanup_error:
logger.error(
"取消接管时等待旧连接清理失败: {}", cleanup_error
)
raise cancelled
if mqtt_connection._closed or self._stopping:
return False
async with self._client_id_lock:
if (
self._client_reservations.get(client_id)
!= mqtt_connection.connection_id
):
return False
self.client_id_map[client_id] = mqtt_connection.connection_id
self._client_reservations.pop(client_id, None)
reservation_active = False
await accept_callback()
except BaseException:
async with self._client_id_lock:
if (
self.client_id_map.get(client_id)
== mqtt_connection.connection_id
):
self.client_id_map.pop(client_id, None)
raise
finally:
if reservation_active:
async with self._client_id_lock:
if (
self._client_reservations.get(client_id)
== mqtt_connection.connection_id
):
self._client_reservations.pop(client_id, None)
logger.info(f"MQTT客户端已连接: {client_id}")
return True
async def on_client_disconnected(self, mqtt_connection: MQTTConnection):
"""客户端断开连接回调"""
logger.info(f"MQTT客户端已断开: {mqtt_connection.client_id}")
# Registry removal belongs to the physical cleanup owner. Scheduling
# rather than awaiting avoids a close -> disconnect -> cleanup cycle.
connection_id = mqtt_connection.connection_id
if (
connection_id in self.connections
and connection_id not in self._cleanup_tasks
):
self._ensure_cleanup_task(connection_id)
async def register_connection_context(self, context, transport) -> bool:
existing = self.connection_registry.resolve_device_now(
getattr(context, "device_id", None)
)
registered = await self.connection_registry.register(
context, transport
)
if (
registered
and existing is not None
and existing.transport is not transport
):
await self.call_manager.end_call(
context.device_id,
"设备连接被接管",
notify_device=False,
notify_peer=True,
expected_session_id=getattr(
existing.transport, "session_id", None
),
)
try:
await existing.transport.close()
except Exception as error:
logger.warning(
"关闭被接管的MQTT连接失败: device_id={}, error={}",
context.device_id,
error,
)
return registered
async def unregister_connection_context(self, context, transport) -> bool:
removed = await self.connection_registry.unregister(context, transport)
if removed:
await self.call_manager.end_call(
context.device_id,
"设备连接已断开",
notify_device=False,
notify_peer=True,
expected_session_id=getattr(transport, "session_id", None),
)
return removed
async def resolve_connection_context(self, client_id: str):
return await self.connection_registry.resolve(client_id)
async def get_connection_status(self, client_ids):
return await self.connection_registry.status(client_ids)
async def request_device_call(
self, caller_mac: str, target_mac: str, caller_nickname: str = ""
):
return await self.call_manager.request_call(
caller_mac, target_mac, caller_nickname
)
async def accept_device_call(self, device_id: str):
return await self.call_manager.accept_call(device_id)
async def end_native_mqtt_call(
self,
device_id: str,
reason: str = "",
notify_device: bool = False,
expected_session_id: str = None,
expected_generation: int = None,
) -> bool:
return await self.call_manager.end_call(
device_id,
reason,
notify_device=notify_device,
notify_peer=True,
expected_session_id=expected_session_id,
expected_generation=expected_generation,
)
async def handle_logical_hello(self, mqtt_connection) -> None:
await self.call_manager.handle_logical_hello(
mqtt_connection.device_id,
mqtt_connection.session_id,
)
def bind_udp_session(
self, mqtt_connection: MQTTConnection, udp_handler: UDPAudioHandler
) -> int:
"""Bind one unpredictable UDP cookie to the current Hello session."""
for session_id, registered in list(self.udp_handlers.items()):
if registered is udp_handler:
self.udp_handlers.pop(session_id, None)
while True:
session_id = secrets.randbits(32)
if session_id and session_id not in self.udp_handlers:
break
udp_handler.connection_id = session_id
self.udp_handlers[session_id] = udp_handler
return session_id
async def send_udp_message(self, data: bytes, remote_addr: tuple):
"""发送UDP消息"""
if not self.udp_server:
raise RuntimeError("UDP server is not running")
transport, protocol = self.udp_server
transport.sendto(data, remote_addr)
async def send_encrypted_audio(self, connection_id: int, audio_data: bytes,
timestamp: int, sequence: int, remote_addr: tuple,
encryption_config: Dict[str, Any]):
"""发送加密音频数据"""
self.send_encrypted_audio_nowait(
connection_id,
audio_data,
timestamp,
sequence,
remote_addr,
encryption_config,
)
def send_encrypted_audio_nowait(
self,
connection_id: int,
audio_data: bytes,
timestamp: int,
sequence: int,
remote_addr: tuple,
encryption_config: Dict[str, Any],
):
header = self._generate_udp_header(connection_id, len(audio_data), timestamp, sequence)
key = encryption_config.get('key') if encryption_config else None
if key:
cipher = Cipher(algorithms.AES(key), modes.CTR(header), backend=default_backend())
encryptor = cipher.encryptor()
encrypted = encryptor.update(audio_data) + encryptor.finalize()
message = header + encrypted
else:
message = header + audio_data
if not self.udp_server:
raise RuntimeError("UDP server is not running")
transport, _ = self.udp_server
transport.sendto(message, remote_addr)
def _generate_udp_header(self, connection_id: int, length: int, timestamp: int, sequence: int) -> bytes:
"""生成UDP消息头"""
header = bytearray(16)
header[0] = 1 # type
header[2:4] = length.to_bytes(2, 'big') # payload length
header[4:8] = connection_id.to_bytes(4, 'big') # connection id
header[8:12] = timestamp.to_bytes(4, 'big') # timestamp
header[12:16] = sequence.to_bytes(4, 'big') # sequence
return bytes(header)
async def _heartbeat_check(self):
"""心跳检查任务"""
try:
while True:
await asyncio.sleep(self.heartbeat_interval)
# 检查所有连接的状态
dead_connections = []
for connection_id, connection in self.connections.items():
if not connection.is_connected():
dead_connections.append(connection_id)
# 清理死连接
for connection_id in dead_connections:
logger.info(f"清理死连接: {connection_id}")
await self._cleanup_connection(connection_id)
await self.call_manager.cleanup_timeouts()
# 记录活跃连接数
active_count = len(self.connections)
if active_count > 0:
logger.info(f"MQTT活跃连接数: {active_count}")
except asyncio.CancelledError:
pass
except Exception as e:
logger.error(f"心跳检查任务出错: {e}")
async def stop(self):
"""停止MQTT服务器"""
async with self._shutdown_lock:
self._prune_finished_lifecycle_tasks()
shutdown_in_progress = (
self._shutdown_task is not None
and not self._shutdown_task.done()
)
if (
not self._is_running
and self._shutdown_resources_released()
and not shutdown_in_progress
):
self._stop_event.set()
return True
logger.info("正在停止MQTT服务器...")
shutdown_task = self._ensure_shutdown_owner()
shutdown_result = await asyncio.shield(shutdown_task)
# Existing lifecycle test doubles and older protocol servers return
# None on success. Only an explicit False means cleanup is partial.
complete = shutdown_result is not False
if complete:
logger.info("MQTT服务器已停止")
else:
status = self.get_server_status()
logger.error(
"MQTT监听器已停止,但仍有未退出任务: handlers={}, draining={}",
status["connection_handler_tasks"],
status["draining_tasks"],
)
return complete
def _ensure_shutdown_owner(self) -> asyncio.Task:
"""Create or reuse the single physical shutdown owner."""
if self._shutdown_task is None or self._shutdown_task.done():
self._shutdown_task = asyncio.create_task(
self._run_shutdown_owner()
)
return self._shutdown_task
async def _await_shutdown_owner(self):
"""Wait for shutdown without letting caller cancellation kill it."""
return await asyncio.shield(self._ensure_shutdown_owner())
async def _run_shutdown_owner(self):
try:
return await self._shutdown_resources()
finally:
# A cancelled stop waiter must not leave start() blocked forever.
self._stop_event.set()
self._is_running = False
self._started_event.clear()
def _prune_finished_lifecycle_tasks(self) -> None:
for connection_id, task in list(self._cleanup_tasks.items()):
if task.done() and self._cleanup_tasks.get(connection_id) is task:
self._cleanup_tasks.pop(connection_id, None)
for task in list(self._connection_handler_tasks):
if task.done():
self._connection_handler_tasks.discard(task)
for task in list(self._draining_tasks):
if task.done():
self._draining_tasks.discard(task)
def _shutdown_resources_released(self) -> bool:
"""Return true only when every MQTT-owned resource is gone."""
heartbeat_alive = (
self.heartbeat_task is not None
and not self.heartbeat_task.done()
)
return not (
self.mqtt_server
or self.udp_server
or heartbeat_alive
or self.connections
or self.client_id_map
or self.udp_handlers
or self.active_transports
or self._client_reservations
or self._pending_connection_ids
or self._cleanup_tasks
or self.connection_registry.count
or self.call_manager.count
or self.call_manager.background_task_count
or any(
not task.done()
for task in self._connection_handler_tasks
)
or any(not task.done() for task in self._draining_tasks)
)
async def _shutdown_resources(self):
"""幂等释放 MQTT/UDP 监听器、连接和后台任务。"""
self._stopping = True
loop = asyncio.get_running_loop()
deadline = loop.time() + self.shutdown_timeout
# Stop accepting before taking any connection/task snapshot.
if self.mqtt_server:
self.mqtt_server.close()
await self.mqtt_server.wait_closed()
self.mqtt_server = None
# Let callbacks already queued by the event loop enter the synchronous
# tracking wrapper before handler/connection snapshots are taken.
await asyncio.sleep(0)
if self.heartbeat_task and not self.heartbeat_task.done():
self.heartbeat_task.cancel()
await asyncio.gather(self.heartbeat_task, return_exceptions=True)
self.heartbeat_task = None
# Start every cleanup owner first, then wait against one global deadline.
connection_ids = (
set(self.connections)
| set(self._cleanup_tasks)
| set(self.client_id_map.values())
| set(self._client_reservations.values())
| set(self._pending_connection_ids)
)
connection_ids.update(
connection_id
for connection_id in (
getattr(
getattr(transport, "_mqtt_connection", None),
"connection_id",
None,
)
for transport in self.active_transports
)
if connection_id is not None
)
cleanup_tasks = [
self._ensure_cleanup_task(connection_id)
for connection_id in connection_ids
if self._connection_has_resources(connection_id)
or connection_id in self._cleanup_tasks
]
pending_cleanup = await self._wait_until_deadline(
cleanup_tasks, deadline
)
if pending_cleanup:
logger.warning(
"MQTT停服清理超过全局时限,强制取消{}个清理任务",
len(pending_cleanup),
)
for task in pending_cleanup:
task.cancel()
await asyncio.gather(*pending_cleanup, return_exceptions=True)
# Accepted callbacks are tracked synchronously by _accept_mqtt_connection.
handler_tasks = [
task
for task in self._connection_handler_tasks
if not task.done()
]
for task in handler_tasks:
task.cancel()
pending_handlers = await self._wait_until_deadline(
handler_tasks, deadline
)
for task in pending_handlers:
self.track_draining_task(task, "MQTT连接处理任务")
draining_tasks = [
task for task in self._draining_tasks if not task.done()
]
for task in draining_tasks:
task.cancel()
pending_draining = await self._wait_until_deadline(
draining_tasks, deadline
)
if pending_draining:
logger.warning(
"MQTT停服后仍有{}个取消不响应任务", len(pending_draining)
)
# Defensive closure for any transport whose connection registration was
# corrupted by an earlier failure.
transport_tasks = [
asyncio.create_task(self._close_transport(transport))
for transport in list(self.active_transports)
]
pending_transports = await self._wait_until_deadline(
transport_tasks, deadline
)
for task in pending_transports:
task.cancel()
if pending_transports:
await asyncio.gather(
*pending_transports, return_exceptions=True
)
# Close any orphaned UDP session left behind by a corrupted registry.
for udp_handler in set(self.udp_handlers.values()):
await udp_handler.close()
self.udp_handlers.clear()
# 关闭UDP服务器
if self.udp_server:
transport, protocol = self.udp_server
transport.close()
self.udp_server = None
await self.call_manager.clear()
await self.connection_registry.clear()
self._prune_finished_lifecycle_tasks()
return self._shutdown_resources_released()
@staticmethod
def _create_silence_frame():
try:
import opuslib_next
encoder = opuslib_next.Encoder(
16000, 1, opuslib_next.APPLICATION_AUDIO
)
return encoder.encode(bytes(960 * 2), 960)
except Exception as error:
logger.warning("Native MQTT通话静音帧初始化失败: {}", error)
return None
async def _wait_until_deadline(self, tasks, deadline):
pending = {task for task in tasks if task is not None and not task.done()}
if not pending:
return set()
remaining = max(0.0, deadline - asyncio.get_running_loop().time())
_, pending = await asyncio.wait(pending, timeout=remaining)
return pending
async def _close_transport(self, transport) -> None:
try:
await asyncio.wait_for(
transport.close(), timeout=self.close_timeout
)
except asyncio.TimeoutError:
connection = getattr(transport, "_mqtt_connection", None)
abort = getattr(
getattr(connection, "protocol", None), "abort", None
)
if callable(abort):
abort()
finally:
self.active_transports.discard(transport)
async def update_config(self, new_config: Dict[str, Any]) -> bool:
"""更新非监听配置;监听地址变化由 MultiProtocolServer 重建实例。"""
try:
self.config = new_config
server_config = new_config.get('mqtt_server', {})
self.public_endpoint = self._resolve_public_host(server_config)
self._udp_bind_host_config = server_config.get('udp_bind_host')
signature_key = self._resolve_signature_key(new_config)
if not signature_key:
raise ValueError(
"启用原生MQTT必须配置mqtt_server.signature_key或"
"server.mqtt_signature_key"
)
self.signature_key = signature_key
self.message_queue_size = int(server_config.get('message_queue_size', 128))
self.business_ready_timeout = float(
server_config.get('business_ready_timeout', 30) or 30
)
self.close_timeout = max(
0.1, float(server_config.get('close_timeout', 2) or 2)
)
self.shutdown_timeout = max(
self.close_timeout,
float(server_config.get('shutdown_timeout', 10) or 10),
)
self.goodbye_timeout = max(
0.1, float(server_config.get('goodbye_timeout', 1) or 1)
)
self.max_connections = int(server_config.get('max_connections', 1000))
self.max_pending_connections = int(
server_config.get('max_pending_connections', 128)
)
self.max_payload_size = int(server_config.get('max_payload_size', 8192))
self.heartbeat_interval = int(server_config.get('heartbeat_interval', 30))
# 已建立连接保留会话快照,新连接使用新配置。
self.connection_service = ConnectionService(new_config)
self.connection_service.server = getattr(self, "management_owner", self)
return True
except Exception as e:
logger.error(f"更新MQTT服务器配置失败: {e}")
return False
@staticmethod
def _resolve_signature_key(config: Dict[str, Any]) -> str:
server_config = config.get('mqtt_server', {})
return (
normalize_signature_key(server_config.get('signature_key'))
or normalize_signature_key(
config.get('server', {}).get('mqtt_signature_key')
)
)
@staticmethod
def _resolve_public_host(server_config: Dict[str, Any]) -> str:
host, _ = parse_mqtt_endpoint(
server_config.get('public_endpoint'),
None,
)
return host or 'localhost'
async def apply_config(self, new_config: Dict[str, Any]) -> bool:
return await self.update_config(new_config)
def get_server_status(self) -> Dict[str, Any]:
"""获取服务器状态"""
return {
'type': 'mqtt',
'host': self.host,
'mqtt_port': self.mqtt_port,
'udp_port': self.udp_port,
'active_connections': len(self.connections),
'active_transports': len(self.active_transports),
'pending_connections': len(self._pending_connection_ids),
'cleanup_tasks': len(self._cleanup_tasks),
'draining_tasks': sum(
not task.done() for task in self._draining_tasks
),
'connection_handler_tasks': sum(
not task.done()
for task in self._connection_handler_tasks
),
'stopping': self._stopping,
}
class UDPProtocol(asyncio.DatagramProtocol):
"""UDP协议处理器"""
def __init__(self, mqtt_server: MQTTServer):
self.mqtt_server = mqtt_server
self.transport = None
def connection_made(self, transport):
self.transport = transport
def datagram_received(self, data: bytes, addr: tuple):
"""接收UDP数据报"""
try:
# 解析UDP消息头
if len(data) < 16:
return
packet_type = data[0]
if packet_type != 1:
return
payload_length = int.from_bytes(data[2:4], 'big')
if len(data) != 16 + payload_length:
logger.warning(
"UDP数据报长度不匹配: actual={}, expected={}",
len(data),
16 + payload_length,
)
return
connection_id = int.from_bytes(data[4:8], 'big')
timestamp = int.from_bytes(data[8:12], 'big')
sequence = int.from_bytes(data[12:16], 'big')
header = data[:16]
encrypted_payload = data[16:16 + payload_length]
# 找到对应的UDP处理器
udp_handler = self.mqtt_server.udp_handlers.get(connection_id)
if udp_handler:
udp_handler.on_udp_message(header, encrypted_payload, payload_length, timestamp, sequence, addr)
except Exception as e:
logger.error(f"处理UDP数据报失败: {e}")
def error_received(self, exc):
logger.error(f"UDP协议错误: {exc}")