Compare commits

..
6 Commits
Author SHA1 Message Date
CGDandGitHub f2b4c7b932 Merge pull request #3308 from CaixyPromise/feature/native-mqtt-clean-pr
feature(mqtt): 新增原生MQTT与UDP传输支持
2026-07-27 16:28:41 +08:00
CGDandGitHub 5b288bb0d2 Merge pull request #3307 from CaixyPromise/feature/connection-runtime-pr
refactor: 传输协议无关的连接运行时
2026-07-27 16:26:18 +08:00
Sakura-RanChenandGitHub 5e0d853256 Merge pull request #3309 from xinnan-tech/py-fix-reportHandle
refactor(reportHandle): 重构上报处理逻辑,适配多格式音频转换
2026-07-27 15:24:15 +08:00
wengzh 2a9f809700 refactor(reportHandle): 重构上报处理逻辑,适配多格式音频转换
1.  重命名参数与变量,明确上报类型与音频格式
2.  拆分PCM和Opus转WAV逻辑,分别处理不同输入格式
3.  移动校验逻辑到上报队列函数开头,提前拦截无效上报
4.  完善Opus转WAV的解码与资源释放流程
2026-07-27 09:59:32 +08:00
caixypromise eac573706d feat: add optional native mqtt and udp transport 2026-07-27 02:10:55 +08:00
caixypromise 0c582ed3b6 refactor: introduce transport-neutral connection runtime 2026-07-27 02:05:26 +08:00
84 changed files with 16324 additions and 444 deletions
+161
View File
@@ -0,0 +1,161 @@
# Native MQTT/UDP integration
The server can accept device MQTT control messages and encrypted UDP audio
without running `xiaozhi-mqtt-gateway`. Native MQTT is optional. The existing
Gateway path and direct WebSocket path remain supported.
## Compatibility modes
| Mode | Device control | Device audio | Server business runtime |
| --- | --- | --- | --- |
| Direct WebSocket | WebSocket JSON | WebSocket binary | Shared runtime |
| MQTT Gateway | MQTT via Gateway WebSocket | UDP via Gateway WebSocket | Shared runtime |
| Native MQTT | MQTT 3.1.1 QoS 0 | AES-128-CTR UDP | Shared runtime |
Native MQTT keeps the physical MQTT connection alive between conversations.
Each device `hello` creates a new logical conversation and UDP encryption
session. `goodbye` ends only that logical conversation; MQTT is closed only by
MQTT disconnect, duplicate client takeover, keepAlive expiry, server shutdown,
or a protocol/authentication error.
## Enable Native MQTT
Native MQTT is disabled by default. Both the protocol and server switches must
be enabled:
```yaml
protocols:
enabled_protocols: [websocket, mqtt]
websocket_enabled: true
mqtt_enabled: true
mqtt_server:
enabled: true
host: 0.0.0.0
port: 1883
udp_port: 1883
udp_bind_host: ""
public_endpoint: mqtt.example.com
signature_key: replace-with-a-strong-secret
max_connections: 1000
max_pending_connections: 128
heartbeat_interval: 30
max_payload_size: 8192
message_queue_size: 128
business_ready_timeout: 30
goodbye_timeout: 1
close_timeout: 2
# Includes the common model and any Agent-specific local ASR variants.
shared_asr_max_models: 3
```
`host` is the local MQTT listen address. `public_endpoint` must be reachable by
the device and may be either `host` or `host:port`; manager-api normalizes it
and does not append the port twice. When `host` is a wildcard and
`public_endpoint` is a local IPv4 address, the UDP socket binds that address so
reply datagrams use the same source IP advertised to the device. Set
`udp_bind_host` only when the local UDP bind address must differ from the
advertised endpoint (for example, behind NAT); leaving it empty enables the
automatic behavior. Open the configured MQTT TCP and UDP ports in the firewall.
When configuration is read from manager-api, configure the equivalent
`protocols.*` and `mqtt_server.*` system parameters. The OTA response prefers a
valid Native endpoint when Native is enabled. If Native is disabled or invalid,
manager-api falls back to `server.mqtt_gateway`. If neither is configured, OTA
returns only WebSocket information.
The current firmware persists the OTA `mqtt` object. The UDP server, key and
nonce used by Native mode are negotiated in the MQTT `hello` response, so
manager-api does not send a separate top-level OTA `udp` object.
Long-lived Native connections refresh Agent-private components at logical
`hello` boundaries. Local ASR variants are shared by effective configuration
instead of loaded once per device. `shared_asr_max_models` bounds resident
models; when capacity is exhausted, the replacement runtime is rejected and
the previous healthy runtime remains active.
The default capacity reserves one transition slot so a single Agent can replace
an active private local model without dropping its healthy runtime first. Idle
variants are cached and evicted before another model is loaded; resident models
never exceed this limit.
MQTT application packets are dispatched in order outside the socket read loop,
so PINGREQ remains responsive while a logical Hello waits for private runtime
refresh. `business_ready_timeout` bounds that wait and closes an unusable
connection instead of leaving the device in a half-negotiated session.
`max_connections` limits authenticated client identities, while
`max_pending_connections` separately bounds sockets that have not completed
CONNECT. A connection using an existing client id may therefore replace its
previous owner even when active capacity is full. The server sends the success
CONNACK only after the previous owner has been reclaimed. `goodbye_timeout` and
`close_timeout` bound best-effort device reset and physical resource cleanup so
a stalled peer cannot block heartbeat or shutdown indefinitely.
## Authentication and topics
manager-api generates the device credentials and topics:
- client id: `<group>@@@<mac_with_underscores>@@@<mac_with_underscores>`
- publish topic: `device-server`
- subscribe topic: `devices/p2p/<mac_with_underscores>`
- username: Base64-encoded JSON metadata
- password: Base64(HMAC-SHA256(`client_id + "|" + username`, signature key))
Use the same secret for manager-api `mqtt_server.signature_key` and
xiaozhi-server `mqtt_server.signature_key`. The legacy
`server.mqtt_signature_key` remains a fallback for Gateway compatibility.
Native MQTT accepts MQTT 3.1.1 CONNECT, SUBSCRIBE, PUBLISH QoS 0, PING and
DISCONNECT. Unsupported protocol versions, QoS levels, malformed packets, or
invalid credentials close the physical connection.
## Session and UDP flow
1. Device connects and subscribes over MQTT.
2. Server validates CONNECT credentials and replaces any older connection with
the same client id.
3. Device publishes protocol version 3 `hello`.
4. Server waits until private configuration and business components are ready.
5. Server creates a logical session and returns the UDP server, AES key, nonce,
output audio parameters and session id.
6. The first valid UDP packet from the MQTT peer binds the UDP source tuple for
that logical session. Source changes are rejected until the next `hello`.
7. MQTT JSON and decrypted Opus frames enter the same message processors used
by WebSocket and Gateway connections.
8. Server `goodbye` returns the device to Idle and finalizes conversation tasks
while preserving the MQTT connection.
The 16-byte UDP header is also the AES-CTR nonce:
| Bytes | Field |
| --- | --- |
| 0 | packet type (`1`) |
| 1 | reserved |
| 2..3 | encrypted payload length, big endian |
| 4..7 | connection id, big endian |
| 8..11 | timestamp, big endian |
| 12..15 | sequence, big endian |
Datagrams with an invalid type, header, payload length, connection id, source
address or stale sequence are discarded. Sequence gaps are tolerated because
UDP can lose packets; late and duplicate packets are discarded. AES-CTR
provides confidentiality but not integrity, so Native MQTT and UDP should be
exposed only with a strong signature key and appropriate network controls.
## Verification
From `main/xiaozhi-server`:
```bash
python -m compileall -q .
```
From `main/manager-api` with JDK 21:
```bash
mvn clean package
```
Before deployment, verify Native and Gateway separately with real hardware:
multiple conversations, TTS interruption, exit to Idle, duplicate reconnect,
server restart, short network loss and a 30-60 minute keepAlive soak.
@@ -106,6 +106,32 @@ public interface Constant {
*/
String SERVER_MQTT_GATEWAY = "server.mqtt_gateway";
/**
* MQTT原生服务开关
*/
String SERVER_MQTT_ENABLED = "server.mqtt_enabled";
/**
* 新版协议配置
*/
String PROTOCOLS_ENABLED = "protocols.enabled_protocols";
String PROTOCOLS_WEBSOCKET_ENABLED = "protocols.websocket_enabled";
String PROTOCOLS_MQTT_ENABLED = "protocols.mqtt_enabled";
/**
* 新版原生MQTT配置
*/
String MQTT_SERVER_ENABLED = "mqtt_server.enabled";
String MQTT_SERVER_HOST = "mqtt_server.host";
String MQTT_SERVER_PORT = "mqtt_server.port";
String MQTT_SERVER_UDP_PORT = "mqtt_server.udp_port";
String MQTT_SERVER_UDP_BIND_HOST = "mqtt_server.udp_bind_host";
String MQTT_SERVER_PUBLIC_ENDPOINT = "mqtt_server.public_endpoint";
String MQTT_SERVER_SIGNATURE_KEY = "mqtt_server.signature_key";
String MQTT_SERVER_MANAGER_API = "mqtt_server.manager_api";
String MQTT_SERVER_MANAGER_API_SECRET = "mqtt_server.manager_api_secret";
String SERVER_MQTT_MANAGER_API = "server.mqtt_manager_api";
/**
* ota地址
*/
@@ -350,4 +376,4 @@ public interface Constant {
return value;
}
}
}
}
@@ -1,6 +1,7 @@
package xiaozhi.modules.device.controller;
import java.nio.charset.StandardCharsets;
import java.util.Arrays;
import org.apache.commons.lang3.StringUtils;
import org.springframework.http.MediaType;
@@ -78,8 +79,8 @@ public class OTAController {
@Hidden
public ResponseEntity<String> getOTA() {
String mqttUdpConfig = sysParamsService.getValue(Constant.SERVER_MQTT_GATEWAY, false);
if (StringUtils.isBlank(mqttUdpConfig)) {
return ResponseEntity.ok("OTA接口不正常,缺少mqtt_gateway地址,请登录智控台,在参数管理找到【server.mqtt_gateway】配置");
if (!isNativeMqttReady() && !isConfiguredValue(mqttUdpConfig)) {
return ResponseEntity.ok("OTA接口不正常,缺少mqtt_gateway地址或未启用原生MQTT,请登录智控台,在参数管理找到【server.mqtt_gateway】或【mqtt_server.enabled/protocols.mqtt_enabled】配置");
}
String wsUrl = sysParamsService.getValue(Constant.SERVER_WEBSOCKET, true);
if (StringUtils.isBlank(wsUrl) || wsUrl.equals("null")) {
@@ -92,6 +93,49 @@ public class OTAController {
return ResponseEntity.ok("OTA接口运行正常,websocket集群数量:" + wsUrl.split(";").length);
}
private boolean isNativeMqttEnabled() {
String enabled = sysParamsService.getValue(Constant.MQTT_SERVER_ENABLED, false);
boolean serverEnabled = isConfiguredValue(enabled)
? isTrue(enabled)
: isTrue(sysParamsService.getValue(Constant.SERVER_MQTT_ENABLED, false));
boolean protocolEnabled = isTrue(
sysParamsService.getValue(Constant.PROTOCOLS_MQTT_ENABLED, false));
if (!protocolEnabled) {
String enabledProtocols = sysParamsService.getValue(Constant.PROTOCOLS_ENABLED, false);
if (StringUtils.isNotBlank(enabledProtocols)) {
String normalized = enabledProtocols
.replace("[", "")
.replace("]", "")
.replace("\"", "");
protocolEnabled = Arrays.stream(normalized.split("[;,\\s]+"))
.anyMatch("mqtt"::equalsIgnoreCase);
}
}
return serverEnabled && protocolEnabled;
}
private boolean isNativeMqttReady() {
if (!isNativeMqttEnabled()) {
return false;
}
String signatureKey = sysParamsService.getValue(
Constant.MQTT_SERVER_SIGNATURE_KEY, false);
if (!isConfiguredValue(signatureKey)) {
signatureKey = sysParamsService.getValue(
Constant.SERVER_MQTT_SECRET, false);
}
return isConfiguredValue(signatureKey);
}
private boolean isTrue(String value) {
return "true".equalsIgnoreCase(value) || "1".equals(value);
}
private boolean isConfiguredValue(String value) {
return StringUtils.isNotBlank(value) && !"null".equalsIgnoreCase(value.trim());
}
@SneakyThrows
private ResponseEntity<String> createResponse(DeviceReportRespDTO deviceReportRespDTO) {
ObjectMapper objectMapper = new ObjectMapper();
@@ -18,6 +18,7 @@ import xiaozhi.common.redis.RedisKeys;
import xiaozhi.common.redis.RedisUtils;
import xiaozhi.modules.device.dao.DeviceAddressBookDao;
import xiaozhi.modules.device.entity.DeviceAddressBookEntity;
import xiaozhi.modules.device.entity.DeviceEntity;
import xiaozhi.modules.device.service.DeviceAddressBookService;
import xiaozhi.modules.device.service.DeviceService;
import xiaozhi.modules.sys.service.SysParamsService;
@@ -59,7 +60,14 @@ public class DeviceAddressBookServiceImpl implements DeviceAddressBookService {
Map<String, Map<String, String>> allBooks = getAllAddressBooks();
if (isAnswer) {
return postToMqtt("/api/call/accept", Map.of("mac", callerMac), "接听");
DeviceEntity callerDevice =
deviceService.getDeviceByMacAddress(callerMac);
if (callerDevice == null) {
return errorResult("接听失败,设备信息不存在");
}
return postCallAccept(
buildMqttClientId(callerDevice),
Map.of("mac", callerMac));
}
// 主动呼叫模式
@@ -79,6 +87,14 @@ public class DeviceAddressBookServiceImpl implements DeviceAddressBookService {
return errorResult("呼叫失败,您没有权限呼叫该设备");
}
DeviceEntity callerDevice =
deviceService.getDeviceByMacAddress(callerMac);
DeviceEntity targetDevice =
deviceService.getDeviceByMacAddress(targetMac);
if (callerDevice == null || targetDevice == null) {
return errorResult("呼叫失败,设备信息不存在");
}
// 获取目标设备如何称呼主叫方
Map<String, String> targetBook = allBooks.get(targetMac.toLowerCase());
String callerNickname = null;
@@ -86,15 +102,19 @@ public class DeviceAddressBookServiceImpl implements DeviceAddressBookService {
callerNickname = targetBook.get(callerMac.toLowerCase());
}
if (StringUtils.isBlank(callerNickname)) {
callerNickname = deviceService.getDeviceByMacAddress(callerMac).getAlias();
callerNickname = callerDevice.getAlias();
if (StringUtils.isBlank(callerNickname)) {
callerNickname = formatMacAsDeviceName(callerMac);
}
}
return postToMqtt("/api/call/request",
Map.of("caller_mac", callerMac, "target_mac", targetMac, "caller_nickname", callerNickname),
"呼叫");
return postCallRequest(
buildMqttClientId(callerDevice),
buildMqttClientId(targetDevice),
Map.of(
"caller_mac", callerMac,
"target_mac", targetMac,
"caller_nickname", callerNickname));
}
@Override
@@ -195,32 +215,50 @@ public class DeviceAddressBookServiceImpl implements DeviceAddressBookService {
return result;
}
private Map<String, Object> postToMqtt(String path, Map<String, Object> body, String action) {
private Map<String, Object> postCallRequest(
String callerClientId,
String targetClientId,
Map<String, Object> body) {
try {
MqttManagementHttpClient.Response response =
createMqttManagementRouter().sendCallRequest(
callerClientId, targetClientId, body);
return parseCallResponse(response, "呼叫");
} catch (Exception e) {
return errorResult("呼叫失败,请稍后再试");
}
}
private Map<String, Object> postCallAccept(
String clientId,
Map<String, Object> body) {
try {
MqttManagementHttpClient.Response response =
createMqttManagementRouter().sendCallAccept(
clientId, body);
return parseCallResponse(response, "接听");
} catch (Exception e) {
return errorResult("接听失败,请稍后再试");
}
}
private Map<String, Object> parseCallResponse(
MqttManagementHttpClient.Response response,
String action) {
Map<String, Object> result = new HashMap<>();
result.put("status", "error");
String mqttGatewayUrl = sysParamsService.getValue("server.mqtt_manager_api", true);
String mqttSignatureKey = sysParamsService.getValue(Constant.SERVER_MQTT_SECRET, true);
if (StringUtils.isBlank(mqttGatewayUrl) || "null".equals(mqttGatewayUrl)
|| MqttGatewayAuthorization.isMissingSignatureKey(mqttSignatureKey)) {
result.put("message", action + "失败,网关配置缺失");
if (response == null || StringUtils.isBlank(response.body())) {
result.put("message", action + "失败,MQTT管理配置缺失");
return result;
}
try {
String url = "http://" + mqttGatewayUrl + path;
String response = MqttGatewayAuthorization.postJson(
url,
JSONUtil.toJsonStr(body),
mqttSignatureKey,
Instant.now(),
5000);
if (StringUtils.isNotBlank(response)) {
Map<String, Object> gwResult = JSONUtil.parseObj(response);
result.put("status", gwResult.get("status"));
result.put("message", gwResult.get("message"));
Map<String, Object> backendResult =
JSONUtil.parseObj(response.body());
result.put("status", backendResult.get("status"));
result.put("message", backendResult.get("message"));
if (backendResult.containsKey("code")) {
result.put("code", backendResult.get("code"));
}
return result;
} catch (Exception e) {
@@ -229,6 +267,17 @@ public class DeviceAddressBookServiceImpl implements DeviceAddressBookService {
}
}
MqttManagementRouter createMqttManagementRouter() {
return new MqttManagementRouter(
new MqttManagementEndpointResolver(sysParamsService),
new MqttManagementHttpClient());
}
private String buildMqttClientId(DeviceEntity device) {
return MqttClientId.build(
device.getBoard(), device.getMacAddress());
}
private String formatMacAsDeviceName(String mac) {
if (StringUtils.isBlank(mac) || mac.length() < 2) {
return mac;
@@ -5,11 +5,13 @@ import java.security.InvalidKeyException;
import java.security.NoSuchAlgorithmException;
import java.time.Instant;
import java.util.ArrayList;
import java.util.Arrays;
import java.util.Base64;
import java.util.Collections;
import java.util.Date;
import java.util.HashMap;
import java.util.List;
import java.util.Locale;
import java.util.Map;
import java.util.Optional;
import java.util.Set;
@@ -157,32 +159,20 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
*/
@Override
public String getDeviceOnlineData(String agentId) {
// 从系统参数中获取MQTT网关地址
String mqttGatewayUrl = sysParamsService.getValue("server.mqtt_manager_api", true);
if (StringUtils.isBlank(mqttGatewayUrl) || "null".equals(mqttGatewayUrl)) {
return "";
}
// 构建完整的URL
String url = StrUtil.format("http://{}/api/devices/status", mqttGatewayUrl);
// 获取当前用户的设备列表
UserDetail user = SecurityUser.getUser();
List<DeviceEntity> devices = getUserDevices(user.getId(), agentId);
// 构建deviceIds数组
Set<String> deviceIds = devices.stream().map(o -> {
String macAddress = Optional.ofNullable(o.getMacAddress()).orElse("unknown").replace(":", "_");
String groupId = Optional.ofNullable(o.getBoard()).orElse("GID_default").replace(":", "_");
return StrUtil.format("{}@@@{}@@@{}", groupId, macAddress, macAddress);
}).collect(Collectors.toSet());
Set<String> deviceIds = devices.stream()
.map(device -> MqttClientId.build(
device.getBoard(), device.getMacAddress()))
.collect(Collectors.toSet());
// 构建请求入参
Map<String, Set<String>> params = MapUtil
.builder(new HashMap<String, Set<String>>())
.put("clientIds", deviceIds).build();
if (ToolUtil.isNotEmpty(deviceIds)) {
return postToMqttGateway(url, params);
return createMqttManagementRouter()
.getMergedStatus(deviceIds);
}
// 返回响应
return "";
@@ -194,6 +184,15 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
response.setServer_time(buildServerTime());
DeviceEntity deviceById = getDeviceByMacAddress(macAddress);
String reportedBoard = deviceReport.getBoard() == null
? null
: deviceReport.getBoard().getType();
if (deviceById != null
&& StringUtils.isBlank(deviceById.getBoard())
&& StringUtils.isNotBlank(reportedBoard)) {
deviceById.setBoard(reportedBoard);
baseDao.updateById(deviceById);
}
// 设备未绑定,则返回当前上传的固件信息(不更新)以此兼容旧固件版本
if (deviceById == null) {
@@ -204,8 +203,7 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
} else {
// 只有在设备已绑定且明确开启自动升级时才返回固件升级信息
if (Integer.valueOf(1).equals(deviceById.getAutoUpdate())) {
String type = deviceReport.getBoard() == null ? null : deviceReport.getBoard().getType();
DeviceReportRespDTO.Firmware firmware = buildFirmwareInfo(type,
DeviceReportRespDTO.Firmware firmware = buildFirmwareInfo(reportedBoard,
deviceReport.getApplication() == null ? null : deviceReport.getApplication().getVersion());
response.setFirmware(firmware);
}
@@ -248,20 +246,43 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
response.setWebsocket(websocket);
// 添加MQTT UDP配置
// 从系统参数获取MQTT Gateway地址,仅在配置有效时使用
String mqttUdpConfig = sysParamsService.getValue(Constant.SERVER_MQTT_GATEWAY, true);
if (mqttUdpConfig != null && !mqttUdpConfig.equals("null") && !mqttUdpConfig.isEmpty()) {
// UDP key/nonce/endpoint is negotiated by the MQTT hello response.
// OTA only needs to provide the MQTT connection parameters.
String groupId = "GID_default";
if (deviceById != null && StringUtils.isNotBlank(deviceById.getBoard())) {
groupId = deviceById.getBoard();
} else if (deviceReport.getBoard() != null && StringUtils.isNotBlank(deviceReport.getBoard().getType())) {
groupId = deviceReport.getBoard().getType();
}
boolean mqttConfigured = false;
if (isNativeMqttEnabled()) {
try {
String groupId = deviceById != null && deviceById.getBoard() != null ? deviceById.getBoard()
: "GID_default";
DeviceReportRespDTO.MQTT mqtt = buildMqttConfig(macAddress, groupId);
if (mqtt != null) {
mqtt.setEndpoint(mqttUdpConfig);
DeviceReportRespDTO.MQTT mqtt = buildNativeMqttConfig(macAddress, groupId, clientId);
String endpoint = buildNativeMqttEndpoint();
if (mqtt != null && StringUtils.isNotBlank(endpoint)) {
mqtt.setEndpoint(endpoint);
response.setMqtt(mqtt);
mqttConfigured = true;
}
} catch (Exception e) {
log.error("生成MQTT配置失败: {}", e.getMessage());
log.error("生成原生MQTT配置失败: {}", e.getMessage());
}
}
if (!mqttConfigured) {
// Native 未启用或配置无效时继续兼容 MQTT Gateway。
String mqttUdpConfig = sysParamsService.getValue(Constant.SERVER_MQTT_GATEWAY, true);
if (mqttUdpConfig != null && !mqttUdpConfig.equals("null") && !mqttUdpConfig.isEmpty()) {
try {
DeviceReportRespDTO.MQTT mqtt = buildMqttConfig(macAddress, groupId);
if (mqtt != null) {
mqtt.setEndpoint(mqttUdpConfig);
response.setMqtt(mqtt);
}
} catch (Exception e) {
log.error("生成MQTT配置失败: {}", e.getMessage());
}
}
}
@@ -633,6 +654,210 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
return String.format("%s.%d", signatureBase64, timestamp);
}
private boolean isNativeMqttEnabled() {
String enabled = sysParamsService.getValue(Constant.MQTT_SERVER_ENABLED, true);
String legacyEnabled = sysParamsService.getValue(Constant.SERVER_MQTT_ENABLED, true);
boolean serverEnabled = isConfiguredValue(enabled)
? isTrue(enabled)
: isTrue(legacyEnabled);
String protoEnabled = sysParamsService.getValue(Constant.PROTOCOLS_MQTT_ENABLED, true);
String enabledProtocols = sysParamsService.getValue(Constant.PROTOCOLS_ENABLED, true);
boolean protocolEnabled = isTrue(protoEnabled);
if (!protocolEnabled && StringUtils.isNotBlank(enabledProtocols)) {
String normalizedProtocols = enabledProtocols
.replace("[", "")
.replace("]", "")
.replace("\"", "");
protocolEnabled = Arrays.stream(normalizedProtocols.split("[;,\\s]+"))
.anyMatch("mqtt"::equalsIgnoreCase);
}
return serverEnabled && protocolEnabled;
}
private boolean isTrue(String value) {
return "true".equalsIgnoreCase(value) || "1".equals(value);
}
private boolean isConfiguredValue(String value) {
return StringUtils.isNotBlank(value)
&& !"null".equalsIgnoreCase(value.trim())
&& !value.contains("");
}
private String buildNativeMqttEndpoint() {
String publicEndpoint = sysParamsService.getValue(Constant.MQTT_SERVER_PUBLIC_ENDPOINT, true);
String host = sysParamsService.getValue(Constant.MQTT_SERVER_HOST, true);
EndpointParts endpointParts = parseEndpoint(publicEndpoint);
EndpointParts hostParts = parseEndpoint(host);
if ((isConfiguredValue(publicEndpoint) && endpointParts == null)
|| (isConfiguredValue(host) && hostParts == null)) {
return null;
}
EndpointParts selectedEndpoint = null;
if (endpointParts != null && isClientReachableMqttHost(endpointParts.host)) {
selectedEndpoint = endpointParts;
} else if (hostParts != null && isClientReachableMqttHost(hostParts.host)) {
selectedEndpoint = hostParts;
}
if (selectedEndpoint == null) {
return null;
}
String mqttHost = selectedEndpoint.host;
Integer port = selectedEndpoint.port;
if (port == null) {
port = parseInt(sysParamsService.getValue(Constant.MQTT_SERVER_PORT, true), 1883);
}
if (StringUtils.isBlank(mqttHost) || port == null) {
return null;
}
return formatEndpoint(mqttHost, port);
}
private boolean isClientReachableMqttHost(String host) {
if (StringUtils.isBlank(host)) {
return false;
}
String normalized = host.trim().toLowerCase(Locale.ROOT);
return !"null".equals(normalized)
&& !"localhost".equals(normalized)
&& !"localhost.localdomain".equals(normalized)
&& !"0.0.0.0".equals(normalized)
&& !normalized.startsWith("127.");
}
private Integer parseInt(String value, Integer defaultValue) {
if (StringUtils.isBlank(value) || "null".equalsIgnoreCase(value)) {
return defaultValue;
}
try {
int port = Integer.parseInt(value);
if (port < 1 || port > 65535) {
return null;
}
return port;
} catch (NumberFormatException e) {
return null;
}
}
private String formatEndpoint(String host, int port) {
String normalizedHost = host.trim();
if (normalizedHost.indexOf(':') >= 0) {
return null;
}
return normalizedHost + ":" + port;
}
private static class EndpointParts {
private final String host;
private final Integer port;
private EndpointParts(String host, Integer port) {
this.host = host;
this.port = port;
}
}
private EndpointParts parseEndpoint(String raw) {
if (StringUtils.isBlank(raw) || "null".equalsIgnoreCase(raw)) {
return new EndpointParts(null, null);
}
String trimmed = raw.trim();
trimmed = trimmed.replaceFirst("(?i)^(mqtt|tcp|ssl|ws|wss|http|https)://", "");
int slashIdx = trimmed.indexOf('/');
if (slashIdx >= 0) {
trimmed = trimmed.substring(0, slashIdx);
}
if (StringUtils.isBlank(trimmed) || trimmed.startsWith("[")
|| trimmed.indexOf(':') != trimmed.lastIndexOf(':')) {
return null;
}
String host = trimmed;
Integer port = null;
int colon = trimmed.lastIndexOf(':');
if (colon >= 0) {
if (colon == 0 || colon == trimmed.length() - 1) {
return null;
}
port = parseInt(trimmed.substring(colon + 1), null);
if (port == null) {
return null;
}
host = trimmed.substring(0, colon);
}
if (!isValidEndpointHost(host)) {
return null;
}
return new EndpointParts(host, port);
}
private boolean isValidEndpointHost(String host) {
if (StringUtils.isBlank(host)
|| host.length() > 253
|| host.chars().anyMatch(Character::isWhitespace)
|| host.indexOf('@') >= 0
|| host.indexOf('?') >= 0
|| host.indexOf('#') >= 0
|| host.indexOf('\\') >= 0) {
return false;
}
return Arrays.stream(host.split("\\.", -1))
.allMatch(label -> !label.isEmpty()
&& label.length() <= 63
&& Character.isLetterOrDigit(label.charAt(0))
&& Character.isLetterOrDigit(label.charAt(label.length() - 1))
&& label.chars().allMatch(ch ->
Character.isLetterOrDigit(ch)
|| ch == '-'
|| ch == '_'));
}
private String buildMqttUsername() throws Exception {
Map<String, String> userData = new HashMap<>();
try {
ServletRequestAttributes attributes = (ServletRequestAttributes) RequestContextHolder
.getRequestAttributes();
if (attributes != null) {
HttpServletRequest request = attributes.getRequest();
String clientIp = request.getRemoteAddr();
userData.put("ip", clientIp);
}
} catch (Exception e) {
userData.put("ip", "unknown");
}
String userDataJson = new com.fasterxml.jackson.databind.ObjectMapper().writeValueAsString(userData);
return Base64.getEncoder().encodeToString(userDataJson.getBytes(StandardCharsets.UTF_8));
}
private DeviceReportRespDTO.MQTT buildNativeMqttConfig(String macAddress, String groupId, String clientId)
throws Exception {
String deviceIdSafeStr = MqttClientId.normalizeDeviceId(macAddress);
String mqttClientId = MqttClientId.build(groupId, macAddress);
String username = buildMqttUsername();
String password = "";
String signatureKey = sysParamsService.getValue(Constant.MQTT_SERVER_SIGNATURE_KEY, true);
if (!isConfiguredValue(signatureKey)) {
signatureKey = sysParamsService.getValue(Constant.SERVER_MQTT_SECRET, true);
}
if (!isConfiguredValue(signatureKey)) {
log.error("原生MQTT已启用但未配置签名密钥,跳过原生MQTT配置下发");
return null;
}
password = generatePasswordSignature(mqttClientId + "|" + username, signatureKey);
DeviceReportRespDTO.MQTT mqtt = new DeviceReportRespDTO.MQTT();
mqtt.setClient_id(mqttClientId);
mqtt.setUsername(username);
mqtt.setPassword(password);
mqtt.setPublish_topic("device-server");
mqtt.setSubscribe_topic("devices/p2p/" + deviceIdSafeStr);
return mqtt;
}
/**
* 构建MQTT配置信息
*
@@ -650,28 +875,10 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
}
// 构建客户端ID格式:groupId@@@macAddress@@@uuid
String groupIdSafeStr = groupId.replace(":", "_");
String deviceIdSafeStr = macAddress.replace(":", "_");
String mqttClientId = String.format("%s@@@%s@@@%s", groupIdSafeStr, deviceIdSafeStr, deviceIdSafeStr);
String deviceIdSafeStr = MqttClientId.normalizeDeviceId(macAddress);
String mqttClientId = MqttClientId.build(groupId, macAddress);
// 构建用户数据(包含IP等信息)
Map<String, String> userData = new HashMap<>();
// 尝试获取客户端IP
try {
ServletRequestAttributes attributes = (ServletRequestAttributes) RequestContextHolder
.getRequestAttributes();
if (attributes != null) {
HttpServletRequest request = attributes.getRequest();
String clientIp = request.getRemoteAddr();
userData.put("ip", clientIp);
}
} catch (Exception e) {
userData.put("ip", "unknown");
}
// 将用户数据编码为Base64 JSON
String userDataJson = new com.fasterxml.jackson.databind.ObjectMapper().writeValueAsString(userData);
String username = Base64.getEncoder().encodeToString(userDataJson.getBytes(StandardCharsets.UTF_8));
String username = buildMqttUsername();
// 生成密码签名
String password = generatePasswordSignature(mqttClientId + "|" + username, signatureKey);
@@ -687,23 +894,14 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
return mqtt;
}
private String postToMqttGateway(String url, Object requestBody) {
String signatureKey = sysParamsService.getValue(Constant.SERVER_MQTT_SECRET, false);
return MqttGatewayAuthorization.postJson(
url,
JSONUtil.toJsonStr(requestBody),
signatureKey,
Instant.now());
MqttManagementRouter createMqttManagementRouter() {
return new MqttManagementRouter(
new MqttManagementEndpointResolver(sysParamsService),
new MqttManagementHttpClient());
}
@Override
public Object getDeviceTools(String deviceId) {
// 从系统参数中获取MQTT网关地址
String mqttGatewayUrl = sysParamsService.getValue("server.mqtt_manager_api", true);
if (StringUtils.isBlank(mqttGatewayUrl) || "null".equals(mqttGatewayUrl)) {
return null;
}
// 获取设备信息
DeviceEntity device = baseDao.selectById(deviceId);
if (device == null) {
@@ -717,19 +915,20 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
}
// 构建clientId
String macAddress = Optional.ofNullable(device.getMacAddress()).orElse("unknown").replace(":", "_");
String groupId = Optional.ofNullable(device.getBoard()).orElse("GID_default").replace(":", "_");
String clientId = StrUtil.format("{}@@@{}@@@{}", groupId, macAddress, macAddress);
// 构建完整的URL
String url = StrUtil.format("http://{}/api/commands/{}", mqttGatewayUrl, clientId);
String clientId = MqttClientId.build(
device.getBoard(), device.getMacAddress());
// 存储所有工具列表
List<Object> allTools = new ArrayList<>();
String cursor = null;
Set<String> seenCursors = new java.util.HashSet<>();
int pageCount = 0;
MqttManagementRouter managementRouter =
createMqttManagementRouter();
MqttManagementEndpointResolver.Backend selectedBackend = null;
// 循环获取分页数据
while (true) {
while (pageCount++ < 32) {
// 构建params
Map<String, Object> paramsMap = MapUtil.builder(new HashMap<String, Object>())
.put("withUserTools", true)
@@ -754,26 +953,44 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
.put("payload", payload)
.build();
String resultMessage = postToMqttGateway(url, requestBody);
MqttManagementHttpClient.Response response =
managementRouter.sendReadOnlyCommand(
clientId, requestBody, selectedBackend);
if (response == null || !response.isSuccessfulHttp()) {
return null;
}
String resultMessage =
response.body();
if (selectedBackend == null) {
selectedBackend = response.backend();
}
// 解析响应
if (StringUtils.isBlank(resultMessage)) {
break;
return null;
}
JSONObject jsonObject = JSONUtil.parseObj(resultMessage);
JSONObject jsonObject;
try {
jsonObject = JSONUtil.parseObj(resultMessage);
} catch (RuntimeException e) {
return null;
}
if (!jsonObject.getBool("success", false)) {
break;
return null;
}
JSONObject data = jsonObject.getJSONObject("data");
if (data == null) {
break;
return null;
}
// 获取当前页的工具列表
JSONArray tools = data.getJSONArray("tools");
if (tools != null && !tools.isEmpty()) {
if (tools == null) {
return null;
}
if (!tools.isEmpty()) {
allTools.addAll(tools);
}
@@ -781,29 +998,23 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
String nextCursor = data.getStr("nextCursor");
if (StringUtils.isBlank(nextCursor)) {
// 没有下一页了
break;
Map<String, Object> resultData = new HashMap<>();
resultData.put("tools", allTools);
return resultData;
}
if (!seenCursors.add(nextCursor)) {
log.warn("MQTT设备工具列表返回重复cursor,终止分页: {}",
nextCursor);
return null;
}
cursor = nextCursor;
}
// 构建返回结果
if (allTools.isEmpty()) {
return null;
}
Map<String, Object> resultData = new HashMap<>();
resultData.put("tools", allTools);
return resultData;
return null;
}
@Override
public Object callDeviceTool(String deviceId, String toolName, Map<String, Object> arguments) {
// 从系统参数中获取MQTT网关地址
String mqttGatewayUrl = sysParamsService.getValue("server.mqtt_manager_api", true);
if (StringUtils.isBlank(mqttGatewayUrl) || "null".equals(mqttGatewayUrl)) {
return null;
}
// 获取设备信息
DeviceEntity device = baseDao.selectById(deviceId);
if (device == null) {
@@ -817,12 +1028,8 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
}
// 构建clientId
String macAddress = Optional.ofNullable(device.getMacAddress()).orElse("unknown").replace(":", "_");
String groupId = Optional.ofNullable(device.getBoard()).orElse("GID_default").replace(":", "_");
String clientId = StrUtil.format("{}@@@{}@@@{}", groupId, macAddress, macAddress);
// 构建完整的URL
String url = StrUtil.format("http://{}/api/commands/{}", mqttGatewayUrl, clientId);
String clientId = MqttClientId.build(
device.getBoard(), device.getMacAddress());
// 构建请求体
Map<String, Object> params = MapUtil
@@ -845,7 +1052,11 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
.put("payload", payload)
.build();
String resultMessage = postToMqttGateway(url, requestBody);
MqttManagementHttpClient.Response response =
createMqttManagementRouter()
.sendMutatingCommand(clientId, requestBody);
String resultMessage =
response == null ? null : response.body();
// 解析响应
if (StringUtils.isNotBlank(resultMessage)) {
@@ -0,0 +1,28 @@
package xiaozhi.modules.device.service.impl;
import java.util.Locale;
import org.apache.commons.lang3.StringUtils;
final class MqttClientId {
private MqttClientId() {
}
static String build(String board, String macAddress) {
String groupId = StringUtils.defaultIfBlank(
board, "GID_default").trim().replace(":", "_");
String deviceId = StringUtils.defaultIfBlank(macAddress, "unknown")
.trim()
.replace(":", "_");
return groupId + "@@@" + deviceId + "@@@" + deviceId;
}
static String normalizeDeviceId(String macAddress) {
return StringUtils.defaultIfBlank(macAddress, "unknown")
.trim()
.toLowerCase(Locale.ROOT)
.replace(":", "_")
.replace("-", "_");
}
}
@@ -27,10 +27,8 @@ final class MqttGatewayAuthorization {
}
static String postJson(String url, String jsonBody, String signatureKey, Instant now, int timeoutMillis) {
GatewayResponse response = executeWithDateFallback(
signatureKey,
now,
token -> executeRequest(url, jsonBody, token, timeoutMillis));
GatewayResponse response = postJsonResponse(
url, jsonBody, signatureKey, now, timeoutMillis);
if (response.statusCode() < 200 || response.statusCode() >= 300) {
throw new GatewayRequestException(
@@ -40,6 +38,19 @@ final class MqttGatewayAuthorization {
return response.body();
}
static GatewayResponse postJsonResponse(
String url,
String jsonBody,
String signatureKey,
Instant now,
int timeoutMillis) {
return executeWithDateFallback(
signatureKey,
now,
token -> executeRequest(
url, jsonBody, token, timeoutMillis));
}
static List<String> generateDailyTokens(String signatureKey, Instant now) {
if (isMissingSignatureKey(signatureKey)) {
throw new GatewayRequestException("MQTT Gateway signature key is empty", null);
@@ -0,0 +1,183 @@
package xiaozhi.modules.device.service.impl;
import java.net.URI;
import java.net.URISyntaxException;
import java.util.ArrayList;
import java.util.Arrays;
import java.util.List;
import java.util.Locale;
import org.apache.commons.lang3.StringUtils;
import xiaozhi.common.constant.Constant;
import xiaozhi.modules.sys.service.SysParamsService;
class MqttManagementEndpointResolver {
enum Backend {
NATIVE,
GATEWAY
}
record Endpoint(Backend backend, String baseUrl, String signatureKey) {
}
private final SysParamsService sysParamsService;
MqttManagementEndpointResolver(SysParamsService sysParamsService) {
this.sysParamsService = sysParamsService;
}
List<Endpoint> resolve() {
List<Endpoint> endpoints = new ArrayList<>(2);
resolveNative().ifPresent(endpoints::add);
resolveGateway().ifPresent(endpoints::add);
return endpoints;
}
java.util.Optional<Endpoint> resolveNative() {
if (!isNativeMqttEnabled()) {
return java.util.Optional.empty();
}
String endpoint = normalizeHttpEndpoint(sysParamsService.getValue(
Constant.MQTT_SERVER_MANAGER_API, true));
String signatureKey = firstConfigured(
sysParamsService.getValue(
Constant.MQTT_SERVER_MANAGER_API_SECRET, true),
sysParamsService.getValue(
Constant.MQTT_SERVER_SIGNATURE_KEY, true),
sysParamsService.getValue(
Constant.SERVER_MQTT_SECRET, true));
if (endpoint == null || signatureKey == null) {
throw new ManagementConfigurationException(
"原生MQTT已启用但管理端点或密钥无效");
}
return java.util.Optional.of(
new Endpoint(Backend.NATIVE, endpoint, signatureKey));
}
java.util.Optional<Endpoint> resolveGateway() {
String rawGateway = sysParamsService.getValue(
Constant.SERVER_MQTT_GATEWAY, true);
String rawEndpoint = sysParamsService.getValue(
Constant.SERVER_MQTT_MANAGER_API, true);
if (configured(rawGateway) == null
&& configured(rawEndpoint) == null) {
return java.util.Optional.empty();
}
String endpoint = normalizeHttpEndpoint(rawEndpoint);
String signatureKey = configured(
sysParamsService.getValue(
Constant.SERVER_MQTT_SECRET, true));
if (endpoint == null || signatureKey == null) {
throw new ManagementConfigurationException(
"MQTT Gateway管理端点或密钥无效");
}
return java.util.Optional.of(
new Endpoint(Backend.GATEWAY, endpoint, signatureKey));
}
boolean isNativeMqttEnabled() {
String enabled = sysParamsService.getValue(
Constant.MQTT_SERVER_ENABLED, true);
String legacyEnabled = sysParamsService.getValue(
Constant.SERVER_MQTT_ENABLED, true);
boolean serverEnabled = configured(enabled) != null
? isTrue(enabled)
: isTrue(legacyEnabled);
String protocolEnabled = sysParamsService.getValue(
Constant.PROTOCOLS_MQTT_ENABLED, true);
String enabledProtocols = sysParamsService.getValue(
Constant.PROTOCOLS_ENABLED, true);
boolean enabledByProtocol = isTrue(protocolEnabled);
if (!enabledByProtocol && StringUtils.isNotBlank(enabledProtocols)) {
String normalized = enabledProtocols
.replace("[", "")
.replace("]", "")
.replace("\"", "");
enabledByProtocol = Arrays.stream(
normalized.split("[;,\\s]+"))
.anyMatch("mqtt"::equalsIgnoreCase);
}
return serverEnabled && enabledByProtocol;
}
static String normalizeHttpEndpoint(String raw) {
String configured = configured(raw);
if (configured == null) {
return null;
}
boolean hasScheme = configured.matches(
"^[A-Za-z][A-Za-z0-9+.-]*://.*");
if (hasScheme
&& !configured.matches("(?i)^https?://.*")) {
return null;
}
String candidate = hasScheme ? configured : "http://" + configured;
try {
URI uri = new URI(candidate);
String scheme = uri.getScheme();
if (scheme == null
|| (!"http".equalsIgnoreCase(scheme)
&& !"https".equalsIgnoreCase(scheme))
|| StringUtils.isBlank(uri.getHost())
|| uri.getUserInfo() != null
|| uri.getQuery() != null
|| uri.getFragment() != null
|| uri.getPort() == 0
|| uri.getPort() > 65535) {
return null;
}
String path = uri.getPath();
if (path == null || "/".equals(path)) {
path = "";
} else {
path = path.replaceAll("/+$", "");
}
return new URI(
scheme.toLowerCase(Locale.ROOT),
null,
uri.getHost(),
uri.getPort(),
path,
null,
null).toString();
} catch (URISyntaxException | IllegalArgumentException e) {
return null;
}
}
private static String firstConfigured(String... values) {
for (String value : values) {
String normalized = configured(value);
if (normalized != null) {
return normalized;
}
}
return null;
}
private static String configured(String value) {
if (StringUtils.isBlank(value)) {
return null;
}
String normalized = value.trim();
if ("null".equalsIgnoreCase(normalized)
|| normalized.contains("")) {
return null;
}
return normalized;
}
private static boolean isTrue(String value) {
return "true".equalsIgnoreCase(value) || "1".equals(value);
}
static final class ManagementConfigurationException
extends RuntimeException {
ManagementConfigurationException(String message) {
super(message);
}
}
}
@@ -0,0 +1,56 @@
package xiaozhi.modules.device.service.impl;
import java.time.Instant;
import cn.hutool.json.JSONUtil;
class MqttManagementHttpClient {
private final int timeoutMillis;
MqttManagementHttpClient() {
this(5000);
}
MqttManagementHttpClient(int timeoutMillis) {
this.timeoutMillis = timeoutMillis;
}
Response post(
MqttManagementEndpointResolver.Endpoint endpoint,
String path,
Object requestBody) {
String url = appendPath(endpoint.baseUrl(), path);
MqttGatewayAuthorization.GatewayResponse response =
MqttGatewayAuthorization.postJsonResponse(
url,
JSONUtil.toJsonStr(requestBody),
endpoint.signatureKey(),
Instant.now(),
timeoutMillis);
return new Response(
endpoint.backend(),
response.statusCode(),
response.body());
}
static String appendPath(String baseUrl, String path) {
String normalizedBase = baseUrl.endsWith("/")
? baseUrl.substring(0, baseUrl.length() - 1)
: baseUrl;
String normalizedPath = path.startsWith("/")
? path
: "/" + path;
return normalizedBase + normalizedPath;
}
record Response(
MqttManagementEndpointResolver.Backend backend,
int statusCode,
String body) {
boolean isSuccessfulHttp() {
return statusCode >= 200 && statusCode < 300;
}
}
}
@@ -0,0 +1,431 @@
package xiaozhi.modules.device.service.impl;
import java.net.URLEncoder;
import java.nio.charset.StandardCharsets;
import java.util.ArrayList;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import java.util.Set;
import cn.hutool.json.JSONObject;
import cn.hutool.json.JSONUtil;
final class MqttManagementRouter {
private final MqttManagementEndpointResolver endpointResolver;
private final MqttManagementHttpClient httpClient;
MqttManagementRouter(
MqttManagementEndpointResolver endpointResolver,
MqttManagementHttpClient httpClient) {
this.endpointResolver = endpointResolver;
this.httpClient = httpClient;
}
String getMergedStatus(Set<String> clientIds) {
List<MqttManagementEndpointResolver.Endpoint> endpoints =
endpointResolver.resolve();
if (endpoints.isEmpty()) {
return "";
}
Map<String, MutableDeviceStatus> merged =
initializeStatuses(clientIds);
int successfulBackends = 0;
RuntimeException lastFailure = null;
for (MqttManagementEndpointResolver.Endpoint endpoint : endpoints) {
try {
MqttManagementHttpClient.Response response =
httpClient.post(
endpoint,
"/api/devices/status",
Map.of("clientIds", clientIds));
if (!response.isSuccessfulHttp()) {
throw new IllegalStateException(
"MQTT管理状态查询失败: "
+ response.statusCode());
}
mergeStatusResponse(
merged, endpoint.backend(), response.body());
successfulBackends++;
} catch (RuntimeException e) {
lastFailure = e;
}
}
if (lastFailure != null
|| successfulBackends != endpoints.size()) {
throw new ManagementUnavailableException(
"无法确认全部MQTT后端状态", lastFailure);
}
Map<String, Object> result = new LinkedHashMap<>();
merged.forEach((clientId, status) ->
result.put(clientId, status.toMap()));
return JSONUtil.toJsonStr(result);
}
MqttManagementHttpClient.Response sendReadOnlyCommand(
String clientId, Object requestBody) {
return sendReadOnlyCommand(clientId, requestBody, null);
}
MqttManagementHttpClient.Response sendReadOnlyCommand(
String clientId,
Object requestBody,
MqttManagementEndpointResolver.Backend preferredBackend) {
List<MqttManagementEndpointResolver.Endpoint> endpoints =
endpointResolver.resolve();
if (endpoints.isEmpty()) {
return null;
}
if (preferredBackend != null) {
MqttManagementEndpointResolver.Endpoint endpoint =
endpoints.stream()
.filter(candidate ->
candidate.backend()
== preferredBackend)
.findFirst()
.orElseThrow(() ->
new ManagementUnavailableException(
"分页MQTT后端配置已变化",
null));
return sendCommand(endpoint, clientId, requestBody);
}
MqttManagementHttpClient.Response lastResponse = null;
RuntimeException lastFailure = null;
for (MqttManagementEndpointResolver.Endpoint endpoint : endpoints) {
try {
MqttManagementHttpClient.Response response =
sendCommand(endpoint, clientId, requestBody);
lastResponse = response;
if (isCommandSuccess(response)) {
return response;
}
} catch (RuntimeException e) {
lastFailure = e;
}
}
if (lastResponse != null) {
return lastResponse;
}
throw new ManagementUnavailableException(
"MQTT管理服务均不可用", lastFailure);
}
MqttManagementHttpClient.Response sendMutatingCommand(
String clientId, Object requestBody) {
List<MqttManagementEndpointResolver.Endpoint> endpoints =
endpointResolver.resolve();
if (endpoints.isEmpty()) {
return null;
}
List<MqttManagementEndpointResolver.Endpoint> online =
new ArrayList<>();
int successfulStatusBackends = 0;
RuntimeException lastStatusFailure = null;
for (MqttManagementEndpointResolver.Endpoint endpoint : endpoints) {
try {
MqttManagementHttpClient.Response statusResponse =
httpClient.post(
endpoint,
"/api/devices/status",
Map.of("clientIds", Set.of(clientId)));
if (!statusResponse.isSuccessfulHttp()) {
throw new IllegalStateException(
"MQTT管理状态查询失败: "
+ statusResponse.statusCode());
}
Map<String, BackendDeviceStatus> statuses =
parseStatusResponse(
statusResponse.body(),
Set.of(clientId));
successfulStatusBackends++;
if (statuses.get(clientId).exists()) {
online.add(endpoint);
}
} catch (RuntimeException e) {
lastStatusFailure = e;
}
}
if (lastStatusFailure != null
|| successfulStatusBackends != endpoints.size()) {
throw new ManagementUnavailableException(
"无法确认设备所在的MQTT后端",
lastStatusFailure);
}
if (online.isEmpty()) {
return offlineResponse(endpoints.get(0).backend());
}
if (online.size() > 1) {
return commandErrorResponse(
endpoints.get(0).backend(),
409,
"设备同时存在于多个MQTT后端",
"DEVICE_BACKEND_AMBIGUOUS");
}
return sendCommand(online.get(0), clientId, requestBody);
}
MqttManagementHttpClient.Response sendCallRequest(
String callerClientId,
String targetClientId,
Object requestBody) {
Set<String> clientIds = new java.util.LinkedHashSet<>();
clientIds.add(callerClientId);
clientIds.add(targetClientId);
return sendCallMutation(
clientIds,
"/api/call/request",
requestBody);
}
MqttManagementHttpClient.Response sendCallAccept(
String clientId, Object requestBody) {
return sendCallMutation(
Set.of(clientId),
"/api/call/accept",
requestBody);
}
private MqttManagementHttpClient.Response sendCallMutation(
Set<String> clientIds,
String path,
Object requestBody) {
List<MqttManagementEndpointResolver.Endpoint> endpoints =
endpointResolver.resolve();
if (endpoints.isEmpty()) {
return null;
}
Map<String, List<MqttManagementEndpointResolver.Endpoint>> owners =
new LinkedHashMap<>();
clientIds.forEach(clientId ->
owners.put(clientId, new ArrayList<>()));
RuntimeException statusFailure = null;
int successfulStatusBackends = 0;
for (MqttManagementEndpointResolver.Endpoint endpoint : endpoints) {
try {
MqttManagementHttpClient.Response response =
httpClient.post(
endpoint,
"/api/devices/status",
Map.of("clientIds", clientIds));
if (!response.isSuccessfulHttp()) {
throw new IllegalStateException(
"MQTT管理状态查询失败: "
+ response.statusCode());
}
Map<String, BackendDeviceStatus> statuses =
parseStatusResponse(response.body(), clientIds);
owners.forEach((clientId, matches) -> {
if (statuses.get(clientId).exists()) {
matches.add(endpoint);
}
});
successfulStatusBackends++;
} catch (RuntimeException e) {
statusFailure = e;
}
}
if (statusFailure != null
|| successfulStatusBackends != endpoints.size()) {
throw new ManagementUnavailableException(
"无法确认通话设备所在的MQTT后端",
statusFailure);
}
for (Map.Entry<String, List<MqttManagementEndpointResolver.Endpoint>>
owner : owners.entrySet()) {
if (owner.getValue().isEmpty()) {
return callErrorResponse(
endpoints.get(0).backend(),
404,
"offline",
"设备不在线",
"DEVICE_OFFLINE");
}
if (owner.getValue().size() > 1) {
return callErrorResponse(
endpoints.get(0).backend(),
409,
"error",
"设备同时存在于多个MQTT后端",
"DEVICE_BACKEND_AMBIGUOUS");
}
}
MqttManagementEndpointResolver.Endpoint selected =
owners.values().iterator().next().get(0);
boolean sameBackend = owners.values().stream()
.allMatch(matches -> matches.get(0).equals(selected));
if (!sameBackend) {
return callErrorResponse(
selected.backend(),
409,
"error",
"跨MQTT后端通话暂不支持",
"CROSS_BACKEND_CALL_UNSUPPORTED");
}
return httpClient.post(selected, path, requestBody);
}
private MqttManagementHttpClient.Response sendCommand(
MqttManagementEndpointResolver.Endpoint endpoint,
String clientId,
Object requestBody) {
String encodedClientId = URLEncoder.encode(
clientId, StandardCharsets.UTF_8)
.replace("+", "%20");
return httpClient.post(
endpoint,
"/api/commands/" + encodedClientId,
requestBody);
}
private static Map<String, MutableDeviceStatus> initializeStatuses(
Set<String> clientIds) {
Map<String, MutableDeviceStatus> result = new LinkedHashMap<>();
clientIds.forEach(clientId ->
result.put(clientId, new MutableDeviceStatus()));
return result;
}
private static void mergeStatusResponse(
Map<String, MutableDeviceStatus> merged,
MqttManagementEndpointResolver.Backend backend,
String body) {
Map<String, BackendDeviceStatus> response =
parseStatusResponse(body, merged.keySet());
merged.forEach((clientId, status) -> {
BackendDeviceStatus value = response.get(clientId);
status.exists |= value.exists();
status.alive |= value.alive();
if (value.exists()) {
status.backends.add(
backend.name().toLowerCase());
}
});
}
private static Map<String, BackendDeviceStatus> parseStatusResponse(
String body, Set<String> clientIds) {
JSONObject response;
try {
response = JSONUtil.parseObj(body);
} catch (RuntimeException e) {
throw new IllegalStateException(
"MQTT管理状态响应不是JSON对象", e);
}
Map<String, BackendDeviceStatus> statuses =
new LinkedHashMap<>();
for (String clientId : clientIds) {
Object rawStatus = response.get(clientId);
if (!(rawStatus instanceof JSONObject status)
|| !status.containsKey("exists")
|| !status.containsKey("isAlive")) {
throw new IllegalStateException(
"MQTT管理状态响应缺少设备状态: " + clientId);
}
Object rawExists = status.get("exists");
Object rawAlive = status.get("isAlive");
if (!(rawExists instanceof Boolean exists)
|| !(rawAlive instanceof Boolean alive)
|| (!exists && alive)) {
throw new IllegalStateException(
"MQTT管理状态响应字段无效: " + clientId);
}
statuses.put(
clientId,
new BackendDeviceStatus(exists, alive));
}
return statuses;
}
private static boolean isCommandSuccess(
MqttManagementHttpClient.Response response) {
if (response == null || !response.isSuccessfulHttp()) {
return false;
}
try {
return JSONUtil.parseObj(response.body())
.getBool("success", false);
} catch (Exception e) {
return false;
}
}
private static MqttManagementHttpClient.Response offlineResponse(
MqttManagementEndpointResolver.Backend backend) {
return new MqttManagementHttpClient.Response(
backend,
404,
JSONUtil.toJsonStr(Map.of(
"success", false,
"error", "设备未连接",
"code", "DEVICE_OFFLINE",
"dispatchAttempted", false)));
}
private static MqttManagementHttpClient.Response callErrorResponse(
MqttManagementEndpointResolver.Backend backend,
int statusCode,
String status,
String message,
String code) {
return new MqttManagementHttpClient.Response(
backend,
statusCode,
JSONUtil.toJsonStr(Map.of(
"status", status,
"message", message,
"code", code)));
}
private static MqttManagementHttpClient.Response commandErrorResponse(
MqttManagementEndpointResolver.Backend backend,
int statusCode,
String error,
String code) {
return new MqttManagementHttpClient.Response(
backend,
statusCode,
JSONUtil.toJsonStr(Map.of(
"success", false,
"error", error,
"code", code,
"dispatchAttempted", false)));
}
static final class ManagementUnavailableException
extends RuntimeException {
ManagementUnavailableException(
String message, Throwable cause) {
super(message, cause);
}
}
private static final class MutableDeviceStatus {
private boolean exists;
private boolean alive;
private final List<String> backends = new ArrayList<>();
Map<String, Object> toMap() {
Map<String, Object> result = new LinkedHashMap<>();
result.put("isAlive", alive);
result.put("exists", exists);
result.put("backends", backends);
return result;
}
}
private record BackendDeviceStatus(boolean exists, boolean alive) {
}
}
@@ -85,6 +85,7 @@ public class SysParamsController {
public Result<Void> save(@RequestBody SysParamsDTO dto) {
// 效验数据
ValidatorUtils.validateEntity(dto, AddGroup.class, DefaultGroup.class);
validateMqttSecretLength(dto.getParamCode(), dto.getParamValue());
sysParamsService.save(dto);
configService.getConfig(false);
@@ -277,7 +278,10 @@ public class SysParamsController {
// 校验mqtt密钥长度和复杂度
private void validateMqttSecretLength(String paramCode, String secret) {
if (!paramCode.equals(Constant.SERVER_MQTT_SECRET)) {
if (!paramCode.equals(Constant.SERVER_MQTT_SECRET)
&& !paramCode.equals(Constant.MQTT_SERVER_SIGNATURE_KEY)
&& !paramCode.equals(
Constant.MQTT_SERVER_MANAGER_API_SECRET)) {
return;
}
if (StringUtils.isBlank(secret) || secret.equals("null")) {
@@ -205,9 +205,14 @@ public class SysParamsServiceImpl extends BaseServiceImpl<SysParamsDao, SysParam
public void initServerSecret() {
// 获取服务器密钥
String secretParam = getValue(Constant.SERVER_SECRET, false);
if (StringUtils.isBlank(secretParam) || "null".equals(secretParam)) {
if (StringUtils.isBlank(secretParam)
|| "null".equalsIgnoreCase(secretParam.trim())) {
String newSecret = UUID.randomUUID().toString();
updateValueByCode(Constant.SERVER_SECRET, newSecret);
} else {
// Startup must repair a stale Redis value because authentication
// reads this parameter from cache while initialization reads DB.
sysParamsRedis.set(Constant.SERVER_SECRET, secretParam);
}
// 初始化SM2密钥对
@@ -0,0 +1,15 @@
insert ignore into `sys_params`
(id, param_code, param_value, value_type, param_type, remark)
values
(775, 'protocols.enabled_protocols', '["websocket"]', 'string', 1, '启用的协议列表'),
(776, 'protocols.websocket_enabled', 'true', 'boolean', 1, 'WebSocket协议开关'),
(777, 'protocols.mqtt_enabled', 'false', 'boolean', 1, 'MQTT协议开关'),
(778, 'mqtt_server.enabled', 'false', 'boolean', 1, '是否启用MQTT服务器'),
(779, 'mqtt_server.host', '0.0.0.0', 'string', 1, 'MQTT服务器监听地址'),
(780, 'mqtt_server.port', '1883', 'number', 1, 'MQTT服务器端口'),
(781, 'mqtt_server.udp_port', '1883', 'number', 1, 'UDP音频端口'),
(782, 'mqtt_server.public_endpoint', '127.0.0.1', 'string', 1, 'MQTT公网/局域网访问地址'),
(783, 'mqtt_server.max_connections', '1000', 'number', 1, '最大连接数'),
(784, 'mqtt_server.heartbeat_interval', '30', 'number', 1, '心跳间隔(秒)'),
(785, 'mqtt_server.max_payload_size', '8192', 'number', 1, '最大消息大小'),
(786, 'mqtt_server.signature_key', 'null', 'string', 1, 'MQTT签名密钥');
@@ -0,0 +1,9 @@
insert ignore into `sys_params`
(id, param_code, param_value, value_type, param_type, remark)
values
(787, 'mqtt_server.udp_bind_host', '', 'string', 1, 'UDP音频监听地址(留空时自动选择)'),
(788, 'mqtt_server.message_queue_size', '128', 'number', 1, 'MQTT应用消息队列上限'),
(789, 'mqtt_server.business_ready_timeout', '30', 'number', 1, 'Hello等待业务运行时就绪超时(秒)'),
(790, 'mqtt_server.max_pending_connections', '128', 'number', 1, 'MQTT待认证连接上限'),
(791, 'mqtt_server.goodbye_timeout', '1', 'number', 1, 'MQTT goodbye发送超时(秒)'),
(792, 'mqtt_server.close_timeout', '2', 'number', 1, 'MQTT连接清理超时(秒)');
@@ -0,0 +1,11 @@
insert ignore into `sys_params`
(id, param_code, param_value, value_type, param_type, remark)
values
(793, 'mqtt_server.manager_api', 'null', 'string', 1, '原生MQTT管理API地址'),
(794, 'mqtt_server.manager_api_secret', 'null', 'string', 1, '原生MQTT管理API签名密钥');
update `sys_params`
set param_value = '',
remark = 'MQTT公网/局域网访问地址(需配置为设备可达地址)'
where param_code = 'mqtt_server.public_endpoint'
and param_value = '127.0.0.1';
@@ -711,3 +711,26 @@ databaseChangeLog:
- sqlFile:
encoding: utf8
path: classpath:db/changelog/202607101200.sql
- changeSet:
id: 202607190300
author: CAIXYPROMISE
validCheckSum:
- "8:5df1e1f3654e6271889d5579a3908f77"
changes:
- sqlFile:
encoding: utf8
path: classpath:db/changelog/202607190300.sql
- changeSet:
id: 202607260900
author: CAIXYPROMISE
changes:
- sqlFile:
encoding: utf8
path: classpath:db/changelog/202607260900.sql
- changeSet:
id: 202607262100
author: CAIXYPROMISE
changes:
- sqlFile:
encoding: utf8
path: classpath:db/changelog/202607262100.sql
+146 -37
View File
@@ -3,13 +3,14 @@ import uuid
import signal
import asyncio
from aioconsole import ainput
from config.settings import load_config
from config.config_loader import load_config
from config.logger import setup_logging
from core.utils.util import get_local_ip, validate_mcp_endpoint
from core.http_server import SimpleHttpServer
from core.websocket_server import WebSocketServer
from core.xiaozhi_server_facade import XiaozhiServerFacade
from core.utils.util import check_ffmpeg_installed
from core.utils.gc_manager import get_gc_manager
from config.manage_api_client import manage_api_http_close
TAG = __name__
logger = setup_logging()
@@ -51,14 +52,14 @@ async def main():
# auth_key用于jwt认证,比如视觉分析接口的jwt认证、ota接口的token生成与websocket认证
# 获取配置文件中的auth_key
auth_key = config["server"].get("auth_key", "")
# 验证auth_key,无效则尝试使用manager-api.secret
if not auth_key or len(auth_key) == 0 or "" in auth_key:
auth_key = config.get("manager-api", {}).get("secret", "")
# 验证secret,无效则生成随机密钥
if not auth_key or len(auth_key) == 0 or "" in auth_key:
auth_key = str(uuid.uuid4().hex)
config["server"]["auth_key"] = auth_key
# 添加 stdin 监控任务
@@ -68,12 +69,49 @@ async def main():
gc_manager = get_gc_manager(interval_seconds=300)
await gc_manager.start()
# 启动 WebSocket 服务器
ws_server = WebSocketServer(config)
ws_task = asyncio.create_task(ws_server.start())
# 启动 Simple http 服务器
ota_server = SimpleHttpServer(config)
ota_task = asyncio.create_task(ota_server.start())
# 启动小智服务器门面(支持WebSocket和MQTT
xiaozhi_server = XiaozhiServerFacade(config)
ota_server = None
ota_task = None
try:
# Facade.start() returns after every listener is ready. Await it so bind
# or configuration failures stop startup instead of leaving a partial process.
await xiaozhi_server.start()
ota_server = SimpleHttpServer(
config, management_owner=xiaozhi_server
)
ota_task = asyncio.create_task(
ota_server.start(), name="xiaozhi-http-server"
)
await ota_server.wait_started(
ota_task,
timeout=float(config.get("server_startup_timeout", 10)),
)
except Exception:
if ota_server:
try:
await ota_server.stop()
except Exception as cleanup_error:
logger.bind(tag=TAG).error(f"停止HTTP服务器失败: {cleanup_error}")
if ota_task:
if not ota_task.done():
ota_task.cancel()
await asyncio.gather(ota_task, return_exceptions=True)
try:
await xiaozhi_server.stop()
except Exception as cleanup_error:
logger.bind(tag=TAG).error(f"停止协议服务器失败: {cleanup_error}")
try:
await gc_manager.stop()
except Exception as cleanup_error:
logger.bind(tag=TAG).error(f"停止GC管理器失败: {cleanup_error}")
try:
await manage_api_http_close()
except Exception as cleanup_error:
logger.bind(tag=TAG).error(f"关闭管理API客户端失败: {cleanup_error}")
stdin_task.cancel()
await asyncio.gather(stdin_task, return_exceptions=True)
raise
read_config_from_api = config.get("read_config_from_api", False)
port = int(config["server"].get("http_port", 8003))
@@ -100,24 +138,51 @@ async def main():
logger.bind(tag=TAG).error("mcp接入点不符合规范")
config["mcp_endpoint"] = "你的接入点 websocket地址"
# 获取WebSocket配置,使用安全的默认值
websocket_port = 8000
server_config = config.get("server", {})
if isinstance(server_config, dict):
websocket_port = int(server_config.get("port", 8000))
# 显示协议连接信息
connection_info = xiaozhi_server.get_connection_info()
logger.bind(tag=TAG).info(
"Websocket地址是\tws://{}:{}/xiaozhi/v1/",
get_local_ip(),
websocket_port,
)
# WebSocket信息
websocket_info = connection_info.get("websocket", {})
if websocket_info.get("enabled", False):
websocket_port = websocket_info.get("port", 8000)
logger.bind(tag=TAG).info(
"WebSocket地址是\tws://{}:{}/xiaozhi/v1/",
get_local_ip(),
websocket_port,
)
logger.bind(tag=TAG).info(
"=======上面的地址是websocket协议地址,请勿用浏览器访问======="
)
logger.bind(tag=TAG).info(
"如想测试websocket请启动digital-human模块,打开浏览器交互测试"
)
# MQTT信息
mqtt_info = connection_info.get("mqtt", {})
if mqtt_info.get("enabled", False):
mqtt_port = mqtt_info.get("port", 1883)
udp_port = mqtt_info.get("udp_port", 1883)
logger.bind(tag=TAG).info(
"MQTT地址是\t\tmqtt://{}:{}",
get_local_ip(),
mqtt_port,
)
logger.bind(tag=TAG).info(
"UDP音频端口是\t{}:{}",
get_local_ip(),
udp_port,
)
# 显示启用的协议
enabled_protocols = xiaozhi_server.config.get("enabled_protocols", [])
logger.bind(tag=TAG).info(f"启用的协议: {', '.join(enabled_protocols)}")
if "websocket" in enabled_protocols:
logger.bind(tag=TAG).info(
"=======上面的WebSocket地址请勿用浏览器访问======="
)
logger.bind(tag=TAG).info(
"如想测试WebSocket请启动digital-human模块,打开浏览器交互测试"
)
if "mqtt" in enabled_protocols:
logger.bind(tag=TAG).info(
"=======MQTT客户端ID格式: GID_test@@@mac_address@@@uuid======="
)
logger.bind(tag=TAG).info(
"=============================================================\n"
)
@@ -127,22 +192,66 @@ async def main():
except asyncio.CancelledError:
print("任务被取消,清理资源中...")
finally:
shutdown_errors = []
if ota_server:
try:
await ota_server.stop()
except Exception as e:
shutdown_errors.append(("HTTP服务器", e))
logger.bind(tag=TAG).error(f"停止HTTP服务器失败: {e}")
# 停止小智服务器
try:
await xiaozhi_server.stop()
except Exception as e:
shutdown_errors.append(("协议服务器", e))
logger.bind(tag=TAG).error(f"停止小智服务器失败: {e}")
# 停止全局GC管理器
await gc_manager.stop()
try:
await gc_manager.stop()
except Exception as e:
shutdown_errors.append(("GC管理器", e))
logger.bind(tag=TAG).error(f"停止GC管理器失败: {e}")
# 取消所有任务(关键修复点)
try:
await manage_api_http_close()
except Exception as e:
shutdown_errors.append(("管理API客户端", e))
logger.bind(tag=TAG).error(f"关闭管理API客户端失败: {e}")
# stdin 可立即取消;HTTP task 先获得窗口完成 runner.cleanup()。
stdin_task.cancel()
ws_task.cancel()
if ota_task:
ota_task.cancel()
await asyncio.gather(stdin_task, return_exceptions=True)
# 等待任务终止(必须加超时)
await asyncio.wait(
[stdin_task, ws_task, ota_task] if ota_task else [stdin_task, ws_task],
timeout=3.0,
return_when=asyncio.ALL_COMPLETED,
)
if ota_task:
done, pending = await asyncio.wait({ota_task}, timeout=3.0)
if pending:
ota_task.cancel()
await asyncio.gather(ota_task, return_exceptions=True)
elif done and not ota_task.cancelled():
try:
ota_task.result()
except Exception as e:
shutdown_errors.append((ota_task.get_name(), e))
logger.bind(tag=TAG).error(
f"后台任务退出失败: {ota_task.get_name()}: {e}"
)
# If task cancellation or its first cleanup attempt retained the
# runner, retry after the task has fully relinquished ownership.
if ota_server:
try:
await ota_server.stop()
except Exception as e:
shutdown_errors.append(("HTTP服务器重试清理", e))
logger.bind(tag=TAG).error(f"重试停止HTTP服务器失败: {e}")
print("服务器已关闭,程序退出。")
if shutdown_errors:
details = ", ".join(
f"{owner}: {error}" for owner, error in shutdown_errors
)
raise RuntimeError(f"服务器清理失败: {details}")
if __name__ == "__main__":
+64 -1
View File
@@ -35,12 +35,75 @@ server:
# 如果属于白名单内的设备,不校验token,直接放行
allowed_devices:
- "11:22:33:44:55:66"
# MQTT网关配置,用于通过OTA下发到设备,根据mqtt_gateway的.env文件配置,格式为host:port
# MQTT网关配置,用于通过OTA下发到设备,根据mqtt_gateway的.env文件配置,格式为host:port
mqtt_gateway: null
# MQTT签名密钥,用于生成MQTT连接密码,根据mqtt_gateway的.env文件配置
mqtt_signature_key: null
# UDP网关配置
udp_gateway: null
# #####################################################################################
# #############################协议配置(Protocol Configuration########################
# 支持WebSocket和MQTT两种协议,可以单独启用或同时启用
protocols:
# 启用的协议列表,可选值: ["websocket", "mqtt"]
enabled_protocols: ["websocket"] # 默认只启用WebSocket
# WebSocket协议开关
websocket_enabled: true
# MQTT协议开关
mqtt_enabled: false
# 同时驻留的本地共享ASR模型上限(包含公共模型)。Agent差异化ASR会按配置复用。
shared_asr_max_models: 3
# MQTT服务器配置(仅在mqtt_enabled为true时生效)
mqtt_server:
# 是否启用MQTT服务器
enabled: false
# MQTT服务器监听地址
host: 0.0.0.0
# MQTT服务器端口
port: 1883
# UDP音频传输端口(通常与MQTT端口相同)
udp_port: 1883
# UDP监听地址。留空时若public_endpoint是本机IPv4,会优先绑定该地址,
# 确保设备收到的UDP回包源地址与hello中声明的地址一致。
udp_bind_host: ""
# 公网/局域网可达地址(设备连接时使用)
public_endpoint: ""
# MQTT签名密钥(用于生成连接密码)
signature_key: ""
# 最大连接数
max_connections: 1000
# 尚未完成CONNECT认证的连接上限;重复clientId替换不占用新的活跃名额
max_pending_connections: 128
# 心跳检查间隔(秒)
heartbeat_interval: 30
# 最大消息载荷大小(字节)
max_payload_size: 8192
# MQTT应用控制消息队列上限
message_queue_size: 128
# Hello等待业务运行时就绪的最长时间(秒)
business_ready_timeout: 30
# 发送goodbye和释放物理连接的最长等待时间(秒)
goodbye_timeout: 1
close_timeout: 2
# MQTT协议使用说明:
# 1. 客户端ID格式:GID_test@@@mac_address@@@uuid 或 GID_test@@@mac_address
# 例如:GID_test@@@aa:bb:cc:dd:ee:ff@@@unique_uuid_123
# 2. 连接地址:mqtt://your.server.ip:1883
# 3. 音频传输:通过UDP加密传输,配置信息在hello消息中返回
# 4. 消息格式:JSON格式,支持hello、音频、文本等消息类型
#
# 启用MQTT的配置示例:
# protocols:
# enabled_protocols: ["websocket", "mqtt"] # 同时启用两种协议
# mqtt_enabled: true
# mqtt_server:
# enabled: true
# port: 1883
# public_endpoint: "your.server.ip" # 替换为实际IP
log:
# 设置控制台输出的日志格式,时间、日志级别、标签、消息
log_format: "<green>{time:YYMMDD HH:mm:ss}</green>[{version}_{selected_module}][<light-blue>{extra[tag]}</light-blue>]-<level>{level}</level>-<light-green>{message}</light-green>"
@@ -80,6 +80,28 @@ class ManageApiClient:
# 如果没有运行中的事件循环,创建一个临时的
raise Exception("必须在异步上下文中调用")
@classmethod
async def close_current_loop_client(cls):
"""Close the client owned by the running loop before that loop exits."""
import asyncio
loop_id = id(asyncio.get_running_loop())
client = cls._async_clients.pop(loop_id, None)
if client is not None:
await client.aclose()
@classmethod
async def close_all_clients(cls):
"""Close remaining clients after session/report workers have stopped."""
clients = list(cls._async_clients.values())
cls._async_clients.clear()
for client in clients:
try:
await client.aclose()
except Exception:
pass
cls._instance = None
@classmethod
async def _async_request(cls, method: str, endpoint: str, **kwargs) -> Dict:
"""发送单次异步HTTP请求并处理响应"""
@@ -266,3 +288,7 @@ def init_service(config):
def manage_api_http_safe_close():
ManageApiClient.safe_close()
async def manage_api_http_close():
await ManageApiClient.close_all_clients()
@@ -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:
+158
View File
@@ -10,6 +10,164 @@ class AuthenticationError(Exception):
pass
class AuthMiddleware:
"""
认证中间件
用于 WebSocket/MQTT 连接认证
集成 AuthManager 的 token 验证逻辑,支持多种认证方式
"""
def __init__(self, config: dict):
"""
初始化认证中间件
Args:
config: 配置字典,包含认证相关配置
"""
self.config = config
server_config = config.get("server", {})
auth_config = server_config.get("auth", {})
self.enabled = auth_config.get("enabled", False)
self.tokens = auth_config.get("tokens", [])
self.allowed_devices = set(auth_config.get("allowed_devices", []))
# 获取 auth_key 用于 HMAC token 验证
self.auth_key = server_config.get("auth_key", "")
expire_seconds = auth_config.get("expire_seconds", None)
# 创建 AuthManager 实例用于 HMAC token 验证
if self.auth_key:
self._auth_manager = AuthManager(
secret_key=self.auth_key,
expire_seconds=expire_seconds
)
else:
self._auth_manager = None
def authenticate(self, device_id: str, token: str = None, client_id: str = None) -> bool:
"""
验证设备认证(同步方法)
Args:
device_id: 设备 ID
token: 认证令牌(可以是静态 token 或 HMAC token
client_id: 客户端 ID(用于 HMAC token 验证)
Returns:
bool: 认证是否通过
"""
from config.logger import setup_logging
logger = setup_logging()
if not self.enabled:
logger.debug("[AuthMiddleware.authenticate] 认证未启用")
return True
# 1. 检查白名单
if device_id and device_id in self.allowed_devices:
logger.debug(f"[AuthMiddleware.authenticate] 设备 {device_id} 在白名单中,放行")
return True
# 2. 检查静态 token
if token:
# 移除 Bearer 前缀(如果有)
if token.startswith("Bearer "):
token = token[7:]
logger.debug(f"[AuthMiddleware.authenticate] 检查静态token, 配置了 {len(self.tokens)} 个token")
for token_config in self.tokens:
configured_token = token_config.get("token")
if configured_token == token:
logger.debug(f"[AuthMiddleware.authenticate] 静态token匹配成功")
return True
logger.debug(f"[AuthMiddleware.authenticate] 静态token不匹配")
# 3. 检查 HMAC token(需要 AuthManager
if token and self._auth_manager and client_id and device_id:
logger.debug(f"[AuthMiddleware.authenticate] 尝试HMAC token验证")
if self._auth_manager.verify_token(token, client_id, device_id):
logger.debug(f"[AuthMiddleware.authenticate] HMAC token验证成功")
return True
logger.debug(f"[AuthMiddleware.authenticate] HMAC token验证失败")
logger.debug(f"[AuthMiddleware.authenticate] 所有认证方式均失败")
return False
async def authenticate_async(self, headers: dict) -> bool:
"""
从 headers 中提取信息并进行异步认证
Args:
headers: HTTP 请求头字典
Returns:
bool: 认证是否通过
Raises:
AuthenticationError: 认证失败时抛出
"""
from config.logger import setup_logging
logger = setup_logging()
logger.debug(f"[AuthMiddleware] 开始认证, enabled={self.enabled}")
if not self.enabled:
logger.debug("[AuthMiddleware] 认证未启用,跳过认证")
return True
device_id = headers.get("device-id")
client_id = headers.get("client-id")
authorization = headers.get("authorization", "")
logger.debug(f"[AuthMiddleware] device_id={device_id}, client_id={client_id}, has_auth={bool(authorization)}")
# 提取 token
token = None
if authorization:
if authorization.startswith("Bearer "):
token = authorization[7:]
else:
token = authorization
# 执行认证
auth_result = self.authenticate(device_id, token, client_id)
logger.debug(f"[AuthMiddleware] 认证结果: {auth_result}")
if auth_result:
return True
raise AuthenticationError(f"认证失败: device_id={device_id}")
def authenticate_websocket(self, websocket) -> bool:
"""
WebSocket 连接认证
Args:
websocket: WebSocket 连接对象
Returns:
bool: 认证是否通过
"""
if not self.enabled:
return True
headers = dict(websocket.request.headers)
device_id = headers.get("device-id")
client_id = headers.get("client-id")
authorization = headers.get("authorization", "")
# 提取 token
token = None
if authorization:
if authorization.startswith("Bearer "):
token = authorization[7:]
else:
token = authorization
return self.authenticate(device_id, token, client_id)
class AuthManager:
"""
统一授权认证管理器
@@ -0,0 +1,116 @@
from typing import Any, Dict
from core.components.component_manager import Component, ComponentType, ComponentFactory
from core.utils import asr
from core.utils.modules_initialize import initialize_asr
from config.logger import setup_logging
logger = setup_logging()
TAG = __name__
class ASRAdapter(Component):
"""
ASR组件适配器:将现有ASR组件包装为新的组件接口
支持两种模式:
1. 共享实例模式:使用 SharedASRManager 的全局共享实例
2. 独立实例模式:每个连接创建独立的 ASR 实例(原有逻辑)
"""
def __init__(self, config: Dict[str, Any]):
super().__init__(ComponentType.ASR, config)
self._asr_instance = None
self._delete_audio = config.get("delete_audio", True)
self._using_shared = False # 是否使用共享实例
self._shared_owner = None
async def _do_initialize(self, context: Any) -> None:
"""初始化ASR组件"""
try:
# 获取ASR配置
selected_module = self.config.get("selected_module", {}).get("ASR")
if not selected_module:
raise ValueError("未配置ASR模块")
# 检查是否有全局共享 ASR 管理器
shared_manager = getattr(context, 'shared_asr_manager', None)
acquired_manager = None
if shared_manager and hasattr(shared_manager, "acquire_for_config"):
acquired_manager = await shared_manager.acquire_for_config(
self.config
)
elif (
shared_manager
and shared_manager.is_ready()
and shared_manager.matches_config(self.config)
):
acquired_manager = shared_manager
if acquired_manager is not None:
# 使用共享实例模式
logger.bind(tag=TAG).info(f"使用共享 ASR 实例: {selected_module}")
from core.providers.asr.shared_asr_proxy import SharedASRProxy
self._asr_instance = SharedASRProxy(acquired_manager)
self._using_shared = True
self._shared_owner = shared_manager
else:
# 使用独立实例模式(原有逻辑)
if shared_manager and shared_manager.is_ready():
logger.bind(tag=TAG).info(
f"连接ASR配置与共享实例不匹配,使用独立实例: {selected_module}"
)
logger.bind(tag=TAG).info(f"使用独立 ASR 实例: {selected_module}")
self._asr_instance = initialize_asr(self.config)
self._using_shared = False
# Legacy providers inspect conn.asr while handling stream state.
context.asr = self._asr_instance
# 打开音频通道(如果需要)
if hasattr(self._asr_instance, 'open_audio_channels'):
await self._asr_instance.open_audio_channels(context)
logger.bind(tag=TAG).info(
f"ASR组件初始化完成: {selected_module}, "
f"共享模式: {self._using_shared}"
)
except Exception as e:
logger.bind(tag=TAG).error(f"ASR组件初始化失败: {e}")
raise
async def _do_cleanup(self) -> None:
"""清理ASR组件"""
if self._asr_instance:
instance = self._asr_instance
if self._using_shared:
await self._close_resource(instance)
if self._shared_owner and hasattr(
self._shared_owner, "release_for_config"
):
await self._shared_owner.release_for_config(instance.manager)
logger.bind(tag=TAG).debug("共享 ASR 代理已释放")
else:
await self._close_resource(instance)
if hasattr(instance, 'cleanup_audio_files'):
instance.cleanup_audio_files()
logger.bind(tag=TAG).info("ASR组件清理完成")
self._asr_instance = None
self._using_shared = False
self._shared_owner = None
@property
def asr_instance(self):
"""获取ASR实例"""
return self._asr_instance
class ASRFactory(ComponentFactory):
"""ASR组件工厂"""
def create(self, config: Dict[str, Any]) -> Component:
return ASRAdapter(config)
def get_component_type(self) -> ComponentType:
return ComponentType.ASR
@@ -0,0 +1,92 @@
from typing import Any, Dict
from core.components.component_manager import Component, ComponentType, ComponentFactory
from core.utils import intent, llm
from config.logger import setup_logging
logger = setup_logging()
class IntentAdapter(Component):
"""Intent组件适配器:将现有Intent组件包装为新的组件接口"""
def __init__(self, config: Dict[str, Any]):
super().__init__(ComponentType.INTENT, config)
self._intent_instance = None
self._owned_llm_instance = None
async def _do_initialize(self, context: Any) -> None:
"""初始化Intent组件"""
try:
# 获取Intent配置
selected_module = self.config.get("selected_module", {}).get("Intent")
if not selected_module:
raise ValueError("未配置Intent模块")
# 获取Intent类型
intent_type = (
selected_module
if "type" not in self.config["Intent"][selected_module]
else self.config["Intent"][selected_module]["type"]
)
# 创建Intent实例
self._intent_instance = intent.create_instance(
intent_type,
self.config["Intent"][selected_module],
)
# intent_llm 可配置独立模型;未配置时回退主 LLM。
if intent_type == "intent_llm":
llm_component = None
if getattr(context, "component_manager", None):
llm_component = await context.component_manager.get_component(ComponentType.LLM, context)
main_llm = getattr(llm_component, "llm_instance", None)
intent_config = self.config["Intent"][selected_module]
dedicated_llm_name = intent_config.get("llm")
selected_llm = main_llm
if (
dedicated_llm_name
and dedicated_llm_name in self.config.get("LLM", {})
):
dedicated_config = self.config["LLM"][dedicated_llm_name]
dedicated_type = dedicated_config.get("type", dedicated_llm_name)
self._owned_llm_instance = llm.create_instance(
dedicated_type, dedicated_config
)
selected_llm = self._owned_llm_instance
logger.info(f"为意图识别创建专用LLM: {dedicated_llm_name}")
if selected_llm and hasattr(self._intent_instance, "set_llm"):
self._intent_instance.set_llm(selected_llm)
logger.info(f"Intent组件初始化完成: {intent_type}")
except Exception as e:
logger.error(f"Intent组件初始化失败: {e}")
raise
async def _do_cleanup(self) -> None:
"""清理Intent组件"""
if self._intent_instance:
instance = self._intent_instance
await self._close_resource(instance)
self._intent_instance = None
logger.info("Intent组件清理完成")
if self._owned_llm_instance:
instance = self._owned_llm_instance
await self._close_resource(instance)
self._owned_llm_instance = None
@property
def intent_instance(self):
"""获取Intent实例"""
return self._intent_instance
class IntentFactory(ComponentFactory):
"""Intent组件工厂"""
def create(self, config: Dict[str, Any]) -> Component:
return IntentAdapter(config)
def get_component_type(self) -> ComponentType:
return ComponentType.INTENT
@@ -0,0 +1,64 @@
from typing import Any, Dict
from core.components.component_manager import Component, ComponentType, ComponentFactory
from core.utils import llm
from config.logger import setup_logging
logger = setup_logging()
class LLMAdapter(Component):
"""LLM组件适配器:将现有LLM组件包装为新的组件接口"""
def __init__(self, config: Dict[str, Any]):
super().__init__(ComponentType.LLM, config)
self._llm_instance = None
async def _do_initialize(self, context: Any) -> None:
"""初始化LLM组件"""
try:
# 获取LLM配置
selected_module = self.config.get("selected_module", {}).get("LLM")
if not selected_module:
raise ValueError("未配置LLM模块")
# 获取LLM类型
llm_type = (
selected_module
if "type" not in self.config["LLM"][selected_module]
else self.config["LLM"][selected_module]["type"]
)
# 创建LLM实例
self._llm_instance = llm.create_instance(
llm_type,
self.config["LLM"][selected_module],
)
logger.info(f"LLM组件初始化完成: {llm_type}")
except Exception as e:
logger.error(f"LLM组件初始化失败: {e}")
raise
async def _do_cleanup(self) -> None:
"""清理LLM组件"""
if self._llm_instance:
instance = self._llm_instance
await self._close_resource(instance)
self._llm_instance = None
logger.info("LLM组件清理完成")
@property
def llm_instance(self):
"""获取LLM实例"""
return self._llm_instance
class LLMFactory(ComponentFactory):
"""LLM组件工厂"""
def create(self, config: Dict[str, Any]) -> Component:
return LLMAdapter(config)
def get_component_type(self) -> ComponentType:
return ComponentType.LLM
@@ -0,0 +1,103 @@
from typing import Any, Dict
from core.components.component_manager import Component, ComponentType, ComponentFactory
from core.utils import llm, memory
from config.logger import setup_logging
logger = setup_logging()
class MemoryAdapter(Component):
"""Memory组件适配器:将现有Memory组件包装为新的组件接口"""
def __init__(self, config: Dict[str, Any]):
super().__init__(ComponentType.MEMORY, config)
self._memory_instance = None
self._owned_llm_instance = None
async def _do_initialize(self, context: Any) -> None:
"""初始化Memory组件"""
try:
# 获取Memory配置
selected_module = self.config.get("selected_module", {}).get("Memory")
if not selected_module:
raise ValueError("未配置Memory模块")
# 获取Memory类型
memory_type = (
selected_module
if "type" not in self.config["Memory"][selected_module]
else self.config["Memory"][selected_module]["type"]
)
# 创建Memory实例
self._memory_instance = memory.create_instance(
memory_type,
self.config["Memory"][selected_module],
self.config.get("summaryMemory", None),
)
# 初始化记忆模块
main_llm = None
if hasattr(self._memory_instance, 'init_memory'):
# 需要LLM实例来初始化记忆
llm_component = None
if getattr(context, "component_manager", None):
llm_component = await context.component_manager.get_component(ComponentType.LLM, context)
if llm_component and hasattr(llm_component, 'llm_instance'):
main_llm = llm_component.llm_instance
self._memory_instance.init_memory(
role_id=context.device_id,
llm=main_llm,
summary_memory=self.config.get("summaryMemory", None),
save_to_file=not self.config.get("read_config_from_api", False),
)
memory_config = self.config["Memory"][selected_module]
dedicated_llm_name = memory_config.get("llm")
if (
memory_type == "mem_local_short"
and dedicated_llm_name
and dedicated_llm_name in self.config.get("LLM", {})
):
dedicated_config = self.config["LLM"][dedicated_llm_name]
dedicated_type = dedicated_config.get("type", dedicated_llm_name)
self._owned_llm_instance = llm.create_instance(
dedicated_type, dedicated_config
)
self._memory_instance.set_llm(self._owned_llm_instance)
logger.info(f"为记忆总结创建专用LLM: {dedicated_llm_name}")
elif main_llm and hasattr(self._memory_instance, "set_llm"):
self._memory_instance.set_llm(main_llm)
logger.info(f"Memory组件初始化完成: {memory_type}")
except Exception as e:
logger.error(f"Memory组件初始化失败: {e}")
raise
async def _do_cleanup(self) -> None:
"""清理Memory组件"""
if self._memory_instance:
instance = self._memory_instance
await self._close_resource(instance)
self._memory_instance = None
logger.info("Memory组件清理完成")
if self._owned_llm_instance:
instance = self._owned_llm_instance
await self._close_resource(instance)
self._owned_llm_instance = None
@property
def memory_instance(self):
"""获取Memory实例"""
return self._memory_instance
class MemoryFactory(ComponentFactory):
"""Memory组件工厂"""
def create(self, config: Dict[str, Any]) -> Component:
return MemoryAdapter(config)
def get_component_type(self) -> ComponentType:
return ComponentType.MEMORY
@@ -0,0 +1,76 @@
from typing import Any, Dict
from core.components.component_manager import Component, ComponentType, ComponentFactory
from core.utils import tts
from config.logger import setup_logging
logger = setup_logging()
class TTSAdapter(Component):
"""TTS组件适配器:将现有TTS组件包装为新的组件接口"""
def __init__(self, config: Dict[str, Any]):
super().__init__(ComponentType.TTS, config)
self._tts_instance = None
self._delete_audio = config.get("delete_audio", True)
async def _do_initialize(self, context: Any) -> None:
"""初始化TTS组件"""
try:
# 获取TTS配置
selected_module = self.config.get("selected_module", {}).get("TTS")
if not selected_module:
raise ValueError("未配置TTS模块")
# 获取TTS类型
tts_type = (
selected_module
if "type" not in self.config["TTS"][selected_module]
else self.config["TTS"][selected_module]["type"]
)
# 创建TTS实例
self._tts_instance = tts.create_instance(
tts_type,
self.config["TTS"][selected_module],
str(self._delete_audio).lower() in ("true", "1", "yes"),
)
# 打开音频通道
if hasattr(self._tts_instance, 'open_audio_channels'):
await self._tts_instance.open_audio_channels(context)
# 设置兼容属性(用于向后兼容)
if hasattr(context, 'tts'):
context.tts = self._tts_instance
logger.info(f"TTS组件初始化完成: {tts_type}")
except Exception as e:
logger.error(f"TTS组件初始化失败: {e}")
raise
async def _do_cleanup(self) -> None:
"""清理TTS组件"""
if self._tts_instance:
instance = self._tts_instance
await self._close_resource(instance)
if hasattr(instance, 'cleanup_audio_files'):
instance.cleanup_audio_files()
self._tts_instance = None
logger.info("TTS组件清理完成")
@property
def tts_instance(self):
"""获取TTS实例"""
return self._tts_instance
class TTSFactory(ComponentFactory):
"""TTS组件工厂"""
def create(self, config: Dict[str, Any]) -> Component:
return TTSAdapter(config)
def get_component_type(self) -> ComponentType:
return ComponentType.TTS
@@ -0,0 +1,64 @@
from typing import Any, Dict
from core.components.component_manager import Component, ComponentType, ComponentFactory
from core.utils import vad
from config.logger import setup_logging
logger = setup_logging()
class VADAdapter(Component):
"""VAD组件适配器:将现有VAD组件包装为新的组件接口"""
def __init__(self, config: Dict[str, Any]):
super().__init__(ComponentType.VAD, config)
self._vad_instance = None
async def _do_initialize(self, context: Any) -> None:
"""初始化VAD组件"""
try:
# 获取VAD配置
selected_module = self.config.get("selected_module", {}).get("VAD")
if not selected_module:
raise ValueError("未配置VAD模块")
# 获取VAD类型
vad_type = (
selected_module
if "type" not in self.config["VAD"][selected_module]
else self.config["VAD"][selected_module]["type"]
)
# 创建VAD实例
self._vad_instance = vad.create_instance(
vad_type,
self.config["VAD"][selected_module],
)
logger.info(f"VAD组件初始化完成: {vad_type}")
except Exception as e:
logger.error(f"VAD组件初始化失败: {e}")
raise
async def _do_cleanup(self) -> None:
"""清理VAD组件"""
if self._vad_instance:
instance = self._vad_instance
await self._close_resource(instance)
self._vad_instance = None
logger.info("VAD组件清理完成")
@property
def vad_instance(self):
"""获取VAD实例"""
return self._vad_instance
class VADFactory(ComponentFactory):
"""VAD组件工厂"""
def create(self, config: Dict[str, Any]) -> Component:
return VADAdapter(config)
def get_component_type(self) -> ComponentType:
return ComponentType.VAD
@@ -0,0 +1,268 @@
import asyncio
import inspect
import weakref
from abc import ABC, abstractmethod
from typing import Any, Callable, Dict, Optional, Type, TypeVar, Generic
from enum import Enum
from config.logger import setup_logging
T = TypeVar('T')
logger = setup_logging()
class ComponentType(Enum):
"""组件类型枚举"""
TTS = "tts"
ASR = "asr"
VAD = "vad"
LLM = "llm"
MEMORY = "memory"
INTENT = "intent"
class ComponentState(Enum):
"""组件状态枚举"""
UNINITIALIZED = "uninitialized"
INITIALIZING = "initializing"
READY = "ready"
ERROR = "error"
CLEANING = "cleaning"
CLEANED = "cleaned"
class Component(ABC):
"""组件基类:定义统一的组件接口和生命周期管理"""
def __init__(self, component_type: ComponentType, config: Dict[str, Any]):
self.component_type = component_type
self.config = config
self.state = ComponentState.UNINITIALIZED
self._initialization_lock = asyncio.Lock()
self._cleanup_lock = asyncio.Lock()
self._dependencies: Dict[str, 'Component'] = {}
self._dependents: weakref.WeakSet['Component'] = weakref.WeakSet()
self._resources: list = [] # 存储需要清理的资源
@abstractmethod
async def _do_initialize(self, context: Any) -> None:
"""子类实现具体的初始化逻辑"""
pass
@abstractmethod
async def _do_cleanup(self) -> None:
"""子类实现具体的清理逻辑"""
pass
async def initialize(self, context: Any) -> None:
"""初始化组件(带锁保护)"""
async with self._initialization_lock:
if self.state != ComponentState.UNINITIALIZED:
return
try:
self.state = ComponentState.INITIALIZING
logger.info(f"正在初始化组件: {self.component_type.value}")
# 初始化依赖组件
await self._initialize_dependencies(context)
# 执行具体初始化
await self._do_initialize(context)
self.state = ComponentState.READY
logger.info(f"组件初始化完成: {self.component_type.value}")
except Exception as e:
self.state = ComponentState.ERROR
logger.error(f"组件初始化失败: {self.component_type.value}, 错误: {e}")
# 初始化可能已经分配线程、连接或文件,失败路径也必须释放。
try:
await self._do_cleanup()
await self._cleanup_resources()
except Exception as cleanup_error:
logger.error(
f"组件初始化回滚失败: {self.component_type.value}, "
f"错误: {cleanup_error}"
)
raise
async def cleanup(self) -> None:
"""清理组件(带锁保护)"""
async with self._cleanup_lock:
if self.state in [ComponentState.CLEANING, ComponentState.CLEANED]:
return
try:
self.state = ComponentState.CLEANING
logger.info(f"正在清理组件: {self.component_type.value}")
# 清理依赖此组件的其他组件
await self._cleanup_dependents()
# 执行具体清理
await self._do_cleanup()
# 清理资源
await self._cleanup_resources()
self.state = ComponentState.CLEANED
logger.info(f"组件清理完成: {self.component_type.value}")
except Exception as e:
logger.error(f"组件清理失败: {self.component_type.value}, 错误: {e}")
self.state = ComponentState.ERROR
raise
def add_dependency(self, name: str, component: 'Component') -> None:
"""添加依赖组件"""
self._dependencies[name] = component
component._dependents.add(self)
def add_resource(self, resource: Any) -> None:
"""添加需要清理的资源"""
self._resources.append(resource)
async def _initialize_dependencies(self, context: Any) -> None:
"""初始化依赖组件"""
for name, dep in self._dependencies.items():
if dep.state == ComponentState.UNINITIALIZED:
await dep.initialize(context)
async def _cleanup_dependents(self) -> None:
"""清理依赖此组件的其他组件"""
for dependent in list(self._dependents):
await dependent.cleanup()
async def _cleanup_resources(self) -> None:
"""清理所有注册的资源"""
for resource in self._resources:
try:
await self._close_resource(resource)
except Exception as e:
logger.warning(f"清理资源时出错: {e}")
self._resources.clear()
@staticmethod
async def _close_resource(resource: Any) -> None:
"""Close either sync or async providers without invoking them twice."""
cleanup = getattr(resource, "close", None) or getattr(resource, "cleanup", None)
if not cleanup:
return
result = cleanup()
if inspect.isawaitable(result):
await result
class ComponentFactory(ABC):
"""组件工厂基类"""
@abstractmethod
def create(self, config: Dict[str, Any]) -> Component:
"""创建组件实例"""
pass
@abstractmethod
def get_component_type(self) -> ComponentType:
"""获取组件类型"""
pass
class ComponentManager:
"""
组件管理器:统一管理连接期内的组件实例生命周期。
支持分类管理、依赖注入、按需懒加载与统一清理。
"""
def __init__(self, config: Dict[str, Any]):
self._config = config
self._components: Dict[str, Component] = {}
self._factories: Dict[ComponentType, ComponentFactory] = {}
self._component_locks: Dict[str, asyncio.Lock] = {}
self._initialization_order: list[ComponentType] = []
self.lazy_enabled = True
def register_factory(self, factory: ComponentFactory) -> None:
"""注册组件工厂"""
component_type = factory.get_component_type()
self._factories[component_type] = factory
logger.debug(f"注册组件工厂: {component_type.value}")
def set_initialization_order(self, order: list[ComponentType]) -> None:
"""设置组件初始化顺序"""
self._initialization_order = order
async def get_component(self, component_type: ComponentType, context: Any) -> Optional[Component]:
"""获取组件实例(按需创建)"""
key = component_type.value
lock = self._component_locks.setdefault(key, asyncio.Lock())
async with lock:
if key in self._components:
return self._components[key]
if not self.lazy_enabled:
logger.warning(f"组件未初始化且已关闭懒加载: {component_type.value}")
return None
factory = self._factories.get(component_type)
if factory is None:
logger.warning(f"未找到组件工厂: {component_type.value}")
return None
try:
instance = factory.create(self._config)
# 先登记所有权,确保初始化中途失败时资源仍然可达。
self._components[key] = instance
await instance.initialize(context)
logger.info(f"组件创建并初始化完成: {component_type.value}")
except Exception as e:
logger.error(f"组件创建失败: {component_type.value}, 错误: {e}")
failed = self._components.pop(key, None)
if failed is not None:
await failed.cleanup()
return None
return self._components.get(key)
def get(self, component_name: str) -> Optional[Component]:
"""获取已初始化的组件实例(兼容接口)"""
return self._components.get(component_name)
async def initialize_all(self, context: Any) -> None:
"""按顺序初始化所有组件"""
for component_type in self._initialization_order:
await self.get_component(component_type, context)
async def cleanup_all(self) -> None:
"""清理所有组件(逆序清理)"""
errors = []
# 按逆序清理,确保依赖关系正确
for component_type in reversed(self._initialization_order):
key = component_type.value
if key in self._components:
component = self._components.pop(key)
try:
await component.cleanup()
except Exception as e:
errors.append((key, e))
# 清理可能遗漏的组件
remaining_components = list(self._components.items())
for key, component in remaining_components:
try:
await component.cleanup()
except Exception as e:
errors.append((key, e))
self._components.clear()
self._component_locks.clear()
logger.info("所有组件已清理完成")
if errors:
details = ", ".join(f"{name}: {error}" for name, error in errors)
raise RuntimeError(f"部分组件清理失败: {details}")
def get_component_status(self) -> Dict[str, str]:
"""获取所有组件状态"""
return {
name: component.state.value
for name, component in self._components.items()
}
@@ -0,0 +1,73 @@
from typing import Dict, Any
from core.components.component_manager import ComponentManager, ComponentType
from core.components.adapters.tts_adapter import TTSFactory
from core.components.adapters.asr_adapter import ASRFactory
from core.components.adapters.vad_adapter import VADFactory
from core.components.adapters.llm_adapter import LLMFactory
from core.components.adapters.memory_adapter import MemoryFactory
from core.components.adapters.intent_adapter import IntentFactory
from config.logger import setup_logging
logger = setup_logging()
class ComponentRegistry:
"""组件注册器:统一管理所有组件工厂的注册"""
_factories_registered = False
@classmethod
def create_component_manager(cls, config: Dict[str, Any]) -> ComponentManager:
"""创建并配置组件管理器"""
manager = ComponentManager(config)
# 只在第一次时记录注册日志
if not cls._factories_registered:
logger.info("注册组件工厂")
cls._factories_registered = True
# 注册所有组件工厂(每个manager都需要注册,但不重复记录日志)
manager.register_factory(TTSFactory())
manager.register_factory(ASRFactory())
manager.register_factory(VADFactory())
manager.register_factory(LLMFactory())
manager.register_factory(MemoryFactory())
manager.register_factory(IntentFactory())
# 设置组件初始化顺序(考虑依赖关系)
# VAD -> ASR -> LLM -> Memory/Intent -> TTS
initialization_order = [
ComponentType.VAD,
ComponentType.ASR,
ComponentType.LLM,
ComponentType.MEMORY,
ComponentType.INTENT,
ComponentType.TTS,
]
manager.set_initialization_order(initialization_order)
if not cls._factories_registered:
logger.info("组件管理器创建完成,已注册所有组件工厂")
return manager
@staticmethod
def get_required_components(config: Dict[str, Any]) -> list[ComponentType]:
"""根据配置获取需要的组件类型"""
required = []
selected_modules = config.get("selected_module", {})
if selected_modules.get("VAD"):
required.append(ComponentType.VAD)
if selected_modules.get("ASR"):
required.append(ComponentType.ASR)
if selected_modules.get("LLM"):
required.append(ComponentType.LLM)
if selected_modules.get("TTS"):
required.append(ComponentType.TTS)
if selected_modules.get("Memory"):
required.append(ComponentType.MEMORY)
if selected_modules.get("Intent"):
required.append(ComponentType.INTENT)
return required
@@ -0,0 +1,505 @@
import copy
import uuid
import time
import queue
import asyncio
import threading
from dataclasses import dataclass, field
from typing import Any, Dict, Optional, List, Callable, Awaitable, Union
from collections import deque
from concurrent.futures import ThreadPoolExecutor
from core.utils.dialogue import Dialogue
from core.auth import AuthMiddleware
from core.utils.prompt_manager import PromptManager
from core.utils.voiceprint_provider import VoiceprintProvider
from config.logger import setup_logging
@dataclass
class SessionContext:
"""
会话上下文:完全替换ConnectionHandler的所有功能
承载单连接生命周期内的状态、组件、资源管理
与传输层解耦,支持WebSocket/MQTT/UDP等多协议
"""
# === 基础标识 ===
session_id: str = field(default_factory=lambda: str(uuid.uuid4()))
device_id: Optional[str] = None
client_ip: Optional[str] = None
headers: Dict[str, str] = field(default_factory=dict)
# === 配置管理 ===
config: Dict[str, Any] = field(default_factory=dict)
common_config: Dict[str, Any] = field(default_factory=dict)
private_config: Dict[str, Any] = field(default_factory=dict)
selected_module_str: str = ""
# === 认证与绑定 ===
is_authenticated: bool = False
need_bind: bool = False
bind_code: Optional[str] = None
read_config_from_api: bool = False
max_output_size: int = 0
chat_history_conf: int = 0
# === 会话状态 ===
is_speaking: bool = False
listen_mode: str = "auto"
abort_requested: bool = False
close_after_chat: bool = False
conversation_active: bool = False
last_finalized_session_id: Optional[str] = None
just_woken_up: bool = False
calling: bool = False
incoming_call: Optional[Any] = None
load_function_plugin: bool = False
intent_type: str = "nointent"
# === 音频相关 ===
audio_format: str = "opus"
# sample_rate/channels/frame_duration remain the output-side compatibility
# fields consumed by existing TTS providers.
sample_rate: int = 24000
channels: int = 1
frame_duration: int = 60
input_sample_rate: int = 16000
input_channels: int = 1
input_frame_duration: int = 60
output_sample_rate: int = 24000
output_channels: int = 1
output_frame_duration: int = 60
client_aec: bool = False
# Native MQTT control and UDP audio use different channels. Keep UDP
# closed until the first listen/start of a logical session. It stays open
# through the bounded listen/stop grace period so queued tail frames are
# included in the same ASR turn.
accepting_input_audio: bool = False
wake_audio_suppression_task: Optional[asyncio.Task] = None
listen_stop_pending: bool = False
listen_stop_task: Optional[asyncio.Task] = None
listen_stop_deadline: float = 0.0
listen_start_task: Optional[asyncio.Task] = None
client_have_voice: bool = False
client_voice_stop: bool = False
client_audio_buffer: bytearray = field(default_factory=bytearray)
client_voice_window: deque = field(default_factory=lambda: deque(maxlen=5))
last_is_voice: bool = False
audio_flow_control: Dict[str, Any] = field(default_factory=dict)
aec_audio_cache: Dict[int, bytes] = field(default_factory=dict)
aec_audio_cache_time: Dict[int, float] = field(default_factory=dict)
# === ASR相关 ===
asr_audio: List[bytes] = field(default_factory=list)
asr_audio_queue: queue.Queue = field(default_factory=queue.Queue)
asr_priority_thread: Optional[threading.Thread] = None
asr_result_handler: Optional[Callable[[str, List[bytes]], Awaitable[None]]] = None
uses_pipeline_runtime: bool = True
# === LLM相关 ===
llm_finish_task: bool = True
dialogue: Optional[Dialogue] = None
current_speaker: Optional[str] = None
introduced_speakers: set = field(default_factory=set)
system_introduced_speakers: set = field(default_factory=set)
sentence_id: Optional[str] = None
# === TTS相关 ===
tts_MessageText: str = ""
# === IoT相关 ===
iot_descriptors: Dict[str, Any] = field(default_factory=dict)
func_handler: Optional[Any] = None
# === 时间管理 ===
last_activity_time_ms: float = field(default_factory=lambda: time.time() * 1000)
created_at: float = field(default_factory=lambda: time.time())
timeout_seconds: int = 180 # 默认超时时间
timeout_task: Optional[asyncio.Task] = None
# === 组件实例 ===
# components属性通过@property方法提供,指向component_manager
# === 其他状态 ===
welcome_msg: Optional[Dict[str, Any]] = None
prompt: Optional[str] = None
features: Optional[Dict[str, Any]] = None
mcp_client: Optional[Any] = None
cmd_exit: List[str] = field(default_factory=list)
# === 线程与并发 ===
loop: Optional[asyncio.AbstractEventLoop] = None
stop_event: Optional[threading.Event] = None
executor: Optional[ThreadPoolExecutor] = None
conversation_tasks: set = field(default_factory=set)
turn_tasks: set = field(default_factory=set)
background_tasks: set = field(default_factory=set)
# === 队列管理 ===
report_queue: queue.Queue = field(
default_factory=lambda: queue.Queue(maxsize=64)
)
report_thread: Optional[threading.Thread] = None
report_asr_enable: bool = False
report_tts_enable: bool = False
# === 组件管理器 ===
component_manager: Optional[Any] = None
# === 兼容属性(用于向后兼容TTS处理) ===
tts: Optional[Any] = None
websocket: Optional[Any] = None # 兼容旧TTS组件
transport: Optional[Any] = None # 新的transport接口
# === 工具类 ===
auth: Optional[AuthMiddleware] = None
prompt_manager: Optional[PromptManager] = None
voiceprint_provider: Optional[VoiceprintProvider] = None
server: Optional[Any] = None # WebSocket服务器引用
# === 会话级清理回调 ===
_cleanup_callbacks: List[Callable[[], Union[None, Awaitable[None]]]] = field(default_factory=list)
def __post_init__(self):
"""初始化后处理"""
# 深拷贝配置避免污染
if self.config:
self.common_config = self.config
self.config = copy.deepcopy(self.config)
# 从配置中读取相关设置
self.read_config_from_api = self.config.get("read_config_from_api", False)
self.max_output_size = self.config.get("max_output_size", 0)
self.chat_history_conf = self.config.get("chat_history_conf", 0)
self.cmd_exit = self.config.get("exit_commands", [])
self.timeout_seconds = int(self.config.get("close_connection_no_voice_time", 120)) + 60
# 初始化认证中间件
self.auth = AuthMiddleware(self.config)
# 初始化提示词管理器
self.prompt_manager = PromptManager(self.config, setup_logging())
# 初始化对话管理
if not self.dialogue:
self.dialogue = Dialogue()
# 初始化线程相关
if not self.loop:
try:
self.loop = asyncio.get_event_loop()
except RuntimeError:
self.loop = asyncio.new_event_loop()
if not self.stop_event:
self.stop_event = threading.Event()
if not self.executor:
self.executor = ThreadPoolExecutor(max_workers=5)
# 初始化上报设置
self.report_asr_enable = self.read_config_from_api
self.report_tts_enable = self.read_config_from_api
def update_activity(self) -> None:
"""刷新最后活跃时间"""
self.last_activity_time_ms = time.time() * 1000
def clearSpeakStatus(self) -> None:
"""清除服务端讲话状态(兼容方法)"""
self.is_speaking = False
logger = setup_logging()
logger.debug("清除服务端讲话状态")
def change_system_prompt(self, prompt: str) -> None:
"""Update the active role prompt for legacy function plugins."""
self.prompt = prompt
self.dialogue.update_system_message(prompt)
def reset_vad_states(self) -> None:
"""重置VAD状态(兼容方法)"""
self.client_audio_buffer = bytearray()
self.client_have_voice = False
self.client_voice_stop = False
self.last_is_voice = False
self.client_voice_window.clear()
logger = setup_logging()
logger.debug("VAD states reset.")
def reset_audio_states(self) -> None:
"""Reset all turn-scoped input state expected by legacy ASR providers."""
self.asr_audio.clear()
self.listen_stop_pending = False
self.reset_vad_states()
def is_timeout(self, timeout_seconds: int) -> bool:
"""检查是否超时"""
now_ms = time.time() * 1000
return (now_ms - self.last_activity_time_ms) > (timeout_seconds * 1000)
def register_cleanup(self, callback: Callable[[], Union[None, Awaitable[None]]]) -> None:
"""注册会话结束时需要执行的清理回调"""
self._cleanup_callbacks.append(callback)
def unregister_cleanup(
self, callback: Callable[[], Union[None, Awaitable[None]]]
) -> None:
"""Remove a callback when its connection-scoped owner is replaced."""
self._cleanup_callbacks = [
registered
for registered in self._cleanup_callbacks
if registered != callback
]
def create_background_task(
self,
coroutine,
*,
conversation_scoped: bool = True,
turn_scoped: bool = False,
) -> asyncio.Task:
"""Create a tracked task with an explicit turn/session/connection owner."""
task = asyncio.create_task(coroutine)
if turn_scoped:
owner = self.turn_tasks
elif conversation_scoped:
owner = self.conversation_tasks
else:
owner = self.background_tasks
owner.add(task)
task.add_done_callback(owner.discard)
return task
async def cancel_turn_tasks(self) -> None:
"""Cancel producers belonging only to the active speech/chat turn."""
current_task = asyncio.current_task()
pending_tasks = [
task
for task in list(self.turn_tasks)
if task is not current_task and not task.done()
]
for task in pending_tasks:
task.cancel()
if pending_tasks:
await asyncio.gather(*pending_tasks, return_exceptions=True)
self.turn_tasks = {
task for task in self.turn_tasks if task is current_task
}
async def cancel_conversation_tasks(self) -> None:
"""Cancel all producers owned by the current logical conversation."""
await self.cancel_turn_tasks()
current_task = asyncio.current_task()
while True:
pending_tasks = [
task
for task in list(self.conversation_tasks)
if task is not current_task and not task.done()
]
if not pending_tasks:
break
for task in pending_tasks:
task.cancel()
await asyncio.gather(*pending_tasks, return_exceptions=True)
self.conversation_tasks = {
task for task in self.conversation_tasks if task is current_task
}
async def run_cleanup(self) -> None:
"""执行所有注册的清理回调"""
logger = setup_logging()
logger.info(f"Session {self.session_id} 开始执行会话级清理 ({len(self._cleanup_callbacks)} 个回调)")
# 停止所有线程
if self.stop_event:
self.stop_event.set()
# 取消超时任务
if self.timeout_task and not self.timeout_task.done():
self.timeout_task.cancel()
await self.cancel_conversation_tasks()
current_task = asyncio.current_task()
pending_tasks = [
task
for task in list(self.background_tasks)
if task is not current_task and not task.done()
]
for task in pending_tasks:
task.cancel()
if pending_tasks:
await asyncio.gather(*pending_tasks, return_exceptions=True)
self.background_tasks.clear()
# 执行清理回调
for callback in reversed(self._cleanup_callbacks):
try:
result = callback()
if asyncio.iscoroutine(result):
await result
except Exception as e:
logger.error(f"Session {self.session_id} 清理回调执行失败: {e}", exc_info=True)
self._cleanup_callbacks.clear()
# 所有异步生产者退出后再关闭线程池,避免任务继续提交工作。
if self.executor:
self.executor.shutdown(wait=False)
self.executor = None
logger.info(f"Session {self.session_id} 会话级清理完成")
# === 兼容旧代码的属性访问 ===
@property
def client_is_speaking(self) -> bool:
"""兼容旧代码的属性名"""
return self.is_speaking
@client_is_speaking.setter
def client_is_speaking(self, value: bool):
self.is_speaking = value
@property
def client_listen_mode(self) -> str:
"""兼容旧代码的属性名"""
return self.listen_mode
@client_listen_mode.setter
def client_listen_mode(self, value: str):
self.listen_mode = value
@property
def client_abort(self) -> bool:
"""兼容旧代码的属性名"""
return self.abort_requested
@client_abort.setter
def client_abort(self, value: bool):
self.abort_requested = value
@property
def components(self):
"""组件访问器(兼容属性)"""
return self.component_manager
@components.setter
def components(self, value):
"""组件设置器(兼容属性)- 实际设置到component_manager"""
# 如果尝试设置components,我们忽略它或者给出警告
# 因为components应该通过component_manager管理
logger = setup_logging()
logger.warning("尝试直接设置components属性,请使用component_manager")
@property
def last_activity_time(self) -> float:
"""兼容旧代码:返回毫秒级时间戳"""
return self.last_activity_time_ms
@last_activity_time.setter
def last_activity_time(self, value: float):
"""兼容旧代码:接受毫秒级时间戳"""
self.last_activity_time_ms = value
# === 日志相关 ===
@property
def logger(self):
"""获取日志记录器"""
return setup_logging()
# === 工具方法 ===
def get_component(self, component_name: str) -> Optional[Any]:
"""获取组件实例"""
return self.components.get(component_name)
def set_component(self, component_name: str, component_instance: Any) -> None:
"""设置组件实例"""
if self.component_manager:
self.component_manager._components[component_name] = component_instance
def has_component(self, component_name: str) -> bool:
"""检查是否有指定组件"""
return component_name in self.components
def clear_audio_buffer(self) -> None:
"""清空音频缓冲区"""
self.client_audio_buffer.clear()
self.asr_audio.clear()
# 清空队列
try:
while not self.asr_audio_queue.empty():
self.asr_audio_queue.get_nowait()
except queue.Empty:
pass
def reset_voice_state(self) -> None:
"""重置语音状态"""
self.client_have_voice = False
self.client_voice_stop = False
self.last_is_voice = False
self.client_voice_window.clear()
def initialize_private_config(self) -> None:
"""初始化差异化配置(从ConnectionHandler迁移)"""
from config.config_loader import get_private_config_from_api
from config.manage_api_client import DeviceNotFoundException, DeviceBindException
if not self.read_config_from_api:
return
try:
# 获取设备私有配置
private_config = get_private_config_from_api(
self.config, self.device_id, self.headers.get("client-id")
)
if private_config:
self.private_config = private_config
# 合并私有配置到主配置
self.config.update(private_config)
except DeviceNotFoundException:
self.logger.error(f"设备 {self.device_id} 未找到")
self.need_bind = True
except DeviceBindException as e:
self.logger.error(f"设备绑定异常: {e}")
self.need_bind = True
self.bind_code = str(e)
except Exception as e:
self.logger.error(f"获取私有配置失败: {e}")
async def initialize_components(self) -> None:
"""异步初始化组件(从ConnectionHandler迁移)"""
if not self.component_manager:
return
try:
# 初始化各个组件
from core.components.component_registry import ComponentType
# 按依赖顺序初始化组件
component_types = [
ComponentType.VAD,
ComponentType.ASR,
ComponentType.LLM,
ComponentType.MEMORY,
ComponentType.INTENT,
ComponentType.TTS
]
for component_type in component_types:
try:
component = await self.component_manager.get_component(component_type, self)
if component:
self.logger.info(f"组件 {component_type} 初始化成功")
except Exception as e:
self.logger.error(f"组件 {component_type} 初始化失败: {e}")
except Exception as e:
self.logger.error(f"组件初始化失败: {e}")
def __str__(self) -> str:
return f"SessionContext(session_id={self.session_id}, device_id={self.device_id})"
def __repr__(self) -> str:
return self.__str__()
+79 -18
View File
@@ -22,39 +22,45 @@ from config.manage_api_client import report as manage_report
TAG = __name__
async def report(conn: "ConnectionHandler", type, text, opus_data, report_time):
async def report(conn: "ConnectionHandler", chat_type, text, audio_data, report_time):
"""执行聊天记录上报操作
Args:
conn: 连接对象
type: 上报类型,1为用户2为智能体3为工具调用
chat_type: 上报类型,1为用户(ASR/PCM)2为智能体(TTS/Opus)3为工具调用
text: 合成文本
opus_data: opus音频数据
audio_data: 音频数据chat_type=1时为PCM格式,chat_type=2时为Opus格式)
report_time: 上报时间
"""
try:
if opus_data:
audio_data = opus_to_wav(conn, opus_data)
if audio_data:
if chat_type == 1:
wav_data = pcm_to_wav(conn, audio_data)
elif chat_type == 2:
wav_data = opus_to_wav(conn, audio_data)
else:
wav_data = None
else:
audio_data = None
wav_data = None
# 执行异步上报
await manage_report(
mac_address=conn.device_id,
session_id=conn.session_id,
chat_type=type,
chat_type=chat_type,
content=text,
audio=audio_data,
audio=wav_data,
report_time=report_time,
)
except Exception as e:
conn.logger.bind(tag=TAG).error(f"聊天记录上报失败: {e}")
def opus_to_wav(conn: "ConnectionHandler", pcm_data):
def pcm_to_wav(conn: "ConnectionHandler", pcm_data):
"""将PCM数据转换为WAV格式的字节流
Args:
output_dir: 输出目录(保留参数以保持接口兼容)
conn: 连接对象
pcm_data: PCM音频数据(可能是列表或bytes)
Returns:
@@ -96,11 +102,62 @@ def opus_to_wav(conn: "ConnectionHandler", pcm_data):
raise
def opus_to_wav(conn: "ConnectionHandler", opus_data):
"""将Opus数据转换为WAV格式的字节流
Args:
conn: 连接对象
opus_data: Opus音频数据(可能是列表或bytes)
Returns:
bytes: WAV格式的音频数据
"""
decoder = None
try:
decoder = opuslib_next.Decoder(16000, 1)
pcm_data = []
if isinstance(opus_data, list):
for opus_packet in opus_data:
try:
pcm_frame = decoder.decode(opus_packet, 960)
pcm_data.append(pcm_frame)
except opuslib_next.OpusError as e:
conn.logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True)
elif isinstance(opus_data, bytes):
pcm_frame = decoder.decode(opus_data, 960)
pcm_data.append(pcm_frame)
if not pcm_data:
raise ValueError("没有有效的音频数据")
pcm_data_bytes = b"".join(pcm_data)
wav_header = bytearray()
wav_header.extend(b"RIFF")
wav_header.extend((36 + len(pcm_data_bytes)).to_bytes(4, "little"))
wav_header.extend(b"WAVE")
wav_header.extend(b"fmt ")
wav_header.extend((16).to_bytes(4, "little"))
wav_header.extend((1).to_bytes(2, "little"))
wav_header.extend((1).to_bytes(2, "little"))
wav_header.extend((16000).to_bytes(4, "little"))
wav_header.extend((32000).to_bytes(4, "little"))
wav_header.extend((2).to_bytes(2, "little"))
wav_header.extend((16).to_bytes(2, "little"))
wav_header.extend(b"data")
wav_header.extend(len(pcm_data_bytes).to_bytes(4, "little"))
return bytes(wav_header) + pcm_data_bytes
finally:
if decoder is not None:
try:
del decoder
except Exception as e:
conn.logger.bind(tag=TAG).debug(f"释放decoder资源时出错: {e}")
def enqueue_tts_report(conn: "ConnectionHandler", text, opus_data):
if not conn.read_config_from_api or conn.need_bind or not conn.report_tts_enable:
return
if conn.chat_history_conf == 0:
return
"""将TTS数据加入上报队列
Args:
@@ -108,6 +165,10 @@ def enqueue_tts_report(conn: "ConnectionHandler", text, opus_data):
text: 合成文本
opus_data: opus音频数据
"""
if not conn.read_config_from_api or conn.need_bind or not conn.report_tts_enable:
return
if conn.chat_history_conf == 0:
return
try:
# 使用连接对象的队列,传入文本和二进制数据而非文件路径
if conn.chat_history_conf == 2:
@@ -164,10 +225,6 @@ def enqueue_tool_report(conn: "ConnectionHandler", tool_name: str, tool_input: d
def enqueue_asr_report(conn: "ConnectionHandler", text, opus_data):
if not conn.read_config_from_api or conn.need_bind or not conn.report_asr_enable:
return
if conn.chat_history_conf == 0:
return
"""将ASR数据加入上报队列
Args:
@@ -175,6 +232,10 @@ def enqueue_asr_report(conn: "ConnectionHandler", text, opus_data):
text: 合成文本
opus_data: opus音频数据
"""
if not conn.read_config_from_api or conn.need_bind or not conn.report_asr_enable:
return
if conn.chat_history_conf == 0:
return
try:
# 使用连接对象的队列,传入文本和二进制数据而非文件路径
if conn.chat_history_conf == 2:
@@ -350,4 +350,16 @@ async def send_display_message(conn: "ConnectionHandler", text):
"text": text,
"session_id": conn.session_id
}
await conn.websocket.send(json.dumps(message))
transport = getattr(conn, "transport", None)
if transport is not None:
send_json = getattr(transport, "send_json", None)
if callable(send_json):
await send_json(message)
else:
await transport.send(json.dumps(message))
return
websocket = getattr(conn, "websocket", None)
if websocket is None:
raise AttributeError("无法找到可用的传输层接口")
await websocket.send(json.dumps(message))
+107 -6
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,11 +11,53 @@ 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())
try:
done, _ = await asyncio.wait(
{task, event_waiter},
timeout=timeout,
return_when=asyncio.FIRST_COMPLETED,
)
if self._started_event.is_set():
if task.done():
await task
return
if task in done:
await task
raise RuntimeError("HTTP服务器在就绪前退出")
raise TimeoutError("等待HTTP服务器启动超时")
finally:
if not event_waiter.done():
event_waiter.cancel()
await asyncio.gather(event_waiter, return_exceptions=True)
def _get_websocket_url(self, local_ip: str, port: int) -> str:
"""获取websocket地址
@@ -33,14 +78,28 @@ class SimpleHttpServer:
return f"ws://{local_ip}:{port}/xiaozhi/v1/"
async def start(self):
runner = None
self._start_active = True
try:
self._started_event.clear()
self._stop_event.clear()
server_config = self.config["server"]
read_config_from_api = self.config.get("read_config_from_api", False)
host = server_config.get("ip", "0.0.0.0")
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接口
@@ -74,19 +133,61 @@ 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
await runner.setup()
site = web.TCPSite(runner, host, port)
await site.start()
# 保持服务运行
while True:
await asyncio.sleep(3600) # 每隔 1 小时检查一次
self._started_event.set()
await self._stop_event.wait()
except Exception as e:
self.logger.bind(tag=TAG).error(f"HTTP服务器启动失败: {e}")
import traceback
self.logger.bind(tag=TAG).error(f"错误堆栈: {traceback.format_exc()}")
raise
finally:
self._started_event.clear()
try:
if runner is not None:
await self._cleanup_runner()
finally:
self._start_active = False
async def _cleanup_runner(self):
"""Cleanup the owned runner and retain it when cleanup must be retried."""
async with self._cleanup_lock:
runner = self._runner
if runner is None:
return
await runner.cleanup()
self._runner = None
async def stop(self):
"""Stop the HTTP loop and retry cleanup left by a failed start task."""
self._stop_event.set()
if not self._start_active:
await self._cleanup_runner()
@@ -0,0 +1,26 @@
from abc import ABC, abstractmethod
from typing import Any, List
class MessageProcessor(ABC):
"""消息处理器接口。返回 True 表示已处理并中止后续处理。"""
@abstractmethod
async def process(self, context: Any, transport: Any, message: Any) -> bool:
raise NotImplementedError
class MessagePipeline:
"""责任链式消息处理管道。"""
def __init__(self) -> None:
self._processors: List[MessageProcessor] = []
def add_processor(self, processor: MessageProcessor) -> None:
self._processors.append(processor)
async def process_message(self, context: Any, transport: Any, message: Any) -> None:
for processor in self._processors:
handled = await processor.process(context, transport, message)
if handled:
return
@@ -0,0 +1,132 @@
import json
import uuid
from typing import Any
from core.pipeline.message_pipeline import MessageProcessor
from core.context.session_context import SessionContext
from core.transport.transport_interface import TransportInterface
from core.components.component_manager import ComponentType
from config.logger import setup_logging
logger = setup_logging()
class AbortProcessor(MessageProcessor):
"""中断消息处理器:完整迁移abortMessageHandler.py和abortHandle.py的所有功能"""
async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
"""处理abort类型的消息"""
msg_json = None
if isinstance(message, str):
try:
msg_json = json.loads(message)
except json.JSONDecodeError:
msg_json = None
elif isinstance(message, dict):
msg_json = message
if isinstance(msg_json, dict) and msg_json.get("type") == "abort":
await self.handle_abort_message(context, transport, msg_json)
return True
return False
async def handle_abort_message(self, context: SessionContext, transport: TransportInterface, msg_json: dict):
"""处理中断消息 - 完整迁移自abortHandle.py的handleAbortMessage"""
msg_session_id = msg_json.get("session_id")
if msg_session_id and msg_session_id != context.session_id:
logger.warning(
f"忽略非当前会话的abort消息: "
f"msg={msg_session_id}, current={context.session_id}"
)
return
logger.info("Abort message received")
# 设置成打断状态,会自动打断llm、tts任务 - 完整迁移原逻辑
context.abort_requested = True
context.close_after_chat = False
# Stop the old LLM/tool producers before a subsequent listen/start can
# clear abort_requested and begin a new turn.
cancel_tasks = getattr(context, "cancel_turn_tasks", None)
if callable(cancel_tasks):
await cancel_tasks()
# 清理队列 - 完整迁移原逻辑
await self._clear_queues(context)
# 打断客户端说话状态 - 完整迁移原逻辑
await transport.send_json({
"type": "tts",
"state": "stop",
"session_id": context.session_id
})
# 清理说话状态 - 完整迁移原逻辑
self._clear_speak_status(context)
logger.info("Abort message received-end")
async def _clear_queues(self, context: SessionContext):
"""清理所有队列 - 完整迁移原clear_queues逻辑"""
try:
# 清理TTS音频队列
tts_component = await self._get_component(context, ComponentType.TTS)
if tts_component and hasattr(tts_component, 'tts_instance'):
tts_instance = tts_component.tts_instance
for queue_name in ('tts_text_queue', 'tts_audio_queue'):
pending_queue = getattr(tts_instance, queue_name, None)
if pending_queue is None:
continue
try:
while not pending_queue.empty():
pending_queue.get_nowait()
except:
pass
# Rotate the turn token so late text/audio from the aborted turn is stale.
context.sentence_id = uuid.uuid4().hex
# 清理ASR音频队列
context.clear_audio_buffer()
# 清理其他可能的队列
if hasattr(context, 'clear_queues'):
context.clear_queues()
except Exception as e:
logger.error(f"清理队列时出错: {e}")
def _clear_speak_status(self, context: SessionContext):
"""清理说话状态 - 完整迁移原clearSpeakStatus逻辑"""
try:
# 清理说话状态
context.is_speaking = False
# 如果有其他说话状态相关的属性,也一并清理
if hasattr(context, 'clearSpeakStatus'):
context.clearSpeakStatus()
# 重置相关状态
context.client_have_voice = False
context.client_voice_stop = True
except Exception as e:
logger.error(f"清理说话状态时出错: {e}")
async def _send_abort_confirmation(self, transport: TransportInterface, session_id: str):
"""发送中断确认响应(可选)"""
response = {
"type": "abort",
"status": "success",
"message": "中断操作已完成",
"session_id": session_id
}
try:
await transport.send(json.dumps(response))
except Exception as e:
logger.error(f"发送中断确认响应失败: {e}")
async def _get_component(self, context: SessionContext, component_type: ComponentType):
if not context.component_manager:
return None
return await context.component_manager.get_component(component_type, context)
@@ -0,0 +1,439 @@
import time
import json
import asyncio
import uuid
from typing import Any
from core.pipeline.message_pipeline import MessageProcessor
from core.context.session_context import SessionContext
from core.transport.transport_interface import TransportInterface
from core.components.component_manager import ComponentType
from core.services.audio_ingress_service import AudioIngressService
from core.utils.util import audio_to_data
from core.utils.output_counter import check_device_output_limit
from config.logger import setup_logging
logger = setup_logging()
def arm_wake_audio_suppression(context: SessionContext) -> None:
"""Ignore reordered wake-word audio for a bounded interval."""
context.just_woken_up = True
previous = getattr(context, "wake_audio_suppression_task", None)
if previous and not previous.done():
previous.cancel()
delay_ms = max(
0,
int(context.config.get("wake_audio_suppression_ms", 2000)),
)
if delay_ms == 0:
context.just_woken_up = False
context.wake_audio_suppression_task = None
return
session_id = context.session_id
context.wake_audio_suppression_task = context.create_background_task(
_release_wake_audio_suppression(
context,
session_id,
delay_ms / 1000,
),
conversation_scoped=True,
)
async def _release_wake_audio_suppression(
context: SessionContext,
session_id: str,
delay_seconds: float,
) -> None:
current_task = asyncio.current_task()
try:
await asyncio.sleep(delay_seconds)
if context.session_id == session_id:
context.just_woken_up = False
finally:
if getattr(context, "wake_audio_suppression_task", None) is current_task:
context.wake_audio_suppression_task = None
class AudioReceiveProcessor(MessageProcessor):
"""音频接收处理器:完整迁移receiveAudioHandle.py的所有功能"""
def __init__(self, audio_ingress_service=None):
self.audio_ingress_service = audio_ingress_service or AudioIngressService()
async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
"""处理音频消息"""
if isinstance(message, bytes):
await self.handle_audio_message(context, transport, message, 0)
return True
if isinstance(message, dict) and message.get("type") == "audio":
audio_data = message.get("data")
if isinstance(audio_data, (bytes, bytearray)):
await self.handle_audio_message(
context,
transport,
bytes(audio_data),
int(message.get("timestamp", 0) or 0),
)
return True
return True
return False
async def handle_audio_message(
self,
context: SessionContext,
transport: TransportInterface,
audio: bytes,
timestamp: int = 0,
):
"""处理音频消息 - 完整迁移自handleAudioMessage"""
if (
getattr(transport, "requires_audio_tail_grace", False)
and not getattr(context, "accepting_input_audio", False)
):
return
pcm_frame = self.audio_ingress_service.process(context, audio, timestamp)
if not pcm_frame:
return
# Do not feed encoded wake-word history into stateful VAD. Native MQTT
# gates the usual pre-start history above; this bounded suppression only
# covers frames that crossed the independent control/audio channels.
if context.just_woken_up:
if not getattr(context, "wake_audio_suppression_task", None):
arm_wake_audio_suppression(context)
context.asr_audio.clear()
context.reset_vad_states()
return
# 获取VAD组件
vad_component = await self._get_component(context, ComponentType.VAD)
if not vad_component or not hasattr(vad_component, 'vad_instance'):
now = time.time()
last_log = getattr(context, "_vad_missing_last_log", 0.0)
if now - last_log > 10:
logger.warning("VAD组件未初始化")
context._vad_missing_last_log = now
return
vad_instance = vad_component.vad_instance
# 当前片段是否有人说话
have_voice = vad_instance.is_vad(context, pcm_frame)
if have_voice:
if (
getattr(context, "client_aec", False)
and context.is_speaking
and context.listen_mode != "manual"
):
await self._handle_abort_message(context, transport)
# 设备长时间空闲检测,用于say goodbye
await self._no_voice_close_connect(context, transport, have_voice)
# 接收音频
asr_component = await self._get_component(context, ComponentType.ASR)
asr_instance = asr_component.asr_instance if asr_component and hasattr(asr_component, 'asr_instance') else None
# 自动模式兜底:没有收到 listen stop 时,基于VAD触发一次停止
if (
not have_voice
and context.client_have_voice
and not context.client_voice_stop
and not getattr(context, "listen_stop_pending", False)
and context.listen_mode != "manual"
and asr_instance is not None
):
context.client_voice_stop = True
try:
from core.providers.asr.dto.dto import InterfaceType
if hasattr(asr_instance, "interface_type") and asr_instance.interface_type == InterfaceType.STREAM:
if hasattr(asr_instance, "_send_stop_request"):
context.create_background_task(
asr_instance._send_stop_request(),
turn_scoped=True,
)
else:
if len(context.asr_audio) > 0:
asr_audio_task = context.asr_audio.copy()
context.asr_audio.clear()
context.reset_vad_states()
await asr_instance.handle_voice_stop(context, asr_audio_task)
except Exception as e:
logger.error(f"自动VAD停止处理失败: {e}")
if asr_instance and hasattr(asr_instance, 'receive_audio'):
await asr_instance.receive_audio(context, pcm_frame, have_voice)
async def start_to_chat(
self,
context: SessionContext,
transport: TransportInterface,
text: str,
*,
preserve_close_after_chat: bool = False,
):
"""开始聊天 - 完整迁移自startToChat"""
# 检查输入是否是JSON格式(包含说话人信息)
speaker_name = None
actual_text = text
try:
# 尝试解析JSON格式的输入
if text.strip().startswith('{') and text.strip().endswith('}'):
data = json.loads(text)
if 'speaker' in data and 'content' in data:
speaker_name = data['speaker']
actual_content = data['content']
logger.info(f"解析到说话人信息: {speaker_name}")
if speaker_name not in context.introduced_speakers:
context.introduced_speakers.add(speaker_name)
actual_text = text
else:
actual_text = actual_content
except (json.JSONDecodeError, KeyError):
# 如果解析失败,继续使用原始文本
pass
# 保存说话人信息到上下文
if speaker_name:
context.current_speaker = speaker_name
else:
context.current_speaker = None
# 检查设备绑定
if context.need_bind:
await self.prompt_bind_device(context, transport)
return
# 如果当日的输出字数大于限定的字数
if context.max_output_size > 0:
if check_device_output_limit(
context.headers.get("device-id"), context.max_output_size
):
await self._max_out_size(context, transport)
return
if context.is_speaking and getattr(context, "listen_mode", "auto") != "manual":
await self._handle_abort_message(context, transport)
context.abort_requested = False
if not preserve_close_after_chat:
context.close_after_chat = False
# 首先进行意图分析,使用实际文本内容
from core.processors.chat_processor import ChatProcessor
chat_processor = ChatProcessor()
intent_handled = await chat_processor.handle_user_intent(context, transport, actual_text)
if intent_handled:
# 如果意图已被处理,不再进行聊天
return
# 意图未被处理,继续常规聊天流程,使用实际文本内容
await self._send_stt_message(context, transport, actual_text)
# 与旧架构一致:将聊天处理放入线程池,避免阻塞事件循环
from core.processors.chat_processor import ChatProcessor
chat_processor = ChatProcessor()
if hasattr(context, "create_background_task"):
context.create_background_task(
chat_processor.handle_chat(
context,
transport,
actual_text,
skip_intent=True,
),
turn_scoped=True,
)
else:
await chat_processor.handle_chat(
context, transport, actual_text, skip_intent=True
)
async def _no_voice_close_connect(self, context: SessionContext, transport: TransportInterface, have_voice: bool):
"""无声音时关闭连接检测 - 完整迁移自no_voice_close_connect"""
if have_voice:
context.update_activity()
return
# 只有在已经初始化过时间戳的情况下才进行超时检查
if context.last_activity_time_ms > 0.0:
no_voice_time = time.time() * 1000 - context.last_activity_time_ms
close_connection_no_voice_time = int(
context.config.get("close_connection_no_voice_time", 120)
)
if (
not context.close_after_chat
and no_voice_time > 1000 * close_connection_no_voice_time
):
context.close_after_chat = True
context.abort_requested = False
end_prompt = context.config.get("end_prompt", {})
if end_prompt and end_prompt.get("enable", True) is False:
logger.info("结束对话,无需发送结束提示语")
if transport.keeps_connection_between_sessions:
session_id = context.session_id
end_conversation = getattr(context, "end_conversation", None)
if callable(end_conversation):
await end_conversation(session_id)
await transport.end_session(session_id)
context.close_after_chat = False
context.reset_vad_states()
else:
await transport.close()
return
prompt = end_prompt.get("prompt")
if not prompt:
prompt = "请你以```时间过得真快```未来头,用富有感情、依依不舍的话来结束这场对话吧。!"
await self.start_to_chat(
context,
transport,
prompt,
preserve_close_after_chat=True,
)
async def _max_out_size(self, context: SessionContext, transport: TransportInterface):
"""超出最大输出字数处理 - 完整迁移自max_out_size"""
# 播放超出最大输出字数的提示
context.abort_requested = False
text = "不好意思,我现在有点事情要忙,明天这个时候我们再聊,约好了哦!明天不见不散,拜拜!"
await self._send_stt_message(context, transport, text)
file_path = "config/assets/max_output_size.wav"
opus_packets = await audio_to_data(file_path)
# 获取TTS组件并添加到队列
tts_component = await self._get_component(context, ComponentType.TTS)
if tts_component and hasattr(tts_component, 'tts_instance'):
tts_instance = tts_component.tts_instance
if hasattr(tts_instance, 'tts_audio_queue'):
from core.providers.tts.dto.dto import SentenceType
tts_instance.tts_audio_queue.put(
(SentenceType.LAST, opus_packets, text, context.sentence_id)
)
context.close_after_chat = True
async def prompt_bind_device(
self,
context: SessionContext,
transport: TransportInterface,
):
"""检查设备绑定 - 完整迁移自check_bind_device"""
bind_code = context.bind_code
if bind_code:
# 确保bind_code是6位数字
if len(bind_code) != 6:
logger.error(f"无效的绑定码格式: {bind_code}")
text = "绑定码格式错误,请检查配置。"
await self._send_stt_message(context, transport, text)
return
text = f"请登录控制面板,输入{bind_code},绑定设备。"
await self._send_stt_message(context, transport, text)
# 获取TTS组件
tts_component = await self._get_component(context, ComponentType.TTS)
if not tts_component or not hasattr(tts_component, 'tts_instance'):
return
tts_instance = tts_component.tts_instance
if not hasattr(tts_instance, 'tts_audio_queue'):
return
# 播放提示音
from core.providers.tts.dto.dto import SentenceType
music_path = "config/assets/bind_code.wav"
opus_packets = await audio_to_data(music_path)
tts_instance.tts_audio_queue.put(
(SentenceType.FIRST, opus_packets, text, context.sentence_id)
)
# 逐个播放数字
for i in range(6): # 确保只播放6位数字
try:
digit = bind_code[i]
num_path = f"config/assets/bind_code/{digit}.wav"
num_packets = await audio_to_data(num_path)
tts_instance.tts_audio_queue.put(
(SentenceType.MIDDLE, num_packets, None, context.sentence_id)
)
except Exception as e:
logger.error(f"播放数字音频失败: {e}")
continue
tts_instance.tts_audio_queue.put(
(SentenceType.LAST, [], None, context.sentence_id)
)
else:
# 播放未绑定提示
context.abort_requested = False
text = f"没有找到该设备的版本信息,请正确配置 OTA地址,然后重新编译固件。"
await self._send_stt_message(context, transport, text)
# 获取TTS组件
tts_component = await self._get_component(context, ComponentType.TTS)
if tts_component and hasattr(tts_component, 'tts_instance'):
tts_instance = tts_component.tts_instance
if hasattr(tts_instance, 'tts_audio_queue'):
from core.providers.tts.dto.dto import SentenceType
music_path = "config/assets/bind_not_found.wav"
opus_packets = await audio_to_data(music_path)
tts_instance.tts_audio_queue.put(
(SentenceType.LAST, opus_packets, text, context.sentence_id)
)
async def _handle_abort_message(self, context: SessionContext, transport: TransportInterface):
"""处理中断消息"""
from core.processors.abort_processor import AbortProcessor
await AbortProcessor().handle_abort_message(
context,
transport,
{"type": "abort", "session_id": context.session_id},
)
async def _clear_queues(self, context: SessionContext):
"""清理所有队列"""
# 清理TTS音频队列
tts_component = await self._get_component(context, ComponentType.TTS)
if tts_component and hasattr(tts_component, 'tts_instance'):
tts_instance = tts_component.tts_instance
for queue_name in ('tts_text_queue', 'tts_audio_queue'):
pending_queue = getattr(tts_instance, queue_name, None)
if pending_queue is None:
continue
try:
while not pending_queue.empty():
pending_queue.get_nowait()
except Exception:
pass
# Invalidate late TTS text/audio generated by the interrupted turn.
context.sentence_id = uuid.uuid4().hex
# 清理ASR音频队列
context.clear_audio_buffer()
async def _get_component(self, context: SessionContext, component_type: ComponentType):
if not context.component_manager:
return None
return await context.component_manager.get_component(component_type, context)
async def _send_stt_message(self, context: SessionContext, transport: TransportInterface, text: str):
"""发送STT消息"""
from core.processors.audio_send_processor import AudioSendProcessor
await AudioSendProcessor().send_stt_message(
context, transport, text
)
@@ -0,0 +1,578 @@
import asyncio
import json
import time
from typing import Any, List
from core.pipeline.message_pipeline import MessageProcessor
from core.context.session_context import SessionContext
from core.transport.transport_interface import TransportInterface
from core.providers.tts.dto.dto import SentenceType
from core.utils import textUtils
from core.components.component_manager import ComponentType
from core.services.audio_ingress_service import AudioIngressService
from config.logger import setup_logging
logger = setup_logging()
class AudioSendProcessor(MessageProcessor):
"""音频发送处理器:完整迁移sendAudioHandle.py的所有功能"""
def __init__(self, audio_ingress_service=None):
self.audio_ingress_service = audio_ingress_service or AudioIngressService()
async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
"""这个处理器不直接处理消息,而是被其他处理器调用"""
return False
async def send_audio_message(self, context: SessionContext, transport: TransportInterface,
sentence_type: SentenceType, audios: bytes, text: str,
sentence_id: str = None):
"""发送音频消息 - 完整迁移自sendAudioMessage"""
if not self._is_current_turn(context, sentence_id):
return
tts_component = await self._get_component(context, ComponentType.TTS)
if not tts_component or not hasattr(tts_component, 'tts_instance'):
return
tts_instance = tts_component.tts_instance
if hasattr(tts_instance, 'tts_audio_first_sentence') and tts_instance.tts_audio_first_sentence:
logger.info(f"发送第一段语音: {text}")
tts_instance.tts_audio_first_sentence = False
await self.start_tts_stream(context, transport)
if sentence_type == SentenceType.FIRST:
await self.send_tts_message(
context, transport, "sentence_start", text, sentence_id=sentence_id
)
await self.send_audio(
context, transport, audios, sentence_id=sentence_id
)
if not self._is_current_turn(context, sentence_id):
return
# 发送句子开始消息
if sentence_type is not SentenceType.MIDDLE:
logger.info(f"发送音频消息: {sentence_type}, {text}")
# 发送结束消息(如果是最后一个文本)
if (
not getattr(context, "calling", False)
and context.llm_finish_task
and sentence_type == SentenceType.LAST
):
# Latch the terminal action before sending tts:stop. The device can
# immediately answer with listen:start, whose handler resets the
# mutable turn flags while this coroutine is still awaiting I/O.
close_after_chat = bool(context.close_after_chat)
closing_session_id = (
getattr(transport, "session_id", None) or context.session_id
)
if close_after_chat:
context.close_after_chat = False
await self.send_tts_message(
context, transport, "stop", None, sentence_id=sentence_id
)
if not self._is_current_turn(context, sentence_id):
return
context.is_speaking = False
if close_after_chat:
# MQTT/UDP:结束后回到Idle,不关闭连接
if transport.keeps_connection_between_sessions:
end_conversation = getattr(context, "end_conversation", None)
if callable(end_conversation):
await end_conversation(closing_session_id)
await transport.end_session(closing_session_id)
else:
await transport.close()
async def send_audio(
self,
context: SessionContext,
transport: TransportInterface,
audios: bytes,
frame_duration: int = 60,
sentence_id: str = None,
):
"""发送单个opus包,支持流控 - 完整迁移自sendAudio"""
if audios is None or len(audios) == 0:
return
if getattr(context, "audio_flow_control", {}).get("send_failed"):
return
# MQTT/UDP:等待UDP远端地址就绪,避免首包丢失
if transport.has_datagram_audio:
if not await transport.wait_audio_ready(timeout=2):
logger.warning("UDP远端地址未就绪,跳过音频发送")
return
audio_list = [audios] if isinstance(audios, bytes) else audios
if not isinstance(audio_list, list):
return
flow = context.audio_flow_control
pre_buffer_count = max(
0, int(context.config.get("tts_pre_buffer_count", 5))
)
for audio in audio_list:
if context.abort_requested or not self._is_current_turn(context, sentence_id):
break
# Match the legacy AudioRateController contract: send only the
# initial pre-buffer inline, then enqueue the remaining frames so
# the TTS consumer can prefetch later sentence chunks.
packet_count = int(flow.get("packet_count", 0))
if packet_count < pre_buffer_count and "_send_queue" not in flow:
sent = await self._send_audio_packet(
context,
transport,
flow,
audio,
frame_duration,
sentence_id,
paced=False,
)
if not sent:
break
continue
queue = self._ensure_audio_sender(
context, transport, flow, sentence_id
)
await queue.put(("audio", audio, frame_duration, sentence_id))
@staticmethod
async def _pace_audio_send(
context: SessionContext,
frame_duration: int,
flow: dict = None,
) -> None:
flow = flow if flow is not None else context.audio_flow_control
packet_count = int(flow.get("packet_count", 0))
pre_buffer_count = max(0, int(context.config.get("tts_pre_buffer_count", 5)))
if packet_count < pre_buffer_count:
return
configured_delay = int(context.config.get("tts_audio_send_delay", -1))
if configured_delay > 0:
await asyncio.sleep(configured_delay / 1000.0)
return
delay_ms = max(0, int(frame_duration))
if delay_ms <= 0:
return
now = time.monotonic()
# packet_count is the number already sent. The first packet after the
# pre-buffer must therefore wait one complete frame, not become an
# extra immediate packet.
paced_packet_index = packet_count - pre_buffer_count + 1
pacing_started_at = flow.get("pacing_started_at")
pacing_delay_ms = flow.get("pacing_delay_ms")
if pacing_started_at is None or pacing_delay_ms != delay_ms:
pacing_started_at = now
flow["pacing_started_at"] = pacing_started_at
flow["pacing_delay_ms"] = delay_ms
delay_seconds = delay_ms / 1000.0
target = pacing_started_at + paced_packet_index * delay_seconds
# TTS producers can pause between sentence chunks. Rebase after a
# full-frame gap instead of bursting stale deadlines to catch up.
if now - target >= delay_seconds:
pacing_started_at = now - paced_packet_index * delay_seconds
flow["pacing_started_at"] = pacing_started_at
target = now
wait_seconds = target - now
if wait_seconds > 0:
await asyncio.sleep(wait_seconds)
def _ensure_audio_sender(
self,
context: SessionContext,
transport: TransportInterface,
flow: dict,
sentence_id: str = None,
) -> asyncio.Queue:
queue = flow.get("_send_queue")
task = flow.get("_send_task")
if queue is not None and task is not None and not task.done():
return queue
queue_size = max(
1, int(context.config.get("tts_audio_queue_size", 256))
)
queue = asyncio.Queue(maxsize=queue_size)
sender = self._audio_send_loop(
context, transport, flow, queue, sentence_id
)
create_task = getattr(context, "create_background_task", None)
if callable(create_task):
task = create_task(sender, turn_scoped=True)
else:
task = asyncio.create_task(sender)
flow["_send_queue"] = queue
flow["_send_task"] = task
return queue
async def _audio_send_loop(
self,
context: SessionContext,
transport: TransportInterface,
flow: dict,
queue: asyncio.Queue,
sentence_id: str = None,
) -> None:
try:
while True:
item = await queue.get()
try:
item_type = item[0]
if not self._is_current_turn(context, sentence_id):
continue
if item_type == "audio":
_, audio, frame_duration, item_sentence_id = item
await self._send_audio_packet(
context,
transport,
flow,
audio,
frame_duration,
item_sentence_id,
paced=True,
)
elif item_type == "json":
_, message, item_sentence_id = item
if self._is_current_turn(context, item_sentence_id):
await transport.send_json(message)
finally:
queue.task_done()
except asyncio.CancelledError:
raise
except Exception as error:
flow["send_failed"] = True
logger.error("后台音频发送循环失败: {}", error)
finally:
while True:
try:
queue.get_nowait()
except asyncio.QueueEmpty:
break
else:
queue.task_done()
async def _send_audio_packet(
self,
context: SessionContext,
transport: TransportInterface,
flow: dict,
audio: bytes,
frame_duration: int,
sentence_id: str = None,
*,
paced: bool,
) -> bool:
if paced:
await self._pace_audio_send(context, frame_duration, flow)
if context.abort_requested or not self._is_current_turn(context, sentence_id):
return False
context.update_activity()
timestamp = self._next_audio_timestamp(
context, frame_duration, flow
)
sent = await self._send_audio_with_retry(
context,
transport,
audio,
timestamp,
sentence_id,
flow,
)
if not sent or not self._is_current_turn(context, sentence_id):
return False
self.audio_ingress_service.cache_output_reference(
context, audio, timestamp
)
flow["packet_count"] = int(flow.get("packet_count", 0)) + 1
return True
async def _send_audio_with_retry(
self,
context: SessionContext,
transport: TransportInterface,
audio: bytes,
timestamp: int,
sentence_id: str = None,
flow: dict = None,
) -> bool:
flow = flow if flow is not None else context.audio_flow_control
retries = max(0, int(context.config.get("audio_send_retries", 2)))
retry_delay = max(
0, int(context.config.get("audio_send_retry_delay_ms", 20))
) / 1000.0
for attempt in range(retries + 1):
if context.abort_requested or not self._is_current_turn(context, sentence_id):
return False
try:
await transport.send_audio(audio, timestamp)
return self._is_current_turn(context, sentence_id)
except Exception as error:
if attempt >= retries:
flow["send_failed"] = True
logger.error(
"Audio send failed after {} attempts: {}",
attempt + 1,
error,
)
raise
await asyncio.sleep(retry_delay)
return False
@staticmethod
def _next_audio_timestamp(
context: SessionContext,
frame_duration: int,
flow: dict = None,
) -> int:
flow = flow if flow is not None else getattr(
context, "audio_flow_control", None
)
if flow is None:
flow = {}
context.audio_flow_control = flow
timestamp_step = int(frame_duration) or int(
getattr(context, "output_frame_duration", 60) or 60
)
timestamp = int(flow.get("output_timestamp", 0)) + timestamp_step
timestamp %= 2 ** 32
flow["output_timestamp"] = timestamp
return timestamp
async def send_stt_message(self, context: SessionContext, transport: TransportInterface, text: str):
"""发送STT消息 - 完整迁移自send_stt_message"""
end_prompt = getattr(context, "config", {}).get("end_prompt", {}).get("prompt")
if end_prompt and text == end_prompt:
await self.start_tts_stream(context, transport)
return
display_text = text
parsed_data = None
if isinstance(text, dict):
parsed_data = text
elif isinstance(text, str):
stripped = text.strip()
if stripped.startswith("{") and stripped.endswith("}"):
try:
parsed_data = json.loads(stripped)
except json.JSONDecodeError:
pass
if isinstance(parsed_data, dict) and "content" in parsed_data:
display_text = parsed_data["content"]
if "speaker" in parsed_data:
context.current_speaker = parsed_data["speaker"]
stt_text = textUtils.get_string_no_punctuation_or_emoji(
str(display_text)
)
await self._send_json(transport, {
"type": "stt",
"text": stt_text,
"session_id": context.session_id
})
logger.info(f"发送STT消息: {stt_text}")
# Legacy WS sends TTS start together with STT, before synthesis begins.
# This lead time is required by MQTT/UDP clients because the device
# changes to Speaking asynchronously and drops UDP audio beforehand.
await self.start_tts_stream(context, transport)
async def start_tts_stream(
self,
context: SessionContext,
transport: TransportInterface,
) -> bool:
"""Open one device-side TTS stream without sending duplicate starts."""
if getattr(context, "is_speaking", False):
return False
context.is_speaking = True
flow_control = getattr(context, "audio_flow_control", None) or {}
await self._stop_audio_sender(flow_control)
context.audio_flow_control = {}
try:
await self._send_json(transport, {
"type": "tts",
"state": "start",
"session_id": context.session_id,
})
except Exception:
context.is_speaking = False
raise
return True
@staticmethod
async def _send_json(
transport: TransportInterface,
message: dict,
) -> None:
send_json = getattr(transport, "send_json", None)
if callable(send_json):
await send_json(message)
return
await transport.send(json.dumps(message))
async def send_tts_message(
self,
context: SessionContext,
transport: TransportInterface,
state: str,
text: str = None,
sentence_id: str = None,
):
"""发送TTS消息 - 完整迁移自send_tts_message"""
if not self._is_current_turn(context, sentence_id):
return False
if state == "sentence_start":
flow = getattr(context, "audio_flow_control", {})
queue = flow.get("_send_queue")
task = flow.get("_send_task")
if queue is not None and task is not None and not task.done():
message = {
"type": "tts",
"state": state,
"session_id": context.session_id,
}
if text:
message["text"] = text
await queue.put(("json", message, sentence_id))
return True
if state == "stop":
if context.config.get("enable_stop_tts_notify", False):
from core.utils.util import audio_to_data
notify_path = context.config.get(
"stop_tts_notify_voice", "config/assets/tts_notify.mp3"
)
notify_audio = await audio_to_data(notify_path, is_opus=True)
if notify_audio:
await self.send_audio(
context,
transport,
notify_audio,
sentence_id=sentence_id,
)
if not await self._wait_for_audio_completion(context, sentence_id):
return False
if not self._is_current_turn(context, sentence_id):
return False
message = {
"type": "tts",
"state": state,
"session_id": context.session_id
}
if text:
message["text"] = text
await transport.send_json(message)
logger.debug(f"发送TTS消息: state={state}, text={text}")
return True
@staticmethod
async def _wait_for_audio_completion(
context: SessionContext, sentence_id: str = None
) -> bool:
flow = context.audio_flow_control
queue = flow.get("_send_queue")
sender_task = flow.get("_send_task")
if queue is not None and sender_task is not None:
join_task = asyncio.create_task(queue.join())
done, _ = await asyncio.wait(
{join_task, sender_task},
return_when=asyncio.FIRST_COMPLETED,
)
if sender_task in done and not join_task.done():
join_task.cancel()
await asyncio.gather(join_task, return_exceptions=True)
return False
await join_task
await AudioSendProcessor._stop_audio_sender(flow)
if flow.get("send_failed"):
return False
packet_count = int(flow.get("packet_count", 0))
if packet_count <= 0:
return AudioSendProcessor._is_current_turn(context, sentence_id)
frame_duration = max(1, int(getattr(context, "output_frame_duration", 60)))
pre_buffer_count = max(0, int(context.config.get("tts_pre_buffer_count", 5)))
tail_frames = max(
0,
int(context.config.get("tts_stop_buffer_frames", pre_buffer_count + 2)),
)
if tail_frames:
await asyncio.sleep(tail_frames * frame_duration / 1000.0)
return AudioSendProcessor._is_current_turn(context, sentence_id)
@staticmethod
async def _stop_audio_sender(flow: dict) -> None:
task = flow.pop("_send_task", None)
flow.pop("_send_queue", None)
if task is not None and not task.done():
task.cancel()
await asyncio.gather(task, return_exceptions=True)
@staticmethod
def _is_current_turn(context: SessionContext, sentence_id: str = None) -> bool:
if sentence_id is None:
return bool(
getattr(context, "conversation_active", True)
and getattr(context, "sentence_id", None) is None
)
return sentence_id == context.sentence_id
async def send_music_message(self, context: SessionContext, transport: TransportInterface,
music_path: str, text: str):
"""发送音乐消息 - 完整迁移自send_music_message"""
from core.utils.util import audio_to_data
try:
# 获取音频数据
opus_packets = await audio_to_data(music_path)
if opus_packets:
# 发送音乐开始消息
await self.send_tts_message(context, transport, "start", text)
# 发送音频数据
await self.send_audio(context, transport, opus_packets)
# 发送音乐结束消息
await self.send_tts_message(context, transport, "stop", None)
logger.info(f"发送音乐: {music_path}")
else:
logger.warning(f"无法加载音乐文件: {music_path}")
except Exception as e:
logger.error(f"发送音乐失败: {e}")
async def _get_component(self, context: SessionContext, component_type: ComponentType):
if not context.component_manager:
return None
return await context.component_manager.get_component(component_type, context)
async def send_welcome_audio(self, context: SessionContext, transport: TransportInterface):
"""发送欢迎音频"""
welcome_audio_path = context.config.get("welcome_audio_path")
if welcome_audio_path:
await self.send_music_message(context, transport, welcome_audio_path, "欢迎使用小智助手")
async def send_goodbye_audio(self, context: SessionContext, transport: TransportInterface):
"""发送告别音频"""
goodbye_audio_path = context.config.get("goodbye_audio_path")
if goodbye_audio_path:
await self.send_music_message(context, transport, goodbye_audio_path, "再见,期待下次相遇")
@@ -0,0 +1,55 @@
from typing import Any
from core.pipeline.message_pipeline import MessageProcessor
from core.context.session_context import SessionContext
from core.transport.transport_interface import TransportInterface
from core.auth import AuthMiddleware, AuthenticationError
from config.logger import setup_logging
logger = setup_logging()
class AuthProcessor(MessageProcessor):
"""认证处理器:处理连接认证逻辑"""
def __init__(self):
self.auth_middleware = None
async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
"""处理认证相关逻辑"""
# 服务器管理连接走server消息处理(基于secret校验),跳过认证
if isinstance(message, str):
try:
import json
msg_json = json.loads(message)
if isinstance(msg_json, dict) and msg_json.get("type") == "server":
return False
except json.JSONDecodeError:
pass
elif isinstance(message, dict) and message.get("type") == "server":
return False
# 如果已经认证,跳过
if context.is_authenticated:
return False
# 初始化认证中间件(延迟初始化)
if self.auth_middleware is None:
self.auth_middleware = AuthMiddleware(context.config)
# 检查是否为认证消息(通过headers进行认证)
if context.headers:
try:
await self.auth_middleware.authenticate_async(context.headers)
context.is_authenticated = True
logger.info(f"设备认证成功: {context.device_id}")
return False # 认证成功,继续处理其他消息
except AuthenticationError as e:
logger.error(f"设备认证失败: {e}")
# 发送认证失败消息
await transport.send("Authentication failed")
await transport.close()
return True # 认证失败,停止处理
# 如果没有认证信息,要求认证
await transport.send("Authentication required")
return True # 停止后续处理
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,46 @@
import json
from typing import Any
from core.pipeline.message_pipeline import MessageProcessor
from core.context.session_context import SessionContext
from core.transport.transport_interface import TransportInterface
from config.logger import setup_logging
logger = setup_logging()
class GoodbyeProcessor(MessageProcessor):
"""Goodbye消息处理器:处理会话结束"""
async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
msg_json = None
if isinstance(message, str):
try:
msg_json = json.loads(message)
except json.JSONDecodeError:
msg_json = None
elif isinstance(message, dict):
msg_json = message
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"))
else:
await transport.close()
return True
return False
@@ -0,0 +1,351 @@
import time
import json
import random
import asyncio
import uuid
from typing import Any
from core.pipeline.message_pipeline import MessageProcessor
from core.context.session_context import SessionContext
from core.transport.transport_interface import TransportInterface
from core.utils.dialogue import Message
from core.utils.util import audio_to_data, remove_punctuation_and_length, opus_datas_to_wav_bytes
from core.providers.tts.dto.dto import SentenceType
from core.processors.audio_receive_processor import arm_wake_audio_suppression
from core.utils.wakeup_word import WakeupWordsConfig
from core.components.component_manager import ComponentType
from core.providers.tools.device_mcp import (
MCPClient,
send_mcp_initialize_message,
)
from config.logger import setup_logging
logger = setup_logging()
# 唤醒词配置
WAKEUP_CONFIG = {
"refresh_time": 5,
"words": ["你好", "你好啊", "嘿,你好", ""],
}
# 创建全局的唤醒词配置管理器
wakeup_words_config = WakeupWordsConfig()
# 用于防止并发调用wakeupWordsResponse的锁
_wakeup_response_lock = asyncio.Lock()
class HelloProcessor(MessageProcessor):
"""Hello消息处理器:完整迁移helloHandle.py的所有功能"""
async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
"""处理hello类型的消息"""
msg_json = None
if isinstance(message, str):
try:
msg_json = json.loads(message)
except json.JSONDecodeError:
msg_json = None
elif isinstance(message, dict):
msg_json = message
if isinstance(msg_json, dict) and msg_json.get("type") == "hello":
await self.handle_hello_message(context, transport, msg_json)
return True
return False
async def handle_hello_message(self, context: SessionContext, transport: TransportInterface, msg_json: dict):
"""处理hello消息 - 完整迁移自handleHelloMessage"""
pending_tasks = []
for task_name in (
"listen_stop_task",
"listen_start_task",
):
pending_task = getattr(context, task_name, None)
if pending_task and not pending_task.done():
pending_task.cancel()
pending_tasks.append(pending_task)
setattr(context, task_name, None)
if pending_tasks:
await asyncio.gather(*pending_tasks, return_exceptions=True)
transport_session_id = getattr(transport, "session_id", None)
if transport_session_id:
if (
context.session_id != transport_session_id
and callable(getattr(context, "end_conversation", None))
):
# A previous finalizer may already have published
# conversation_active=False while persistence is still in
# progress. Always join it before exposing the new session.
await context.end_conversation(context.session_id)
context.session_id = transport_session_id
if context.welcome_msg:
context.welcome_msg["session_id"] = context.session_id
dialogue = getattr(context, "dialogue", None)
if getattr(context, "prompt", None) and dialogue and not dialogue.dialogue:
dialogue.update_system_message(context.prompt)
inject_fewshot = getattr(context, "inject_tool_call_fewshot", None)
if callable(inject_fewshot):
inject_fewshot()
context.conversation_active = True
context.listen_stop_pending = False
context.listen_stop_deadline = 0.0
if transport.requires_audio_tail_grace:
context.accepting_input_audio = False
# 处理音频参数
audio_params = msg_json.get("audio_params")
if audio_params:
format = audio_params.get("format")
logger.info(f"客户端音频格式: {format}")
context.audio_format = format
context.input_sample_rate = int(
audio_params.get("sample_rate", getattr(context, "input_sample_rate", 16000))
)
context.input_channels = int(
audio_params.get("channels", getattr(context, "input_channels", 1))
)
context.input_frame_duration = int(
audio_params.get(
"frame_duration", getattr(context, "input_frame_duration", 60)
)
)
# 处理客户端特性
features = msg_json.get("features") or {}
context.features = dict(features)
context.client_aec = bool(features.get("aec"))
mcp_enabled = bool(features.get("mcp"))
previous_mcp_client = getattr(context, "mcp_client", None)
if not mcp_enabled and previous_mcp_client:
initialize_task = getattr(
context, "mcp_initialize_task", None
)
if initialize_task and not initialize_task.done():
initialize_task.cancel()
await asyncio.gather(
initialize_task, return_exceptions=True
)
context.mcp_initialize_task = None
previous_mcp_cleanup = getattr(
context, "_mcp_cleanup_callback", None
)
if previous_mcp_cleanup:
context.unregister_cleanup(previous_mcp_cleanup)
context._mcp_cleanup_callback = None
if hasattr(previous_mcp_client, "close"):
await previous_mcp_client.close()
context.mcp_client = None
if features:
logger.info(f"客户端特性: {features}")
if mcp_enabled:
logger.info("客户端支持MCP")
if context.mcp_client is None:
context.mcp_client = MCPClient()
context._mcp_cleanup_callback = (
context.mcp_client.close
)
context.register_cleanup(
context._mcp_cleanup_callback
)
initialize_task = getattr(
context, "mcp_initialize_task", None
)
if (
not await context.mcp_client.is_ready()
and (
initialize_task is None
or initialize_task.done()
)
):
context.mcp_initialize_task = (
context.create_background_task(
send_mcp_initialize_message(
context, transport
),
conversation_scoped=False,
)
)
if features.get("aec"):
logger.info("客户端启用了服务端AEC")
# MQTT/UDP 场景已由MQTT连接层发送hello回复,避免重复发送websocket格式
if transport.has_datagram_audio:
return
# 发送欢迎消息
if context.welcome_msg:
await transport.send_json(context.welcome_msg)
else:
# 默认欢迎消息
welcome_msg = {
"type": "hello",
"session_id": context.session_id,
"version": 1,
"transport": "websocket"
}
await transport.send_json(welcome_msg)
async def check_wakeup_words(self, context: SessionContext, transport: TransportInterface, text: str) -> bool:
"""检查唤醒词 - 完整迁移自checkWakeupWords"""
enable_wakeup_words_response_cache = context.config.get("enable_wakeup_words_response_cache", False)
# 等待tts初始化,最多等待3秒
tts_component = await self._get_component(context, ComponentType.TTS)
start_time = time.time()
while time.time() - start_time < 3:
if tts_component and hasattr(tts_component, 'tts_instance'):
break
await asyncio.sleep(0.1)
else:
return False
if not enable_wakeup_words_response_cache:
return False
_, filtered_text = remove_punctuation_and_length(text)
if filtered_text not in context.config.get("wakeup_words", []):
return False
sentence_id = uuid.uuid4().hex
context.sentence_id = sentence_id
context.just_woken_up = True
await self._send_stt_message(context, transport, text)
# 获取当前音色
tts_instance = getattr(tts_component, 'tts_instance', None) if tts_component else None
voice = getattr(tts_instance, "voice", "default") if tts_instance else "default"
if not voice:
voice = "default"
# 获取唤醒词回复配置
response = wakeup_words_config.get_wakeup_response(voice)
if not response or not response.get("file_path"):
response = {
"voice": "default",
"file_path": "config/assets/wakeup_words.wav",
"time": 0,
"text": "哈啰啊,我是小智啦,声音好听的台湾女孩一枚,超开心认识你耶,最近在忙啥,别忘了给我来点有趣的料哦,我超爱听八卦的啦",
}
# 获取音频数据
opus_packets = await audio_to_data(response.get("file_path"))
# 播放唤醒词回复
context.abort_requested = False
context.llm_finish_task = True
if tts_instance is not None:
# A prior turn may have consumed the one-shot start marker. Cached
# playback is a new TTS stream and must always emit start first.
tts_instance.tts_audio_first_sentence = True
logger.info(f"播放唤醒词回复: {response.get('text')}")
try:
await self._send_audio_message(
context,
transport,
SentenceType.FIRST,
opus_packets,
response.get("text"),
sentence_id,
)
await self._send_audio_message(
context,
transport,
SentenceType.LAST,
[],
None,
sentence_id,
)
finally:
# Cached wake playback bypasses ListenProcessor, so it must arm
# the same bounded release used by the regular wake-word path.
arm_wake_audio_suppression(context)
# 补充对话
if context.dialogue:
context.dialogue.put(Message(role="assistant", content=response.get("text")))
# 检查是否需要更新唤醒词回复
if time.time() - response.get("time", 0) > WAKEUP_CONFIG["refresh_time"]:
if not _wakeup_response_lock.locked():
context.create_background_task(
self._wakeup_words_response(context, transport)
)
return True
async def _wakeup_words_response(self, context: SessionContext, transport: TransportInterface):
"""生成唤醒词回复 - 完整迁移自wakeupWordsResponse"""
tts_component = await self._get_component(context, ComponentType.TTS)
llm_component = await self._get_component(context, ComponentType.LLM)
tts_instance = getattr(tts_component, 'tts_instance', None) if tts_component else None
llm_instance = getattr(llm_component, 'llm_instance', None) if llm_component else None
if not tts_instance or not llm_instance or not hasattr(llm_instance, 'response_no_stream'):
return
try:
# 尝试获取锁,如果获取不到就返回
async with _wakeup_response_lock:
# 生成唤醒词回复
wakeup_word = random.choice(WAKEUP_CONFIG["words"])
question = (
"此刻用户正在和你说```"
+ wakeup_word
+ "```。\n请你根据以上用户的内容进行20-30字回复。要符合系统设置的角色情感和态度,不要像机器人一样说话。\n"
+ "请勿对这条内容本身进行任何解释和回应,请勿返回表情符号,仅返回对用户的内容的回复。"
)
result = await asyncio.to_thread(
llm_instance.response_no_stream,
context.config.get("prompt", ""),
question,
)
if not result or len(result) == 0:
return
# 生成TTS音频
tts_result = await asyncio.to_thread(tts_instance.to_tts, result)
if not tts_result:
return
# 获取当前音色
voice = getattr(tts_instance, "voice", "default")
wav_bytes = opus_datas_to_wav_bytes(tts_result, sample_rate=16000)
file_path = wakeup_words_config.generate_file_path(voice)
with open(file_path, "wb") as f:
f.write(wav_bytes)
# 更新配置
wakeup_words_config.update_wakeup_response(voice, file_path, result)
except Exception as e:
logger.error(f"生成唤醒词回复失败: {e}")
async def _send_stt_message(self, context: SessionContext, transport: TransportInterface, text: str):
"""发送STT消息"""
from core.processors.audio_send_processor import AudioSendProcessor
await AudioSendProcessor().send_stt_message(
context, transport, text
)
async def _send_audio_message(self, context: SessionContext, transport: TransportInterface,
sentence_type: SentenceType, audios: bytes, text: str,
sentence_id: str = None):
"""发送音频消息"""
# 这里应该调用AudioSendProcessor
from core.processors.audio_send_processor import AudioSendProcessor
audio_send_processor = AudioSendProcessor()
await audio_send_processor.send_audio_message(
context,
transport,
sentence_type,
audios,
text,
sentence_id=sentence_id,
)
async def _get_component(self, context: SessionContext, component_type: ComponentType):
if not context.component_manager:
return None
return await context.component_manager.get_component(component_type, context)
@@ -0,0 +1,124 @@
import json
from typing import Any
from core.pipeline.message_pipeline import MessageProcessor
from core.context.session_context import SessionContext
from core.transport.transport_interface import TransportInterface
from core.providers.tools.device_iot import handleIotStatus, handleIotDescriptors
from config.logger import setup_logging
logger = setup_logging()
class IotProcessor(MessageProcessor):
"""IoT消息处理器:完整迁移iotMessageHandler.py的所有功能"""
async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
"""处理iot类型的消息"""
msg_json = None
if isinstance(message, str):
try:
msg_json = json.loads(message)
except json.JSONDecodeError:
msg_json = None
elif isinstance(message, dict):
msg_json = message
if isinstance(msg_json, dict) and msg_json.get("type") == "iot":
await self.handle_iot_message(context, transport, msg_json)
return True
return False
async def handle_iot_message(self, context: SessionContext, transport: TransportInterface, msg_json: dict):
"""处理IoT消息 - 完整迁移自iotMessageHandler.py"""
tasks = []
# 处理设备描述符 - 完整迁移原逻辑
if "descriptors" in msg_json:
logger.debug("处理IoT设备描述符")
task = context.create_background_task(
self._handle_iot_descriptors(context, transport, msg_json["descriptors"])
)
tasks.append(task)
# 处理设备状态 - 完整迁移原逻辑
if "states" in msg_json:
logger.debug("处理IoT设备状态")
task = context.create_background_task(
self._handle_iot_status(context, transport, msg_json["states"])
)
tasks.append(task)
# 如果没有有效的IoT数据
if not tasks:
logger.warning("IoT消息缺少descriptors或states字段")
await self._send_error_response(
transport,
context.session_id,
"IoT消息格式错误:缺少descriptors或states字段"
)
return
# 任务由 SessionContext 持有,结束会话时统一取消。
async def _handle_iot_descriptors(self, context: SessionContext, transport: TransportInterface, descriptors: Any):
"""处理IoT设备描述符 - 包装原handleIotDescriptors函数"""
try:
# 调用原有的handleIotDescriptors函数
# 注意:这里需要传入context而不是conn,因为handleIotDescriptors可能需要适配
await handleIotDescriptors(context, descriptors)
logger.debug("IoT设备描述符处理完成")
except Exception as e:
logger.error(f"处理IoT设备描述符失败: {e}", exc_info=True)
await self._send_error_response(
transport,
context.session_id,
f"IoT设备描述符处理失败: {str(e)}"
)
async def _handle_iot_status(self, context: SessionContext, transport: TransportInterface, states: Any):
"""处理IoT设备状态 - 包装原handleIotStatus函数"""
try:
# 调用原有的handleIotStatus函数
# 注意:这里需要传入context而不是conn,因为handleIotStatus可能需要适配
await handleIotStatus(context, states)
logger.debug("IoT设备状态处理完成")
except Exception as e:
logger.error(f"处理IoT设备状态失败: {e}", exc_info=True)
await self._send_error_response(
transport,
context.session_id,
f"IoT设备状态处理失败: {str(e)}"
)
async def _send_error_response(self, transport: TransportInterface, session_id: str, message: str):
"""发送IoT错误响应"""
response = {
"type": "iot",
"status": "error",
"message": message,
"session_id": session_id
}
try:
await transport.send(json.dumps(response))
except Exception as e:
logger.error(f"发送IoT错误响应失败: {e}")
async def _send_success_response(self, transport: TransportInterface, session_id: str,
message: str, data: dict = None):
"""发送IoT成功响应"""
response = {
"type": "iot",
"status": "success",
"message": message,
"session_id": session_id
}
if data:
response["data"] = data
try:
await transport.send(json.dumps(response))
except Exception as e:
logger.error(f"发送IoT成功响应失败: {e}")
@@ -0,0 +1,433 @@
import time
import json
import asyncio
import uuid
from typing import Any
from core.pipeline.message_pipeline import MessageProcessor
from core.context.session_context import SessionContext
from core.transport.transport_interface import TransportInterface
from core.utils.util import remove_punctuation_and_length
from core.utils.dialogue import Message
from core.components.component_manager import ComponentType
from core.providers.tts.dto.dto import ContentType, TTSMessageDTO, SentenceType
from core.processors.audio_receive_processor import arm_wake_audio_suppression
from config.logger import setup_logging
logger = setup_logging()
class ListenProcessor(MessageProcessor):
"""Listen消息处理器:完整迁移listenMessageHandler.py的所有功能"""
async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
"""处理listen类型的消息"""
msg_json = None
if isinstance(message, str):
try:
msg_json = json.loads(message)
except json.JSONDecodeError:
msg_json = None
elif isinstance(message, dict):
msg_json = message
if isinstance(msg_json, dict) and msg_json.get("type") == "listen":
await self.handle_listen_message(context, transport, msg_json)
return True
return False
async def handle_listen_message(self, context: SessionContext, transport: TransportInterface, msg_json: dict):
"""处理listen消息 - 完整迁移自listenMessageHandler.py"""
msg_session_id = msg_json.get("session_id")
if msg_session_id and msg_session_id != context.session_id:
logger.warning(
f"忽略非当前会话的listen消息: "
f"msg={msg_session_id}, current={context.session_id}"
)
return
state = msg_json.get("state")
has_datagram_audio = getattr(
transport, "has_datagram_audio", False
)
if (
has_datagram_audio
and state in {"start", "detect"}
and not context.conversation_active
):
if transport.keeps_connection_between_sessions:
await transport.send_json(
{
"type": "goodbye",
"session_id": msg_session_id or context.session_id,
}
)
logger.info(
"忽略已结束会话的listen {}: session_id={}",
state,
context.session_id,
)
return
# 设置拾音模式
if "mode" in msg_json:
context.listen_mode = msg_json["mode"]
logger.debug(f"客户端拾音模式:{context.listen_mode}")
# 处理不同的状态
if state == "start":
pending_start = getattr(context, "listen_start_task", None)
if pending_start and not pending_start.done():
logger.debug("忽略尾包隔离期内的重复listen start")
return
tail_quarantine = await self._flush_pending_listen_stop(
context, transport
)
# 开始监听语音
logger.info(f"listen start: session_id={context.session_id}, mode={context.listen_mode}")
context.reset_audio_states()
if (
getattr(transport, "requires_audio_tail_grace", False)
and tail_quarantine > 0
):
context.accepting_input_audio = False
context.listen_start_task = context.create_background_task(
self._open_input_after_tail_quarantine(
context,
context.session_id,
tail_quarantine,
),
turn_scoped=True,
)
else:
context.accepting_input_audio = True
context.abort_requested = False
context.close_after_chat = False
logger.debug("开始语音监听")
elif state == "stop":
# 停止监听语音
logger.info(f"listen stop: session_id={context.session_id}")
context.client_have_voice = True
context.listen_stop_pending = True
pending_start = getattr(context, "listen_start_task", None)
if pending_start and not pending_start.done():
pending_start.cancel()
context.listen_start_task = None
if getattr(transport, "requires_audio_tail_grace", False):
pending = getattr(context, "listen_stop_task", None)
if pending and not pending.done():
logger.debug("忽略重复的listen stop")
return
delay_ms = max(
0,
min(
1000,
int(
context.config.get(
"mqtt_udp_tail_grace_ms", 180
)
),
),
)
context.listen_stop_deadline = (
time.monotonic() + delay_ms / 1000
)
context.listen_stop_task = context.create_background_task(
self._finalize_listen_stop(
context,
transport,
context.session_id,
delay_ms / 1000,
),
turn_scoped=True,
)
else:
context.listen_stop_deadline = 0.0
await self._finalize_listen_stop(
context,
transport,
context.session_id,
0,
)
logger.debug("停止语音监听")
elif state == "detect":
# 检测到文本输入
logger.info(f"listen detect: session_id={context.session_id}, text={msg_json.get('text')}")
context.client_have_voice = False
context.asr_audio.clear()
if "text" in msg_json:
context.update_activity()
original_text = msg_json["text"] # 保留原始文本
filtered_len, filtered_text = remove_punctuation_and_length(original_text)
if original_text.startswith("[device_call]"):
await self._handle_device_call(
context,
transport,
original_text[len("[device_call]"):].strip(),
)
return
# 识别是否是唤醒词
is_wakeup_words = filtered_text in context.config.get("wakeup_words", [])
# 是否开启唤醒词回复
enable_greeting = context.config.get("enable_greeting", True)
if is_wakeup_words and not enable_greeting:
# 如果是唤醒词,且关闭了唤醒词回复,就不用回答
# Native already drops wake history until listen/start.
# Keeping the greeting suppression window armed here would
# discard the beginning of the user's first real utterance.
await self._send_stt_message(context, transport, original_text)
await self._send_tts_message(context, transport, "stop", None)
context.is_speaking = False
elif is_wakeup_words:
# 处理唤醒词
self._arm_wake_audio_suppression(context)
# 上报纯文字数据(复用ASR上报功能,但不提供音频数据)
await self._enqueue_asr_report(context, "嘿,你好呀", [])
await self._start_to_chat(context, transport, "嘿,你好呀", skip_intent=True)
else:
# 处理普通文本
# 上报纯文字数据(复用ASR上报功能,但不提供音频数据)
await self._enqueue_asr_report(context, original_text, [])
# 否则需要LLM对文字内容进行答复
await self._start_to_chat(context, transport, original_text)
async def _flush_pending_listen_stop(
self,
context: SessionContext,
transport: TransportInterface,
) -> float:
"""Finalize buffered speech before a new listen/start resets the turn."""
pending = getattr(context, "listen_stop_task", None)
had_pending_stop = bool(getattr(context, "listen_stop_pending", False))
tail_quarantine = max(
0.0,
float(getattr(context, "listen_stop_deadline", 0.0) or 0.0)
- time.monotonic(),
)
if pending and not pending.done():
pending.cancel()
try:
await pending
except asyncio.CancelledError:
pass
context.listen_stop_task = None
if had_pending_stop:
await self._finalize_listen_stop(
context,
transport,
context.session_id,
0,
)
else:
context.listen_stop_pending = False
if tail_quarantine <= 0:
context.listen_stop_deadline = 0.0
return tail_quarantine
@staticmethod
async def _open_input_after_tail_quarantine(
context: SessionContext,
session_id: str,
delay_seconds: float,
) -> None:
current_task = asyncio.current_task()
try:
await asyncio.sleep(delay_seconds)
if (
context.session_id == session_id
and context.conversation_active
and not context.listen_stop_pending
):
context.accepting_input_audio = True
finally:
if getattr(context, "listen_start_task", None) is current_task:
context.listen_start_task = None
context.listen_stop_deadline = 0.0
async def _finalize_listen_stop(
self,
context: SessionContext,
transport: TransportInterface,
session_id: str,
delay_seconds: float,
) -> None:
"""Finalize ASR after datagram tail frames have crossed the router."""
current_task = asyncio.current_task()
try:
if delay_seconds > 0:
await asyncio.sleep(delay_seconds)
if context.session_id != session_id:
return
if (
getattr(transport, "has_datagram_audio", False)
and not context.conversation_active
):
return
# Close the tail-grace ingress gate before resolving ASR.
# Component lookup may fail; leaving it open would leak subsequent
# packets into a turn that has already stopped.
context.listen_stop_pending = False
context.client_voice_stop = True
if getattr(transport, "requires_audio_tail_grace", False):
context.accepting_input_audio = False
asr_component = await self._get_component(
context, ComponentType.ASR
)
if context.session_id != session_id:
return
asr_instance = (
getattr(asr_component, "asr_instance", None)
if asr_component
else None
)
if not asr_instance or not hasattr(
asr_instance, "interface_type"
):
return
from core.providers.asr.dto.dto import InterfaceType
if asr_instance.interface_type == InterfaceType.STREAM:
if hasattr(asr_instance, "_send_stop_request"):
await asr_instance._send_stop_request()
return
if context.asr_audio:
asr_audio_task = context.asr_audio.copy()
context.asr_audio.clear()
context.reset_vad_states()
await asr_instance.handle_voice_stop(
context, asr_audio_task
)
except asyncio.CancelledError:
raise
except Exception as error:
logger.error(
"listen stop收尾失败: session_id={}, error={}",
session_id,
error,
)
finally:
if getattr(context, "listen_stop_task", None) is current_task:
context.listen_stop_task = None
context.listen_stop_deadline = 0.0
if context.session_id == session_id:
context.listen_stop_pending = False
if getattr(transport, "requires_audio_tail_grace", False):
context.accepting_input_audio = False
def _arm_wake_audio_suppression(self, context: SessionContext) -> None:
"""Bound the UDP/TCP reorder window used by the legacy gateway flow."""
arm_wake_audio_suppression(context)
async def _handle_device_call(
self,
context: SessionContext,
transport: TransportInterface,
call_text: str,
) -> None:
"""Restore the legacy device-call announcement and call-state handoff."""
logger.info(f"收到设备呼叫指令: {call_text}")
context.incoming_call = True
context.sentence_id = uuid.uuid4().hex
await self._send_stt_message(context, transport, call_text)
tts_component = await self._get_component(context, ComponentType.TTS)
tts_instance = (
getattr(tts_component, "tts_instance", None)
if tts_component
else None
)
if tts_instance:
tts_instance.store_tts_text(context.sentence_id, call_text)
tts_instance.tts_text_queue.put(
TTSMessageDTO(
sentence_id=context.sentence_id,
sentence_type=SentenceType.FIRST,
content_type=ContentType.ACTION,
)
)
tts_instance.tts_one_sentence(
context,
ContentType.TEXT,
content_detail=call_text,
)
tts_instance.tts_text_queue.put(
TTSMessageDTO(
sentence_id=context.sentence_id,
sentence_type=SentenceType.LAST,
content_type=ContentType.ACTION,
)
)
context.dialogue.put(Message(role="assistant", content=call_text))
async def _handle_audio_message(self, context: SessionContext, transport: TransportInterface, audio: bytes):
"""处理音频消息 - 调用AudioReceiveProcessor"""
# 这里应该调用AudioReceiveProcessor来处理音频
from core.processors.audio_receive_processor import AudioReceiveProcessor
audio_processor = AudioReceiveProcessor()
await audio_processor.handle_audio_message(context, transport, audio)
async def _send_stt_message(self, context: SessionContext, transport: TransportInterface, text: str):
"""发送STT消息"""
from core.processors.audio_send_processor import AudioSendProcessor
await AudioSendProcessor().send_stt_message(
context, transport, text
)
async def _send_tts_message(self, context: SessionContext, transport: TransportInterface, state: str, text: str = None):
"""发送TTS消息"""
message = {
"type": "tts",
"state": state,
"session_id": context.session_id
}
if text:
message["text"] = text
await transport.send(json.dumps(message))
logger.debug(f"发送TTS消息: state={state}, text={text}")
async def _get_component(self, context: SessionContext, component_type: ComponentType):
if not context.component_manager:
return None
return await context.component_manager.get_component(component_type, context)
async def _enqueue_asr_report(self, context: SessionContext, text: str, audio_data: list):
"""ASR上报队列"""
if context.report_asr_enable:
from core.processors.report_processor import ReportProcessor
report_processor = ReportProcessor()
report_processor.enqueue_asr_report(context, text, audio_data)
async def _start_to_chat(self, context: SessionContext, transport: TransportInterface, text: str, skip_intent: bool = False):
"""开始聊天 - 调用ChatProcessor"""
# 与旧架构一致:先发送STT,再异步触发聊天,避免阻塞事件循环
await self._send_stt_message(context, transport, text)
from core.processors.chat_processor import ChatProcessor
chat_processor = ChatProcessor()
if hasattr(context, "create_background_task"):
context.create_background_task(
chat_processor.handle_chat(
context, transport, text, skip_intent=skip_intent
),
turn_scoped=True,
)
logger.info(f"chat任务已提交: session_id={context.session_id}")
else:
await chat_processor.handle_chat(context, transport, text, skip_intent=skip_intent)
@@ -0,0 +1,96 @@
import json
from typing import Any
from core.pipeline.message_pipeline import MessageProcessor
from core.context.session_context import SessionContext
from core.transport.transport_interface import TransportInterface
from core.providers.tools.device_mcp import handle_mcp_message
from config.logger import setup_logging
logger = setup_logging()
class McpProcessor(MessageProcessor):
"""MCP消息处理器:完整迁移mcpMessageHandler.py的所有功能"""
async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
"""处理mcp类型的消息"""
msg_json = None
if isinstance(message, str):
try:
msg_json = json.loads(message)
except json.JSONDecodeError:
msg_json = None
elif isinstance(message, dict):
msg_json = message
if isinstance(msg_json, dict) and msg_json.get("type") == "mcp":
await self.handle_mcp_message(context, transport, msg_json)
return True
return False
async def handle_mcp_message(self, context: SessionContext, transport: TransportInterface, msg_json: dict):
"""处理MCP消息 - 完整迁移自mcpMessageHandler.py"""
if "payload" in msg_json:
# 检查MCP客户端是否存在
if not context.mcp_client:
logger.warning("MCP客户端未初始化,无法处理MCP消息")
await self._send_error_response(transport, context.session_id, "MCP客户端未初始化")
return
# 创建异步任务处理MCP消息 - 完整迁移原逻辑
context.create_background_task(
self._handle_mcp_payload(
context, transport, msg_json["payload"]
),
conversation_scoped=False,
)
else:
logger.warning("MCP消息缺少payload字段")
await self._send_error_response(transport, context.session_id, "MCP消息格式错误:缺少payload")
async def _handle_mcp_payload(self, context: SessionContext, transport: TransportInterface, payload: dict):
"""处理MCP payload - 包装原handle_mcp_message函数"""
try:
# 调用原有的handle_mcp_message函数
# 注意:这里需要传入context而不是conn,因为handle_mcp_message可能需要适配
await handle_mcp_message(context, context.mcp_client, payload, transport)
logger.debug("MCP消息处理完成")
except Exception as e:
logger.error(f"处理MCP消息失败: {e}", exc_info=True)
await self._send_error_response(
transport,
context.session_id,
f"MCP消息处理失败: {str(e)}"
)
async def _send_error_response(self, transport: TransportInterface, session_id: str, message: str):
"""发送MCP错误响应"""
response = {
"type": "mcp",
"status": "error",
"message": message,
"session_id": session_id
}
try:
await transport.send(json.dumps(response))
except Exception as e:
logger.error(f"发送MCP错误响应失败: {e}")
async def _send_success_response(self, transport: TransportInterface, session_id: str,
message: str, data: dict = None):
"""发送MCP成功响应"""
response = {
"type": "mcp",
"status": "success",
"message": message,
"session_id": session_id
}
if data:
response["data"] = data
try:
await transport.send(json.dumps(response))
except Exception as e:
logger.error(f"发送MCP成功响应失败: {e}")
@@ -0,0 +1,127 @@
import json
from typing import Any, List
from core.pipeline.message_pipeline import MessageProcessor
from core.context.session_context import SessionContext
from core.transport.transport_interface import TransportInterface
from core.processors.hello_processor import HelloProcessor
from core.processors.listen_processor import ListenProcessor
from core.processors.audio_receive_processor import AudioReceiveProcessor
from core.processors.auth_processor import AuthProcessor
from core.processors.timeout_processor import TimeoutProcessor
from core.processors.server_processor import ServerProcessor
from core.processors.mcp_processor import McpProcessor
from core.processors.iot_processor import IotProcessor
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()
class MessageRouter(MessageProcessor):
"""
消息路由器协调所有独立的processor
按功能职责分离避免耦合每个processor专注单一职责
"""
def __init__(self):
# 初始化所有独立的processor
self.auth_processor = AuthProcessor()
self.timeout_processor = TimeoutProcessor()
self.abort_processor = AbortProcessor()
self.goodbye_processor = GoodbyeProcessor()
self.hello_processor = HelloProcessor()
self.listen_processor = ListenProcessor()
self.server_processor = ServerProcessor()
self.mcp_processor = McpProcessor()
self.iot_processor = IotProcessor()
self.audio_receive_processor = AudioReceiveProcessor()
self.text_processor = TextProcessor()
self.ping_processor = PingProcessor()
self.udp_timeout_processor = UdpTimeoutProcessor()
# 按优先级排序的processor列表
self.processors: List[MessageProcessor] = [
self.timeout_processor, # 首先检查超时
self.auth_processor, # 然后检查认证
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消息
self.mcp_processor, # MCP消息
self.iot_processor, # IoT消息
self.audio_receive_processor, # 音频消息
self.text_processor, # 纯文本消息(放在最后,作为兜底处理)
]
async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
"""
路由消息到合适的processor
每个processor专注处理自己的消息类型避免耦合
"""
# Audio activity is updated only after VAD confirms voice. Treating
# every silent frame as activity prevents logical-session timeouts.
is_audio = isinstance(message, bytes) or (
isinstance(message, dict) and message.get("type") == "audio"
)
if not is_audio:
context.update_activity()
# 按优先级顺序尝试每个processor
for processor in self.processors:
try:
if await processor.process(context, transport, message):
# 消息已被处理,记录日志并返回
logger.debug(f"消息被 {processor.__class__.__name__} 处理")
return True
except Exception as e:
logger.error(f"{processor.__class__.__name__} 处理消息时出错: {e}", exc_info=True)
continue
# 如果没有processor处理该消息,记录警告
if isinstance(message, str):
try:
msg_json = json.loads(message)
msg_type = msg_json.get("type", "unknown") if isinstance(msg_json, dict) else "non-dict"
logger.warning(f"未处理的消息类型: {msg_type}, 内容: {message[:100]}...")
except json.JSONDecodeError:
logger.warning(f"未处理的非JSON消息: {message[:100]}...")
elif isinstance(message, bytes):
logger.warning(f"未处理的二进制消息,大小: {len(message)} bytes")
else:
logger.warning(f"未处理的消息类型: {type(message)}, {message}")
return False
def add_processor(self, processor: MessageProcessor, priority: int = None):
"""
添加新的processor
priority: 优先级数字越小优先级越高None表示添加到末尾
"""
if priority is None:
self.processors.append(processor)
else:
self.processors.insert(priority, processor)
logger.info(f"添加processor: {processor.__class__.__name__}")
def remove_processor(self, processor_class):
"""移除指定类型的processor"""
self.processors = [p for p in self.processors if not isinstance(p, processor_class)]
logger.info(f"移除processor: {processor_class.__name__}")
def get_processor(self, processor_class):
"""获取指定类型的processor"""
for processor in self.processors:
if isinstance(processor, processor_class):
return processor
return None
def list_processors(self) -> List[str]:
"""列出所有processor的名称"""
return [processor.__class__.__name__ for processor in self.processors]
@@ -0,0 +1,40 @@
import json
import time
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 PingProcessor(MessageProcessor):
"""Handle the optional legacy JSON ping/pong control contract."""
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") != "ping":
return False
if not context.config.get("enable_websocket_ping", False):
logger.debug("WebSocket心跳功能未启用,忽略PING消息")
return True
await transport.send_json(
{
"type": "pong",
"timestamp": time.strftime("%Y-%m-%d %H:%M:%S", time.localtime()),
}
)
return True
@@ -0,0 +1,301 @@
import asyncio
import time
import queue
import threading
import json
from typing import Any, List
from core.pipeline.message_pipeline import MessageProcessor
from core.context.session_context import SessionContext
from core.transport.transport_interface import TransportInterface
from config.manage_api_client import ManageApiClient, report as manage_report
from config.logger import setup_logging
logger = setup_logging()
class ReportProcessor(MessageProcessor):
"""上报处理器:完整迁移reportHandle.py的所有功能"""
async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
"""这个处理器不直接处理消息,而是被其他处理器调用"""
return False
def enqueue_asr_report(self, context: SessionContext, text: str, audio_data: List[bytes]):
"""ASR上报队列 - 完整迁移自enqueue_asr_report"""
if not self._should_report(context, context.report_asr_enable):
return
report_time = int(time.time() * 1000)
if context.chat_history_conf != 2:
audio_data = None
# 将上报任务放入队列
self._enqueue_report(context, {
"type": 1, # 用户类型
"text": text,
"audio_data": audio_data,
"report_time": report_time,
"session_id": context.session_id,
"device_id": context.device_id,
})
# 确保上报线程已启动
self._ensure_report_thread(context)
def enqueue_tts_report(self, context: SessionContext, text: str, opus_data: bytes):
"""TTS上报队列 - 完整迁移自enqueue_tts_report"""
if not self._should_report(context, context.report_tts_enable):
return
report_time = int(time.time() * 1000)
if context.chat_history_conf != 2:
opus_data = None
# 将上报任务放入队列
self._enqueue_report(context, {
"type": 2, # 智能体类型
"text": text,
"audio_data": opus_data,
"report_time": report_time,
"session_id": context.session_id,
"device_id": context.device_id,
})
# 确保上报线程已启动
self._ensure_report_thread(context)
def enqueue_tool_report(
self,
context: SessionContext,
tool_name: str,
tool_input: dict,
tool_result: str = None,
report_tool_call: bool = True,
):
if not self._should_report(context, True):
return
timestamp = int(time.time() * 1000)
entries = []
if report_tool_call:
entries.append(
(
json.dumps(
[{"type": "tool", "text": f"{tool_name}({json.dumps(tool_input, ensure_ascii=False)})"}],
ensure_ascii=False,
),
timestamp,
)
)
if tool_result is not None:
entries.append(
(
json.dumps(
[{"type": "tool_result", "text": json.dumps({"result": str(tool_result)}, ensure_ascii=False)}],
ensure_ascii=False,
),
timestamp + 1,
)
)
for text, report_time in entries:
self._enqueue_report(
context,
{
"type": 3,
"text": text,
"audio_data": None,
"report_time": report_time,
"session_id": context.session_id,
"device_id": context.device_id,
},
)
if entries:
self._ensure_report_thread(context)
def _ensure_report_thread(self, context: SessionContext):
"""确保上报线程已启动"""
if context.report_thread is None or not context.report_thread.is_alive():
if not getattr(context, "_report_cleanup_registered", False):
context.register_cleanup(
lambda: asyncio.to_thread(self.cleanup_session, context)
)
context._report_cleanup_registered = True
context.report_thread = threading.Thread(
target=self._report_worker,
args=(context,),
daemon=True
)
context.report_thread.start()
logger.info(f"上报线程已启动: {context.session_id}")
def _enqueue_report(self, context: SessionContext, report_task: dict) -> None:
try:
context.report_queue.put_nowait(report_task)
except queue.Full:
try:
context.report_queue.get_nowait()
context.report_queue.put_nowait(report_task)
logger.warning("上报队列已满,丢弃最旧任务: {}", context.session_id)
except (queue.Empty, queue.Full):
logger.warning("上报队列持续拥塞,丢弃当前任务: {}", context.session_id)
def _should_report(self, context: SessionContext, enabled: bool) -> bool:
return bool(
enabled
and context.read_config_from_api
and not context.need_bind
and context.chat_history_conf != 0
)
def _report_worker(self, context: SessionContext):
"""上报工作线程 - 完整迁移自ConnectionHandler中的上报逻辑"""
logger.info(f"上报工作线程启动: {context.session_id}")
event_loop = asyncio.new_event_loop()
asyncio.set_event_loop(event_loop)
try:
while not context.stop_event.is_set():
try:
# 从队列获取上报任务
report_task = context.report_queue.get(timeout=1)
if report_task is None:
break
# 执行上报
self._execute_report(context, report_task, event_loop)
except queue.Empty:
continue
except Exception as e:
logger.error(f"上报工作线程异常: {e}")
finally:
try:
event_loop.run_until_complete(
ManageApiClient.close_current_loop_client()
)
except Exception as exc:
logger.debug("关闭上报HTTP客户端失败: {}", exc)
event_loop.close()
logger.info(f"上报工作线程退出: {context.session_id}")
def _execute_report(
self, context: SessionContext, report_task: dict,
event_loop=None
):
"""执行聊天记录上报操作 - 完整迁移自report函数"""
try:
report_type = report_task["type"]
text = report_task["text"]
audio_data = report_task["audio_data"]
report_time = report_task["report_time"]
# 处理音频数据
processed_audio = None
if audio_data:
if isinstance(audio_data, list):
# ASR音频数据(多个音频片段)
processed_audio = self._process_asr_audio(audio_data)
elif isinstance(audio_data, bytes):
# TTS音频数据(opus格式)
processed_audio = self._opus_to_wav(audio_data)
# 执行上报
coroutine = manage_report(
mac_address=report_task["device_id"],
session_id=report_task["session_id"],
chat_type=report_type,
content=text,
audio=processed_audio,
report_time=report_time,
)
if event_loop is None:
asyncio.run(coroutine)
else:
event_loop.run_until_complete(coroutine)
logger.debug(f"上报成功: type={report_type}, text={text[:50]}...")
except Exception as e:
logger.error(f"聊天记录上报失败: {e}")
def _process_asr_audio(self, audio_data_list: List[bytes]) -> bytes:
"""处理ASR音频数据"""
try:
# 将多个音频片段合并
combined_audio = b''.join(audio_data_list)
return self._pcm_to_wav(combined_audio)
except Exception as e:
logger.error(f"处理ASR音频数据失败: {e}")
return b''
def _pcm_to_wav(self, pcm_data: bytes) -> bytes:
import io
import wave
wav_buffer = io.BytesIO()
with wave.open(wav_buffer, 'wb') as wav_file:
wav_file.setnchannels(1)
wav_file.setsampwidth(2)
wav_file.setframerate(16000)
wav_file.writeframes(pcm_data)
return wav_buffer.getvalue()
def _opus_to_wav(self, opus_data: bytes) -> bytes:
"""将Opus数据转换为WAV格式的字节流 - 完整迁移自opus_to_wav"""
try:
import opuslib_next
import io
import wave
# Opus解码器配置
sample_rate = 16000
channels = 1
# 创建Opus解码器
decoder = opuslib_next.Decoder(sample_rate, channels)
# 解码Opus数据
pcm_data = decoder.decode(opus_data, frame_size=960)
# 创建WAV文件
wav_buffer = io.BytesIO()
with wave.open(wav_buffer, 'wb') as wav_file:
wav_file.setnchannels(channels)
wav_file.setsampwidth(2) # 16-bit
wav_file.setframerate(sample_rate)
wav_file.writeframes(pcm_data)
return wav_buffer.getvalue()
except Exception as e:
logger.error(f"Opus转WAV失败: {e}")
return b''
def cleanup_session(self, context: SessionContext):
"""清理会话上报资源"""
# 停止上报线程
if context.report_thread and context.report_thread.is_alive():
context.stop_event.set()
try:
context.report_queue.put_nowait(None)
except queue.Full:
try:
context.report_queue.get_nowait()
context.report_queue.put_nowait(None)
except (queue.Empty, queue.Full):
pass
context.report_thread.join(timeout=5)
if context.report_thread.is_alive():
logger.warning(f"上报线程未能按时退出: {context.session_id}")
else:
context.report_thread = None
# 清理上报队列
try:
while not context.report_queue.empty():
context.report_queue.get_nowait()
except queue.Empty:
pass
logger.info(f"上报资源清理完成: {context.session_id}")
@@ -0,0 +1,177 @@
import json
import os
import sys
import time
import threading
import subprocess
from pathlib import Path
from typing import Any
from core.pipeline.message_pipeline import MessageProcessor
from core.context.session_context import SessionContext
from core.transport.transport_interface import TransportInterface
from config.logger import setup_logging
from core.utils.restart import restart_server
logger = setup_logging()
class ServerProcessor(MessageProcessor):
"""服务器消息处理器:完整迁移serverMessageHandler.py的所有功能"""
async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
"""处理server类型的消息"""
msg_json = None
if isinstance(message, str):
try:
msg_json = json.loads(message)
except json.JSONDecodeError:
msg_json = None
elif isinstance(message, dict):
msg_json = message
if isinstance(msg_json, dict) and msg_json.get("type") == "server":
await self.handle_server_message(context, transport, msg_json)
return True
return False
async def handle_server_message(self, context: SessionContext, transport: TransportInterface, msg_json: dict):
"""处理server消息 - 完整迁移自serverMessageHandler.py"""
# 如果配置是从API读取的,则需要验证secret
if not context.read_config_from_api:
return
# 获取post请求的secret
post_secret = msg_json.get("content", {}).get("secret", "")
secret = context.config.get("manager-api", {}).get("secret", "")
# 如果secret不匹配,则返回
if post_secret != secret:
await self._send_error_response(
transport,
context.session_id,
"服务器密钥验证失败"
)
await self._close_transport(transport)
return
# 处理不同的action
action = msg_json.get("action")
if action == "update_config":
await self._handle_update_config(context, transport, msg_json)
elif action == "restart":
await self._handle_restart(context, transport, msg_json)
else:
await self._send_error_response(
transport,
context.session_id,
f"未知的服务器操作: {action}"
)
await self._close_transport(transport)
async def _handle_update_config(self, context: SessionContext, transport: TransportInterface, msg_json: dict):
"""处理配置更新 - 完整迁移自update_config逻辑"""
try:
# 检查是否有服务器实例
if not context.server:
await self._send_error_response(
transport,
context.session_id,
"无法获取服务器实例",
{"action": "update_config"}
)
await self._close_transport(transport)
return
# 更新WebSocketServer的配置
if not await context.server.update_config():
error_msg = ""
if hasattr(context.server, "get_last_update_error"):
error_msg = context.server.get_last_update_error()
await self._send_error_response(
transport,
context.session_id,
error_msg or "更新服务器配置失败",
{"action": "update_config"}
)
await self._close_transport(transport)
return
# 发送成功响应
await self._send_success_response(
transport,
context.session_id,
"配置更新成功",
{"action": "update_config"}
)
await self._close_transport(transport)
except Exception as e:
logger.error(f"更新配置失败: {str(e)}")
await self._send_error_response(
transport,
context.session_id,
f"更新配置失败: {str(e)}",
{"action": "update_config"}
)
await self._close_transport(transport)
async def _handle_restart(self, context: SessionContext, transport: TransportInterface, msg_json: dict):
"""处理服务器重启 - 完整迁移自handle_restart逻辑"""
try:
# 发送确认响应
await self._send_success_response(
transport,
context.session_id,
"服务器重启中...",
{"action": "restart"}
)
await self._close_transport(transport)
# 异步执行重启操作
threading.Thread(target=lambda : restart_server(logger), daemon=True).start()
except Exception as e:
logger.error(f"处理重启请求失败: {str(e)}")
await self._send_error_response(
transport,
context.session_id,
f"重启失败: {str(e)}",
{"action": "restart"}
)
await self._close_transport(transport)
async def _send_success_response(self, transport: TransportInterface, session_id: str,
message: str, content: dict = None):
"""发送成功响应"""
response = {
"type": "server",
"status": "success",
"message": message,
"session_id": session_id
}
if content:
response["content"] = content
await transport.send(json.dumps(response))
logger.info(f"服务器操作成功: {message}")
async def _send_error_response(self, transport: TransportInterface, session_id: str,
message: str, content: dict = None):
"""发送错误响应"""
response = {
"type": "server",
"status": "fail",
"message": message,
"session_id": session_id
}
if content:
response["content"] = content
await transport.send(json.dumps(response))
logger.error(f"服务器操作失败: {message}")
async def _close_transport(self, transport: TransportInterface):
try:
await transport.close()
except Exception:
pass
@@ -0,0 +1,54 @@
import json
from typing import Any
from core.pipeline.message_pipeline import MessageProcessor
from core.context.session_context import SessionContext
from core.transport.transport_interface import TransportInterface
from config.logger import setup_logging
logger = setup_logging()
class TextProcessor(MessageProcessor):
"""
纯文本消息处理器处理非JSON格式的文本消息
这是新架构中缺失的重要组件用于处理直接发送的文本聊天内容
"""
async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
"""处理纯文本消息"""
if isinstance(message, str):
try:
# 尝试解析为JSON,如果成功则不是纯文本消息
json.loads(message)
return False # JSON消息由其他processor处理
except json.JSONDecodeError:
# 确实是纯文本消息,进行聊天处理
await self.handle_text_message(context, transport, message)
return True
return False
async def handle_text_message(self, context: SessionContext, transport: TransportInterface, text: str):
"""处理纯文本消息 - 直接调用ChatProcessor进行聊天"""
try:
# 记录收到纯文本消息
logger.info(f"收到纯文本消息: {text[:100]}...")
# 使用ChatProcessor处理聊天
from core.processors.chat_processor import ChatProcessor
chat_processor = ChatProcessor()
await chat_processor.handle_chat(context, transport, text)
except Exception as e:
logger.error(f"处理纯文本消息失败: {e}")
# 发送错误响应
await self._send_error_response(transport, "文本处理失败,请重试")
async def _send_error_response(self, transport: TransportInterface, error_message: str):
"""发送错误响应"""
try:
await transport.send(json.dumps({
"type": "error",
"message": error_message
}))
except Exception as e:
logger.error(f"发送错误响应失败: {e}")
@@ -0,0 +1,48 @@
from typing import Any
from core.pipeline.message_pipeline import MessageProcessor
from core.context.session_context import SessionContext
from core.transport.transport_interface import TransportInterface
from config.logger import setup_logging
logger = setup_logging()
class TimeoutProcessor(MessageProcessor):
"""超时检查处理器:检查会话是否超时"""
async def process(self, context: SessionContext, transport: TransportInterface, message: Any) -> bool:
"""检查会话超时"""
return await self.handle_timeout(context, transport)
async def handle_timeout(
self,
context: SessionContext,
transport: TransportInterface,
) -> bool:
"""End an expired logical conversation without conflating it with MQTT."""
if not getattr(context, "conversation_active", False):
return False
timeout_seconds = context.config.get("close_connection_no_voice_time", 120)
if not context.is_timeout(timeout_seconds):
return False
try:
if transport.keeps_connection_between_sessions:
logger.info(f"会话超时,结束MQTT逻辑会话: {context.session_id}")
from core.processors.audio_receive_processor import (
AudioReceiveProcessor,
)
await AudioReceiveProcessor()._no_voice_close_connect(
context,
transport,
have_voice=False,
)
else:
logger.info(f"会话超时,准备关闭连接: {context.session_id}")
await transport.close()
except Exception as e:
logger.error(f"处理超时连接失败: {e}")
return True
@@ -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()
@@ -109,7 +109,7 @@ class ASRProvider(ASRProviderBase):
self.token, expire_time_str = AccessToken.create_token(self.access_key_id, self.access_key_secret)
if not self.token:
raise ValueError("无法获取有效的访问Token")
try:
expire_str = str(expire_time_str).strip()
if expire_str.isdigit():
@@ -151,7 +151,7 @@ class ASRProvider(ASRProviderBase):
"""开始识别会话"""
if self._is_token_expired():
self._refresh_token()
# 建立连接
headers = {"X-NLS-Token": self.token}
self.asr_ws = await websockets.connect(
@@ -169,7 +169,10 @@ class ASRProvider(ASRProviderBase):
self.is_processing = True
self.server_ready = False # 重置服务器准备状态
self.forward_task = asyncio.create_task(self._forward_results(conn))
session_id = getattr(conn, "session_id", None)
self.forward_task = self._create_session_task(
conn, self._forward_results(conn, session_id)
)
# 发送开始请求
start_request = {
@@ -193,10 +196,13 @@ class ASRProvider(ASRProviderBase):
await self.asr_ws.send(json.dumps(start_request, ensure_ascii=False))
logger.bind(tag=TAG).debug("已发送开始请求,等待服务器准备...")
async def _forward_results(self, conn: "ConnectionHandler"):
async def _forward_results(self, conn: "ConnectionHandler", session_id=None):
"""转发识别结果"""
try:
while not conn.stop_event.is_set():
while (
not conn.stop_event.is_set()
and self._session_is_current(conn, session_id)
):
# 获取当前连接的音频数据
audio_data = conn.asr_audio
try:
@@ -276,7 +282,7 @@ class ASRProvider(ASRProviderBase):
finally:
# 清理连接的音频缓存
await self._cleanup()
conn.reset_audio_states()
self._reset_audio_if_current(conn, session_id)
async def _send_stop_request(self):
"""发送停止识别请求(不关闭连接)"""
@@ -308,6 +314,21 @@ class ASRProvider(ASRProviderBase):
self.server_ready = False
logger.bind(tag=TAG).debug("ASR状态已重置")
forward_task = self.forward_task
current_task = asyncio.current_task()
if (
forward_task
and forward_task is not current_task
and not forward_task.done()
):
forward_task.cancel()
try:
await forward_task
except asyncio.CancelledError:
pass
except Exception as e:
logger.bind(tag=TAG).warning(f"等待ASR转发任务退出失败: {e}")
# 关闭连接
if self.asr_ws:
try:
@@ -319,8 +340,10 @@ class ASRProvider(ASRProviderBase):
finally:
self.asr_ws = None
# 清理任务引用
self.forward_task = None
# Never discard a live task reference. The current forward task reaches
# this branch from its own finally block and is already completing.
if self.forward_task is forward_task:
self.forward_task = None
logger.bind(tag=TAG).debug("ASR会话清理完成")
@@ -103,7 +103,10 @@ class ASRProvider(ASRProviderBase):
logger.bind(tag=TAG).debug("WebSocket连接建立成功")
self.server_ready = False
self.forward_task = asyncio.create_task(self._forward_results(conn))
session_id = getattr(conn, "session_id", None)
self.forward_task = self._create_session_task(
conn, self._forward_results(conn, session_id)
)
# 发送run-task指令
run_task_msg = self._build_run_task_message()
@@ -154,10 +157,13 @@ class ASRProvider(ASRProviderBase):
return message
async def _forward_results(self, conn: "ConnectionHandler"):
async def _forward_results(self, conn: "ConnectionHandler", session_id=None):
"""转发识别结果"""
try:
while not conn.stop_event.is_set():
while (
not conn.stop_event.is_set()
and self._session_is_current(conn, session_id)
):
# 获取当前连接的音频数据
audio_data = conn.asr_audio
try:
@@ -243,7 +249,7 @@ class ASRProvider(ASRProviderBase):
finally:
# 清理连接的音频缓存
await self._cleanup()
conn.reset_audio_states()
self._reset_audio_if_current(conn, session_id)
async def _send_stop_request(self):
"""发送停止请求(用于手动模式停止录音)"""
@@ -285,6 +291,21 @@ class ASRProvider(ASRProviderBase):
self.server_ready = False
logger.bind(tag=TAG).debug("ASR状态已重置")
forward_task = self.forward_task
current_task = asyncio.current_task()
if (
forward_task
and forward_task is not current_task
and not forward_task.done()
):
forward_task.cancel()
try:
await forward_task
except asyncio.CancelledError:
pass
except Exception as e:
logger.bind(tag=TAG).warning(f"等待ASR转发任务退出失败: {e}")
# 关闭连接
if self.asr_ws:
try:
@@ -301,8 +322,8 @@ class ASRProvider(ASRProviderBase):
finally:
self.asr_ws = None
# 清理任务引用
self.forward_task = None
if self.forward_task is forward_task:
self.forward_task = None
self.task_id = None
logger.bind(tag=TAG).debug("ASR会话清理完成")
@@ -315,4 +336,4 @@ class ASRProvider(ASRProviderBase):
async def close(self):
"""关闭资源"""
await self._cleanup()
await self._cleanup()
+54 -9
View File
@@ -32,8 +32,37 @@ class ASRProviderBase(ABC):
def __init__(self):
pass
@staticmethod
def _session_is_current(conn: "ConnectionHandler", session_id: str) -> bool:
"""Return whether an async ASR result still belongs to this session."""
return session_id is None or getattr(conn, "session_id", None) == session_id
def _create_session_task(self, conn: "ConnectionHandler", coroutine):
"""Bind streaming ASR work to the logical session when supported."""
creator = getattr(conn, "create_background_task", None)
if callable(creator):
try:
return creator(coroutine, turn_scoped=True)
except TypeError:
return creator(coroutine)
return asyncio.create_task(coroutine)
def _reset_audio_if_current(
self, conn: "ConnectionHandler", session_id: str
) -> None:
if not self._session_is_current(conn, session_id):
return
reset_audio_states = getattr(conn, "reset_audio_states", None)
if callable(reset_audio_states):
reset_audio_states()
# 打开音频通道
async def open_audio_channels(self, conn: "ConnectionHandler"):
# The pipeline runtime feeds PCM directly and does not use the legacy
# priority queue. Starting that worker would route frames back through
# ConnectionHandler-only functions.
if getattr(conn, "uses_pipeline_runtime", False):
return
conn.asr_priority_thread = threading.Thread(
target=self.asr_text_priority_thread, args=(conn,), daemon=True
)
@@ -72,18 +101,26 @@ class ASRProviderBase(ABC):
return
# 自动模式下通过VAD检测到语音停止时触发识别
if conn.asr.interface_type != InterfaceType.STREAM and conn.client_voice_stop:
interface_type = getattr(
self,
"interface_type",
getattr(getattr(conn, "asr", None), "interface_type", None),
)
if interface_type != InterfaceType.STREAM and conn.client_voice_stop:
# 直接使用asr_audio中的PCM数据
pcm_bytes = b"".join(conn.asr_audio)
# 检查是否有足够的音频数据(每帧1920字节,15帧约28800字节)
if len(pcm_bytes) > 1920 * 15:
await self.handle_voice_stop(conn, [pcm_bytes])
conn.reset_audio_states()
reset_audio_states = getattr(conn, "reset_audio_states", None)
if callable(reset_audio_states):
reset_audio_states()
# 处理语音停止
async def handle_voice_stop(self, conn: "ConnectionHandler", asr_audio_task: List[bytes]):
"""并行处理ASR和声纹识别"""
try:
session_id = getattr(conn, "session_id", None)
total_start_time = time.monotonic()
# 数据已经是PCM直接使用
@@ -96,13 +133,11 @@ class ASRProviderBase(ABC):
wav_data = self._pcm_to_wav(combined_pcm_data)
# 定义ASR任务
asr_task = self.speech_to_text_wrapper(
asr_audio_task, conn.session_id
)
asr_task = self.speech_to_text_wrapper(asr_audio_task, session_id)
if conn.voiceprint_provider and wav_data:
voiceprint_task = conn.voiceprint_provider.identify_speaker(
wav_data, conn.session_id
wav_data, session_id
)
# 并发等待两个结果
asr_result, voiceprint_result = await asyncio.gather(
@@ -112,6 +147,12 @@ class ASRProviderBase(ABC):
asr_result = await asr_task
voiceprint_result = None
if not self._session_is_current(conn, session_id):
logger.bind(tag=TAG).info(
f"丢弃旧会话ASR结果: session_id={session_id}"
)
return
# 记录识别结果 - 检查是否为异常
if isinstance(asr_result, Exception):
logger.bind(tag=TAG).error(f"ASR识别失败: {asr_result}")
@@ -165,9 +206,13 @@ class ASRProviderBase(ABC):
if text_len > 0:
audio_snapshot = asr_audio_task.copy()
enqueue_asr_report(conn, enhanced_text, audio_snapshot)
# 使用自定义模块进行上报
await startToChat(conn, enhanced_text)
result_handler = getattr(conn, "asr_result_handler", None)
if callable(result_handler):
await result_handler(enhanced_text, audio_snapshot)
else:
enqueue_asr_report(conn, enhanced_text, audio_snapshot)
# Legacy ConnectionHandler path.
await startToChat(conn, enhanced_text)
except Exception as e:
logger.bind(tag=TAG).error(f"处理语音停止失败: {e}")
import traceback
@@ -117,7 +117,10 @@ class ASRProvider(ASRProviderBase):
raise e
# 启动接收ASR结果的异步任务
self.forward_task = asyncio.create_task(self._forward_asr_results(conn))
session_id = getattr(conn, "session_id", None)
self.forward_task = self._create_session_task(
conn, self._forward_asr_results(conn, session_id)
)
# 发送缓存的音频数据
if conn.asr_audio and len(conn.asr_audio) > 0:
@@ -156,9 +159,13 @@ class ASRProvider(ASRProviderBase):
except Exception as e:
logger.bind(tag=TAG).info(f"发送音频数据时发生错误: {e}")
async def _forward_asr_results(self, conn: "ConnectionHandler"):
async def _forward_asr_results(self, conn: "ConnectionHandler", session_id=None):
try:
while self.asr_ws and not conn.stop_event.is_set():
while (
self.asr_ws
and not conn.stop_event.is_set()
and self._session_is_current(conn, session_id)
):
# 获取当前连接的音频数据
audio_data = conn.asr_audio
try:
@@ -249,21 +256,47 @@ class ASRProvider(ASRProviderBase):
if hasattr(e, "__cause__") and e.__cause__:
logger.bind(tag=TAG).error(f"错误原因: {str(e.__cause__)}")
finally:
if self.asr_ws:
await self.asr_ws.close()
self.asr_ws = None
self.is_processing = False
self._is_stopping = False
await self._cleanup()
# 重置所有音频相关状态
conn.reset_audio_states()
self._reset_audio_if_current(conn, session_id)
def stop_ws_connection(self):
if self.asr_ws:
asyncio.create_task(self.asr_ws.close())
self.asr_ws = None
# The forward task owns the WebSocket and closes it from _cleanup().
# Scheduling an untracked close here races with that cleanup path.
self.is_processing = False
self._is_stopping = False
async def _cleanup(self):
"""取消转发任务并关闭流式 ASR 连接。"""
self.is_processing = False
self._is_stopping = False
forward_task = self.forward_task
current_task = asyncio.current_task()
if (
forward_task
and forward_task is not current_task
and not forward_task.done()
):
forward_task.cancel()
try:
await forward_task
except asyncio.CancelledError:
pass
except Exception as e:
logger.bind(tag=TAG).warning(f"等待ASR转发任务退出失败: {e}")
if self.asr_ws:
try:
await asyncio.wait_for(self.asr_ws.close(), timeout=2.0)
except Exception as e:
logger.bind(tag=TAG).warning(f"关闭ASR WebSocket连接失败: {e}")
finally:
self.asr_ws = None
if self.forward_task is forward_task:
self.forward_task = None
async def _send_stop_request(self):
"""发送最后一个音频帧以通知服务器结束"""
self._is_stopping = True # 先标记为停止状态,阻止后续音频发送
@@ -417,14 +450,4 @@ class ASRProvider(ASRProviderBase):
async def close(self):
"""资源清理方法"""
if self.asr_ws:
await self.asr_ws.close()
self.asr_ws = None
if self.forward_task:
self.forward_task.cancel()
try:
await self.forward_task
except asyncio.CancelledError:
pass
self.forward_task = None
self.is_processing = False
await self._cleanup()
@@ -90,8 +90,11 @@ class ASRProvider(ASRProviderBase):
batch_size_s=60,
)
text = lang_tag_filter(result[0]["text"])
recognized_content = (
text.get("content", "") if isinstance(text, dict) else text
)
logger.bind(tag=TAG).debug(
f"语音识别耗时: {time.time() - start_time:.3f}s | 结果: {text['content']}"
f"语音识别耗时: {time.time() - start_time:.3f}s | 结果: {recognized_content}"
)
return text, artifacts.file_path
@@ -0,0 +1,522 @@
"""
SharedASRManager: 全局 ASR 管理器
实现单例模型 + 单推理执行器 + 队列限流
单例的原因是推理是 CPU/GPU-bound不是 I/O-bound多实例不仅会占用内存还会降低吞吐能力
"""
import asyncio
import copy
import json
from concurrent.futures import ThreadPoolExecutor
from typing import Dict, Any, Optional, Tuple, List
from config.logger import setup_logging
logger = setup_logging()
TAG = __name__
class SharedASRManager:
"""
全局共享 ASR 管理器
"""
# 支持预加载的本地模型类型
LOCAL_MODEL_TYPES = [
"fun_local", # FunASR 本地
"sherpa_onnx_local", # Sherpa ONNX
"sense_voice" # SenseVoice
]
def __init__(self, config: Dict[str, Any], asr_type: str = None):
"""
初始化 ASR 管理器
Args:
config: 服务器配置
asr_type: ASR 类型Optional用于显式指定
"""
self.config = config
self.asr_type = asr_type
# 模型实例(全局单例)
self.model_instance = None
# 任务队列(限流)
queue_max_size = self._get_queue_max_size()
self.task_queue: asyncio.Queue = asyncio.Queue(maxsize=queue_max_size)
# 推理锁(使得推理串行化)
self.inference_lock = asyncio.Lock()
# 线程池执行器,用于阻塞调用
self.executor: Optional[ThreadPoolExecutor] = None
# 运行状态
self.running = False
self._inference_task: Optional[asyncio.Task] = None
self.is_local_model = self._check_local_model()
self._variant_lock = asyncio.Lock()
self._variant_managers = {}
self._max_shared_models = max(
1, int(config.get("shared_asr_max_models", 3) or 3)
)
logger.bind(tag=TAG).info(
f"SharedASRManager 初始化完成, "
f"类型: {self.asr_type}, "
f"本地模型: {self.is_local_model}, "
f"队列大小: {queue_max_size}"
)
def _get_queue_max_size(self) -> int:
"""获取队列最大大小"""
# 尝试从配置获取
selected_asr = self.config.get("selected_module", {}).get("ASR")
if selected_asr:
asr_config = self.config.get("ASR", {}).get(selected_asr, {})
return asr_config.get("queue_max_size", 100)
return 100
def _check_local_model(self) -> bool:
"""检查是否为本地模型"""
if self.asr_type:
return self.asr_type in self.LOCAL_MODEL_TYPES
# 从配置推断
selected_asr = self.config.get("selected_module", {}).get("ASR")
if not selected_asr:
return False
asr_config = self.config.get("ASR", {}).get(selected_asr, {})
asr_type = asr_config.get("type", selected_asr)
self.asr_type = asr_type
return asr_type in self.LOCAL_MODEL_TYPES
async def initialize(self):
"""
初始化管理器
- 预加载模型
- 启动推理执行器
"""
if not self.is_local_model:
logger.bind(tag=TAG).info("非本地模型,跳过预加载")
return
if self.running:
logger.bind(tag=TAG).warning("管理器已在运行中")
return
try:
logger.bind(tag=TAG).info(f"开始预加载 ASR 模型: {self.asr_type}")
# 预加载模型
await self._preload_model()
# 启动推理执行器
self.running = True
self._inference_task = asyncio.create_task(self._inference_loop())
logger.bind(tag=TAG).info("ASR 模型预加载完成,推理执行器已启动")
except Exception as e:
logger.bind(tag=TAG).error(f"ASR 模型预加载失败: {e}")
raise
async def _preload_model(self):
"""在线程池中预加载模型"""
# 创建线程池
self.executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix="asr_worker")
loop = asyncio.get_event_loop()
self.model_instance = await loop.run_in_executor(
self.executor,
self._create_model_instance
)
logger.bind(tag=TAG).info("模型实例创建完成")
def _create_model_instance(self):
"""
实际创建模型实例在线程中执行
Returns:
ASR Provider 实例
"""
from core.utils.modules_initialize import initialize_asr
logger.bind(tag=TAG).info("正在创建 ASR 模型实例...")
instance = initialize_asr(self.config)
logger.bind(tag=TAG).info("ASR 模型实例创建成功")
return instance
async def submit_task(
self,
opus_data: List[bytes],
session_id: str,
audio_format: str = "opus"
) -> Tuple[Optional[str], Optional[str]]:
"""
提交推理任务
Args:
opus_data: 音频数据
session_id: 会话 ID
audio_format: 音频格式
Returns:
(识别文本, 文件路径)
Raises:
RuntimeError: 队列满或服务未运行
"""
if not self.running:
raise RuntimeError("ASR 服务未运行")
# 检查队列是否满(限流)
if self.task_queue.full():
queue_status = self.get_queue_status()
logger.bind(tag=TAG).warning(
f"ASR 队列已满: {queue_status}"
)
raise RuntimeError("ASR 服务繁忙,请稍后重试")
# 创建 Future 用于返回结果
result_future: asyncio.Future = asyncio.Future()
# 构造任务
task = {
'opus_data': opus_data,
'session_id': session_id,
'audio_format': audio_format,
'future': result_future
}
# 放入队列
await self.task_queue.put(task)
logger.bind(tag=TAG).debug(
f"任务已提交, session: {session_id}, "
f"队列大小: {self.task_queue.qsize()}"
)
# 等待结果
return await result_future
async def _inference_loop(self):
"""
单个推理执行器循环
核心原则
- 只有一个执行器
- 串行处理任务
- 带超时的队列获取
"""
logger.bind(tag=TAG).info("推理执行器启动")
while self.running:
task = None
try:
# 带超时的队列获取,避免关闭时卡住
try:
task = await asyncio.wait_for(
self.task_queue.get(),
timeout=1.0
)
except asyncio.TimeoutError:
# 超时后检查 running 状态,继续循环
continue
if task['future'].done():
continue
# 执行推理(加锁保证串行)
async with self.inference_lock:
if task['future'].done():
continue
result = await self._run_inference(
task['opus_data'],
task['session_id'],
task['audio_format']
)
# 返回结果
if not task['future'].done():
task['future'].set_result(result)
logger.bind(tag=TAG).debug(
f"推理完成, session: {task['session_id']}"
)
except asyncio.CancelledError:
logger.bind(tag=TAG).info("推理执行器被取消")
break
except Exception as e:
logger.bind(tag=TAG).error(f"推理执行失败: {e}")
if task and 'future' in task and not task['future'].done():
task['future'].set_exception(e)
finally:
if task is not None:
self.task_queue.task_done()
logger.bind(tag=TAG).info("推理执行器已停止")
async def _run_inference(
self,
opus_data: List[bytes],
session_id: str,
audio_format: str
) -> Tuple[Optional[str], Optional[str]]:
"""
执行实际推理在线程池中
Args:
opus_data: 音频数据
session_id: 会话 ID
audio_format: 音频格式
Returns:
(识别文本, 文件路径)
"""
loop = asyncio.get_event_loop()
# 在线程池中执行推理
result = await loop.run_in_executor(
self.executor,
lambda: self._sync_wrapper(opus_data, session_id)
)
return result
def _sync_wrapper(
self,
opus_data: List[bytes],
session_id: str,
) -> Tuple[Optional[str], Optional[str]]:
"""Run the provider's complete artifact wrapper in the worker thread."""
import asyncio
async def _call():
return await self.model_instance.speech_to_text_wrapper(
opus_data, session_id
)
# 创建新的事件循环执行
loop = None
try:
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
result = loop.run_until_complete(_call())
return result
finally:
if loop:
loop.close()
def matches_config(self, config: Dict[str, Any]) -> bool:
"""Return whether this manager owns the exact selected ASR config."""
selected = config.get("selected_module", {}).get("ASR")
manager_selected = self.config.get("selected_module", {}).get("ASR")
if not selected or selected != manager_selected:
return False
return (
config.get("ASR", {}).get(selected)
== self.config.get("ASR", {}).get(manager_selected)
and bool(config.get("delete_audio", True))
== bool(self.config.get("delete_audio", True))
)
@classmethod
def _config_fingerprint(cls, config: Dict[str, Any]) -> str:
selected = config.get("selected_module", {}).get("ASR")
payload = {
"selected": selected,
"config": config.get("ASR", {}).get(selected),
"delete_audio": bool(config.get("delete_audio", True)),
}
return json.dumps(payload, sort_keys=True, separators=(",", ":"))
@classmethod
def _config_is_local(cls, config: Dict[str, Any]) -> bool:
selected = config.get("selected_module", {}).get("ASR")
asr_config = config.get("ASR", {}).get(selected, {})
asr_type = asr_config.get("type", selected)
return asr_type in cls.LOCAL_MODEL_TYPES
async def acquire_for_config(
self, config: Dict[str, Any]
) -> Optional["SharedASRManager"]:
"""Acquire a bounded shared local model matching one Agent config."""
if self.matches_config(config):
return self if self.is_ready() else None
if not self._config_is_local(config):
return None
fingerprint = self._config_fingerprint(config)
async with self._variant_lock:
entry = self._variant_managers.get(fingerprint)
if entry is not None:
entry["references"] += 1
return entry["manager"]
if len(self._variant_managers) >= self._max_shared_models - 1:
idle_fingerprint = next(
(
key
for key, candidate in self._variant_managers.items()
if candidate["references"] == 0
),
None,
)
if idle_fingerprint is None:
raise RuntimeError(
"ASR共享模型容量已满,请提高shared_asr_max_models或统一Agent ASR配置"
)
idle_manager = self._variant_managers.pop(
idle_fingerprint
)["manager"]
# Keep the lock until shutdown completes so a replacement can
# never overlap the retiring model and exceed the hard limit.
await idle_manager.shutdown()
manager = SharedASRManager(copy.deepcopy(config))
try:
await manager.initialize()
except Exception:
await manager.shutdown()
raise
self._variant_managers[fingerprint] = {
"manager": manager,
"references": 1,
}
logger.bind(tag=TAG).info(
"已加载Agent专用共享ASR模型,当前模型数: {}",
len(self._variant_managers) + 1,
)
return manager
async def release_for_config(self, manager: "SharedASRManager") -> None:
"""Release an Agent model; idle variants stay cached for safe reuse."""
if manager is self:
return
async with self._variant_lock:
for entry in self._variant_managers.values():
if entry["manager"] is not manager:
continue
entry["references"] = max(0, entry["references"] - 1)
break
async def shutdown(self):
"""
优雅停机
步骤
1. 停止接收新任务
2. 等待当前任务完成带超时
3. 取消未完成的任务
4. 关闭线程池
"""
async with self._variant_lock:
variants = [
entry["manager"] for entry in self._variant_managers.values()
]
self._variant_managers.clear()
if variants:
await asyncio.gather(
*(manager.shutdown() for manager in variants),
return_exceptions=True,
)
if (
not self.running
and self.executor is None
and self.model_instance is None
):
return
logger.bind(tag=TAG).info("开始关闭 ASR 管理器...")
# 停止接收新任务
self.running = False
# 等待推理任务完成
if self._inference_task and not self._inference_task.done():
try:
# 最多等待 5 秒
await asyncio.wait_for(
self._inference_task,
timeout=5.0
)
except asyncio.TimeoutError:
logger.bind(tag=TAG).warning("推理任务超时,强制取消")
self._inference_task.cancel()
try:
await self._inference_task
except asyncio.CancelledError:
pass
except asyncio.CancelledError:
pass
# 取消所有队列中未完成的任务
cancelled_count = 0
while not self.task_queue.empty():
try:
task = self.task_queue.get_nowait()
if not task['future'].done():
task['future'].set_exception(
RuntimeError("ASR 服务正在关闭")
)
cancelled_count += 1
except asyncio.QueueEmpty:
break
if cancelled_count > 0:
logger.bind(tag=TAG).info(f"已取消 {cancelled_count} 个待处理任务")
# 关闭线程池
if self.executor:
self.executor.shutdown(wait=False)
self.executor = None
logger.bind(tag=TAG).info("线程池已关闭")
# 清理模型实例
self.model_instance = None
logger.bind(tag=TAG).info("ASR 管理器已关闭")
def get_queue_status(self) -> Dict[str, Any]:
"""
获取队列状态用于监控
Returns:
队列状态字典
"""
queue_size = self.task_queue.qsize()
queue_max = self.task_queue.maxsize
return {
'queue_size': queue_size,
'queue_max': queue_max,
'is_busy': queue_size > queue_max * 0.8,
'utilization': queue_size / queue_max if queue_max > 0 else 0,
'running': self.running
}
def is_ready(self) -> bool:
"""检查管理器是否就绪"""
return (
self.running and
self.model_instance is not None and
self.executor is not None
)
@classmethod
def is_local_model_type(cls, asr_type: str) -> bool:
"""
检查 ASR 类型是否为本地模型
Args:
asr_type: ASR 类型
Returns:
是否为本地模型
"""
return asr_type in cls.LOCAL_MODEL_TYPES
@@ -0,0 +1,119 @@
"""
SharedASRProxy: 共享 ASR 管理器的代理类
功能
- 包装 SharedASRManager
- 提供与原 ASR Provider 相同的接口
- 处理队列满等异常情况
"""
from typing import List, Tuple, Optional, Dict, Any
from core.providers.asr.base import ASRProviderBase
from core.providers.asr.dto.dto import InterfaceType
from config.logger import setup_logging
logger = setup_logging()
TAG = __name__
class SharedASRProxy(ASRProviderBase):
"""
共享 ASR 管理器的代理类
该类提供与原 ASR Provider 相同的接口
但实际推理工作由 SharedASRManager 完成
"""
def __init__(self, manager):
"""
初始化代理
Args:
manager: SharedASRManager 实例
"""
super().__init__()
self.manager = manager
# 从共享管理器获取接口类型
if manager.model_instance and hasattr(manager.model_instance, 'interface_type'):
self.interface_type = manager.model_instance.interface_type
else:
self.interface_type = InterfaceType.LOCAL
logger.bind(tag=TAG).info("SharedASRProxy 初始化完成")
async def speech_to_text(
self,
opus_data: List[bytes],
session_id: str,
audio_format: str = "opus"
) -> Tuple[Optional[str], Optional[str]]:
"""
语音转文本通过共享管理器
Args:
opus_data: 音频数据Opus 编码的字节列表
session_id: 会话 ID
audio_format: 音频格式默认 "opus"
Returns:
Tuple[str, str]: (识别的文本, 音频文件路径)
"""
try:
# 检查管理器状态
if not self.manager.is_ready():
logger.bind(tag=TAG).error("ASR 管理器未就绪")
return "", None
# 提交任务到共享管理器
result = await self.manager.submit_task(
opus_data,
session_id,
audio_format
)
return result
except RuntimeError as e:
# 队列满或服务未运行
logger.bind(tag=TAG).warning(f"ASR 服务繁忙: {e}")
# 返回友好提示,而不是空字符串
return "服务繁忙,请稍后重试", None
except Exception as e:
logger.bind(tag=TAG).error(f"ASR 推理失败: {e}")
return "", None
async def speech_to_text_wrapper(
self, pcm_data: List[bytes], session_id: str
) -> Tuple[Optional[str], Optional[str]]:
"""Delegate PCM inference without using uninitialized proxy file state."""
return await self.manager.submit_task(pcm_data, session_id, "pcm")
def get_queue_status(self) -> Dict[str, Any]:
"""
获取队列状态
Returns:
队列状态字典
"""
return self.manager.get_queue_status()
def is_ready(self) -> bool:
"""
检查代理是否就绪
Returns:
是否就绪
"""
return self.manager.is_ready()
async def close(self):
"""
关闭代理
注意不关闭共享管理器由服务器统一管理
"""
logger.bind(tag=TAG).debug("SharedASRProxy 关闭")
# 代理不负责关闭共享管理器
pass
@@ -141,7 +141,10 @@ class ASRProvider(ASRProviderBase):
logger.bind(tag=TAG).info("ASR WebSocket连接已建立")
self.server_ready = False
self.forward_task = asyncio.create_task(self._forward_results(conn))
session_id = getattr(conn, "session_id", None)
self.forward_task = self._create_session_task(
conn, self._forward_results(conn, session_id)
)
# 发送首帧音频
if conn.asr_audio and len(conn.asr_audio) > 0:
@@ -185,10 +188,13 @@ class ASRProvider(ASRProviderBase):
await self.asr_ws.send(json.dumps(frame_data, ensure_ascii=False))
async def _forward_results(self, conn: "ConnectionHandler"):
async def _forward_results(self, conn: "ConnectionHandler", session_id=None):
"""转发识别结果"""
try:
while not conn.stop_event.is_set():
while (
not conn.stop_event.is_set()
and self._session_is_current(conn, session_id)
):
try:
response = await asyncio.wait_for(self.asr_ws.recv(), timeout=60)
result = json.loads(response)
@@ -247,7 +253,7 @@ class ASRProvider(ASRProviderBase):
finally:
# 清理连接资源
await self._cleanup()
conn.reset_audio_states()
self._reset_audio_if_current(conn, session_id)
async def handle_voice_stop(
self, conn: "ConnectionHandler", asr_audio_task: List[bytes]
@@ -272,9 +278,8 @@ class ASRProvider(ASRProviderBase):
logger.bind(tag=TAG).debug(f"异常详情: {traceback.format_exc()}")
def stop_ws_connection(self):
if self.asr_ws:
asyncio.create_task(self.asr_ws.close())
self.asr_ws = None
# The forward task owns the WebSocket and closes it from _cleanup().
# Scheduling an untracked close here races with that cleanup path.
self.is_processing = False
async def _send_stop_request(self):
@@ -299,6 +304,21 @@ class ASRProvider(ASRProviderBase):
self.server_ready = False
logger.bind(tag=TAG).debug("ASR状态已重置")
forward_task = self.forward_task
current_task = asyncio.current_task()
if (
forward_task
and forward_task is not current_task
and not forward_task.done()
):
forward_task.cancel()
try:
await forward_task
except asyncio.CancelledError:
pass
except Exception as e:
logger.bind(tag=TAG).warning(f"等待ASR转发任务退出失败: {e}")
# 关闭连接
if self.asr_ws:
try:
@@ -310,8 +330,8 @@ class ASRProvider(ASRProviderBase):
finally:
self.asr_ws = None
# 清理任务引用
self.forward_task = None
if self.forward_task is forward_task:
self.forward_task = None
logger.bind(tag=TAG).debug("ASR会话清理完成")
@@ -323,15 +343,4 @@ class ASRProvider(ASRProviderBase):
async def close(self):
"""资源清理方法"""
if self.asr_ws:
await self.asr_ws.close()
self.asr_ws = None
if self.forward_task:
self.forward_task.cancel()
try:
await self.forward_task
except asyncio.CancelledError:
pass
self.forward_task = None
self.is_processing = False
await self._cleanup()
@@ -127,7 +127,16 @@ class DeviceIoTExecutor(ToolExecutor):
send_message = json.dumps(
{"type": "iot", "commands": [command]}
)
await self.conn.websocket.send(send_message)
# 使用transport接口发送消息
if hasattr(self.conn, 'transport') and self.conn.transport:
await self.conn.transport.send(send_message)
elif hasattr(self.conn, 'websocket') and self.conn.websocket:
# 兼容旧版本
logger.warning("未找到SessionContext的传输层接口, 回退使用旧版conn.websocket发送消息")
await self.conn.websocket.send(send_message)
else:
raise AttributeError("无法找到可用的传输层接口")
return
raise Exception(f"未找到设备{device_name}的方法{method_name}")
@@ -17,7 +17,7 @@ class MCPClient:
self.name_mapping = {}
self.ready = False
self.call_results = {} # To store Futures for tool call responses
self.next_id = 1
self.next_id = 10000
self.lock = asyncio.Lock()
self._cached_available_tools = None # Cache for get_available_tools
@@ -91,3 +91,12 @@ class MCPClient:
async with self.lock:
if id in self.call_results:
self.call_results.pop(id)
async def close(self):
async with self.lock:
pending = list(self.call_results.values())
self.call_results.clear()
self.ready = False
for future in pending:
if not future.done():
future.set_exception(ConnectionError("MCP会话已关闭"))
@@ -3,11 +3,11 @@
import json
import asyncio
import re
from concurrent.futures import Future
from core.utils.util import get_vision_url, sanitize_tool_name
from core.utils.util import get_vision_url
from core.utils.auth import AuthToken
from config.logger import setup_logging
from typing import TYPE_CHECKING
from .mcp_client import MCPClient
if TYPE_CHECKING:
from core.connection import ConnectionHandler
@@ -15,108 +15,45 @@ if TYPE_CHECKING:
TAG = __name__
logger = setup_logging()
class MCPClient:
"""设备端MCP客户端,用于管理MCP状态和工具"""
def __init__(self):
self.tools = {} # sanitized_name -> tool_data
self.name_mapping = {}
self.ready = False
self.call_results = {} # To store Futures for tool call responses
self.next_id = 1
self.lock = asyncio.Lock()
self._cached_available_tools = None # Cache for get_available_tools
def has_tool(self, name: str) -> bool:
return name in self.tools
def get_available_tools(self) -> list:
# Check if the cache is valid
if self._cached_available_tools is not None:
return self._cached_available_tools
# If cache is not valid, regenerate the list
result = []
for tool_name, tool_data in self.tools.items():
function_def = {
"name": tool_name,
"description": tool_data["description"],
"parameters": {
"type": tool_data["inputSchema"].get("type", "object"),
"properties": tool_data["inputSchema"].get("properties", {}),
"required": tool_data["inputSchema"].get("required", []),
},
}
result.append({"type": "function", "function": function_def})
self._cached_available_tools = result # Store the generated list in cache
return result
async def is_ready(self) -> bool:
async with self.lock:
return self.ready
async def set_ready(self, status: bool):
async with self.lock:
self.ready = status
async def add_tool(self, tool_data: dict):
async with self.lock:
sanitized_name = sanitize_tool_name(tool_data["name"])
self.tools[sanitized_name] = tool_data
self.name_mapping[sanitized_name] = tool_data["name"]
self._cached_available_tools = (
None # Invalidate the cache when a tool is added
)
async def get_next_id(self) -> int:
async with self.lock:
current_id = self.next_id
self.next_id += 1
return current_id
async def register_call_result_future(self, id: int, future: Future):
async with self.lock:
self.call_results[id] = future
async def resolve_call_result(self, id: int, result: any):
async with self.lock:
if id in self.call_results:
future = self.call_results.pop(id)
if not future.done():
future.set_result(result)
async def reject_call_result(self, id: int, exception: Exception):
async with self.lock:
if id in self.call_results:
future = self.call_results.pop(id)
if not future.done():
future.set_exception(exception)
async def cleanup_call_result(self, id: int):
async with self.lock:
if id in self.call_results:
self.call_results.pop(id)
async def send_mcp_message(conn: "ConnectionHandler", payload: dict):
async def send_mcp_message(
conn: "ConnectionHandler",
payload: dict,
transport=None,
*,
raise_on_error: bool = False,
):
"""Helper to send MCP messages, encapsulating common logic."""
if not conn.features.get("mcp"):
features = getattr(conn, "features", {}) or {}
if not features.get("mcp"):
logger.bind(tag=TAG).warning("客户端不支持MCP,无法发送MCP消息")
return
message = json.dumps({"type": "mcp", "payload": payload})
try:
await conn.websocket.send(message)
# 优先使用传入的transport,否则尝试从conn获取
if transport:
await transport.send(message)
elif getattr(conn, "transport", None):
# 新架构
await conn.transport.send(message)
elif getattr(conn, "websocket", None):
# 兼容旧版本
await conn.websocket.send(message)
else:
raise AttributeError("无法找到可用的传输层接口")
logger.bind(tag=TAG).debug(f"成功发送MCP消息: {message}")
except Exception as e:
logger.bind(tag=TAG).error(f"发送MCP消息失败: {e}")
if raise_on_error:
raise
async def handle_mcp_message(
conn: "ConnectionHandler", mcp_client: MCPClient, payload: dict
conn: "ConnectionHandler",
mcp_client: MCPClient,
payload: dict,
transport=None,
):
"""处理MCP消息,包括初始化、工具列表和工具调用响应等"""
logger.bind(tag=TAG).debug(f"处理MCP消息: {str(payload)[:100]}")
@@ -150,7 +87,7 @@ async def handle_mcp_message(
await asyncio.sleep(1)
logger.bind(tag=TAG).debug("初始化完成,开始请求MCP工具列表")
await send_mcp_tools_list_request(conn)
await send_mcp_tools_list_request(conn, transport)
return
@@ -207,7 +144,7 @@ async def handle_mcp_message(
next_cursor = result.get("nextCursor", "")
if next_cursor:
logger.bind(tag=TAG).debug(f"有更多工具,nextCursor: {next_cursor}")
await send_mcp_tools_list_continue_request(conn, next_cursor)
await send_mcp_tools_list_continue_request(conn, next_cursor, transport)
else:
await mcp_client.set_ready(True)
logger.bind(tag=TAG).debug("所有工具已获取,MCP客户端准备就绪")
@@ -235,7 +172,9 @@ async def handle_mcp_message(
)
async def send_mcp_initialize_message(conn: "ConnectionHandler"):
async def send_mcp_initialize_message(
conn: "ConnectionHandler", transport=None
):
"""发送MCP初始化消息"""
vision_url = get_vision_url(conn.config)
@@ -267,10 +206,12 @@ async def send_mcp_initialize_message(conn: "ConnectionHandler"):
},
}
logger.bind(tag=TAG).debug("发送MCP初始化消息")
await send_mcp_message(conn, payload)
await send_mcp_message(conn, payload, transport)
async def send_mcp_tools_list_request(conn: "ConnectionHandler"):
async def send_mcp_tools_list_request(
conn: "ConnectionHandler", transport=None
):
"""发送MCP工具列表请求"""
payload = {
"jsonrpc": "2.0",
@@ -278,10 +219,12 @@ async def send_mcp_tools_list_request(conn: "ConnectionHandler"):
"method": "tools/list",
}
logger.bind(tag=TAG).debug("发送MCP工具列表请求")
await send_mcp_message(conn, payload)
await send_mcp_message(conn, payload, transport)
async def send_mcp_tools_list_continue_request(conn: "ConnectionHandler", cursor: str):
async def send_mcp_tools_list_continue_request(
conn: "ConnectionHandler", cursor: str, transport=None
):
"""发送带有cursor的MCP工具列表请求"""
payload = {
"jsonrpc": "2.0",
@@ -290,7 +233,7 @@ async def send_mcp_tools_list_continue_request(conn: "ConnectionHandler", cursor
"params": {"cursor": cursor},
}
logger.bind(tag=TAG).info(f"发送带cursor的MCP工具列表请求: {cursor}")
await send_mcp_message(conn, payload)
await send_mcp_message(conn, payload, transport)
async def call_mcp_tool(
@@ -299,6 +242,7 @@ async def call_mcp_tool(
tool_name: str,
args: str = "{}",
timeout: int = 30,
return_raw: bool = False,
):
"""
调用指定的工具并等待响应
@@ -359,6 +303,7 @@ async def call_mcp_tool(
raise ValueError(f"参数必须是字典类型,实际类型: {type(arguments)}")
except Exception as e:
await mcp_client.cleanup_call_result(tool_call_id)
if not isinstance(e, ValueError):
raise ValueError(f"参数处理失败: {str(e)}")
raise e
@@ -372,7 +317,11 @@ async def call_mcp_tool(
}
logger.bind(tag=TAG).info(f"发送客户端mcp工具调用请求: {actual_name},参数: {args}")
await send_mcp_message(conn, payload)
try:
await send_mcp_message(conn, payload, raise_on_error=True)
except Exception:
await mcp_client.cleanup_call_result(tool_call_id)
raise
try:
# Wait for response or timeout
@@ -388,13 +337,21 @@ async def call_mcp_tool(
)
raise RuntimeError(f"工具调用错误: {error_msg}")
if return_raw:
return raw_result
content = raw_result.get("content")
if isinstance(content, list) and len(content) > 0:
if isinstance(content[0], dict) and "text" in content[0]:
# 直接返回文本内容,不进行JSON解析
return content[0]["text"]
# 如果结果不是预期的格式,将其转换为字符串
if return_raw:
return raw_result
return str(raw_result)
except asyncio.CancelledError:
await mcp_client.cleanup_call_result(tool_call_id)
raise
except asyncio.TimeoutError:
await mcp_client.cleanup_call_result(tool_call_id)
raise TimeoutError("工具调用请求超时")
@@ -22,6 +22,7 @@ class MCPEndpointClient:
self.lock = asyncio.Lock()
self._cached_available_tools = None # Cache for get_available_tools
self.websocket = None # WebSocket连接
self.listener_task = None
def has_tool(self, name: str) -> bool:
return name in self.tools
@@ -107,6 +108,18 @@ class MCPEndpointClient:
async def close(self):
"""关闭WebSocket连接"""
current_task = asyncio.current_task()
if (
self.listener_task
and self.listener_task is not current_task
and not self.listener_task.done()
):
self.listener_task.cancel()
try:
await self.listener_task
except asyncio.CancelledError:
pass
self.listener_task = None
if self.websocket:
await self.websocket.close()
self.websocket = None
@@ -16,6 +16,8 @@ async def connect_mcp_endpoint(mcp_endpoint_url: str, conn=None) -> MCPEndpointC
if not mcp_endpoint_url or "你的" in mcp_endpoint_url or mcp_endpoint_url == "null":
return None
websocket = None
mcp_client = None
try:
websocket = await websockets.connect(mcp_endpoint_url)
@@ -23,7 +25,9 @@ async def connect_mcp_endpoint(mcp_endpoint_url: str, conn=None) -> MCPEndpointC
mcp_client.set_websocket(websocket)
# 启动消息监听器
asyncio.create_task(_message_listener(mcp_client))
mcp_client.listener_task = asyncio.create_task(
_message_listener(mcp_client), name="xiaozhi-mcp-endpoint-listener"
)
# 发送初始化消息
await send_mcp_endpoint_initialize(mcp_client)
@@ -39,6 +43,10 @@ async def connect_mcp_endpoint(mcp_endpoint_url: str, conn=None) -> MCPEndpointC
except Exception as e:
logger.bind(tag=TAG).error(f"连接MCP接入点失败: {e}")
if mcp_client is not None:
await mcp_client.close()
elif websocket is not None:
await websocket.close()
return None
@@ -72,16 +72,25 @@ class ServerMCPClient:
async def cleanup(self):
"""清理MCP客户端资源"""
if not self._worker_task:
task = self._worker_task
if not task:
return
self._shutdown_evt.set()
try:
await asyncio.wait_for(self._worker_task, timeout=20)
except (asyncio.TimeoutError, Exception) as e:
await asyncio.wait_for(asyncio.shield(task), timeout=20)
except asyncio.TimeoutError:
self.logger.bind(tag=TAG).warning("服务端MCP关闭超时,取消工作任务")
task.cancel()
done, _ = await asyncio.wait({task}, timeout=5)
if task not in done:
self.logger.bind(tag=TAG).error("服务端MCP工作任务取消超时")
return
except Exception as e:
self.logger.bind(tag=TAG).error(f"服务端MCP客户端关闭错误: {e}")
finally:
self._worker_task = None
if task.done():
self._worker_task = None
def has_tool(self, name: str) -> bool:
"""检查是否包含指定工具
@@ -197,7 +206,7 @@ class ServerMCPClient:
if "API_ACCESS_TOKEN" in self.config:
headers["Authorization"] = f"Bearer {self.config['API_ACCESS_TOKEN']}"
self.logger.bind(tag=TAG).warning(f"你正在使用旧过时的配置 API_ACCESS_TOKEN ,请在.mcp_server_settings.json中将API_ACCESS_TOKEN直接设置在headers中,例如 'Authorization': 'Bearer API_ACCESS_TOKEN'")
# 根据transport类型选择不同的客户端,默认为SSE
transport_type = self.config.get("transport", "sse")
+96 -7
View File
@@ -1,4 +1,5 @@
import os
import json
import re
import uuid
import queue
@@ -16,8 +17,6 @@ from config.logger import setup_logging
from core.utils import opus_encoder_utils
from core.utils.tts import MarkdownCleaner, convert_percentage_to_range
from core.utils.output_counter import add_device_output
from core.handle.reportHandle import enqueue_tts_report
from core.handle.sendAudioHandle import sendAudioMessage
from core.utils.util import audio_bytes_to_data_stream, audio_to_data_stream
from core.providers.tts.dto.dto import (
TTSMessageDTO,
@@ -30,6 +29,58 @@ TAG = __name__
logger = setup_logging()
async def sendAudioMessage(conn, sentenceType, audios, text, sentence_id=None):
"""兼容函数:使用新的processor发送音频消息"""
try:
# 获取transport接口
transport = getattr(conn, 'transport', None)
if not transport:
logger.error("SessionContext中没有transport接口")
return
# 使用AudioSendProcessor发送音频
from core.processors.audio_send_processor import AudioSendProcessor
processor = AudioSendProcessor()
await processor.send_audio_message(
conn,
transport,
sentenceType,
audios,
text,
sentence_id=sentence_id,
)
except Exception as e:
logger.error(f"发送音频消息失败: {e}")
import traceback
traceback.print_exc()
def enqueue_tts_report(conn, text, audio_data):
"""兼容函数:使用新的processor处理TTS报告"""
try:
# 获取transport接口
transport = getattr(conn, 'transport', None)
if not transport:
logger.error("SessionContext中没有transport接口")
return
# 使用ReportProcessor处理报告
from core.processors.report_processor import ReportProcessor
processor = ReportProcessor()
# 异步执行报告
if hasattr(conn, 'loop') and conn.loop:
# 直接调用同步方法
processor.enqueue_tts_report(conn, text, audio_data)
else:
logger.warning("SessionContext中没有事件循环,跳过TTS报告")
except Exception as e:
logger.error(f"TTS报告处理失败: {e}")
class TTSProviderBase(ABC):
def __init__(self, config, delete_audio_file):
self.interface_type = InterfaceType.NON_STREAM
@@ -189,7 +240,7 @@ class TTSProviderBase(ABC):
except Exception as e:
logger.bind(tag=TAG).error(f"Failed to generate TTS file: {e}")
return None
def to_tts(self, text):
# 保留原始文本用于日志/显示
original_text = text
@@ -312,13 +363,17 @@ class TTSProviderBase(ABC):
# tts 消化线程
self.tts_priority_thread = threading.Thread(
target=self.tts_text_priority_thread, daemon=True
target=self.tts_text_priority_thread,
name=f"xiaozhi-tts-text-{id(self)}",
daemon=True,
)
self.tts_priority_thread.start()
# 音频播放 消化线程
self.audio_play_priority_thread = threading.Thread(
target=self._audio_play_priority_thread, daemon=True
target=self._audio_play_priority_thread,
name=f"xiaozhi-tts-audio-{id(self)}",
daemon=True,
)
self.audio_play_priority_thread.start()
@@ -369,6 +424,8 @@ class TTSProviderBase(ABC):
while not self.conn.stop_event.is_set():
try:
message = self.tts_text_queue.get(timeout=1)
if message is None:
break
if self.conn.client_abort:
logger.bind(tag=TAG).info("收到打断信息,终止TTS文本处理线程")
continue
@@ -417,11 +474,15 @@ class TTSProviderBase(ABC):
try:
try:
item = self.tts_audio_queue.get(timeout=0.1)
if item is None:
break
if len(item) == 4:
sentence_type, audio_datas, text, sentence_id = item
else:
sentence_type, audio_datas, text = item
sentence_id = None
sentence_id = getattr(
self, "current_sentence_id", None
) or getattr(self.conn, "sentence_id", None)
except queue.Empty:
if self.conn.stop_event.is_set():
break
@@ -459,7 +520,19 @@ class TTSProviderBase(ABC):
sendAudioMessage(self.conn, sentence_type, audio_datas, text, sentence_id),
self.conn.loop,
)
future.result()
self._pending_audio_future = future
try:
future.result(timeout=max(1, self.tts_timeout))
except concurrent.futures.CancelledError:
break
except concurrent.futures.TimeoutError:
future.cancel()
logger.bind(tag=TAG).warning(
"TTS音频发送超时,取消当前发送任务"
)
finally:
if getattr(self, "_pending_audio_future", None) is future:
self._pending_audio_future = None
# 记录输出和报告
if self.conn.max_output_size > 0 and text:
@@ -476,6 +549,22 @@ class TTSProviderBase(ABC):
async def close(self):
"""资源清理方法"""
self.tts_stop_request = True
pending_future = getattr(self, "_pending_audio_future", None)
if pending_future and not pending_future.done():
pending_future.cancel()
# Wake workers immediately; relying on queue timeouts leaves Provider
# threads alive after the component has released its ownership.
self.tts_text_queue.put(None)
self.tts_audio_queue.put(None)
for thread_name in ("tts_priority_thread", "audio_play_priority_thread"):
thread = getattr(self, thread_name, None)
if thread and thread.is_alive() and thread is not threading.current_thread():
await asyncio.to_thread(thread.join, 2)
if thread.is_alive():
logger.bind(tag=TAG).warning(
"TTS工作线程未能按时退出: {}", thread.name
)
self._sentence_text_map.clear()
if hasattr(self, "ws") and self.ws:
await self.ws.close()
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,464 @@
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服务器
提供统一的启动停止和状态监控接口
"""
def __init__(self, config: Dict[str, Any]):
self.config = config
self.logger = setup_logging()
# 服务器实例
self.servers: Dict[str, Any] = {}
self.server_tasks: Dict[str, asyncio.Task] = {}
self.monitor_task: Optional[asyncio.Task] = None
self._lifecycle_lock = asyncio.Lock()
self.management_owner = None
# 服务器状态
self.is_running = False
self.startup_complete = False
# 初始化服务器
self._initialize_servers()
def _initialize_servers(self):
"""初始化所有协议服务器"""
try:
self.servers.clear()
# 检查配置中启用的协议
enabled_protocols = self.config.get('enabled_protocols', ['websocket'])
# 初始化WebSocket服务器
if 'websocket' in enabled_protocols:
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("没有启用任何协议服务器")
for server in self.servers.values():
self._bind_management_owner(server)
except Exception as e:
logger.error(f"初始化服务器失败: {e}")
raise
def set_management_owner(self, owner: Any) -> None:
"""Route management commands through the facade that owns all protocols."""
self.management_owner = owner
for server in self.servers.values():
self._bind_management_owner(server)
def _bind_management_owner(self, server: Any) -> None:
owner = self.management_owner or server
server.management_owner = owner
connection_service = getattr(server, "connection_service", None)
if connection_service is not None:
connection_service.server = owner
async def start(self):
"""启动所有服务器,并由管理器持有协议任务。"""
async with self._lifecycle_lock:
await self._start_unlocked()
async def _start_unlocked(self):
if self.is_running:
logger.warning("服务器已经在运行中")
return
if self.server_tasks:
raise RuntimeError("上次协议停止尚未完成,请先重试 stop 清理残留资源")
try:
logger.info("开始启动多协议服务器...")
self.is_running = True
self.startup_complete = False
# 启动所有服务器
for protocol, server in self.servers.items():
try:
logger.info(f"启动{protocol}服务器...")
task = asyncio.create_task(
server.start(), name=f"xiaozhi-{protocol}-server"
)
self.server_tasks[protocol] = task
await self._wait_for_server_started(protocol, server, task)
logger.info(f"{protocol}服务器启动成功")
except Exception as e:
logger.error(f"启动{protocol}服务器失败: {e}")
raise
self.startup_complete = True
logger.info("多协议服务器启动完成")
# 启动监控任务;必须保存引用,以便停止和重启时回收。
self.monitor_task = asyncio.create_task(
self._monitor_servers(), name="xiaozhi-protocol-monitor"
)
except Exception as e:
logger.error(f"启动多协议服务器失败: {e}")
try:
await self._stop_unlocked()
except Exception as cleanup_error:
logger.error(f"启动失败后的协议清理失败: {cleanup_error}")
raise
async def _wait_for_server_started(self, protocol, server, task):
"""等待协议监听器真正就绪,而不是依赖固定 sleep。"""
started_event = getattr(server, "_started_event", None)
if not isinstance(started_event, asyncio.Event):
await asyncio.sleep(0)
if task.done():
await task
return
timeout = float(self.config.get("server_startup_timeout", 10))
event_waiter = asyncio.create_task(started_event.wait())
try:
done, _ = await asyncio.wait(
{task, event_waiter},
timeout=timeout,
return_when=asyncio.FIRST_COMPLETED,
)
if started_event.is_set():
if task.done():
await task
return
if task in done:
await task
raise RuntimeError(f"{protocol}服务器在就绪前退出")
raise TimeoutError(f"等待{protocol}服务器启动超时")
finally:
if not event_waiter.done():
event_waiter.cancel()
try:
await event_waiter
except asyncio.CancelledError:
pass
async def stop(self):
"""停止所有服务器"""
async with self._lifecycle_lock:
await self._stop_unlocked()
async def _stop_unlocked(self):
if not self.is_running and not self.server_tasks and not self.monitor_task:
return
logger.info("开始停止多协议服务器...")
self.is_running = False
self.startup_complete = False
if self.monitor_task and not self.monitor_task.done():
self.monitor_task.cancel()
try:
await self.monitor_task
except asyncio.CancelledError:
pass
self.monitor_task = None
# 停止所有服务器
errors = []
stopped_protocols = []
for protocol, server in self.servers.items():
try:
logger.info(f"停止{protocol}服务器...")
stopped = await server.stop()
if stopped is False:
raise RuntimeError("协议仍有取消不响应的后台任务")
logger.info(f"{protocol}服务器已停止")
stopped_protocols.append(protocol)
except Exception as e:
logger.error(f"停止{protocol}服务器失败: {e}")
errors.append((protocol, e))
# 只释放已确认停止的任务。失败协议保留所有权,允许再次 stop。
for protocol in stopped_protocols:
task = self.server_tasks.get(protocol)
if task is None:
continue
if not task.done():
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
except Exception as e:
errors.append((f"{protocol}任务", e))
logger.error(f"回收{protocol}服务器任务失败: {e}")
continue
self.server_tasks.pop(protocol, None)
if errors:
details = ", ".join(
f"{protocol}: {error}" for protocol, error in errors
)
raise RuntimeError(f"停止协议服务器时发生错误: {details}")
logger.info("多协议服务器已停止")
async def restart(self):
"""重启所有服务器"""
async with self._lifecycle_lock:
logger.info("重启多协议服务器...")
await self._stop_unlocked()
await self._start_unlocked()
async def restart_server(self, protocol: str):
"""重启指定协议的服务器"""
async with self._lifecycle_lock:
return await self._restart_server_unlocked(protocol)
async def _restart_server_unlocked(self, protocol: str):
if protocol not in self.servers:
logger.error(f"未找到协议服务器: {protocol}")
return False
try:
logger.info(f"重启{protocol}服务器...")
# 停止指定服务器
server = self.servers[protocol]
stopped = await server.stop()
if stopped is False:
raise RuntimeError("协议仍有取消不响应的后台任务")
# 取消任务
if protocol in self.server_tasks:
task = self.server_tasks[protocol]
if not task.done():
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
# 重新启动
task = asyncio.create_task(
server.start(), name=f"xiaozhi-{protocol}-server"
)
self.server_tasks[protocol] = task
await self._wait_for_server_started(protocol, server, task)
logger.info(f"{protocol}服务器重启成功")
return True
except Exception as e:
logger.error(f"重启{protocol}服务器失败: {e}")
return False
async def update_config(self, new_config: Dict[str, Any]):
"""更新配置"""
async with self._lifecycle_lock:
return await self._update_config_unlocked(new_config)
async def _update_config_unlocked(self, new_config: Dict[str, Any]):
old_config = self.config
was_running = self.is_running
config_changed = False
old_runtime_stopped = False
try:
logger.info("更新多协议服务器配置...")
# 检查配置变化
config_changed = self._check_config_changes(self.config, new_config)
# 如果配置有重大变化,重新初始化服务器
if config_changed:
logger.info("配置有重大变化,重新初始化服务器...")
await self._stop_unlocked()
old_runtime_stopped = True
self.config = new_config
self._initialize_servers()
if was_running:
await self._start_unlocked()
else:
# 更新各个服务器的配置
for protocol, server in self.servers.items():
updated = True
if hasattr(server, 'apply_config'):
updated = await server.apply_config(new_config)
elif hasattr(server, 'update_config'):
updated = await server.update_config(new_config)
if updated is False:
raise RuntimeError(f"{protocol}配置更新失败")
self._bind_management_owner(server)
self.config = new_config
logger.info("配置更新完成")
return True
except Exception as e:
logger.error(f"更新配置失败: {e}")
if config_changed and not old_runtime_stopped:
# A listener still owns resources. Preserve the old server
# objects/tasks so stop can be retried safely.
self.config = old_config
return False
try:
await self._stop_unlocked()
self.config = old_config
self._initialize_servers()
if was_running:
await self._start_unlocked()
except Exception:
logger.exception("回滚多协议服务器配置失败")
return False
def _check_config_changes(self, old_config: Dict[str, Any], new_config: Dict[str, Any]) -> bool:
"""检查配置是否有重大变化"""
# 检查启用的协议是否变化
old_protocols = set(old_config.get('enabled_protocols', ['websocket']))
new_protocols = set(new_config.get('enabled_protocols', ['websocket']))
if old_protocols != new_protocols:
logger.info(f"启用协议发生变化: {old_protocols} -> {new_protocols}")
return True
if old_config.get("_shared_asr_manager") is not new_config.get(
"_shared_asr_manager"
):
logger.info("共享ASR管理器发生变化")
return True
# 检查服务器端口配置
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, {})
new_server_config = new_config.get(config_key, {})
# 检查端口和主机配置
for key in listener_keys:
if old_server_config.get(key) != new_server_config.get(key):
logger.info(f"服务器配置{config_key}.{key}发生变化")
return True
return False
async def _monitor_servers(self):
"""监控服务器状态"""
try:
while self.is_running:
await asyncio.sleep(30) # 每30秒检查一次
# 检查服务器任务状态
for protocol, task in list(self.server_tasks.items()):
if task.done():
exception = None if task.cancelled() else task.exception()
if exception:
logger.error(f"{protocol}服务器异常退出: {exception}")
else:
logger.warning(f"{protocol}服务器任务意外退出")
# 监控任务本身不能重入 lifecycle lock。
async with self._lifecycle_lock:
if self.is_running:
await self._restart_server_unlocked(protocol)
except asyncio.CancelledError:
pass
except Exception as e:
logger.error(f"服务器监控任务出错: {e}")
def get_server_status(self) -> Dict[str, Any]:
"""获取所有服务器状态"""
status = {
'is_running': self.is_running,
'startup_complete': self.startup_complete,
'enabled_protocols': list(self.servers.keys()),
'servers': {}
}
# 获取各个服务器的状态
for protocol, server in self.servers.items():
try:
if hasattr(server, 'get_server_status'):
server_status = server.get_server_status()
else:
server_status = {'type': protocol, 'status': 'unknown'}
# 添加任务状态
task = self.server_tasks.get(protocol)
if task:
server_status['task_status'] = 'running' if not task.done() else 'stopped'
if task.done() and not task.cancelled() and task.exception():
server_status['task_error'] = str(task.exception())
status['servers'][protocol] = server_status
except Exception as e:
status['servers'][protocol] = {
'type': protocol,
'status': 'error',
'error': str(e)
}
return status
def get_active_connections_count(self) -> Dict[str, int]:
"""获取各协议的活跃连接数"""
connections = {}
for protocol, server in self.servers.items():
try:
if hasattr(server, 'get_active_connections_count'):
connections[protocol] = server.get_active_connections_count()
elif hasattr(server, 'connections'):
connections[protocol] = len(server.connections)
else:
connections[protocol] = 0
except Exception as e:
logger.error(f"获取{protocol}连接数失败: {e}")
connections[protocol] = -1
return connections
async def broadcast_message(self, message: Dict[str, Any], protocol: Optional[str] = None):
"""向所有连接广播消息"""
try:
if protocol:
# 向指定协议广播
if protocol in self.servers:
server = self.servers[protocol]
if hasattr(server, 'broadcast_message'):
await server.broadcast_message(message)
else:
# 向所有协议广播
for server in self.servers.values():
if hasattr(server, 'broadcast_message'):
await server.broadcast_message(message)
except Exception as e:
logger.error(f"广播消息失败: {e}")
def get_supported_protocols(self) -> List[str]:
"""获取支持的协议列表"""
return ['websocket', 'mqtt']
def is_protocol_enabled(self, protocol: str) -> bool:
"""检查协议是否启用"""
return protocol in self.servers
@@ -0,0 +1,229 @@
import time
from typing import Any, Optional
import numpy as np
import opuslib_next
from config.logger import setup_logging
logger = setup_logging()
class AudioIngressService:
"""Decode transport audio once and apply optional server-side AEC."""
def __init__(self, decoder_factory=None):
self._decoder_factory = decoder_factory or opuslib_next.Decoder
def process(self, context: Any, opus_packet: bytes, timestamp: int = 0) -> Optional[bytes]:
pcm_frame = self.decode(context, opus_packet)
if pcm_frame and timestamp > 0 and getattr(context, "client_aec", False):
return self.apply_aec(context, timestamp, pcm_frame)
return pcm_frame
def decode(self, context: Any, opus_packet: bytes) -> Optional[bytes]:
if not opus_packet:
return None
sample_rate = int(getattr(context, "input_sample_rate", 16000) or 16000)
channels = int(getattr(context, "input_channels", 1) or 1)
frame_duration = int(
getattr(context, "input_frame_duration", 60) or 60
)
decoder_config = (sample_rate, channels)
try:
if getattr(context, "_audio_ingress_decoder_config", None) != decoder_config:
context._audio_ingress_decoder = self._decoder_factory(sample_rate, channels)
context._audio_ingress_decoder_config = decoder_config
self._register_cleanup(context)
frame_size = max(1, sample_rate * frame_duration // 1000)
pcm_frame = context._audio_ingress_decoder.decode(bytes(opus_packet), frame_size)
return bytes(pcm_frame)
except Exception as exc:
logger.debug("Opus decode failed: {}", exc)
return None
def apply_aec(self, context: Any, timestamp: int, pcm_frame: bytes) -> bytes:
"""Apply the AEC algorithm used by the upstream gateway path."""
try:
self._expire_aec_references(context)
audio_cache = getattr(context, "aec_audio_cache", None)
if not pcm_frame or not audio_cache:
return pcm_frame
mic_audio = np.frombuffer(pcm_frame, dtype=np.int16).astype(np.float32)
mic_rms = np.sqrt(np.mean(mic_audio**2))
if mic_rms < 100:
return pcm_frame
sorted_timestamps = sorted(audio_cache.keys())
if len(sorted_timestamps) < 2:
return pcm_frame
sample_count = len(mic_audio)
closest_idx = min(
range(len(sorted_timestamps)),
key=lambda index: abs(sorted_timestamps[index] - timestamp),
)
mic_window = np.hanning(sample_count)
mic_fft = np.fft.rfft(mic_audio * mic_window)
mic_psd = np.abs(mic_fft) ** 2
mic_log_psd = 10 * np.log10(mic_psd + 1e-8)
mic_power = np.dot(mic_log_psd, mic_log_psd)
best_corr = -1
best_ref_idx = closest_idx
best_ref_rms = 0.0
for offset in range(-2, 3):
test_idx = closest_idx + offset
if test_idx < 0 or test_idx >= len(sorted_timestamps):
continue
test_ref = np.frombuffer(
audio_cache[sorted_timestamps[test_idx]], dtype=np.int16
).astype(np.float32)
test_ref_rms = np.sqrt(np.mean(test_ref**2))
if test_ref_rms < 50:
continue
test_fft = np.fft.rfft(test_ref * np.hanning(len(test_ref)))
test_log_psd = 10 * np.log10(np.abs(test_fft) ** 2 + 1e-8)
cross_power = np.dot(mic_log_psd, test_log_psd)
reference_power = np.dot(test_log_psd, test_log_psd)
corr = abs(cross_power) / (
np.sqrt(mic_power) * np.sqrt(reference_power) + 1e-8
)
if corr > best_corr:
best_corr = corr
best_ref_idx = test_idx
best_ref_rms = test_ref_rms
best_ref = np.frombuffer(
audio_cache[sorted_timestamps[best_ref_idx]], dtype=np.int16
).astype(np.float32)
if best_ref_rms < 50:
return pcm_frame
aligned_ref = best_ref[:sample_count]
if len(aligned_ref) < sample_count:
aligned_ref = np.pad(aligned_ref, (0, sample_count - len(aligned_ref)))
mic_mag = np.abs(mic_fft)
mic_phase = np.angle(mic_fft)
ref_mag = np.abs(np.fft.rfft(aligned_ref * np.hanning(sample_count)))
scale = np.sum(mic_mag * ref_mag) / (np.dot(ref_mag, ref_mag) + 1e-8)
coefficient = max(
0.5, min(3.0, 1.0 + scale * 3 + (best_corr - 0.97) * 30)
)
result_mag = np.maximum(
mic_mag - ref_mag * scale * coefficient * 1.5, mic_mag * 0.1
)
output = np.fft.irfft(result_mag * np.exp(1j * mic_phase), sample_count)
if best_corr >= 0.97 and best_ref_rms > 500:
output *= 0.3
return np.clip(output, -32768, 32767).astype(np.int16).tobytes()
except Exception as exc:
logger.warning("AEC processing failed: {}", exc)
return pcm_frame
def cache_output_reference(
self, context: Any, opus_packet: bytes, timestamp: int
) -> None:
if not getattr(context, "client_aec", False) or timestamp <= 0:
return
try:
sample_rate = int(
getattr(
context,
"output_sample_rate",
getattr(context, "sample_rate", 24000),
)
or 24000
)
channels = int(getattr(context, "output_channels", 1) or 1)
frame_duration = int(
getattr(context, "output_frame_duration", 60) or 60
)
decoder_config = (sample_rate, channels)
decoder = getattr(context, "_audio_reference_decoder", None)
if (
decoder is None
or getattr(context, "_audio_reference_decoder_config", None)
!= decoder_config
):
decoder = self._decoder_factory(sample_rate, channels)
context._audio_reference_decoder = decoder
context._audio_reference_decoder_config = decoder_config
self._register_cleanup(context)
frame_size = max(1, sample_rate * frame_duration // 1000)
pcm_data = bytes(decoder.decode(bytes(opus_packet), frame_size))
input_rate = int(getattr(context, "input_sample_rate", 16000) or 16000)
if sample_rate != input_rate:
pcm_data = self._resample_pcm(pcm_data, sample_rate, input_rate)
self._expire_aec_references(context)
context.aec_audio_cache[timestamp] = pcm_data
context.aec_audio_cache_time[timestamp] = time.time()
except Exception as exc:
logger.debug("AEC reference decode failed: {}", exc)
@staticmethod
def _expire_aec_references(context: Any) -> None:
cache = getattr(context, "aec_audio_cache", {})
cache_times = getattr(context, "aec_audio_cache_time", {})
config = getattr(context, "config", {}) or {}
max_age = max(1, int(config.get("aec_reference_max_age_seconds", 120)))
max_frames = max(2, int(config.get("aec_reference_max_frames", 256)))
now = time.time()
expired = [
timestamp
for timestamp, cached_at in list(cache_times.items())
if now - cached_at > max_age
]
overflow = max(0, len(cache) - max_frames + 1)
if overflow:
expired.extend(sorted(cache_times, key=cache_times.get)[:overflow])
for timestamp in set(expired):
cache.pop(timestamp, None)
cache_times.pop(timestamp, None)
def cleanup(self, context: Any) -> None:
for name in (
"_audio_ingress_decoder",
"_audio_ingress_decoder_config",
"_audio_reference_decoder",
"_audio_reference_decoder_config",
):
if hasattr(context, name):
delattr(context, name)
if hasattr(context, "aec_audio_cache"):
context.aec_audio_cache.clear()
if hasattr(context, "aec_audio_cache_time"):
context.aec_audio_cache_time.clear()
def _register_cleanup(self, context: Any) -> None:
if getattr(context, "_audio_ingress_cleanup_registered", False):
return
register_cleanup = getattr(context, "register_cleanup", None)
if callable(register_cleanup):
register_cleanup(lambda: self.cleanup(context))
context._audio_ingress_cleanup_registered = True
@staticmethod
def _resample_pcm(pcm_data: bytes, source_rate: int, target_rate: int) -> bytes:
if not pcm_data or source_rate == target_rate:
return pcm_data
source = np.frombuffer(pcm_data, dtype=np.int16)
if len(source) < 2:
return pcm_data
target_size = max(1, round(len(source) * target_rate / source_rate))
source_positions = np.linspace(0.0, 1.0, len(source), endpoint=False)
target_positions = np.linspace(0.0, 1.0, target_size, endpoint=False)
target = np.interp(target_positions, source_positions, source)
return np.clip(target, -32768, 32767).astype(np.int16).tobytes()
File diff suppressed because it is too large Load Diff
@@ -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
@@ -0,0 +1,90 @@
from abc import ABC, abstractmethod
import json
from typing import Any, AsyncGenerator
class TransportInterface(ABC):
"""
传输层抽象接口
"""
@abstractmethod
async def send(self, data: Any) -> None:
"""发送一条消息。"""
raise NotImplementedError
async def send_json(self, message: Any) -> None:
"""Send a control message without exposing transport framing to callers."""
payload = message if isinstance(message, str) else json.dumps(message)
await self.send(payload)
async def send_audio(self, audio: bytes, timestamp: int = 0) -> None:
"""Send one encoded audio frame over the transport's audio channel."""
await self.send(audio)
async def prepare_audio_channel(self, audio_params=None, version: int = 3) -> None:
"""Prepare transport-specific audio negotiation when required."""
async def wait_audio_ready(self, timeout: float = 0) -> bool:
"""Wait until encoded audio can be delivered to the peer."""
return True
async def mark_business_ready(self) -> None:
"""Allow a transport handshake to proceed after runtime initialization."""
async def mark_session_ready(self, session_id: str = None) -> None:
"""Release a logical-session handshake after its runtime is ready."""
async def end_session(self, session_id: str) -> None:
"""Tell the device to return to Idle without closing the connection."""
await self.send_json({"type": "goodbye", "session_id": session_id})
@abstractmethod
async def receive(self) -> AsyncGenerator[Any, None]:
"""异步消息流。"""
yield # pragma: no cover
@abstractmethod
async def close(self) -> None:
"""关闭底层连接。"""
raise NotImplementedError
@property
@abstractmethod
def is_connected(self) -> bool:
"""连接是否存活。"""
raise NotImplementedError
@property
def transport_type(self) -> str:
"""Stable transport identifier used by shared connection logic."""
return "unknown"
@property
def has_datagram_audio(self) -> bool:
"""Whether audio is carried by a channel separate from control messages."""
return False
@property
def requires_audio_tail_grace(self) -> bool:
"""Whether control can overtake the last audio frames of a turn."""
return False
@property
def keeps_connection_between_sessions(self) -> bool:
"""Whether ending a conversation should leave the transport connected."""
return False
@property
def is_protocol_authenticated(self) -> bool:
"""Whether the transport already authenticated the peer during handshake."""
return False
@property
def raw_connection(self):
"""Underlying connection for temporary compatibility with legacy code."""
return None
@property
def session_id(self):
return None
@@ -0,0 +1,154 @@
import struct
import time
from typing import Any, AsyncGenerator
from .transport_interface import TransportInterface
class WebSocketTransport(TransportInterface):
"""
WebSocket 传输实现包装 websockets 库的协议对象
提供统一的 send/receive/close 接口
"""
def __init__(
self,
websocket,
from_mqtt_gateway: bool = False,
protocol_version: int = 1,
):
self._ws = websocket
self._from_mqtt_gateway = from_mqtt_gateway
self._protocol_version = protocol_version if protocol_version in (1, 2, 3) else 1
self._gateway_sequence = 0
async def send(self, data: Any) -> None:
if isinstance(data, bytes):
await self.send_audio(data)
return
if isinstance(data, (str, bytes)):
await self._ws.send(data)
else:
await self._ws.send(str(data))
async def send_audio(self, audio: bytes, timestamp: int = 0) -> None:
if self._from_mqtt_gateway:
await self._ws.send(self._frame_gateway_audio(audio, timestamp))
return
await self._ws.send(self._frame_websocket_audio(audio, timestamp))
def _frame_websocket_audio(self, audio: bytes, timestamp: int = 0) -> bytes:
if self._protocol_version == 2:
return struct.pack("!HHIII", 2, 0, 0, timestamp, len(audio)) + audio
if self._protocol_version == 3:
return struct.pack("!BBH", 0, 0, len(audio)) + audio
return audio
def _parse_websocket_audio(self, message: bytes):
if self._protocol_version == 2:
if len(message) < 16:
return None
version, message_type, _, timestamp, audio_length = struct.unpack(
"!HHIII", message[:16]
)
if (
version != 2
or message_type != 0
or audio_length <= 0
or len(message) != 16 + audio_length
):
return None
return {
"type": "audio",
"data": message[16:],
"timestamp": timestamp,
"_transport_type": "websocket",
}
if self._protocol_version == 3:
if len(message) < 4:
return None
message_type, _, audio_length = struct.unpack("!BBH", message[:4])
if (
message_type != 0
or audio_length <= 0
or len(message) != 4 + audio_length
):
return None
return {
"type": "audio",
"data": message[4:],
"timestamp": 0,
"_transport_type": "websocket",
}
return message
def _frame_gateway_audio(self, audio: bytes, timestamp: int = 0) -> bytes:
self._gateway_sequence += 1
if timestamp <= 0:
timestamp = int(time.time() * 1000) % (2 ** 32)
header = bytearray(16)
header[0] = 1
header[2:4] = len(audio).to_bytes(2, "big")
header[4:8] = self._gateway_sequence.to_bytes(4, "big")
header[8:12] = timestamp.to_bytes(4, "big")
header[12:16] = len(audio).to_bytes(4, "big")
return bytes(header) + audio
async def receive(self) -> AsyncGenerator[Any, None]:
async for message in self._ws:
if self._from_mqtt_gateway and isinstance(message, bytes):
if len(message) < 16:
continue
if message[:8] != b"\x00" * 8:
continue
timestamp = int.from_bytes(message[8:12], "big")
audio_length = int.from_bytes(message[12:16], "big")
if audio_length <= 0 or len(message) != 16 + audio_length:
continue
message = {
"type": "audio",
"data": message[16:],
"timestamp": timestamp,
"_transport_type": "gateway",
}
elif isinstance(message, bytes):
message = self._parse_websocket_audio(message)
if message is None:
continue
yield message
async def close(self) -> None:
try:
if hasattr(self._ws, "closed") and not self._ws.closed:
await self._ws.close()
elif hasattr(self._ws, "state") and self._ws.state.name != "CLOSED":
await self._ws.close()
else:
await self._ws.close()
except Exception:
raise RuntimeError("WebSocket close failed")
@property
def is_connected(self) -> bool:
try:
if hasattr(self._ws, "closed"):
return not self._ws.closed
if hasattr(self._ws, "state"):
return getattr(self._ws.state, "name", "CLOSED") != "CLOSED"
except Exception:
raise RuntimeError("WebSocket connection check failed")
return False
@property
def transport_type(self) -> str:
return "gateway" if self._from_mqtt_gateway else "websocket"
@property
def requires_audio_tail_grace(self) -> bool:
# The gateway receives MQTT control and UDP audio independently before
# serializing both streams onto this WebSocket.
return self._from_mqtt_gateway
@property
def raw_connection(self):
return self._ws
@@ -0,0 +1,106 @@
import asyncio
from typing import Any, Dict, Tuple
from config.logger import setup_logging
from core.utils.modules_initialize import initialize_modules
from core.providers.asr.shared_asr_manager import SharedASRManager
def _validate_selected_modules(config: Dict[str, Any]) -> Tuple[bool, str]:
selected = config.get("selected_module", {})
required_sections = {
"VAD": "VAD",
"ASR": "ASR",
"LLM": "LLM",
"TTS": "TTS",
"Memory": "Memory",
"Intent": "Intent",
}
for module_key, section_key in required_sections.items():
module_name = selected.get(module_key)
if not module_name:
continue
section = config.get(section_key, {})
if module_name not in section:
return False, f"{section_key}配置缺失: {module_name}"
return True, ""
async def _cleanup_instance(instance: Any) -> None:
if instance is None:
return
try:
if hasattr(instance, "close"):
result = instance.close()
if asyncio.iscoroutine(result):
await result
if hasattr(instance, "cleanup"):
result = instance.cleanup()
if asyncio.iscoroutine(result):
await result
if hasattr(instance, "cleanup_audio_files"):
instance.cleanup_audio_files()
except Exception:
# 校验阶段的清理异常不影响结果
pass
async def _cleanup_modules(modules: Dict[str, Any]) -> None:
for instance in modules.values():
await _cleanup_instance(instance)
async def validate_config_components(
config: Dict[str, Any], logger=None
) -> Tuple[bool, str]:
"""
预初始化选中组件用于校验配置捕获配置/密钥类错误
返回 (是否通过, 错误信息)
"""
logger = logger or setup_logging()
ok, msg = _validate_selected_modules(config)
if not ok:
return False, msg
selected = config.get("selected_module", {})
init_vad = bool(selected.get("VAD"))
init_asr = bool(selected.get("ASR"))
asr_type = None
selected_asr = selected.get("ASR")
if selected_asr:
asr_config = config.get("ASR", {}).get(selected_asr, {})
asr_type = asr_config.get("type", selected_asr)
if SharedASRManager.is_local_model_type(asr_type):
init_asr = False
init_llm = bool(selected.get("LLM"))
init_tts = bool(selected.get("TTS"))
init_memory = bool(selected.get("Memory"))
init_intent = bool(selected.get("Intent"))
modules: Dict[str, Any] = {}
manager = None
try:
modules = initialize_modules(
logger,
config,
init_vad,
init_asr,
init_llm,
init_tts,
init_memory,
init_intent,
)
if selected_asr and asr_type and SharedASRManager.is_local_model_type(asr_type):
manager = SharedASRManager(config, asr_type)
await manager.initialize()
return True, ""
except Exception as e:
logger.error(f"配置校验失败: {e}")
return False, str(e)
finally:
if manager:
try:
await manager.shutdown()
except Exception:
pass
await _cleanup_modules(modules)
+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("密码签名验证失败")
+18
View File
@@ -0,0 +1,18 @@
import os
import subprocess
import sys
import time
def restart_server(logger):
"""实际执行重启的方法"""
time.sleep(1)
logger.info("执行服务器重启...")
subprocess.Popen(
[sys.executable, "app.py"],
stdin=sys.stdin,
stdout=sys.stdout,
stderr=sys.stderr,
start_new_session=True,
)
os._exit(0)
+20 -10
View File
@@ -1,5 +1,6 @@
import json
from typing import TYPE_CHECKING
from config.logger import setup_logging
if TYPE_CHECKING:
from core.connection import ConnectionHandler
@@ -91,18 +92,27 @@ async def get_emotion(conn: "ConnectionHandler", text):
emotion = EMOJI_MAP[char]
break
try:
await conn.websocket.send(
json.dumps(
{
"type": "llm",
"text": emoji,
"emotion": emotion,
"session_id": conn.session_id,
}
)
message = json.dumps(
{
"type": "llm",
"text": emoji,
"emotion": emotion,
"session_id": conn.session_id,
}
)
# 使用transport接口发送消息
if hasattr(conn, 'transport') and conn.transport:
await conn.transport.send(message)
elif hasattr(conn, 'websocket') and conn.websocket:
# 兼容旧版本
await conn.websocket.send(message)
else:
raise AttributeError("无法找到可用的传输层接口")
except Exception as e:
conn.logger.bind(tag=TAG).warning(f"发送情绪表情失败,错误:{e}")
logger = setup_logging()
logger.warning(f"发送情绪表情失败,错误:{e}")
return
@@ -0,0 +1,452 @@
import asyncio
import logging
import json
import websockets
from typing import Dict, Any, Optional
from config.logger import setup_logging
from core.services.connection_service import ConnectionService
from core.transport.websocket_transport import WebSocketTransport
from config.config_loader import get_config_from_api_async
from core.utils.util import check_vad_update, check_asr_update
from core.auth import AuthMiddleware, AuthenticationError
from core.utils.config_validation import validate_config_components
from core.providers.asr.shared_asr_manager import SharedASRManager
class SuppressInvalidHandshakeFilter(logging.Filter):
"""过滤掉无效握手错误日志(如HTTPS访问WS端口)"""
def filter(self, record):
msg = record.getMessage()
suppress_keywords = [
"opening handshake failed",
"did not receive a valid HTTP request",
"connection closed while reading HTTP request",
"line without CRLF",
]
return not any(keyword in msg for keyword in suppress_keywords)
def _setup_websockets_logger():
"""配置 websockets 相关的所有 logger,过滤无效握手错误"""
filter_instance = SuppressInvalidHandshakeFilter()
for logger_name in ["websockets", "websockets.server", "websockets.client"]:
ws_logger = logging.getLogger(logger_name)
ws_logger.addFilter(filter_instance)
_setup_websockets_logger()
logger = setup_logging()
TAG = __name__
class NewWebSocketServer:
"""
新的WebSocket服务器使用新架构替代旧的ConnectionHandler
集成ConnectionServiceMessageRouter和新的Processor架构
"""
def __init__(self, config: Dict[str, Any]):
self.config = config
self.logger = setup_logging()
self.config_lock = asyncio.Lock()
self.last_update_error = None
# 创建连接服务
self.connection_service = ConnectionService(config)
self.connection_service.server = self
# 活跃连接管理
self.active_connections = set()
# 认证中间件
self.auth_middleware = AuthMiddleware(config)
# 服务器实例和控制
self._server = None
self._stop_event = asyncio.Event()
self._started_event = asyncio.Event()
self._is_running = False
async def start(self):
"""启动WebSocket服务器"""
server_config = self.config["server"]
host = server_config.get("ip", "0.0.0.0")
port = int(server_config.get("port", 8000))
logger.bind(tag=TAG).info(f"启动新架构WebSocket服务器: {host}:{port}")
self._stop_event.clear()
self._started_event.clear()
try:
self._server = await websockets.serve(
self._handle_connection,
host,
port,
process_request=self._http_response
)
self._is_running = True
self._started_event.set()
logger.bind(tag=TAG).info("WebSocket服务器启动成功")
# 等待停止信号
await self._stop_event.wait()
except Exception as e:
logger.bind(tag=TAG).error(f"WebSocket服务器启动失败: {e}")
raise
finally:
self._is_running = False
self._started_event.clear()
async def stop(self):
"""停止WebSocket服务器"""
if not self._is_running:
logger.bind(tag=TAG).debug("WebSocket服务器未运行,无需停止")
return
logger.bind(tag=TAG).info("正在停止WebSocket服务器...")
# 关闭所有活跃连接
for transport in list(self.active_connections):
try:
await transport.close()
except Exception as e:
logger.bind(tag=TAG).error(f"关闭连接失败: {e}")
self.active_connections.clear()
# 关闭服务器
if self._server:
self._server.close()
try:
await asyncio.wait_for(self._server.wait_closed(), timeout=5.0)
except asyncio.TimeoutError:
logger.bind(tag=TAG).warning("等待服务器关闭超时")
self._server = None
# 发送停止信号
self._stop_event.set()
self._is_running = False
self._started_event.clear()
logger.bind(tag=TAG).info("WebSocket服务器已停止")
async def _handle_connection(self, websocket):
"""处理新连接 - 使用新架构"""
# 提取连接头信息
headers = self._extract_headers(websocket)
device_id = headers.get('device-id')
# 如果没有 device-id,提示并关闭连接
if not device_id:
await websocket.send("端口正常,如需测试连接,请使用test_page.html")
await websocket.close()
return
# 连接时认证
try:
await self._handle_auth(headers)
except AuthenticationError as e:
logger.bind(tag=TAG).warning(f"认证失败: {e}")
await websocket.send("认证失败")
await websocket.close()
return
# 创建WebSocket传输层
try:
protocol_version = int(headers.get('protocol-version', 1) or 1)
except (TypeError, ValueError):
protocol_version = 1
transport = WebSocketTransport(
websocket,
from_mqtt_gateway=headers.get('from_mqtt_gateway') == 'true',
protocol_version=protocol_version,
)
# 记录活跃连接
self.active_connections.add(transport)
try:
logger.bind(tag=TAG).info(
f"新连接建立: {device_id} from {headers.get('x-real-ip', 'unknown')}"
)
# 使用ConnectionService处理连接
await self.connection_service.handle_connection(transport, headers)
except websockets.exceptions.ConnectionClosed:
logger.bind(tag=TAG).info("WebSocket连接正常关闭")
except Exception as e:
logger.bind(tag=TAG).error(f"处理WebSocket连接时出错: {e}", exc_info=True)
# 将错误反馈给管理端,避免长时间等待
try:
if hasattr(websocket, "closed") and not websocket.closed:
await websocket.send(
json.dumps(
{
"type": "server",
"status": "error",
"message": f"Server error: {e}",
"content": {"action": "unknown"},
}
)
)
except Exception:
pass
finally:
# 确保从活动连接集合中移除
self.active_connections.discard(transport)
# 强制关闭连接(如果还没有关闭的话)
try:
if hasattr(websocket, "closed") and not websocket.closed:
await websocket.close()
elif hasattr(websocket, "state") and websocket.state.name != "CLOSED":
await websocket.close()
except Exception as close_error:
logger.bind(tag=TAG).error(f"强制关闭WebSocket连接时出错: {close_error}")
async def _handle_auth(self, headers: Dict[str, str]):
"""
连接时认证
Args:
headers: HTTP 请求头
Raises:
AuthenticationError: 认证失败时抛出
"""
await self.auth_middleware.authenticate_async(headers)
def _extract_headers(self, websocket) -> Dict[str, str]:
"""
从WebSocket请求中提取头信息
支持从以下来源提取信息
1. HTTP 请求头
2. URL 查询参数device-id, client-id, authorization
3. 路径参数 ?from=mqtt_gateway
"""
headers = {}
# 1. 提取 HTTP 请求头
if hasattr(websocket, 'request') and hasattr(websocket.request, 'headers'):
for name, value in websocket.request.headers.items():
headers[name.lower()] = value
elif hasattr(websocket, 'request_headers'):
for name, value in websocket.request_headers.items():
headers[name.lower()] = value
# 2. 提取路径参数(如果有的话)
request_path = None
if hasattr(websocket, 'request') and hasattr(websocket.request, 'path'):
request_path = websocket.request.path
elif hasattr(websocket, 'path'):
request_path = websocket.path
if request_path:
from urllib.parse import urlparse, parse_qs
parsed = urlparse(request_path)
query_params = parse_qs(parsed.query)
# 处理关键参数:device-id, client-id, authorization
key_params = ['device-id', 'client-id', 'authorization']
for key in key_params:
if key in query_params and query_params[key]:
# URL 参数优先级低于 header
if key not in headers or not headers[key]:
headers[key] = query_params[key][0]
# 处理其他参数
for key, values in query_params.items():
if values and key not in headers:
headers[key] = values[0]
# 检查是否来自 MQTT 网关
if request_path.endswith("?from=mqtt_gateway") or "from=mqtt_gateway" in request_path:
headers['from_mqtt_gateway'] = 'true'
# 3. 提取远程地址
if hasattr(websocket, 'remote_address'):
# 如果 headers 中没有 x-real-ip,使用 remote_address
if 'x-real-ip' not in headers:
headers['x-real-ip'] = websocket.remote_address[0]
return headers
async def _http_response(self, websocket, request_headers):
"""处理HTTP请求"""
# 检查是否为 WebSocket 升级请求
if request_headers.headers.get("connection", "").lower() == "upgrade":
# 如果是 WebSocket 请求,返回 None 允许握手继续
return None
else:
# 如果是普通 HTTP 请求,返回服务器状态
return websocket.respond(200, "New Architecture WebSocket Server is running\n")
async def apply_config(self, new_config: Dict[str, Any]) -> bool:
"""Apply a facade-validated config to future WebSocket connections."""
self.config = new_config
self.connection_service = ConnectionService(new_config)
self.connection_service.server = getattr(self, "management_owner", self)
self.auth_middleware = AuthMiddleware(new_config)
return True
async def update_config(
self, new_config: Optional[Dict[str, Any]] = None
) -> bool:
"""
更新服务器配置并重新初始化组件
Returns:
bool: 更新是否成功
"""
try:
async with self.config_lock:
logger.bind(tag=TAG).info("开始更新服务器配置")
self.last_update_error = None
old_config = self.config
old_connection_service = self.connection_service
old_auth_middleware = self.auth_middleware
# 管理命令可自行拉取配置;多协议管理器则直接传入已合并配置。
if new_config is None:
new_config = await get_config_from_api_async(self.config)
if new_config is None:
logger.bind(tag=TAG).error("获取新配置失败")
self.last_update_error = "获取新配置失败"
return False
logger.bind(tag=TAG).info("获取新配置成功")
new_shared_manager = None
reuse_manager = False
# 校验新配置(预初始化组件以发现配置错误)
ok, error_msg = await validate_config_components(new_config, logger)
if not ok:
logger.bind(tag=TAG).error(f"配置校验失败: {error_msg}")
self.last_update_error = f"配置校验失败: {error_msg}"
return False
# 准备共享 ASR 管理器(本地模型走共享预加载)
old_shared_manager = old_config.get("_shared_asr_manager")
selected_asr = new_config.get("selected_module", {}).get("ASR")
if selected_asr:
asr_config = new_config.get("ASR", {}).get(selected_asr, {})
asr_type = asr_config.get("type", selected_asr)
if SharedASRManager.is_local_model_type(asr_type):
if (
old_shared_manager
and getattr(old_shared_manager, "asr_type", None) == asr_type
and old_shared_manager.is_ready()
):
new_shared_manager = old_shared_manager
reuse_manager = True
else:
new_shared_manager = SharedASRManager(new_config, asr_type)
await new_shared_manager.initialize()
# 非本地模型时不立即关闭旧共享管理器,待更新成功后统一处理
# 检查 VAD 和 ASR 类型是否需要更新
update_vad = check_vad_update(self.config, new_config)
update_asr = check_asr_update(self.config, new_config)
logger.bind(tag=TAG).info(
f"检查VAD和ASR类型是否需要更新: VAD={update_vad}, ASR={update_asr}"
)
# 检查配置是否有重大变化
changed_configs = self._get_changed_configs(self.config, new_config)
# 更新配置
self.config = new_config
if new_shared_manager:
self.config["_shared_asr_manager"] = new_shared_manager
elif "_shared_asr_manager" in self.config:
del self.config["_shared_asr_manager"]
# 根据变化类型进行更新
if changed_configs:
logger.bind(tag=TAG).info(f"配置项变化: {', '.join(changed_configs)}")
# 重新创建连接服务,使用新配置
# 注意:已建立的连接会继续使用旧配置,只有新连接使用新配置
self.connection_service = ConnectionService(new_config)
self.connection_service.server = self
# 如果 ASR 配置变化且复用旧共享管理器,提示重启
if update_asr and reuse_manager:
logger.bind(tag=TAG).warning(
"ASR 配置已变化,但仍复用已有共享 ASR 管理器,建议重启服务"
)
else:
# 即使没有重大变化,也更新 ConnectionService 的配置引用
self.connection_service.config = new_config
self.connection_service.server = self
# 更新认证中间件
self.auth_middleware = AuthMiddleware(new_config)
# 更新成功后再关闭旧共享管理器(避免失败回滚时不可用)
if old_shared_manager and old_shared_manager is not new_shared_manager:
await old_shared_manager.shutdown()
logger.bind(tag=TAG).info("配置更新任务执行完毕")
return True
except Exception as e:
logger.bind(tag=TAG).error(f"更新服务器配置失败: {str(e)}", exc_info=True)
self.last_update_error = f"更新服务器配置失败: {str(e)}"
try:
if new_shared_manager and not reuse_manager:
await new_shared_manager.shutdown()
self.config = old_config
self.connection_service = old_connection_service
self.auth_middleware = old_auth_middleware
except Exception:
pass
return False
def get_last_update_error(self) -> str:
return self.last_update_error or ""
def _get_changed_configs(self, old_config: Dict[str, Any], new_config: Dict[str, Any]) -> list:
"""
获取变化的配置项列表
Returns:
list: 变化的配置项名称列表
"""
changed = []
key_configs = [
"selected_module",
"VAD",
"ASR",
"LLM",
"TTS",
"Memory",
"Intent"
]
for key in key_configs:
old_value = old_config.get(key)
new_value = new_config.get(key)
if old_value != new_value:
changed.append(key)
return changed
def get_active_connections_count(self) -> int:
"""获取活跃连接数"""
return len(self.active_connections)
def get_server_status(self) -> Dict[str, Any]:
"""获取服务器状态"""
return {
"active_connections": self.get_active_connections_count(),
"server_type": "new_architecture",
"processors": self.connection_service.message_router.list_processors()
}
@@ -0,0 +1,523 @@
#!/usr/bin/env python3
"""
小智服务器门面类
统一管理所有协议服务器的启动和停止
"""
import asyncio
from typing import Dict, Any, Optional
from config.logger import setup_logging
from core.servers.multi_protocol_server import MultiProtocolServer
logger = setup_logging()
TAG = __name__
class XiaozhiServerFacade:
"""
小智服务器门面类
提供统一的服务器管理接口屏蔽内部协议复杂性
功能
- 协议管理WebSocketMQTT
- 本地 ASR 模型预加载
- 优雅启动和停止
"""
def __init__(self, config: Dict[str, Any]):
"""
初始化服务器门面
Args:
config: 服务器配置字典
"""
self.config = config
self.multi_protocol_server: Optional[MultiProtocolServer] = None
self.shared_asr_manager = None # 共享 ASR 管理器
self._retired_shared_asr_managers = []
self.is_initialized = False
self.is_running = False
self.last_update_error = None
self._cleanup_pending = False
# 处理协议配置
self._setup_protocol_config()
def _setup_protocol_config(self):
"""设置协议配置"""
try:
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}")
except Exception as e:
logger.error(f"设置协议配置失败: {e}")
# 使用最基本的配置
self.config["enabled_protocols"] = ["websocket"]
async def initialize(self):
"""初始化服务器"""
if self.is_initialized:
logger.bind(tag=TAG).warning("服务器已经初始化")
return
try:
logger.bind(tag=TAG).info("正在初始化小智服务器...")
# 检查并预加载本地 ASR 模型(关键步骤)
await self._preload_asr_if_needed()
# 创建多协议服务器
self.multi_protocol_server = MultiProtocolServer(self.config)
self.multi_protocol_server.set_management_owner(self)
self.is_initialized = True
logger.bind(tag=TAG).info("小智服务器初始化完成")
except Exception as e:
logger.bind(tag=TAG).error(f"初始化服务器失败: {e}")
raise
async def _preload_asr_if_needed(self):
"""
检查并预加载本地 ASR 模型
如果配置使用本地 ASR 模型 FunASR则在服务器启动时预加载
避免首次语音识别时的延迟导致客户端超时
"""
try:
# 获取 ASR 配置
selected_asr = self.config.get("selected_module", {}).get("ASR")
if not selected_asr:
logger.bind(tag=TAG).info("未配置 ASR 模块,跳过预加载")
return
# 获取 ASR 类型
asr_config = self.config.get("ASR", {}).get(selected_asr, {})
asr_type = asr_config.get("type", selected_asr)
# 导入 SharedASRManager 检查是否为本地模型
from core.providers.asr.shared_asr_manager import SharedASRManager
if SharedASRManager.is_local_model_type(asr_type):
logger.bind(tag=TAG).info(
f"检测到本地 ASR 模型: {asr_type},开始预加载..."
)
# 创建全局 ASR 管理器
self.shared_asr_manager = SharedASRManager(self.config, asr_type)
# 预加载模型
await self.shared_asr_manager.initialize()
# 将管理器放入配置中供后续使用
self.config['_shared_asr_manager'] = self.shared_asr_manager
logger.bind(tag=TAG).info(
f"ASR 模型预加载完成,类型: {asr_type}"
)
else:
logger.bind(tag=TAG).info(
f"ASR 类型为远程服务: {asr_type},无需预加载"
)
except Exception as e:
logger.bind(tag=TAG).error(f"ASR 预加载失败: {e}")
# 预加载失败不影响服务器启动,继续使用懒加载模式
logger.bind(tag=TAG).warning("将回退到懒加载模式")
async def start(self):
"""启动服务器"""
if self._cleanup_pending:
raise RuntimeError("上次停止尚未完成,请先重试 stop 清理残留资源")
if not self.is_initialized:
await self.initialize()
if self.is_running:
logger.warning("服务器已经在运行中")
return
try:
logger.info("正在启动小智服务器...")
# 启动多协议服务器
await self.multi_protocol_server.start()
self.is_running = True
logger.info("小智服务器启动成功")
except Exception as e:
logger.error(f"启动服务器失败: {e}")
self.is_running = False
# initialize() may already own a shared ASR manager and partially
# started listeners. Release both before propagating startup failure.
try:
await self.stop()
except Exception as cleanup_error:
logger.bind(tag=TAG).error(
f"启动失败后的资源清理失败: {cleanup_error}"
)
raise
async def stop(self):
"""停止服务器"""
if (
not self.is_running
and self.multi_protocol_server is None
and self.shared_asr_manager is None
and not self._retired_shared_asr_managers
):
logger.bind(tag=TAG).info("服务器未在运行")
return
logger.bind(tag=TAG).info("正在停止小智服务器...")
errors = []
protocols_stopped = self.multi_protocol_server is None
if self.multi_protocol_server:
try:
await self.multi_protocol_server.stop()
except Exception as e:
errors.append(("多协议服务器", e))
logger.bind(tag=TAG).error(f"停止多协议服务器失败: {e}")
else:
self.multi_protocol_server = None
protocols_stopped = True
if self.shared_asr_manager and protocols_stopped:
logger.bind(tag=TAG).info("正在关闭共享 ASR 管理器...")
try:
await self.shared_asr_manager.shutdown()
except Exception as e:
errors.append(("共享 ASR", e))
logger.bind(tag=TAG).error(f"关闭共享 ASR 管理器失败: {e}")
else:
self.shared_asr_manager = None
self.config.pop('_shared_asr_manager', None)
elif self.shared_asr_manager:
logger.bind(tag=TAG).warning(
"协议服务器仍持有连接,延后关闭共享 ASR 管理器"
)
if protocols_stopped and self._retired_shared_asr_managers:
remaining_retired = []
for manager in self._retired_shared_asr_managers:
try:
await manager.shutdown()
except Exception as e:
remaining_retired.append(manager)
errors.append(("旧共享 ASR", e))
logger.bind(tag=TAG).error(
f"关闭旧共享ASR管理器失败: {e}"
)
self._retired_shared_asr_managers = remaining_retired
self.is_running = False
self.is_initialized = bool(
self.multi_protocol_server
or self.shared_asr_manager
or self._retired_shared_asr_managers
)
if not self.is_initialized:
self.shared_asr_manager = None
self.config.pop('_shared_asr_manager', None)
logger.bind(tag=TAG).info("小智服务器已停止")
else:
logger.bind(tag=TAG).warning(
"服务器部分资源停止失败,已保留所有权供重试清理"
)
if errors:
self._cleanup_pending = True
details = ", ".join(f"{owner}: {error}" for owner, error in errors)
raise RuntimeError(f"停止服务器时发生错误: {details}")
self._cleanup_pending = False
async def restart(self):
"""重启服务器"""
logger.info("重启小智服务器...")
await self.stop()
await asyncio.sleep(1) # 等待清理完成
await self.start()
async def update_config(
self, new_config: Optional[Dict[str, Any]] = None
) -> bool:
"""
更新服务器配置
Args:
new_config: 新的配置字典
Returns:
bool: 更新是否成功
"""
old_config = self.config
old_shared_manager = self.shared_asr_manager
prepared_shared_manager = old_shared_manager
owns_prepared_manager = False
try:
logger.info("更新服务器配置...")
self.last_update_error = None
if new_config is None:
from config.config_loader import get_config_from_api_async
new_config = await get_config_from_api_async(self.config)
if new_config is None:
raise RuntimeError("获取新配置失败")
# 使用新顶层对象,避免 MultiProtocolServer 的旧配置引用被原地改写,
# 从而导致协议/端口变化无法被检测。
merged_config = dict(self.config)
merged_config.update(new_config)
self.config = merged_config
self._setup_protocol_config()
from core.utils.config_validation import validate_config_components
from core.utils.util import check_asr_update
from core.providers.asr.shared_asr_manager import SharedASRManager
ok, error_msg = await validate_config_components(self.config, logger)
if not ok:
raise RuntimeError(f"配置校验失败: {error_msg}")
selected_asr = self.config.get("selected_module", {}).get("ASR")
asr_config = self.config.get("ASR", {}).get(selected_asr, {})
asr_type = asr_config.get("type", selected_asr) if selected_asr else None
needs_new_asr = check_asr_update(old_config, self.config)
if asr_type and SharedASRManager.is_local_model_type(asr_type):
if not (
old_shared_manager
and not needs_new_asr
and old_shared_manager.is_ready()
):
prepared_shared_manager = SharedASRManager(self.config, asr_type)
await prepared_shared_manager.initialize()
owns_prepared_manager = True
self.config["_shared_asr_manager"] = prepared_shared_manager
else:
prepared_shared_manager = None
self.config.pop("_shared_asr_manager", None)
# 已初始化时即更新实例集合;运行中会完成监听器切换。
if self.multi_protocol_server:
success = await self.multi_protocol_server.update_config(self.config)
if success:
logger.info("服务器配置更新成功")
else:
raise RuntimeError("多协议服务器配置更新失败")
self.shared_asr_manager = prepared_shared_manager
if (
old_shared_manager
and old_shared_manager is not prepared_shared_manager
):
try:
await old_shared_manager.shutdown()
except Exception as e:
# 新配置已经提交,保留旧资源所有权供 stop 重试。
self._retired_shared_asr_managers.append(
old_shared_manager
)
logger.bind(tag=TAG).error(f"关闭旧共享ASR管理器失败: {e}")
logger.info("配置更新完成")
return True
except Exception as e:
logger.error(f"更新配置失败: {e}")
self.last_update_error = str(e)
if owns_prepared_manager and prepared_shared_manager:
try:
await prepared_shared_manager.shutdown()
except Exception as cleanup_error:
self._retired_shared_asr_managers.append(
prepared_shared_manager
)
logger.bind(tag=TAG).error(
f"回滚新共享ASR管理器失败: {cleanup_error}"
)
self.config = old_config
self.shared_asr_manager = old_shared_manager
if self.multi_protocol_server:
self.is_running = self.multi_protocol_server.is_running
if (
self.multi_protocol_server.server_tasks
and not self.multi_protocol_server.is_running
):
self._cleanup_pending = True
return False
def get_last_update_error(self) -> str:
return self.last_update_error or ""
def get_server_status(self) -> Dict[str, Any]:
"""获取服务器状态"""
base_status = {
'is_initialized': self.is_initialized,
'is_running': self.is_running,
'enabled_protocols': self.config.get('enabled_protocols', [])
}
if self.multi_protocol_server:
server_status = self.multi_protocol_server.get_server_status()
base_status.update(server_status)
# 添加 ASR 状态
if self.shared_asr_manager:
base_status['asr'] = {
'mode': 'shared',
'ready': self.shared_asr_manager.is_ready(),
'queue_status': self.shared_asr_manager.get_queue_status()
}
else:
base_status['asr'] = {'mode': 'lazy_load'}
return base_status
def get_active_connections_count(self) -> Dict[str, int]:
"""获取各协议的活跃连接数"""
if self.multi_protocol_server:
return self.multi_protocol_server.get_active_connections_count()
return {}
def get_supported_protocols(self) -> list:
"""获取支持的协议列表"""
if self.multi_protocol_server:
return self.multi_protocol_server.get_supported_protocols()
return ['websocket', 'mqtt']
def is_protocol_enabled(self, protocol: str) -> bool:
"""检查协议是否启用"""
enabled_protocols = self.config.get('enabled_protocols', [])
return protocol in enabled_protocols
async def broadcast_message(self, message: Dict[str, Any], protocol: Optional[str] = None):
"""
向所有连接广播消息
Args:
message: 要广播的消息
protocol: 指定协议None表示向所有协议广播
"""
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'):
return {'enabled': False}
server_config = self.config.get('server', {})
return {
'enabled': True,
'host': server_config.get('ip', '0.0.0.0'),
'port': server_config.get('port', 8000),
'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()
}