mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-27 17:43:55 +08:00
Compare commits
6
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f2b4c7b932 | ||
|
|
5b288bb0d2 | ||
|
|
5e0d853256 | ||
|
|
2a9f809700 | ||
|
|
eac573706d | ||
|
|
0c582ed3b6 |
@@ -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;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+46
-2
@@ -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();
|
||||
|
||||
+74
-25
@@ -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;
|
||||
|
||||
+310
-99
@@ -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("-", "_");
|
||||
}
|
||||
}
|
||||
+15
-4
@@ -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);
|
||||
|
||||
+183
@@ -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);
|
||||
}
|
||||
}
|
||||
}
|
||||
+56
@@ -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;
|
||||
}
|
||||
}
|
||||
}
|
||||
+431
@@ -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) {
|
||||
}
|
||||
}
|
||||
+5
-1
@@ -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")) {
|
||||
|
||||
+6
-1
@@ -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
@@ -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__":
|
||||
|
||||
@@ -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})
|
||||
@@ -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:
|
||||
|
||||
@@ -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__()
|
||||
@@ -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))
|
||||
|
||||
@@ -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主题名称无效")
|
||||
|
||||
# 消息ID(QoS > 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()
|
||||
|
||||
@@ -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")
|
||||
|
||||
|
||||
@@ -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)
|
||||
@@ -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("密码签名验证失败")
|
||||
@@ -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)
|
||||
@@ -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
|
||||
集成ConnectionService、MessageRouter和新的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:
|
||||
"""
|
||||
小智服务器门面类
|
||||
提供统一的服务器管理接口,屏蔽内部协议复杂性
|
||||
|
||||
功能:
|
||||
- 协议管理(WebSocket、MQTT)
|
||||
- 本地 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()
|
||||
}
|
||||
Reference in New Issue
Block a user