feat: add optional native mqtt and udp transport

This commit is contained in:
caixypromise
2026-07-27 02:10:55 +08:00
parent 0c582ed3b6
commit eac573706d
35 changed files with 6325 additions and 200 deletions
@@ -0,0 +1,323 @@
import copy
import hashlib
import hmac
from datetime import datetime, timedelta, timezone
from typing import Any, Dict, Optional
from aiohttp import web
from config.logger import setup_logging
from core.providers.tools.device_mcp import call_mcp_tool
from core.utils.mqtt_auth import normalize_signature_key
from core.utils.util import sanitize_tool_name
TAG = __name__
class NativeMqttManagementHandler:
def __init__(self, config: Dict[str, Any], management_owner: Any):
self.config = config
self.management_owner = management_owner
self.logger = setup_logging()
mqtt_config = config.get("mqtt_server", {})
server_config = config.get("server", {})
self.signature_key = normalize_signature_key(
mqtt_config.get("manager_api_secret")
or mqtt_config.get("signature_key")
or server_config.get("mqtt_signature_key")
)
self.command_timeout = max(
0.1, float(mqtt_config.get("manager_command_timeout", 5) or 5)
)
self.max_status_ids = max(
1, int(mqtt_config.get("manager_max_status_ids", 1000) or 1000)
)
@staticmethod
def generate_daily_tokens(
signature_key: str, now: Optional[datetime] = None
) -> set[str]:
normalized = normalize_signature_key(signature_key)
if not normalized:
return set()
current = now or datetime.now(timezone.utc)
utc_date = current.astimezone(timezone.utc).date()
return {
hashlib.sha256(
f"{utc_date + timedelta(days=offset)}{normalized}".encode(
"utf-8"
)
).hexdigest()
for offset in (-1, 0, 1)
}
def _is_authorized(self, authorization: str) -> bool:
if not self.signature_key:
return False
if not authorization or not authorization.startswith("Bearer "):
return False
provided = authorization[len("Bearer ") :].strip()
return any(
hmac.compare_digest(provided, expected)
for expected in self.generate_daily_tokens(self.signature_key)
)
def _authorize(self, request) -> Optional[web.Response]:
if not self.signature_key:
return self._error(
503,
"Native MQTT管理API未配置签名密钥",
"MANAGEMENT_AUTH_NOT_CONFIGURED",
False,
)
if not self._is_authorized(request.headers.get("Authorization", "")):
return self._error(
401, "无效的授权令牌", "UNAUTHORIZED", False
)
return None
@staticmethod
def _error(
status: int,
message: str,
code: str,
dispatch_attempted: bool,
) -> web.Response:
return web.json_response(
{
"success": False,
"error": message,
"code": code,
"dispatchAttempted": dispatch_attempted,
},
status=status,
)
async def _read_json_object(self, request) -> Optional[Dict[str, Any]]:
try:
body = await request.json()
except Exception:
return None
return body if isinstance(body, dict) else None
async def handle_device_status(self, request) -> web.Response:
unauthorized = self._authorize(request)
if unauthorized is not None:
return unauthorized
body = await self._read_json_object(request)
client_ids = body.get("clientIds") if body else None
if (
not isinstance(client_ids, list)
or not client_ids
or len(client_ids) > self.max_status_ids
or any(
not isinstance(client_id, str) or not client_id
for client_id in client_ids
)
):
return self._error(
400,
"clientIds必须是非空字符串数组且未超过数量限制",
"INVALID_CLIENT_IDS",
False,
)
get_status = getattr(
self.management_owner, "get_native_mqtt_status", None
)
if not callable(get_status):
return self._error(
503,
"Native MQTT管理服务未就绪",
"MANAGEMENT_NOT_READY",
False,
)
return web.json_response(await get_status(client_ids))
async def handle_command(self, request) -> web.Response:
unauthorized = self._authorize(request)
if unauthorized is not None:
return unauthorized
body = await self._read_json_object(request)
payload = (
body.get("payload")
if body and body.get("type") == "mcp"
else None
)
if not isinstance(payload, dict):
return self._error(
400, "指令类型无效", "INVALID_COMMAND", False
)
resolver = getattr(
self.management_owner, "resolve_native_mqtt_connection", None
)
if not callable(resolver):
return self._error(
503,
"Native MQTT管理服务未就绪",
"MANAGEMENT_NOT_READY",
False,
)
connection = await resolver(request.match_info.get("client_id", ""))
if connection is None:
return self._error(
404, "设备未连接", "DEVICE_OFFLINE", False
)
method = payload.get("method")
params = payload.get("params") or {}
if not isinstance(params, dict):
return self._error(
400, "MCP参数格式无效", "INVALID_MCP_PARAMS", False
)
if method == "tools/list":
return await self._list_tools(connection.context)
if method == "tools/call":
return await self._call_tool(connection.context, params)
return self._error(
422, "不支持的MCP方法", "UNSUPPORTED_MCP_METHOD", False
)
async def handle_call_request(self, request) -> web.Response:
unauthorized = self._authorize(request)
if unauthorized is not None:
return unauthorized
body = await self._read_json_object(request)
caller_mac = body.get("caller_mac") if body else None
target_mac = body.get("target_mac") if body else None
caller_nickname = body.get("caller_nickname", "") if body else ""
if (
not isinstance(caller_mac, str)
or not caller_mac.strip()
or not isinstance(target_mac, str)
or not target_mac.strip()
or not isinstance(caller_nickname, str)
):
return web.json_response(
{
"status": "error",
"message": "缺少必要参数: caller_mac, target_mac",
},
status=400,
)
request_call = getattr(
self.management_owner, "request_native_mqtt_call", None
)
if not callable(request_call):
return web.json_response(
{"status": "error", "message": "Native MQTT呼叫服务未就绪"},
status=503,
)
result = await request_call(
caller_mac, target_mac, caller_nickname
)
return web.json_response(result)
async def handle_call_accept(self, request) -> web.Response:
unauthorized = self._authorize(request)
if unauthorized is not None:
return unauthorized
body = await self._read_json_object(request)
device_id = body.get("mac") if body else None
if not isinstance(device_id, str) or not device_id.strip():
return web.json_response(
{"status": "error", "message": "缺少必要参数: mac"},
status=400,
)
accept_call = getattr(
self.management_owner, "accept_native_mqtt_call", None
)
if not callable(accept_call):
return web.json_response(
{"status": "error", "message": "Native MQTT呼叫服务未就绪"},
status=503,
)
return web.json_response(await accept_call(device_id))
async def _list_tools(self, context) -> web.Response:
mcp_client = getattr(context, "mcp_client", None)
if mcp_client is None or not await mcp_client.is_ready():
return self._error(
503,
"设备MCP尚未准备就绪",
"MCP_NOT_READY",
False,
)
async with mcp_client.lock:
tools = [
copy.deepcopy(tool)
for tool in mcp_client.tools.values()
]
return web.json_response(
{"success": True, "data": {"tools": tools}}
)
async def _call_tool(self, context, params: Dict[str, Any]) -> web.Response:
tool_name = params.get("name")
arguments = params.get("arguments", {})
if not isinstance(tool_name, str) or not tool_name:
return self._error(
422, "工具名称不能为空", "INVALID_TOOL_NAME", False
)
if not isinstance(arguments, dict):
return self._error(
422, "工具参数必须是对象", "INVALID_TOOL_ARGUMENTS", False
)
mcp_client = getattr(context, "mcp_client", None)
if mcp_client is None or not await mcp_client.is_ready():
return self._error(
503,
"设备MCP尚未准备就绪",
"MCP_NOT_READY",
False,
)
sanitized_name = sanitize_tool_name(tool_name)
if not mcp_client.has_tool(sanitized_name):
return self._error(
422, "设备不存在该工具", "TOOL_NOT_FOUND", False
)
try:
result = await call_mcp_tool(
context,
mcp_client,
sanitized_name,
arguments,
timeout=self.command_timeout,
return_raw=True,
)
except TimeoutError:
return self._error(
504, "工具调用请求超时", "COMMAND_TIMEOUT", True
)
except ConnectionError:
return self._error(
503, "设备连接已关闭", "DEVICE_DISCONNECTED", True
)
except ValueError as error:
return self._error(422, str(error), "INVALID_TOOL_CALL", False)
except Exception as error:
self.logger.bind(tag=TAG).warning(
"Native MQTT设备工具调用失败: {}", error
)
return self._error(
502, str(error), "TOOL_CALL_FAILED", True
)
data = (
result
if isinstance(result, dict)
else {"content": [{"type": "text", "text": str(result)}]}
)
return web.json_response({"success": True, "data": data})
+118 -54
View File
@@ -1,8 +1,6 @@
import json
import time
import base64
import hashlib
import hmac
import os
import re
import glob
@@ -11,6 +9,11 @@ from aiohttp import web
from core.auth import AuthManager
from core.utils.util import get_local_ip, get_vision_url
from core.utils.mqtt_auth import (
generate_password_signature,
normalize_signature_key,
parse_mqtt_endpoint,
)
from core.api.base_handler import BaseHandler
TAG = __name__
@@ -102,26 +105,6 @@ class OTAHandler(BaseHandler):
self.logger.bind(tag=TAG).error(f"刷新固件缓存失败: {e}")
# keep previous cache if any
def generate_password_signature(self, content: str, secret_key: str) -> str:
"""生成MQTT密码签名
Args:
content: 签名内容 (clientId + '|' + username)
secret_key: 密钥
Returns:
str: Base64编码的HMAC-SHA256签名
"""
try:
hmac_obj = hmac.new(
secret_key.encode("utf-8"), content.encode("utf-8"), hashlib.sha256
)
signature = hmac_obj.digest()
return base64.b64encode(signature).decode("utf-8")
except Exception as e:
self.logger.bind(tag=TAG).error(f"生成MQTT密码签名失败: {e}")
return ""
def _get_websocket_url(self, local_ip: str, port: int) -> str:
"""获取websocket地址
@@ -231,11 +214,32 @@ class OTAHandler(BaseHandler):
},
}
# existing mqtt/websocket logic (unchanged)
mqtt_gateway_endpoint = server_config.get("mqtt_gateway")
# ========== 协议下发逻辑 ==========
# 按照原版逻辑:总是下发 WebSocket,如果启用了 MQTT 则额外下发 MQTT 和 UDP
# 这样设备有回退能力:如果 MQTT 连接失败,还可以使用 WebSocket
if mqtt_gateway_endpoint: # 如果配置了非空字符串
# 尝试从请求数据中获取设备型号(已解析 above)
mqtt_server_config = self.config.get("mqtt_server", {})
enabled_protocols = self.config.get("enabled_protocols")
if isinstance(enabled_protocols, list):
mqtt_protocol_enabled = "mqtt" in enabled_protocols
else:
protocol_config = self.config.get("protocols", {})
requested_protocols = protocol_config.get(
"enabled_protocols", []
)
mqtt_protocol_enabled = (
protocol_config.get("mqtt_enabled") is True
or "mqtt" in requested_protocols
)
mqtt_server_enabled = bool(
mqtt_server_config.get("enabled") and mqtt_protocol_enabled
)
mqtt_gateway_endpoint = server_config.get("mqtt_gateway")
if not mqtt_gateway_endpoint or str(mqtt_gateway_endpoint).lower() == "null":
mqtt_gateway_endpoint = None
# 生成通用的 MQTT 凭证信息
def _build_mqtt_credentials():
try:
group_id = f"GID_{device_model}".replace(":", "_").replace(" ", "_")
except Exception as e:
@@ -246,56 +250,116 @@ class OTAHandler(BaseHandler):
mqtt_client_id = f"{group_id}@@@{mac_address_safe}@@@{mac_address_safe}"
# 构建用户数据
user_data = {"ip": "unknown"}
user_data = {"ip": local_ip}
try:
user_data_json = json.dumps(user_data)
username = base64.b64encode(user_data_json.encode("utf-8")).decode(
"utf-8"
)
username = base64.b64encode(user_data_json.encode("utf-8")).decode("utf-8")
except Exception as e:
self.logger.bind(tag=TAG).error(f"生成用户名失败: {e}")
username = ""
return group_id, mac_address_safe, mqtt_client_id, username
# ========== 1. 总是下发 WebSocket 配置(作为基础/回退方案)==========
ws_token = ""
if self.auth_enable:
if self.allowed_devices:
if device_id not in self.allowed_devices:
ws_token = self.auth.generate_token(client_id, device_id)
else:
ws_token = self.auth.generate_token(client_id, device_id)
return_json["websocket"] = {
"url": self._get_websocket_url(local_ip, websocket_port),
"token": ws_token,
}
# ========== 2. 如果启用了原生 MQTT 服务器,额外下发 MQTT 配置 ==========
signature_key = normalize_signature_key(
mqtt_server_config.get("signature_key")
or server_config.get("mqtt_signature_key")
)
native_mqtt_ready = bool(mqtt_server_enabled and signature_key)
if mqtt_server_enabled and not signature_key:
self.logger.bind(tag=TAG).error(
"原生MQTT已启用但未配置签名密钥,跳过Native配置下发"
)
if native_mqtt_ready:
try:
mqtt_host, mqtt_port = parse_mqtt_endpoint(
mqtt_server_config.get("public_endpoint"),
int(mqtt_server_config.get("port", 1883)),
)
except (TypeError, ValueError) as exc:
self.logger.bind(tag=TAG).error(f"MQTT endpoint配置无效: {exc}")
mqtt_host, mqtt_port = "", None
placeholder_keywords = ("localhost", "0.0.0.0", "your", "example")
if not mqtt_host or any(
keyword in mqtt_host.lower() for keyword in placeholder_keywords
):
mqtt_host = local_ip
self.logger.bind(tag=TAG).info(
f"检测到 public_endpoint 为占位符,自动使用本地IP: {local_ip}"
)
if mqtt_port is None:
native_mqtt_ready = False
if native_mqtt_ready:
group_id, mac_address_safe, mqtt_client_id, username = _build_mqtt_credentials()
mqtt_password = generate_password_signature(
mqtt_client_id + "|" + username, signature_key
)
return_json["mqtt"] = {
"endpoint": f"{mqtt_host}:{mqtt_port}",
"client_id": mqtt_client_id,
"username": username,
"password": mqtt_password,
"publish_topic": "device-server",
"subscribe_topic": f"devices/p2p/{mac_address_safe}",
}
self.logger.bind(tag=TAG).info(
f"为设备 {device_id} 下发原生MQTT配置: {mqtt_host}:{mqtt_port}"
)
# ========== 3. 如果配置了外部 MQTT 网关,额外下发 MQTT 配置 ==========
elif mqtt_gateway_endpoint:
group_id, mac_address_safe, mqtt_client_id, username = _build_mqtt_credentials()
# 生成密码
password = ""
mqtt_password = ""
signature_key = server_config.get("mqtt_signature_key", "")
if signature_key:
password = self.generate_password_signature(
mqtt_password = generate_password_signature(
mqtt_client_id + "|" + username, signature_key
)
if not password:
password = "" # 签名失败则留空,由设备决定是否允许无密码
if not mqtt_password:
mqtt_password = ""
else:
self.logger.bind(tag=TAG).warning("缺少MQTT签名密钥,密码留空")
# 构建MQTT配置(直接使用 mqtt_gateway 字符串)
return_json["mqtt"] = {
"endpoint": mqtt_gateway_endpoint,
"client_id": mqtt_client_id,
"username": username,
"password": password,
"password": mqtt_password,
"publish_topic": "device-server",
"subscribe_topic": f"devices/p2p/{mac_address_safe}",
}
self.logger.bind(tag=TAG).info(f"为设备 {device_id} 下发MQTT网关配置")
else: # 未配置 mqtt_gateway,下发 WebSocket
# 如果开启了认证,则进行认证校验
token = ""
if self.auth_enable:
if self.allowed_devices:
if device_id not in self.allowed_devices:
token = self.auth.generate_token(client_id, device_id)
else:
token = self.auth.generate_token(client_id, device_id)
# NOTE: use websocket_port here
return_json["websocket"] = {
"url": self._get_websocket_url(local_ip, websocket_port),
"token": token,
}
self.logger.bind(tag=TAG).info(
f"未配置MQTT网关,为设备 {device_id} 下发WebSocket配置"
)
self.logger.bind(tag=TAG).info(f"为设备 {device_id} 下发MQTT网关配置: {mqtt_gateway_endpoint}")
# 记录最终下发的协议
protocols = ["websocket"]
if "mqtt" in return_json:
protocols.append("mqtt")
self.logger.bind(tag=TAG).info(f"为设备 {device_id} 下发协议配置: {', '.join(protocols)}")
# Now check firmware files for updates
try:
+52 -2
View File
@@ -1,6 +1,9 @@
import asyncio
from aiohttp import web
from config.logger import setup_logging
from core.api.native_mqtt_management_handler import (
NativeMqttManagementHandler,
)
from core.api.ota_handler import OTAHandler
from core.api.vision_handler import VisionHandler
@@ -8,17 +11,32 @@ TAG = __name__
class SimpleHttpServer:
def __init__(self, config: dict):
def __init__(self, config: dict, management_owner=None):
self.config = config
self.management_owner = management_owner
self.logger = setup_logging()
self.ota_handler = OTAHandler(config)
self.vision_handler = VisionHandler(config)
self.native_mqtt_management_handler = (
NativeMqttManagementHandler(config, management_owner)
if management_owner is not None and self._native_mqtt_enabled()
else None
)
self._started_event = asyncio.Event()
self._stop_event = asyncio.Event()
self._runner = None
self._start_active = False
self._cleanup_lock = asyncio.Lock()
def _native_mqtt_enabled(self) -> bool:
mqtt_config = self.config.get("mqtt_server", {})
enabled_protocols = self.config.get("enabled_protocols", [])
return (
isinstance(mqtt_config, dict)
and mqtt_config.get("enabled") is True
and "mqtt" in enabled_protocols
)
async def wait_started(self, task: asyncio.Task, timeout: float = 10) -> None:
"""Wait until the HTTP listener is bound or surface startup failure."""
event_waiter = asyncio.create_task(self._started_event.wait())
@@ -71,7 +89,17 @@ class SimpleHttpServer:
port = int(server_config.get("http_port", 8003))
if port:
app = web.Application()
mqtt_config = self.config.get("mqtt_server", {})
client_max_size = max(
1024,
int(
mqtt_config.get(
"manager_max_request_size", 64 * 1024
)
or 64 * 1024
),
)
app = web.Application(client_max_size=client_max_size)
if not read_config_from_api:
# 如果没有开启智控台,只是单模块运行,就需要再添加简单OTA接口,用于下发websocket接口
@@ -105,6 +133,28 @@ class SimpleHttpServer:
),
]
)
if self.native_mqtt_management_handler is not None:
app.add_routes(
[
web.post(
"/api/devices/status",
self.native_mqtt_management_handler.handle_device_status,
),
web.post(
"/api/commands/{client_id}",
self.native_mqtt_management_handler.handle_command,
),
web.post(
"/api/call/request",
self.native_mqtt_management_handler.handle_call_request,
),
web.post(
"/api/call/accept",
self.native_mqtt_management_handler.handle_call_accept,
),
]
)
# 运行服务
runner = web.AppRunner(app)
self._runner = runner
@@ -23,8 +23,20 @@ class GoodbyeProcessor(MessageProcessor):
if isinstance(msg_json, dict) and msg_json.get("type") == "goodbye":
logger.info(f"收到goodbye: session_id={msg_json.get('session_id')}")
# 长连接传输仅结束逻辑会话,短连接传输关闭连接
# WebSocket 直接关闭连接;MQTT/UDP 仅结束音频会话,保持连接
if transport.keeps_connection_between_sessions:
end_call = getattr(
getattr(context, "server", None),
"end_native_mqtt_call",
None,
)
if callable(end_call):
await end_call(
context.device_id,
"设备结束通话",
notify_device=False,
expected_session_id=msg_json.get("session_id"),
)
end_conversation = getattr(context, "end_conversation", None)
if callable(end_conversation):
await end_conversation(msg_json.get("session_id"))
@@ -15,6 +15,7 @@ from core.processors.abort_processor import AbortProcessor
from core.processors.goodbye_processor import GoodbyeProcessor
from core.processors.text_processor import TextProcessor
from core.processors.ping_processor import PingProcessor
from core.processors.udp_timeout_processor import UdpTimeoutProcessor
from config.logger import setup_logging
logger = setup_logging()
@@ -40,6 +41,7 @@ class MessageRouter(MessageProcessor):
self.audio_receive_processor = AudioReceiveProcessor()
self.text_processor = TextProcessor()
self.ping_processor = PingProcessor()
self.udp_timeout_processor = UdpTimeoutProcessor()
# 按优先级排序的processor列表
self.processors: List[MessageProcessor] = [
@@ -48,6 +50,7 @@ class MessageRouter(MessageProcessor):
self.server_processor, # 服务器消息(管理端下发动作)
self.abort_processor, # 中断消息
self.goodbye_processor, # goodbye消息
self.udp_timeout_processor,
self.hello_processor, # hello消息
self.ping_processor, # 可选JSON ping/pong心跳
self.listen_processor, # listen消息
@@ -19,7 +19,7 @@ class TimeoutProcessor(MessageProcessor):
context: SessionContext,
transport: TransportInterface,
) -> bool:
"""End an expired logical conversation without closing a long connection."""
"""End an expired logical conversation without conflating it with MQTT."""
if not getattr(context, "conversation_active", False):
return False
@@ -29,7 +29,7 @@ class TimeoutProcessor(MessageProcessor):
try:
if transport.keeps_connection_between_sessions:
logger.info(f"会话超时,结束长连接逻辑会话: {context.session_id}")
logger.info(f"会话超时,结束MQTT逻辑会话: {context.session_id}")
from core.processors.audio_receive_processor import (
AudioReceiveProcessor,
)
@@ -0,0 +1,52 @@
import json
from typing import Any
from core.context.session_context import SessionContext
from core.pipeline.message_pipeline import MessageProcessor
from core.transport.transport_interface import TransportInterface
from config.logger import setup_logging
logger = setup_logging()
class UdpTimeoutProcessor(MessageProcessor):
async def process(
self,
context: SessionContext,
transport: TransportInterface,
message: Any,
) -> bool:
if isinstance(message, str):
try:
message = json.loads(message)
except json.JSONDecodeError:
return False
if not isinstance(message, dict) or message.get("type") != "udp_timeout":
return False
if not transport.keeps_connection_between_sessions:
return False
logger.info(
"收到udp_timeout: session_id={}", message.get("session_id")
)
end_call = getattr(
getattr(context, "server", None),
"end_native_mqtt_call",
None,
)
call_ended = False
if callable(end_call):
call_ended = await end_call(
context.device_id,
"设备UDP接收超时",
notify_device=True,
expected_session_id=message.get("session_id"),
)
if not call_ended:
end_conversation = getattr(context, "end_conversation", None)
if callable(end_conversation):
await end_conversation(message.get("session_id"))
await transport.end_session(
message.get("session_id") or context.session_id
)
return True
@@ -0,0 +1,649 @@
import asyncio
import json
import time
import uuid
from typing import Dict, Any, Optional, Callable
from config.logger import setup_logging
from core.utils.mqtt_auth import validate_mqtt_credentials
logger = setup_logging()
class MQTTConnection:
"""
MQTT连接处理类:管理单个MQTT客户端连接
处理MQTT协议消息和会话管理
"""
def __init__(self, socket, connection_id: int, mqtt_server, udp_handler=None, reader=None, writer=None):
self.socket = socket
self.reader = reader
self.writer = writer
self.connection_id = connection_id
self.mqtt_server = mqtt_server
self.udp_handler = udp_handler
# 连接信息
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.monotonic()
# 消息处理
self.message_callback = None
self.reply_topic = None
# UDP相关
self.udp_config = None
self.message_queue_size = int(
getattr(mqtt_server, "message_queue_size", 128) or 128
)
self.business_ready_timeout = float(
getattr(mqtt_server, "business_ready_timeout", 30) or 30
)
self.close_timeout = max(
0.1, float(getattr(mqtt_server, "close_timeout", 2) or 2)
)
self.goodbye_timeout = max(
0.1, float(getattr(mqtt_server, "goodbye_timeout", 1) or 1)
)
# 任务管理
self.keep_alive_task = None
self.business_task = None
self._closed = False
self._close_task = None
self._close_initiator = None
self._close_complete = asyncio.Event()
self.connect_processed_event = asyncio.Event()
self.connect_accepted = False
self.business_ready_event = asyncio.Event()
self._hello_business_ready_event = None
self._hello_business_session_id = None
self._logical_hello_received = False
self._startup_recovery_task = None
self._session_transition_lock = asyncio.Lock()
self._last_goodbye_session_id = None
self._goodbye_lock = asyncio.Lock()
# 创建MQTT协议处理器
from core.protocols.mqtt_protocol import MQTTProtocol
self.protocol = MQTTProtocol(
socket=socket,
reader=reader,
writer=writer,
max_payload_size=getattr(mqtt_server, "max_payload_size", 8192),
event_queue_size=self.message_queue_size,
close_timeout=self.close_timeout,
)
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)
self.protocol.on('protocolError', self._handle_error)
self.protocol.on('activity', self._handle_activity)
def _handle_activity(self):
"""更新最近活动时间(用于心跳保活)"""
self.last_activity = time.monotonic()
async def _handle_connect(self, connect_data: Dict[str, Any]):
"""处理CONNECT消息"""
try:
if self.connect_processed_event.is_set() or self.is_connected_flag:
logger.warning("同一TCP连接收到重复CONNECT,关闭连接")
await self.close()
return
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}")
try:
validate_mqtt_credentials(
self.client_id,
self.username,
self.password,
self.mqtt_server.signature_key,
)
except ValueError as e:
logger.warning(f"MQTT客户端认证失败: {self.client_id}, {e}")
await self.protocol.send_connack(4)
self.connect_processed_event.set()
await self.close()
return False
# 解析客户端ID获取设备信息
if not self._parse_client_id():
await self.protocol.send_connack(1) # 连接被拒绝
self.connect_processed_event.set()
await self.close()
return False
# 生成会话ID
self.session_id = str(uuid.uuid4())
# 设置回复主题
self.reply_topic = f"devices/p2p/{self.device_id.replace(':', '_')}"
async def complete_acceptance():
# Keep this inside the server's per-client takeover lock. A
# newer connection cannot reclaim this clientId between the
# success CONNACK and publishing the active owner.
self.is_connected_flag = True
try:
await self.protocol.send_connack(0)
except Exception:
self.is_connected_flag = False
raise
if self.keep_alive_interval > 0:
self.keep_alive_task = asyncio.create_task(
self._keep_alive_check()
)
self.connect_accepted = True
accepted = await self.mqtt_server.on_client_connected(
self, complete_acceptance
)
if not accepted:
if not self._closed:
await self.protocol.send_connack(3) # 服务端暂不可用
self.connect_processed_event.set()
await self.close()
return False
self.connect_processed_event.set()
return True
except Exception as e:
logger.error(f"处理CONNECT消息失败: {e}")
self.connect_processed_event.set()
await self.close()
return False
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('_', ':').replace('-', ':').lower()
)
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:
if publish_data.get('qos', 0) != 0:
logger.warning(
f"不支持的MQTT QoS级别: {publish_data.get('qos')}"
)
await self.close()
return
topic = publish_data['topic']
payload = publish_data['payload']
logger.debug(f"收到MQTT发布消息: topic={topic}, payload={payload}")
# 更新活动时间
self.last_activity = time.monotonic()
# 解析JSON消息
try:
message_data = json.loads(payload)
# 处理不同类型的消息
if message_data.get('type') == 'hello':
if message_data.get('version', 3) != 3:
logger.warning(
f"不支持的MQTT协议版本: {message_data.get('version')}"
)
await self.close()
return
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配置"""
async with self._session_transition_lock:
self._logical_hello_received = True
try:
handle_logical_hello = getattr(
self.mqtt_server, "handle_logical_hello", None
)
if callable(handle_logical_hello):
await handle_logical_hello(self)
# Match the gateway contract: do not advertise a usable audio
# channel until private config and runtime components are ready.
await asyncio.wait_for(
self.business_ready_event.wait(),
timeout=self.business_ready_timeout,
)
if self._closed or not self.is_connected_flag:
return
hello_reply = self._prepare_hello_reply(
message_data.get('audio_params', {}),
message_data.get('version', 3),
)
hello_ready = asyncio.Event()
self._hello_business_ready_event = hello_ready
self._hello_business_session_id = self.session_id
# Enqueue the logical-session boundary before the device can react
# to the reply with UDP audio.
if self.message_callback:
try:
self.message_callback(self.reply_topic, json.dumps(message_data))
except Exception as e:
logger.error(f"转发hello消息失败: {e}")
# Long-lived MQTT connections can change Agent configuration at
# each logical Hello. Do not expose the new UDP session until the
# business runtime has either refreshed or deliberately retained
# the previous healthy runtime.
await asyncio.wait_for(
hello_ready.wait(),
timeout=self.business_ready_timeout,
)
if self._closed or not self.is_connected_flag:
return
await self.send_message(self.reply_topic, json.dumps(hello_reply))
logger.info(f"MQTT Hello消息处理完成: {self.client_id}")
except asyncio.TimeoutError:
logger.error(
f"MQTT Hello等待业务运行时超时: {self.client_id}, "
f"timeout={self.business_ready_timeout}s"
)
await self.close()
except Exception as e:
logger.error(f"处理hello消息失败: {e}")
finally:
self._hello_business_ready_event = None
self._hello_business_session_id = None
def schedule_stale_session_recovery(self, delay: float = 1.0) -> None:
"""Return a reconnected device with a stale UDP session to Idle."""
if self._startup_recovery_task is not None:
return
self._startup_recovery_task = asyncio.create_task(
self._recover_stale_session(delay)
)
async def _recover_stale_session(self, delay: float) -> None:
try:
await asyncio.sleep(max(0.0, delay))
async with self._session_transition_lock:
if (
self._closed
or not self.is_connected_flag
or self._logical_hello_received
or not self.reply_topic
):
return
# No session id is intentional: firmware accepts this as a
# connection-level reset and discards an UDP session owned by
# a previous server process. Serialize it with Hello so this
# reset can never overtake a newly negotiated session.
await self.send_message(
self.reply_topic,
json.dumps({"type": "goodbye"}),
)
logger.info("已通知重连MQTT设备清理旧UDP会话: {}", self.client_id)
except asyncio.CancelledError:
pass
except Exception as e:
logger.warning("通知重连MQTT设备清理旧会话失败: {}", e)
def mark_business_session_ready(self, session_id: str = None) -> None:
"""Acknowledge readiness for the currently pending logical Hello."""
event = self._hello_business_ready_event
pending_session_id = self._hello_business_session_id
if event is None:
return
if session_id is not None and pending_session_id != session_id:
return
event.set()
async def send_hello_reply(self, audio_params: Dict[str, Any], version: int = 3):
"""发送hello回复(可在未收到设备hello时调用)"""
hello_reply = self._prepare_hello_reply(audio_params, version)
await self.send_message(self.reply_topic, json.dumps(hello_reply))
def _prepare_hello_reply(self, audio_params: Dict[str, Any], version: int = 3):
"""Create and install one UDP session without publishing it yet."""
import os
self.session_id = str(uuid.uuid4())
self._last_goodbye_session_id = None
udp_session_id = self.mqtt_server.bind_udp_session(
self, self.udp_handler
)
nonce = self._generate_udp_header(
0, 0, 0, connection_id=udp_session_id
)
self.udp_config = {
'key': os.urandom(16),
'encryption': 'aes-128-ctr',
'server': self.mqtt_server.public_endpoint,
'port': self.mqtt_server.udp_port,
'nonce': nonce,
'local_sequence': 0,
'remote_sequence': 0
}
if self.udp_handler:
self.udp_handler.configure_encryption(self.udp_config)
configured_audio_params = (
getattr(self.mqtt_server, 'config', {})
.get('xiaozhi', {})
.get('audio_params', {})
)
hello_reply = {
'type': 'hello',
'version': version,
'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': nonce.hex()
},
'audio_params': configured_audio_params or audio_params or {}
}
return hello_reply
def _generate_udp_header(
self, length: int, timestamp: int, sequence: int,
connection_id: int = None
) -> bytes:
header = bytearray(16)
header[0] = 1 # type
header[2:4] = length.to_bytes(2, 'big')
udp_connection_id = connection_id or self.connection_id
header[4:8] = udp_connection_id.to_bytes(4, 'big')
header[8:12] = timestamp.to_bytes(4, 'big')
header[12:16] = sequence.to_bytes(4, 'big')
return bytes(header)
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.monotonic()
if current_time - self.last_activity > self.keep_alive_interval / 1000 * 1.5:
logger.info(f"MQTT客户端心跳超时: {self.client_id}")
try:
await asyncio.wait_for(
self.notify_device_idle(),
timeout=self.goodbye_timeout,
)
except Exception as e:
logger.warning(
"MQTT心跳超时发送goodbye失败,继续关闭连接: {}", e
)
finally:
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:
raise RuntimeError("MQTT connection is closed")
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}")
raise
async def notify_device_idle(self, session_id: str = None) -> bool:
"""Send one session-scoped goodbye before the physical MQTT close."""
async with self._goodbye_lock:
target_session_id = session_id or self.session_id
if (
self._closed
or not self.is_connected_flag
or not self.udp_config
or not self.reply_topic
or not target_session_id
or self._last_goodbye_session_id == target_session_id
):
return False
await self.send_message(
self.reply_topic,
json.dumps(
{"type": "goodbye", "session_id": target_session_id}
),
)
self._last_goodbye_session_id = target_session_id
return True
def is_connected(self) -> bool:
"""检查连接状态"""
return self.is_connected_flag and not self._closed
async def close(self):
"""关闭连接"""
if not self._closed:
self._closed = True
self.is_connected_flag = False
self.connect_processed_event.set()
self.business_ready_event.set()
if self._hello_business_ready_event is not None:
self._hello_business_ready_event.set()
# Run cleanup in a dedicated task. Shielding it lets a cancelled
# first caller leave without falsely completing the close barrier;
# later callers can still await the same cleanup owner.
self._close_initiator = asyncio.current_task()
self._close_task = asyncio.create_task(self._close_impl())
close_task = self._close_task
if close_task is None:
return
current_task = asyncio.current_task()
dependency_tasks = {
close_task,
getattr(self.protocol, "_processing_task", None),
getattr(self.protocol, "_dispatch_task", None),
}
if self.business_task is not self._close_initiator:
dependency_tasks.add(self.business_task)
if self.keep_alive_task is not self._close_initiator:
dependency_tasks.add(self.keep_alive_task)
# The dedicated closer can be waiting for these tasks. Let them unwind
# instead of creating a reverse wait cycle.
if current_task in dependency_tasks:
return
await asyncio.shield(close_task)
async def _close_impl(self):
"""Own and complete physical cleanup independently of caller lifetime."""
try:
current_task = asyncio.current_task()
# The socket/protocol task can time out while ConnectionService is
# blocked in private config or component initialization. Cancel the
# owning server task so its finally block releases the SessionContext
# and any partially initialized runtime instead of leaking per retry.
if (
self.business_task
and self.business_task is not current_task
and self.business_task is not self._close_initiator
and not self.business_task.done()
):
self.business_task.cancel()
# The business owner's finally block calls back into server
# cleanup. Waiting for it here would create a close cycle:
# close -> business finally -> transport.close -> close.
# Give cancellation one loop turn, then let it unwind
# independently while the physical socket is released.
await asyncio.sleep(0)
if not self.business_task.done():
tracker = getattr(
self.mqtt_server, "track_draining_task", None
)
if callable(tracker):
tracker(self.business_task, "MQTT业务任务")
else:
self.business_task.add_done_callback(
lambda task: self._consume_background_task(
task, "MQTT业务任务"
)
)
# 取消心跳检查任务
if (
self.keep_alive_task
and self.keep_alive_task is not current_task
and self.keep_alive_task is not self._close_initiator
and not self.keep_alive_task.done()
):
self.keep_alive_task.cancel()
await self._wait_cancelled_task(
self.keep_alive_task, "MQTT心跳任务"
)
if (
self._startup_recovery_task
and self._startup_recovery_task is not current_task
and not self._startup_recovery_task.done()
):
self._startup_recovery_task.cancel()
await self._wait_cancelled_task(
self._startup_recovery_task, "MQTT会话恢复任务"
)
# 关闭协议处理器
try:
await asyncio.wait_for(
self.protocol.close(),
timeout=self.close_timeout,
)
except asyncio.TimeoutError:
logger.warning("关闭MQTT协议处理器超时,强制中止socket")
abort = getattr(self.protocol, "abort", None)
if callable(abort):
abort()
except (asyncio.CancelledError, Exception) as e:
logger.error(f"关闭MQTT协议处理器失败: {e}")
# Publish disconnect only after the physical protocol close barrier.
try:
await self.mqtt_server.on_client_disconnected(self)
except (asyncio.CancelledError, Exception) as e:
logger.error(f"通知服务器连接关闭失败: {e}")
logger.info(f"MQTT连接已关闭: {self.client_id}")
finally:
self._close_complete.set()
async def _wait_cancelled_task(self, task: asyncio.Task, label: str) -> None:
done, _ = await asyncio.wait({task}, timeout=self.close_timeout)
if task not in done:
logger.warning(
"{}取消后{}秒仍未退出,继续释放物理连接",
label,
self.close_timeout,
)
return
try:
task.result()
except asyncio.CancelledError:
pass
except Exception as exc:
logger.error(f"{label}退出失败: {exc}")
@staticmethod
def _consume_background_task(task: asyncio.Task, label: str) -> None:
try:
task.result()
except asyncio.CancelledError:
pass
except Exception as exc:
logger.error(f"{label}退出失败: {exc}")
@@ -0,0 +1,651 @@
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=None,
reader=None,
writer=None,
max_payload_size=8192,
event_queue_size=128,
close_timeout=2,
):
self.socket = socket
self.reader = reader
self.writer = writer
self.buffer = b''
self.event_handlers = {}
self.is_connected = False
self.keep_alive_interval = 0
self.last_activity = 0
self.max_payload_size = int(max_payload_size or 0)
self.close_timeout = max(0.1, float(close_timeout or 2))
self._closed = False
self._application_queue = asyncio.Queue(
maxsize=max(1, int(event_queue_size or 128))
)
# Application publishes stay ordered, while PINGREQ remains on the read
# loop so a slow Hello/runtime refresh cannot starve MQTT keepalive.
self._dispatch_task = asyncio.create_task(
self._dispatch_application_messages()
)
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 emit_async(self, event: str, *args, **kwargs):
"""Emit protocol events in packet order."""
handler = self.event_handlers.get(event)
if not handler:
return None
result = handler(*args, **kwargs)
if asyncio.iscoroutine(result):
return await result
return result
async def _process_messages(self):
"""处理消息的主循环"""
try:
while not self._closed:
# 从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:
if not self._closed:
if self._dispatch_task.done():
await self.emit_async('close')
else:
# Preserve parsed QoS0 publishes before a normal peer EOF.
# The Hello/runtime barrier is separately time-bounded.
await self._application_queue.put({'type': 'peer_close'})
async def _dispatch_application_messages(self):
"""Dispatch non-heartbeat packets sequentially outside the read loop."""
try:
while not self._closed:
message = await self._application_queue.get()
try:
await self._dispatch_application_message(message)
finally:
self._application_queue.task_done()
except asyncio.CancelledError:
pass
except Exception as e:
logger.error(f"MQTT应用消息处理循环出错: {e}")
self.emit('error', e)
async def _dispatch_application_message(self, message: Dict[str, Any]):
message_type = message.get('type')
if message_type == 'publish':
await self.emit_async('publish', message)
elif message_type == 'disconnect':
await self.emit_async('disconnect')
self.is_connected = False
elif message_type == 'peer_close':
await self.emit_async('close')
else:
raise ValueError(f"不支持的MQTT应用消息类型: {message_type}")
def _enqueue_application_message(self, message: Dict[str, Any]) -> None:
try:
self._application_queue.put_nowait(message)
except asyncio.QueueFull as exc:
raise ValueError("MQTT应用消息队列已满") from exc
async def _read_socket(self) -> bytes:
"""从socket读取数据"""
try:
if self.reader is not None:
return await self.reader.read(4096)
# 使用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)
client_packet_types = {
PacketType.CONNECT,
PacketType.PUBLISH,
PacketType.SUBSCRIBE,
PacketType.PINGREQ,
PacketType.DISCONNECT,
}
if packet_type not in client_packet_types:
raise ValueError(f"不支持的客户端MQTT消息类型: {packet_type}")
fixed_flags = first_byte & 0x0F
expected_flags = {
PacketType.CONNECT: 0,
PacketType.SUBSCRIBE: 2,
PacketType.PINGREQ: 0,
PacketType.DISCONNECT: 0,
}
if (
packet_type in expected_flags
and fixed_flags != expected_flags[packet_type]
):
raise ValueError(
f"MQTT packet type {packet_type} has invalid fixed-header flags "
f"0x{fixed_flags:x}"
)
if packet_type == PacketType.PUBLISH and ((first_byte >> 1) & 0x03) == 3:
raise ValueError("MQTT PUBLISH QoS 3 is invalid")
# 解析剩余长度
remaining_length, bytes_read = self._decode_remaining_length()
if remaining_length == -1:
return 0, {} # 长度解析失败,等待更多数据
max_payload_size = getattr(self, "max_payload_size", 0)
if max_payload_size > 0 and remaining_length > max_payload_size:
raise ValueError(
f"MQTT remaining length {remaining_length} exceeds limit {max_payload_size}"
)
# 计算完整消息长度
total_length = 1 + bytes_read + remaining_length
if len(self.buffer) < total_length:
return 0, {} # 消息不完整
if (
packet_type in (PacketType.PINGREQ, PacketType.DISCONNECT)
and remaining_length != 0
):
raise ValueError(
f"MQTT packet type {packet_type} requires remaining length 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:
if bytes_read + 1 >= len(self.buffer):
return -1, 0
digit = self.buffer[bytes_read + 1]
bytes_read += 1
value += (digit & 127) * multiplier
multiplier *= 128
if (digit & 128) == 0:
return value, bytes_read
if bytes_read == 4:
raise ValueError("MQTT remaining length字段超过4字节")
raise ValueError("MQTT remaining length字段无效")
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
def read_bytes(binary=False):
nonlocal pos
if pos + 2 > len(message_data):
raise ValueError("MQTT CONNECT字符串长度字段不完整")
value_length = int.from_bytes(message_data[pos:pos + 2], 'big')
pos += 2
if pos + value_length > len(message_data):
raise ValueError("MQTT CONNECT字符串内容不完整")
value = message_data[pos:pos + value_length]
pos += value_length
return value if binary else value.decode('utf-8')
protocol = read_bytes()
# 协议级别
if pos + 4 > len(message_data):
raise ValueError("MQTT CONNECT可变头部不完整")
protocol_level = message_data[pos]
pos += 1
if protocol != 'MQTT' or protocol_level != 4:
raise ValueError(
f"不支持的MQTT协议: {protocol}/{protocol_level}"
)
# 连接标志
connect_flags = message_data[pos]
if connect_flags & 0x01:
raise ValueError("MQTT CONNECT保留标志必须为0")
has_username = (connect_flags & 0x80) != 0
has_password = (connect_flags & 0x40) != 0
will_retain = (connect_flags & 0x20) != 0
will_qos = (connect_flags >> 3) & 0x03
has_will = (connect_flags & 0x04) != 0
clean_session = (connect_flags & 0x02) != 0
if has_password and not has_username:
raise ValueError("MQTT CONNECT密码标志要求用户名标志")
if will_qos == 3:
raise ValueError("MQTT CONNECT Will QoS 3无效")
if not has_will and (will_retain or will_qos):
raise ValueError("MQTT CONNECT未启用Will但设置了Will标志")
pos += 1
# 保持连接时间
keep_alive = int.from_bytes(message_data[pos:pos+2], 'big')
pos += 2
client_id = read_bytes()
if not client_id and not clean_session:
raise ValueError("MQTT CONNECT空clientId必须启用clean session")
if has_will:
read_bytes() # Will topic
read_bytes(binary=True) # Will payload
# 用户名(如果存在)
username = ''
if has_username:
username = read_bytes()
# 密码(如果存在)
password = ''
if has_password:
password = read_bytes(binary=True).decode('utf-8')
if pos != len(message_data):
raise ValueError("MQTT CONNECT包含未解析的尾部数据")
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
# 主题长度
if pos + 2 > len(message_data):
raise ValueError("MQTT PUBLISH缺少主题长度")
topic_length = int.from_bytes(message_data[pos:pos+2], 'big')
pos += 2
if topic_length == 0 or pos + topic_length > len(message_data):
raise ValueError("MQTT PUBLISH主题为空或不完整")
# 主题
topic = message_data[pos:pos+topic_length].decode('utf-8')
pos += topic_length
if "\x00" in topic or "+" in topic or "#" in topic:
raise ValueError("MQTT PUBLISH主题名称无效")
# 消息IDQoS > 0时存在)
packet_id = None
if qos > 0:
if pos + 2 > len(message_data):
raise ValueError("MQTT PUBLISH缺少packetId")
packet_id = int.from_bytes(message_data[pos:pos+2], 'big')
pos += 2
if packet_id == 0:
raise ValueError("MQTT PUBLISH packetId不能为0")
# 有效载荷
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
if pos + 2 > len(message_data):
raise ValueError("MQTT SUBSCRIBE缺少packetId")
packet_id = int.from_bytes(message_data[pos:pos+2], 'big')
pos += 2
if packet_id == 0:
raise ValueError("MQTT SUBSCRIBE packetId不能为0")
# 主题长度
if pos + 2 > len(message_data):
raise ValueError("MQTT SUBSCRIBE缺少主题长度")
topic_length = int.from_bytes(message_data[pos:pos+2], 'big')
pos += 2
if topic_length == 0 or pos + topic_length > len(message_data):
raise ValueError("MQTT SUBSCRIBE主题为空或不完整")
# 主题
topic = message_data[pos:pos+topic_length].decode('utf-8')
pos += topic_length
if "\x00" in topic:
raise ValueError("MQTT SUBSCRIBE主题过滤器无效")
# QoS
if pos >= len(message_data):
raise ValueError("MQTT SUBSCRIBE缺少请求QoS")
qos = message_data[pos]
pos += 1
if qos > 2:
raise ValueError("MQTT SUBSCRIBE请求QoS无效")
if pos != len(message_data):
raise ValueError("MQTT SUBSCRIBE当前仅支持单个主题过滤器")
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':
if self.is_connected:
raise ValueError("MQTT连接只能发送一次CONNECT")
self.keep_alive_interval = message.get('keepAlive', 0)
accepted = await self.emit_async('connect', message)
self.is_connected = accepted is not False
if self.is_connected:
self.emit('activity')
return
if not self.is_connected:
raise ValueError("MQTT客户端必须先发送CONNECT")
if message_type == 'publish':
self.emit('activity')
self._enqueue_application_message(message)
elif message_type == 'subscribe':
self.emit('activity')
await self.emit_async('subscribe', message)
elif message_type == 'pingreq':
self.emit('activity')
await self.send_pingresp()
elif message_type == 'disconnect':
self.emit('activity')
self._enqueue_application_message(message)
else:
raise ValueError(f"不支持的MQTT消息类型: {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:
if self.writer is not None:
self.writer.write(bytes(packet))
await self.writer.drain()
else:
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):
"""关闭协议处理器"""
self._closed = True
current_task = asyncio.current_task()
if (
hasattr(self, '_processing_task')
and self._processing_task is not current_task
and not self._processing_task.done()
):
self._processing_task.cancel()
try:
await self._processing_task
except asyncio.CancelledError:
pass
if (
hasattr(self, '_dispatch_task')
and self._dispatch_task is not current_task
and not self._dispatch_task.done()
):
self._dispatch_task.cancel()
try:
await self._dispatch_task
except asyncio.CancelledError:
pass
try:
if self.writer is not None:
self.writer.close()
try:
await asyncio.wait_for(
self.writer.wait_closed(),
timeout=self.close_timeout,
)
except asyncio.TimeoutError:
logger.warning(
"等待MQTT socket关闭超时,强制中止transport"
)
self.abort()
except Exception:
pass
elif self.socket:
self.socket.close()
except Exception as e:
logger.error(f"关闭socket失败: {e}")
def abort(self):
"""Force-close the underlying transport when graceful close stalls."""
if self.writer is not None:
transport = getattr(self.writer, "transport", None)
if transport is not None:
transport.abort()
return
if self.socket:
self.socket.close()
File diff suppressed because it is too large Load Diff
@@ -2,13 +2,14 @@ import asyncio
from typing import Dict, Any, List, Optional
from config.logger import setup_logging
from core.websocket_server_new import NewWebSocketServer
from core.servers.mqtt_server import MQTTServer
logger = setup_logging()
class MultiProtocolServer:
"""
协议服务器管理器
协议服务器管理器:统一管理WebSocket和MQTT服务器
提供统一的启动、停止和状态监控接口
"""
@@ -42,6 +43,11 @@ class MultiProtocolServer:
self.servers['websocket'] = NewWebSocketServer(self.config)
logger.info("WebSocket服务器已初始化")
# 初始化MQTT服务器
if 'mqtt' in enabled_protocols:
self.servers['mqtt'] = MQTTServer(self.config)
logger.info("MQTT服务器已初始化")
if not self.servers:
logger.warning("没有启用任何协议服务器")
@@ -333,6 +339,14 @@ class MultiProtocolServer:
# 检查服务器端口配置
server_configs = {
'server': ('port', 'host', 'ip'),
'mqtt_server': (
'port',
'udp_port',
'host',
'ip',
'public_endpoint',
'udp_bind_host',
),
}
for config_key, listener_keys in server_configs.items():
old_server_config = old_config.get(config_key, {})
@@ -443,7 +457,7 @@ class MultiProtocolServer:
def get_supported_protocols(self) -> List[str]:
"""获取支持的协议列表"""
return ['websocket']
return ['websocket', 'mqtt']
def is_protocol_enabled(self, protocol: str) -> bool:
"""检查协议是否启用"""
@@ -229,6 +229,23 @@ class ConnectionService:
await initialization_task
initialization_task = None
is_native_hello = (
transport.keeps_connection_between_sessions
and is_hello
)
if is_native_hello:
try:
# Every Native Hello is a logical-session boundary.
# Always query Agent config here; the refresh helper
# cheaply retains an unchanged healthy runtime.
component_manager = await self._refresh_private_runtime(
context,
component_manager,
bind_completed_event,
)
finally:
await transport.mark_session_ready(transport.session_id)
if getattr(context, "init_error", None):
# 配置错误时,允许hello/listen触发默认语音(节流)
msg_type = msg_json.get("type") if isinstance(msg_json, dict) else None
@@ -337,7 +354,9 @@ class ConnectionService:
except Exception as e:
logger.bind(tag=TAG).error(f"会话清理失败: {e}")
# Clean the manager currently owned by the context.
# A cancelled Native runtime refresh may have already installed a
# healthy replacement before this local variable was reassigned.
# Always clean the manager currently owned by the context.
active_component_manager = (
getattr(context, "component_manager", None) or component_manager
)
@@ -785,6 +804,201 @@ class ConnectionService:
return None
return None
async def _refresh_private_runtime(
self,
context: SessionContext,
component_manager,
bind_completed_event: asyncio.Event,
):
"""Refresh Agent config at a Native MQTT logical-session boundary."""
if not context.read_config_from_api:
return component_manager
try:
private_config = await get_private_config_from_api(
context.common_config,
context.device_id,
context.headers.get("client-id", context.device_id),
)
except DeviceNotFoundException:
context.need_bind = True
return component_manager
except DeviceBindException as exc:
context.need_bind = True
context.bind_code = getattr(exc, "bind_code", None)
return component_manager
except Exception as exc:
logger.bind(tag=TAG).warning(f"刷新差异化配置失败,继续使用当前配置: {exc}")
return component_manager
if not private_config:
return component_manager
private_config["delete_audio"] = bool(
context.common_config.get("delete_audio", True)
)
config_changed = private_config != context.private_config
if not config_changed and getattr(context, "init_error", None) is None:
context.need_bind = False
context.bind_code = None
bind_completed_event.set()
return component_manager
await self._finalize_conversation_session(context, context.session_id)
await context.cancel_conversation_tasks()
try:
staged_context = copy.copy(context)
staged_context.config = copy.deepcopy(context.common_config)
staged_context.private_config = copy.deepcopy(private_config)
self._reset_config_derived_state(staged_context)
staged_context.config.update(copy.deepcopy(private_config))
self._reset_config_derived_state(staged_context)
self._merge_private_modules(
staged_context, copy.deepcopy(private_config)
)
except Exception as exc:
logger.bind(tag=TAG).error(
f"物化刷新配置失败,继续使用旧运行时: {exc}"
)
return component_manager
try:
new_manager = ComponentRegistry.create_component_manager(
staged_context.config
)
except Exception as exc:
logger.bind(tag=TAG).error(f"创建刷新业务组件失败,保留旧运行时: {exc}")
return component_manager
runtime_fields = (
"config",
"private_config",
"component_manager",
"asr",
"tts",
"func_handler",
"voiceprint_provider",
"prompt",
"intent_type",
"load_function_plugin",
"inject_tool_call_fewshot",
"dialogue",
"max_output_size",
"chat_history_conf",
"cmd_exit",
"need_bind",
"bind_code",
"init_error",
"init_error_notified",
"_init_error_last_audio_ts",
)
missing = object()
previous = {
field: getattr(context, field, missing) for field in runtime_fields
}
previous_callbacks = list(context._cleanup_callbacks)
old_func_handler = context.func_handler
context.config = staged_context.config
context.private_config = staged_context.private_config
context.max_output_size = staged_context.max_output_size
context.chat_history_conf = staged_context.chat_history_conf
context.cmd_exit = staged_context.cmd_exit
context.component_manager = new_manager
context.asr = None
context.tts = None
context.func_handler = None
context.voiceprint_provider = None
context.prompt = None
context.intent_type = "nointent"
context.load_function_plugin = False
context.inject_tool_call_fewshot = None
context.dialogue = Dialogue()
context.init_error = None
try:
await self._initialize_components(context, new_manager)
except (asyncio.CancelledError, Exception) as exc:
was_cancelled = isinstance(exc, asyncio.CancelledError)
retry_callbacks = []
new_callbacks = [
callback
for callback in context._cleanup_callbacks
if callback not in previous_callbacks
]
for callback in reversed(new_callbacks):
try:
result = callback()
if asyncio.iscoroutine(result):
await result
except (asyncio.CancelledError, Exception) as cleanup_exc:
retry_callbacks.append(callback)
logger.bind(tag=TAG).warning(
f"回滚刷新工具处理器失败: {cleanup_exc}"
)
try:
await new_manager.cleanup_all()
except (asyncio.CancelledError, Exception) as cleanup_exc:
retry_callbacks.append(new_manager.cleanup_all)
logger.bind(tag=TAG).warning(
f"回滚刷新业务组件失败: {cleanup_exc}"
)
for field, value in previous.items():
if value is missing:
if hasattr(context, field):
delattr(context, field)
else:
setattr(context, field, value)
context._cleanup_callbacks = previous_callbacks + retry_callbacks
if was_cancelled:
logger.bind(tag=TAG).info(
"刷新业务组件被取消,已回滚并保留旧运行时"
)
raise
logger.bind(tag=TAG).error(
f"刷新业务组件失败,继续使用旧运行时: {exc}"
)
return component_manager
# Only retire the previous runtime after the replacement is healthy.
# Register the old manager before the first await so cancellation at
# any point leaves a cleanup owner for handle_connection.finally.
if (
component_manager
and component_manager.cleanup_all not in context._cleanup_callbacks
):
context.register_cleanup(component_manager.cleanup_all)
if old_func_handler and hasattr(old_func_handler, "cleanup"):
try:
await old_func_handler.cleanup()
except asyncio.CancelledError:
raise
except Exception as exc:
logger.bind(tag=TAG).warning(f"清理旧工具处理器失败: {exc}")
else:
context.unregister_cleanup(old_func_handler.cleanup)
if component_manager:
try:
await component_manager.cleanup_all()
except asyncio.CancelledError:
raise
except Exception as exc:
logger.bind(tag=TAG).warning(f"清理旧业务组件失败: {exc}")
else:
context.unregister_cleanup(component_manager.cleanup_all)
context.init_error = None
context.init_error_notified = False
context._init_error_last_audio_ts = 0.0
context.need_bind = False
context.bind_code = None
bind_completed_event.set()
logger.bind(tag=TAG).info(
"Native MQTT逻辑会话已原子刷新私有配置和业务组件"
)
return new_manager
@staticmethod
def _reset_config_derived_state(context: SessionContext) -> None:
"""Reset fields derived from config before applying a replacement."""
@@ -0,0 +1,726 @@
import asyncio
import time
from dataclasses import dataclass
from typing import Any, Dict, Optional
from config.logger import setup_logging
logger = setup_logging()
def normalize_device_id(device_id: Optional[str]) -> Optional[str]:
if not isinstance(device_id, str):
return None
normalized = device_id.strip().lower().replace("-", ":")
return normalized or None
@dataclass(frozen=True)
class PendingCall:
caller_mac: str
target_mac: str
caller_nickname: str
created_at: float
generation: int
class NativeMqttCallManager:
def __init__(
self,
connection_registry,
timeout_seconds: float = 60,
silence_frame: Optional[bytes] = None,
clock=time.monotonic,
):
self.connection_registry = connection_registry
self.timeout_seconds = max(1.0, float(timeout_seconds))
self.silence_frame = silence_frame
self.clock = clock
self.pending_calls: Dict[str, PendingCall] = {}
self.active_calls: Dict[str, str] = {}
self.call_session_ids: Dict[str, str] = {}
self.call_generations: Dict[str, int] = {}
self._next_generation = 1
self._generation_end_events: Dict[int, asyncio.Event] = {}
self._end_tasks: Dict[tuple[str, int], asyncio.Task] = {}
self._lock = asyncio.Lock()
async def request_call(
self,
caller_mac: str,
target_mac: str,
caller_nickname: str = "",
) -> Dict[str, Any]:
caller = normalize_device_id(caller_mac)
target = normalize_device_id(target_mac)
if not caller or not target or caller == target:
return {"status": "error", "message": "呼叫设备参数无效"}
caller_entry = self.connection_registry.resolve_device_now(caller)
target_entry = self.connection_registry.resolve_device_now(target)
if target_entry is None:
return {"status": "offline", "message": "对方设备不在线,请稍后重试"}
if caller_entry is None:
return {"status": "error", "message": "主叫设备不在线"}
async with self._lock:
if caller in self.active_calls or target in self.active_calls:
return {"status": "error", "message": "设备已在通话中"}
caller_pending = self._pending_owner_locked(caller)
target_pending = self._pending_owner_locked(target)
reverse = (
caller_pending == target
and target_pending == target
and self.pending_calls[target].target_mac == caller
)
if reverse:
pending = self.pending_calls.pop(target)
self.active_calls[caller] = target
self.active_calls[target] = caller
generation = pending.generation
status = "bridged"
elif caller_pending or target_pending:
return {"status": "error", "message": "设备已有等待中的通话"}
else:
generation = self._allocate_generation_locked()
self.pending_calls[caller] = PendingCall(
caller_mac=caller,
target_mac=target,
caller_nickname=caller_nickname or "",
created_at=self.clock(),
generation=generation,
)
status = "pending"
self.call_generations[caller] = generation
self.call_generations[target] = generation
self._capture_session(caller, caller_entry)
self._capture_session(target, target_entry)
self._set_call_state(caller_entry, True)
if status == "bridged":
self._set_call_state(target_entry, True)
try:
await self._stop_ai_session(
caller_entry, self.call_session_ids.get(caller)
)
if status == "bridged":
await self._stop_ai_session(
target_entry, self.call_session_ids.get(target)
)
except asyncio.CancelledError:
await self.end_call(
caller,
notify_device=False,
notify_peer=False,
expected_generation=generation,
)
raise
except Exception as error:
await self.end_call(
caller,
notify_device=True,
notify_peer=status == "bridged",
expected_generation=generation,
)
logger.warning(
"Native MQTT停止AI会话失败: caller={}, target={}, error={}",
caller,
target,
error,
)
return {"status": "error", "message": "停止AI会话失败"}
async with self._lock:
if status == "pending":
valid = self._pending_matches_locked(
caller, target, generation
)
else:
valid = self._active_matches_locked(
caller, target, generation
)
if not valid:
return {"status": "error", "message": "通话状态已变化"}
if status == "pending":
try:
sent = await self._send_while_generation_active(
target_entry.transport,
{
"type": "mcp",
"payload": {
"jsonrpc": "2.0",
"id": 9999,
"method": "tools/call",
"params": {
"name": "self.remote_wakeup",
"arguments": {
"reason": (
"[device_call]您收到来自"
f"{caller_nickname or '未知'}的来电,是否接听?"
),
"action": "listen",
},
},
},
},
generation,
)
if not sent:
return {"status": "error", "message": "通话状态已变化"}
except asyncio.CancelledError:
await self.end_call(
caller,
notify_device=False,
notify_peer=False,
expected_generation=generation,
)
raise
except Exception as error:
await self.end_call(
caller,
"发送来电通知失败",
notify_device=True,
notify_peer=False,
expected_generation=generation,
)
logger.warning(
"Native MQTT来电通知发送失败: caller={}, target={}, error={}",
caller,
target,
error,
)
return {"status": "error", "message": "发送来电通知失败"}
return {"status": status}
async def accept_call(self, callee_mac: str) -> Dict[str, Any]:
callee = normalize_device_id(callee_mac)
if not callee:
return {"status": "error", "message": "接听设备参数无效"}
callee_entry = self.connection_registry.resolve_device_now(callee)
if callee_entry is None:
return {"status": "offline", "message": "接听设备不在线"}
async with self._lock:
if callee in self.active_calls:
return {"status": "error", "message": "设备已在通话中"}
pending = next(
(
entry
for entry in self.pending_calls.values()
if entry.target_mac == callee
),
None,
)
if pending is None:
return {"status": "no_pending", "message": "没有等待中的通话"}
caller = pending.caller_mac
caller_entry = self.connection_registry.resolve_device_now(caller)
if caller_entry is None:
self._remove_call_locked(caller)
self._set_call_state(callee_entry, False)
return {
"status": "caller_gone",
"message": "主叫方已离开或通话已超时",
}
self.pending_calls.pop(caller, None)
self.active_calls[caller] = callee
self.active_calls[callee] = caller
generation = pending.generation
self.call_generations[caller] = generation
self.call_generations[callee] = generation
self._capture_session(caller, caller_entry)
self._capture_session(callee, callee_entry)
self._set_call_state(caller_entry, True)
self._set_call_state(callee_entry, True)
try:
await self._stop_ai_session(
callee_entry, self.call_session_ids.get(callee)
)
except asyncio.CancelledError:
await self.end_call(
callee,
notify_device=False,
notify_peer=False,
expected_generation=generation,
)
raise
except Exception as error:
await self.end_call(
callee,
notify_device=True,
notify_peer=True,
expected_generation=generation,
)
logger.warning(
"Native MQTT停止接听方AI会话失败: caller={}, callee={}, error={}",
caller,
callee,
error,
)
return {"status": "error", "message": "停止AI会话失败"}
async with self._lock:
valid = self._active_matches_locked(
caller, callee, generation
)
if not valid:
return {"status": "error", "message": "通话状态已变化"}
try:
sent = await self._send_while_generation_active(
caller_entry.transport,
{"type": "call_accepted", "from": callee},
generation,
)
if not sent:
return {"status": "error", "message": "通话状态已变化"}
except asyncio.CancelledError:
await self.end_call(
callee,
notify_device=False,
notify_peer=False,
expected_generation=generation,
)
raise
except Exception as error:
await self.end_call(
callee,
"发送接听确认失败",
notify_device=True,
notify_peer=True,
expected_generation=generation,
)
logger.warning(
"Native MQTT接听确认发送失败: caller={}, callee={}, error={}",
caller,
callee,
error,
)
return {"status": "error", "message": "发送接听确认失败"}
return {"status": "bridged", "peerMac": caller}
def route_audio(
self, source_device_id: str, payload: bytes, timestamp: int
) -> bool:
source = normalize_device_id(source_device_id)
if not source:
return False
peer = self.active_calls.get(source)
if peer:
target = self.connection_registry.resolve_device_now(peer)
if target is None:
self._schedule_end(source, "对方已离开")
return True
handler = getattr(target.transport, "_udp_handler", None)
try:
sent = (
handler is not None
and handler.send_audio_nowait(payload, 0)
)
except Exception:
sent = False
if not sent:
self._schedule_end(source, "对方音频通道不可用")
return True
if source in self.pending_calls:
source_entry = self.connection_registry.resolve_device_now(source)
handler = (
getattr(source_entry.transport, "_udp_handler", None)
if source_entry
else None
)
if handler is not None and self.silence_frame:
try:
handler.send_audio_nowait(self.silence_frame, 0)
except Exception:
self._schedule_end(source, "主叫音频通道不可用")
return True
return False
async def end_call(
self,
device_id: str,
reason: str = "",
notify_device: bool = False,
notify_peer: bool = True,
expected_session_id: Optional[str] = None,
expected_generation: Optional[int] = None,
) -> bool:
device = normalize_device_id(device_id)
if not device:
return False
observed_generation = self.call_generations.get(device)
generation = (
expected_generation
if expected_generation is not None
else observed_generation
)
if observed_generation is None or observed_generation != generation:
return False
if (
expected_session_id is not None
and self.call_session_ids.get(device) != expected_session_id
):
return False
end_event = self._generation_end_events.get(generation)
if end_event is not None:
end_event.set()
async with self._lock:
current_generation = self.call_generations.get(device)
if current_generation != generation:
return False
if expected_session_id is not None:
current_session_id = self.call_session_ids.get(device)
if current_session_id != expected_session_id:
return False
related = self._remove_call_locked(device)
if related is None:
return False
peer = related.get("peer")
device_session = related.get("device_session")
peer_session = related.get("peer_session")
device_entry = self.connection_registry.resolve_device_now(device)
peer_entry = (
self.connection_registry.resolve_device_now(peer) if peer else None
)
self._set_call_state(device_entry, False)
self._set_call_state(peer_entry, False)
notifications = []
if notify_device and device_entry is not None:
notifications.append(
self._notify_idle(device_entry, device_session, reason)
)
if notify_peer and peer_entry is not None:
notifications.append(
self._notify_idle(peer_entry, peer_session, reason)
)
if notifications:
results = await asyncio.gather(
*notifications, return_exceptions=True
)
for result in results:
if isinstance(result, Exception):
logger.warning(
"Native MQTT通话结束通知失败: error={}", result
)
return True
async def cleanup_timeouts(self) -> int:
expired = []
now = self.clock()
async with self._lock:
for caller, pending in list(self.pending_calls.items()):
if now - pending.created_at >= self.timeout_seconds:
expired.append((caller, pending.generation))
for caller, generation in expired:
await self.end_call(
caller,
"呼叫等待超时",
notify_device=True,
notify_peer=False,
expected_generation=generation,
)
return len(expired)
async def clear(self) -> None:
tasks = list(self._end_tasks.values())
self._end_tasks.clear()
for task in tasks:
task.cancel()
if tasks:
await asyncio.gather(*tasks, return_exceptions=True)
async with self._lock:
devices = set(self.pending_calls)
devices.update(self.active_calls)
devices.update(self.call_generations)
for device in devices:
await self.end_call(
device,
"服务停止",
notify_device=False,
notify_peer=False,
)
async with self._lock:
remaining = set(self.call_generations)
remaining.update(self.call_session_ids)
self.pending_calls.clear()
self.active_calls.clear()
self.call_generations.clear()
self.call_session_ids.clear()
end_events = list(self._generation_end_events.values())
self._generation_end_events.clear()
for end_event in end_events:
end_event.set()
for device in remaining:
self._set_call_state(
self.connection_registry.resolve_device_now(device),
False,
)
def contains(self, device_id: str) -> bool:
device = normalize_device_id(device_id)
return bool(
device
and (
self._pending_owner_now(device) is not None
or device in self.active_calls
)
)
@property
def count(self) -> int:
devices = set(self.call_generations)
devices.update(self.call_session_ids)
devices.update(self.active_calls)
devices.update(self.pending_calls)
devices.update(
pending.target_mac for pending in self.pending_calls.values()
)
generations = set(self.call_generations.values())
orphan_events = set(self._generation_end_events) - generations
return len(devices) + len(orphan_events)
@property
def background_task_count(self) -> int:
return sum(not task.done() for task in self._end_tasks.values())
async def handle_logical_hello(
self, device_id: str, session_id: Optional[str]
) -> bool:
device = normalize_device_id(device_id)
if not device:
return False
async with self._lock:
pending_owner = self._pending_owner_locked(device)
if (
pending_owner is not None
and pending_owner != device
and device not in self.active_calls
):
if session_id:
self.call_session_ids[device] = session_id
return False
generation = self.call_generations.get(device)
if generation is None:
return False
return await self.end_call(
device,
"设备重新进入AI会话",
notify_device=False,
notify_peer=True,
expected_session_id=session_id,
expected_generation=generation,
)
def _capture_session(self, device_id: str, entry) -> None:
session_id = getattr(entry.transport, "session_id", None)
if session_id:
self.call_session_ids[device_id] = session_id
def _drop_session(self, device_id: str) -> Optional[str]:
return self.call_session_ids.pop(device_id, None)
def _remove_call_locked(self, device: str) -> Optional[Dict[str, Any]]:
generation = self.call_generations.get(device)
end_event = self._generation_end_events.get(generation)
if end_event is not None:
end_event.set()
pending = self.pending_calls.pop(device, None)
if pending:
self._drop_generation(device)
self._drop_generation(pending.target_mac)
self._generation_end_events.pop(pending.generation, None)
return {
"peer": pending.target_mac,
"device_session": self._drop_session(device),
"peer_session": self._drop_session(pending.target_mac),
}
pending_owner = next(
(
caller
for caller, entry in self.pending_calls.items()
if entry.target_mac == device
),
None,
)
if pending_owner:
pending = self.pending_calls.pop(pending_owner)
self._drop_generation(device)
self._drop_generation(pending_owner)
self._generation_end_events.pop(pending.generation, None)
return {
"peer": pending_owner,
"device_session": self._drop_session(device),
"peer_session": self._drop_session(pending_owner),
}
peer = self.active_calls.pop(device, None)
if peer:
self.active_calls.pop(peer, None)
self._drop_generation(device)
self._drop_generation(peer)
if generation is not None:
self._generation_end_events.pop(generation, None)
return {
"peer": peer,
"device_session": self._drop_session(device),
"peer_session": self._drop_session(peer),
}
return None
def _schedule_end(self, device_id: str, reason: str) -> None:
device = normalize_device_id(device_id)
generation = self.call_generations.get(device) if device else None
if not device or generation is None:
return
key = (device, generation)
existing = self._end_tasks.get(key)
if existing is not None and not existing.done():
return
task = asyncio.create_task(
self.end_call(
device,
reason,
notify_device=True,
notify_peer=True,
expected_generation=generation,
)
)
self._end_tasks[key] = task
task.add_done_callback(
lambda completed, task_key=key: self._discard_end_task(
task_key, completed
)
)
def _discard_end_task(
self, key: tuple[str, int], task: asyncio.Task
) -> None:
if self._end_tasks.get(key) is task:
self._end_tasks.pop(key, None)
if not task.cancelled():
task.exception()
def _allocate_generation_locked(self) -> int:
generation = self._next_generation
self._next_generation += 1
self._generation_end_events[generation] = asyncio.Event()
return generation
async def _send_while_generation_active(
self, transport, message: Dict[str, Any], generation: int
) -> bool:
end_event = self._generation_end_events.get(generation)
if end_event is None or end_event.is_set():
return False
send_task = asyncio.create_task(transport.send_json(message))
end_task = asyncio.create_task(end_event.wait())
try:
done, _ = await asyncio.wait(
{send_task, end_task},
return_when=asyncio.FIRST_COMPLETED,
)
if end_task in done:
send_task.cancel()
await asyncio.gather(send_task, return_exceptions=True)
return False
await send_task
return not end_event.is_set()
finally:
if not send_task.done():
send_task.cancel()
if not end_task.done():
end_task.cancel()
await asyncio.gather(
send_task, end_task, return_exceptions=True
)
def _pending_owner_locked(self, device: str) -> Optional[str]:
if device in self.pending_calls:
return device
return next(
(
caller
for caller, pending in self.pending_calls.items()
if pending.target_mac == device
),
None,
)
def _pending_owner_now(self, device: str) -> Optional[str]:
return self._pending_owner_locked(device)
def _pending_matches_locked(
self, caller: str, target: str, generation: int
) -> bool:
pending = self.pending_calls.get(caller)
return bool(
pending
and pending.target_mac == target
and pending.generation == generation
and self.call_generations.get(caller) == generation
and self.call_generations.get(target) == generation
)
def _active_matches_locked(
self, caller: str, target: str, generation: int
) -> bool:
return (
self.active_calls.get(caller) == target
and self.active_calls.get(target) == caller
and self.call_generations.get(caller) == generation
and self.call_generations.get(target) == generation
)
def _drop_generation(self, device_id: str) -> Optional[int]:
return self.call_generations.pop(device_id, None)
@staticmethod
async def _stop_ai_session(entry, session_id: Optional[str]) -> None:
if entry is None:
return
end_conversation = getattr(entry.context, "end_conversation", None)
if callable(end_conversation):
await end_conversation(session_id)
return
cancel_tasks = getattr(entry.context, "cancel_conversation_tasks", None)
if callable(cancel_tasks):
await cancel_tasks()
@staticmethod
def _set_call_state(entry, active: bool) -> None:
if entry is None:
return
entry.context.calling = active
if not active:
entry.context.incoming_call = None
@staticmethod
async def _notify_idle(entry, session_id: Optional[str], reason: str) -> None:
raw_connection = getattr(entry.transport, "raw_connection", None)
if raw_connection is None:
return
try:
await raw_connection.notify_device_idle(session_id)
finally:
end_conversation = getattr(entry.context, "end_conversation", None)
if callable(end_conversation):
await end_conversation(session_id)
if reason:
logger.info(
"Native MQTT通话结束: device_id={}, reason={}",
entry.context.device_id,
reason,
)
@@ -0,0 +1,137 @@
import asyncio
from dataclasses import dataclass
from typing import Any, Dict, Iterable, Optional
@dataclass(frozen=True)
class NativeMqttConnection:
client_id: str
device_id: Optional[str]
connection_id: int
context: Any
transport: Any
@property
def is_alive(self) -> bool:
return bool(getattr(self.transport, "is_connected", False))
class NativeMqttConnectionRegistry:
def __init__(self):
self._connections: Dict[str, NativeMqttConnection] = {}
self._devices: Dict[str, NativeMqttConnection] = {}
self._lock = asyncio.Lock()
async def register(self, context: Any, transport: Any) -> bool:
client_id = getattr(transport, "client_id", None)
raw_connection = getattr(transport, "raw_connection", None)
connection_id = getattr(raw_connection, "connection_id", None)
if not client_id or connection_id is None:
return False
entry = NativeMqttConnection(
client_id=client_id,
device_id=self._normalize_device_id(
getattr(context, "device_id", None)
),
connection_id=connection_id,
context=context,
transport=transport,
)
async with self._lock:
previous_client = self._connections.get(client_id)
previous_device = (
self._devices.get(entry.device_id)
if entry.device_id
else None
)
for previous in (previous_client, previous_device):
if previous is None or previous is entry:
continue
if self._connections.get(previous.client_id) is previous:
self._connections.pop(previous.client_id, None)
if (
previous.device_id
and self._devices.get(previous.device_id) is previous
):
self._devices.pop(previous.device_id, None)
self._connections[client_id] = entry
if entry.device_id:
self._devices[entry.device_id] = entry
return True
async def unregister(self, context: Any, transport: Any) -> bool:
client_id = getattr(transport, "client_id", None)
if not client_id:
return False
async with self._lock:
entry = self._connections.get(client_id)
if (
entry is None
or entry.context is not context
or entry.transport is not transport
):
return False
self._connections.pop(client_id, None)
if (
entry.device_id
and self._devices.get(entry.device_id) is entry
):
self._devices.pop(entry.device_id, None)
return True
async def resolve(self, client_id: str) -> Optional[NativeMqttConnection]:
async with self._lock:
entry = self._connections.get(client_id)
if entry is None or not entry.is_alive:
return None
return entry
async def status(self, client_ids: Iterable[str]) -> Dict[str, Dict[str, Any]]:
async with self._lock:
result = {}
for client_id in client_ids:
entry = self._connections.get(client_id)
exists = entry is not None
result[client_id] = {
"isAlive": bool(entry and entry.is_alive),
"exists": exists,
"backend": "native",
}
return result
async def resolve_device(
self, device_id: str
) -> Optional[NativeMqttConnection]:
async with self._lock:
return self.resolve_device_now(device_id)
def resolve_device_now(
self, device_id: str
) -> Optional[NativeMqttConnection]:
normalized = self._normalize_device_id(device_id)
entry = self._devices.get(normalized) if normalized else None
if entry is None or not entry.is_alive:
return None
return entry
async def clear(self) -> None:
async with self._lock:
self._connections.clear()
self._devices.clear()
async def size(self) -> int:
async with self._lock:
return len(self._connections)
@property
def count(self) -> int:
return len(self._connections)
@staticmethod
def _normalize_device_id(device_id: Optional[str]) -> Optional[str]:
if not isinstance(device_id, str):
return None
normalized = device_id.strip().lower().replace("-", ":")
return normalized or None
@@ -0,0 +1,525 @@
import asyncio
import json
import time
from collections import deque
from typing import Any, AsyncGenerator, Dict, Optional
from .transport_interface import TransportInterface
from config.logger import setup_logging
from cryptography.hazmat.primitives.ciphers import Cipher, algorithms, modes
from cryptography.hazmat.backends import default_backend
logger = setup_logging()
class MQTTTransport(TransportInterface):
"""
MQTT传输层实现:直接处理MQTT协议消息
支持JSON消息和二进制音频数据传输
"""
def __init__(self, mqtt_connection, udp_handler=None):
"""
初始化MQTT传输层
Args:
mqtt_connection: MQTT连接对象,包含协议处理器
udp_handler: UDP处理器,用于音频数据传输
"""
self._mqtt_connection = mqtt_connection
self._udp_handler = udp_handler
queue_size = int(getattr(mqtt_connection, "message_queue_size", 128) or 128)
self._audio_queue = deque(maxlen=max(1, queue_size))
self._control_queue = deque(maxlen=max(32, min(queue_size, 128)))
self._urgent_queue = deque(maxlen=32)
self._arrival_sequence = 0
self._message_event = asyncio.Event()
self._closed = False
# 设置MQTT连接的消息回调
self._setup_message_handlers()
def _setup_message_handlers(self):
"""设置消息处理回调"""
# 设置MQTT消息接收回调
self._mqtt_connection.set_message_callback(self._on_mqtt_message)
# 设置UDP消息接收回调(如果有UDP处理器)
if self._udp_handler:
self._udp_handler.set_message_callback(self._on_udp_message)
def _on_mqtt_message(self, topic: str, payload: str):
"""处理接收到的MQTT消息"""
try:
# 解析JSON消息
message_data = json.loads(payload)
message_data['_transport_type'] = 'mqtt'
message_data['_topic'] = topic
# Hello is a complete logical-session barrier. Discard all queued
# work from the previous session before admitting the new Hello.
if message_data.get("type") == "hello":
self._audio_queue.clear()
self._control_queue.clear()
self._urgent_queue.clear()
self._enqueue_message(message_data)
except json.JSONDecodeError as e:
logger.error(f"MQTT消息JSON解析失败: {e}, payload: {payload}")
except Exception as e:
logger.error(f"处理MQTT消息失败: {e}")
def _on_udp_message(self, audio_data: bytes, timestamp: int):
"""处理接收到的UDP音频消息"""
try:
# 构造音频消息格式
message_data = {
'type': 'audio',
'data': audio_data,
'timestamp': timestamp,
'_transport_type': 'udp'
}
self._enqueue_message(message_data)
except Exception as e:
logger.error(f"处理UDP音频消息失败: {e}")
def _enqueue_message(self, message_data: Dict[str, Any]) -> None:
self._arrival_sequence += 1
queued_message = (self._arrival_sequence, message_data)
if message_data.get("type") == "abort":
self._urgent_queue.append(queued_message)
self._message_event.set()
return
if message_data.get("type") == "audio":
if len(self._audio_queue) >= self._audio_queue.maxlen:
self._audio_queue.popleft()
logger.warning("MQTT audio receive queue is full; evicted oldest frame")
self._audio_queue.append(queued_message)
else:
if len(self._control_queue) >= self._control_queue.maxlen:
boundary_types = {"hello", "goodbye"}
evict_index = next(
(
index
for index, (_, queued) in enumerate(self._control_queue)
if queued.get("type") not in boundary_types
),
None,
)
if evict_index is None:
if message_data.get("type") not in boundary_types:
logger.warning(
"MQTT control queue contains only session boundaries; "
"dropping non-boundary frame"
)
return
self._control_queue.popleft()
else:
del self._control_queue[evict_index]
logger.warning(
"MQTT control queue is full; evicted non-boundary control frame"
)
self._control_queue.append(queued_message)
self._message_event.set()
async def _next_message(self):
while not self._closed:
# Hello establishes the logical-session boundary. An Abort sent
# immediately after the Hello reply must not overtake it and be
# compared against the previous session.
if (
self._control_queue
and self._control_queue[0][1].get("type") == "hello"
):
return self._control_queue.popleft()[1]
if self._urgent_queue:
return self._urgent_queue.popleft()[1]
if self._control_queue and self._audio_queue:
queue = (
self._control_queue
if self._control_queue[0][0] < self._audio_queue[0][0]
else self._audio_queue
)
return queue.popleft()[1]
if self._control_queue:
return self._control_queue.popleft()[1]
if self._audio_queue:
return self._audio_queue.popleft()[1]
self._message_event.clear()
if self._urgent_queue or self._control_queue or self._audio_queue:
continue
await asyncio.wait_for(self._message_event.wait(), timeout=1.0)
return None
async def send(self, data: Any) -> None:
"""发送消息"""
if self._closed:
raise RuntimeError("Transport is closed")
try:
if isinstance(data, dict):
# 根据消息类型选择传输方式
if data.get('type') == 'audio' and self._udp_handler:
# 音频数据通过UDP发送
audio_data = data.get('data')
timestamp = data.get('timestamp', 0)
await self._udp_handler.send_audio(audio_data, timestamp)
else:
# JSON消息通过MQTT发送
topic = data.get('_topic', self._mqtt_connection.reply_topic)
payload = json.dumps(data)
await self._mqtt_connection.send_message(topic, payload)
elif isinstance(data, str):
# 字符串消息通过MQTT发送
await self._mqtt_connection.send_message(
self._mqtt_connection.reply_topic,
data
)
elif isinstance(data, bytes):
# 二进制数据通过UDP发送(如果有UDP处理器)
if self._udp_handler:
await self._udp_handler.send_audio(data, 0)
else:
logger.warning("尝试发送二进制数据但没有UDP处理器")
else:
# 其他类型转换为字符串通过MQTT发送
await self._mqtt_connection.send_message(
self._mqtt_connection.reply_topic,
str(data)
)
except Exception as e:
logger.error(f"MQTT传输发送消息失败: {e}")
raise
async def send_json(self, message: Any) -> None:
if isinstance(message, str):
await self.send(message)
return
await self.send(dict(message))
async def send_audio(self, audio: bytes, timestamp: int = 0) -> None:
if not self._udp_handler:
raise RuntimeError("UDP audio channel is not available")
await self._udp_handler.send_audio(audio, timestamp)
@property
def requires_audio_tail_grace(self) -> bool:
return True
async def prepare_audio_channel(self, audio_params=None, version: int = 3) -> None:
if not self._udp_handler:
return
if getattr(self._mqtt_connection, "udp_config", None) is None:
await self._mqtt_connection.send_hello_reply(audio_params or {}, version)
async def wait_audio_ready(self, timeout: float = 0) -> bool:
if not self._udp_handler:
return False
deadline = time.monotonic() + max(timeout, 0)
while getattr(self._udp_handler, "remote_address", None) is None:
if time.monotonic() >= deadline:
return False
await asyncio.sleep(min(0.05, max(deadline - time.monotonic(), 0)))
return True
async def mark_business_ready(self) -> None:
self._mqtt_connection.business_ready_event.set()
schedule_recovery = getattr(
self._mqtt_connection, "schedule_stale_session_recovery", None
)
if callable(schedule_recovery):
schedule_recovery()
async def mark_session_ready(self, session_id: str = None) -> None:
self._mqtt_connection.mark_business_session_ready(session_id)
async def end_session(self, session_id: str) -> None:
await self._mqtt_connection.notify_device_idle(session_id)
async def receive(self) -> AsyncGenerator[Any, None]:
"""异步消息流"""
while not self._closed:
try:
# 等待消息,设置超时避免无限等待
message = await self._next_message()
if message is None:
break
yield message
except asyncio.TimeoutError:
# 超时继续循环,检查连接状态
if not self.is_connected:
break
continue
except Exception as e:
logger.error(f"MQTT传输接收消息失败: {e}")
break
async def close(self) -> None:
"""关闭传输层"""
if self._closed:
return
self._closed = True
try:
# 关闭MQTT连接
if self._mqtt_connection:
await self._mqtt_connection.close()
# 关闭UDP处理器
if self._udp_handler:
await self._udp_handler.close()
except Exception as e:
logger.error(f"关闭MQTT传输层失败: {e}")
raise RuntimeError("MQTT transport close failed")
@property
def is_connected(self) -> bool:
"""检查连接状态"""
if self._closed:
return False
try:
# 检查MQTT连接状态
mqtt_connected = (
self._mqtt_connection and
self._mqtt_connection.is_connected()
)
return mqtt_connected
except Exception as e:
logger.error(f"检查MQTT连接状态失败: {e}")
return False
@property
def device_id(self) -> Optional[str]:
"""获取设备ID"""
return getattr(self._mqtt_connection, 'device_id', None)
@property
def client_id(self) -> Optional[str]:
"""获取客户端ID"""
return getattr(self._mqtt_connection, 'client_id', None)
@property
def username(self) -> Optional[str]:
"""获取MQTT用户名"""
return getattr(self._mqtt_connection, 'username', None)
@property
def password(self) -> Optional[str]:
"""获取MQTT密码"""
return getattr(self._mqtt_connection, 'password', None)
@property
def session_id(self) -> Optional[str]:
"""获取会话ID"""
return getattr(self._mqtt_connection, 'session_id', None)
@property
def transport_type(self) -> str:
return "mqtt"
@property
def has_datagram_audio(self) -> bool:
return self._udp_handler is not None
@property
def keeps_connection_between_sessions(self) -> bool:
return True
@property
def is_protocol_authenticated(self) -> bool:
return bool(getattr(self._mqtt_connection, "connect_accepted", False))
@property
def raw_connection(self):
return self._mqtt_connection
class UDPAudioHandler:
"""
UDP音频处理器:处理加密音频数据传输
"""
def __init__(
self,
connection_id: int,
udp_server,
encryption_config: Dict[str, Any],
allowed_remote_ip: Optional[str] = None,
):
self.connection_id = connection_id
self.udp_server = udp_server
self.encryption_config = encryption_config
self.allowed_remote_ip = allowed_remote_ip
self.remote_address = None
self.message_callback = None
self.audio_interceptor = None
self._closed = False
self.local_sequence = 0
self.remote_sequence = 0
self.audio_sequence_start = None
self.audio_start_time = None
self.frame_ms = 60
def set_message_callback(self, callback):
"""设置消息接收回调"""
self.message_callback = callback
def set_audio_interceptor(self, callback):
self.audio_interceptor = callback
def configure_encryption(self, udp_config: Dict[str, Any]):
"""设置UDP加密参数"""
if not udp_config:
return
self.encryption_config = udp_config
# A new Hello creates a new UDP session. Allow the first valid packet
# from the MQTT peer to establish the new source tuple.
self.remote_address = None
self.reset_sequence()
def reset_sequence(self):
"""重置UDP序列号(本地/远端)"""
self.reset_local_sequence()
self.reset_remote_sequence()
def reset_local_sequence(self):
"""重置本地发送序列号"""
self.local_sequence = 0
def reset_remote_sequence(self):
"""重置远端接收序列号"""
self.remote_sequence = 0
self.audio_sequence_start = None
self.audio_start_time = None
async def send_audio(self, audio_data: bytes, timestamp: int):
"""发送音频数据"""
if self._closed:
raise RuntimeError("UDP audio handler is closed")
if not self.remote_address:
raise RuntimeError("UDP remote address is not ready")
next_sequence = self.local_sequence + 1
await self.udp_server.send_encrypted_audio(
self.connection_id,
audio_data,
timestamp,
next_sequence,
self.remote_address,
self.encryption_config
)
self.local_sequence = next_sequence
def send_audio_nowait(self, audio_data: bytes, timestamp: int) -> bool:
if self._closed or not self.remote_address:
return False
next_sequence = self.local_sequence + 1
self.udp_server.send_encrypted_audio_nowait(
self.connection_id,
audio_data,
timestamp,
next_sequence,
self.remote_address,
self.encryption_config,
)
self.local_sequence = next_sequence
return True
def on_udp_message(self, header: bytes, encrypted_payload: bytes, payload_length: int,
timestamp: int, sequence: int, remote_addr):
"""处理接收到的UDP消息"""
if self._closed:
return
if self.allowed_remote_ip and remote_addr[0] != self.allowed_remote_ip:
logger.warning(
"Rejected UDP packet from non-MQTT peer: {}, expected IP: {}",
remote_addr,
self.allowed_remote_ip,
)
return
if self.remote_address is not None and remote_addr != self.remote_address:
logger.warning(
"Rejected UDP source rebind during active session: {}, bound: {}",
remote_addr,
self.remote_address,
)
return
if self.audio_sequence_start is not None and sequence <= self.remote_sequence:
return
if sequence != self.remote_sequence + 1:
logger.warning(
"Received UDP packet with wrong sequence: {}, expected: {}",
sequence,
self.remote_sequence + 1
)
if len(encrypted_payload) != payload_length:
logger.warning(
"UDP payload length mismatch: {} != {}",
len(encrypted_payload),
payload_length,
)
return
try:
key = self.encryption_config.get('key') if self.encryption_config else None
if key:
cipher = Cipher(algorithms.AES(key), modes.CTR(header), backend=default_backend())
decryptor = cipher.decryptor()
payload = decryptor.update(encrypted_payload) + decryptor.finalize()
else:
payload = encrypted_payload
except Exception as e:
logger.error("UDP decrypt failed: {}", e)
return
if self.remote_address is None:
self.remote_address = remote_addr
self.audio_start_time = time.time()
self.audio_sequence_start = sequence
self.remote_sequence = sequence - 1
if self.audio_sequence_start is None:
self.audio_sequence_start = sequence
self.remote_sequence = sequence
normalized_timestamp = timestamp
if timestamp == 0 and self.audio_sequence_start is not None:
normalized_timestamp = (
(sequence - self.audio_sequence_start) * self.frame_ms
) % (2 ** 32)
if self.audio_interceptor and self.audio_interceptor(
payload, normalized_timestamp
):
return
if self.message_callback:
self.message_callback(payload, normalized_timestamp)
async def close(self):
"""关闭UDP处理器"""
self._closed = True
self.message_callback = None
self.audio_interceptor = None
self.remote_address = None
+100
View File
@@ -0,0 +1,100 @@
import base64
import hashlib
import hmac
import json
import re
from config.logger import setup_logging
logger = setup_logging()
TAG = __name__
_MAC_ADDRESS_PATTERN = re.compile(r"^(?:[0-9A-Fa-f]{2}[:-]){5}[0-9A-Fa-f]{2}$")
_ENDPOINT_SCHEME_PATTERN = re.compile(
r"^(?:mqtt|tcp|ssl|ws|wss|http|https)://", re.IGNORECASE
)
def normalize_signature_key(secret_key: str) -> str:
"""Treat values emitted by manager-api for an unset parameter as empty."""
if secret_key is None:
return ""
value = str(secret_key).strip()
if not value or value.lower() == "null" or "" in value:
return ""
return value
def generate_password_signature(content: str, secret_key: str) -> str:
"""生成MQTT密码签名(HMAC-SHA256 + Base64"""
try:
hmac_obj = hmac.new(
secret_key.encode("utf-8"), content.encode("utf-8"), hashlib.sha256
)
signature = hmac_obj.digest()
return base64.b64encode(signature).decode("utf-8")
except Exception as e:
logger.bind(tag=TAG).error(f"生成MQTT密码签名失败: {e}")
return ""
def parse_mqtt_endpoint(endpoint: str, default_port: int = None) -> tuple[str, int]:
"""Parse the host[:port] syntax supported by the current ESP firmware."""
if endpoint is None:
return "", default_port
value = str(endpoint).strip()
if not value or value.lower() == "null" or "" in value:
return "", default_port
value = _ENDPOINT_SCHEME_PATTERN.sub("", value, count=1).split("/", 1)[0]
if not value or value.startswith("[") or value.count(":") > 1:
raise ValueError("MQTT endpoint格式无效")
host = value
port = default_port
if ":" in value:
host, port_text = value.rsplit(":", 1)
if not port_text.isdigit():
raise ValueError("MQTT endpoint端口无效")
port = int(port_text)
if not host or any(char.isspace() for char in host):
raise ValueError("MQTT endpoint主机无效")
if port is not None and not 1 <= int(port) <= 65535:
raise ValueError("MQTT endpoint端口超出范围")
return host, int(port) if port is not None else None
def validate_mqtt_credentials(
client_id: str, username: str, password: str, secret_key: str
) -> None:
"""Validate the gateway-compatible MQTT client id and HMAC credentials."""
if not client_id or not isinstance(client_id, str):
raise ValueError("clientId必须是非空字符串")
parts = client_id.split("@@@")
if len(parts) not in (2, 3) or not parts[0] or not parts[1]:
raise ValueError("clientId格式错误")
mac_address = parts[1].replace("_", ":")
if not _MAC_ADDRESS_PATTERN.fullmatch(mac_address):
raise ValueError("clientId中的MAC地址无效")
normalized_key = normalize_signature_key(secret_key)
if len(parts) == 2:
if normalized_key:
raise ValueError("启用签名时clientId必须包含UUID")
return
if not username or not isinstance(username, str):
raise ValueError("username必须是非空字符串")
try:
user_data = json.loads(base64.b64decode(username, validate=True).decode("utf-8"))
if not isinstance(user_data, dict):
raise ValueError
except Exception as exc:
raise ValueError("username不是有效的base64编码JSON") from exc
if normalized_key:
expected = generate_password_signature(client_id + "|" + username, normalized_key)
if not password or not hmac.compare_digest(password, expected):
raise ValueError("密码签名验证失败")
@@ -19,7 +19,7 @@ class XiaozhiServerFacade:
提供统一的服务器管理接口,屏蔽内部协议复杂性
功能:
- 协议管理
- 协议管理WebSocket、MQTT
- 本地 ASR 模型预加载
- 优雅启动和停止
"""
@@ -49,15 +49,24 @@ class XiaozhiServerFacade:
protocols = self.config.get("protocols", {})
if not isinstance(protocols, dict):
protocols = {}
mqtt_config = self.config.get("mqtt_server", {})
if not isinstance(mqtt_config, dict):
mqtt_config = {}
requested = protocols.get("enabled_protocols")
requested = requested if isinstance(requested, list) else []
websocket_enabled = protocols.get("websocket_enabled")
if websocket_enabled is None:
websocket_enabled = not protocols or "websocket" in requested
mqtt_requested = protocols.get("mqtt_enabled") is True or "mqtt" in requested
mqtt_enabled = mqtt_config.get("enabled") is True and mqtt_requested
enabled_protocols = []
if websocket_enabled:
enabled_protocols.append("websocket")
if mqtt_enabled:
enabled_protocols.append("mqtt")
self.config["enabled_protocols"] = enabled_protocols
logger.info(f"启用的协议: {enabled_protocols}")
@@ -398,7 +407,7 @@ class XiaozhiServerFacade:
"""获取支持的协议列表"""
if self.multi_protocol_server:
return self.multi_protocol_server.get_supported_protocols()
return ['websocket']
return ['websocket', 'mqtt']
def is_protocol_enabled(self, protocol: str) -> bool:
"""检查协议是否启用"""
@@ -416,6 +425,68 @@ class XiaozhiServerFacade:
if self.multi_protocol_server:
await self.multi_protocol_server.broadcast_message(message, protocol)
def _get_protocol_server(self, protocol: str):
if not self.multi_protocol_server:
return None
return self.multi_protocol_server.servers.get(protocol)
async def register_connection_context(self, context, transport) -> bool:
if getattr(transport, "transport_type", None) != "mqtt":
return False
server = self._get_protocol_server("mqtt")
if server is None:
return False
return await server.register_connection_context(context, transport)
async def unregister_connection_context(self, context, transport) -> bool:
if getattr(transport, "transport_type", None) != "mqtt":
return False
server = self._get_protocol_server("mqtt")
if server is None:
return False
return await server.unregister_connection_context(context, transport)
async def resolve_native_mqtt_connection(self, client_id: str):
server = self._get_protocol_server("mqtt")
if server is None:
return None
return await server.resolve_connection_context(client_id)
async def get_native_mqtt_status(self, client_ids):
server = self._get_protocol_server("mqtt")
if server is None:
return {
client_id: {
"isAlive": False,
"exists": False,
"backend": "native",
}
for client_id in client_ids
}
return await server.get_connection_status(client_ids)
async def request_native_mqtt_call(
self, caller_mac: str, target_mac: str, caller_nickname: str = ""
):
server = self._get_protocol_server("mqtt")
if server is None:
return {
"status": "error",
"message": "Native MQTT服务未启动",
}
return await server.request_device_call(
caller_mac, target_mac, caller_nickname
)
async def accept_native_mqtt_call(self, device_id: str):
server = self._get_protocol_server("mqtt")
if server is None:
return {
"status": "error",
"message": "Native MQTT服务未启动",
}
return await server.accept_device_call(device_id)
def get_websocket_info(self) -> Dict[str, Any]:
"""获取WebSocket连接信息"""
if not self.is_protocol_enabled('websocket'):
@@ -429,9 +500,24 @@ class XiaozhiServerFacade:
'path': '/xiaozhi/v1/'
}
def get_mqtt_info(self) -> Dict[str, Any]:
"""获取MQTT连接信息"""
if not self.is_protocol_enabled('mqtt'):
return {'enabled': False}
mqtt_config = self.config.get('mqtt_server', {})
return {
'enabled': True,
'host': mqtt_config.get('host', '0.0.0.0'),
'port': mqtt_config.get('port', 1883),
'udp_port': mqtt_config.get('udp_port', 1883),
'public_endpoint': mqtt_config.get('public_endpoint', '')
}
def get_connection_info(self) -> Dict[str, Any]:
"""获取所有协议的连接信息"""
return {
'websocket': self.get_websocket_info(),
'mqtt': self.get_mqtt_info(),
'active_connections': self.get_active_connections_count()
}